diff --git a/IMPLEMENTATION_PLAN.md b/IMPLEMENTATION_PLAN.md new file mode 100644 index 0000000..ec1b3ff --- /dev/null +++ b/IMPLEMENTATION_PLAN.md @@ -0,0 +1,148 @@ +# 限速功能重构实施计划 + +## 一、需求概述 + +**原始需求**: 限速功能当前绑定到具体隧道,需要改为不绑定隧道,创建限速后可以自由在隧道上限速,也可以在转发上限速。 + +**核心变更**: +1. 限速规则(SpeedLimit)与隧道的绑定关系改为可选 +2. 转发(Forward)支持独立的限速规则 + +--- + +## 二、实施计划清单 + +### 2.0 计划状态(审计更新:2026-02-26) + +- 总体状态:**进行中(未验收通过)** +- 已完成:模型、仓储查询、限速 CRUD、控制面优先级、限速页与类型改造、编译与测试通过 +- 未完成:**Forward 独立限速写入链路**(前端表单 -> API handler -> repository 落库 `forward.speed_id`) + +### 2.1 后端模型层 (Model) + +| 序号 | 任务 | 文件 | 状态 | +|------|------|------|------| +| M1 | SpeedLimit.TunnelID 改为 sql.NullInt64 (可空) | `go-backend/internal/store/model/model.go` | ✅ 完成 | +| M2 | SpeedLimit.TunnelName 改为 sql.NullString (可空) | `go-backend/internal/store/model/model.go` | ✅ 完成 | +| M3 | Forward 添加 SpeedID sql.NullInt64 字段 | `go-backend/internal/store/model/model.go` | ✅ 完成 | +| M4 | ForwardRecord 添加 SpeedID sql.NullInt64 字段 | `go-backend/internal/store/model/model.go` | ✅ 完成 | +| M5 | SpeedLimitBackup.TunnelID 改为指针类型 | `go-backend/internal/store/model/model.go` | ✅ 完成 | +| M6 | ForwardBackup 添加 SpeedID *int64 字段 | `go-backend/internal/store/model/model.go` | ✅ 完成 | + +### 2.2 后端仓储层 (Repository) + +| 序号 | 任务 | 文件 | 状态 | +|------|------|------|------| +| R1 | ListSpeedLimits() 返回可空 tunnelId/tunnelName | `go-backend/internal/store/repo/repository.go` | ✅ 完成 | +| R2 | ListForwards() 返回 speedId 字段 | `go-backend/internal/store/repo/repository.go` | ✅ 完成 | +| R3 | CreateSpeedLimit() 参数 tunnelID 改为 *int64 | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 | +| R4 | UpdateSpeedLimit() 参数 tunnelID 改为 *int64 | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 | +| R5 | GetSpeedLimitTunnelID() 返回 sql.NullInt64 | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 | +| R6 | exportSpeedLimits() 处理可空字段 | `go-backend/internal/store/repo/repository.go` | ✅ 完成 | +| R7 | importSpeedLimits() 处理可空字段 | `go-backend/internal/store/repo/repository.go` | ✅ 完成 | +| R8 | GetSpeedLimitSpeed() 新增方法 | `go-backend/internal/store/repo/repository_flow.go` | ✅ 完成 | +| R9 | ListForwardsByTunnel() 返回 SpeedID | `go-backend/internal/store/repo/repository_control.go` | ✅ 完成 | +| R10 | ListActiveForwardsByUser() 返回 SpeedID | `go-backend/internal/store/repo/repository_flow.go` | ✅ 完成 | +| R11 | ListActiveForwardsByUserTunnel() 返回 SpeedID | `go-backend/internal/store/repo/repository_flow.go` | ✅ 完成 | +| R12 | GetForwardRecord() 返回 SpeedID | `go-backend/internal/store/repo/repository_flow.go` | ✅ 完成 | + +### 2.3 后端处理器层 (Handler) + +| 序号 | 任务 | 文件 | 状态 | +|------|------|------|------| +| H1 | speedLimitCreate 处理可选 tunnelId | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 | +| H2 | speedLimitUpdate 处理可选 tunnelId | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 | +| H3 | speedLimitDelete 处理可空 tunnelID | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 | + +### 2.4 后端控制平面 (Control Plane) + +| 序号 | 任务 | 文件 | 状态 | +|------|------|------|------| +| C1 | syncForwardServices 优先使用 Forward.SpeedID | `go-backend/internal/http/handler/control_plane.go` | ✅ 完成 | +| C2 | 回退到 UserTunnel 的 speed limit | `go-backend/internal/http/handler/control_plane.go` | ✅ 完成 | + +### 2.5 前端类型定义 (TypeScript Types) + +| 序号 | 任务 | 文件 | 状态 | +|------|------|------|------| +| T1 | SpeedLimitApiItem.tunnelId 改为可选 | `vite-frontend/src/api/types.ts` | ✅ 完成 | +| T2 | ForwardApiItem 添加 speedId 字段 | `vite-frontend/src/api/types.ts` | ✅ 完成 | +| T3 | ForwardMutationPayload 添加 speedId 字段 | `vite-frontend/src/api/types.ts` | ✅ 完成 | +| T4 | SpeedLimitMutationPayload.tunnelId 改为可选 | `vite-frontend/src/api/types.ts` | ✅ 完成 | + +### 2.6 前端页面组件 + +| 序号 | 任务 | 文件 | 状态 | +|------|------|------|------| +| F1 | SpeedLimitRule 接口更新 | `vite-frontend/src/pages/limit.tsx` | ✅ 完成 | +| F2 | SpeedLimitForm 接口更新 | `vite-frontend/src/pages/limit.tsx` | ✅ 完成 | +| F3 | validateForm 移除 tunnelId 必填校验 | `vite-frontend/src/pages/limit.tsx` | ✅ 完成 | +| F4 | Select 组件改为可选 | `vite-frontend/src/pages/limit.tsx` | ✅ 完成 | +| F5 | 显示"未绑定"状态 | `vite-frontend/src/pages/limit.tsx` | ✅ 完成 | + +### 2.7 编译验证 + +| 序号 | 任务 | 状态 | +|------|------|------| +| 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` | ✅ 完成 | + +--- + +## 三、优先级说明 + +限速规则应用优先级: +1. **Forward.SpeedID** - 转发级别的限速 (最高优先) +2. **UserTunnel.SpeedID** - 用户隧道权限级别的限速 (回退) + +--- + +## 四、数据库兼容性 + +- SpeedLimit 表: `tunnel_id` 和 `tunnel_name` 字段改为可空 (GORM AutoMigrate 自动处理) +- Forward 表: 新增 `speed_id` 可空字段 (GORM AutoMigrate 自动处理) + +--- + +## 五、验证检查项 + +### 5.1 功能验证(审计后) + +- [x] 创建不限速规则的限速 (不绑定隧道) +- [x] 创建绑定隧道的限速 (兼容旧逻辑) +- [x] 编辑限速规则,切换隧道绑定状态 +- [ ] 删除限速规则 +- [ ] 转发列表正确显示 speedId + +### 5.2 API 验证(审计后) + +- [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 兼容性验证(审计后) + +- [x] 现有绑定隧道的限速规则继续正常工作 +- [ ] 现有 UserTunnel 的限速继续正常工作 +- [ ] 备份/恢复功能正常 + +### 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 e34aff3..17b7c23 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -152,11 +152,32 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all return errors.New("转发入口端口不存在") } - userTunnelID, limiterID, speed, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) - if err != nil { - return err + // Determine limiter from forward's SpeedID first, fallback to UserTunnel's limiter + var limiterID *int64 + var speed *int + + if forward.SpeedID.Valid && forward.SpeedID.Int64 > 0 { + // Forward has its own speed limit + speedVal, err := h.repo.GetSpeedLimitSpeed(forward.SpeedID.Int64) + if err == nil && speedVal > 0 { + limiterID = &forward.SpeedID.Int64 + speed = &speedVal + } } - serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID) + + if limiterID == nil { + // Fall back to UserTunnel speed limit + var utLimiterID *int64 + var utSpeed *int + _, utLimiterID, utSpeed, err = h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) + if err != nil { + return err + } + limiterID = utLimiterID + speed = utSpeed + } + + serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, 0) tunnelTLSProtocol, err := h.isTunnelSelectedTLSProtocol(forward.TunnelID) if err != nil { return err @@ -164,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) @@ -1030,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 c0132d5..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 } @@ -1586,29 +1650,37 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("请求参数错误")) return } - tunnelID := asInt64(req["tunnelId"], 0) - if tunnelID <= 0 { - response.WriteJSON(w, response.ErrDefault("隧道ID不能为空")) - return - } + name := asString(req["name"]) if name == "" { response.WriteJSON(w, response.ErrDefault("名称不能为空")) return } - tunnelName := h.repo.GetTunnelNameByID(tunnelID) - if tunnelName == "" { - response.WriteJSON(w, response.ErrDefault("隧道不存在")) - return - } - now := time.Now().UnixMilli() + speed := asInt(req["speed"], 100) + + var tunnelID *int64 + var tunnelName string + if tid := asInt64(req["tunnelId"], 0); tid > 0 { + tunnelID = &tid + tunnelName = h.repo.GetTunnelNameByID(tid) + if tunnelName == "" { + response.WriteJSON(w, response.ErrDefault("隧道不存在")) + return + } + } + + now := time.Now().UnixMilli() id, err := h.repo.CreateSpeedLimit(name, speed, tunnelID, tunnelName, now, asInt(req["status"], 1)) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - _ = h.sendLimiterConfig(id, speed, tunnelID) + + if tunnelID != nil && *tunnelID > 0 { + _ = h.sendLimiterConfig(id, speed, *tunnelID) + } + response.WriteJSON(w, response.OKEmpty()) } @@ -1618,23 +1690,41 @@ func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("请求参数错误")) return } + id := asInt64(req["id"], 0) - tunnelID := asInt64(req["tunnelId"], 0) - if id <= 0 || tunnelID <= 0 { - response.WriteJSON(w, response.ErrDefault("请求参数错误")) + if id <= 0 { + response.WriteJSON(w, response.ErrDefault("限速规则ID不能为空")) return } - tunnelName := h.repo.GetTunnelNameByID(tunnelID) - if tunnelName == "" { - response.WriteJSON(w, response.ErrDefault("隧道不存在")) + + name := asString(req["name"]) + if name == "" { + response.WriteJSON(w, response.ErrDefault("名称不能为空")) return } + speed := asInt(req["speed"], 100) - if err := h.repo.UpdateSpeedLimit(id, asString(req["name"]), speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli()); err != nil { + + var tunnelID *int64 + var tunnelName string + if tid := asInt64(req["tunnelId"], 0); tid > 0 { + tunnelID = &tid + tunnelName = h.repo.GetTunnelNameByID(tid) + if tunnelName == "" { + response.WriteJSON(w, response.ErrDefault("隧道不存在")) + return + } + } + + if err := h.repo.UpdateSpeedLimit(id, name, speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli()); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - _ = h.sendLimiterConfig(id, speed, tunnelID) + + if tunnelID != nil && *tunnelID > 0 { + _ = h.sendLimiterConfig(id, speed, *tunnelID) + } + response.WriteJSON(w, response.OKEmpty()) } @@ -1643,15 +1733,18 @@ func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) { if id <= 0 { return } + tunnelID := h.repo.GetSpeedLimitTunnelID(id) if err := h.repo.DeleteSpeedLimit(id); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if tunnelID > 0 { - _ = h.sendDeleteLimiterConfig(id, tunnelID) + + if tunnelID.Valid && tunnelID.Int64 > 0 { + _ = h.sendDeleteLimiterConfig(id, tunnelID.Int64) } + response.WriteJSON(w, response.OKEmpty()) } @@ -2977,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(), ) @@ -2998,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) @@ -3038,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 @@ -3076,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/model/model.go b/go-backend/internal/store/model/model.go index dd2bc10..dd06073 100644 --- a/go-backend/internal/store/model/model.go +++ b/go-backend/internal/store/model/model.go @@ -29,19 +29,20 @@ func (User) TableName() string { return "user" } // Forward maps to the "forward" table. type Forward struct { - ID int64 `gorm:"primaryKey;autoIncrement"` - UserID int64 `gorm:"column:user_id;not null"` - UserName string `gorm:"column:user_name;type:varchar(100);not null"` - Name string `gorm:"type:varchar(100);not null"` - TunnelID int64 `gorm:"column:tunnel_id;not null"` - RemoteAddr string `gorm:"column:remote_addr;type:text;not null"` - Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"` - InFlow int64 `gorm:"column:in_flow;not null;default:0"` - OutFlow int64 `gorm:"column:out_flow;not null;default:0"` - CreatedTime int64 `gorm:"column:created_time;not null"` - UpdatedTime int64 `gorm:"column:updated_time;not null"` - Status int `gorm:"not null"` - Inx int `gorm:"not null;default:0"` + ID int64 `gorm:"primaryKey;autoIncrement"` + UserID int64 `gorm:"column:user_id;not null"` + UserName string `gorm:"column:user_name;type:varchar(100);not null"` + Name string `gorm:"type:varchar(100);not null"` + TunnelID int64 `gorm:"column:tunnel_id;not null"` + RemoteAddr string `gorm:"column:remote_addr;type:text;not null"` + Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"` + InFlow int64 `gorm:"not null;default:0"` + OutFlow int64 `gorm:"column:out_flow;not null;default:0"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime int64 `gorm:"column:updated_time;not null"` + Status int `gorm:"not null"` + Inx int `gorm:"not null;default:0"` + SpeedID sql.NullInt64 `gorm:"column:speed_id"` } func (Forward) TableName() string { return "forward" } @@ -83,14 +84,14 @@ type Node struct { func (Node) TableName() string { return "node" } type SpeedLimit struct { - ID int64 `gorm:"primaryKey;autoIncrement"` - Name string `gorm:"type:varchar(100);not null"` - Speed int `gorm:"not null"` - TunnelID int64 `gorm:"column:tunnel_id;not null"` - TunnelName string `gorm:"column:tunnel_name;type:varchar(100);not null"` - CreatedTime int64 `gorm:"column:created_time;not null"` - UpdatedTime sql.NullInt64 `gorm:"column:updated_time"` - Status int `gorm:"not null"` + ID int64 `gorm:"primaryKey;autoIncrement"` + Name string `gorm:"type:varchar(100);not null"` + Speed int `gorm:"not null"` + TunnelID sql.NullInt64 `gorm:"column:tunnel_id"` + TunnelName sql.NullString `gorm:"column:tunnel_name;type:varchar(100)"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime sql.NullInt64 `gorm:"column:updated_time"` + Status int `gorm:"not null"` } func (SpeedLimit) TableName() string { return "speed_limit" } @@ -395,6 +396,7 @@ type ForwardBackup struct { UpdatedTime int64 `json:"updatedTime"` Status int `json:"status"` Inx int `json:"inx"` + SpeedID *int64 `json:"speedId,omitempty"` ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"` } @@ -421,8 +423,8 @@ type SpeedLimitBackup struct { ID int64 `json:"id"` Name string `json:"name"` Speed int64 `json:"speed"` - TunnelID int64 `json:"tunnelId"` - TunnelName string `json:"tunnelName"` + TunnelID *int64 `json:"tunnelId,omitempty"` + TunnelName string `json:"tunnelName,omitempty"` CreatedTime int64 `json:"createdTime"` UpdatedTime int64 `json:"updatedTime,omitempty"` Status int `json:"status"` @@ -492,6 +494,7 @@ type ForwardRecord struct { RemoteAddr string Strategy string Status int + SpeedID sql.NullInt64 } // TunnelRecord is a minimal tunnel view used by control plane. diff --git a/go-backend/internal/store/repo/repository.go b/go-backend/internal/store/repo/repository.go index c1d2b04..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,17 +678,23 @@ 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)) for _, sl := range limits { - items = append(items, map[string]interface{}{ + item := map[string]interface{}{ "id": sl.ID, "name": sl.Name, "speed": sl.Speed, - "tunnelId": sl.TunnelID, "tunnelName": sl.TunnelName, "status": sl.Status, "createdTime": sl.CreatedTime, "updatedTime": nullableInt64(sl.UpdatedTime), - }) + } + if sl.TunnelID.Valid { + item["tunnelId"] = sl.TunnelID.Int64 + } + if sl.TunnelName.Valid { + item["tunnelName"] = sl.TunnelName.String + } + items = append(items, item) } return items, nil } @@ -712,11 +718,12 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) { CreatedTime int64 Status int Inx int + SpeedID sql.NullInt64 } var rows []fwdRow err := r.db.Model(&model.Forward{}). - Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx"). + Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id"). Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id"). Order("forward.inx ASC, forward.id ASC"). Find(&rows).Error @@ -730,14 +737,18 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) { if err != nil { return nil, err } - items = append(items, map[string]interface{}{ + item := map[string]interface{}{ "id": row.ID, "userId": row.UserID, "userName": row.UserName, "name": row.Name, "tunnelId": row.TunnelID, "tunnelName": row.TunnelName, "inIp": nullableForwardIngress(inIP), "inPort": nullableInt64(inPort), "remoteAddr": row.RemoteAddr, "strategy": row.Strategy, "inFlow": row.InFlow, "outFlow": row.OutFlow, "createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx), - }) + } + if row.SpeedID.Valid { + item["speedId"] = row.SpeedID.Int64 + } + items = append(items, item) } return items, nil } @@ -1308,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") @@ -1813,9 +1823,15 @@ func (r *Repository) exportSpeedLimits() ([]model.SpeedLimitBackup, error) { for _, sl := range sls { b := model.SpeedLimitBackup{ ID: sl.ID, Name: sl.Name, Speed: int64(sl.Speed), - TunnelID: sl.TunnelID, TunnelName: sl.TunnelName, CreatedTime: sl.CreatedTime, Status: sl.Status, } + if sl.TunnelID.Valid { + tid := sl.TunnelID.Int64 + b.TunnelID = &tid + } + if sl.TunnelName.Valid { + b.TunnelName = sl.TunnelName.String + } if sl.UpdatedTime.Valid { b.UpdatedTime = sl.UpdatedTime.Int64 } @@ -2186,12 +2202,18 @@ func importSpeedLimits(tx *gorm.DB, speedLimits []model.SpeedLimitBackup, now in ID: sl.ID, Name: sl.Name, Speed: int(sl.Speed), - TunnelID: sl.TunnelID, - TunnelName: sl.TunnelName, + TunnelID: sql.NullInt64{Int64: 0, Valid: false}, + TunnelName: sql.NullString{String: "", Valid: false}, CreatedTime: sl.CreatedTime, UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, Status: sl.Status, } + if sl.TunnelID != nil { + item.TunnelID = sql.NullInt64{Int64: *sl.TunnelID, Valid: true} + } + if sl.TunnelName != "" { + item.TunnelName = sql.NullString{String: sl.TunnelName, Valid: true} + } err := tx.Clauses(clause.OnConflict{ Columns: []clause.Column{{Name: "id"}}, DoUpdates: clause.AssignmentColumns([]string{ diff --git a/go-backend/internal/store/repo/repository_control.go b/go-backend/internal/store/repo/repository_control.go index 2aed13c..680fae8 100644 --- a/go-backend/internal/store/repo/repository_control.go +++ b/go-backend/internal/store/repo/repository_control.go @@ -46,6 +46,7 @@ func (r *Repository) ListForwardsByTunnel(tunnelID int64) ([]model.ForwardRecord RemoteAddr: f.RemoteAddr, Strategy: f.Strategy, Status: f.Status, + SpeedID: f.SpeedID, }) } for i := range rows { 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_flow.go b/go-backend/internal/store/repo/repository_flow.go index cacc55c..15c9762 100644 --- a/go-backend/internal/store/repo/repository_flow.go +++ b/go-backend/internal/store/repo/repository_flow.go @@ -38,6 +38,7 @@ func (r *Repository) ListActiveForwardsByUser(userID int64) ([]model.ForwardReco RemoteAddr: f.RemoteAddr, Strategy: f.Strategy, Status: f.Status, + SpeedID: f.SpeedID, }) } for i := range rows { @@ -68,6 +69,7 @@ func (r *Repository) ListActiveForwardsByUserTunnel(userID, tunnelID int64) ([]m RemoteAddr: f.RemoteAddr, Strategy: f.Strategy, Status: f.Status, + SpeedID: f.SpeedID, }) } for i := range rows { @@ -99,6 +101,7 @@ func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, er RemoteAddr: f.RemoteAddr, Strategy: f.Strategy, Status: f.Status, + SpeedID: f.SpeedID, } if strings.TrimSpace(fr.Strategy) == "" { fr.Strategy = "fifo" @@ -169,3 +172,15 @@ func (r *Repository) SpeedLimitExists(id int64) (bool, error) { } return count > 0, nil } + +func (r *Repository) GetSpeedLimitSpeed(id int64) (int, error) { + if r == nil || r.db == nil { + return 0, errors.New("repository not initialized") + } + var sl model.SpeedLimit + err := r.db.Select("speed").Where("id = ?", id).First(&sl).Error + if err != nil { + return 0, err + } + return sl.Speed, nil +} diff --git a/go-backend/internal/store/repo/repository_mutations.go b/go-backend/internal/store/repo/repository_mutations.go index 6332e1e..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 } @@ -764,51 +766,66 @@ func (r *Repository) GetUsedPortsOnNodeAsMap(nodeID int64) (map[int]bool, error) return used, nil } -func (r *Repository) CreateSpeedLimit(name string, speed int, tunnelID int64, tunnelName string, now int64, status int) (int64, error) { +func (r *Repository) CreateSpeedLimit(name string, speed int, tunnelID *int64, tunnelName string, now int64, status int) (int64, error) { if r == nil || r.db == nil { return 0, errors.New("repository not initialized") } sl := model.SpeedLimit{ Name: name, Speed: speed, - TunnelID: tunnelID, - TunnelName: tunnelName, + TunnelID: sql.NullInt64{Int64: 0, Valid: false}, + TunnelName: sql.NullString{String: "", Valid: false}, CreatedTime: now, UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, Status: status, } + if tunnelID != nil { + sl.TunnelID = sql.NullInt64{Int64: *tunnelID, Valid: true} + } + if tunnelName != "" { + sl.TunnelName = sql.NullString{String: tunnelName, Valid: true} + } if err := r.db.Create(&sl).Error; err != nil { return 0, err } return sl.ID, nil } -func (r *Repository) UpdateSpeedLimit(id int64, name string, speed int, tunnelID int64, tunnelName string, status int, now int64) error { +func (r *Repository) UpdateSpeedLimit(id int64, name string, speed int, tunnelID *int64, tunnelName string, status int, now int64) error { if r == nil || r.db == nil { return errors.New("repository not initialized") } + updates := map[string]interface{}{ + "name": name, + "speed": speed, + "status": status, + "updated_time": sql.NullInt64{ + Int64: now, + Valid: true, + }, + } + if tunnelID != nil { + updates["tunnel_id"] = sql.NullInt64{Int64: *tunnelID, Valid: true} + } else { + updates["tunnel_id"] = sql.NullInt64{Int64: 0, Valid: false} + } + if tunnelName != "" { + updates["tunnel_name"] = sql.NullString{String: tunnelName, Valid: true} + } else { + updates["tunnel_name"] = sql.NullString{String: "", Valid: false} + } return r.db.Model(&model.SpeedLimit{}). Where("id = ?", id). - Updates(map[string]interface{}{ - "name": name, - "speed": speed, - "tunnel_id": tunnelID, - "tunnel_name": tunnelName, - "status": status, - "updated_time": sql.NullInt64{ - Int64: now, - Valid: true, - }, - }).Error + Updates(updates).Error } -func (r *Repository) GetSpeedLimitTunnelID(speedLimitID int64) int64 { +func (r *Repository) GetSpeedLimitTunnelID(speedLimitID int64) sql.NullInt64 { if r == nil || r.db == nil { - return 0 + return sql.NullInt64{Valid: false} } var sl model.SpeedLimit if err := r.db.Select("tunnel_id").Where("id = ?", speedLimitID).First(&sl).Error; err != nil { - return 0 + return sql.NullInt64{Valid: false} } return sl.TunnelID } @@ -1190,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") } @@ -1209,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/api/types.ts b/vite-frontend/src/api/types.ts index f1a7454..55cce65 100644 --- a/vite-frontend/src/api/types.ts +++ b/vite-frontend/src/api/types.ts @@ -51,6 +51,7 @@ export interface ForwardApiItem { outFlow?: number; userId?: number; tunnelId?: number; + speedId?: number | null; inx?: number; [key: string]: unknown; } @@ -96,10 +97,10 @@ export interface StatisticsFlowApiItem { export interface SpeedLimitApiItem { id: number; name: string; - tunnelId: number; + tunnelId?: number | null; speed: number; status: number; - tunnelName: string; + tunnelName?: string; createdTime: string; updatedTime: string; uploadSpeed?: number; @@ -284,6 +285,7 @@ export interface ForwardMutationPayload { inPort?: number | null; remoteAddr?: string; strategy?: string; + speedId?: number | null; } export interface SpeedLimitMutationPayload { 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} -