From 46394388b1e959b9645192085d6353da7ab8b72a Mon Sep 17 00:00:00 2001 From: sagitchu Date: Mon, 27 Apr 2026 22:32:35 +0800 Subject: [PATCH] feat: sync per-IP runtime limiters --- .../internal/http/handler/control_plane.go | 149 ++++++++++++---- .../http/handler/control_plane_test.go | 54 +++++- .../handler/forward_proxy_protocol_test.go | 2 +- .../contract/max_conn_limit_contract_test.go | 9 +- .../per_ip_speed_limit_contract_test.go | 168 ++++++++++++++++++ 5 files changed, 337 insertions(+), 45 deletions(-) create mode 100644 go-backend/tests/contract/per_ip_speed_limit_contract_test.go diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 8fd2c33..3bb7b25 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -27,6 +27,16 @@ type nodeRecord = model.NodeRecord type chainNodeRecord = model.ChainNodeRecord +type forwardRuntimeLimiters struct { + TrafficLimiter string + ConnLimiter string +} + +type forwardLimiterConfig struct { + Name string + Limits []string +} + type diagnosisTarget struct { Address string IP string @@ -264,6 +274,13 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method speed = utSpeed } + var ipSpeed *int + if forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0 { + if speedVal, err := h.repo.GetSpeedLimitSpeed(forward.IPSpeedID.Int64); err == nil && speedVal > 0 { + ipSpeed = &speedVal + } + } + serviceBase := buildForwardServiceBaseWithResolvedUserTunnel(forward.ID, forward.UserID, userTunnelID) user, err := h.repo.GetUserByID(forward.UserID) @@ -271,19 +288,31 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method 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) + userMaxConn := 0 + if user != nil && user.MaxConn > 0 { + userMaxConn = user.MaxConn } + connLimiterConfig := buildConnLimiterConfig(forward, userMaxConn) for _, fp := range ports { - if limiterID != nil && speed != nil { + runtimeLimiters := forwardRuntimeLimiters{ConnLimiter: connLimiterConfig.Name} + if ipSpeed != nil { + runtimeLimiters.TrafficLimiter = fmt.Sprintf("rule_traffic_limit_%d", forward.ID) + if err := h.ensureTrafficLimiterOnNode(fp.NodeID, runtimeLimiters.TrafficLimiter, speed, ipSpeed); 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 + } + } else if limiterID != nil && speed != nil { + runtimeLimiters.TrafficLimiter = strconv.FormatInt(*limiterID, 10) 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) { @@ -299,8 +328,8 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method } } - if cLimiterName != "" { - if err := h.ensureConnLimiterOnNode(fp.NodeID, cLimiterName, maxConnToSet); err != nil { + if connLimiterConfig.Name != "" { + if err := h.ensureConnLimiterOnNode(fp.NodeID, connLimiterConfig); err != nil { warnings = append(warnings, fmt.Sprintf("节点 %d 连接限制器下发失败: %v", fp.NodeID, err)) } } @@ -309,7 +338,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method if err != nil { return nil, err } - services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID, cLimiterName) + services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), runtimeLimiters) _, err = h.sendNodeCommand(node.ID, method, services, true, false) if err != nil && allowFallbackAdd && method == "UpdateService" { if isNotFoundError(err) { @@ -324,7 +353,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, cLimiterName) + warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, runtimeLimiters) if err == nil && warning != "" { warnings = append(warnings, warning) } @@ -350,7 +379,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, cLimiterName string) (string, error) { +func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, runtimeLimiters forwardRuntimeLimiters) (string, error) { if h == nil || forward == nil || tunnel == nil || node == nil { return "", errors.New("invalid bind fallback context") } @@ -367,7 +396,7 @@ func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunne } time.Sleep(150 * time.Millisecond) - defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID, cLimiterName) + defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", runtimeLimiters) if _, err := h.sendNodeCommand(node.ID, "AddService", defaultServices, true, false); err != nil { return "", err } @@ -1659,7 +1688,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, cLimiterName string) []map[string]interface{} { +func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, runtimeLimiters forwardRuntimeLimiters) []map[string]interface{} { protocols := []string{"tcp", "udp"} services := make([]map[string]interface{}, 0, 2) targets := splitRemoteTargets(forward.RemoteAddr) @@ -1702,8 +1731,11 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel }, }, } - if cLimiterName != "" { - service["climiter"] = cLimiterName + if runtimeLimiters.ConnLimiter != "" { + service["climiter"] = runtimeLimiters.ConnLimiter + } + if runtimeLimiters.TrafficLimiter != "" { + service["limiter"] = runtimeLimiters.TrafficLimiter } if forward.ProxyProtocol > 0 { handlerConfig := service["handler"].(map[string]interface{}) @@ -1728,9 +1760,6 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel } service["metadata"].(map[string]interface{})["interface"] = node.InterfaceName } - if limiterID != nil && *limiterID > 0 { - service["limiter"] = strconv.FormatInt(*limiterID, 10) - } services = append(services, service) } @@ -1830,22 +1859,16 @@ 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}, +func (h *Handler) ensureConnLimiterOnNode(nodeID int64, cfg forwardLimiterConfig) error { + if cfg.Name == "" || len(cfg.Limits) == 0 { + return nil } - + payload := map[string]interface{}{"name": cfg.Name, "limits": cfg.Limits} 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, - } + updatePayload := map[string]interface{}{"limiter": cfg.Name, "data": payload} if _, updateErr := h.sendNodeCommand(nodeID, "UpdateCLimiters", updatePayload, false, false); updateErr != nil { return fmt.Errorf("连接限制器更新失败: %w", updateErr) } @@ -1853,14 +1876,51 @@ func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxC return nil } -func buildLimiterAddPayload(limiterID int64, speed int) (string, map[string]interface{}) { +func buildConnLimiterConfig(forward *forwardRecord, userMaxConn int) forwardLimiterConfig { + if forward == nil { + return forwardLimiterConfig{} + } + limits := make([]string, 0, 2) + if forward.MaxConn > 0 { + limits = append(limits, fmt.Sprintf("$ %d", forward.MaxConn)) + } else if userMaxConn > 0 { + limits = append(limits, fmt.Sprintf("$ %d", userMaxConn)) + } + if forward.IPMaxConn > 0 { + limits = append(limits, fmt.Sprintf("$$ %d", forward.IPMaxConn)) + } + if len(limits) == 0 { + return forwardLimiterConfig{} + } + name := fmt.Sprintf("user_conn_limit_%d", forward.UserID) + if forward.MaxConn > 0 || forward.IPMaxConn > 0 { + name = fmt.Sprintf("rule_conn_limit_%d", forward.ID) + } + return forwardLimiterConfig{Name: name, Limits: limits} +} + +func speedToLimitLine(key string, speed int) string { rate := float64(speed) / 8.0 - limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate) + return fmt.Sprintf("%s %.1fMB %.1fMB", key, rate, rate) +} + +func buildTrafficLimiterPayload(name string, totalSpeed *int, ipSpeed *int) map[string]interface{} { + limits := make([]string, 0, 3) + if totalSpeed != nil && *totalSpeed > 0 { + limits = append(limits, speedToLimitLine("$", *totalSpeed)) + } + if ipSpeed != nil && *ipSpeed > 0 { + limits = append(limits, speedToLimitLine("0.0.0.0/0", *ipSpeed), speedToLimitLine("::/0", *ipSpeed)) + } + return map[string]interface{}{"name": name, "limits": limits} +} + +func buildLimiterAddPayload(limiterID int64, speed int) (string, map[string]interface{}) { name := strconv.FormatInt(limiterID, 10) return name, map[string]interface{}{ "name": name, - "limits": []string{limitStr}, + "limits": []string{speedToLimitLine("$", speed)}, } } @@ -1888,3 +1948,20 @@ func (h *Handler) upsertLimiterOnNode(nodeID int64, limiterID int64, speed int) return nil } + +func (h *Handler) ensureTrafficLimiterOnNode(nodeID int64, name string, totalSpeed *int, ipSpeed *int) error { + payload := buildTrafficLimiterPayload(name, totalSpeed, ipSpeed) + limits, _ := payload["limits"].([]string) + if name == "" || len(limits) == 0 { + return nil + } + if _, err := h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false); err != nil { + if !isAlreadyExistsMessage(err.Error()) { + return fmt.Errorf("限速规则下发失败: %w", err) + } + if _, updateErr := h.sendNodeCommand(nodeID, "UpdateLimiters", buildLimiterUpdatePayload(name, payload), false, false); updateErr != nil { + return fmt.Errorf("限速规则更新失败: %w", updateErr) + } + } + return nil +} diff --git a/go-backend/internal/http/handler/control_plane_test.go b/go-backend/internal/http/handler/control_plane_test.go index 97bd80a..c3fb79a 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", forwardRuntimeLimiters{}) 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, "", forwardRuntimeLimiters{}) 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", forwardRuntimeLimiters{}) 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, forwardRuntimeLimiters{}) if len(services) != 2 { t.Fatalf("expected 2 services, got %d", len(services)) } @@ -478,6 +478,52 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) { } } +func TestBuildConnLimiterConfigCombinesTotalAndPerIP(t *testing.T) { + cfg := buildConnLimiterConfig(&forwardRecord{ID: 42, UserID: 9, MaxConn: 100, IPMaxConn: 5}, 37) + want := forwardLimiterConfig{Name: "rule_conn_limit_42", Limits: []string{"$ 100", "$$ 5"}} + if !reflect.DeepEqual(cfg, want) { + t.Fatalf("expected %+v, got %+v", want, cfg) + } +} + +func TestBuildConnLimiterConfigUsesUserTotalWithRulePerIP(t *testing.T) { + cfg := buildConnLimiterConfig(&forwardRecord{ID: 42, UserID: 9, IPMaxConn: 5}, 37) + want := forwardLimiterConfig{Name: "rule_conn_limit_42", Limits: []string{"$ 37", "$$ 5"}} + if !reflect.DeepEqual(cfg, want) { + t.Fatalf("expected %+v, got %+v", want, cfg) + } +} + +func TestBuildTrafficLimiterPayloadCombinesTotalAndPerIP(t *testing.T) { + payload := buildTrafficLimiterPayload("rule_traffic_limit_42", intPtr(80), intPtr(40)) + wantLimits := []string{"$ 10.0MB 10.0MB", "0.0.0.0/0 5.0MB 5.0MB", "::/0 5.0MB 5.0MB"} + if payload["name"] != "rule_traffic_limit_42" { + t.Fatalf("expected name rule_traffic_limit_42, got %v", payload["name"]) + } + if !reflect.DeepEqual(payload["limits"], wantLimits) { + t.Fatalf("expected limits %v, got %v", wantLimits, payload["limits"]) + } +} + +func TestBuildForwardServiceConfigsUsesRuntimeLimiterNames(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, "", forwardRuntimeLimiters{TrafficLimiter: "rule_traffic_limit_42", ConnLimiter: "rule_conn_limit_42"}) + if len(services) != 2 { + t.Fatalf("expected 2 services, got %d", len(services)) + } + for _, service := range services { + if service["limiter"] != "rule_traffic_limit_42" { + t.Fatalf("expected traffic limiter rule_traffic_limit_42, got %v", service["limiter"]) + } + if service["climiter"] != "rule_conn_limit_42" { + t.Fatalf("expected conn limiter rule_conn_limit_42, got %v", service["climiter"]) + } + } +} + +func intPtr(v int) *int { return &v } + func TestProcessServerAddress_StripsURLSchemeAndPath(t *testing.T) { tests := []struct { name string diff --git a/go-backend/internal/http/handler/forward_proxy_protocol_test.go b/go-backend/internal/http/handler/forward_proxy_protocol_test.go index c5fc9e9..b54df23 100644 --- a/go-backend/internal/http/handler/forward_proxy_protocol_test.go +++ b/go-backend/internal/http/handler/forward_proxy_protocol_test.go @@ -25,7 +25,7 @@ func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing UDPListenAddr: "0.0.0.0", } - services := buildForwardServiceConfigs("1_2_3", forward, tunnel, node, 4001, "", nil, "") + services := buildForwardServiceConfigs("1_2_3", forward, tunnel, node, 4001, "", forwardRuntimeLimiters{}) if len(services) != 2 { t.Fatalf("expected 2 services, got %d", len(services)) } diff --git a/go-backend/tests/contract/max_conn_limit_contract_test.go b/go-backend/tests/contract/max_conn_limit_contract_test.go index c55e7c5..b356be1 100644 --- a/go-backend/tests/contract/max_conn_limit_contract_test.go +++ b/go-backend/tests/contract/max_conn_limit_contract_test.go @@ -97,6 +97,7 @@ func TestMaxConnLimit(t *testing.T) { "remoteAddr": "1.1.1.1:443", "strategy": "fifo", "maxConn": 42, + "ipMaxConn": 7, "proxyProtocol": 2, } body, err := json.Marshal(payload) @@ -195,8 +196,8 @@ func TestMaxConnLimit(t *testing.T) { t.Fatalf("expected limiter name %s, got %v", expectedName, addData["name"]) } if limits, ok := addData["limits"].([]interface{}); ok { - if len(limits) != 1 || limits[0] != "$ 42" { - t.Fatalf("expected limits to contain '$ 42', got %v", limits) + if len(limits) != 2 || limits[0] != "$ 42" || limits[1] != "$$ 7" { + t.Fatalf("expected limits to contain '$ 42' and '$$ 7', got %v", limits) } } else { t.Fatalf("invalid limits type in AddCLimiters data: %v", addData) @@ -218,8 +219,8 @@ func TestMaxConnLimit(t *testing.T) { t.Fatalf("expected nested name %s, got %v", expectedName, nestedData["name"]) } if nestedLimits, ok := nestedData["limits"].([]interface{}); ok { - if len(nestedLimits) != 1 || nestedLimits[0] != "$ 42" { - t.Fatalf("expected nested limits to contain '$ 42', got %v", nestedLimits) + if len(nestedLimits) != 2 || nestedLimits[0] != "$ 42" || nestedLimits[1] != "$$ 7" { + t.Fatalf("expected nested limits to contain '$ 42' and '$$ 7', got %v", nestedLimits) } } else { t.Fatalf("invalid limits type in UpdateCLimiters nested data: %v", nestedData) diff --git a/go-backend/tests/contract/per_ip_speed_limit_contract_test.go b/go-backend/tests/contract/per_ip_speed_limit_contract_test.go new file mode 100644 index 0000000..faaa1ec --- /dev/null +++ b/go-backend/tests/contract/per_ip_speed_limit_contract_test.go @@ -0,0 +1,168 @@ +package contract_test + +import ( + "bytes" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "reflect" + "sync" + "testing" + "time" + + "go-backend/internal/auth" + "go-backend/internal/http/response" +) + +func TestPerIPSpeedLimitRuntimePayload(t *testing.T) { + secret := "contract-jwt-secret" + router, r := setupContractRouter(t, secret) + server := httptest.NewServer(router) + defer server.Close() + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + now := time.Now().UnixMilli() + if err := r.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "per-ip-speed-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil { + t.Fatalf("insert tunnel: %v", err) + } + var tunnelID int64 + if err := r.DB().Raw("SELECT id FROM tunnel WHERE name = ?", "per-ip-speed-tunnel").Scan(&tunnelID).Error; err != nil { + t.Fatalf("get tunnel ID: %v", err) + } + + if err := r.DB().Exec(` + INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "per-ip-speed-node", "per-ip-speed-secret", "10.22.0.1", "10.22.0.1", "", "32200-32210", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { + t.Fatalf("insert node: %v", err) + } + var nodeID int64 + if err := r.DB().Raw("SELECT id FROM node WHERE name = ?", "per-ip-speed-node").Scan(&nodeID).Error; err != nil { + t.Fatalf("get node ID: %v", err) + } + + if err := r.DB().Exec(` + INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, 1, ?, 32201, 'round', 1, 'tls') + `, tunnelID, nodeID).Error; err != nil { + t.Fatalf("insert chain_tunnel: %v", err) + } + if err := r.DB().Exec(` + INSERT INTO user_tunnel(user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) + VALUES(1, ?, 10, 99999, 0, 0, 1, ?, 1) + `, tunnelID, now+365*24*3600*1000).Error; err != nil { + t.Fatalf("insert user_tunnel: %v", err) + } + + totalSpeedID, err := r.CreateSpeedLimit("per-ip-total-speed", 80, now, 1) + if err != nil { + t.Fatalf("create total speed limit: %v", err) + } + ipSpeedID, err := r.CreateSpeedLimit("per-ip-client-speed", 40, now, 1) + if err != nil { + t.Fatalf("create per-ip speed limit: %v", err) + } + + var commandMu sync.Mutex + receivedCommands := make([]string, 0) + var addLimitersData json.RawMessage + var updateServiceData json.RawMessage + + stopNode := startMockSessionForMaxConn(t, server.URL, "per-ip-speed-secret", func(cmdType string, data json.RawMessage) (bool, string) { + commandMu.Lock() + defer commandMu.Unlock() + receivedCommands = append(receivedCommands, cmdType) + if cmdType == "AddLimiters" { + addLimitersData = append([]byte(nil), data...) + } + if cmdType == "UpdateService" { + updateServiceData = append([]byte(nil), data...) + } + return false, "" + }) + defer stopNode() + + waitNodeStatus(t, r, nodeID, 1) + + payload := map[string]interface{}{ + "name": "per-ip-speed-forward", + "tunnelId": tunnelID, + "remoteAddr": "1.1.1.1:443", + "strategy": "fifo", + "speedId": totalSpeedID, + "ipSpeedId": ipSpeedID, + } + body, err := json.Marshal(payload) + if err != nil { + t.Fatalf("marshal payload: %v", err) + } + req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(body)) + req.Header.Set("Authorization", adminToken) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + router.ServeHTTP(res, req) + + var out response.R + if err := json.NewDecoder(res.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code != 0 { + t.Fatalf("expected create success, got code=%d msg=%s", out.Code, out.Msg) + } + + var forwardID int64 + if err := r.DB().Raw("SELECT id FROM forward WHERE name = ?", "per-ip-speed-forward").Scan(&forwardID).Error; err != nil { + t.Fatalf("get forward ID: %v", err) + } + expectedName := fmt.Sprintf("rule_traffic_limit_%d", forwardID) + expectedLimits := []string{"$ 10.0MB 10.0MB", "0.0.0.0/0 5.0MB 5.0MB", "::/0 5.0MB 5.0MB"} + + commandMu.Lock() + defer commandMu.Unlock() + if addLimitersData == nil { + t.Fatalf("expected AddLimiters to be sent. Received: %v", receivedCommands) + } + if updateServiceData == nil { + t.Fatalf("expected UpdateService to be sent. Received: %v", receivedCommands) + } + + var addData map[string]interface{} + if err := json.Unmarshal(addLimitersData, &addData); err != nil { + t.Fatalf("unmarshal AddLimiters data: %v", err) + } + if addData["name"] != expectedName { + t.Fatalf("expected limiter name %s, got %v", expectedName, addData["name"]) + } + limits, ok := addData["limits"].([]interface{}) + if !ok { + t.Fatalf("expected limits array, got %T", addData["limits"]) + } + gotLimits := make([]string, 0, len(limits)) + for _, limit := range limits { + gotLimits = append(gotLimits, fmt.Sprint(limit)) + } + if !reflect.DeepEqual(gotLimits, expectedLimits) { + t.Fatalf("expected limits %v, got %v", expectedLimits, gotLimits) + } + + var services []map[string]interface{} + if err := json.Unmarshal(updateServiceData, &services); err != nil { + t.Fatalf("unmarshal UpdateService data: %v", err) + } + if len(services) == 0 { + t.Fatalf("expected services in UpdateService") + } + for _, service := range services { + if service["limiter"] != expectedName { + t.Fatalf("expected service limiter %s, got %v", expectedName, service["limiter"]) + } + } +}