From a2000e4d986a0cf984f63f0ccf5fbed5a33c60e5 Mon Sep 17 00:00:00 2001 From: sagitchu Date: Mon, 27 Apr 2026 22:16:38 +0800 Subject: [PATCH] fix: preserve per-IP forward limits during rollback --- .../handler/forward_proxy_protocol_test.go | 23 ++++++++++---- go-backend/internal/http/handler/mutations.go | 4 +-- .../repository_forward_proxy_protocol_test.go | 30 +++++++++++++++++++ .../store/repo/repository_mutations.go | 4 ++- 4 files changed, 53 insertions(+), 8 deletions(-) diff --git a/go-backend/internal/http/handler/forward_proxy_protocol_test.go b/go-backend/internal/http/handler/forward_proxy_protocol_test.go index e5db61c..c5fc9e9 100644 --- a/go-backend/internal/http/handler/forward_proxy_protocol_test.go +++ b/go-backend/internal/http/handler/forward_proxy_protocol_test.go @@ -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) } } diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 2189ce2..09781f6 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -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(), ) diff --git a/go-backend/internal/store/repo/repository_forward_proxy_protocol_test.go b/go-backend/internal/store/repo/repository_forward_proxy_protocol_test.go index c11a095..19e208a 100644 --- a/go-backend/internal/store/repo/repository_forward_proxy_protocol_test.go +++ b/go-backend/internal/store/repo/repository_forward_proxy_protocol_test.go @@ -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 diff --git a/go-backend/internal/store/repo/repository_mutations.go b/go-backend/internal/store/repo/repository_mutations.go index 0f78d09..394ea86 100644 --- a/go-backend/internal/store/repo/repository_mutations.go +++ b/go-backend/internal/store/repo/repository_mutations.go @@ -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