fix: preserve per-IP forward limits during rollback

This commit is contained in:
sagitchu
2026-04-27 22:16:38 +08:00
parent 2b76a9f0be
commit a2000e4d98
4 changed files with 53 additions and 8 deletions
@@ -1,6 +1,7 @@
package handler
import (
"database/sql"
"testing"
"time"
@@ -73,6 +74,8 @@ func TestRollbackForwardMutationRestoresProxyProtocol(t *testing.T) {
CreatedTime: now,
UpdatedTime: now,
Status: 1,
IPMaxConn: 5,
IPSpeedID: sql.NullInt64{Int64: 21, Valid: true},
ProxyProtocol: 2,
}).Error; err != nil {
t.Fatalf("create forward: %v", err)
@@ -81,6 +84,8 @@ func TestRollbackForwardMutationRestoresProxyProtocol(t *testing.T) {
forwardID := mustLastInsertID(t, r, "rollback-forward")
if err := r.DB().Model(&model.Forward{}).Where("id = ?", forwardID).Updates(map[string]interface{}{
"name": "changed-forward",
"ip_max_conn": 0,
"ip_speed_id": nil,
"proxy_protocol": 0,
"updated_time": now + 1,
}).Error; err != nil {
@@ -97,14 +102,22 @@ func TestRollbackForwardMutationRestoresProxyProtocol(t *testing.T) {
RemoteAddr: "9.9.9.9:443",
Strategy: "fifo",
Status: 1,
IPMaxConn: 5,
IPSpeedID: sql.NullInt64{Int64: 21, Valid: true},
ProxyProtocol: 2,
}, nil)
var proxyProtocol int
if err := r.DB().Raw("SELECT proxy_protocol FROM forward WHERE id = ?", forwardID).Row().Scan(&proxyProtocol); err != nil {
t.Fatalf("query proxy_protocol: %v", err)
var record model.Forward
if err := r.DB().Where("id = ?", forwardID).First(&record).Error; err != nil {
t.Fatalf("query forward: %v", err)
}
if proxyProtocol != 2 {
t.Fatalf("expected proxyProtocol restored to 2, got %d", proxyProtocol)
if record.ProxyProtocol != 2 {
t.Fatalf("expected proxyProtocol restored to 2, got %d", record.ProxyProtocol)
}
if record.IPMaxConn != 5 {
t.Fatalf("expected ipMaxConn restored to 5, got %d", record.IPMaxConn)
}
if !record.IPSpeedID.Valid || record.IPSpeedID.Int64 != 21 {
t.Fatalf("expected ipSpeedId restored to 21, got %+v", record.IPSpeedID)
}
}
@@ -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, 0, nil, proxyProtocol); err != nil {
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn, forward.IPMaxConn, forward.IPSpeedID, proxyProtocol); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
@@ -4110,7 +4110,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []
h.repo.RollbackForwardFields(
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
oldForward.SpeedID, oldForward.MaxConn, oldForward.ProxyProtocol,
oldForward.SpeedID, oldForward.MaxConn, oldForward.IPMaxConn, oldForward.IPSpeedID, oldForward.ProxyProtocol,
time.Now().UnixMilli(),
)
@@ -216,6 +216,36 @@ func TestForwardRepositoryPersistsPerIPLimits(t *testing.T) {
}
}
func TestRollbackForwardFieldsRestoresPerIPLimits(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", "rollback-per-ip-forward", 2, "1.1.1.1:443", "fifo", now, 1, nil, 0, "", nil, 7, 5, int64(21), 2)
if err != nil {
t.Fatalf("CreateForwardTx: %v", err)
}
if err := r.UpdateForward(forwardID, "rollback-per-ip-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 0, nil, 0); err != nil {
t.Fatalf("UpdateForward: %v", err)
}
r.RollbackForwardFields(forwardID, 1, "admin", "rollback-per-ip-forward", 2, "1.1.1.1:443", "fifo", 1, nil, 7, 5, int64(21), 2, now+2)
record, err := r.GetForwardRecord(forwardID)
if err != nil {
t.Fatalf("GetForwardRecord: %v", err)
}
if record.IPMaxConn != 5 {
t.Fatalf("expected rollback ipMaxConn 5, got %d", record.IPMaxConn)
}
if !record.IPSpeedID.Valid || record.IPSpeedID.Int64 != 21 {
t.Fatalf("expected rollback ipSpeedId 21, got %+v", record.IPSpeedID)
}
}
func mustRepoLastInsertID(t *testing.T, r *Repository) int64 {
t.Helper()
var id int64
@@ -785,7 +785,7 @@ func (r *Repository) UpdateForwardPortBindIP(forwardID, nodeID int64, port int,
Update("in_ip", sql.NullString{String: inIP, Valid: strings.TrimSpace(inIP) != ""}).Error
}
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, proxyProtocol int, now int64) {
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int, now int64) {
if r == nil || r.db == nil {
return
}
@@ -801,6 +801,8 @@ func (r *Repository) RollbackForwardFields(id, userID int64, userName, name stri
"status": status,
"speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn,
"ip_max_conn": ipMaxConn,
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
"proxy_protocol": proxyProtocol,
"updated_time": now,
}).Error