From 61c5b5e75961f17ccd735d27d00d03665fbe43c2 Mon Sep 17 00:00:00 2001 From: sagitchu Date: Thu, 26 Feb 2026 20:18:29 +0800 Subject: [PATCH] feat: add speed limit contract tests and refine limit/user UI --- IMPLEMENTATION_PLAN.md | 48 +- .../internal/http/handler/control_plane.go | 12 +- go-backend/internal/http/handler/mutations.go | 152 +++++- go-backend/internal/store/repo/repository.go | 5 +- .../store/repo/repository_federation.go | 8 +- .../store/repo/repository_mutations.go | 9 +- .../tests/contract/forward_contract_test.go | 139 ++++++ .../limiter_sync_failure_contract_test.go | 402 +++++++++++++++ .../contract/speed_limit_contract_test.go | 462 ++++++++++++++++++ vite-frontend/src/pages/forward.tsx | 72 ++- vite-frontend/src/pages/limit.tsx | 78 +-- vite-frontend/src/pages/user.tsx | 8 +- 12 files changed, 1304 insertions(+), 91 deletions(-) create mode 100644 go-backend/tests/contract/limiter_sync_failure_contract_test.go create mode 100644 go-backend/tests/contract/speed_limit_contract_test.go diff --git a/IMPLEMENTATION_PLAN.md b/IMPLEMENTATION_PLAN.md index 28b5ab2..ec1b3ff 100644 --- a/IMPLEMENTATION_PLAN.md +++ b/IMPLEMENTATION_PLAN.md @@ -12,6 +12,12 @@ ## 二、实施计划清单 +### 2.0 计划状态(审计更新:2026-02-26) + +- 总体状态:**进行中(未验收通过)** +- 已完成:模型、仓储查询、限速 CRUD、控制面优先级、限速页与类型改造、编译与测试通过 +- 未完成:**Forward 独立限速写入链路**(前端表单 -> API handler -> repository 落库 `forward.speed_id`) + ### 2.1 后端模型层 (Model) | 序号 | 任务 | 文件 | 状态 | @@ -80,6 +86,19 @@ |------|------|------| | B1 | Go 后端编译通过 | ✅ 完成 | | B2 | TypeScript 类型检查通过 | ✅ 完成 | +| B3 | `go test ./...` 全量通过 | ✅ 完成 | +| B4 | `go test ./tests/contract/... -run SpeedLimit` 通过 | ✅ 完成 | + +### 2.8 Forward 独立限速写入链路补全(新增) + +| 序号 | 任务 | 文件 | 状态 | +|------|------|------|------| +| N1 | forwardCreate 支持接收并校验可选 speedId,写入 Forward.SpeedID | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 | +| N2 | forwardUpdate 支持更新/清空 speedId,并触发服务重下发 | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 | +| N3 | CreateForwardTx 支持落库 speed_id | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 | +| N4 | UpdateForward 支持更新 speed_id | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 | +| N5 | Forward 页面新增限速选择并透传 speedId | `vite-frontend/src/pages/forward.tsx` | ✅ 完成 | +| N6 | Forward 相关契约测试补充 speedId 写入/清空断言 | `go-backend/tests/contract/forward_contract_test.go` | ✅ 完成 | --- @@ -100,23 +119,30 @@ ## 五、验证检查项 -### 5.1 功能验证 (待手动测试) +### 5.1 功能验证(审计后) -- [ ] 创建不限速规则的限速 (不绑定隧道) -- [ ] 创建绑定隧道的限速 (兼容旧逻辑) -- [ ] 编辑限速规则,切换隧道绑定状态 +- [x] 创建不限速规则的限速 (不绑定隧道) +- [x] 创建绑定隧道的限速 (兼容旧逻辑) +- [x] 编辑限速规则,切换隧道绑定状态 - [ ] 删除限速规则 - [ ] 转发列表正确显示 speedId -### 5.2 API 验证 (待手动测试) +### 5.2 API 验证(审计后) -- [ ] GET /api/speed-limit/list 返回可选 tunnelId -- [ ] POST /api/speed-limit/create 接受可选 tunnelId -- [ ] POST /api/speed-limit/update 接受可选 tunnelId +- [x] GET /api/speed-limit/list 返回可选 tunnelId +- [x] POST /api/speed-limit/create 接受可选 tunnelId +- [x] POST /api/speed-limit/update 接受可选 tunnelId - [ ] GET /api/forward/list 返回 speedId -### 5.3 兼容性验证 (待手动测试) +### 5.3 兼容性验证(审计后) -- [ ] 现有绑定隧道的限速规则继续正常工作 +- [x] 现有绑定隧道的限速规则继续正常工作 - [ ] 现有 UserTunnel 的限速继续正常工作 -- [ ] 备份/恢复功能正常 \ No newline at end of file +- [ ] 备份/恢复功能正常 + +### 5.4 Forward 独立限速闭环验证(新增) + +- [x] POST /api/forward/create 接受 speedId 并写入 `forward.speed_id` +- [x] POST /api/forward/update 可更新/清空 speedId +- [x] Forward 表单可选择限速并提交 speedId +- [ ] `syncForwardServices` 实际使用 Forward.SpeedID 而非仅回退 UserTunnel.SpeedID diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 8244958..17b7c23 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -185,7 +185,9 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all for _, fp := range ports { if limiterID != nil && speed != nil { - h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed) + if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil { + return err + } } node, err := h.getNodeRecord(fp.NodeID) @@ -1051,12 +1053,16 @@ func (h *Handler) sendDeleteLimiterConfig(limiterID int64, tunnelID int64) error return nil } -func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) { +func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) error { 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) + if _, err := h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false); err != nil { + return fmt.Errorf("限速规则下发失败: %w", err) + } + + return nil } diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 4d9ef31..9308287 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -1048,21 +1048,55 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("权限ID不能为空")) return } + + speedID := asAnyToInt64Ptr(req["speedId"]) + if err := h.validateSpeedLimitReference(speedID); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + userID, tunnelID, utErr := h.repo.GetUserTunnelUserAndTunnel(id) + if utErr != nil { + response.WriteJSON(w, response.Err(-2, utErr.Error())) + return + } + + _, oldFlow, oldNum, oldExpTime, oldFlowReset, oldSpeedID, oldStatus, oldErr := + h.repo.GetExistingUserTunnel(userID, tunnelID) + if oldErr != nil { + response.WriteJSON(w, response.Err(-2, oldErr.Error())) + return + } + if err := h.repo.UpdateUserTunnel(id, asInt64(req["flow"], 0), asInt(req["num"], 0), asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()), asInt64(req["flowResetTime"], 1), - nullableInt(asAnyToInt64Ptr(req["speedId"])), + nullableInt(speedID), asInt(req["status"], 1), ); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - userID, tunnelID, utErr := h.repo.GetUserTunnelUserAndTunnel(id) - if utErr == nil { - h.syncUserTunnelForwards(userID, tunnelID) + if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil { + rollbackErr := h.repo.UpdateUserTunnel( + id, + oldFlow, + int(oldNum), + oldExpTime, + oldFlowReset, + oldSpeedID, + oldStatus, + ) + if rollbackErr != nil { + response.WriteJSON(w, response.Err(-2, fmt.Sprintf("下发失败且回滚失败: %v; 回滚错误: %v", syncErr, rollbackErr))) + return + } + + response.WriteJSON(w, response.Err(-2, fmt.Sprintf("下发失败,已回滚: %v", syncErr))) + return } response.WriteJSON(w, response.OKEmpty()) @@ -1103,6 +1137,18 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空")) return } + speedID := asAnyToInt64Ptr(req["speedId"]) + if speedID != nil { + exists, speedErr := h.repo.SpeedLimitExists(*speedID) + if speedErr != nil { + response.WriteJSON(w, response.Err(-2, speedErr.Error())) + return + } + if !exists { + response.WriteJSON(w, response.ErrDefault("限速规则不存在")) + return + } + } port := asInt(req["inPort"], 0) if port <= 0 { port = h.pickTunnelPort(tunnelID) @@ -1127,7 +1173,7 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) { if userName == "" { userName = "user" } - forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port) + forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, nullableInt(speedID)) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -1202,6 +1248,24 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) { if strategy == "" { strategy = forward.Strategy } + speedID := asAnyToInt64Ptr(req["speedId"]) + if speedID != nil { + exists, speedErr := h.repo.SpeedLimitExists(*speedID) + if speedErr != nil { + response.WriteJSON(w, response.Err(-2, speedErr.Error())) + return + } + if !exists { + response.WriteJSON(w, response.ErrDefault("限速规则不存在")) + return + } + } + newSpeedID := forward.SpeedID + if speedID != nil { + newSpeedID = sql.NullInt64{Int64: *speedID, Valid: true} + } else if _, ok := req["speedId"]; ok { + newSpeedID = sql.NullInt64{Valid: false} + } port := asInt(req["inPort"], 0) if port <= 0 { @@ -1225,7 +1289,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) { } } now := time.Now().UnixMilli() - if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now); err != nil { + if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -3006,6 +3070,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts [] h.repo.RollbackForwardFields( oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name, oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status, + oldForward.SpeedID, time.Now().UnixMilli(), ) @@ -3027,6 +3092,10 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error { h.repo.GetExistingUserTunnel(userID, tunnelID) speedID := asAnyToInt64Ptr(req["speedId"]) + if err := h.validateSpeedLimitReference(speedID); err != nil { + return err + } + reqFlow := asInt64(req["flow"], -1) reqNum := asInt(req["num"], -1) reqExpTime := asInt64(req["expTime"], -1) @@ -3067,7 +3136,24 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error { reqStatus = 1 } - return h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus) + if err := h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus); err != nil { + return err + } + + if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil { + insertedID, _, _, _, _, _, _, lookupErr := h.repo.GetExistingUserTunnel(userID, tunnelID) + if lookupErr != nil { + return fmt.Errorf("下发失败且回滚失败: %v; 回滚查询错误: %w", syncErr, lookupErr) + } + + if rollbackErr := h.repo.DeleteUserTunnel(insertedID); rollbackErr != nil { + return fmt.Errorf("下发失败且回滚失败: %v; 回滚删除错误: %w", syncErr, rollbackErr) + } + + return fmt.Errorf("下发失败,已回滚: %w", syncErr) + } + + return nil } if err != nil { return err @@ -3105,25 +3191,61 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error { newSpeedID = sql.NullInt64{Valid: false} } - err = h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus) - - if err == nil { - h.syncUserTunnelForwards(userID, tunnelID) + if err := h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus); err != nil { + return err } - return err + + if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil { + rollbackErr := h.repo.UpdateUserTunnelFields( + existingID, + currentSpeedID, + currentFlow, + int(currentNum), + currentExpTime, + currentFlowReset, + currentStatus, + ) + if rollbackErr != nil { + return fmt.Errorf("下发失败且回滚失败: %v; 回滚错误: %w", syncErr, rollbackErr) + } + + return fmt.Errorf("下发失败,已回滚: %w", syncErr) + } + + return nil } -func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) { +func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) error { forwards, err := h.listForwardsByTunnel(tunnelID) if err != nil { - return + return err } for i := range forwards { f := &forwards[i] if f.UserID == userID { - _ = h.syncForwardServices(f, "UpdateService", true) + if err := h.syncForwardServices(f, "UpdateService", true); err != nil { + return err + } } } + + return nil +} + +func (h *Handler) validateSpeedLimitReference(speedID *int64) error { + if speedID == nil { + return nil + } + + exists, err := h.repo.SpeedLimitExists(*speedID) + if err != nil { + return err + } + if !exists { + return errors.New("限速规则不存在") + } + + return nil } func asAnySlice(v interface{}) []interface{} { diff --git a/go-backend/internal/store/repo/repository.go b/go-backend/internal/store/repo/repository.go index e4893f9..0e569eb 100644 --- a/go-backend/internal/store/repo/repository.go +++ b/go-backend/internal/store/repo/repository.go @@ -656,7 +656,7 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) { return nil, errors.New("repository not initialized") } var users []model.User - if err := r.db.Where("role_id != ?", 0).Order("id ASC").Find(&users).Error; err != nil { + if err := r.db.Where("role_id != ?", 0).Order("id DESC").Find(&users).Error; err != nil { return nil, err } items := make([]map[string]interface{}, 0, len(users)) @@ -678,7 +678,7 @@ func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) { return nil, errors.New("repository not initialized") } var limits []model.SpeedLimit - if err := r.db.Order("id ASC").Find(&limits).Error; err != nil { + if err := r.db.Order("id DESC").Find(&limits).Error; err != nil { return nil, err } items := make([]map[string]interface{}, 0, len(limits)) @@ -1319,7 +1319,6 @@ func (r *Repository) ListActiveForwardPeerShareRuntimesByNodeAndServiceName(node return items, nil } - func (r *Repository) ListActiveForwardPeerShareRuntimeServiceNamesByNode(nodeID int64) ([]string, error) { if r == nil || r.db == nil { return nil, errors.New("repository not initialized") diff --git a/go-backend/internal/store/repo/repository_federation.go b/go-backend/internal/store/repo/repository_federation.go index 438e7c0..8ec45e8 100644 --- a/go-backend/internal/store/repo/repository_federation.go +++ b/go-backend/internal/store/repo/repository_federation.go @@ -228,7 +228,6 @@ func (r *Repository) ListTunnelIDsByNamePrefix(prefix string) ([]int64, error) { return ids, nil } -// NextIndex returns COALESCE(MAX(inx), -1) + 1 for the given table. func (r *Repository) NextIndex(table string) int { if r == nil || r.db == nil { return 0 @@ -251,7 +250,7 @@ func (r *Repository) NextIndex(table string) int { var row inxRow err := r.db.Model(modelRef). Select("inx"). - Order("inx DESC"). + Order("inx ASC, id ASC"). Limit(1). Take(&row).Error if errors.Is(err, gorm.ErrRecordNotFound) { @@ -260,10 +259,7 @@ func (r *Repository) NextIndex(table string) int { if err != nil { return 0 } - if row.Inx < 0 { - return 0 - } - return row.Inx + 1 + return row.Inx - 1 } // CreateRemoteNode inserts a new remote node. diff --git a/go-backend/internal/store/repo/repository_mutations.go b/go-backend/internal/store/repo/repository_mutations.go index ca70a08..18f0b8e 100644 --- a/go-backend/internal/store/repo/repository_mutations.go +++ b/go-backend/internal/store/repo/repository_mutations.go @@ -657,7 +657,7 @@ func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 { return p } -func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64) error { +func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}) error { if r == nil || r.db == nil { return errors.New("repository not initialized") } @@ -668,6 +668,7 @@ func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remote "tunnel_id": tunnelID, "remote_addr": remoteAddr, "strategy": strategy, + "speed_id": nullInt64FromInterface(speedID), "updated_time": now, }).Error } @@ -724,7 +725,7 @@ func (r *Repository) ReplaceForwardPorts(forwardID int64, entries []struct { }) } -func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, now int64) { +func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, now int64) { if r == nil || r.db == nil { return } @@ -738,6 +739,7 @@ func (r *Repository) RollbackForwardFields(id, userID int64, userName, name stri "remote_addr": remoteAddr, "strategy": strategy, "status": status, + "speed_id": nullInt64FromInterface(speedID), "updated_time": now, }).Error } @@ -1205,7 +1207,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool, return ut.ID, true, nil } -func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int) (int64, error) { +func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, speedID interface{}) (int64, error) { if r == nil || r.db == nil { return 0, errors.New("repository not initialized") } @@ -1224,6 +1226,7 @@ func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnel UpdatedTime: now, Status: 1, Inx: inx, + SpeedID: nullInt64FromInterface(speedID), } if err := tx.Create(&fwd).Error; err != nil { return err diff --git a/go-backend/tests/contract/forward_contract_test.go b/go-backend/tests/contract/forward_contract_test.go index ffe0a56..490440d 100644 --- a/go-backend/tests/contract/forward_contract_test.go +++ b/go-backend/tests/contract/forward_contract_test.go @@ -2,6 +2,7 @@ package contract_test import ( "bytes" + "database/sql" "encoding/json" "net/http" "net/http/httptest" @@ -471,6 +472,144 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) { } } +func TestForwardSpeedIDWriteAndClearContracts(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(2, 'speed_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) + `, now, now).Error; err != nil { + t.Fatalf("insert user: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "forward-speed-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID := mustLastInsertID(t, repo, "forward-speed-tunnel") + + if err := repo.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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "forward-speed-node", "forward-speed-secret", "10.30.0.1", "10.30.0.1", "", "31000-31010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { + t.Fatalf("insert node: %v", err) + } + nodeID := mustLastInsertID(t, repo, "forward-speed-node") + + if err := repo.DB().Exec(` + INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, 1, ?, 31001, 'round', 1, 'tls') + `, tunnelID, nodeID).Error; err != nil { + t.Fatalf("insert chain_tunnel: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) + VALUES(?, ?, NULL, NULL, ?, NULL, ?) + `, "forward-speed-limit-a", 2048, now, 1).Error; err != nil { + t.Fatalf("insert speed limit a: %v", err) + } + speedIDA := mustLastInsertID(t, repo, "forward-speed-limit-a") + + if err := repo.DB().Exec(` + INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) + VALUES(?, ?, NULL, NULL, ?, NULL, ?) + `, "forward-speed-limit-b", 4096, now, 1).Error; err != nil { + t.Fatalf("insert speed limit b: %v", err) + } + speedIDB := mustLastInsertID(t, repo, "forward-speed-limit-b") + + server := httptest.NewServer(router) + defer server.Close() + stopNode := startMockNodeSession(t, server.URL, "forward-speed-secret") + defer stopNode() + + createPayload := map[string]interface{}{ + "name": "forward-speed-target", + "tunnelId": tunnelID, + "remoteAddr": "1.1.1.1:443", + "strategy": "fifo", + "speedId": speedIDA, + } + createBody, err := json.Marshal(createPayload) + if err != nil { + t.Fatalf("marshal create payload: %v", err) + } + createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody)) + createReq.Header.Set("Authorization", adminToken) + createReq.Header.Set("Content-Type", "application/json") + createRes := httptest.NewRecorder() + router.ServeHTTP(createRes, createReq) + assertCode(t, createRes, 0) + + forwardID := mustLastInsertID(t, repo, "forward-speed-target") + storedSpeed := repo.DB().Raw(`SELECT speed_id FROM forward WHERE id = ?`, forwardID).Row() + var createdSpeed sql.NullInt64 + if err := storedSpeed.Scan(&createdSpeed); err != nil { + t.Fatalf("query created forward speed_id: %v", err) + } + if !createdSpeed.Valid || createdSpeed.Int64 != speedIDA { + t.Fatalf("expected created speed_id=%d, got valid=%v value=%d", speedIDA, createdSpeed.Valid, createdSpeed.Int64) + } + + updateToBPayload := map[string]interface{}{ + "id": forwardID, + "speedId": speedIDB, + } + updateToBBody, err := json.Marshal(updateToBPayload) + if err != nil { + t.Fatalf("marshal update-to-b payload: %v", err) + } + updateToBReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateToBBody)) + updateToBReq.Header.Set("Authorization", adminToken) + updateToBReq.Header.Set("Content-Type", "application/json") + updateToBRes := httptest.NewRecorder() + router.ServeHTTP(updateToBRes, updateToBReq) + assertCode(t, updateToBRes, 0) + + storedSpeed = repo.DB().Raw(`SELECT speed_id FROM forward WHERE id = ?`, forwardID).Row() + var updatedSpeed sql.NullInt64 + if err := storedSpeed.Scan(&updatedSpeed); err != nil { + t.Fatalf("query updated forward speed_id: %v", err) + } + if !updatedSpeed.Valid || updatedSpeed.Int64 != speedIDB { + t.Fatalf("expected updated speed_id=%d, got valid=%v value=%d", speedIDB, updatedSpeed.Valid, updatedSpeed.Int64) + } + + clearPayload := map[string]interface{}{ + "id": forwardID, + "speedId": nil, + } + clearBody, err := json.Marshal(clearPayload) + if err != nil { + t.Fatalf("marshal clear payload: %v", err) + } + clearReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(clearBody)) + clearReq.Header.Set("Authorization", adminToken) + clearReq.Header.Set("Content-Type", "application/json") + clearRes := httptest.NewRecorder() + router.ServeHTTP(clearRes, clearReq) + assertCode(t, clearRes, 0) + + storedSpeed = repo.DB().Raw(`SELECT speed_id FROM forward WHERE id = ?`, forwardID).Row() + var clearedSpeed sql.NullInt64 + if err := storedSpeed.Scan(&clearedSpeed); err != nil { + t.Fatalf("query cleared forward speed_id: %v", err) + } + if clearedSpeed.Valid { + t.Fatalf("expected cleared speed_id to be NULL, got %d", clearedSpeed.Int64) + } +} + func jsonNumber(v int64) string { return strconv.FormatInt(v, 10) } diff --git a/go-backend/tests/contract/limiter_sync_failure_contract_test.go b/go-backend/tests/contract/limiter_sync_failure_contract_test.go new file mode 100644 index 0000000..72694ff --- /dev/null +++ b/go-backend/tests/contract/limiter_sync_failure_contract_test.go @@ -0,0 +1,402 @@ +package contract_test + +import ( + "bytes" + "database/sql" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "sync" + "testing" + "time" + + "github.com/gorilla/websocket" + + "go-backend/internal/auth" + "go-backend/internal/http/response" + "go-backend/internal/security" +) + +func TestForwardCreateRollbackWhenLimiterDispatchFailsContract(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "limiter-fail-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID := mustLastInsertID(t, r, "limiter-fail-tunnel") + + 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "limiter-fail-node", "limiter-fail-secret", "10.20.0.1", "10.20.0.1", "", "32000-32010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { + t.Fatalf("insert node: %v", err) + } + nodeID := mustLastInsertID(t, r, "limiter-fail-node") + + if err := r.DB().Exec(` + INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, 1, ?, 32001, 'round', 1, 'tls') + `, tunnelID, nodeID).Error; err != nil { + t.Fatalf("insert chain_tunnel: %v", err) + } + + if err := r.DB().Exec(` + INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) + VALUES(?, ?, NULL, NULL, ?, NULL, ?) + `, "limiter-fail-rule", 1024, now, 1).Error; err != nil { + t.Fatalf("insert speed limit: %v", err) + } + speedID := mustLastInsertID(t, r, "limiter-fail-rule") + + stopNode := startMockNodeSessionWithCommandFailures(t, server.URL, "limiter-fail-secret", map[string]string{ + "addlimiters": "mock add limiters failed", + }) + defer stopNode() + + payload := map[string]interface{}{ + "name": "limiter-fail-forward", + "tunnelId": tunnelID, + "remoteAddr": "1.1.1.1:443", + "strategy": "fifo", + "speedId": speedID, + } + 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 failure on limiter dispatch, got code=0") + } + + forwardCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM forward WHERE name = ?`, "limiter-fail-forward") + if forwardCount != 0 { + t.Fatalf("expected forward rollback delete on limiter failure, got count=%d", forwardCount) + } +} + +func TestBatchAssignRollbackWhenLimiterDispatchFailsContract(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 user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(2, 'assign_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) + `, now, now).Error; err != nil { + t.Fatalf("insert user: %v", err) + } + + if err := r.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "assign-limiter-fail-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID := mustLastInsertID(t, r, "assign-limiter-fail-tunnel") + + 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "assign-limiter-fail-node", "assign-limiter-fail-secret", "10.21.0.1", "10.21.0.1", "", "33000-33010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { + t.Fatalf("insert node: %v", err) + } + nodeID := mustLastInsertID(t, r, "assign-limiter-fail-node") + + if err := r.DB().Exec(` + INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, 1, ?, 33001, 'round', 1, 'tls') + `, tunnelID, nodeID).Error; err != nil { + t.Fatalf("insert chain_tunnel: %v", err) + } + + if err := r.DB().Exec(` + INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) + VALUES(?, ?, NULL, NULL, ?, NULL, ?) + `, "assign-limiter-fail-rule", 2048, now, 1).Error; err != nil { + t.Fatalf("insert speed limit: %v", err) + } + speedID := mustLastInsertID(t, r, "assign-limiter-fail-rule") + + if err := r.DB().Exec(` + INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) + VALUES(21, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1) + `, tunnelID).Error; err != nil { + t.Fatalf("insert user_tunnel: %v", err) + } + + if err := r.DB().Exec(` + INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(2, 'assign_user', 'assign-limiter-fail-forward', ?, '9.9.9.9:53', 'fifo', 0, 0, ?, ?, 1, 0) + `, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert forward: %v", err) + } + forwardID := mustLastInsertID(t, r, "assign-limiter-fail-forward") + + if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, 33001).Error; err != nil { + t.Fatalf("insert forward_port: %v", err) + } + + stopNode := startMockNodeSessionWithCommandFailures(t, server.URL, "assign-limiter-fail-secret", map[string]string{ + "addlimiters": "mock add limiters failed", + }) + defer stopNode() + + assignPayload := map[string]interface{}{ + "userId": 2, + "tunnels": []map[string]interface{}{{ + "tunnelId": tunnelID, + "speedId": speedID, + }}, + } + body, err := json.Marshal(assignPayload) + if err != nil { + t.Fatalf("marshal assign payload: %v", err) + } + req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/batch-assign", 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 assign failure on limiter dispatch, got code=0") + } + + var persistedSpeedID sql.NullInt64 + if err := r.DB().Raw(`SELECT speed_id FROM user_tunnel WHERE user_id = 2 AND tunnel_id = ?`, tunnelID).Row().Scan(&persistedSpeedID); err != nil { + t.Fatalf("query user_tunnel speed_id: %v", err) + } + if persistedSpeedID.Valid { + t.Fatalf("expected speed_id rollback to NULL, got %d", persistedSpeedID.Int64) + } +} + +func TestBatchAssignInsertRollbackWhenLimiterDispatchFailsContract(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 user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(3, 'assign_insert_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) + `, now, now).Error; err != nil { + t.Fatalf("insert user: %v", err) + } + + if err := r.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "assign-insert-limiter-fail-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID := mustLastInsertID(t, r, "assign-insert-limiter-fail-tunnel") + + 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "assign-insert-limiter-fail-node", "assign-insert-limiter-fail-secret", "10.22.0.1", "10.22.0.1", "", "34000-34010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { + t.Fatalf("insert node: %v", err) + } + nodeID := mustLastInsertID(t, r, "assign-insert-limiter-fail-node") + + if err := r.DB().Exec(` + INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, 1, ?, 34001, 'round', 1, 'tls') + `, tunnelID, nodeID).Error; err != nil { + t.Fatalf("insert chain_tunnel: %v", err) + } + + if err := r.DB().Exec(` + INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) + VALUES(?, ?, NULL, NULL, ?, NULL, ?) + `, "assign-insert-limiter-fail-rule", 3072, now, 1).Error; err != nil { + t.Fatalf("insert speed limit: %v", err) + } + speedID := mustLastInsertID(t, r, "assign-insert-limiter-fail-rule") + + if err := r.DB().Exec(` + INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(3, 'assign_insert_user', 'assign-insert-limiter-fail-forward', ?, '8.8.4.4:53', 'fifo', 0, 0, ?, ?, 1, 0) + `, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert forward: %v", err) + } + forwardID := mustLastInsertID(t, r, "assign-insert-limiter-fail-forward") + + if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, 34001).Error; err != nil { + t.Fatalf("insert forward_port: %v", err) + } + + stopNode := startMockNodeSessionWithCommandFailures(t, server.URL, "assign-insert-limiter-fail-secret", map[string]string{ + "addlimiters": "mock add limiters failed", + }) + defer stopNode() + + assignPayload := map[string]interface{}{ + "userId": 3, + "tunnels": []map[string]interface{}{{ + "tunnelId": tunnelID, + "speedId": speedID, + }}, + } + body, err := json.Marshal(assignPayload) + if err != nil { + t.Fatalf("marshal assign payload: %v", err) + } + req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/batch-assign", 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 assign(insert) failure on limiter dispatch, got code=0") + } + + insertedCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM user_tunnel WHERE user_id = 3 AND tunnel_id = ?`, tunnelID) + if insertedCount != 0 { + t.Fatalf("expected inserted user_tunnel rollback delete, got count=%d", insertedCount) + } +} + +func startMockNodeSessionWithCommandFailures(t *testing.T, baseURL string, nodeSecret string, failCommands map[string]string) func() { + t.Helper() + + u, err := url.Parse(baseURL) + if err != nil { + t.Fatalf("parse provider url: %v", err) + } + if strings.EqualFold(u.Scheme, "https") { + u.Scheme = "wss" + } else { + u.Scheme = "ws" + } + u.Path = "/system-info" + q := u.Query() + q.Set("type", "1") + q.Set("secret", nodeSecret) + q.Set("version", "v1") + q.Set("http", "1") + q.Set("tls", "1") + q.Set("socks", "1") + u.RawQuery = q.Encode() + + conn, _, err := websocket.DefaultDialer.Dial(u.String(), nil) + if err != nil { + t.Fatalf("dial mock node websocket: %v", err) + } + + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + for { + _, raw, readErr := conn.ReadMessage() + if readErr != nil { + return + } + + plain := raw + var wrap struct { + Encrypted bool `json:"encrypted"` + Data string `json:"data"` + } + if err := json.Unmarshal(raw, &wrap); err == nil && wrap.Encrypted && strings.TrimSpace(wrap.Data) != "" { + crypto, cryptoErr := security.NewAESCrypto(nodeSecret) + if cryptoErr == nil { + if dec, decErr := crypto.Decrypt(wrap.Data); decErr == nil { + plain = []byte(dec) + } + } + } + + var cmd struct { + Type string `json:"type"` + RequestID string `json:"requestId"` + } + if err := json.Unmarshal(plain, &cmd); err != nil { + continue + } + if strings.TrimSpace(cmd.RequestID) == "" { + continue + } + + cmdType := strings.TrimSpace(cmd.Type) + failMsg, shouldFail := failCommands[strings.ToLower(cmdType)] + + respType := fmt.Sprintf("%sResponse", cmdType) + respPayload := map[string]interface{}{ + "type": respType, + "success": !shouldFail, + "message": "OK", + "requestId": cmd.RequestID, + } + if shouldFail { + if strings.TrimSpace(failMsg) == "" { + failMsg = "mock command failed" + } + respPayload["message"] = failMsg + } + + respBytes, err := json.Marshal(respPayload) + if err != nil { + continue + } + _ = conn.WriteMessage(websocket.TextMessage, respBytes) + } + }() + + var stopOnce sync.Once + return func() { + stopOnce.Do(func() { + _ = conn.Close() + wg.Wait() + }) + } +} diff --git a/go-backend/tests/contract/speed_limit_contract_test.go b/go-backend/tests/contract/speed_limit_contract_test.go new file mode 100644 index 0000000..00f945b --- /dev/null +++ b/go-backend/tests/contract/speed_limit_contract_test.go @@ -0,0 +1,462 @@ +package contract_test + +import ( + "bytes" + "database/sql" + "encoding/json" + "net/http" + "net/http/httptest" + "path/filepath" + "testing" + "time" + + "go-backend/internal/auth" + "go-backend/internal/http/response" + "go-backend/internal/store/repo" +) + +// TestSpeedLimitWithoutTunnelContract tests that speed limits can be created without binding to a tunnel +func TestSpeedLimitWithoutTunnelContract(t *testing.T) { + secret := "contract-jwt-secret" + router, _ := setupContractRouter(t, secret) + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + // Create a speed limit without tunnel binding + t.Run("create speed limit without tunnel", func(t *testing.T) { + body := `{"name":"test-limit-no-tunnel","speed":100,"status":1}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/create", bytes.NewBufferString(body)) + req.Header.Set("Authorization", adminToken) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + router.ServeHTTP(res, req) + + assertCode(t, res, 0) + }) + + // Verify the speed limit has null tunnelId + t.Run("list speed limits shows null tunnelId", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil) + req.Header.Set("Authorization", adminToken) + 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 code 0, got %d", out.Code) + } + + data, ok := out.Data.([]interface{}) + if !ok { + t.Fatalf("expected data to be array, got %T", out.Data) + } + + // Find our speed limit + var found bool + for _, item := range data { + m, ok := item.(map[string]interface{}) + if !ok { + continue + } + if m["name"] == "test-limit-no-tunnel" { + found = true + // tunnelId should be nil/not present for unbound speed limits + if tunnelID, exists := m["tunnelId"]; exists && tunnelID != nil { + t.Fatalf("expected tunnelId to be nil for unbound speed limit, got %v", tunnelID) + } + break + } + } + + if !found { + t.Fatal("speed limit 'test-limit-no-tunnel' not found in list") + } + }) +} + +// TestSpeedLimitWithTunnelContract tests that speed limits can still be bound to tunnels +func TestSpeedLimitWithTunnelContract(t *testing.T) { + secret := "contract-jwt-secret" + router, r := setupContractRouter(t, secret) + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + // First create a tunnel + tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-tunnel-for-limit") + + // Create a speed limit with tunnel binding + t.Run("create speed limit with tunnel", func(t *testing.T) { + body := `{"name":"test-limit-with-tunnel","speed":200,"tunnelId":` + jsonInt(tunnelID) + `,"status":1}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/create", bytes.NewBufferString(body)) + req.Header.Set("Authorization", adminToken) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + router.ServeHTTP(res, req) + + assertCode(t, res, 0) + }) + + // Verify the speed limit has the tunnelId + t.Run("list speed limits shows tunnelId", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil) + req.Header.Set("Authorization", adminToken) + 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 code 0, got %d", out.Code) + } + + data, ok := out.Data.([]interface{}) + if !ok { + t.Fatalf("expected data to be array, got %T", out.Data) + } + + var found bool + for _, item := range data { + m, ok := item.(map[string]interface{}) + if !ok { + continue + } + if m["name"] == "test-limit-with-tunnel" { + found = true + tunnelIDVal, exists := m["tunnelId"] + if !exists || tunnelIDVal == nil { + t.Fatal("expected tunnelId to be present for bound speed limit") + } + // Verify tunnelId matches + if tunnelIDFloat, ok := tunnelIDVal.(float64); ok { + if int64(tunnelIDFloat) != tunnelID { + t.Fatalf("expected tunnelId %d, got %d", tunnelID, int64(tunnelIDFloat)) + } + } + break + } + } + + if !found { + t.Fatal("speed limit 'test-limit-with-tunnel' not found in list") + } + }) +} + +// TestSpeedLimitUpdateTunnelBindingContract tests updating speed limit tunnel binding +func TestSpeedLimitUpdateTunnelBindingContract(t *testing.T) { + secret := "contract-jwt-secret" + router, r := setupContractRouter(t, secret) + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + // Create a tunnel + tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-tunnel-update") + + // Create a speed limit without tunnel + speedLimitID := mustCreateSpeedLimitRepo(t, r, "test-limit-update", 0) + + // Update to bind to tunnel + t.Run("update speed limit to bind tunnel", func(t *testing.T) { + body := `{"id":` + jsonInt(speedLimitID) + `,"name":"test-limit-update","speed":150,"tunnelId":` + jsonInt(tunnelID) + `,"status":1}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/update", bytes.NewBufferString(body)) + req.Header.Set("Authorization", adminToken) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + router.ServeHTTP(res, req) + + assertCode(t, res, 0) + }) + + // Verify binding + t.Run("verify tunnel binding after update", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil) + req.Header.Set("Authorization", adminToken) + 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 code 0, got %d", out.Code) + } + + data, ok := out.Data.([]interface{}) + if !ok { + t.Fatalf("expected data to be array, got %T", out.Data) + } + + for _, item := range data { + m, ok := item.(map[string]interface{}) + if !ok { + continue + } + if m["name"] == "test-limit-update" { + tunnelIDVal, exists := m["tunnelId"] + if !exists || tunnelIDVal == nil { + t.Fatal("expected tunnelId to be present after update") + } + return + } + } + t.Fatal("speed limit 'test-limit-update' not found") + }) + + // Update to unbind from tunnel (set tunnelId to null) + t.Run("update speed limit to unbind tunnel", func(t *testing.T) { + body := `{"id":` + jsonInt(speedLimitID) + `,"name":"test-limit-update","speed":150,"status":1}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/update", bytes.NewBufferString(body)) + req.Header.Set("Authorization", adminToken) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + router.ServeHTTP(res, req) + + assertCode(t, res, 0) + }) + + // Verify unbinding + t.Run("verify tunnel unbinding after update", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil) + req.Header.Set("Authorization", adminToken) + 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 code 0, got %d", out.Code) + } + + data, ok := out.Data.([]interface{}) + if !ok { + t.Fatalf("expected data to be array, got %T", out.Data) + } + + for _, item := range data { + m, ok := item.(map[string]interface{}) + if !ok { + continue + } + if m["name"] == "test-limit-update" { + if tunnelIDVal, exists := m["tunnelId"]; exists && tunnelIDVal != nil { + t.Fatalf("expected tunnelId to be nil after unbinding, got %v", tunnelIDVal) + } + return + } + } + t.Fatal("speed limit 'test-limit-update' not found") + }) +} + +// TestSpeedLimitDatabaseNullableFields tests database-level nullable fields +func TestSpeedLimitDatabaseNullableFields(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "speed-limit-null.db") + r, err := repo.Open(dbPath) + if err != nil { + t.Fatalf("open sqlite: %v", err) + } + t.Cleanup(func() { _ = r.Close() }) + + // Create speed limit via repository + t.Run("repository create speed limit without tunnel", func(t *testing.T) { + id, err := r.CreateSpeedLimit("db-test-limit", 100, nil, "", 1, 1) + if err != nil { + t.Fatalf("CreateSpeedLimit failed: %v", err) + } + if id <= 0 { + t.Fatalf("expected valid id, got %d", id) + } + }) + + // Verify TunnelID is null in database + t.Run("verify null TunnelID in database", func(t *testing.T) { + var tunnelID sql.NullInt64 + var tunnelName sql.NullString + err := r.DB().Raw("SELECT tunnel_id, tunnel_name FROM speed_limit WHERE name = ?", "db-test-limit").Row().Scan(&tunnelID, &tunnelName) + if err != nil { + t.Fatalf("query failed: %v", err) + } + if tunnelID.Valid { + t.Fatalf("expected TunnelID to be NULL, got %d", tunnelID.Int64) + } + if tunnelName.Valid && tunnelName.String != "" { + t.Fatalf("expected TunnelName to be NULL or empty, got %s", tunnelName.String) + } + }) + + // Create a tunnel for binding test + tunnelID := mustCreateSpeedLimitTunnel(t, r, "db-test-tunnel") + + // Create speed limit with tunnel + t.Run("repository create speed limit with tunnel", func(t *testing.T) { + id, err := r.CreateSpeedLimit("db-test-limit-with-tunnel", 200, &tunnelID, "db-test-tunnel", 1, 1) + if err != nil { + t.Fatalf("CreateSpeedLimit failed: %v", err) + } + if id <= 0 { + t.Fatalf("expected valid id, got %d", id) + } + }) + + // Verify TunnelID is set + t.Run("verify TunnelID is set in database", func(t *testing.T) { + var dbTunnelID sql.NullInt64 + var dbTunnelName sql.NullString + err := r.DB().Raw("SELECT tunnel_id, tunnel_name FROM speed_limit WHERE name = ?", "db-test-limit-with-tunnel").Row().Scan(&dbTunnelID, &dbTunnelName) + if err != nil { + t.Fatalf("query failed: %v", err) + } + if !dbTunnelID.Valid { + t.Fatal("expected TunnelID to be valid") + } + if dbTunnelID.Int64 != tunnelID { + t.Fatalf("expected TunnelID %d, got %d", tunnelID, dbTunnelID.Int64) + } + if !dbTunnelName.Valid || dbTunnelName.String != "db-test-tunnel" { + t.Fatalf("expected TunnelName 'db-test-tunnel', got %v", dbTunnelName.String) + } + }) + + // Test GetSpeedLimitTunnelID returns correct nullability + t.Run("GetSpeedLimitTunnelID returns null for unbound limit", func(t *testing.T) { + result := r.GetSpeedLimitTunnelID(1) // First speed limit (db-test-limit) + if result.Valid { + t.Fatalf("expected GetSpeedLimitTunnelID to return invalid/null, got valid with value %d", result.Int64) + } + }) + + t.Run("GetSpeedLimitTunnelID returns value for bound limit", func(t *testing.T) { + result := r.GetSpeedLimitTunnelID(2) // Second speed limit (db-test-limit-with-tunnel) + if !result.Valid { + t.Fatal("expected GetSpeedLimitTunnelID to return valid result for bound limit") + } + if result.Int64 != tunnelID { + t.Fatalf("expected TunnelID %d, got %d", tunnelID, result.Int64) + } + }) +} + +// TestSpeedLimitUpdateUnbindFromTunnel tests unbinding a speed limit from a tunnel +func TestSpeedLimitUpdateUnbindFromTunnel(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "speed-limit-unbind.db") + r, err := repo.Open(dbPath) + if err != nil { + t.Fatalf("open sqlite: %v", err) + } + t.Cleanup(func() { _ = r.Close() }) + + // Create tunnel + tunnelID := mustCreateSpeedLimitTunnel(t, r, "unbind-test-tunnel") + + // Create speed limit bound to tunnel + speedLimitID, err := r.CreateSpeedLimit("unbind-test-limit", 300, &tunnelID, "unbind-test-tunnel", 1, 1) + if err != nil { + t.Fatalf("create speed limit: %v", err) + } + + // Verify initial binding + t.Run("verify initial binding", func(t *testing.T) { + result := r.GetSpeedLimitTunnelID(speedLimitID) + if !result.Valid { + t.Fatal("expected initial binding to tunnel") + } + if result.Int64 != tunnelID { + t.Fatalf("expected tunnel ID %d, got %d", tunnelID, result.Int64) + } + }) + + // Update to unbind + t.Run("unbind speed limit from tunnel via UpdateSpeedLimit", func(t *testing.T) { + err := r.UpdateSpeedLimit(speedLimitID, "unbind-test-limit", 300, nil, "", 1, time.Now().UnixMilli()) + if err != nil { + t.Fatalf("UpdateSpeedLimit failed: %v", err) + } + }) + + // Verify unbinding + t.Run("verify unbinding after update", func(t *testing.T) { + result := r.GetSpeedLimitTunnelID(speedLimitID) + if result.Valid { + t.Fatalf("expected GetSpeedLimitTunnelID to return invalid/null after unbind, got valid with value %d", result.Int64) + } + }) +} + +// TestSpeedLimitGetSpeed tests the GetSpeedLimitSpeed function +func TestSpeedLimitGetSpeed(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "speed-limit-getspeed.db") + r, err := repo.Open(dbPath) + if err != nil { + t.Fatalf("open sqlite: %v", err) + } + t.Cleanup(func() { _ = r.Close() }) + + // Create speed limit + speedLimitID, err := r.CreateSpeedLimit("get-speed-test", 500, nil, "", 1, 1) + if err != nil { + t.Fatalf("create speed limit: %v", err) + } + + // Test GetSpeedLimitSpeed + t.Run("GetSpeedLimitSpeed returns correct speed", func(t *testing.T) { + speed, err := r.GetSpeedLimitSpeed(speedLimitID) + if err != nil { + t.Fatalf("GetSpeedLimitSpeed failed: %v", err) + } + if speed != 500 { + t.Fatalf("expected speed 500, got %d", speed) + } + }) + + t.Run("GetSpeedLimitSpeed returns error for non-existent id", func(t *testing.T) { + _, err := r.GetSpeedLimitSpeed(99999) + if err == nil { + t.Fatal("expected error for non-existent speed limit ID") + } + }) +} + +// Helper functions + +func mustCreateSpeedLimitTunnel(t *testing.T, r *repo.Repository, name string) int64 { + t.Helper() + 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(?, 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) + `, name, now, now).Error; err != nil { + t.Fatalf("create tunnel failed: %v", err) + } + return mustLastInsertID(t, r, name) +} + +func mustCreateSpeedLimitRepo(t *testing.T, r *repo.Repository, name string, tunnelID int64) int64 { + t.Helper() + now := time.Now().UnixMilli() + var tid *int64 + if tunnelID > 0 { + tid = &tunnelID + } + id, err := r.CreateSpeedLimit(name, 100, tid, "", now, 1) + if err != nil { + t.Fatalf("create speed limit failed: %v", err) + } + return id +} diff --git a/vite-frontend/src/pages/forward.tsx b/vite-frontend/src/pages/forward.tsx index 67fb061..b9d1226 100644 --- a/vite-frontend/src/pages/forward.tsx +++ b/vite-frontend/src/pages/forward.tsx @@ -1,3 +1,5 @@ +import type { SpeedLimitApiItem } from "@/api/types"; + import { useState, useEffect, useMemo } from "react"; import toast from "react-hot-toast"; import { @@ -50,6 +52,7 @@ import { Checkbox } from "@/shadcn-bridge/heroui/checkbox"; import { createForward, getForwardList, + getSpeedLimitList, getPeerShareList, getPeerRemoteUsageList, updateForward, @@ -105,6 +108,7 @@ interface Forward { userName?: string; userId?: number; inx?: number; + speedId?: number | null; } interface Tunnel { @@ -123,12 +127,14 @@ interface ForwardForm { remoteAddr: string; interfaceName?: string; strategy: string; + speedId: number | null; } export default function ForwardPage() { const [loading, setLoading] = useState(true); const [forwards, setForwards] = useState([]); const [tunnels, setTunnels] = useState([]); + const [speedLimits, setSpeedLimits] = useState([]); const isMobile = useMobileBreakpoint(); const [searchKeyword, setSearchKeyword] = useLocalStorageState( "forward-search-keyword", @@ -206,6 +212,7 @@ export default function ForwardPage() { remoteAddr: "", interfaceName: "", strategy: "fifo", + speedId: null, }); // 表单验证错误 @@ -327,7 +334,9 @@ export default function ForwardPage() { const resolveShareIdForForward = (forward: Forward): number | null => { const candidates = new Set(); - const shareIdFromName = parseShareIdFromTunnelName(forward.tunnelName || ""); + const shareIdFromName = parseShareIdFromTunnelName( + forward.tunnelName || "", + ); if (shareIdFromName) { candidates.add(shareIdFromName); @@ -445,9 +454,10 @@ export default function ForwardPage() { const loadData = async (lod = true) => { setLoading(lod); try { - const [forwardsRes, tunnelsRes] = await Promise.all([ + const [forwardsRes, tunnelsRes, speedLimitsRes] = await Promise.all([ getForwardList(), userTunnel(), + getSpeedLimitList(), ]); if (forwardsRes.code === 0) { @@ -481,6 +491,10 @@ export default function ForwardPage() { setTunnels(tunnelsRes.data || []); } else { } + + if (speedLimitsRes.code === 0) { + setSpeedLimits(speedLimitsRes.data || []); + } } catch { toast.error("加载数据失败"); } finally { @@ -489,6 +503,10 @@ export default function ForwardPage() { }; // 表单验证 + const availableSpeedLimits = useMemo(() => { + return speedLimits; + }, [speedLimits]); + const validateForm = (): boolean => { const newErrors: { [key: string]: string } = {}; @@ -555,6 +573,7 @@ export default function ForwardPage() { remoteAddr: "", interfaceName: "", strategy: "fifo", + speedId: null, }); setErrors({}); setModalOpen(true); @@ -572,6 +591,7 @@ export default function ForwardPage() { remoteAddr: forward.remoteAddr.split(",").join("\n"), interfaceName: forward.interfaceName || "", strategy: forward.strategy || "fifo", + speedId: forward.speedId ?? null, }); setErrors({}); setModalOpen(true); @@ -651,6 +671,7 @@ export default function ForwardPage() { inPort: form.inPort, remoteAddr: processedRemoteAddr, strategy: addressCount > 1 ? form.strategy : "fifo", + speedId: form.speedId, }; res = await updateForward(updateData); @@ -662,6 +683,7 @@ export default function ForwardPage() { inPort: form.inPort, remoteAddr: processedRemoteAddr, strategy: addressCount > 1 ? form.strategy : "fifo", + speedId: form.speedId, }; res = await createForward(createData); @@ -1346,7 +1368,11 @@ export default function ForwardPage() { const aInx = a.inx ?? 0; const bInx = b.inx ?? 0; - return aInx - bInx; + if (aInx !== bInx) { + return aInx - bInx; + } + + return (a.id ?? 0) - (b.id ?? 0); }); // 如果数据库中没有排序信息,则使用本地存储的顺序 @@ -1511,6 +1537,9 @@ export default function ForwardPage() { {forward.userName || "未知用户"} + + {forward.name} + - - {forward.name} -