refactor(backend): reimplement speed limit logic

1. Refactor speed limit CRUD to sync with agents immediately via WebSocket (AddLimiters/DeleteLimiters).
2. Update unit conversion to match GOST v3 requirements (Mbps -> MB/s).
3. Update service config generation to reference Limiter IDs instead of hardcoded values.
This commit is contained in:
sagit
2026-02-09 08:12:10 +00:00
parent 630ed969d3
commit 565d732967
2 changed files with 61 additions and 17 deletions
@@ -234,9 +234,9 @@ func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) {
return &n, nil return &n, nil
} }
func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int, error) { func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int64, error) {
row := h.repo.DB().QueryRow(` row := h.repo.DB().QueryRow(`
SELECT ut.id, sl.speed SELECT ut.id, sl.id
FROM user_tunnel ut FROM user_tunnel ut
LEFT JOIN speed_limit sl ON sl.id = ut.speed_id LEFT JOIN speed_limit sl ON sl.id = ut.speed_id
WHERE ut.user_id = ? AND ut.tunnel_id = ? WHERE ut.user_id = ? AND ut.tunnel_id = ?
@@ -244,18 +244,18 @@ func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *i
LIMIT 1 LIMIT 1
`, userID, tunnelID) `, userID, tunnelID)
var userTunnelID int64 var userTunnelID int64
var speed sql.NullInt64 var limiterID sql.NullInt64
err := row.Scan(&userTunnelID, &speed) err := row.Scan(&userTunnelID, &limiterID)
if err != nil { if err != nil {
if errors.Is(err, sql.ErrNoRows) { if errors.Is(err, sql.ErrNoRows) {
return 0, nil, nil return 0, nil, nil
} }
return 0, nil, err return 0, nil, err
} }
if !speed.Valid || speed.Int64 <= 0 { if !limiterID.Valid || limiterID.Int64 <= 0 {
return userTunnelID, nil, nil return userTunnelID, nil, nil
} }
v := int(speed.Int64) v := limiterID.Int64
return userTunnelID, &v, nil return userTunnelID, &v, nil
} }
@@ -328,7 +328,7 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
return errors.New("转发入口端口不存在") return errors.New("转发入口端口不存在")
} }
userTunnelID, limiter, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) userTunnelID, limiterID, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
if err != nil { if err != nil {
return err return err
} }
@@ -339,7 +339,7 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
if err != nil { if err != nil {
return err return err
} }
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiter) services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiterID)
_, err = h.sendNodeCommand(node.ID, method, services, true, false) _, err = h.sendNodeCommand(node.ID, method, services, true, false)
if err != nil && allowFallbackAdd && method == "UpdateService" { if err != nil && allowFallbackAdd && method == "UpdateService" {
_, err = h.sendNodeCommand(node.ID, "AddService", services, true, false) _, err = h.sendNodeCommand(node.ID, "AddService", services, true, false)
@@ -1003,7 +1003,7 @@ func isNotFoundError(err error) bool {
return strings.Contains(msg, "not found") || strings.Contains(msg, "不存在") return strings.Contains(msg, "not found") || strings.Contains(msg, "不存在")
} }
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiter *int) []map[string]interface{} { func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiterID *int64) []map[string]interface{} {
protocols := []string{"tcp", "udp"} protocols := []string{"tcp", "udp"}
services := make([]map[string]interface{}, 0, 2) services := make([]map[string]interface{}, 0, 2)
targets := splitRemoteTargets(forward.RemoteAddr) targets := splitRemoteTargets(forward.RemoteAddr)
@@ -1044,11 +1044,8 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" { if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" {
service["metadata"] = map[string]interface{}{"interface": node.InterfaceName} service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
} }
if limiter != nil && *limiter > 0 { if limiterID != nil && *limiterID > 0 {
// Convert Mbps to Bytes/s service["limiter"] = strconv.FormatInt(*limiterID, 10)
// 1 Mbps = 1,000,000 bits/s = 125,000 Bytes/s
// We use decimal Mbps standard as is common in networking
service["limiter"] = strconv.Itoa(*limiter * 125000)
} }
services = append(services, service) services = append(services, service)
} }
@@ -1111,3 +1108,39 @@ func asBool(v interface{}, def bool) bool {
return def return def
} }
} }
func (h *Handler) sendLimiterConfig(limiterID int64, speedMbps int, tunnelID int64) error {
rate := float64(speedMbps) / 8.0
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
payload := map[string]interface{}{
"name": strconv.FormatInt(limiterID, 10),
"limits": []string{limitStr},
}
nodes, err := h.tunnelEntryNodeIDs(tunnelID)
if err != nil {
return err
}
for _, nodeID := range nodes {
_, _ = h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false)
}
return nil
}
func (h *Handler) sendDeleteLimiterConfig(limiterID int64, tunnelID int64) error {
payload := map[string]interface{}{
"limiter": strconv.FormatInt(limiterID, 10),
}
nodes, err := h.tunnelEntryNodeIDs(tunnelID)
if err != nil {
return err
}
for _, nodeID := range nodes {
_, _ = h.sendNodeCommand(nodeID, "DeleteLimiters", payload, false, true)
}
return nil
}
+14 -3
View File
@@ -1469,12 +1469,15 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
return return
} }
now := time.Now().UnixMilli() now := time.Now().UnixMilli()
_, err := h.repo.DB().Exec(`INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?, ?, ?)`, speed := asInt(req["speed"], 100)
name, asInt(req["speed"], 100), tunnelID, tunnelName, now, now, asInt(req["status"], 1)) res, err := h.repo.DB().Exec(`INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?, ?, ?)`,
name, speed, tunnelID, tunnelName, now, now, asInt(req["status"], 1))
if err != nil { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
id, _ := res.LastInsertId()
_ = h.sendLimiterConfig(id, speed, tunnelID)
response.WriteJSON(w, response.OKEmpty()) response.WriteJSON(w, response.OKEmpty())
} }
@@ -1496,12 +1499,14 @@ func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("隧道不存在")) response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return return
} }
speed := asInt(req["speed"], 100)
_, err := h.repo.DB().Exec(`UPDATE speed_limit SET name=?, speed=?, tunnel_id=?, tunnel_name=?, status=?, updated_time=? WHERE id=?`, _, err := h.repo.DB().Exec(`UPDATE speed_limit SET name=?, speed=?, tunnel_id=?, tunnel_name=?, status=?, updated_time=? WHERE id=?`,
asString(req["name"]), asInt(req["speed"], 100), tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli(), id) asString(req["name"]), speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli(), id)
if err != nil { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
_ = h.sendLimiterConfig(id, speed, tunnelID)
response.WriteJSON(w, response.OKEmpty()) response.WriteJSON(w, response.OKEmpty())
} }
@@ -1510,11 +1515,17 @@ func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) {
if id <= 0 { if id <= 0 {
return return
} }
var tunnelID int64
_ = h.repo.DB().QueryRow(`SELECT tunnel_id FROM speed_limit WHERE id = ?`, id).Scan(&tunnelID)
_, err := h.repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, id) _, err := h.repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, id)
if err != nil { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
if tunnelID > 0 {
_ = h.sendDeleteLimiterConfig(id, tunnelID)
}
response.WriteJSON(w, response.OKEmpty()) response.WriteJSON(w, response.OKEmpty())
} }