From 77321d4fa1f0f1b41ab50c3386680e62e5c9057e Mon Sep 17 00:00:00 2001 From: sagitchu Date: Thu, 23 Apr 2026 15:55:14 +0800 Subject: [PATCH] feat: implement max conn limiter dispatching --- .../internal/http/handler/control_plane.go | 59 ++++++- .../http/handler/control_plane_test.go | 8 +- go-backend/internal/store/model/model.go | 1 + go-backend/patch.py | 163 ++++++++++++++++++ 4 files changed, 222 insertions(+), 9 deletions(-) create mode 100644 go-backend/patch.py diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index baec361..e4f8d1e 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -266,6 +266,22 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method serviceBase := buildForwardServiceBaseWithResolvedUserTunnel(forward.ID, forward.UserID, userTunnelID) + user, err := h.repo.GetUserByID(forward.UserID) + if err != nil { + return nil, err + } + + var cLimiterName string + var maxConnToSet int + + if forward.MaxConn > 0 { + maxConnToSet = forward.MaxConn + cLimiterName = fmt.Sprintf("rule_conn_limit_%d", forward.ID) + } else if user != nil && user.MaxConn > 0 { + maxConnToSet = user.MaxConn + cLimiterName = fmt.Sprintf("user_conn_limit_%d", user.ID) + } + for _, fp := range ports { if limiterID != nil && speed != nil { if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil { @@ -283,11 +299,17 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method } } + if cLimiterName != "" { + if err := h.ensureConnLimiterOnNode(fp.NodeID, cLimiterName, maxConnToSet); err != nil { + warnings = append(warnings, fmt.Sprintf("节点 %d 连接限制器下发失败: %v", fp.NodeID, err)) + } + } + node, err := h.getNodeRecord(fp.NodeID) if err != nil { return nil, err } - services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID) + services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID, cLimiterName) _, err = h.sendNodeCommand(node.ID, method, services, true, false) if err != nil && allowFallbackAdd && method == "UpdateService" { if isNotFoundError(err) { @@ -302,7 +324,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method } if err != nil && strings.EqualFold(strings.TrimSpace(method), "UpdateService") && isCannotAssignRequestedAddressError(err) { var warning string - warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, limiterID) + warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, limiterID, cLimiterName) if err == nil && warning != "" { warnings = append(warnings, warning) } @@ -328,7 +350,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method return warnings, nil } -func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, limiterID *int64) (string, error) { +func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, limiterID *int64, cLimiterName string) (string, error) { if h == nil || forward == nil || tunnel == nil || node == nil { return "", errors.New("invalid bind fallback context") } @@ -345,7 +367,7 @@ func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunne } time.Sleep(150 * time.Millisecond) - defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID) + defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID, cLimiterName) if _, err := h.sendNodeCommand(node.ID, "AddService", defaultServices, true, false); err != nil { return "", err } @@ -1637,7 +1659,7 @@ func compactErrorMessage(msg string) string { return strings.Join(strings.Fields(strings.ToLower(msg)), "") } -func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64) []map[string]interface{} { +func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64, cLimiterName string) []map[string]interface{} { protocols := []string{"tcp", "udp"} services := make([]map[string]interface{}, 0, 2) targets := splitRemoteTargets(forward.RemoteAddr) @@ -1680,6 +1702,9 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel }, }, } + if cLimiterName != "" { + service["climiter"] = cLimiterName + } if protocol == "udp" { listenerMetadata := map[string]interface{}{ "keepAlive": true, @@ -1795,6 +1820,30 @@ func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) return nil } +func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxConn int) error { + limitStr := fmt.Sprintf("$ %d", maxConn) + + payload := map[string]interface{}{ + "name": limiterName, + "limits": []string{limitStr}, + } + + if _, err := h.sendNodeCommand(nodeID, "AddCLimiters", payload, false, false); err != nil { + if !isAlreadyExistsMessage(err.Error()) { + return fmt.Errorf("连接限制器下发失败: %w", err) + } + updatePayload := map[string]interface{}{ + "limiter": limiterName, + "data": payload, + } + if _, updateErr := h.sendNodeCommand(nodeID, "UpdateCLimiters", updatePayload, false, false); updateErr != nil { + return fmt.Errorf("连接限制器更新失败: %w", updateErr) + } + } + return nil +} + + func buildLimiterAddPayload(limiterID int64, speed int) (string, map[string]interface{}) { rate := float64(speed) / 8.0 limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate) diff --git a/go-backend/internal/http/handler/control_plane_test.go b/go-backend/internal/http/handler/control_plane_test.go index 81d79bc..97bd80a 100644 --- a/go-backend/internal/http/handler/control_plane_test.go +++ b/go-backend/internal/http/handler/control_plane_test.go @@ -378,7 +378,7 @@ func TestRetryTunnelServiceAddWithCleanupReturnsCleanupError(t *testing.T) { func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) { forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7} node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"} - services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", nil) + services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", nil, "") if len(services) != 2 { t.Fatalf("expected 2 services, got %d", len(services)) } @@ -393,7 +393,7 @@ func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) { func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.T) { forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7} node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"} - services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", nil) + services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", nil, "") if len(services) != 2 { t.Fatalf("expected 2 services, got %d", len(services)) } @@ -409,7 +409,7 @@ func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing. func TestBuildForwardServiceConfigs_BindIPAlreadyContainsPort(t *testing.T) { forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7} node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"} - services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", nil) + services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", nil, "") if len(services) != 2 { t.Fatalf("expected 2 services, got %d", len(services)) } @@ -464,7 +464,7 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) { t.Run(tt.name, func(t *testing.T) { forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7} node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"} - services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, nil) + services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, nil, "") if len(services) != 2 { t.Fatalf("expected 2 services, got %d", len(services)) } diff --git a/go-backend/internal/store/model/model.go b/go-backend/internal/store/model/model.go index e3507d6..d403c7d 100644 --- a/go-backend/internal/store/model/model.go +++ b/go-backend/internal/store/model/model.go @@ -539,6 +539,7 @@ type ForwardRecord struct { Strategy string Status int SpeedID sql.NullInt64 + MaxConn int } // TunnelRecord is a minimal tunnel view used by control plane. diff --git a/go-backend/patch.py b/go-backend/patch.py new file mode 100644 index 0000000..10be01e --- /dev/null +++ b/go-backend/patch.py @@ -0,0 +1,163 @@ +import re + +with open('internal/http/handler/control_plane.go', 'r') as f: + content = f.read() + +# 1. Update ensureLimiterOnNode and add ensureConnLimiterOnNode +ensure_conn_limiter = """ +func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxConn int) error { +\tlimitStr := fmt.Sprintf("$ %d", maxConn) +\t +\tpayload := map[string]interface{}{ +\t\t"name": limiterName, +\t\t"limits": []string{limitStr}, +\t} +\t +\tif _, err := h.sendNodeCommand(nodeID, "AddCLimiters", payload, false, false); err != nil { +\t\tif !isAlreadyExistsMessage(err.Error()) { +\t\t\treturn fmt.Errorf("连接限制器下发失败: %w", err) +\t\t} +\t\tupdatePayload := map[string]interface{}{ +\t\t\t"limiter": limiterName, +\t\t\t"data": payload, +\t\t} +\t\tif _, updateErr := h.sendNodeCommand(nodeID, "UpdateCLimiters", updatePayload, false, false); updateErr != nil { +\t\t\treturn fmt.Errorf("连接限制器更新失败: %w", updateErr) +\t\t} +\t} +\treturn nil +} +""" + +content = content.replace('func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) error {\n\tif err := h.upsertLimiterOnNode(nodeID, limiterID, speed); err != nil {\n\t\treturn fmt.Errorf("限速规则下发失败: %w", err)\n\t}\n\n\treturn nil\n}', +'func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) error {\n\tif err := h.upsertLimiterOnNode(nodeID, limiterID, speed); err != nil {\n\t\treturn fmt.Errorf("限速规则下发失败: %w", err)\n\t}\n\n\treturn nil\n}\n' + ensure_conn_limiter) + + +# 2. Update buildForwardServiceConfigs declaration +content = content.replace( + 'func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64) []map[string]interface{} {', + 'func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64, cLimiterName string) []map[string]interface{} {' +) + + +# 3. Inject climiter into generated service +service_map_end = """ } + if protocol == "udp" {""" +service_map_end_new = """ } + if cLimiterName != "" { + service["climiter"] = cLimiterName + } + if protocol == "udp" {""" +content = content.replace(service_map_end, service_map_end_new) + + +# 4. Update syncForwardServicesWithWarnings +# Find user tunnel resolution +resolution = """ serviceBase := buildForwardServiceBaseWithResolvedUserTunnel(forward.ID, forward.UserID, userTunnelID) + + for _, fp := range ports {""" + +resolution_new = """ serviceBase := buildForwardServiceBaseWithResolvedUserTunnel(forward.ID, forward.UserID, userTunnelID) + + user, err := h.repo.GetUserByID(forward.UserID) + if err != nil { + return nil, err + } + + var cLimiterName string + var maxConnToSet int + + if forward.MaxConn > 0 { + maxConnToSet = forward.MaxConn + cLimiterName = fmt.Sprintf("rule_conn_limit_%d", forward.ID) + } else if user != nil && user.MaxConn > 0 { + maxConnToSet = user.MaxConn + cLimiterName = fmt.Sprintf("user_conn_limit_%d", user.ID) + } + + for _, fp := range ports {""" +content = content.replace(resolution, resolution_new) + +# Inject ensureConnLimiterOnNode inside loop +loop_inner = """ if limiterID != nil && speed != nil { + if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil { + // If the limiter push fails because the node is offline, skip it with a warning + if isNodeOfflineOrTimeoutError(err) { + node, _ := h.getNodeRecord(fp.NodeID) + nodeName := fmt.Sprintf("%d", fp.NodeID) + if node != nil && strings.TrimSpace(node.Name) != "" { + nodeName = strings.TrimSpace(node.Name) + } + warnings = append(warnings, fmt.Sprintf("节点 %s 不在线,已跳过下发", nodeName)) + continue + } + return nil, err + } + } + + node, err := h.getNodeRecord(fp.NodeID)""" + +loop_inner_new = """ if limiterID != nil && speed != nil { + if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil { + // If the limiter push fails because the node is offline, skip it with a warning + if isNodeOfflineOrTimeoutError(err) { + node, _ := h.getNodeRecord(fp.NodeID) + nodeName := fmt.Sprintf("%d", fp.NodeID) + if node != nil && strings.TrimSpace(node.Name) != "" { + nodeName = strings.TrimSpace(node.Name) + } + warnings = append(warnings, fmt.Sprintf("节点 %s 不在线,已跳过下发", nodeName)) + continue + } + return nil, err + } + } + + if cLimiterName != "" { + if err := h.ensureConnLimiterOnNode(fp.NodeID, cLimiterName, maxConnToSet); err != nil { + warnings = append(warnings, fmt.Sprintf("节点 %d 连接限制器下发失败: %v", fp.NodeID, err)) + } + } + + node, err := h.getNodeRecord(fp.NodeID)""" +content = content.replace(loop_inner, loop_inner_new) + +# Update buildForwardServiceConfigs call in syncForwardServicesWithWarnings +content = content.replace( + 'services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID)', + 'services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID, cLimiterName)' +) + +# Update fallbackForwardPortToDefaultBind call +content = content.replace( + 'warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, limiterID)', + 'warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, limiterID, cLimiterName)' +) + +# 5. Update fallbackForwardPortToDefaultBind declaration and logic +content = content.replace( + 'func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, limiterID *int64) (string, error) {', + 'func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, limiterID *int64, cLimiterName string) (string, error) {' +) + +content = content.replace( + 'defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID)', + 'defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID, cLimiterName)' +) + +with open('internal/http/handler/control_plane.go', 'w') as f: + f.write(content) + + +# Update control_plane_test.go +with open('internal/http/handler/control_plane_test.go', 'r') as f: + test_content = f.read() + +test_content = re.sub( + r'buildForwardServiceConfigs\((.*?),(.*?),(.*?),(.*?),(.*?),(.*?),(.*?)\)', + r'buildForwardServiceConfigs(\1,\2,\3,\4,\5,\6,\7, "")', + test_content +) + +with open('internal/http/handler/control_plane_test.go', 'w') as f: + f.write(test_content)