diff --git a/IMPLEMENTATION_PLAN.md b/IMPLEMENTATION_PLAN.md new file mode 100644 index 0000000..28b5ab2 --- /dev/null +++ b/IMPLEMENTATION_PLAN.md @@ -0,0 +1,122 @@ +# 限速功能重构实施计划 + +## 一、需求概述 + +**原始需求**: 限速功能当前绑定到具体隧道,需要改为不绑定隧道,创建限速后可以自由在隧道上限速,也可以在转发上限速。 + +**核心变更**: +1. 限速规则(SpeedLimit)与隧道的绑定关系改为可选 +2. 转发(Forward)支持独立的限速规则 + +--- + +## 二、实施计划清单 + +### 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 类型检查通过 | ✅ 完成 | + +--- + +## 三、优先级说明 + +限速规则应用优先级: +1. **Forward.SpeedID** - 转发级别的限速 (最高优先) +2. **UserTunnel.SpeedID** - 用户隧道权限级别的限速 (回退) + +--- + +## 四、数据库兼容性 + +- SpeedLimit 表: `tunnel_id` 和 `tunnel_name` 字段改为可空 (GORM AutoMigrate 自动处理) +- Forward 表: 新增 `speed_id` 可空字段 (GORM AutoMigrate 自动处理) + +--- + +## 五、验证检查项 + +### 5.1 功能验证 (待手动测试) + +- [ ] 创建不限速规则的限速 (不绑定隧道) +- [ ] 创建绑定隧道的限速 (兼容旧逻辑) +- [ ] 编辑限速规则,切换隧道绑定状态 +- [ ] 删除限速规则 +- [ ] 转发列表正确显示 speedId + +### 5.2 API 验证 (待手动测试) + +- [ ] GET /api/speed-limit/list 返回可选 tunnelId +- [ ] POST /api/speed-limit/create 接受可选 tunnelId +- [ ] POST /api/speed-limit/update 接受可选 tunnelId +- [ ] GET /api/forward/list 返回 speedId + +### 5.3 兼容性验证 (待手动测试) + +- [ ] 现有绑定隧道的限速规则继续正常工作 +- [ ] 现有 UserTunnel 的限速继续正常工作 +- [ ] 备份/恢复功能正常 \ No newline at end of file diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index e34aff3..8244958 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 diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index c0132d5..4d9ef31 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -1586,29 +1586,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 +1626,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 +1669,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()) } 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..e4893f9 100644 --- a/go-backend/internal/store/repo/repository.go +++ b/go-backend/internal/store/repo/repository.go @@ -683,12 +683,18 @@ func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) { } 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 } @@ -1813,9 +1824,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 +2203,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_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..ca70a08 100644 --- a/go-backend/internal/store/repo/repository_mutations.go +++ b/go-backend/internal/store/repo/repository_mutations.go @@ -764,51 +764,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 } 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/limit.tsx b/vite-frontend/src/pages/limit.tsx index 8726040..ce595da 100644 --- a/vite-frontend/src/pages/limit.tsx +++ b/vite-frontend/src/pages/limit.tsx @@ -34,8 +34,8 @@ interface SpeedLimitRule { name: string; speed: number; status: number; - tunnelId: number; - tunnelName: string; + tunnelId?: number | null; + tunnelName?: string; createdTime: string; updatedTime: string; } @@ -139,9 +139,7 @@ export default function LimitPage() { newErrors.speed = "请输入有效的速度限制(≥1 Mbps)"; } - if (!form.tunnelId) { - newErrors.tunnelId = "请选择要绑定的隧道"; - } + // tunnelId is optional - speed limits can be created without binding to a tunnel setErrors(newErrors); @@ -169,8 +167,8 @@ export default function LimitPage() { id: rule.id, name: rule.name, speed: rule.speed, - tunnelId: rule.tunnelId, - tunnelName: rule.tunnelName, + tunnelId: rule.tunnelId ?? null, + tunnelName: rule.tunnelName ?? "", status: rule.status, }); setErrors({}); @@ -436,12 +434,11 @@ export default function LimitPage() { />