diff --git a/docker-compose-v4.yml b/docker-compose-v4.yml index f13eb1b..1169509 100644 --- a/docker-compose-v4.yml +++ b/docker-compose-v4.yml @@ -12,6 +12,7 @@ services: JWT_SECRET: ${JWT_SECRET} LOG_DIR: /app/logs SERVER_ADDR: :6365 + TZ: Asia/Shanghai ports: - "${BACKEND_PORT}:6365" volumes: diff --git a/docker-compose-v6.yml b/docker-compose-v6.yml index 15ba827..e0832eb 100644 --- a/docker-compose-v6.yml +++ b/docker-compose-v6.yml @@ -12,6 +12,7 @@ services: JWT_SECRET: ${JWT_SECRET} LOG_DIR: /app/logs SERVER_ADDR: :6365 + TZ: Asia/Shanghai ports: - "${BACKEND_PORT}:6365" volumes: diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 174405a..75452eb 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, *int, error) { row := h.repo.DB().QueryRow(` - SELECT ut.id, sl.speed + SELECT ut.id, sl.id, sl.speed FROM user_tunnel ut LEFT JOIN speed_limit sl ON sl.id = ut.speed_id WHERE ut.user_id = ? AND ut.tunnel_id = ? @@ -244,19 +244,21 @@ func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *i LIMIT 1 `, userID, tunnelID) var userTunnelID int64 + var limiterID sql.NullInt64 var speed sql.NullInt64 - err := row.Scan(&userTunnelID, &speed) + err := row.Scan(&userTunnelID, &limiterID, &speed) if err != nil { if errors.Is(err, sql.ErrNoRows) { - return 0, nil, nil + return 0, nil, nil, nil } - return 0, nil, err + return 0, nil, nil, err } - if !speed.Valid || speed.Int64 <= 0 { - return userTunnelID, nil, nil + if !limiterID.Valid || limiterID.Int64 <= 0 { + return userTunnelID, nil, nil, nil } - v := int(speed.Int64) - return userTunnelID, &v, nil + v := limiterID.Int64 + s := int(speed.Int64) + return userTunnelID, &v, &s, nil } func (h *Handler) listUserTunnelIDs(userID, tunnelID int64) ([]int64, error) { @@ -328,18 +330,22 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all return errors.New("转发入口端口不存在") } - userTunnelID, limiter, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) + userTunnelID, limiterID, speed, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) if err != nil { return err } serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID) for _, fp := range ports { + if limiterID != nil && speed != nil { + h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed) + } + node, err := h.getNodeRecord(fp.NodeID) 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) @@ -362,7 +368,7 @@ func (h *Handler) controlForwardServices(forward *forwardRecord, commandType str if len(ports) == 0 { return nil } - userTunnelID, _, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) + userTunnelID, _, _, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) if err != nil { return err } @@ -1003,7 +1009,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 +1050,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 +1114,49 @@ 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 +} + +func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) { + rate := float64(speed) / 8.0 + limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate) + payload := map[string]interface{}{ + "name": strconv.FormatInt(limiterID, 10), + "limits": []string{limitStr}, + } + _, _ = h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false) +} diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index fcf060f..6324371 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -1518,12 +1518,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()) } @@ -1545,12 +1548,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()) } @@ -1559,11 +1564,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()) }