mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-10 19:36:36 +08:00
feat: add speed limit contract tests and refine limit/user UI
This commit is contained in:
+37
-11
@@ -12,6 +12,12 @@
|
|||||||
|
|
||||||
## 二、实施计划清单
|
## 二、实施计划清单
|
||||||
|
|
||||||
|
### 2.0 计划状态(审计更新:2026-02-26)
|
||||||
|
|
||||||
|
- 总体状态:**进行中(未验收通过)**
|
||||||
|
- 已完成:模型、仓储查询、限速 CRUD、控制面优先级、限速页与类型改造、编译与测试通过
|
||||||
|
- 未完成:**Forward 独立限速写入链路**(前端表单 -> API handler -> repository 落库 `forward.speed_id`)
|
||||||
|
|
||||||
### 2.1 后端模型层 (Model)
|
### 2.1 后端模型层 (Model)
|
||||||
|
|
||||||
| 序号 | 任务 | 文件 | 状态 |
|
| 序号 | 任务 | 文件 | 状态 |
|
||||||
@@ -80,6 +86,19 @@
|
|||||||
|------|------|------|
|
|------|------|------|
|
||||||
| B1 | Go 后端编译通过 | ✅ 完成 |
|
| B1 | Go 后端编译通过 | ✅ 完成 |
|
||||||
| B2 | TypeScript 类型检查通过 | ✅ 完成 |
|
| B2 | TypeScript 类型检查通过 | ✅ 完成 |
|
||||||
|
| B3 | `go test ./...` 全量通过 | ✅ 完成 |
|
||||||
|
| B4 | `go test ./tests/contract/... -run SpeedLimit` 通过 | ✅ 完成 |
|
||||||
|
|
||||||
|
### 2.8 Forward 独立限速写入链路补全(新增)
|
||||||
|
|
||||||
|
| 序号 | 任务 | 文件 | 状态 |
|
||||||
|
|------|------|------|------|
|
||||||
|
| N1 | forwardCreate 支持接收并校验可选 speedId,写入 Forward.SpeedID | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 |
|
||||||
|
| N2 | forwardUpdate 支持更新/清空 speedId,并触发服务重下发 | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 |
|
||||||
|
| N3 | CreateForwardTx 支持落库 speed_id | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 |
|
||||||
|
| N4 | UpdateForward 支持更新 speed_id | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 |
|
||||||
|
| N5 | Forward 页面新增限速选择并透传 speedId | `vite-frontend/src/pages/forward.tsx` | ✅ 完成 |
|
||||||
|
| N6 | Forward 相关契约测试补充 speedId 写入/清空断言 | `go-backend/tests/contract/forward_contract_test.go` | ✅ 完成 |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -100,23 +119,30 @@
|
|||||||
|
|
||||||
## 五、验证检查项
|
## 五、验证检查项
|
||||||
|
|
||||||
### 5.1 功能验证 (待手动测试)
|
### 5.1 功能验证(审计后)
|
||||||
|
|
||||||
- [ ] 创建不限速规则的限速 (不绑定隧道)
|
- [x] 创建不限速规则的限速 (不绑定隧道)
|
||||||
- [ ] 创建绑定隧道的限速 (兼容旧逻辑)
|
- [x] 创建绑定隧道的限速 (兼容旧逻辑)
|
||||||
- [ ] 编辑限速规则,切换隧道绑定状态
|
- [x] 编辑限速规则,切换隧道绑定状态
|
||||||
- [ ] 删除限速规则
|
- [ ] 删除限速规则
|
||||||
- [ ] 转发列表正确显示 speedId
|
- [ ] 转发列表正确显示 speedId
|
||||||
|
|
||||||
### 5.2 API 验证 (待手动测试)
|
### 5.2 API 验证(审计后)
|
||||||
|
|
||||||
- [ ] GET /api/speed-limit/list 返回可选 tunnelId
|
- [x] GET /api/speed-limit/list 返回可选 tunnelId
|
||||||
- [ ] POST /api/speed-limit/create 接受可选 tunnelId
|
- [x] POST /api/speed-limit/create 接受可选 tunnelId
|
||||||
- [ ] POST /api/speed-limit/update 接受可选 tunnelId
|
- [x] POST /api/speed-limit/update 接受可选 tunnelId
|
||||||
- [ ] GET /api/forward/list 返回 speedId
|
- [ ] GET /api/forward/list 返回 speedId
|
||||||
|
|
||||||
### 5.3 兼容性验证 (待手动测试)
|
### 5.3 兼容性验证(审计后)
|
||||||
|
|
||||||
- [ ] 现有绑定隧道的限速规则继续正常工作
|
- [x] 现有绑定隧道的限速规则继续正常工作
|
||||||
- [ ] 现有 UserTunnel 的限速继续正常工作
|
- [ ] 现有 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
|
||||||
|
|||||||
@@ -185,7 +185,9 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
|
|||||||
|
|
||||||
for _, fp := range ports {
|
for _, fp := range ports {
|
||||||
if limiterID != nil && speed != nil {
|
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)
|
node, err := h.getNodeRecord(fp.NodeID)
|
||||||
@@ -1051,12 +1053,16 @@ func (h *Handler) sendDeleteLimiterConfig(limiterID int64, tunnelID int64) error
|
|||||||
return nil
|
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
|
rate := float64(speed) / 8.0
|
||||||
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
|
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
|
||||||
payload := map[string]interface{}{
|
payload := map[string]interface{}{
|
||||||
"name": strconv.FormatInt(limiterID, 10),
|
"name": strconv.FormatInt(limiterID, 10),
|
||||||
"limits": []string{limitStr},
|
"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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1048,21 +1048,55 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
|||||||
response.WriteJSON(w, response.ErrDefault("权限ID不能为空"))
|
response.WriteJSON(w, response.ErrDefault("权限ID不能为空"))
|
||||||
return
|
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,
|
if err := h.repo.UpdateUserTunnel(id,
|
||||||
asInt64(req["flow"], 0),
|
asInt64(req["flow"], 0),
|
||||||
asInt(req["num"], 0),
|
asInt(req["num"], 0),
|
||||||
asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()),
|
asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()),
|
||||||
asInt64(req["flowResetTime"], 1),
|
asInt64(req["flowResetTime"], 1),
|
||||||
nullableInt(asAnyToInt64Ptr(req["speedId"])),
|
nullableInt(speedID),
|
||||||
asInt(req["status"], 1),
|
asInt(req["status"], 1),
|
||||||
); err != nil {
|
); err != nil {
|
||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
userID, tunnelID, utErr := h.repo.GetUserTunnelUserAndTunnel(id)
|
if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil {
|
||||||
if utErr == nil {
|
rollbackErr := h.repo.UpdateUserTunnel(
|
||||||
h.syncUserTunnelForwards(userID, tunnelID)
|
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())
|
response.WriteJSON(w, response.OKEmpty())
|
||||||
@@ -1103,6 +1137,18 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
|||||||
response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空"))
|
response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空"))
|
||||||
return
|
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)
|
port := asInt(req["inPort"], 0)
|
||||||
if port <= 0 {
|
if port <= 0 {
|
||||||
port = h.pickTunnelPort(tunnelID)
|
port = h.pickTunnelPort(tunnelID)
|
||||||
@@ -1127,7 +1173,7 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
|||||||
if userName == "" {
|
if userName == "" {
|
||||||
userName = "user"
|
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 {
|
if err != nil {
|
||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
return
|
return
|
||||||
@@ -1202,6 +1248,24 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
|||||||
if strategy == "" {
|
if strategy == "" {
|
||||||
strategy = forward.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)
|
port := asInt(req["inPort"], 0)
|
||||||
if port <= 0 {
|
if port <= 0 {
|
||||||
@@ -1225,7 +1289,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
now := time.Now().UnixMilli()
|
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()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -3006,6 +3070,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []
|
|||||||
h.repo.RollbackForwardFields(
|
h.repo.RollbackForwardFields(
|
||||||
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
|
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
|
||||||
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
|
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
|
||||||
|
oldForward.SpeedID,
|
||||||
time.Now().UnixMilli(),
|
time.Now().UnixMilli(),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -3027,6 +3092,10 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
|||||||
h.repo.GetExistingUserTunnel(userID, tunnelID)
|
h.repo.GetExistingUserTunnel(userID, tunnelID)
|
||||||
|
|
||||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||||
|
if err := h.validateSpeedLimitReference(speedID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
reqFlow := asInt64(req["flow"], -1)
|
reqFlow := asInt64(req["flow"], -1)
|
||||||
reqNum := asInt(req["num"], -1)
|
reqNum := asInt(req["num"], -1)
|
||||||
reqExpTime := asInt64(req["expTime"], -1)
|
reqExpTime := asInt64(req["expTime"], -1)
|
||||||
@@ -3067,7 +3136,24 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
|||||||
reqStatus = 1
|
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 {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -3105,25 +3191,61 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
|||||||
newSpeedID = sql.NullInt64{Valid: false}
|
newSpeedID = sql.NullInt64{Valid: false}
|
||||||
}
|
}
|
||||||
|
|
||||||
err = h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus)
|
if err := h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus); err != nil {
|
||||||
|
return err
|
||||||
if err == nil {
|
|
||||||
h.syncUserTunnelForwards(userID, tunnelID)
|
|
||||||
}
|
}
|
||||||
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)
|
forwards, err := h.listForwardsByTunnel(tunnelID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return
|
return err
|
||||||
}
|
}
|
||||||
for i := range forwards {
|
for i := range forwards {
|
||||||
f := &forwards[i]
|
f := &forwards[i]
|
||||||
if f.UserID == userID {
|
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{} {
|
func asAnySlice(v interface{}) []interface{} {
|
||||||
|
|||||||
@@ -656,7 +656,7 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
|
|||||||
return nil, errors.New("repository not initialized")
|
return nil, errors.New("repository not initialized")
|
||||||
}
|
}
|
||||||
var users []model.User
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
items := make([]map[string]interface{}, 0, len(users))
|
items := make([]map[string]interface{}, 0, len(users))
|
||||||
@@ -678,7 +678,7 @@ func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) {
|
|||||||
return nil, errors.New("repository not initialized")
|
return nil, errors.New("repository not initialized")
|
||||||
}
|
}
|
||||||
var limits []model.SpeedLimit
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
items := make([]map[string]interface{}, 0, len(limits))
|
items := make([]map[string]interface{}, 0, len(limits))
|
||||||
@@ -1319,7 +1319,6 @@ func (r *Repository) ListActiveForwardPeerShareRuntimesByNodeAndServiceName(node
|
|||||||
return items, nil
|
return items, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
func (r *Repository) ListActiveForwardPeerShareRuntimeServiceNamesByNode(nodeID int64) ([]string, error) {
|
func (r *Repository) ListActiveForwardPeerShareRuntimeServiceNamesByNode(nodeID int64) ([]string, error) {
|
||||||
if r == nil || r.db == nil {
|
if r == nil || r.db == nil {
|
||||||
return nil, errors.New("repository not initialized")
|
return nil, errors.New("repository not initialized")
|
||||||
|
|||||||
@@ -228,7 +228,6 @@ func (r *Repository) ListTunnelIDsByNamePrefix(prefix string) ([]int64, error) {
|
|||||||
return ids, nil
|
return ids, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// NextIndex returns COALESCE(MAX(inx), -1) + 1 for the given table.
|
|
||||||
func (r *Repository) NextIndex(table string) int {
|
func (r *Repository) NextIndex(table string) int {
|
||||||
if r == nil || r.db == nil {
|
if r == nil || r.db == nil {
|
||||||
return 0
|
return 0
|
||||||
@@ -251,7 +250,7 @@ func (r *Repository) NextIndex(table string) int {
|
|||||||
var row inxRow
|
var row inxRow
|
||||||
err := r.db.Model(modelRef).
|
err := r.db.Model(modelRef).
|
||||||
Select("inx").
|
Select("inx").
|
||||||
Order("inx DESC").
|
Order("inx ASC, id ASC").
|
||||||
Limit(1).
|
Limit(1).
|
||||||
Take(&row).Error
|
Take(&row).Error
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
@@ -260,10 +259,7 @@ func (r *Repository) NextIndex(table string) int {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
if row.Inx < 0 {
|
return row.Inx - 1
|
||||||
return 0
|
|
||||||
}
|
|
||||||
return row.Inx + 1
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// CreateRemoteNode inserts a new remote node.
|
// CreateRemoteNode inserts a new remote node.
|
||||||
|
|||||||
@@ -657,7 +657,7 @@ func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 {
|
|||||||
return p
|
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 {
|
if r == nil || r.db == nil {
|
||||||
return errors.New("repository not initialized")
|
return errors.New("repository not initialized")
|
||||||
}
|
}
|
||||||
@@ -668,6 +668,7 @@ func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remote
|
|||||||
"tunnel_id": tunnelID,
|
"tunnel_id": tunnelID,
|
||||||
"remote_addr": remoteAddr,
|
"remote_addr": remoteAddr,
|
||||||
"strategy": strategy,
|
"strategy": strategy,
|
||||||
|
"speed_id": nullInt64FromInterface(speedID),
|
||||||
"updated_time": now,
|
"updated_time": now,
|
||||||
}).Error
|
}).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 {
|
if r == nil || r.db == nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -738,6 +739,7 @@ func (r *Repository) RollbackForwardFields(id, userID int64, userName, name stri
|
|||||||
"remote_addr": remoteAddr,
|
"remote_addr": remoteAddr,
|
||||||
"strategy": strategy,
|
"strategy": strategy,
|
||||||
"status": status,
|
"status": status,
|
||||||
|
"speed_id": nullInt64FromInterface(speedID),
|
||||||
"updated_time": now,
|
"updated_time": now,
|
||||||
}).Error
|
}).Error
|
||||||
}
|
}
|
||||||
@@ -1205,7 +1207,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
|
|||||||
return ut.ID, true, nil
|
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 {
|
if r == nil || r.db == nil {
|
||||||
return 0, errors.New("repository not initialized")
|
return 0, errors.New("repository not initialized")
|
||||||
}
|
}
|
||||||
@@ -1224,6 +1226,7 @@ func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnel
|
|||||||
UpdatedTime: now,
|
UpdatedTime: now,
|
||||||
Status: 1,
|
Status: 1,
|
||||||
Inx: inx,
|
Inx: inx,
|
||||||
|
SpeedID: nullInt64FromInterface(speedID),
|
||||||
}
|
}
|
||||||
if err := tx.Create(&fwd).Error; err != nil {
|
if err := tx.Create(&fwd).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package contract_test
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"database/sql"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"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 {
|
func jsonNumber(v int64) string {
|
||||||
return strconv.FormatInt(v, 10)
|
return strconv.FormatInt(v, 10)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -1,3 +1,5 @@
|
|||||||
|
import type { SpeedLimitApiItem } from "@/api/types";
|
||||||
|
|
||||||
import { useState, useEffect, useMemo } from "react";
|
import { useState, useEffect, useMemo } from "react";
|
||||||
import toast from "react-hot-toast";
|
import toast from "react-hot-toast";
|
||||||
import {
|
import {
|
||||||
@@ -50,6 +52,7 @@ import { Checkbox } from "@/shadcn-bridge/heroui/checkbox";
|
|||||||
import {
|
import {
|
||||||
createForward,
|
createForward,
|
||||||
getForwardList,
|
getForwardList,
|
||||||
|
getSpeedLimitList,
|
||||||
getPeerShareList,
|
getPeerShareList,
|
||||||
getPeerRemoteUsageList,
|
getPeerRemoteUsageList,
|
||||||
updateForward,
|
updateForward,
|
||||||
@@ -105,6 +108,7 @@ interface Forward {
|
|||||||
userName?: string;
|
userName?: string;
|
||||||
userId?: number;
|
userId?: number;
|
||||||
inx?: number;
|
inx?: number;
|
||||||
|
speedId?: number | null;
|
||||||
}
|
}
|
||||||
|
|
||||||
interface Tunnel {
|
interface Tunnel {
|
||||||
@@ -123,12 +127,14 @@ interface ForwardForm {
|
|||||||
remoteAddr: string;
|
remoteAddr: string;
|
||||||
interfaceName?: string;
|
interfaceName?: string;
|
||||||
strategy: string;
|
strategy: string;
|
||||||
|
speedId: number | null;
|
||||||
}
|
}
|
||||||
|
|
||||||
export default function ForwardPage() {
|
export default function ForwardPage() {
|
||||||
const [loading, setLoading] = useState(true);
|
const [loading, setLoading] = useState(true);
|
||||||
const [forwards, setForwards] = useState<Forward[]>([]);
|
const [forwards, setForwards] = useState<Forward[]>([]);
|
||||||
const [tunnels, setTunnels] = useState<Tunnel[]>([]);
|
const [tunnels, setTunnels] = useState<Tunnel[]>([]);
|
||||||
|
const [speedLimits, setSpeedLimits] = useState<SpeedLimitApiItem[]>([]);
|
||||||
const isMobile = useMobileBreakpoint();
|
const isMobile = useMobileBreakpoint();
|
||||||
const [searchKeyword, setSearchKeyword] = useLocalStorageState(
|
const [searchKeyword, setSearchKeyword] = useLocalStorageState(
|
||||||
"forward-search-keyword",
|
"forward-search-keyword",
|
||||||
@@ -206,6 +212,7 @@ export default function ForwardPage() {
|
|||||||
remoteAddr: "",
|
remoteAddr: "",
|
||||||
interfaceName: "",
|
interfaceName: "",
|
||||||
strategy: "fifo",
|
strategy: "fifo",
|
||||||
|
speedId: null,
|
||||||
});
|
});
|
||||||
|
|
||||||
// 表单验证错误
|
// 表单验证错误
|
||||||
@@ -327,7 +334,9 @@ export default function ForwardPage() {
|
|||||||
|
|
||||||
const resolveShareIdForForward = (forward: Forward): number | null => {
|
const resolveShareIdForForward = (forward: Forward): number | null => {
|
||||||
const candidates = new Set<number>();
|
const candidates = new Set<number>();
|
||||||
const shareIdFromName = parseShareIdFromTunnelName(forward.tunnelName || "");
|
const shareIdFromName = parseShareIdFromTunnelName(
|
||||||
|
forward.tunnelName || "",
|
||||||
|
);
|
||||||
|
|
||||||
if (shareIdFromName) {
|
if (shareIdFromName) {
|
||||||
candidates.add(shareIdFromName);
|
candidates.add(shareIdFromName);
|
||||||
@@ -445,9 +454,10 @@ export default function ForwardPage() {
|
|||||||
const loadData = async (lod = true) => {
|
const loadData = async (lod = true) => {
|
||||||
setLoading(lod);
|
setLoading(lod);
|
||||||
try {
|
try {
|
||||||
const [forwardsRes, tunnelsRes] = await Promise.all([
|
const [forwardsRes, tunnelsRes, speedLimitsRes] = await Promise.all([
|
||||||
getForwardList(),
|
getForwardList(),
|
||||||
userTunnel(),
|
userTunnel(),
|
||||||
|
getSpeedLimitList(),
|
||||||
]);
|
]);
|
||||||
|
|
||||||
if (forwardsRes.code === 0) {
|
if (forwardsRes.code === 0) {
|
||||||
@@ -481,6 +491,10 @@ export default function ForwardPage() {
|
|||||||
setTunnels(tunnelsRes.data || []);
|
setTunnels(tunnelsRes.data || []);
|
||||||
} else {
|
} else {
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if (speedLimitsRes.code === 0) {
|
||||||
|
setSpeedLimits(speedLimitsRes.data || []);
|
||||||
|
}
|
||||||
} catch {
|
} catch {
|
||||||
toast.error("加载数据失败");
|
toast.error("加载数据失败");
|
||||||
} finally {
|
} finally {
|
||||||
@@ -489,6 +503,10 @@ export default function ForwardPage() {
|
|||||||
};
|
};
|
||||||
|
|
||||||
// 表单验证
|
// 表单验证
|
||||||
|
const availableSpeedLimits = useMemo(() => {
|
||||||
|
return speedLimits;
|
||||||
|
}, [speedLimits]);
|
||||||
|
|
||||||
const validateForm = (): boolean => {
|
const validateForm = (): boolean => {
|
||||||
const newErrors: { [key: string]: string } = {};
|
const newErrors: { [key: string]: string } = {};
|
||||||
|
|
||||||
@@ -555,6 +573,7 @@ export default function ForwardPage() {
|
|||||||
remoteAddr: "",
|
remoteAddr: "",
|
||||||
interfaceName: "",
|
interfaceName: "",
|
||||||
strategy: "fifo",
|
strategy: "fifo",
|
||||||
|
speedId: null,
|
||||||
});
|
});
|
||||||
setErrors({});
|
setErrors({});
|
||||||
setModalOpen(true);
|
setModalOpen(true);
|
||||||
@@ -572,6 +591,7 @@ export default function ForwardPage() {
|
|||||||
remoteAddr: forward.remoteAddr.split(",").join("\n"),
|
remoteAddr: forward.remoteAddr.split(",").join("\n"),
|
||||||
interfaceName: forward.interfaceName || "",
|
interfaceName: forward.interfaceName || "",
|
||||||
strategy: forward.strategy || "fifo",
|
strategy: forward.strategy || "fifo",
|
||||||
|
speedId: forward.speedId ?? null,
|
||||||
});
|
});
|
||||||
setErrors({});
|
setErrors({});
|
||||||
setModalOpen(true);
|
setModalOpen(true);
|
||||||
@@ -651,6 +671,7 @@ export default function ForwardPage() {
|
|||||||
inPort: form.inPort,
|
inPort: form.inPort,
|
||||||
remoteAddr: processedRemoteAddr,
|
remoteAddr: processedRemoteAddr,
|
||||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||||
|
speedId: form.speedId,
|
||||||
};
|
};
|
||||||
|
|
||||||
res = await updateForward(updateData);
|
res = await updateForward(updateData);
|
||||||
@@ -662,6 +683,7 @@ export default function ForwardPage() {
|
|||||||
inPort: form.inPort,
|
inPort: form.inPort,
|
||||||
remoteAddr: processedRemoteAddr,
|
remoteAddr: processedRemoteAddr,
|
||||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||||
|
speedId: form.speedId,
|
||||||
};
|
};
|
||||||
|
|
||||||
res = await createForward(createData);
|
res = await createForward(createData);
|
||||||
@@ -1346,7 +1368,11 @@ export default function ForwardPage() {
|
|||||||
const aInx = a.inx ?? 0;
|
const aInx = a.inx ?? 0;
|
||||||
const bInx = b.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.userName || "未知用户"}
|
||||||
</span>
|
</span>
|
||||||
</TableCell>
|
</TableCell>
|
||||||
|
<TableCell className="whitespace-nowrap font-semibold text-foreground">
|
||||||
|
{forward.name}
|
||||||
|
</TableCell>
|
||||||
<TableCell className="whitespace-nowrap">
|
<TableCell className="whitespace-nowrap">
|
||||||
<Chip
|
<Chip
|
||||||
className="border-none bg-secondary/10 px-2"
|
className="border-none bg-secondary/10 px-2"
|
||||||
@@ -1522,9 +1551,6 @@ export default function ForwardPage() {
|
|||||||
</span>
|
</span>
|
||||||
</Chip>
|
</Chip>
|
||||||
</TableCell>
|
</TableCell>
|
||||||
<TableCell className="whitespace-nowrap font-semibold text-foreground">
|
|
||||||
{forward.name}
|
|
||||||
</TableCell>
|
|
||||||
<TableCell className="max-w-[220px]">
|
<TableCell className="max-w-[220px]">
|
||||||
<button
|
<button
|
||||||
className={`w-full truncate rounded-md bg-default-100/50 px-2.5 py-1.5 text-left font-mono text-xs font-medium text-default-700 transition-all ${
|
className={`w-full truncate rounded-md bg-default-100/50 px-2.5 py-1.5 text-left font-mono text-xs font-medium text-default-700 transition-all ${
|
||||||
@@ -2181,8 +2207,8 @@ export default function ForwardPage() {
|
|||||||
)}
|
)}
|
||||||
<TableColumn className="w-10 pl-4" />
|
<TableColumn className="w-10 pl-4" />
|
||||||
<TableColumn>用户</TableColumn>
|
<TableColumn>用户</TableColumn>
|
||||||
<TableColumn>隧道</TableColumn>
|
|
||||||
<TableColumn>名称</TableColumn>
|
<TableColumn>名称</TableColumn>
|
||||||
|
<TableColumn>隧道</TableColumn>
|
||||||
<TableColumn>入口</TableColumn>
|
<TableColumn>入口</TableColumn>
|
||||||
<TableColumn>目标</TableColumn>
|
<TableColumn>目标</TableColumn>
|
||||||
<TableColumn>策略</TableColumn>
|
<TableColumn>策略</TableColumn>
|
||||||
@@ -2301,6 +2327,38 @@ export default function ForwardPage() {
|
|||||||
}
|
}
|
||||||
/>
|
/>
|
||||||
|
|
||||||
|
<Select
|
||||||
|
label="限速规则"
|
||||||
|
placeholder="不限速"
|
||||||
|
selectedKeys={
|
||||||
|
form.speedId !== null && form.speedId !== undefined
|
||||||
|
? [form.speedId.toString()]
|
||||||
|
: ["null"]
|
||||||
|
}
|
||||||
|
variant="bordered"
|
||||||
|
onSelectionChange={(keys) => {
|
||||||
|
const selectedKey = Array.from(keys)[0] as string;
|
||||||
|
|
||||||
|
setForm((prev) => ({
|
||||||
|
...prev,
|
||||||
|
speedId:
|
||||||
|
selectedKey === "null" ? null : Number(selectedKey),
|
||||||
|
}));
|
||||||
|
}}
|
||||||
|
>
|
||||||
|
<SelectItem key="null" textValue="不限速">
|
||||||
|
不限速
|
||||||
|
</SelectItem>
|
||||||
|
{availableSpeedLimits.map((speedLimit) => (
|
||||||
|
<SelectItem
|
||||||
|
key={speedLimit.id.toString()}
|
||||||
|
textValue={speedLimit.name}
|
||||||
|
>
|
||||||
|
{speedLimit.name}
|
||||||
|
</SelectItem>
|
||||||
|
))}
|
||||||
|
</Select>
|
||||||
|
|
||||||
<Select
|
<Select
|
||||||
description={
|
description={
|
||||||
isEdit
|
isEdit
|
||||||
|
|||||||
@@ -217,6 +217,8 @@ export default function LimitPage() {
|
|||||||
const createData = { ...form };
|
const createData = { ...form };
|
||||||
|
|
||||||
delete createData.id;
|
delete createData.id;
|
||||||
|
createData.tunnelId = null;
|
||||||
|
createData.tunnelName = "";
|
||||||
|
|
||||||
res = await createSpeedLimit(createData);
|
res = await createSpeedLimit(createData);
|
||||||
}
|
}
|
||||||
@@ -391,9 +393,7 @@ export default function LimitPage() {
|
|||||||
{isEdit ? "编辑限速规则" : "新增限速规则"}
|
{isEdit ? "编辑限速规则" : "新增限速规则"}
|
||||||
</h2>
|
</h2>
|
||||||
<p className="text-small text-default-500">
|
<p className="text-small text-default-500">
|
||||||
{isEdit
|
{isEdit ? "修改现有限速规则的配置信息" : "创建新的限速规则"}
|
||||||
? "修改现有限速规则的配置信息"
|
|
||||||
: "创建新的限速规则并绑定到隧道"}
|
|
||||||
</p>
|
</p>
|
||||||
</ModalHeader>
|
</ModalHeader>
|
||||||
<ModalBody>
|
<ModalBody>
|
||||||
@@ -433,42 +433,44 @@ export default function LimitPage() {
|
|||||||
}
|
}
|
||||||
/>
|
/>
|
||||||
|
|
||||||
<Select
|
{isEdit && (
|
||||||
description="绑定隧道为可选项,不绑定则创建通用限速规则"
|
<Select
|
||||||
errorMessage={errors.tunnelId}
|
description="仅编辑时可调整绑定隧道"
|
||||||
isInvalid={!!errors.tunnelId}
|
errorMessage={errors.tunnelId}
|
||||||
label="绑定隧道"
|
isInvalid={!!errors.tunnelId}
|
||||||
placeholder="可选择要绑定的隧道(可选)"
|
label="绑定隧道"
|
||||||
selectedKeys={
|
placeholder="可选择要绑定的隧道(可选)"
|
||||||
form.tunnelId ? [form.tunnelId.toString()] : []
|
selectedKeys={
|
||||||
}
|
form.tunnelId ? [form.tunnelId.toString()] : []
|
||||||
variant="bordered"
|
|
||||||
onSelectionChange={(keys) => {
|
|
||||||
const selectedKey = Array.from(keys)[0] as string;
|
|
||||||
|
|
||||||
if (selectedKey) {
|
|
||||||
const selectedTunnel = tunnels.find(
|
|
||||||
(tunnel) => tunnel.id === parseInt(selectedKey),
|
|
||||||
);
|
|
||||||
|
|
||||||
setForm((prev) => ({
|
|
||||||
...prev,
|
|
||||||
tunnelId: parseInt(selectedKey),
|
|
||||||
tunnelName: selectedTunnel?.name || "",
|
|
||||||
}));
|
|
||||||
} else {
|
|
||||||
setForm((prev) => ({
|
|
||||||
...prev,
|
|
||||||
tunnelId: null,
|
|
||||||
tunnelName: "",
|
|
||||||
}));
|
|
||||||
}
|
}
|
||||||
}}
|
variant="bordered"
|
||||||
>
|
onSelectionChange={(keys) => {
|
||||||
{tunnels.map((tunnel) => (
|
const selectedKey = Array.from(keys)[0] as string;
|
||||||
<SelectItem key={tunnel.id}>{tunnel.name}</SelectItem>
|
|
||||||
))}
|
if (selectedKey) {
|
||||||
</Select>
|
const selectedTunnel = tunnels.find(
|
||||||
|
(tunnel) => tunnel.id === parseInt(selectedKey),
|
||||||
|
);
|
||||||
|
|
||||||
|
setForm((prev) => ({
|
||||||
|
...prev,
|
||||||
|
tunnelId: parseInt(selectedKey),
|
||||||
|
tunnelName: selectedTunnel?.name || "",
|
||||||
|
}));
|
||||||
|
} else {
|
||||||
|
setForm((prev) => ({
|
||||||
|
...prev,
|
||||||
|
tunnelId: null,
|
||||||
|
tunnelName: "",
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
}}
|
||||||
|
>
|
||||||
|
{tunnels.map((tunnel) => (
|
||||||
|
<SelectItem key={tunnel.id}>{tunnel.name}</SelectItem>
|
||||||
|
))}
|
||||||
|
</Select>
|
||||||
|
)}
|
||||||
</div>
|
</div>
|
||||||
</ModalBody>
|
</ModalBody>
|
||||||
<ModalFooter>
|
<ModalFooter>
|
||||||
|
|||||||
@@ -582,12 +582,10 @@ export default function UserPage() {
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
const editAvailableSpeedLimits = speedLimits.filter(
|
const editAvailableSpeedLimits = speedLimits;
|
||||||
(speedLimit) => speedLimit.tunnelId === editTunnelForm?.tunnelId,
|
|
||||||
);
|
|
||||||
|
|
||||||
const getSpeedLimitsForTunnel = (tunnelId: number) => {
|
const getSpeedLimitsForTunnel = (_tunnelId: number) => {
|
||||||
return speedLimits.filter((sl) => sl.tunnelId === tunnelId);
|
return speedLimits;
|
||||||
};
|
};
|
||||||
|
|
||||||
const toggleTunnelSelection = (tunnelId: number) => {
|
const toggleTunnelSelection = (tunnelId: number) => {
|
||||||
|
|||||||
Reference in New Issue
Block a user