diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 0ece85d..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, *int64, error) { +func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int64, *int, error) { row := h.repo.DB().QueryRow(` - SELECT ut.id, sl.id + 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 = ? @@ -245,18 +245,20 @@ func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *i `, userID, tunnelID) var userTunnelID int64 var limiterID sql.NullInt64 - err := row.Scan(&userTunnelID, &limiterID) + var speed sql.NullInt64 + 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 !limiterID.Valid || limiterID.Int64 <= 0 { - return userTunnelID, nil, nil + return userTunnelID, nil, nil, nil } v := limiterID.Int64 - return userTunnelID, &v, nil + s := int(speed.Int64) + return userTunnelID, &v, &s, nil } func (h *Handler) listUserTunnelIDs(userID, tunnelID int64) ([]int64, error) { @@ -328,13 +330,17 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all return errors.New("转发入口端口不存在") } - userTunnelID, limiterID, 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 @@ -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 } @@ -1144,3 +1150,13 @@ func (h *Handler) sendDeleteLimiterConfig(limiterID int64, tunnelID int64) error } 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) +}