feat: sync per-IP runtime limiters

This commit is contained in:
sagitchu
2026-04-27 22:32:35 +08:00
parent dec337d46b
commit 46394388b1
5 changed files with 337 additions and 45 deletions
+113 -36
View File
@@ -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
}
@@ -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
@@ -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))
}
@@ -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)
@@ -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"])
}
}
}