mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-08 02:36: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 {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user