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)
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 {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
@@ -1958,7 +1958,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
maxConn := asInt(req["maxConn"], forward.MaxConn)
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()))
return
}
+25 -19
View File
@@ -45,6 +45,8 @@ type Forward struct {
Inx int `gorm:"not null;default:0"`
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
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"`
}
@@ -428,20 +430,22 @@ type ChainTunnelBackup struct {
}
type ForwardBackup struct {
ID int64 `json:"id"`
UserID int64 `json:"userId"`
UserName string `json:"userName"`
Name string `json:"name"`
TunnelID int64 `json:"tunnelId"`
RemoteAddr string `json:"remoteAddr"`
Strategy string `json:"strategy"`
InFlow int64 `json:"inFlow"`
OutFlow int64 `json:"outFlow"`
CreatedTime int64 `json:"createdTime"`
UpdatedTime int64 `json:"updatedTime"`
Status int `json:"status"`
ID int64 `json:"id"`
UserID int64 `json:"userId"`
UserName string `json:"userName"`
Name string `json:"name"`
TunnelID int64 `json:"tunnelId"`
RemoteAddr string `json:"remoteAddr"`
Strategy string `json:"strategy"`
InFlow int64 `json:"inFlow"`
OutFlow int64 `json:"outFlow"`
CreatedTime int64 `json:"createdTime"`
UpdatedTime int64 `json:"updatedTime"`
Status int `json:"status"`
Inx int `json:"inx"`
SpeedID *int64 `json:"speedId,omitempty"`
IPMaxConn int `json:"ipMaxConn,omitempty"`
IPSpeedID *int64 `json:"ipSpeedId,omitempty"`
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
ProxyProtocol int `json:"proxyProtocol"`
}
@@ -532,16 +536,18 @@ type ImportResult struct {
// ForwardRecord is a minimal forward view used by control plane and flow policy.
type ForwardRecord struct {
ID int64
UserID int64
UserName string
Name string
TunnelID int64
RemoteAddr string
Strategy string
ID int64
UserID int64
UserName string
Name string
TunnelID int64
RemoteAddr string
Strategy string
Status int
SpeedID sql.NullInt64
MaxConn int
IPMaxConn int
IPSpeedID sql.NullInt64
ProxyProtocol int
}
+49 -19
View File
@@ -865,29 +865,33 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
}
type fwdRow struct {
ID int64
UserID int64
UserName string
Name string
TunnelID int64
TunnelName string
TrafficRatio float64
RemoteAddr string
Strategy string
InFlow int64
OutFlow int64
CreatedTime int64
Status int
Inx int
SpeedID sql.NullInt64
MaxConn int
ProxyProtocol int
ID int64
UserID int64
UserName string
Name string
TunnelID int64
TunnelName string
TrafficRatio float64
RemoteAddr string
Strategy string
InFlow int64
OutFlow int64
CreatedTime int64
Status int
Inx int
SpeedID sql.NullInt64
MaxConn int
IPMaxConn int
IPSpeedID sql.NullInt64
IPSpeedLimitName string
ProxyProtocol int
}
var rows []fwdRow
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 speed_limit AS ip_speed_limit ON ip_speed_limit.id = forward.ip_speed_id").
Order("forward.inx ASC, forward.id ASC").
Find(&rows).Error
if err != nil {
@@ -909,11 +913,18 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
"inFlow": row.InFlow, "outFlow": row.OutFlow,
"createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx),
"maxConn": row.MaxConn,
"ipMaxConn": row.IPMaxConn,
"proxyProtocol": row.ProxyProtocol,
}
if row.SpeedID.Valid {
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)
}
return items, nil
@@ -2102,8 +2113,17 @@ func (r *Repository) exportForwards() ([]model.ForwardBackup, error) {
TunnelID: f.TunnelID, RemoteAddr: f.RemoteAddr, Strategy: f.Strategy,
InFlow: f.InFlow, OutFlow: f.OutFlow, CreatedTime: f.CreatedTime,
UpdatedTime: f.UpdatedTime, Status: f.Status, Inx: f.Inx,
IPMaxConn: f.IPMaxConn,
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)
if err != nil {
return nil, err
@@ -2483,6 +2503,13 @@ func importTunnels(tx *gorm.DB, tunnels []model.TunnelBackup, now int64) (int, e
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) {
count := 0
for _, f := range forwards {
@@ -2500,13 +2527,16 @@ func importForwards(tx *gorm.DB, forwards []model.ForwardBackup, now int64) (int
UpdatedTime: now,
Status: f.Status,
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,
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
"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
if err != nil {
@@ -55,6 +55,8 @@ func (r *Repository) ListForwardsByTunnelTx(tx *gorm.DB, tunnelID int64) ([]mode
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
})
}
@@ -124,6 +124,8 @@ func (r *Repository) ListActiveForwardsByUser(userID int64) ([]model.ForwardReco
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
})
}
@@ -157,6 +159,8 @@ func (r *Repository) ListActiveForwardsByUserTunnel(userID, tunnelID int64) ([]m
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
})
}
@@ -190,6 +194,8 @@ func (r *Repository) ListForwardsByUserAndTunnel(userID, tunnelID int64) ([]mode
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
})
}
@@ -224,6 +230,8 @@ func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, er
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
}
if strings.TrimSpace(fr.Strategy) == "" {
@@ -1,6 +1,7 @@
package repo
import (
"database/sql"
"testing"
"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 {
t.Helper()
var id int64
@@ -695,7 +695,7 @@ func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 {
return p
}
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, 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 {
return errors.New("repository not initialized")
}
@@ -708,6 +708,8 @@ func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remote
"strategy": strategy,
"speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn,
"ip_max_conn": ipMaxConn,
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
"proxy_protocol": proxyProtocol,
"updated_time": now,
}).Error
@@ -790,17 +792,17 @@ func (r *Repository) RollbackForwardFields(id, userID int64, userName, name stri
_ = r.db.Model(&model.Forward{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"user_id": userID,
"user_name": userName,
"name": name,
"tunnel_id": tunnelID,
"remote_addr": remoteAddr,
"strategy": strategy,
"status": status,
"speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn,
"user_id": userID,
"user_name": userName,
"name": name,
"tunnel_id": tunnelID,
"remote_addr": remoteAddr,
"strategy": strategy,
"status": status,
"speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn,
"proxy_protocol": proxyProtocol,
"updated_time": now,
"updated_time": now,
}).Error
}
@@ -1260,7 +1262,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
return ut.ID, true, nil
}
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, 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 {
return 0, errors.New("repository not initialized")
}
@@ -1281,6 +1283,8 @@ func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnel
Inx: inx,
MaxConn: maxConn,
SpeedID: nullInt64FromInterface(speedID),
IPMaxConn: ipMaxConn,
IPSpeedID: nullInt64FromInterface(ipSpeedID),
ProxyProtocol: proxyProtocol,
}
if err := tx.Create(&fwd).Error; err != nil {