mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-30 16:26: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)
|
||||
|
||||
| 序号 | 任务 | 文件 | 状态 |
|
||||
@@ -80,6 +86,19 @@
|
||||
|------|------|------|
|
||||
| B1 | Go 后端编译通过 | ✅ 完成 |
|
||||
| B2 | TypeScript 类型检查通过 | ✅ 完成 |
|
||||
| B3 | `go test ./...` 全量通过 | ✅ 完成 |
|
||||
| B4 | `go test ./tests/contract/... -run SpeedLimit` 通过 | ✅ 完成 |
|
||||
|
||||
### 2.8 Forward 独立限速写入链路补全(新增)
|
||||
|
||||
| 序号 | 任务 | 文件 | 状态 |
|
||||
|------|------|------|------|
|
||||
| N1 | forwardCreate 支持接收并校验可选 speedId,写入 Forward.SpeedID | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 |
|
||||
| N2 | forwardUpdate 支持更新/清空 speedId,并触发服务重下发 | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 |
|
||||
| N3 | CreateForwardTx 支持落库 speed_id | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 |
|
||||
| N4 | UpdateForward 支持更新 speed_id | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 |
|
||||
| N5 | Forward 页面新增限速选择并透传 speedId | `vite-frontend/src/pages/forward.tsx` | ✅ 完成 |
|
||||
| N6 | Forward 相关契约测试补充 speedId 写入/清空断言 | `go-backend/tests/contract/forward_contract_test.go` | ✅ 完成 |
|
||||
|
||||
---
|
||||
|
||||
@@ -100,23 +119,30 @@
|
||||
|
||||
## 五、验证检查项
|
||||
|
||||
### 5.1 功能验证 (待手动测试)
|
||||
### 5.1 功能验证(审计后)
|
||||
|
||||
- [ ] 创建不限速规则的限速 (不绑定隧道)
|
||||
- [ ] 创建绑定隧道的限速 (兼容旧逻辑)
|
||||
- [ ] 编辑限速规则,切换隧道绑定状态
|
||||
- [x] 创建不限速规则的限速 (不绑定隧道)
|
||||
- [x] 创建绑定隧道的限速 (兼容旧逻辑)
|
||||
- [x] 编辑限速规则,切换隧道绑定状态
|
||||
- [ ] 删除限速规则
|
||||
- [ ] 转发列表正确显示 speedId
|
||||
|
||||
### 5.2 API 验证 (待手动测试)
|
||||
### 5.2 API 验证(审计后)
|
||||
|
||||
- [ ] GET /api/speed-limit/list 返回可选 tunnelId
|
||||
- [ ] POST /api/speed-limit/create 接受可选 tunnelId
|
||||
- [ ] POST /api/speed-limit/update 接受可选 tunnelId
|
||||
- [x] GET /api/speed-limit/list 返回可选 tunnelId
|
||||
- [x] POST /api/speed-limit/create 接受可选 tunnelId
|
||||
- [x] POST /api/speed-limit/update 接受可选 tunnelId
|
||||
- [ ] GET /api/forward/list 返回 speedId
|
||||
|
||||
### 5.3 兼容性验证 (待手动测试)
|
||||
### 5.3 兼容性验证(审计后)
|
||||
|
||||
- [ ] 现有绑定隧道的限速规则继续正常工作
|
||||
- [x] 现有绑定隧道的限速规则继续正常工作
|
||||
- [ ] 现有 UserTunnel 的限速继续正常工作
|
||||
- [ ] 备份/恢复功能正常
|
||||
- [ ] 备份/恢复功能正常
|
||||
|
||||
### 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 {
|
||||
if limiterID != nil && speed != nil {
|
||||
h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed)
|
||||
if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
node, err := h.getNodeRecord(fp.NodeID)
|
||||
@@ -1051,12 +1053,16 @@ func (h *Handler) sendDeleteLimiterConfig(limiterID int64, tunnelID int64) error
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) {
|
||||
func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) error {
|
||||
rate := float64(speed) / 8.0
|
||||
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
|
||||
payload := map[string]interface{}{
|
||||
"name": strconv.FormatInt(limiterID, 10),
|
||||
"limits": []string{limitStr},
|
||||
}
|
||||
_, _ = h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false)
|
||||
if _, err := h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false); err != nil {
|
||||
return fmt.Errorf("限速规则下发失败: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1048,21 +1048,55 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("权限ID不能为空"))
|
||||
return
|
||||
}
|
||||
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if err := h.validateSpeedLimitReference(speedID); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
userID, tunnelID, utErr := h.repo.GetUserTunnelUserAndTunnel(id)
|
||||
if utErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, utErr.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
_, oldFlow, oldNum, oldExpTime, oldFlowReset, oldSpeedID, oldStatus, oldErr :=
|
||||
h.repo.GetExistingUserTunnel(userID, tunnelID)
|
||||
if oldErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, oldErr.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateUserTunnel(id,
|
||||
asInt64(req["flow"], 0),
|
||||
asInt(req["num"], 0),
|
||||
asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()),
|
||||
asInt64(req["flowResetTime"], 1),
|
||||
nullableInt(asAnyToInt64Ptr(req["speedId"])),
|
||||
nullableInt(speedID),
|
||||
asInt(req["status"], 1),
|
||||
); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
userID, tunnelID, utErr := h.repo.GetUserTunnelUserAndTunnel(id)
|
||||
if utErr == nil {
|
||||
h.syncUserTunnelForwards(userID, tunnelID)
|
||||
if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil {
|
||||
rollbackErr := h.repo.UpdateUserTunnel(
|
||||
id,
|
||||
oldFlow,
|
||||
int(oldNum),
|
||||
oldExpTime,
|
||||
oldFlowReset,
|
||||
oldSpeedID,
|
||||
oldStatus,
|
||||
)
|
||||
if rollbackErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("下发失败且回滚失败: %v; 回滚错误: %v", syncErr, rollbackErr)))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("下发失败,已回滚: %v", syncErr)))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
@@ -1103,6 +1137,18 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空"))
|
||||
return
|
||||
}
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if speedID != nil {
|
||||
exists, speedErr := h.repo.SpeedLimitExists(*speedID)
|
||||
if speedErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, speedErr.Error()))
|
||||
return
|
||||
}
|
||||
if !exists {
|
||||
response.WriteJSON(w, response.ErrDefault("限速规则不存在"))
|
||||
return
|
||||
}
|
||||
}
|
||||
port := asInt(req["inPort"], 0)
|
||||
if port <= 0 {
|
||||
port = h.pickTunnelPort(tunnelID)
|
||||
@@ -1127,7 +1173,7 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
if userName == "" {
|
||||
userName = "user"
|
||||
}
|
||||
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port)
|
||||
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, nullableInt(speedID))
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1202,6 +1248,24 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
if strategy == "" {
|
||||
strategy = forward.Strategy
|
||||
}
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if speedID != nil {
|
||||
exists, speedErr := h.repo.SpeedLimitExists(*speedID)
|
||||
if speedErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, speedErr.Error()))
|
||||
return
|
||||
}
|
||||
if !exists {
|
||||
response.WriteJSON(w, response.ErrDefault("限速规则不存在"))
|
||||
return
|
||||
}
|
||||
}
|
||||
newSpeedID := forward.SpeedID
|
||||
if speedID != nil {
|
||||
newSpeedID = sql.NullInt64{Int64: *speedID, Valid: true}
|
||||
} else if _, ok := req["speedId"]; ok {
|
||||
newSpeedID = sql.NullInt64{Valid: false}
|
||||
}
|
||||
|
||||
port := asInt(req["inPort"], 0)
|
||||
if port <= 0 {
|
||||
@@ -1225,7 +1289,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now); err != nil {
|
||||
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -3006,6 +3070,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []
|
||||
h.repo.RollbackForwardFields(
|
||||
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
|
||||
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
|
||||
oldForward.SpeedID,
|
||||
time.Now().UnixMilli(),
|
||||
)
|
||||
|
||||
@@ -3027,6 +3092,10 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
h.repo.GetExistingUserTunnel(userID, tunnelID)
|
||||
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if err := h.validateSpeedLimitReference(speedID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
reqFlow := asInt64(req["flow"], -1)
|
||||
reqNum := asInt(req["num"], -1)
|
||||
reqExpTime := asInt64(req["expTime"], -1)
|
||||
@@ -3067,7 +3136,24 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
reqStatus = 1
|
||||
}
|
||||
|
||||
return h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus)
|
||||
if err := h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil {
|
||||
insertedID, _, _, _, _, _, _, lookupErr := h.repo.GetExistingUserTunnel(userID, tunnelID)
|
||||
if lookupErr != nil {
|
||||
return fmt.Errorf("下发失败且回滚失败: %v; 回滚查询错误: %w", syncErr, lookupErr)
|
||||
}
|
||||
|
||||
if rollbackErr := h.repo.DeleteUserTunnel(insertedID); rollbackErr != nil {
|
||||
return fmt.Errorf("下发失败且回滚失败: %v; 回滚删除错误: %w", syncErr, rollbackErr)
|
||||
}
|
||||
|
||||
return fmt.Errorf("下发失败,已回滚: %w", syncErr)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -3105,25 +3191,61 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
newSpeedID = sql.NullInt64{Valid: false}
|
||||
}
|
||||
|
||||
err = h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus)
|
||||
|
||||
if err == nil {
|
||||
h.syncUserTunnelForwards(userID, tunnelID)
|
||||
if err := h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus); err != nil {
|
||||
return err
|
||||
}
|
||||
return err
|
||||
|
||||
if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil {
|
||||
rollbackErr := h.repo.UpdateUserTunnelFields(
|
||||
existingID,
|
||||
currentSpeedID,
|
||||
currentFlow,
|
||||
int(currentNum),
|
||||
currentExpTime,
|
||||
currentFlowReset,
|
||||
currentStatus,
|
||||
)
|
||||
if rollbackErr != nil {
|
||||
return fmt.Errorf("下发失败且回滚失败: %v; 回滚错误: %w", syncErr, rollbackErr)
|
||||
}
|
||||
|
||||
return fmt.Errorf("下发失败,已回滚: %w", syncErr)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) {
|
||||
func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) error {
|
||||
forwards, err := h.listForwardsByTunnel(tunnelID)
|
||||
if err != nil {
|
||||
return
|
||||
return err
|
||||
}
|
||||
for i := range forwards {
|
||||
f := &forwards[i]
|
||||
if f.UserID == userID {
|
||||
_ = h.syncForwardServices(f, "UpdateService", true)
|
||||
if err := h.syncForwardServices(f, "UpdateService", true); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) validateSpeedLimitReference(speedID *int64) error {
|
||||
if speedID == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
exists, err := h.repo.SpeedLimitExists(*speedID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !exists {
|
||||
return errors.New("限速规则不存在")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func asAnySlice(v interface{}) []interface{} {
|
||||
|
||||
@@ -656,7 +656,7 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var users []model.User
|
||||
if err := r.db.Where("role_id != ?", 0).Order("id ASC").Find(&users).Error; err != nil {
|
||||
if err := r.db.Where("role_id != ?", 0).Order("id DESC").Find(&users).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items := make([]map[string]interface{}, 0, len(users))
|
||||
@@ -678,7 +678,7 @@ func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var limits []model.SpeedLimit
|
||||
if err := r.db.Order("id ASC").Find(&limits).Error; err != nil {
|
||||
if err := r.db.Order("id DESC").Find(&limits).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items := make([]map[string]interface{}, 0, len(limits))
|
||||
@@ -1319,7 +1319,6 @@ func (r *Repository) ListActiveForwardPeerShareRuntimesByNodeAndServiceName(node
|
||||
return items, nil
|
||||
}
|
||||
|
||||
|
||||
func (r *Repository) ListActiveForwardPeerShareRuntimeServiceNamesByNode(nodeID int64) ([]string, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
|
||||
@@ -228,7 +228,6 @@ func (r *Repository) ListTunnelIDsByNamePrefix(prefix string) ([]int64, error) {
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
// NextIndex returns COALESCE(MAX(inx), -1) + 1 for the given table.
|
||||
func (r *Repository) NextIndex(table string) int {
|
||||
if r == nil || r.db == nil {
|
||||
return 0
|
||||
@@ -251,7 +250,7 @@ func (r *Repository) NextIndex(table string) int {
|
||||
var row inxRow
|
||||
err := r.db.Model(modelRef).
|
||||
Select("inx").
|
||||
Order("inx DESC").
|
||||
Order("inx ASC, id ASC").
|
||||
Limit(1).
|
||||
Take(&row).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
@@ -260,10 +259,7 @@ func (r *Repository) NextIndex(table string) int {
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
if row.Inx < 0 {
|
||||
return 0
|
||||
}
|
||||
return row.Inx + 1
|
||||
return row.Inx - 1
|
||||
}
|
||||
|
||||
// CreateRemoteNode inserts a new remote node.
|
||||
|
||||
@@ -657,7 +657,7 @@ func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 {
|
||||
return p
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64) error {
|
||||
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -668,6 +668,7 @@ func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remote
|
||||
"tunnel_id": tunnelID,
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
@@ -724,7 +725,7 @@ func (r *Repository) ReplaceForwardPorts(forwardID int64, entries []struct {
|
||||
})
|
||||
}
|
||||
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, now int64) {
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, now int64) {
|
||||
if r == nil || r.db == nil {
|
||||
return
|
||||
}
|
||||
@@ -738,6 +739,7 @@ func (r *Repository) RollbackForwardFields(id, userID int64, userName, name stri
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"status": status,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
@@ -1205,7 +1207,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
|
||||
return ut.ID, true, nil
|
||||
}
|
||||
|
||||
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int) (int64, error) {
|
||||
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, speedID interface{}) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
@@ -1224,6 +1226,7 @@ func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnel
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
Inx: inx,
|
||||
SpeedID: nullInt64FromInterface(speedID),
|
||||
}
|
||||
if err := tx.Create(&fwd).Error; err != nil {
|
||||
return err
|
||||
|
||||
@@ -2,6 +2,7 @@ package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -471,6 +472,144 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardSpeedIDWriteAndClearContracts(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'speed_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-speed-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "forward-speed-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-speed-node", "forward-speed-secret", "10.30.0.1", "10.30.0.1", "", "31000-31010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, repo, "forward-speed-node")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 31001, 'round', 1, 'tls')
|
||||
`, tunnelID, nodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, NULL, NULL, ?, NULL, ?)
|
||||
`, "forward-speed-limit-a", 2048, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit a: %v", err)
|
||||
}
|
||||
speedIDA := mustLastInsertID(t, repo, "forward-speed-limit-a")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, NULL, NULL, ?, NULL, ?)
|
||||
`, "forward-speed-limit-b", 4096, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit b: %v", err)
|
||||
}
|
||||
speedIDB := mustLastInsertID(t, repo, "forward-speed-limit-b")
|
||||
|
||||
server := httptest.NewServer(router)
|
||||
defer server.Close()
|
||||
stopNode := startMockNodeSession(t, server.URL, "forward-speed-secret")
|
||||
defer stopNode()
|
||||
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "forward-speed-target",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.1.1.1:443",
|
||||
"strategy": "fifo",
|
||||
"speedId": speedIDA,
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
assertCode(t, createRes, 0)
|
||||
|
||||
forwardID := mustLastInsertID(t, repo, "forward-speed-target")
|
||||
storedSpeed := repo.DB().Raw(`SELECT speed_id FROM forward WHERE id = ?`, forwardID).Row()
|
||||
var createdSpeed sql.NullInt64
|
||||
if err := storedSpeed.Scan(&createdSpeed); err != nil {
|
||||
t.Fatalf("query created forward speed_id: %v", err)
|
||||
}
|
||||
if !createdSpeed.Valid || createdSpeed.Int64 != speedIDA {
|
||||
t.Fatalf("expected created speed_id=%d, got valid=%v value=%d", speedIDA, createdSpeed.Valid, createdSpeed.Int64)
|
||||
}
|
||||
|
||||
updateToBPayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"speedId": speedIDB,
|
||||
}
|
||||
updateToBBody, err := json.Marshal(updateToBPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update-to-b payload: %v", err)
|
||||
}
|
||||
updateToBReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateToBBody))
|
||||
updateToBReq.Header.Set("Authorization", adminToken)
|
||||
updateToBReq.Header.Set("Content-Type", "application/json")
|
||||
updateToBRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(updateToBRes, updateToBReq)
|
||||
assertCode(t, updateToBRes, 0)
|
||||
|
||||
storedSpeed = repo.DB().Raw(`SELECT speed_id FROM forward WHERE id = ?`, forwardID).Row()
|
||||
var updatedSpeed sql.NullInt64
|
||||
if err := storedSpeed.Scan(&updatedSpeed); err != nil {
|
||||
t.Fatalf("query updated forward speed_id: %v", err)
|
||||
}
|
||||
if !updatedSpeed.Valid || updatedSpeed.Int64 != speedIDB {
|
||||
t.Fatalf("expected updated speed_id=%d, got valid=%v value=%d", speedIDB, updatedSpeed.Valid, updatedSpeed.Int64)
|
||||
}
|
||||
|
||||
clearPayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"speedId": nil,
|
||||
}
|
||||
clearBody, err := json.Marshal(clearPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal clear payload: %v", err)
|
||||
}
|
||||
clearReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(clearBody))
|
||||
clearReq.Header.Set("Authorization", adminToken)
|
||||
clearReq.Header.Set("Content-Type", "application/json")
|
||||
clearRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(clearRes, clearReq)
|
||||
assertCode(t, clearRes, 0)
|
||||
|
||||
storedSpeed = repo.DB().Raw(`SELECT speed_id FROM forward WHERE id = ?`, forwardID).Row()
|
||||
var clearedSpeed sql.NullInt64
|
||||
if err := storedSpeed.Scan(&clearedSpeed); err != nil {
|
||||
t.Fatalf("query cleared forward speed_id: %v", err)
|
||||
}
|
||||
if clearedSpeed.Valid {
|
||||
t.Fatalf("expected cleared speed_id to be NULL, got %d", clearedSpeed.Int64)
|
||||
}
|
||||
}
|
||||
|
||||
func jsonNumber(v int64) string {
|
||||
return strconv.FormatInt(v, 10)
|
||||
}
|
||||
|
||||
@@ -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 toast from "react-hot-toast";
|
||||
import {
|
||||
@@ -50,6 +52,7 @@ import { Checkbox } from "@/shadcn-bridge/heroui/checkbox";
|
||||
import {
|
||||
createForward,
|
||||
getForwardList,
|
||||
getSpeedLimitList,
|
||||
getPeerShareList,
|
||||
getPeerRemoteUsageList,
|
||||
updateForward,
|
||||
@@ -105,6 +108,7 @@ interface Forward {
|
||||
userName?: string;
|
||||
userId?: number;
|
||||
inx?: number;
|
||||
speedId?: number | null;
|
||||
}
|
||||
|
||||
interface Tunnel {
|
||||
@@ -123,12 +127,14 @@ interface ForwardForm {
|
||||
remoteAddr: string;
|
||||
interfaceName?: string;
|
||||
strategy: string;
|
||||
speedId: number | null;
|
||||
}
|
||||
|
||||
export default function ForwardPage() {
|
||||
const [loading, setLoading] = useState(true);
|
||||
const [forwards, setForwards] = useState<Forward[]>([]);
|
||||
const [tunnels, setTunnels] = useState<Tunnel[]>([]);
|
||||
const [speedLimits, setSpeedLimits] = useState<SpeedLimitApiItem[]>([]);
|
||||
const isMobile = useMobileBreakpoint();
|
||||
const [searchKeyword, setSearchKeyword] = useLocalStorageState(
|
||||
"forward-search-keyword",
|
||||
@@ -206,6 +212,7 @@ export default function ForwardPage() {
|
||||
remoteAddr: "",
|
||||
interfaceName: "",
|
||||
strategy: "fifo",
|
||||
speedId: null,
|
||||
});
|
||||
|
||||
// 表单验证错误
|
||||
@@ -327,7 +334,9 @@ export default function ForwardPage() {
|
||||
|
||||
const resolveShareIdForForward = (forward: Forward): number | null => {
|
||||
const candidates = new Set<number>();
|
||||
const shareIdFromName = parseShareIdFromTunnelName(forward.tunnelName || "");
|
||||
const shareIdFromName = parseShareIdFromTunnelName(
|
||||
forward.tunnelName || "",
|
||||
);
|
||||
|
||||
if (shareIdFromName) {
|
||||
candidates.add(shareIdFromName);
|
||||
@@ -445,9 +454,10 @@ export default function ForwardPage() {
|
||||
const loadData = async (lod = true) => {
|
||||
setLoading(lod);
|
||||
try {
|
||||
const [forwardsRes, tunnelsRes] = await Promise.all([
|
||||
const [forwardsRes, tunnelsRes, speedLimitsRes] = await Promise.all([
|
||||
getForwardList(),
|
||||
userTunnel(),
|
||||
getSpeedLimitList(),
|
||||
]);
|
||||
|
||||
if (forwardsRes.code === 0) {
|
||||
@@ -481,6 +491,10 @@ export default function ForwardPage() {
|
||||
setTunnels(tunnelsRes.data || []);
|
||||
} else {
|
||||
}
|
||||
|
||||
if (speedLimitsRes.code === 0) {
|
||||
setSpeedLimits(speedLimitsRes.data || []);
|
||||
}
|
||||
} catch {
|
||||
toast.error("加载数据失败");
|
||||
} finally {
|
||||
@@ -489,6 +503,10 @@ export default function ForwardPage() {
|
||||
};
|
||||
|
||||
// 表单验证
|
||||
const availableSpeedLimits = useMemo(() => {
|
||||
return speedLimits;
|
||||
}, [speedLimits]);
|
||||
|
||||
const validateForm = (): boolean => {
|
||||
const newErrors: { [key: string]: string } = {};
|
||||
|
||||
@@ -555,6 +573,7 @@ export default function ForwardPage() {
|
||||
remoteAddr: "",
|
||||
interfaceName: "",
|
||||
strategy: "fifo",
|
||||
speedId: null,
|
||||
});
|
||||
setErrors({});
|
||||
setModalOpen(true);
|
||||
@@ -572,6 +591,7 @@ export default function ForwardPage() {
|
||||
remoteAddr: forward.remoteAddr.split(",").join("\n"),
|
||||
interfaceName: forward.interfaceName || "",
|
||||
strategy: forward.strategy || "fifo",
|
||||
speedId: forward.speedId ?? null,
|
||||
});
|
||||
setErrors({});
|
||||
setModalOpen(true);
|
||||
@@ -651,6 +671,7 @@ export default function ForwardPage() {
|
||||
inPort: form.inPort,
|
||||
remoteAddr: processedRemoteAddr,
|
||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||
speedId: form.speedId,
|
||||
};
|
||||
|
||||
res = await updateForward(updateData);
|
||||
@@ -662,6 +683,7 @@ export default function ForwardPage() {
|
||||
inPort: form.inPort,
|
||||
remoteAddr: processedRemoteAddr,
|
||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||
speedId: form.speedId,
|
||||
};
|
||||
|
||||
res = await createForward(createData);
|
||||
@@ -1346,7 +1368,11 @@ export default function ForwardPage() {
|
||||
const aInx = a.inx ?? 0;
|
||||
const bInx = b.inx ?? 0;
|
||||
|
||||
return aInx - bInx;
|
||||
if (aInx !== bInx) {
|
||||
return aInx - bInx;
|
||||
}
|
||||
|
||||
return (a.id ?? 0) - (b.id ?? 0);
|
||||
});
|
||||
|
||||
// 如果数据库中没有排序信息,则使用本地存储的顺序
|
||||
@@ -1511,6 +1537,9 @@ export default function ForwardPage() {
|
||||
{forward.userName || "未知用户"}
|
||||
</span>
|
||||
</TableCell>
|
||||
<TableCell className="whitespace-nowrap font-semibold text-foreground">
|
||||
{forward.name}
|
||||
</TableCell>
|
||||
<TableCell className="whitespace-nowrap">
|
||||
<Chip
|
||||
className="border-none bg-secondary/10 px-2"
|
||||
@@ -1522,9 +1551,6 @@ export default function ForwardPage() {
|
||||
</span>
|
||||
</Chip>
|
||||
</TableCell>
|
||||
<TableCell className="whitespace-nowrap font-semibold text-foreground">
|
||||
{forward.name}
|
||||
</TableCell>
|
||||
<TableCell className="max-w-[220px]">
|
||||
<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 ${
|
||||
@@ -2181,8 +2207,8 @@ export default function ForwardPage() {
|
||||
)}
|
||||
<TableColumn className="w-10 pl-4" />
|
||||
<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
|
||||
description={
|
||||
isEdit
|
||||
|
||||
@@ -217,6 +217,8 @@ export default function LimitPage() {
|
||||
const createData = { ...form };
|
||||
|
||||
delete createData.id;
|
||||
createData.tunnelId = null;
|
||||
createData.tunnelName = "";
|
||||
|
||||
res = await createSpeedLimit(createData);
|
||||
}
|
||||
@@ -391,9 +393,7 @@ export default function LimitPage() {
|
||||
{isEdit ? "编辑限速规则" : "新增限速规则"}
|
||||
</h2>
|
||||
<p className="text-small text-default-500">
|
||||
{isEdit
|
||||
? "修改现有限速规则的配置信息"
|
||||
: "创建新的限速规则并绑定到隧道"}
|
||||
{isEdit ? "修改现有限速规则的配置信息" : "创建新的限速规则"}
|
||||
</p>
|
||||
</ModalHeader>
|
||||
<ModalBody>
|
||||
@@ -433,42 +433,44 @@ export default function LimitPage() {
|
||||
}
|
||||
/>
|
||||
|
||||
<Select
|
||||
description="绑定隧道为可选项,不绑定则创建通用限速规则"
|
||||
errorMessage={errors.tunnelId}
|
||||
isInvalid={!!errors.tunnelId}
|
||||
label="绑定隧道"
|
||||
placeholder="可选择要绑定的隧道(可选)"
|
||||
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: "",
|
||||
}));
|
||||
{isEdit && (
|
||||
<Select
|
||||
description="仅编辑时可调整绑定隧道"
|
||||
errorMessage={errors.tunnelId}
|
||||
isInvalid={!!errors.tunnelId}
|
||||
label="绑定隧道"
|
||||
placeholder="可选择要绑定的隧道(可选)"
|
||||
selectedKeys={
|
||||
form.tunnelId ? [form.tunnelId.toString()] : []
|
||||
}
|
||||
}}
|
||||
>
|
||||
{tunnels.map((tunnel) => (
|
||||
<SelectItem key={tunnel.id}>{tunnel.name}</SelectItem>
|
||||
))}
|
||||
</Select>
|
||||
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: "",
|
||||
}));
|
||||
}
|
||||
}}
|
||||
>
|
||||
{tunnels.map((tunnel) => (
|
||||
<SelectItem key={tunnel.id}>{tunnel.name}</SelectItem>
|
||||
))}
|
||||
</Select>
|
||||
)}
|
||||
</div>
|
||||
</ModalBody>
|
||||
<ModalFooter>
|
||||
|
||||
@@ -582,12 +582,10 @@ export default function UserPage() {
|
||||
}
|
||||
};
|
||||
|
||||
const editAvailableSpeedLimits = speedLimits.filter(
|
||||
(speedLimit) => speedLimit.tunnelId === editTunnelForm?.tunnelId,
|
||||
);
|
||||
const editAvailableSpeedLimits = speedLimits;
|
||||
|
||||
const getSpeedLimitsForTunnel = (tunnelId: number) => {
|
||||
return speedLimits.filter((sl) => sl.tunnelId === tunnelId);
|
||||
const getSpeedLimitsForTunnel = (_tunnelId: number) => {
|
||||
return speedLimits;
|
||||
};
|
||||
|
||||
const toggleTunnelSelection = (tunnelId: number) => {
|
||||
|
||||
Reference in New Issue
Block a user