feat: persist per-IP forward limits

This commit is contained in:
sagitchu
2026-04-27 22:07:04 +08:00
parent 54d7dfb7c9
commit 2b76a9f0be
7 changed files with 167 additions and 52 deletions
@@ -1804,7 +1804,7 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
maxConn := asInt(req["maxConn"], 0) maxConn := asInt(req["maxConn"], 0)
proxyProtocol := asInt(req["proxyProtocol"], 0) proxyProtocol := asInt(req["proxyProtocol"], 0)
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID), maxConn, proxyProtocol) forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID), maxConn, 0, nil, proxyProtocol)
if err != nil { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
@@ -1958,7 +1958,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
maxConn := asInt(req["maxConn"], forward.MaxConn) maxConn := asInt(req["maxConn"], forward.MaxConn)
proxyProtocol := asInt(req["proxyProtocol"], forward.ProxyProtocol) proxyProtocol := asInt(req["proxyProtocol"], forward.ProxyProtocol)
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn, proxyProtocol); err != nil { if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn, 0, nil, proxyProtocol); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
+25 -19
View File
@@ -45,6 +45,8 @@ type Forward struct {
Inx int `gorm:"not null;default:0"` Inx int `gorm:"not null;default:0"`
SpeedID sql.NullInt64 `gorm:"column:speed_id"` SpeedID sql.NullInt64 `gorm:"column:speed_id"`
MaxConn int `gorm:"column:max_conn;not null;default:0"` MaxConn int `gorm:"column:max_conn;not null;default:0"`
IPMaxConn int `gorm:"column:ip_max_conn;not null;default:0"`
IPSpeedID sql.NullInt64 `gorm:"column:ip_speed_id"`
ProxyProtocol int `gorm:"column:proxy_protocol;not null;default:0"` ProxyProtocol int `gorm:"column:proxy_protocol;not null;default:0"`
} }
@@ -428,20 +430,22 @@ type ChainTunnelBackup struct {
} }
type ForwardBackup struct { type ForwardBackup struct {
ID int64 `json:"id"` ID int64 `json:"id"`
UserID int64 `json:"userId"` UserID int64 `json:"userId"`
UserName string `json:"userName"` UserName string `json:"userName"`
Name string `json:"name"` Name string `json:"name"`
TunnelID int64 `json:"tunnelId"` TunnelID int64 `json:"tunnelId"`
RemoteAddr string `json:"remoteAddr"` RemoteAddr string `json:"remoteAddr"`
Strategy string `json:"strategy"` Strategy string `json:"strategy"`
InFlow int64 `json:"inFlow"` InFlow int64 `json:"inFlow"`
OutFlow int64 `json:"outFlow"` OutFlow int64 `json:"outFlow"`
CreatedTime int64 `json:"createdTime"` CreatedTime int64 `json:"createdTime"`
UpdatedTime int64 `json:"updatedTime"` UpdatedTime int64 `json:"updatedTime"`
Status int `json:"status"` Status int `json:"status"`
Inx int `json:"inx"` Inx int `json:"inx"`
SpeedID *int64 `json:"speedId,omitempty"` SpeedID *int64 `json:"speedId,omitempty"`
IPMaxConn int `json:"ipMaxConn,omitempty"`
IPSpeedID *int64 `json:"ipSpeedId,omitempty"`
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"` ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
ProxyProtocol int `json:"proxyProtocol"` ProxyProtocol int `json:"proxyProtocol"`
} }
@@ -532,16 +536,18 @@ type ImportResult struct {
// ForwardRecord is a minimal forward view used by control plane and flow policy. // ForwardRecord is a minimal forward view used by control plane and flow policy.
type ForwardRecord struct { type ForwardRecord struct {
ID int64 ID int64
UserID int64 UserID int64
UserName string UserName string
Name string Name string
TunnelID int64 TunnelID int64
RemoteAddr string RemoteAddr string
Strategy string Strategy string
Status int Status int
SpeedID sql.NullInt64 SpeedID sql.NullInt64
MaxConn int MaxConn int
IPMaxConn int
IPSpeedID sql.NullInt64
ProxyProtocol int ProxyProtocol int
} }
+49 -19
View File
@@ -865,29 +865,33 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
} }
type fwdRow struct { type fwdRow struct {
ID int64 ID int64
UserID int64 UserID int64
UserName string UserName string
Name string Name string
TunnelID int64 TunnelID int64
TunnelName string TunnelName string
TrafficRatio float64 TrafficRatio float64
RemoteAddr string RemoteAddr string
Strategy string Strategy string
InFlow int64 InFlow int64
OutFlow int64 OutFlow int64
CreatedTime int64 CreatedTime int64
Status int Status int
Inx int Inx int
SpeedID sql.NullInt64 SpeedID sql.NullInt64
MaxConn int MaxConn int
ProxyProtocol int IPMaxConn int
IPSpeedID sql.NullInt64
IPSpeedLimitName string
ProxyProtocol int
} }
var rows []fwdRow var rows []fwdRow
err := r.db.Model(&model.Forward{}). err := r.db.Model(&model.Forward{}).
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id, forward.max_conn, forward.proxy_protocol"). Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id, forward.max_conn, forward.ip_max_conn, forward.ip_speed_id, COALESCE(ip_speed_limit.name, '') AS ip_speed_limit_name, forward.proxy_protocol").
Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id"). Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id").
Joins("LEFT JOIN speed_limit AS ip_speed_limit ON ip_speed_limit.id = forward.ip_speed_id").
Order("forward.inx ASC, forward.id ASC"). Order("forward.inx ASC, forward.id ASC").
Find(&rows).Error Find(&rows).Error
if err != nil { if err != nil {
@@ -909,11 +913,18 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
"inFlow": row.InFlow, "outFlow": row.OutFlow, "inFlow": row.InFlow, "outFlow": row.OutFlow,
"createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx), "createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx),
"maxConn": row.MaxConn, "maxConn": row.MaxConn,
"ipMaxConn": row.IPMaxConn,
"proxyProtocol": row.ProxyProtocol, "proxyProtocol": row.ProxyProtocol,
} }
if row.SpeedID.Valid { if row.SpeedID.Valid {
item["speedId"] = row.SpeedID.Int64 item["speedId"] = row.SpeedID.Int64
} }
if row.IPSpeedID.Valid {
item["ipSpeedId"] = row.IPSpeedID.Int64
}
if strings.TrimSpace(row.IPSpeedLimitName) != "" {
item["ipSpeedLimitName"] = row.IPSpeedLimitName
}
items = append(items, item) items = append(items, item)
} }
return items, nil return items, nil
@@ -2102,8 +2113,17 @@ func (r *Repository) exportForwards() ([]model.ForwardBackup, error) {
TunnelID: f.TunnelID, RemoteAddr: f.RemoteAddr, Strategy: f.Strategy, TunnelID: f.TunnelID, RemoteAddr: f.RemoteAddr, Strategy: f.Strategy,
InFlow: f.InFlow, OutFlow: f.OutFlow, CreatedTime: f.CreatedTime, InFlow: f.InFlow, OutFlow: f.OutFlow, CreatedTime: f.CreatedTime,
UpdatedTime: f.UpdatedTime, Status: f.Status, Inx: f.Inx, UpdatedTime: f.UpdatedTime, Status: f.Status, Inx: f.Inx,
IPMaxConn: f.IPMaxConn,
ProxyProtocol: f.ProxyProtocol, ProxyProtocol: f.ProxyProtocol,
} }
if f.SpeedID.Valid {
v := f.SpeedID.Int64
b.SpeedID = &v
}
if f.IPSpeedID.Valid {
v := f.IPSpeedID.Int64
b.IPSpeedID = &v
}
ports, err := r.exportForwardPorts(f.ID) ports, err := r.exportForwardPorts(f.ID)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -2483,6 +2503,13 @@ func importTunnels(tx *gorm.DB, tunnels []model.TunnelBackup, now int64) (int, e
return count, nil return count, nil
} }
func nullableBackupInt64(v *int64) int64 {
if v == nil {
return 0
}
return *v
}
func importForwards(tx *gorm.DB, forwards []model.ForwardBackup, now int64) (int, error) { func importForwards(tx *gorm.DB, forwards []model.ForwardBackup, now int64) (int, error) {
count := 0 count := 0
for _, f := range forwards { for _, f := range forwards {
@@ -2500,13 +2527,16 @@ func importForwards(tx *gorm.DB, forwards []model.ForwardBackup, now int64) (int
UpdatedTime: now, UpdatedTime: now,
Status: f.Status, Status: f.Status,
Inx: f.Inx, Inx: f.Inx,
SpeedID: sql.NullInt64{Int64: nullableBackupInt64(f.SpeedID), Valid: f.SpeedID != nil && *f.SpeedID > 0},
IPMaxConn: f.IPMaxConn,
IPSpeedID: sql.NullInt64{Int64: nullableBackupInt64(f.IPSpeedID), Valid: f.IPSpeedID != nil && *f.IPSpeedID > 0},
ProxyProtocol: f.ProxyProtocol, ProxyProtocol: f.ProxyProtocol,
} }
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{
"user_id", "user_name", "name", "tunnel_id", "remote_addr", "strategy", "user_id", "user_name", "name", "tunnel_id", "remote_addr", "strategy",
"in_flow", "out_flow", "updated_time", "status", "inx", "proxy_protocol", "in_flow", "out_flow", "updated_time", "status", "inx", "speed_id", "ip_max_conn", "ip_speed_id", "proxy_protocol",
}), }),
}).Create(&item).Error }).Create(&item).Error
if err != nil { if err != nil {
@@ -55,6 +55,8 @@ func (r *Repository) ListForwardsByTunnelTx(tx *gorm.DB, tunnelID int64) ([]mode
Status: f.Status, Status: f.Status,
SpeedID: f.SpeedID, SpeedID: f.SpeedID,
MaxConn: f.MaxConn, MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol, ProxyProtocol: f.ProxyProtocol,
}) })
} }
@@ -124,6 +124,8 @@ func (r *Repository) ListActiveForwardsByUser(userID int64) ([]model.ForwardReco
Status: f.Status, Status: f.Status,
SpeedID: f.SpeedID, SpeedID: f.SpeedID,
MaxConn: f.MaxConn, MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol, ProxyProtocol: f.ProxyProtocol,
}) })
} }
@@ -157,6 +159,8 @@ func (r *Repository) ListActiveForwardsByUserTunnel(userID, tunnelID int64) ([]m
Status: f.Status, Status: f.Status,
SpeedID: f.SpeedID, SpeedID: f.SpeedID,
MaxConn: f.MaxConn, MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol, ProxyProtocol: f.ProxyProtocol,
}) })
} }
@@ -190,6 +194,8 @@ func (r *Repository) ListForwardsByUserAndTunnel(userID, tunnelID int64) ([]mode
Status: f.Status, Status: f.Status,
SpeedID: f.SpeedID, SpeedID: f.SpeedID,
MaxConn: f.MaxConn, MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol, ProxyProtocol: f.ProxyProtocol,
}) })
} }
@@ -224,6 +230,8 @@ func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, er
Status: f.Status, Status: f.Status,
SpeedID: f.SpeedID, SpeedID: f.SpeedID,
MaxConn: f.MaxConn, MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol, ProxyProtocol: f.ProxyProtocol,
} }
if strings.TrimSpace(fr.Strategy) == "" { if strings.TrimSpace(fr.Strategy) == "" {
@@ -1,6 +1,7 @@
package repo package repo
import ( import (
"database/sql"
"testing" "testing"
"time" "time"
@@ -151,6 +152,70 @@ func TestListActiveForwardsByUserTunnelIncludesMaxConn(t *testing.T) {
} }
} }
func TestForwardRepositoryPersistsPerIPLimits(t *testing.T) {
r, err := Open(":memory:")
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
forwardID, err := r.CreateForwardTx(1, "admin", "per-ip-forward", 2, "1.1.1.1:443", "fifo", now, 1, []int64{3}, 24000, "", nil, 0, 5, int64(21), 0)
if err != nil {
t.Fatalf("CreateForwardTx: %v", err)
}
record, err := r.GetForwardRecord(forwardID)
if err != nil {
t.Fatalf("GetForwardRecord after create: %v", err)
}
if record.IPMaxConn != 5 {
t.Fatalf("expected created ipMaxConn 5, got %d", record.IPMaxConn)
}
if !record.IPSpeedID.Valid || record.IPSpeedID.Int64 != 21 {
t.Fatalf("expected created ipSpeedId 21, got %+v", record.IPSpeedID)
}
if err := r.UpdateForward(forwardID, "per-ip-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 9, int64(22), 0); err != nil {
t.Fatalf("UpdateForward: %v", err)
}
record, err = r.GetForwardRecord(forwardID)
if err != nil {
t.Fatalf("GetForwardRecord after update: %v", err)
}
if record.IPMaxConn != 9 {
t.Fatalf("expected updated ipMaxConn 9, got %d", record.IPMaxConn)
}
if !record.IPSpeedID.Valid || record.IPSpeedID.Int64 != 22 {
t.Fatalf("expected updated ipSpeedId 22, got %+v", record.IPSpeedID)
}
if err := r.DB().Create(&model.Forward{
UserID: 4,
UserName: "user",
Name: "listed-per-ip-forward",
TunnelID: 8,
RemoteAddr: "3.3.3.3:443",
Strategy: "fifo",
CreatedTime: now,
UpdatedTime: now,
Status: 1,
IPMaxConn: 11,
IPSpeedID: sql.NullInt64{Int64: 33, Valid: true},
}).Error; err != nil {
t.Fatalf("create listed forward: %v", err)
}
records, err := r.ListForwardsByTunnel(8)
if err != nil {
t.Fatalf("ListForwardsByTunnel: %v", err)
}
if len(records) != 1 {
t.Fatalf("expected 1 listed record, got %d", len(records))
}
if records[0].IPMaxConn != 11 || !records[0].IPSpeedID.Valid || records[0].IPSpeedID.Int64 != 33 {
t.Fatalf("expected listed per-IP limits 11/33, got ipMaxConn=%d ipSpeedId=%+v", records[0].IPMaxConn, records[0].IPSpeedID)
}
}
func mustRepoLastInsertID(t *testing.T, r *Repository) int64 { func mustRepoLastInsertID(t *testing.T, r *Repository) int64 {
t.Helper() t.Helper()
var id int64 var id int64
@@ -695,7 +695,7 @@ func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 {
return p return p
} }
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int, proxyProtocol int) error { func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int) 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")
} }
@@ -708,6 +708,8 @@ func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remote
"strategy": strategy, "strategy": strategy,
"speed_id": nullInt64FromInterface(speedID), "speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn, "max_conn": maxConn,
"ip_max_conn": ipMaxConn,
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
"proxy_protocol": proxyProtocol, "proxy_protocol": proxyProtocol,
"updated_time": now, "updated_time": now,
}).Error }).Error
@@ -790,17 +792,17 @@ func (r *Repository) RollbackForwardFields(id, userID int64, userName, name stri
_ = r.db.Model(&model.Forward{}). _ = r.db.Model(&model.Forward{}).
Where("id = ?", id). Where("id = ?", id).
Updates(map[string]interface{}{ Updates(map[string]interface{}{
"user_id": userID, "user_id": userID,
"user_name": userName, "user_name": userName,
"name": name, "name": name,
"tunnel_id": tunnelID, "tunnel_id": tunnelID,
"remote_addr": remoteAddr, "remote_addr": remoteAddr,
"strategy": strategy, "strategy": strategy,
"status": status, "status": status,
"speed_id": nullInt64FromInterface(speedID), "speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn, "max_conn": maxConn,
"proxy_protocol": proxyProtocol, "proxy_protocol": proxyProtocol,
"updated_time": now, "updated_time": now,
}).Error }).Error
} }
@@ -1260,7 +1262,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
return ut.ID, true, nil return ut.ID, true, nil
} }
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int, proxyProtocol int) (int64, error) { func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol 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")
} }
@@ -1281,6 +1283,8 @@ func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnel
Inx: inx, Inx: inx,
MaxConn: maxConn, MaxConn: maxConn,
SpeedID: nullInt64FromInterface(speedID), SpeedID: nullInt64FromInterface(speedID),
IPMaxConn: ipMaxConn,
IPSpeedID: nullInt64FromInterface(ipSpeedID),
ProxyProtocol: proxyProtocol, ProxyProtocol: proxyProtocol,
} }
if err := tx.Create(&fwd).Error; err != nil { if err := tx.Create(&fwd).Error; err != nil {