diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index c4b3f76..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 } @@ -1145,29 +1151,12 @@ func (h *Handler) sendDeleteLimiterConfig(limiterID int64, tunnelID int64) error return nil } -func (h *Handler) onNodeConnected(nodeID int64) { - // Sync limiters - rows, err := h.repo.DB().Query(` - SELECT DISTINCT sl.id, sl.speed - FROM speed_limit sl - JOIN chain_tunnel ct ON ct.tunnel_id = sl.tunnel_id - WHERE ct.node_id = ? AND ct.chain_type = 1 AND sl.status = 1 - `, nodeID) - - if err == nil { - defer rows.Close() - for rows.Next() { - var id int64 - var speed int - if err := rows.Scan(&id, &speed); err == nil { - rate := float64(speed) / 8.0 - limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate) - payload := map[string]interface{}{ - "name": strconv.FormatInt(id, 10), - "limits": []string{limitStr}, - } - _, _ = h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false) - } - } +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/handler.go b/go-backend/internal/http/handler/handler.go index ac48973..8236cfe 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -62,13 +62,11 @@ type flowItem struct { } func New(repo *sqlite.Repository, jwtSecret string) *Handler { - h := &Handler{ + return &Handler{ repo: repo, jwtSecret: jwtSecret, wsServer: ws.NewServer(repo, jwtSecret), } - h.wsServer.OnNodeConnected = h.onNodeConnected - return h } func (h *Handler) WebSocketHandler() http.Handler { diff --git a/go-backend/internal/ws/server.go b/go-backend/internal/ws/server.go index 1ff4fd6..f36ec97 100644 --- a/go-backend/internal/ws/server.go +++ b/go-backend/internal/ws/server.go @@ -71,8 +71,6 @@ type Server struct { nodes map[int64]*nodeSession byConn map[*websocket.Conn]*nodeSession pending map[string]pendingRequest - - OnNodeConnected func(nodeID int64) } func NewServer(repo *sqlite.Repository, jwtSecret string) *Server { @@ -166,10 +164,6 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64 _ = s.repo.UpdateNodeOnline(nodeID, 1, version, httpVal, tlsVal, socksVal) s.broadcastStatus(nodeID, 1) - if s.OnNodeConnected != nil { - go s.OnNodeConnected(nodeID) - } - defer func() { needOfflineBroadcast := false s.mu.Lock()