From 6e8406f43911cc5fc99a03e247d68e14ca20ab82 Mon Sep 17 00:00:00 2001 From: sagitchu Date: Fri, 27 Feb 2026 19:30:02 +0800 Subject: [PATCH] feat: remove speed limit tunnel binding and add migration cleanup - Remove tunnel binding UI from speed limit page (no more Select component) - Remove /api/v1/speed-limit/tunnels route alias - Simplify CreateSpeedLimit/UpdateSpeedLimit to not accept tunnel parameters - Add schema migration v4 to clear historical tunnel_id/tunnel_name bindings - Update contract tests to verify tunnel binding is ignored - Add limiter sync failure tests for forward-level rate limiting --- .../internal/http/handler/control_plane.go | 80 ++-- .../internal/http/handler/federation.go | 2 +- go-backend/internal/http/handler/handler.go | 1 - go-backend/internal/http/handler/mutations.go | 40 +- go-backend/internal/store/repo/repository.go | 49 +- .../store/repo/repository_migrate_test.go | 75 ++++ .../store/repo/repository_mutations.go | 53 +-- .../federation_dual_panel_contract_test.go | 15 + .../limiter_sync_failure_contract_test.go | 162 +++++++ .../tests/contract/migration_contract_test.go | 39 +- .../contract/speed_limit_contract_test.go | 421 ++++++------------ vite-frontend/src/api/types.ts | 4 - vite-frontend/src/pages/limit.tsx | 108 +---- vite-frontend/src/types/index.ts | 1 - 14 files changed, 491 insertions(+), 559 deletions(-) diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index a91bdc9..a366b06 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -1028,52 +1028,46 @@ func asBool(v interface{}, def bool) bool { } } -func (h *Handler) sendLimiterConfig(limiterID int64, speedMbps int, tunnelID int64) error { - rate := float64(speedMbps) / 8.0 - limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate) - - payload := map[string]interface{}{ - "name": strconv.FormatInt(limiterID, 10), - "limits": []string{limitStr}, - } - - nodes, err := h.tunnelEntryNodeIDs(tunnelID) - if err != nil { - return err - } - - for _, nodeID := range nodes { - _, _ = h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false) - } - return nil -} - -func (h *Handler) sendDeleteLimiterConfig(limiterID int64, tunnelID int64) error { - payload := map[string]interface{}{ - "limiter": strconv.FormatInt(limiterID, 10), - } - - nodes, err := h.tunnelEntryNodeIDs(tunnelID) - if err != nil { - return err - } - - for _, nodeID := range nodes { - _, _ = h.sendNodeCommand(nodeID, "DeleteLimiters", payload, false, true) - } - return nil -} - 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}, - } - if _, err := h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false); err != nil { + if err := h.upsertLimiterOnNode(nodeID, limiterID, speed); err != nil { return fmt.Errorf("限速规则下发失败: %w", err) } return nil } + +func buildLimiterAddPayload(limiterID int64, speed int) (string, map[string]interface{}) { + rate := float64(speed) / 8.0 + limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate) + name := strconv.FormatInt(limiterID, 10) + + return name, map[string]interface{}{ + "name": name, + "limits": []string{limitStr}, + } +} + +func buildLimiterUpdatePayload(name string, data map[string]interface{}) map[string]interface{} { + return map[string]interface{}{ + "limiter": name, + "data": data, + } +} + +func (h *Handler) upsertLimiterOnNode(nodeID int64, limiterID int64, speed int) error { + name, addPayload := buildLimiterAddPayload(limiterID, speed) + if _, err := h.sendNodeCommand(nodeID, "AddLimiters", addPayload, false, false); err != nil { + if !isAlreadyExistsMessage(err.Error()) { + return err + } + payload := map[string]interface{}{ + "name": name, + "limits": addPayload["limits"], + } + if _, updateErr := h.sendNodeCommand(nodeID, "UpdateLimiters", buildLimiterUpdatePayload(name, payload), false, false); updateErr != nil { + return updateErr + } + } + + return nil +} diff --git a/go-backend/internal/http/handler/federation.go b/go-backend/internal/http/handler/federation.go index fe2cf57..ee33d3f 100644 --- a/go-backend/internal/http/handler/federation.go +++ b/go-backend/internal/http/handler/federation.go @@ -1472,7 +1472,7 @@ func (h *Handler) releasePeerShareForwardRuntimeServices(share *repo.PeerShare, func isFederationRuntimeCommandAllowed(commandType string) bool { switch strings.ToLower(strings.TrimSpace(commandType)) { - case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice", "addchains", "deletechains", "addlimiters", "deletelimiters", "tcpping", "reload": + case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice", "addchains", "deletechains", "addlimiters", "updatelimiters", "deletelimiters", "tcpping", "reload": return true default: return false diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 264f654..38aaac3 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -159,7 +159,6 @@ func (h *Handler) Register(mux *http.ServeMux) { mux.HandleFunc("/api/v1/speed-limit/create", h.speedLimitCreate) mux.HandleFunc("/api/v1/speed-limit/update", h.speedLimitUpdate) mux.HandleFunc("/api/v1/speed-limit/delete", h.speedLimitDelete) - mux.HandleFunc("/api/v1/speed-limit/tunnels", h.tunnelList) mux.HandleFunc("/api/v1/tunnel/user/tunnel", h.userTunnelVisibleList) mux.HandleFunc("/api/v1/tunnel/user/list", h.userTunnelList) mux.HandleFunc("/api/v1/group/tunnel/list", h.tunnelGroupList) diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 5b09d74..a17e9a5 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -1659,28 +1659,13 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) { speed := asInt(req["speed"], 100) - var tunnelID *int64 - var tunnelName string - if tid := asInt64(req["tunnelId"], 0); tid > 0 { - tunnelID = &tid - tunnelName = h.repo.GetTunnelNameByID(tid) - if tunnelName == "" { - response.WriteJSON(w, response.ErrDefault("隧道不存在")) - return - } - } - now := time.Now().UnixMilli() - id, err := h.repo.CreateSpeedLimit(name, speed, tunnelID, tunnelName, now, asInt(req["status"], 1)) + _, err := h.repo.CreateSpeedLimit(name, speed, now, asInt(req["status"], 1)) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if tunnelID != nil && *tunnelID > 0 { - _ = h.sendLimiterConfig(id, speed, *tunnelID) - } - response.WriteJSON(w, response.OKEmpty()) } @@ -1705,26 +1690,11 @@ func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) { speed := asInt(req["speed"], 100) - var tunnelID *int64 - var tunnelName string - if tid := asInt64(req["tunnelId"], 0); tid > 0 { - tunnelID = &tid - tunnelName = h.repo.GetTunnelNameByID(tid) - if tunnelName == "" { - response.WriteJSON(w, response.ErrDefault("隧道不存在")) - return - } - } - - if err := h.repo.UpdateSpeedLimit(id, name, speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli()); err != nil { + if err := h.repo.UpdateSpeedLimit(id, name, speed, asInt(req["status"], 1), time.Now().UnixMilli()); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if tunnelID != nil && *tunnelID > 0 { - _ = h.sendLimiterConfig(id, speed, *tunnelID) - } - response.WriteJSON(w, response.OKEmpty()) } @@ -1734,17 +1704,11 @@ func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) { return } - tunnelID := h.repo.GetSpeedLimitTunnelID(id) - if err := h.repo.DeleteSpeedLimit(id); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if tunnelID.Valid && tunnelID.Int64 > 0 { - _ = h.sendDeleteLimiterConfig(id, tunnelID.Int64) - } - response.WriteJSON(w, response.OKEmpty()) } diff --git a/go-backend/internal/store/repo/repository.go b/go-backend/internal/store/repo/repository.go index 49e04fb..6f42cc1 100644 --- a/go-backend/internal/store/repo/repository.go +++ b/go-backend/internal/store/repo/repository.go @@ -688,12 +688,6 @@ func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) { "status": sl.Status, "createdTime": sl.CreatedTime, "updatedTime": nullableInt64(sl.UpdatedTime), } - if sl.TunnelID.Valid { - item["tunnelId"] = sl.TunnelID.Int64 - } - if sl.TunnelName.Valid { - item["tunnelName"] = sl.TunnelName.String - } items = append(items, item) } return items, nil @@ -1825,13 +1819,6 @@ func (r *Repository) exportSpeedLimits() ([]model.SpeedLimitBackup, error) { ID: sl.ID, Name: sl.Name, Speed: int64(sl.Speed), CreatedTime: sl.CreatedTime, Status: sl.Status, } - if sl.TunnelID.Valid { - tid := sl.TunnelID.Int64 - b.TunnelID = &tid - } - if sl.TunnelName.Valid { - b.TunnelName = sl.TunnelName.String - } if sl.UpdatedTime.Valid { b.UpdatedTime = sl.UpdatedTime.Int64 } @@ -2208,12 +2195,6 @@ func importSpeedLimits(tx *gorm.DB, speedLimits []model.SpeedLimitBackup, now in UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, Status: sl.Status, } - if sl.TunnelID != nil { - item.TunnelID = sql.NullInt64{Int64: *sl.TunnelID, Valid: true} - } - if sl.TunnelName != "" { - item.TunnelName = sql.NullString{String: sl.TunnelName, Valid: true} - } err := tx.Clauses(clause.OnConflict{ Columns: []clause.Column{{Name: "id"}}, DoUpdates: clause.AssignmentColumns([]string{ @@ -2482,10 +2463,11 @@ func (r *Repository) GetUserTunnelByID(id int64) (*model.UserTunnel, error) { // ─── Migration ─────────────────────────────────────────────────────── -const currentSchemaVersion = 3 +const currentSchemaVersion = 4 var ensurePostgresIDDefaultsFn = ensurePostgresIDDefaults var migrateViteConfigValueColumnTypeFn = migrateViteConfigValueColumnType +var migrateSpeedLimitTunnelBindingFn = migrateSpeedLimitTunnelBinding func getSchemaVersion(db *gorm.DB) int { var v model.SchemaVersion @@ -2543,6 +2525,12 @@ func migrateSchema(db *gorm.DB) error { } } + if ver < 4 { + if err := migrateSpeedLimitTunnelBindingFn(db); err != nil { + return err + } + } + setSchemaVersion(db, currentSchemaVersion) return nil } @@ -2586,6 +2574,27 @@ func migrateViteConfigValueColumnType(db *gorm.DB) error { return nil } +func migrateSpeedLimitTunnelBinding(db *gorm.DB) error { + if db == nil { + return errors.New("nil db") + } + + if !db.Migrator().HasTable(&model.SpeedLimit{}) { + return nil + } + + if err := db.Model(&model.SpeedLimit{}). + Where("tunnel_id IS NOT NULL OR tunnel_name IS NOT NULL"). + UpdateColumns(map[string]interface{}{ + "tunnel_id": nil, + "tunnel_name": nil, + }).Error; err != nil { + return fmt.Errorf("clear speed_limit tunnel binding: %w", err) + } + + return nil +} + func ensurePostgresIDDefaults(db *gorm.DB) error { if db.Dialector.Name() != "postgres" { return nil diff --git a/go-backend/internal/store/repo/repository_migrate_test.go b/go-backend/internal/store/repo/repository_migrate_test.go index a5c53ec..663a68c 100644 --- a/go-backend/internal/store/repo/repository_migrate_test.go +++ b/go-backend/internal/store/repo/repository_migrate_test.go @@ -1,6 +1,7 @@ package repo import ( + "database/sql" "errors" "testing" @@ -175,3 +176,77 @@ func TestMigrateSchemaReturnsViteConfigMigrationError(t *testing.T) { t.Fatalf("expected error %v, got %v", wantErr, err) } } + +func TestMigrateSchemaClearsSpeedLimitTunnelBinding(t *testing.T) { + db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) + if err != nil { + t.Fatalf("open sqlite: %v", err) + } + t.Cleanup(func() { + sqlDB, _ := db.DB() + if sqlDB != nil { + _ = sqlDB.Close() + } + }) + + if err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`).Error; err != nil { + t.Fatalf("create schema_version: %v", err) + } + if err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, 3).Error; err != nil { + t.Fatalf("seed schema_version: %v", err) + } + if err := db.Exec(` + CREATE TABLE speed_limit ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name VARCHAR(100) NOT NULL, + speed INTEGER NOT NULL, + tunnel_id INTEGER, + tunnel_name VARCHAR(100), + created_time INTEGER NOT NULL, + updated_time INTEGER, + status INTEGER NOT NULL + ) + `).Error; err != nil { + t.Fatalf("create speed_limit: %v", err) + } + if err := db.Exec(` + INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) + VALUES(?, ?, ?, ?, ?, ?, ?) + `, "legacy-speed-limit", 100, 101, "legacy-tunnel", 1, 1, 1).Error; err != nil { + t.Fatalf("seed speed_limit: %v", err) + } + + originalIDRepair := ensurePostgresIDDefaultsFn + ensurePostgresIDDefaultsFn = func(db *gorm.DB) error { + return nil + } + t.Cleanup(func() { + ensurePostgresIDDefaultsFn = originalIDRepair + }) + + if err := migrateSchema(db); err != nil { + t.Fatalf("migrateSchema: %v", err) + } + + var tunnelID sql.NullInt64 + var tunnelName sql.NullString + if err := db.Raw(`SELECT tunnel_id, tunnel_name FROM speed_limit WHERE name = ?`, "legacy-speed-limit").Row().Scan(&tunnelID, &tunnelName); err != nil { + t.Fatalf("query speed_limit: %v", err) + } + if tunnelID.Valid { + t.Fatalf("expected tunnel_id cleared to NULL, got %d", tunnelID.Int64) + } + if tunnelName.Valid { + t.Fatalf("expected tunnel_name cleared to NULL, got %q", tunnelName.String) + } + + var schemaVersion int + if err := db.Raw(`SELECT version FROM schema_version LIMIT 1`).Row().Scan(&schemaVersion); err != nil { + t.Fatalf("query schema_version: %v", err) + } + if schemaVersion != currentSchemaVersion { + t.Fatalf("expected schema version %d, got %d", currentSchemaVersion, schemaVersion) + } +} diff --git a/go-backend/internal/store/repo/repository_mutations.go b/go-backend/internal/store/repo/repository_mutations.go index 18f0b8e..c7fb53b 100644 --- a/go-backend/internal/store/repo/repository_mutations.go +++ b/go-backend/internal/store/repo/repository_mutations.go @@ -522,9 +522,6 @@ func (r *Repository) DeleteTunnelCascade(tunnelID int64) error { if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.UserTunnel{}).Error; err != nil { return err } - if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.SpeedLimit{}).Error; err != nil { - return err - } if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.ChainTunnel{}).Error; err != nil { return err } @@ -535,17 +532,6 @@ func (r *Repository) DeleteTunnelCascade(tunnelID int64) error { }) } -func (r *Repository) GetTunnelNameByID(tunnelID int64) string { - if r == nil || r.db == nil { - return "" - } - var tunnel model.Tunnel - if err := r.db.Select("name").Where("id = ?", tunnelID).First(&tunnel).Error; err != nil { - return "" - } - return tunnel.Name -} - func (r *Repository) TunnelEntryNodeIDs(tunnelID int64) ([]int64, error) { if r == nil || r.db == nil { return nil, errors.New("repository not initialized") @@ -766,7 +752,7 @@ func (r *Repository) GetUsedPortsOnNodeAsMap(nodeID int64) (map[int]bool, error) return used, nil } -func (r *Repository) CreateSpeedLimit(name string, speed int, tunnelID *int64, tunnelName string, now int64, status int) (int64, error) { +func (r *Repository) CreateSpeedLimit(name string, speed int, now int64, status int) (int64, error) { if r == nil || r.db == nil { return 0, errors.New("repository not initialized") } @@ -779,57 +765,32 @@ func (r *Repository) CreateSpeedLimit(name string, speed int, tunnelID *int64, t UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, Status: status, } - if tunnelID != nil { - sl.TunnelID = sql.NullInt64{Int64: *tunnelID, Valid: true} - } - if tunnelName != "" { - sl.TunnelName = sql.NullString{String: tunnelName, Valid: true} - } if err := r.db.Create(&sl).Error; err != nil { return 0, err } return sl.ID, nil } -func (r *Repository) UpdateSpeedLimit(id int64, name string, speed int, tunnelID *int64, tunnelName string, status int, now int64) error { +func (r *Repository) UpdateSpeedLimit(id int64, name string, speed int, status int, now int64) error { if r == nil || r.db == nil { return errors.New("repository not initialized") } updates := map[string]interface{}{ - "name": name, - "speed": speed, - "status": status, + "name": name, + "speed": speed, + "status": status, + "tunnel_id": nil, + "tunnel_name": nil, "updated_time": sql.NullInt64{ Int64: now, Valid: true, }, } - if tunnelID != nil { - updates["tunnel_id"] = sql.NullInt64{Int64: *tunnelID, Valid: true} - } else { - updates["tunnel_id"] = sql.NullInt64{Int64: 0, Valid: false} - } - if tunnelName != "" { - updates["tunnel_name"] = sql.NullString{String: tunnelName, Valid: true} - } else { - updates["tunnel_name"] = sql.NullString{String: "", Valid: false} - } return r.db.Model(&model.SpeedLimit{}). Where("id = ?", id). Updates(updates).Error } -func (r *Repository) GetSpeedLimitTunnelID(speedLimitID int64) sql.NullInt64 { - if r == nil || r.db == nil { - return sql.NullInt64{Valid: false} - } - var sl model.SpeedLimit - if err := r.db.Select("tunnel_id").Where("id = ?", speedLimitID).First(&sl).Error; err != nil { - return sql.NullInt64{Valid: false} - } - return sl.TunnelID -} - func (r *Repository) DeleteSpeedLimit(id int64) error { if r == nil || r.db == nil { return errors.New("repository not initialized") diff --git a/go-backend/tests/contract/federation_dual_panel_contract_test.go b/go-backend/tests/contract/federation_dual_panel_contract_test.go index 3e9cbfb..b73c277 100644 --- a/go-backend/tests/contract/federation_dual_panel_contract_test.go +++ b/go-backend/tests/contract/federation_dual_panel_contract_test.go @@ -759,6 +759,21 @@ func TestFederationRuntimeCommandPortRangeEnforcement(t *testing.T) { } // Test: Non-service commands should pass through without port validation + res = sendCommand("share-portrange-token", "UpdateLimiters", map[string]interface{}{ + "limiter": "federation-limit-test", + "data": map[string]interface{}{ + "name": "federation-limit-test", + "limits": []string{"$ 1MB 1MB"}, + }, + }) + 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 for UpdateLimiters command, got %d (msg: %s)", out.Code, out.Msg) + } + res = sendCommand("share-portrange-token", "reload", nil) out = response.R{} if err := json.NewDecoder(res.Body).Decode(&out); err != nil { diff --git a/go-backend/tests/contract/limiter_sync_failure_contract_test.go b/go-backend/tests/contract/limiter_sync_failure_contract_test.go index 568b5d9..88bbda5 100644 --- a/go-backend/tests/contract/limiter_sync_failure_contract_test.go +++ b/go-backend/tests/contract/limiter_sync_failure_contract_test.go @@ -99,6 +99,168 @@ func TestForwardCreateRollbackWhenLimiterDispatchFailsContract(t *testing.T) { } } +func TestForwardCreateSucceedsWhenLimiterAlreadyExistsAndUpdateSucceedsContract(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-exists-update-ok-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-exists-update-ok-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-exists-update-ok-node", "limiter-exists-update-ok-secret", "10.20.1.1", "10.20.1.1", "", "32200-32210", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { + t.Fatalf("insert node: %v", err) + } + nodeID := mustLastInsertID(t, r, "limiter-exists-update-ok-node") + + if err := r.DB().Exec(` + INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, 1, ?, 32201, '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-exists-update-ok-rule", 1024, now, 1).Error; err != nil { + t.Fatalf("insert speed limit: %v", err) + } + speedID := mustLastInsertID(t, r, "limiter-exists-update-ok-rule") + + stopNode := startMockNodeSessionWithCommandFailures(t, server.URL, "limiter-exists-update-ok-secret", map[string]string{ + "addlimiters": "limiter 8 already exists", + }) + defer stopNode() + + payload := map[string]interface{}{ + "name": "limiter-exists-update-ok-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 success when updater succeeds, got code=%d msg=%s", out.Code, out.Msg) + } + + forwardCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM forward WHERE name = ?`, "limiter-exists-update-ok-forward") + if forwardCount != 1 { + t.Fatalf("expected forward kept when update limiter succeeds, got count=%d", forwardCount) + } +} + +func TestForwardCreateRollbackWhenLimiterAlreadyExistsAndUpdateFailsContract(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-exists-update-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-exists-update-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-exists-update-fail-node", "limiter-exists-update-fail-secret", "10.20.2.1", "10.20.2.1", "", "32300-32310", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { + t.Fatalf("insert node: %v", err) + } + nodeID := mustLastInsertID(t, r, "limiter-exists-update-fail-node") + + if err := r.DB().Exec(` + INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, 1, ?, 32301, '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-exists-update-fail-rule", 1024, now, 1).Error; err != nil { + t.Fatalf("insert speed limit: %v", err) + } + speedID := mustLastInsertID(t, r, "limiter-exists-update-fail-rule") + + stopNode := startMockNodeSessionWithCommandFailures(t, server.URL, "limiter-exists-update-fail-secret", map[string]string{ + "addlimiters": "limiter 9 already exists", + "updatelimiters": "mock update limiters failed", + }) + defer stopNode() + + payload := map[string]interface{}{ + "name": "limiter-exists-update-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 when update limiter fails, got code=0") + } + if !strings.Contains(out.Msg, "mock update limiters failed") { + t.Fatalf("expected update failure message, got %q", out.Msg) + } + + forwardCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM forward WHERE name = ?`, "limiter-exists-update-fail-forward") + if forwardCount != 0 { + t.Fatalf("expected forward rollback delete when update limiter fails, got count=%d", forwardCount) + } +} + func TestForwardCreateRollbackWhenServiceDispatchReturnsAddressInUseContract(t *testing.T) { secret := "contract-jwt-secret" router, r := setupContractRouter(t, secret) diff --git a/go-backend/tests/contract/migration_contract_test.go b/go-backend/tests/contract/migration_contract_test.go index 9391cd5..6292e7a 100644 --- a/go-backend/tests/contract/migration_contract_test.go +++ b/go-backend/tests/contract/migration_contract_test.go @@ -209,39 +209,24 @@ func TestOpenAPISubStoreContracts(t *testing.T) { }) } -func TestSpeedLimitTunnelsRouteAlias(t *testing.T) { +func TestSpeedLimitTunnelsRouteRemoved(t *testing.T) { secret := "contract-jwt-secret" router, _ := setupContractRouter(t, secret) - t.Run("missing token blocked", func(t *testing.T) { - req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil) - resp := httptest.NewRecorder() + token, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate token: %v", err) + } - router.ServeHTTP(resp, req) + req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil) + req.Header.Set("Authorization", token) + resp := httptest.NewRecorder() - assertCodeMsg(t, resp, 401, "未登录或token已过期") - }) + router.ServeHTTP(resp, req) - t.Run("admin token receives success envelope", func(t *testing.T) { - token, err := auth.GenerateToken(1, "admin_user", 0, secret) - if err != nil { - t.Fatalf("generate token: %v", err) - } - - req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil) - req.Header.Set("Authorization", token) - resp := httptest.NewRecorder() - - router.ServeHTTP(resp, req) - - var out response.R - if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { - t.Fatalf("decode response: %v", err) - } - if out.Code != 0 { - t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg) - } - }) + if resp.Code != http.StatusNotFound { + t.Fatalf("expected status 404 after route removal, got %d", resp.Code) + } } func TestBackupExportImportRestoreContracts(t *testing.T) { diff --git a/go-backend/tests/contract/speed_limit_contract_test.go b/go-backend/tests/contract/speed_limit_contract_test.go index 00f945b..075667b 100644 --- a/go-backend/tests/contract/speed_limit_contract_test.go +++ b/go-backend/tests/contract/speed_limit_contract_test.go @@ -15,7 +15,6 @@ import ( "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) @@ -25,8 +24,7 @@ func TestSpeedLimitWithoutTunnelContract(t *testing.T) { 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) { + t.Run("create speed limit", 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) @@ -37,8 +35,7 @@ func TestSpeedLimitWithoutTunnelContract(t *testing.T) { assertCode(t, res, 0) }) - // Verify the speed limit has null tunnelId - t.Run("list speed limits shows null tunnelId", func(t *testing.T) { + t.Run("list does not expose tunnel binding fields", func(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil) req.Header.Set("Authorization", adminToken) res := httptest.NewRecorder() @@ -57,31 +54,28 @@ func TestSpeedLimitWithoutTunnelContract(t *testing.T) { 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 m["name"] != "test-limit-no-tunnel" { + continue } + if tunnelID, exists := m["tunnelId"]; exists && tunnelID != nil { + t.Fatalf("expected tunnelId to be absent or nil, got %v", tunnelID) + } + if tunnelName, exists := m["tunnelName"]; exists && tunnelName != nil && tunnelName != "" { + t.Fatalf("expected tunnelName to be absent or empty, got %v", tunnelName) + } + return } - if !found { - t.Fatal("speed limit 'test-limit-no-tunnel' not found in list") - } + 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) { +func TestSpeedLimitCreateIgnoresTunnelBindingContract(t *testing.T) { secret := "contract-jwt-secret" router, r := setupContractRouter(t, secret) @@ -90,71 +84,56 @@ func TestSpeedLimitWithTunnelContract(t *testing.T) { t.Fatalf("generate admin token: %v", err) } - // First create a tunnel - tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-tunnel-for-limit") + tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-speed-limit-create-ignore-tunnel") - // 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) + body := `{"name":"test-limit-ignore-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) - }) + 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) + 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) - } + 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{}) + 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 { - t.Fatalf("expected data to be array, got %T", out.Data) + continue } + if m["name"] != "test-limit-ignore-tunnel" { + continue + } + if tunnelIDVal, exists := m["tunnelId"]; exists && tunnelIDVal != nil { + t.Fatalf("expected tunnelId ignored and nil, got %v", tunnelIDVal) + } + if tunnelNameVal, exists := m["tunnelName"]; exists && tunnelNameVal != nil && tunnelNameVal != "" { + t.Fatalf("expected tunnelName ignored and empty, got %v", tunnelNameVal) + } + return + } - 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") - } - }) + t.Fatal("speed limit 'test-limit-ignore-tunnel' not found in list") } -// TestSpeedLimitUpdateTunnelBindingContract tests updating speed limit tunnel binding -func TestSpeedLimitUpdateTunnelBindingContract(t *testing.T) { +func TestSpeedLimitUpdateIgnoresTunnelBindingContract(t *testing.T) { secret := "contract-jwt-secret" router, r := setupContractRouter(t, secret) @@ -163,109 +142,56 @@ func TestSpeedLimitUpdateTunnelBindingContract(t *testing.T) { t.Fatalf("generate admin token: %v", err) } - // Create a tunnel - tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-tunnel-update") + tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-speed-limit-update-ignore-tunnel") + speedLimitID := mustCreateSpeedLimitRepo(t, r, "test-limit-update-ignore-tunnel") - // Create a speed limit without tunnel - speedLimitID := mustCreateSpeedLimitRepo(t, r, "test-limit-update", 0) + body := `{"id":` + jsonInt(speedLimitID) + `,"name":"test-limit-update-ignore-tunnel","speed":256,"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) - // 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) - assertCode(t, res, 0) - }) + req = httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil) + req.Header.Set("Authorization", adminToken) + res = httptest.NewRecorder() + router.ServeHTTP(res, req) - // 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) + } - 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) + } - data, ok := out.Data.([]interface{}) + for _, item := range data { + m, ok := item.(map[string]interface{}) if !ok { - t.Fatalf("expected data to be array, got %T", out.Data) + continue } - - 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 - } + if m["name"] != "test-limit-update-ignore-tunnel" { + continue } - 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 tunnelIDVal, exists := m["tunnelId"]; exists && tunnelIDVal != nil { + t.Fatalf("expected tunnelId ignored and nil after update, got %v", tunnelIDVal) } - if out.Code != 0 { - t.Fatalf("expected code 0, got %d", out.Code) + if speedVal, ok := m["speed"].(float64); !ok || int(speedVal) != 256 { + t.Fatalf("expected speed 256 after update, got %v", m["speed"]) } + return + } - 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") - }) + t.Fatal("speed limit 'test-limit-update-ignore-tunnel' not found in list") } -// 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) @@ -274,132 +200,65 @@ func TestSpeedLimitDatabaseNullableFields(t *testing.T) { } 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) - } - }) + id, err := r.CreateSpeedLimit("db-test-limit", 100, 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) - } - }) + var tunnelID sql.NullInt64 + var tunnelName sql.NullString + err = r.DB().Raw("SELECT tunnel_id, tunnel_name FROM speed_limit WHERE id = ?", id).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) + } } -// TestSpeedLimitUpdateUnbindFromTunnel tests unbinding a speed limit from a tunnel -func TestSpeedLimitUpdateUnbindFromTunnel(t *testing.T) { - dbPath := filepath.Join(t.TempDir(), "speed-limit-unbind.db") +func TestSpeedLimitUpdateClearsHistoricalBinding(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "speed-limit-update-clear.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") + tunnelID := mustCreateSpeedLimitTunnel(t, r, "speed-limit-update-clear-tunnel") + now := time.Now().UnixMilli() + if err := r.DB().Exec(` + INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) + VALUES(?, ?, ?, ?, ?, ?, ?) + `, "speed-limit-update-clear", 300, tunnelID, "speed-limit-update-clear-tunnel", now, now, 1).Error; err != nil { + t.Fatalf("insert speed limit with tunnel binding: %v", err) + } + speedLimitID := mustLastInsertID(t, r, "speed-limit-update-clear") - // Create speed limit bound to tunnel - speedLimitID, err := r.CreateSpeedLimit("unbind-test-limit", 300, &tunnelID, "unbind-test-tunnel", 1, 1) + err = r.UpdateSpeedLimit(speedLimitID, "speed-limit-update-clear", 512, 1, time.Now().UnixMilli()) if err != nil { - t.Fatalf("create speed limit: %v", err) + t.Fatalf("UpdateSpeedLimit failed: %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) - } - }) + var dbTunnelID sql.NullInt64 + var dbTunnelName sql.NullString + err = r.DB().Raw("SELECT tunnel_id, tunnel_name FROM speed_limit WHERE id = ?", speedLimitID).Row().Scan(&dbTunnelID, &dbTunnelName) + if err != nil { + t.Fatalf("query updated speed limit failed: %v", err) + } + if dbTunnelID.Valid { + t.Fatalf("expected tunnel_id cleared after update, got %d", dbTunnelID.Int64) + } + if dbTunnelName.Valid && dbTunnelName.String != "" { + t.Fatalf("expected tunnel_name cleared after update, got %q", dbTunnelName.String) + } } -// TestSpeedLimitGetSpeed tests the GetSpeedLimitSpeed function func TestSpeedLimitGetSpeed(t *testing.T) { dbPath := filepath.Join(t.TempDir(), "speed-limit-getspeed.db") r, err := repo.Open(dbPath) @@ -408,13 +267,11 @@ func TestSpeedLimitGetSpeed(t *testing.T) { } t.Cleanup(func() { _ = r.Close() }) - // Create speed limit - speedLimitID, err := r.CreateSpeedLimit("get-speed-test", 500, nil, "", 1, 1) + speedLimitID, err := r.CreateSpeedLimit("get-speed-test", 500, 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 { @@ -433,8 +290,6 @@ func TestSpeedLimitGetSpeed(t *testing.T) { }) } -// Helper functions - func mustCreateSpeedLimitTunnel(t *testing.T, r *repo.Repository, name string) int64 { t.Helper() now := time.Now().UnixMilli() @@ -447,14 +302,10 @@ func mustCreateSpeedLimitTunnel(t *testing.T, r *repo.Repository, name string) i return mustLastInsertID(t, r, name) } -func mustCreateSpeedLimitRepo(t *testing.T, r *repo.Repository, name string, tunnelID int64) int64 { +func mustCreateSpeedLimitRepo(t *testing.T, r *repo.Repository, name string) int64 { t.Helper() now := time.Now().UnixMilli() - var tid *int64 - if tunnelID > 0 { - tid = &tunnelID - } - id, err := r.CreateSpeedLimit(name, 100, tid, "", now, 1) + id, err := r.CreateSpeedLimit(name, 100, now, 1) if err != nil { t.Fatalf("create speed limit failed: %v", err) } diff --git a/vite-frontend/src/api/types.ts b/vite-frontend/src/api/types.ts index 55cce65..94516c1 100644 --- a/vite-frontend/src/api/types.ts +++ b/vite-frontend/src/api/types.ts @@ -97,10 +97,8 @@ export interface StatisticsFlowApiItem { export interface SpeedLimitApiItem { id: number; name: string; - tunnelId?: number | null; speed: number; status: number; - tunnelName?: string; createdTime: string; updatedTime: string; uploadSpeed?: number; @@ -293,8 +291,6 @@ export interface SpeedLimitMutationPayload { name?: string; speed?: number; status?: number; - tunnelId?: number | null; - tunnelName?: string; } export interface UpdatePasswordPayload { diff --git a/vite-frontend/src/pages/limit.tsx b/vite-frontend/src/pages/limit.tsx index 13dc7d3..664ea3b 100644 --- a/vite-frontend/src/pages/limit.tsx +++ b/vite-frontend/src/pages/limit.tsx @@ -10,7 +10,6 @@ import { SearchBar } from "@/components/search-bar"; import { Card, CardBody, CardHeader } from "@/shadcn-bridge/heroui/card"; import { Button } from "@/shadcn-bridge/heroui/button"; import { Input } from "@/shadcn-bridge/heroui/input"; -import { Select, SelectItem } from "@/shadcn-bridge/heroui/select"; import { Modal, ModalContent, @@ -24,7 +23,6 @@ import { getSpeedLimitList, updateSpeedLimit, deleteSpeedLimit, - getTunnelList, } from "@/api"; import { PageLoadingState } from "@/components/page-state"; import { useLocalStorageState } from "@/hooks/use-local-storage-state"; @@ -34,30 +32,20 @@ interface SpeedLimitRule { name: string; speed: number; status: number; - tunnelId?: number | null; - tunnelName?: string; createdTime: string; updatedTime: string; } -interface Tunnel { - id: number; - name: string; -} - interface SpeedLimitForm { id?: number; name: string; speed: number; - tunnelId: number | null; - tunnelName: string; status: number; } export default function LimitPage() { const [loading, setLoading] = useState(true); const [rules, setRules] = useState([]); - const [tunnels, setTunnels] = useState([]); const [searchKeyword, setSearchKeyword] = useLocalStorageState( "limit-search-keyword", "", @@ -69,9 +57,7 @@ export default function LimitPage() { const lowerKeyword = searchKeyword.toLowerCase(); return rules.filter( - (r) => - (r.name && r.name.toLowerCase().includes(lowerKeyword)) || - (r.tunnelName && r.tunnelName.toLowerCase().includes(lowerKeyword)), + (r) => r.name && r.name.toLowerCase().includes(lowerKeyword), ); }, [rules, searchKeyword]); @@ -87,8 +73,6 @@ export default function LimitPage() { const [form, setForm] = useState({ name: "", speed: 100, - tunnelId: null, - tunnelName: "", status: 1, }); @@ -103,21 +87,13 @@ export default function LimitPage() { const loadData = async () => { setLoading(true); try { - const [rulesRes, tunnelsRes] = await Promise.all([ - getSpeedLimitList(), - getTunnelList(), - ]); + const rulesRes = await getSpeedLimitList(); if (rulesRes.code === 0) { setRules(rulesRes.data || []); } else { toast.error(rulesRes.msg || "获取限速规则失败"); } - - if (tunnelsRes.code === 0) { - setTunnels(tunnelsRes.data || []); - } else { - } } catch { toast.error("加载数据失败"); } finally { @@ -139,8 +115,6 @@ export default function LimitPage() { newErrors.speed = "请输入有效的速度限制(≥1 Mbps)"; } - // tunnelId is optional - speed limits can be created without binding to a tunnel - setErrors(newErrors); return Object.keys(newErrors).length === 0; @@ -152,8 +126,6 @@ export default function LimitPage() { setForm({ name: "", speed: 100, - tunnelId: null, - tunnelName: "", status: 1, }); setErrors({}); @@ -167,8 +139,6 @@ export default function LimitPage() { id: rule.id, name: rule.name, speed: rule.speed, - tunnelId: rule.tunnelId ?? null, - tunnelName: rule.tunnelName ?? "", status: rule.status, }); setErrors({}); @@ -210,16 +180,21 @@ export default function LimitPage() { setSubmitLoading(true); try { let res: { code: number; msg: string }; + const payload = { + id: form.id, + name: form.name, + speed: form.speed, + status: form.status, + }; if (isEdit) { - res = await updateSpeedLimit(form); + res = await updateSpeedLimit(payload); } else { - const createData = { ...form }; - - delete createData.id; - createData.tunnelId = null; - createData.tunnelName = ""; - + const createData = { + name: payload.name, + speed: payload.speed, + status: payload.status, + }; res = await createSpeedLimit(createData); } @@ -247,7 +222,7 @@ export default function LimitPage() {
setIsSearchVisible(false)} @@ -292,20 +267,6 @@ export default function LimitPage() { {rule.speed} Mbps
-
- - 绑定隧道 - - {rule.tunnelName ? ( - - {rule.tunnelName} - - ) : ( - - 未绑定 - - )} -
@@ -432,45 +393,6 @@ export default function LimitPage() { })) } /> - - {isEdit && ( - - )}
diff --git a/vite-frontend/src/types/index.ts b/vite-frontend/src/types/index.ts index f5b3975..8d52797 100644 --- a/vite-frontend/src/types/index.ts +++ b/vite-frontend/src/types/index.ts @@ -88,7 +88,6 @@ export interface Tunnel { export interface SpeedLimit { id: number; name: string; - tunnelId?: number | null; speed?: number; uploadSpeed: number; downloadSpeed: number;