mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-09 03:06:37 +08:00
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
This commit is contained in:
@@ -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 {
|
func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) error {
|
||||||
rate := float64(speed) / 8.0
|
if err := h.upsertLimiterOnNode(nodeID, limiterID, speed); err != nil {
|
||||||
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 {
|
|
||||||
return fmt.Errorf("限速规则下发失败: %w", err)
|
return fmt.Errorf("限速规则下发失败: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -1472,7 +1472,7 @@ func (h *Handler) releasePeerShareForwardRuntimeServices(share *repo.PeerShare,
|
|||||||
|
|
||||||
func isFederationRuntimeCommandAllowed(commandType string) bool {
|
func isFederationRuntimeCommandAllowed(commandType string) bool {
|
||||||
switch strings.ToLower(strings.TrimSpace(commandType)) {
|
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
|
return true
|
||||||
default:
|
default:
|
||||||
return false
|
return false
|
||||||
|
|||||||
@@ -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/create", h.speedLimitCreate)
|
||||||
mux.HandleFunc("/api/v1/speed-limit/update", h.speedLimitUpdate)
|
mux.HandleFunc("/api/v1/speed-limit/update", h.speedLimitUpdate)
|
||||||
mux.HandleFunc("/api/v1/speed-limit/delete", h.speedLimitDelete)
|
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/tunnel", h.userTunnelVisibleList)
|
||||||
mux.HandleFunc("/api/v1/tunnel/user/list", h.userTunnelList)
|
mux.HandleFunc("/api/v1/tunnel/user/list", h.userTunnelList)
|
||||||
mux.HandleFunc("/api/v1/group/tunnel/list", h.tunnelGroupList)
|
mux.HandleFunc("/api/v1/group/tunnel/list", h.tunnelGroupList)
|
||||||
|
|||||||
@@ -1659,28 +1659,13 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
|
|||||||
|
|
||||||
speed := asInt(req["speed"], 100)
|
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()
|
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 {
|
if err != nil {
|
||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if tunnelID != nil && *tunnelID > 0 {
|
|
||||||
_ = h.sendLimiterConfig(id, speed, *tunnelID)
|
|
||||||
}
|
|
||||||
|
|
||||||
response.WriteJSON(w, response.OKEmpty())
|
response.WriteJSON(w, response.OKEmpty())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1705,26 +1690,11 @@ func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) {
|
|||||||
|
|
||||||
speed := asInt(req["speed"], 100)
|
speed := asInt(req["speed"], 100)
|
||||||
|
|
||||||
var tunnelID *int64
|
if err := h.repo.UpdateSpeedLimit(id, name, speed, asInt(req["status"], 1), time.Now().UnixMilli()); err != nil {
|
||||||
var tunnelName string
|
|
||||||
if tid := asInt64(req["tunnelId"], 0); tid > 0 {
|
|
||||||
tunnelID = &tid
|
|
||||||
tunnelName = h.repo.GetTunnelNameByID(tid)
|
|
||||||
if tunnelName == "" {
|
|
||||||
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := h.repo.UpdateSpeedLimit(id, name, speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli()); err != nil {
|
|
||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if tunnelID != nil && *tunnelID > 0 {
|
|
||||||
_ = h.sendLimiterConfig(id, speed, *tunnelID)
|
|
||||||
}
|
|
||||||
|
|
||||||
response.WriteJSON(w, response.OKEmpty())
|
response.WriteJSON(w, response.OKEmpty())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1734,17 +1704,11 @@ func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
tunnelID := h.repo.GetSpeedLimitTunnelID(id)
|
|
||||||
|
|
||||||
if err := h.repo.DeleteSpeedLimit(id); err != nil {
|
if err := h.repo.DeleteSpeedLimit(id); err != nil {
|
||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if tunnelID.Valid && tunnelID.Int64 > 0 {
|
|
||||||
_ = h.sendDeleteLimiterConfig(id, tunnelID.Int64)
|
|
||||||
}
|
|
||||||
|
|
||||||
response.WriteJSON(w, response.OKEmpty())
|
response.WriteJSON(w, response.OKEmpty())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -688,12 +688,6 @@ func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) {
|
|||||||
"status": sl.Status, "createdTime": sl.CreatedTime,
|
"status": sl.Status, "createdTime": sl.CreatedTime,
|
||||||
"updatedTime": nullableInt64(sl.UpdatedTime),
|
"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)
|
items = append(items, item)
|
||||||
}
|
}
|
||||||
return items, nil
|
return items, nil
|
||||||
@@ -1825,13 +1819,6 @@ func (r *Repository) exportSpeedLimits() ([]model.SpeedLimitBackup, error) {
|
|||||||
ID: sl.ID, Name: sl.Name, Speed: int64(sl.Speed),
|
ID: sl.ID, Name: sl.Name, Speed: int64(sl.Speed),
|
||||||
CreatedTime: sl.CreatedTime, Status: sl.Status,
|
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 {
|
if sl.UpdatedTime.Valid {
|
||||||
b.UpdatedTime = sl.UpdatedTime.Int64
|
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},
|
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
|
||||||
Status: sl.Status,
|
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{
|
err := tx.Clauses(clause.OnConflict{
|
||||||
Columns: []clause.Column{{Name: "id"}},
|
Columns: []clause.Column{{Name: "id"}},
|
||||||
DoUpdates: clause.AssignmentColumns([]string{
|
DoUpdates: clause.AssignmentColumns([]string{
|
||||||
@@ -2482,10 +2463,11 @@ func (r *Repository) GetUserTunnelByID(id int64) (*model.UserTunnel, error) {
|
|||||||
|
|
||||||
// ─── Migration ───────────────────────────────────────────────────────
|
// ─── Migration ───────────────────────────────────────────────────────
|
||||||
|
|
||||||
const currentSchemaVersion = 3
|
const currentSchemaVersion = 4
|
||||||
|
|
||||||
var ensurePostgresIDDefaultsFn = ensurePostgresIDDefaults
|
var ensurePostgresIDDefaultsFn = ensurePostgresIDDefaults
|
||||||
var migrateViteConfigValueColumnTypeFn = migrateViteConfigValueColumnType
|
var migrateViteConfigValueColumnTypeFn = migrateViteConfigValueColumnType
|
||||||
|
var migrateSpeedLimitTunnelBindingFn = migrateSpeedLimitTunnelBinding
|
||||||
|
|
||||||
func getSchemaVersion(db *gorm.DB) int {
|
func getSchemaVersion(db *gorm.DB) int {
|
||||||
var v model.SchemaVersion
|
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)
|
setSchemaVersion(db, currentSchemaVersion)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -2586,6 +2574,27 @@ func migrateViteConfigValueColumnType(db *gorm.DB) error {
|
|||||||
return nil
|
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 {
|
func ensurePostgresIDDefaults(db *gorm.DB) error {
|
||||||
if db.Dialector.Name() != "postgres" {
|
if db.Dialector.Name() != "postgres" {
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package repo
|
package repo
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"database/sql"
|
||||||
"errors"
|
"errors"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -175,3 +176,77 @@ func TestMigrateSchemaReturnsViteConfigMigrationError(t *testing.T) {
|
|||||||
t.Fatalf("expected error %v, got %v", wantErr, err)
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -522,9 +522,6 @@ func (r *Repository) DeleteTunnelCascade(tunnelID int64) error {
|
|||||||
if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.UserTunnel{}).Error; err != nil {
|
if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.UserTunnel{}).Error; err != nil {
|
||||||
return err
|
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 {
|
if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.ChainTunnel{}).Error; err != nil {
|
||||||
return err
|
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) {
|
func (r *Repository) TunnelEntryNodeIDs(tunnelID int64) ([]int64, error) {
|
||||||
if r == nil || r.db == nil {
|
if r == nil || r.db == nil {
|
||||||
return nil, errors.New("repository not initialized")
|
return nil, errors.New("repository not initialized")
|
||||||
@@ -766,7 +752,7 @@ func (r *Repository) GetUsedPortsOnNodeAsMap(nodeID int64) (map[int]bool, error)
|
|||||||
return used, nil
|
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 {
|
if r == nil || r.db == nil {
|
||||||
return 0, errors.New("repository not initialized")
|
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},
|
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
|
||||||
Status: status,
|
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 {
|
if err := r.db.Create(&sl).Error; err != nil {
|
||||||
return 0, err
|
return 0, err
|
||||||
}
|
}
|
||||||
return sl.ID, nil
|
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 {
|
if r == nil || r.db == nil {
|
||||||
return errors.New("repository not initialized")
|
return errors.New("repository not initialized")
|
||||||
}
|
}
|
||||||
updates := map[string]interface{}{
|
updates := map[string]interface{}{
|
||||||
"name": name,
|
"name": name,
|
||||||
"speed": speed,
|
"speed": speed,
|
||||||
"status": status,
|
"status": status,
|
||||||
|
"tunnel_id": nil,
|
||||||
|
"tunnel_name": nil,
|
||||||
"updated_time": sql.NullInt64{
|
"updated_time": sql.NullInt64{
|
||||||
Int64: now,
|
Int64: now,
|
||||||
Valid: true,
|
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{}).
|
return r.db.Model(&model.SpeedLimit{}).
|
||||||
Where("id = ?", id).
|
Where("id = ?", id).
|
||||||
Updates(updates).Error
|
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 {
|
func (r *Repository) DeleteSpeedLimit(id int64) error {
|
||||||
if r == nil || r.db == nil {
|
if r == nil || r.db == nil {
|
||||||
return errors.New("repository not initialized")
|
return errors.New("repository not initialized")
|
||||||
|
|||||||
@@ -759,6 +759,21 @@ func TestFederationRuntimeCommandPortRangeEnforcement(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Test: Non-service commands should pass through without port validation
|
// 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)
|
res = sendCommand("share-portrange-token", "reload", nil)
|
||||||
out = response.R{}
|
out = response.R{}
|
||||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||||
|
|||||||
@@ -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) {
|
func TestForwardCreateRollbackWhenServiceDispatchReturnsAddressInUseContract(t *testing.T) {
|
||||||
secret := "contract-jwt-secret"
|
secret := "contract-jwt-secret"
|
||||||
router, r := setupContractRouter(t, secret)
|
router, r := setupContractRouter(t, secret)
|
||||||
|
|||||||
@@ -209,39 +209,24 @@ func TestOpenAPISubStoreContracts(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSpeedLimitTunnelsRouteAlias(t *testing.T) {
|
func TestSpeedLimitTunnelsRouteRemoved(t *testing.T) {
|
||||||
secret := "contract-jwt-secret"
|
secret := "contract-jwt-secret"
|
||||||
router, _ := setupContractRouter(t, secret)
|
router, _ := setupContractRouter(t, secret)
|
||||||
|
|
||||||
t.Run("missing token blocked", func(t *testing.T) {
|
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil)
|
if err != nil {
|
||||||
resp := httptest.NewRecorder()
|
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) {
|
if resp.Code != http.StatusNotFound {
|
||||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
t.Fatalf("expected status 404 after route removal, got %d", resp.Code)
|
||||||
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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBackupExportImportRestoreContracts(t *testing.T) {
|
func TestBackupExportImportRestoreContracts(t *testing.T) {
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ import (
|
|||||||
"go-backend/internal/store/repo"
|
"go-backend/internal/store/repo"
|
||||||
)
|
)
|
||||||
|
|
||||||
// TestSpeedLimitWithoutTunnelContract tests that speed limits can be created without binding to a tunnel
|
|
||||||
func TestSpeedLimitWithoutTunnelContract(t *testing.T) {
|
func TestSpeedLimitWithoutTunnelContract(t *testing.T) {
|
||||||
secret := "contract-jwt-secret"
|
secret := "contract-jwt-secret"
|
||||||
router, _ := setupContractRouter(t, secret)
|
router, _ := setupContractRouter(t, secret)
|
||||||
@@ -25,8 +24,7 @@ func TestSpeedLimitWithoutTunnelContract(t *testing.T) {
|
|||||||
t.Fatalf("generate admin token: %v", err)
|
t.Fatalf("generate admin token: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Create a speed limit without tunnel binding
|
t.Run("create speed limit", func(t *testing.T) {
|
||||||
t.Run("create speed limit without tunnel", func(t *testing.T) {
|
|
||||||
body := `{"name":"test-limit-no-tunnel","speed":100,"status":1}`
|
body := `{"name":"test-limit-no-tunnel","speed":100,"status":1}`
|
||||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/create", bytes.NewBufferString(body))
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/create", bytes.NewBufferString(body))
|
||||||
req.Header.Set("Authorization", adminToken)
|
req.Header.Set("Authorization", adminToken)
|
||||||
@@ -37,8 +35,7 @@ func TestSpeedLimitWithoutTunnelContract(t *testing.T) {
|
|||||||
assertCode(t, res, 0)
|
assertCode(t, res, 0)
|
||||||
})
|
})
|
||||||
|
|
||||||
// Verify the speed limit has null tunnelId
|
t.Run("list does not expose tunnel binding fields", func(t *testing.T) {
|
||||||
t.Run("list speed limits shows null tunnelId", func(t *testing.T) {
|
|
||||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil)
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil)
|
||||||
req.Header.Set("Authorization", adminToken)
|
req.Header.Set("Authorization", adminToken)
|
||||||
res := httptest.NewRecorder()
|
res := httptest.NewRecorder()
|
||||||
@@ -57,31 +54,28 @@ func TestSpeedLimitWithoutTunnelContract(t *testing.T) {
|
|||||||
t.Fatalf("expected data to be array, got %T", out.Data)
|
t.Fatalf("expected data to be array, got %T", out.Data)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Find our speed limit
|
|
||||||
var found bool
|
|
||||||
for _, item := range data {
|
for _, item := range data {
|
||||||
m, ok := item.(map[string]interface{})
|
m, ok := item.(map[string]interface{})
|
||||||
if !ok {
|
if !ok {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
if m["name"] == "test-limit-no-tunnel" {
|
if m["name"] != "test-limit-no-tunnel" {
|
||||||
found = true
|
continue
|
||||||
// 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 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 TestSpeedLimitCreateIgnoresTunnelBindingContract(t *testing.T) {
|
||||||
func TestSpeedLimitWithTunnelContract(t *testing.T) {
|
|
||||||
secret := "contract-jwt-secret"
|
secret := "contract-jwt-secret"
|
||||||
router, r := setupContractRouter(t, secret)
|
router, r := setupContractRouter(t, secret)
|
||||||
|
|
||||||
@@ -90,71 +84,56 @@ func TestSpeedLimitWithTunnelContract(t *testing.T) {
|
|||||||
t.Fatalf("generate admin token: %v", err)
|
t.Fatalf("generate admin token: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// First create a tunnel
|
tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-speed-limit-create-ignore-tunnel")
|
||||||
tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-tunnel-for-limit")
|
|
||||||
|
|
||||||
// Create a speed limit with tunnel binding
|
body := `{"name":"test-limit-ignore-tunnel","speed":200,"tunnelId":` + jsonInt(tunnelID) + `,"status":1}`
|
||||||
t.Run("create speed limit with tunnel", func(t *testing.T) {
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/create", bytes.NewBufferString(body))
|
||||||
body := `{"name":"test-limit-with-tunnel","speed":200,"tunnelId":` + jsonInt(tunnelID) + `,"status":1}`
|
req.Header.Set("Authorization", adminToken)
|
||||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/create", bytes.NewBufferString(body))
|
req.Header.Set("Content-Type", "application/json")
|
||||||
req.Header.Set("Authorization", adminToken)
|
res := httptest.NewRecorder()
|
||||||
req.Header.Set("Content-Type", "application/json")
|
router.ServeHTTP(res, req)
|
||||||
res := httptest.NewRecorder()
|
|
||||||
router.ServeHTTP(res, req)
|
|
||||||
|
|
||||||
assertCode(t, res, 0)
|
assertCode(t, res, 0)
|
||||||
})
|
|
||||||
|
|
||||||
// Verify the speed limit has the tunnelId
|
req = httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil)
|
||||||
t.Run("list speed limits shows tunnelId", func(t *testing.T) {
|
req.Header.Set("Authorization", adminToken)
|
||||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil)
|
res = httptest.NewRecorder()
|
||||||
req.Header.Set("Authorization", adminToken)
|
router.ServeHTTP(res, req)
|
||||||
res := httptest.NewRecorder()
|
|
||||||
router.ServeHTTP(res, req)
|
|
||||||
|
|
||||||
var out response.R
|
var out response.R
|
||||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||||
t.Fatalf("decode response: %v", err)
|
t.Fatalf("decode response: %v", err)
|
||||||
}
|
}
|
||||||
if out.Code != 0 {
|
if out.Code != 0 {
|
||||||
t.Fatalf("expected code 0, got %d", out.Code)
|
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 {
|
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
|
t.Fatal("speed limit 'test-limit-ignore-tunnel' not found in list")
|
||||||
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 TestSpeedLimitUpdateIgnoresTunnelBindingContract(t *testing.T) {
|
||||||
func TestSpeedLimitUpdateTunnelBindingContract(t *testing.T) {
|
|
||||||
secret := "contract-jwt-secret"
|
secret := "contract-jwt-secret"
|
||||||
router, r := setupContractRouter(t, secret)
|
router, r := setupContractRouter(t, secret)
|
||||||
|
|
||||||
@@ -163,109 +142,56 @@ func TestSpeedLimitUpdateTunnelBindingContract(t *testing.T) {
|
|||||||
t.Fatalf("generate admin token: %v", err)
|
t.Fatalf("generate admin token: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Create a tunnel
|
tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-speed-limit-update-ignore-tunnel")
|
||||||
tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-tunnel-update")
|
speedLimitID := mustCreateSpeedLimitRepo(t, r, "test-limit-update-ignore-tunnel")
|
||||||
|
|
||||||
// Create a speed limit without tunnel
|
body := `{"id":` + jsonInt(speedLimitID) + `,"name":"test-limit-update-ignore-tunnel","speed":256,"tunnelId":` + jsonInt(tunnelID) + `,"status":1}`
|
||||||
speedLimitID := mustCreateSpeedLimitRepo(t, r, "test-limit-update", 0)
|
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
|
assertCode(t, res, 0)
|
||||||
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)
|
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
|
var out response.R
|
||||||
t.Run("verify tunnel binding after update", func(t *testing.T) {
|
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil)
|
t.Fatalf("decode response: %v", err)
|
||||||
req.Header.Set("Authorization", adminToken)
|
}
|
||||||
res := httptest.NewRecorder()
|
if out.Code != 0 {
|
||||||
router.ServeHTTP(res, req)
|
t.Fatalf("expected code 0, got %d", out.Code)
|
||||||
|
}
|
||||||
|
|
||||||
var out response.R
|
data, ok := out.Data.([]interface{})
|
||||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
if !ok {
|
||||||
t.Fatalf("decode response: %v", err)
|
t.Fatalf("expected data to be array, got %T", out.Data)
|
||||||
}
|
}
|
||||||
if out.Code != 0 {
|
|
||||||
t.Fatalf("expected code 0, got %d", out.Code)
|
|
||||||
}
|
|
||||||
|
|
||||||
data, ok := out.Data.([]interface{})
|
for _, item := range data {
|
||||||
|
m, ok := item.(map[string]interface{})
|
||||||
if !ok {
|
if !ok {
|
||||||
t.Fatalf("expected data to be array, got %T", out.Data)
|
continue
|
||||||
}
|
}
|
||||||
|
if m["name"] != "test-limit-update-ignore-tunnel" {
|
||||||
for _, item := range data {
|
continue
|
||||||
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")
|
if tunnelIDVal, exists := m["tunnelId"]; exists && tunnelIDVal != nil {
|
||||||
})
|
t.Fatalf("expected tunnelId ignored and nil after update, got %v", tunnelIDVal)
|
||||||
|
|
||||||
// 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 {
|
if speedVal, ok := m["speed"].(float64); !ok || int(speedVal) != 256 {
|
||||||
t.Fatalf("expected code 0, got %d", out.Code)
|
t.Fatalf("expected speed 256 after update, got %v", m["speed"])
|
||||||
}
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
data, ok := out.Data.([]interface{})
|
t.Fatal("speed limit 'test-limit-update-ignore-tunnel' not found in list")
|
||||||
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) {
|
func TestSpeedLimitDatabaseNullableFields(t *testing.T) {
|
||||||
dbPath := filepath.Join(t.TempDir(), "speed-limit-null.db")
|
dbPath := filepath.Join(t.TempDir(), "speed-limit-null.db")
|
||||||
r, err := repo.Open(dbPath)
|
r, err := repo.Open(dbPath)
|
||||||
@@ -274,132 +200,65 @@ func TestSpeedLimitDatabaseNullableFields(t *testing.T) {
|
|||||||
}
|
}
|
||||||
t.Cleanup(func() { _ = r.Close() })
|
t.Cleanup(func() { _ = r.Close() })
|
||||||
|
|
||||||
// Create speed limit via repository
|
id, err := r.CreateSpeedLimit("db-test-limit", 100, 1, 1)
|
||||||
t.Run("repository create speed limit without tunnel", func(t *testing.T) {
|
if err != nil {
|
||||||
id, err := r.CreateSpeedLimit("db-test-limit", 100, nil, "", 1, 1)
|
t.Fatalf("CreateSpeedLimit failed: %v", err)
|
||||||
if err != nil {
|
}
|
||||||
t.Fatalf("CreateSpeedLimit failed: %v", err)
|
if id <= 0 {
|
||||||
}
|
t.Fatalf("expected valid id, got %d", id)
|
||||||
if id <= 0 {
|
}
|
||||||
t.Fatalf("expected valid id, got %d", id)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
// Verify TunnelID is null in database
|
var tunnelID sql.NullInt64
|
||||||
t.Run("verify null TunnelID in database", func(t *testing.T) {
|
var tunnelName sql.NullString
|
||||||
var tunnelID sql.NullInt64
|
err = r.DB().Raw("SELECT tunnel_id, tunnel_name FROM speed_limit WHERE id = ?", id).Row().Scan(&tunnelID, &tunnelName)
|
||||||
var tunnelName sql.NullString
|
if err != nil {
|
||||||
err := r.DB().Raw("SELECT tunnel_id, tunnel_name FROM speed_limit WHERE name = ?", "db-test-limit").Row().Scan(&tunnelID, &tunnelName)
|
t.Fatalf("query failed: %v", err)
|
||||||
if err != nil {
|
}
|
||||||
t.Fatalf("query failed: %v", err)
|
if tunnelID.Valid {
|
||||||
}
|
t.Fatalf("expected TunnelID to be NULL, got %d", tunnelID.Int64)
|
||||||
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)
|
||||||
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 TestSpeedLimitUpdateClearsHistoricalBinding(t *testing.T) {
|
||||||
func TestSpeedLimitUpdateUnbindFromTunnel(t *testing.T) {
|
dbPath := filepath.Join(t.TempDir(), "speed-limit-update-clear.db")
|
||||||
dbPath := filepath.Join(t.TempDir(), "speed-limit-unbind.db")
|
|
||||||
r, err := repo.Open(dbPath)
|
r, err := repo.Open(dbPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("open sqlite: %v", err)
|
t.Fatalf("open sqlite: %v", err)
|
||||||
}
|
}
|
||||||
t.Cleanup(func() { _ = r.Close() })
|
t.Cleanup(func() { _ = r.Close() })
|
||||||
|
|
||||||
// Create tunnel
|
tunnelID := mustCreateSpeedLimitTunnel(t, r, "speed-limit-update-clear-tunnel")
|
||||||
tunnelID := mustCreateSpeedLimitTunnel(t, r, "unbind-test-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
|
err = r.UpdateSpeedLimit(speedLimitID, "speed-limit-update-clear", 512, 1, time.Now().UnixMilli())
|
||||||
speedLimitID, err := r.CreateSpeedLimit("unbind-test-limit", 300, &tunnelID, "unbind-test-tunnel", 1, 1)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("create speed limit: %v", err)
|
t.Fatalf("UpdateSpeedLimit failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Verify initial binding
|
var dbTunnelID sql.NullInt64
|
||||||
t.Run("verify initial binding", func(t *testing.T) {
|
var dbTunnelName sql.NullString
|
||||||
result := r.GetSpeedLimitTunnelID(speedLimitID)
|
err = r.DB().Raw("SELECT tunnel_id, tunnel_name FROM speed_limit WHERE id = ?", speedLimitID).Row().Scan(&dbTunnelID, &dbTunnelName)
|
||||||
if !result.Valid {
|
if err != nil {
|
||||||
t.Fatal("expected initial binding to tunnel")
|
t.Fatalf("query updated speed limit failed: %v", err)
|
||||||
}
|
}
|
||||||
if result.Int64 != tunnelID {
|
if dbTunnelID.Valid {
|
||||||
t.Fatalf("expected tunnel ID %d, got %d", tunnelID, result.Int64)
|
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)
|
||||||
// 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) {
|
func TestSpeedLimitGetSpeed(t *testing.T) {
|
||||||
dbPath := filepath.Join(t.TempDir(), "speed-limit-getspeed.db")
|
dbPath := filepath.Join(t.TempDir(), "speed-limit-getspeed.db")
|
||||||
r, err := repo.Open(dbPath)
|
r, err := repo.Open(dbPath)
|
||||||
@@ -408,13 +267,11 @@ func TestSpeedLimitGetSpeed(t *testing.T) {
|
|||||||
}
|
}
|
||||||
t.Cleanup(func() { _ = r.Close() })
|
t.Cleanup(func() { _ = r.Close() })
|
||||||
|
|
||||||
// Create speed limit
|
speedLimitID, err := r.CreateSpeedLimit("get-speed-test", 500, 1, 1)
|
||||||
speedLimitID, err := r.CreateSpeedLimit("get-speed-test", 500, nil, "", 1, 1)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("create speed limit: %v", err)
|
t.Fatalf("create speed limit: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Test GetSpeedLimitSpeed
|
|
||||||
t.Run("GetSpeedLimitSpeed returns correct speed", func(t *testing.T) {
|
t.Run("GetSpeedLimitSpeed returns correct speed", func(t *testing.T) {
|
||||||
speed, err := r.GetSpeedLimitSpeed(speedLimitID)
|
speed, err := r.GetSpeedLimitSpeed(speedLimitID)
|
||||||
if err != nil {
|
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 {
|
func mustCreateSpeedLimitTunnel(t *testing.T, r *repo.Repository, name string) int64 {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
now := time.Now().UnixMilli()
|
now := time.Now().UnixMilli()
|
||||||
@@ -447,14 +302,10 @@ func mustCreateSpeedLimitTunnel(t *testing.T, r *repo.Repository, name string) i
|
|||||||
return mustLastInsertID(t, r, name)
|
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()
|
t.Helper()
|
||||||
now := time.Now().UnixMilli()
|
now := time.Now().UnixMilli()
|
||||||
var tid *int64
|
id, err := r.CreateSpeedLimit(name, 100, now, 1)
|
||||||
if tunnelID > 0 {
|
|
||||||
tid = &tunnelID
|
|
||||||
}
|
|
||||||
id, err := r.CreateSpeedLimit(name, 100, tid, "", now, 1)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("create speed limit failed: %v", err)
|
t.Fatalf("create speed limit failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -97,10 +97,8 @@ export interface StatisticsFlowApiItem {
|
|||||||
export interface SpeedLimitApiItem {
|
export interface SpeedLimitApiItem {
|
||||||
id: number;
|
id: number;
|
||||||
name: string;
|
name: string;
|
||||||
tunnelId?: number | null;
|
|
||||||
speed: number;
|
speed: number;
|
||||||
status: number;
|
status: number;
|
||||||
tunnelName?: string;
|
|
||||||
createdTime: string;
|
createdTime: string;
|
||||||
updatedTime: string;
|
updatedTime: string;
|
||||||
uploadSpeed?: number;
|
uploadSpeed?: number;
|
||||||
@@ -293,8 +291,6 @@ export interface SpeedLimitMutationPayload {
|
|||||||
name?: string;
|
name?: string;
|
||||||
speed?: number;
|
speed?: number;
|
||||||
status?: number;
|
status?: number;
|
||||||
tunnelId?: number | null;
|
|
||||||
tunnelName?: string;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface UpdatePasswordPayload {
|
export interface UpdatePasswordPayload {
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ import { SearchBar } from "@/components/search-bar";
|
|||||||
import { Card, CardBody, CardHeader } from "@/shadcn-bridge/heroui/card";
|
import { Card, CardBody, CardHeader } from "@/shadcn-bridge/heroui/card";
|
||||||
import { Button } from "@/shadcn-bridge/heroui/button";
|
import { Button } from "@/shadcn-bridge/heroui/button";
|
||||||
import { Input } from "@/shadcn-bridge/heroui/input";
|
import { Input } from "@/shadcn-bridge/heroui/input";
|
||||||
import { Select, SelectItem } from "@/shadcn-bridge/heroui/select";
|
|
||||||
import {
|
import {
|
||||||
Modal,
|
Modal,
|
||||||
ModalContent,
|
ModalContent,
|
||||||
@@ -24,7 +23,6 @@ import {
|
|||||||
getSpeedLimitList,
|
getSpeedLimitList,
|
||||||
updateSpeedLimit,
|
updateSpeedLimit,
|
||||||
deleteSpeedLimit,
|
deleteSpeedLimit,
|
||||||
getTunnelList,
|
|
||||||
} from "@/api";
|
} from "@/api";
|
||||||
import { PageLoadingState } from "@/components/page-state";
|
import { PageLoadingState } from "@/components/page-state";
|
||||||
import { useLocalStorageState } from "@/hooks/use-local-storage-state";
|
import { useLocalStorageState } from "@/hooks/use-local-storage-state";
|
||||||
@@ -34,30 +32,20 @@ interface SpeedLimitRule {
|
|||||||
name: string;
|
name: string;
|
||||||
speed: number;
|
speed: number;
|
||||||
status: number;
|
status: number;
|
||||||
tunnelId?: number | null;
|
|
||||||
tunnelName?: string;
|
|
||||||
createdTime: string;
|
createdTime: string;
|
||||||
updatedTime: string;
|
updatedTime: string;
|
||||||
}
|
}
|
||||||
|
|
||||||
interface Tunnel {
|
|
||||||
id: number;
|
|
||||||
name: string;
|
|
||||||
}
|
|
||||||
|
|
||||||
interface SpeedLimitForm {
|
interface SpeedLimitForm {
|
||||||
id?: number;
|
id?: number;
|
||||||
name: string;
|
name: string;
|
||||||
speed: number;
|
speed: number;
|
||||||
tunnelId: number | null;
|
|
||||||
tunnelName: string;
|
|
||||||
status: number;
|
status: number;
|
||||||
}
|
}
|
||||||
|
|
||||||
export default function LimitPage() {
|
export default function LimitPage() {
|
||||||
const [loading, setLoading] = useState(true);
|
const [loading, setLoading] = useState(true);
|
||||||
const [rules, setRules] = useState<SpeedLimitRule[]>([]);
|
const [rules, setRules] = useState<SpeedLimitRule[]>([]);
|
||||||
const [tunnels, setTunnels] = useState<Tunnel[]>([]);
|
|
||||||
const [searchKeyword, setSearchKeyword] = useLocalStorageState(
|
const [searchKeyword, setSearchKeyword] = useLocalStorageState(
|
||||||
"limit-search-keyword",
|
"limit-search-keyword",
|
||||||
"",
|
"",
|
||||||
@@ -69,9 +57,7 @@ export default function LimitPage() {
|
|||||||
const lowerKeyword = searchKeyword.toLowerCase();
|
const lowerKeyword = searchKeyword.toLowerCase();
|
||||||
|
|
||||||
return rules.filter(
|
return rules.filter(
|
||||||
(r) =>
|
(r) => r.name && r.name.toLowerCase().includes(lowerKeyword),
|
||||||
(r.name && r.name.toLowerCase().includes(lowerKeyword)) ||
|
|
||||||
(r.tunnelName && r.tunnelName.toLowerCase().includes(lowerKeyword)),
|
|
||||||
);
|
);
|
||||||
}, [rules, searchKeyword]);
|
}, [rules, searchKeyword]);
|
||||||
|
|
||||||
@@ -87,8 +73,6 @@ export default function LimitPage() {
|
|||||||
const [form, setForm] = useState<SpeedLimitForm>({
|
const [form, setForm] = useState<SpeedLimitForm>({
|
||||||
name: "",
|
name: "",
|
||||||
speed: 100,
|
speed: 100,
|
||||||
tunnelId: null,
|
|
||||||
tunnelName: "",
|
|
||||||
status: 1,
|
status: 1,
|
||||||
});
|
});
|
||||||
|
|
||||||
@@ -103,21 +87,13 @@ export default function LimitPage() {
|
|||||||
const loadData = async () => {
|
const loadData = async () => {
|
||||||
setLoading(true);
|
setLoading(true);
|
||||||
try {
|
try {
|
||||||
const [rulesRes, tunnelsRes] = await Promise.all([
|
const rulesRes = await getSpeedLimitList();
|
||||||
getSpeedLimitList(),
|
|
||||||
getTunnelList(),
|
|
||||||
]);
|
|
||||||
|
|
||||||
if (rulesRes.code === 0) {
|
if (rulesRes.code === 0) {
|
||||||
setRules(rulesRes.data || []);
|
setRules(rulesRes.data || []);
|
||||||
} else {
|
} else {
|
||||||
toast.error(rulesRes.msg || "获取限速规则失败");
|
toast.error(rulesRes.msg || "获取限速规则失败");
|
||||||
}
|
}
|
||||||
|
|
||||||
if (tunnelsRes.code === 0) {
|
|
||||||
setTunnels(tunnelsRes.data || []);
|
|
||||||
} else {
|
|
||||||
}
|
|
||||||
} catch {
|
} catch {
|
||||||
toast.error("加载数据失败");
|
toast.error("加载数据失败");
|
||||||
} finally {
|
} finally {
|
||||||
@@ -139,8 +115,6 @@ export default function LimitPage() {
|
|||||||
newErrors.speed = "请输入有效的速度限制(≥1 Mbps)";
|
newErrors.speed = "请输入有效的速度限制(≥1 Mbps)";
|
||||||
}
|
}
|
||||||
|
|
||||||
// tunnelId is optional - speed limits can be created without binding to a tunnel
|
|
||||||
|
|
||||||
setErrors(newErrors);
|
setErrors(newErrors);
|
||||||
|
|
||||||
return Object.keys(newErrors).length === 0;
|
return Object.keys(newErrors).length === 0;
|
||||||
@@ -152,8 +126,6 @@ export default function LimitPage() {
|
|||||||
setForm({
|
setForm({
|
||||||
name: "",
|
name: "",
|
||||||
speed: 100,
|
speed: 100,
|
||||||
tunnelId: null,
|
|
||||||
tunnelName: "",
|
|
||||||
status: 1,
|
status: 1,
|
||||||
});
|
});
|
||||||
setErrors({});
|
setErrors({});
|
||||||
@@ -167,8 +139,6 @@ export default function LimitPage() {
|
|||||||
id: rule.id,
|
id: rule.id,
|
||||||
name: rule.name,
|
name: rule.name,
|
||||||
speed: rule.speed,
|
speed: rule.speed,
|
||||||
tunnelId: rule.tunnelId ?? null,
|
|
||||||
tunnelName: rule.tunnelName ?? "",
|
|
||||||
status: rule.status,
|
status: rule.status,
|
||||||
});
|
});
|
||||||
setErrors({});
|
setErrors({});
|
||||||
@@ -210,16 +180,21 @@ export default function LimitPage() {
|
|||||||
setSubmitLoading(true);
|
setSubmitLoading(true);
|
||||||
try {
|
try {
|
||||||
let res: { code: number; msg: string };
|
let res: { code: number; msg: string };
|
||||||
|
const payload = {
|
||||||
|
id: form.id,
|
||||||
|
name: form.name,
|
||||||
|
speed: form.speed,
|
||||||
|
status: form.status,
|
||||||
|
};
|
||||||
|
|
||||||
if (isEdit) {
|
if (isEdit) {
|
||||||
res = await updateSpeedLimit(form);
|
res = await updateSpeedLimit(payload);
|
||||||
} else {
|
} else {
|
||||||
const createData = { ...form };
|
const createData = {
|
||||||
|
name: payload.name,
|
||||||
delete createData.id;
|
speed: payload.speed,
|
||||||
createData.tunnelId = null;
|
status: payload.status,
|
||||||
createData.tunnelName = "";
|
};
|
||||||
|
|
||||||
res = await createSpeedLimit(createData);
|
res = await createSpeedLimit(createData);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -247,7 +222,7 @@ export default function LimitPage() {
|
|||||||
<div className="flex-1 max-w-sm flex items-center gap-2">
|
<div className="flex-1 max-w-sm flex items-center gap-2">
|
||||||
<SearchBar
|
<SearchBar
|
||||||
isVisible={isSearchVisible}
|
isVisible={isSearchVisible}
|
||||||
placeholder="搜索规则名称或绑定隧道"
|
placeholder="搜索规则名称"
|
||||||
value={searchKeyword}
|
value={searchKeyword}
|
||||||
onChange={setSearchKeyword}
|
onChange={setSearchKeyword}
|
||||||
onClose={() => setIsSearchVisible(false)}
|
onClose={() => setIsSearchVisible(false)}
|
||||||
@@ -292,20 +267,6 @@ export default function LimitPage() {
|
|||||||
{rule.speed} Mbps
|
{rule.speed} Mbps
|
||||||
</Chip>
|
</Chip>
|
||||||
</div>
|
</div>
|
||||||
<div className="flex justify-between items-center">
|
|
||||||
<span className="text-small text-default-600">
|
|
||||||
绑定隧道
|
|
||||||
</span>
|
|
||||||
{rule.tunnelName ? (
|
|
||||||
<Chip color="primary" size="sm" variant="flat">
|
|
||||||
{rule.tunnelName}
|
|
||||||
</Chip>
|
|
||||||
) : (
|
|
||||||
<span className="text-default-400 text-small">
|
|
||||||
未绑定
|
|
||||||
</span>
|
|
||||||
)}
|
|
||||||
</div>
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<div className="flex gap-2 mt-4">
|
<div className="flex gap-2 mt-4">
|
||||||
@@ -432,45 +393,6 @@ export default function LimitPage() {
|
|||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
/>
|
/>
|
||||||
|
|
||||||
{isEdit && (
|
|
||||||
<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: "",
|
|
||||||
}));
|
|
||||||
}
|
|
||||||
}}
|
|
||||||
>
|
|
||||||
{tunnels.map((tunnel) => (
|
|
||||||
<SelectItem key={tunnel.id}>{tunnel.name}</SelectItem>
|
|
||||||
))}
|
|
||||||
</Select>
|
|
||||||
)}
|
|
||||||
</div>
|
</div>
|
||||||
</ModalBody>
|
</ModalBody>
|
||||||
<ModalFooter>
|
<ModalFooter>
|
||||||
|
|||||||
@@ -88,7 +88,6 @@ export interface Tunnel {
|
|||||||
export interface SpeedLimit {
|
export interface SpeedLimit {
|
||||||
id: number;
|
id: number;
|
||||||
name: string;
|
name: string;
|
||||||
tunnelId?: number | null;
|
|
||||||
speed?: number;
|
speed?: number;
|
||||||
uploadSpeed: number;
|
uploadSpeed: number;
|
||||||
downloadSpeed: number;
|
downloadSpeed: number;
|
||||||
|
|||||||
Reference in New Issue
Block a user