diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 174405a..0ece85d 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -234,9 +234,9 @@ func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) { 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(` - SELECT ut.id, sl.speed + SELECT ut.id, sl.id FROM user_tunnel ut LEFT JOIN speed_limit sl ON sl.id = ut.speed_id WHERE ut.user_id = ? AND ut.tunnel_id = ? @@ -244,18 +244,18 @@ func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *i LIMIT 1 `, userID, tunnelID) var userTunnelID int64 - var speed sql.NullInt64 - err := row.Scan(&userTunnelID, &speed) + var limiterID sql.NullInt64 + err := row.Scan(&userTunnelID, &limiterID) if err != nil { if errors.Is(err, sql.ErrNoRows) { return 0, nil, nil } return 0, nil, err } - if !speed.Valid || speed.Int64 <= 0 { + if !limiterID.Valid || limiterID.Int64 <= 0 { return userTunnelID, nil, nil } - v := int(speed.Int64) + v := limiterID.Int64 return userTunnelID, &v, nil } @@ -328,7 +328,7 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all return errors.New("转发入口端口不存在") } - userTunnelID, limiter, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) + userTunnelID, limiterID, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) if err != nil { return err } @@ -339,7 +339,7 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all if err != nil { 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) if err != nil && allowFallbackAdd && method == "UpdateService" { _, 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, "不存在") } -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"} services := make([]map[string]interface{}, 0, 2) 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) != "" { service["metadata"] = map[string]interface{}{"interface": node.InterfaceName} } - if limiter != nil && *limiter > 0 { - // Convert Mbps to Bytes/s - // 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) + if limiterID != nil && *limiterID > 0 { + service["limiter"] = strconv.FormatInt(*limiterID, 10) } services = append(services, service) } @@ -1111,3 +1108,39 @@ func asBool(v interface{}, def bool) bool { 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 +} diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 31da320..61849a3 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -1469,12 +1469,15 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) { return } 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(?, ?, ?, ?, ?, ?, ?)`, - name, asInt(req["speed"], 100), tunnelID, tunnelName, now, now, asInt(req["status"], 1)) + speed := asInt(req["speed"], 100) + 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 { response.WriteJSON(w, response.Err(-2, err.Error())) return } + id, _ := res.LastInsertId() + _ = h.sendLimiterConfig(id, speed, tunnelID) response.WriteJSON(w, response.OKEmpty()) } @@ -1496,12 +1499,14 @@ func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("隧道不存在")) 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=?`, - 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 { response.WriteJSON(w, response.Err(-2, err.Error())) return } + _ = h.sendLimiterConfig(id, speed, tunnelID) response.WriteJSON(w, response.OKEmpty()) } @@ -1510,11 +1515,17 @@ func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) { if id <= 0 { 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) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } + if tunnelID > 0 { + _ = h.sendDeleteLimiterConfig(id, tunnelID) + } response.WriteJSON(w, response.OKEmpty()) }