From c952d2fb3a8257f02627a9513a7adadc291f7a55 Mon Sep 17 00:00:00 2001 From: sagit <36596628+Sagit-chu@users.noreply.github.com> Date: Tue, 29 Sep 2026 14:05:49 +0800 Subject: [PATCH] fix: protect shared rules across delivery and recovery (#560) --- .../internal/http/handler/federation.go | 327 ++++---- .../federation_consumer_cleanup_test.go | 275 +++++++ .../http/handler/federation_resources.go | 703 ++++++++++++++++++ .../http/handler/federation_resources_test.go | 538 ++++++++++++++ .../http/handler/federation_share_test.go | 16 +- .../handler/flow_cleanup_regression_test.go | 300 ++++++++ .../http/handler/flow_config_cleanup_test.go | 145 ++++ .../internal/http/handler/flow_ownership.go | 118 +++ .../http/handler/flow_ownership_test.go | 235 ++++++ .../handler/flow_pending_deletion_test.go | 57 ++ .../internal/http/handler/flow_policy.go | 397 +++++----- .../handler/flow_policy_federation_test.go | 12 +- .../http/handler/flow_upload_batch.go | 17 +- .../http/handler/flow_upload_batch_test.go | 2 + go-backend/internal/http/handler/handler.go | 20 +- go-backend/internal/http/handler/jobs.go | 45 +- go-backend/internal/http/handler/mutations.go | 421 ++++++++--- .../handler/peer_share_auth_cleanup_test.go | 160 ++++ .../handler/peer_share_runtime_lifecycle.go | 190 +++++ .../peer_share_runtime_lifecycle_test.go | 356 +++++++++ go-backend/internal/http/handler/upgrade.go | 100 ++- .../internal/http/handler/upgrade_test.go | 15 +- .../store/model/federation_release.go | 15 + go-backend/internal/store/model/model.go | 55 +- .../store/repo/peer_share_resources.go | 55 ++ .../repo/peer_share_runtime_lifecycle.go | 32 + go-backend/internal/store/repo/repository.go | 41 +- .../repo/repository_federation_cleanup.go | 55 ++ .../repo/repository_federation_reconcile.go | 35 + .../federation_dual_panel_contract_test.go | 4 +- .../flow_upload_batch_contract_test.go | 3 + .../tunnel_metrics_ingestion_contract_test.go | 3 + go-gost/main.go | 8 +- go-gost/program.go | 123 ++- go-gost/tests/lifecycle/lifecycle_test.go | 313 ++++++++ go-gost/x/api/api.go | 2 +- go-gost/x/api/config_reload.go | 16 +- go-gost/x/api/config_transaction.go | 90 +++ go-gost/x/api/config_transaction_test.go | 76 ++ go-gost/x/api/service/service.go | 13 +- go-gost/x/config/config.go | 14 +- go-gost/x/config/loader/loader.go | 37 + go-gost/x/config/loader/reload_test.go | 84 +++ go-gost/x/config/parsing/service/parse.go | 10 + go-gost/x/metrics/service/service.go | 4 +- go-gost/x/service/traffic_reporter.go | 60 +- go-gost/x/service/traffic_reporter_test.go | 32 + go-gost/x/socket/command_dispatch_test.go | 36 + go-gost/x/socket/websocket_reporter.go | 31 +- 49 files changed, 5053 insertions(+), 643 deletions(-) create mode 100644 go-backend/internal/http/handler/federation_consumer_cleanup_test.go create mode 100644 go-backend/internal/http/handler/federation_resources.go create mode 100644 go-backend/internal/http/handler/federation_resources_test.go create mode 100644 go-backend/internal/http/handler/flow_cleanup_regression_test.go create mode 100644 go-backend/internal/http/handler/flow_config_cleanup_test.go create mode 100644 go-backend/internal/http/handler/flow_ownership.go create mode 100644 go-backend/internal/http/handler/flow_ownership_test.go create mode 100644 go-backend/internal/http/handler/flow_pending_deletion_test.go create mode 100644 go-backend/internal/http/handler/peer_share_auth_cleanup_test.go create mode 100644 go-backend/internal/http/handler/peer_share_runtime_lifecycle.go create mode 100644 go-backend/internal/http/handler/peer_share_runtime_lifecycle_test.go create mode 100644 go-backend/internal/store/model/federation_release.go create mode 100644 go-backend/internal/store/repo/peer_share_resources.go create mode 100644 go-backend/internal/store/repo/peer_share_runtime_lifecycle.go create mode 100644 go-backend/internal/store/repo/repository_federation_cleanup.go create mode 100644 go-backend/internal/store/repo/repository_federation_reconcile.go create mode 100644 go-gost/tests/lifecycle/lifecycle_test.go create mode 100644 go-gost/x/api/config_transaction.go create mode 100644 go-gost/x/api/config_transaction_test.go create mode 100644 go-gost/x/config/loader/reload_test.go diff --git a/go-backend/internal/http/handler/federation.go b/go-backend/internal/http/handler/federation.go index 51e1171..08c7fcf 100644 --- a/go-backend/internal/http/handler/federation.go +++ b/go-backend/internal/http/handler/federation.go @@ -1,8 +1,11 @@ package handler import ( + "bytes" "encoding/json" + "errors" "fmt" + "io" "net" "net/http" "net/url" @@ -408,9 +411,26 @@ func (h *Handler) federationShareDelete(w http.ResponseWriter, r *http.Request) return } - share, _ := h.repo.GetPeerShare(req.ID) + share, err := h.repo.GetPeerShare(req.ID) + if err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + if share != nil { + // Revoke new allocations before cleanup; failed deletions keep the + // disabled share and its reservations available for retry. + share.IsActive = 0 + share.UpdatedTime = time.Now().UnixMilli() + if err := h.repo.UpdatePeerShare(share); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + } - h.cleanupPeerShareRuntimes(req.ID) + if err := h.cleanupPeerShareRuntimes(req.ID); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } h.cleanupFederationTunnels(req.ID) if err := h.repo.DeletePeerShare(req.ID); err != nil { @@ -418,10 +438,6 @@ func (h *Handler) federationShareDelete(w http.ResponseWriter, r *http.Request) return } - if share != nil && h.wsServer != nil { - h.wsServer.SendCommand(share.NodeID, "reload", nil, time.Second*5) - } - response.WriteJSON(w, response.OKEmpty()) } @@ -805,16 +821,6 @@ func (h *Handler) authPeer(next http.HandlerFunc) http.HandlerFunc { return } - if share.IsActive == 0 { - response.WriteJSON(w, response.Err(403, "Share is disabled")) - return - } - - if share.ExpiryTime > 0 && share.ExpiryTime < time.Now().UnixMilli() { - response.WriteJSON(w, response.Err(403, "Share expired")) - return - } - if strings.TrimSpace(share.AllowedIPs) != "" { clientIP := resolvePeerClientIP(r) if clientIP == nil { @@ -847,10 +853,48 @@ func (h *Handler) authPeer(next http.HandlerFunc) http.HandlerFunc { } } + if share.IsActive != 1 || (share.ExpiryTime > 0 && share.ExpiryTime <= time.Now().UnixMilli()) || isPeerShareFlowExceeded(share) { + if !isFederationCleanupRequest(r) { + response.WriteJSON(w, response.Err(403, "Share is inactive, expired, or over quota")) + return + } + } + next(w, r) } } +// Revoked allocation privileges must not revoke cleanup privileges. Inspect +// only supported destructive commands, preserving the body for the handler. +func isFederationCleanupRequest(r *http.Request) bool { + if r.Method != http.MethodPost { + return false + } + switch r.URL.Path { + case "/api/v1/federation/runtime/release-role": + return true + case "/api/v1/federation/runtime/command": + if r.Body == nil { + return false + } + var copied bytes.Buffer + original := r.Body + var request federationRuntimeCommandRequest + err := json.NewDecoder(io.TeeReader(io.LimitReader(original, 1<<20), &copied)).Decode(&request) + r.Body = struct { + io.Reader + io.Closer + }{Reader: io.MultiReader(&copied, original), Closer: original} + if err != nil || !isFederationRuntimeCommandAllowed(request.CommandType) { + return false + } + _, action := federationResourceCommandKind(request.CommandType) + return action == "delete" + default: + return false + } +} + func (h *Handler) federationConnect(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { response.WriteJSON(w, response.ErrDefault("Invalid method")) @@ -982,8 +1026,6 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request) return } - h.wsServer.SendCommand(share.NodeID, "reload", nil, time.Second*5) - response.WriteJSON(w, response.OK(map[string]interface{}{ "tunnelId": tunnelID, })) @@ -1013,11 +1055,28 @@ func (h *Handler) federationRuntimeReservePort(w http.ResponseWriter, r *http.Re return } + peerRoleRuntimeMu.Lock() + defer peerRoleRuntimeMu.Unlock() + share, err = h.repo.GetPeerShare(share.ID) + if err != nil || share == nil || share.IsActive != 1 || (share.ExpiryTime > 0 && share.ExpiryTime <= time.Now().UnixMilli()) { + response.WriteJSON(w, response.Err(403, "Share is unavailable")) + return + } + + if isPeerShareFlowExceeded(share) { + response.WriteJSON(w, response.Err(403, "Share traffic limit exceeded")) + return + } + existing, err := h.repo.GetPeerShareRuntimeByResourceKey(share.ID, req.ResourceKey) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } + if existing != nil && existing.ReleasePending != 0 { + response.WriteJSON(w, response.ErrDefault("Runtime release is pending")) + return + } if existing != nil && existing.Status == 1 { response.WriteJSON(w, response.OK(map[string]interface{}{ "reservationId": existing.ReservationID, @@ -1026,10 +1085,6 @@ func (h *Handler) federationRuntimeReservePort(w http.ResponseWriter, r *http.Re })) return } - if isPeerShareFlowExceeded(share) { - response.WriteJSON(w, response.Err(403, "Share traffic limit exceeded")) - return - } allocatedPort, err := h.pickPeerSharePort(share, req.RequestedPort) if err != nil { @@ -1039,6 +1094,7 @@ func (h *Handler) federationRuntimeReservePort(w http.ResponseWriter, r *http.Re now := time.Now().UnixMilli() if existing != nil { + existing.ReservationID = randomToken(24) existing.Protocol = defaultString(req.Protocol, "tls") existing.Port = allocatedPort existing.BindingID = "" @@ -1116,6 +1172,14 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ return } + peerRoleRuntimeMu.Lock() + defer peerRoleRuntimeMu.Unlock() + share, err = h.repo.GetPeerShare(share.ID) + if err != nil || share == nil || share.IsActive != 1 || (share.ExpiryTime > 0 && share.ExpiryTime <= time.Now().UnixMilli()) { + response.WriteJSON(w, response.ErrDefault("Share is unavailable")) + return + } + var runtime *repo.PeerShareRuntime if strings.TrimSpace(req.ReservationID) != "" { runtime, err = h.repo.GetPeerShareRuntimeByReservationID(share.ID, strings.TrimSpace(req.ReservationID)) @@ -1131,108 +1195,37 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ return } - node, err := h.getNodeRecord(share.NodeID) - if err != nil { - response.WriteJSON(w, response.ErrDefault(err.Error())) - return - } - - protocol := defaultString(req.Protocol, runtime.Protocol) - strategy := defaultString(req.Strategy, "round") - chainName := defaultString(runtime.ChainName, federationRuntimeChainName(runtime.BindingID)) - if chainName == "" { - chainName = federationRuntimeChainName(fmt.Sprintf("%d", runtime.ID)) - } - serviceName := fmt.Sprintf("fed_svc_%d", runtime.ID) - if runtime.Applied == 1 && strings.TrimSpace(runtime.BindingID) != "" { - if req.Role == "middle" && len(req.Targets) > 0 { - chainData, buildErr := buildFederationMiddleChainConfig(chainName, runtime.ID, protocol, strategy, req.Targets, node.InterfaceName) - if buildErr != nil { - response.WriteJSON(w, response.ErrDefault(buildErr.Error())) - return - } - if _, err := h.sendNodeCommand(share.NodeID, "UpdateChains", updateChainPayload(chainName, chainData), false, false); err != nil { - response.WriteJSON(w, response.ErrDefault(err.Error())) - return - } - targetBytes, _ := json.Marshal(req.Targets) - runtime.Role = req.Role - runtime.ChainName = chainName - runtime.Protocol = protocol - runtime.Strategy = strategy - runtime.Target = string(targetBytes) - runtime.Status = 1 - runtime.UpdatedTime = time.Now().UnixMilli() - if err := h.repo.UpdatePeerShareRuntime(runtime); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - } - response.WriteJSON(w, response.OK(map[string]interface{}{ - "bindingId": runtime.BindingID, - "allocatedPort": runtime.Port, - "reservationId": runtime.ReservationID, - })) + if runtime.ReleasePending != 0 { + response.WriteJSON(w, response.ErrDefault("Runtime release is pending")) return } if isPeerShareFlowExceeded(share) { response.WriteJSON(w, response.Err(403, "Share traffic limit exceeded")) return } - - if share.PortRangeStart > 0 && share.PortRangeEnd > 0 && runtime.Port > 0 { - if runtime.Port < share.PortRangeStart || runtime.Port > share.PortRangeEnd { - response.WriteJSON(w, response.Err(403, fmt.Sprintf("port %d out of allowed range %d-%d", runtime.Port, share.PortRangeStart, share.PortRangeEnd))) - return - } - } - - if req.Role == "middle" { - chainData, buildErr := buildFederationMiddleChainConfig(chainName, runtime.ID, protocol, strategy, req.Targets, node.InterfaceName) - if buildErr != nil { - response.WriteJSON(w, response.ErrDefault(buildErr.Error())) - return - } - if _, err := h.sendNodeCommand(share.NodeID, "AddChains", chainData, true, false); err != nil { - response.WriteJSON(w, response.ErrDefault(err.Error())) - return - } - } - - targetCount := len(req.Targets) - service := buildFederationServiceConfig( - serviceName, - fmt.Sprintf("%s:%d", node.TCPListenAddr, runtime.Port), - protocol, - req.Role, - chainName, - targetCount, - node.InterfaceName, - ) - if _, err := h.sendNodeCommand(share.NodeID, "AddService", []map[string]interface{}{service}, true, false); err != nil { - if req.Role == "middle" { - _, _ = h.sendNodeCommand(share.NodeID, "DeleteChains", map[string]interface{}{"chain": chainName}, false, true) - } - response.WriteJSON(w, response.ErrDefault(err.Error())) + if runtime.Role != "" && runtime.Role != req.Role { + response.WriteJSON(w, response.ErrDefault("Runtime role cannot change without release")) return } - - targetBytes, _ := json.Marshal(req.Targets) - runtime.BindingID = fmt.Sprintf("%d", runtime.ID) + if share.PortRangeStart > 0 && share.PortRangeEnd > 0 && (runtime.Port < share.PortRangeStart || runtime.Port > share.PortRangeEnd) { + response.WriteJSON(w, response.ErrDefault("Reserved port is outside share range")) + return + } + if strings.TrimSpace(runtime.BindingID) == "" { + runtime.BindingID = randomToken(24) + } runtime.Role = req.Role + runtime.ServiceName = fmt.Sprintf("fed_svc_%d", runtime.ID) runtime.ChainName = "" if req.Role == "middle" { - runtime.ChainName = chainName + runtime.ChainName = federationRuntimeChainName(runtime.BindingID) } - runtime.ServiceName = serviceName - runtime.Protocol = protocol - runtime.Strategy = strategy + runtime.Protocol = defaultString(req.Protocol, runtime.Protocol) + runtime.Strategy = defaultString(req.Strategy, "round") + targetBytes, _ := json.Marshal(req.Targets) runtime.Target = string(targetBytes) - runtime.Applied = 1 - runtime.Status = 1 - runtime.UpdatedTime = time.Now().UnixMilli() - if err := h.repo.UpdatePeerShareRuntime(runtime); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) + if err := h.applyPeerShareRoleRuntime(runtime); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) return } @@ -1282,17 +1275,8 @@ func (h *Handler) federationRuntimeReleaseRole(w http.ResponseWriter, r *http.Re return } - if runtime.Applied == 1 { - if strings.TrimSpace(runtime.ServiceName) != "" { - _, _ = h.sendNodeCommand(share.NodeID, "DeleteService", map[string]interface{}{"services": []string{runtime.ServiceName}}, false, true) - } - if strings.TrimSpace(runtime.Role) == "middle" && strings.TrimSpace(runtime.ChainName) != "" { - _, _ = h.sendNodeCommand(share.NodeID, "DeleteChains", map[string]interface{}{"chain": runtime.ChainName}, false, true) - } - } - - if err := h.repo.MarkPeerShareRuntimeReleased(runtime.ID, time.Now().UnixMilli()); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) + if err := h.releasePeerShareRuntime(runtime); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) return } @@ -1385,6 +1369,15 @@ func (h *Handler) federationRuntimeCommand(w http.ResponseWriter, r *http.Reques return } + h.peerResourceMu.Lock() + defer h.peerResourceMu.Unlock() + // Recheck after acquiring the mutation lock: an earlier authentication + // decision cannot authorize a recreation after quota/expiry cleanup. + share, err = h.repo.GetPeerShare(share.ID) + if err != nil || share == nil { + response.WriteJSON(w, response.ErrDefault("share ownership unavailable")) + return + } if isFederationServiceCommand(cmd) { if err := validateFederationCommandPorts(share, req.Data); err != nil { response.WriteJSON(w, response.Err(403, err.Error())) @@ -1392,17 +1385,35 @@ func (h *Handler) federationRuntimeCommand(w http.ResponseWriter, r *http.Reques } } - res, err := h.sendNodeCommand(share.NodeID, cmd, req.Data, false, false) + _, action := federationResourceCommandKind(cmd) + if action != "delete" && (share.IsActive != 1 || (share.ExpiryTime > 0 && share.ExpiryTime <= time.Now().UnixMilli()) || isPeerShareFlowExceeded(share)) { + response.WriteJSON(w, response.Err(403, "share is inactive, expired, or over quota")) + return + } + if strings.EqualFold(cmd, "tcpping") { + res, err := h.sendNodeCommand(share.NodeID, "TcpPing", req.Data, false, false) + if err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + response.WriteJSON(w, response.OK(res)) + return + } + items, err := h.preparePeerResourceCommand(share, cmd, req.Data) if err != nil { response.WriteJSON(w, response.ErrDefault(err.Error())) return } - if strings.EqualFold(cmd, "addservice") || strings.EqualFold(cmd, "updateservice") { - h.bindPeerShareForwardRuntimeServices(share, req.Data) - } else if strings.EqualFold(cmd, "deleteservice") { - h.releasePeerShareForwardRuntimeServices(share, req.Data) + var result interface{} = map[string]interface{}{"success": true} + for _, item := range items { + res, err := h.applyPeerShareResource(item) + if err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + result = res } - response.WriteJSON(w, response.OK(res)) + response.WriteJSON(w, response.OK(result)) } type federationForwardServiceBinding struct { @@ -1436,10 +1447,15 @@ func parseFederationForwardServiceBindings(data interface{}) []federationForward bindings := make([]federationForwardServiceBinding, 0, len(serviceList)) for _, svcMap := range serviceList { name := normalizeForwardRuntimeServiceName(asString(svcMap["name"])) + originalName := name + if shareID, original, ok := parsePeerShareServiceName(asString(svcMap["name"])); ok { + originalName = normalizeForwardRuntimeServiceName(original) + name = peerShareResourceName(shareID, "service", originalName) + } if name == "" { continue } - if _, _, _, ok := parseFlowServiceIDs(name); !ok { + if _, _, _, ok := parseFlowServiceIDs(originalName); !ok { continue } addr := strings.TrimSpace(asString(svcMap["addr"])) @@ -1498,29 +1514,29 @@ func parseFederationForwardServiceNamesForRelease(data interface{}) []string { return out } -func (h *Handler) bindPeerShareForwardRuntimeServices(share *repo.PeerShare, data interface{}) { +func (h *Handler) bindPeerShareForwardRuntimeServices(share *repo.PeerShare, data interface{}) error { if h == nil || h.repo == nil || share == nil { - return + return nil } bindings := parseFederationForwardServiceBindings(data) if len(bindings) == 0 { - return + return nil } now := time.Now().UnixMilli() for _, binding := range bindings { runtime, err := h.repo.GetActiveForwardPeerShareRuntimeByPort(share.ID, binding.Port) if err != nil { - continue + return err } if runtime == nil { runtime, err = h.repo.GetActiveForwardPeerShareRuntimeByServiceName(share.ID, binding.Name) if err != nil { - continue + return err } } if runtime == nil { - _ = h.repo.CreatePeerShareRuntime(&repo.PeerShareRuntime{ + if err := h.repo.CreatePeerShareRuntime(&repo.PeerShareRuntime{ ShareID: share.ID, NodeID: share.NodeID, ReservationID: randomToken(24), @@ -1537,9 +1553,14 @@ func (h *Handler) bindPeerShareForwardRuntimeServices(share *repo.PeerShare, dat Status: 1, CreatedTime: now, UpdatedTime: now, - }) + }); err != nil { + return err + } continue } + if runtime.ReleasePending != 0 { + return fmt.Errorf("runtime release is pending") + } if runtime.ServiceName == binding.Name && runtime.Applied == 1 && runtime.Port == binding.Port && runtime.Status == 1 { continue } @@ -1554,8 +1575,11 @@ func (h *Handler) bindPeerShareForwardRuntimeServices(share *repo.PeerShare, dat if strings.TrimSpace(runtime.Strategy) == "" { runtime.Strategy = "fifo" } - _ = h.repo.UpdatePeerShareRuntime(runtime) + if err := h.repo.UpdatePeerShareRuntime(runtime); err != nil { + return err + } } + return nil } func (h *Handler) releasePeerShareForwardRuntimeServices(share *repo.PeerShare, data interface{}) { @@ -1575,7 +1599,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", "updatelimiters", "deletelimiters", "tcpping", "reload": + case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice", "addchains", "updatechains", "deletechains", "addlimiters", "updatelimiters", "deletelimiters", "addclimiters", "updateclimiters", "deleteclimiters", "tcpping": return true default: return false @@ -1893,27 +1917,24 @@ func (h *Handler) syncRemoteNodeStatuses(items []map[string]interface{}) { } } -func (h *Handler) cleanupPeerShareRuntimes(shareID int64) { +func (h *Handler) cleanupPeerShareRuntimes(shareID int64) error { if h == nil || h.repo == nil || shareID <= 0 { - return + return nil + } + if err := h.releasePeerShareResources(shareID); err != nil { + return err } runtimes, err := h.repo.ListActivePeerShareRuntimesByShareID(shareID) - if err != nil || len(runtimes) == 0 { - return + if err != nil { + return err } - - now := time.Now().UnixMilli() - for _, runtime := range runtimes { - if h.wsServer != nil && runtime.Applied == 1 { - if strings.TrimSpace(runtime.ServiceName) != "" { - _, _ = h.sendNodeCommand(runtime.NodeID, "DeleteService", map[string]interface{}{"services": []string{runtime.ServiceName}}, false, true) - } - if strings.TrimSpace(runtime.Role) == "middle" && strings.TrimSpace(runtime.ChainName) != "" { - _, _ = h.sendNodeCommand(runtime.NodeID, "DeleteChains", map[string]interface{}{"chain": runtime.ChainName}, false, true) - } + var cleanupErr error + for i := range runtimes { + if err := h.releasePeerShareRuntime(&runtimes[i]); err != nil { + cleanupErr = errors.Join(cleanupErr, err) } - _ = h.repo.MarkPeerShareRuntimeReleased(runtime.ID, now) } + return cleanupErr } func (h *Handler) cleanupFederationTunnels(shareID int64) { diff --git a/go-backend/internal/http/handler/federation_consumer_cleanup_test.go b/go-backend/internal/http/handler/federation_consumer_cleanup_test.go new file mode 100644 index 0000000..343a10a --- /dev/null +++ b/go-backend/internal/http/handler/federation_consumer_cleanup_test.go @@ -0,0 +1,275 @@ +package handler + +import ( + "bytes" + "database/sql" + "encoding/json" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "sync/atomic" + "testing" + + "go-backend/internal/http/response" + "go-backend/internal/store/model" + "go-backend/internal/store/repo" + "go-backend/internal/ws" +) + +func consumerCleanupFixture(t *testing.T, remoteURL string) (*repo.Repository, *Handler) { + t.Helper() + r, err := repo.Open(filepath.Join(t.TempDir(), "consumer-cleanup.db")) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = r.Close() }) + if err := r.DB().Create(&model.Node{ID: 1, Name: "remote", Secret: "remote", ServerIP: "127.0.0.1", Port: "31000-31010", IsRemote: 1, Status: 1, RemoteURL: sql.NullString{String: remoteURL, Valid: true}, RemoteToken: sql.NullString{String: "token", Valid: true}}).Error; err != nil { + t.Fatal(err) + } + if err := r.UpsertFederationTunnelBinding(&repo.FederationTunnelBinding{TunnelID: 42, NodeID: 1, ChainType: 3, RemoteURL: remoteURL, ResourceKey: "tunnel:42:node:1:type:3:hop:0", RemoteBindingID: "binding-old", AllocatedPort: 31001, Status: 1}); err != nil { + t.Fatal(err) + } + if err := r.DB().Create(&model.Tunnel{ID: 42, Name: "live tunnel", Type: 2, Protocol: "tls", Flow: 1, TrafficRatio: 1, Status: 1}).Error; err != nil { + t.Fatal(err) + } + return r, &Handler{repo: r, wsServer: ws.NewServer(r, "consumer-cleanup")} +} + +func TestTunnelUpdateValidationPreservesSharedRuntime(t *testing.T) { + for _, body := range []string{ + `{"id":42,"type":2,"name":"invalid update","inNodeId":[]}`, + `{"id":42,"type":2,"name":"invalid update","inNodeId":[{"nodeId":999}],"outNodeId":[{"nodeId":1,"port":31001}]}`, + } { + t.Run(body, func(t *testing.T) { + var releases atomic.Int64 + remote := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + releases.Add(1) + response.WriteJSON(w, response.OKEmpty()) + })) + defer remote.Close() + r, h := consumerCleanupFixture(t, remote.URL) + req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", bytes.NewBufferString(body)) + rec := httptest.NewRecorder() + h.tunnelUpdate(rec, req) + var payload response.R + if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil { + t.Fatal(err) + } + bindings, err := r.ListActiveFederationTunnelBindingsByTunnel(42) + if err != nil { + t.Fatal(err) + } + name, err := r.GetTunnelName(42) + if err != nil { + t.Fatal(err) + } + if payload.Code == 0 || releases.Load() != 0 || len(bindings) != 1 || name != "live tunnel" { + t.Fatalf("invalid update changed runtime: code=%d releases=%d bindings=%v name=%s", payload.Code, releases.Load(), bindings, name) + } + }) + } +} + +func TestFederationCleanupRetainsFailedBindingAndRetriesOnlyPending(t *testing.T) { + var unavailable atomic.Bool + unavailable.Store(true) + var releases atomic.Int64 + remote := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + releases.Add(1) + if unavailable.Load() { + http.Error(w, "peer unavailable", http.StatusServiceUnavailable) + return + } + response.WriteJSON(w, response.OKEmpty()) + })) + defer remote.Close() + r, h := consumerCleanupFixture(t, remote.URL) + if err := r.UpsertFederationTunnelBinding(&repo.FederationTunnelBinding{TunnelID: 43, NodeID: 1, ChainType: 3, RemoteURL: remote.URL, RemoteBindingID: "unrelated", ResourceKey: "unrelated", Status: 1}); err != nil { + t.Fatal(err) + } + if err := h.cleanupFederationRuntime(42); err == nil { + t.Fatal("expected release failure") + } + pending, err := r.ListPendingFederationTunnelBindings() + if err != nil || len(pending) != 1 { + t.Fatalf("lost pending cleanup: %v %v", pending, err) + } + if err := h.cleanupFederationRuntime(42); err == nil { + t.Fatal("expected second release failure") + } + if releases.Load() != 2 { + t.Fatalf("cleanup was not retried: %d", releases.Load()) + } + unavailable.Store(false) + if err := h.retryPendingFederationRuntimeCleanup(); err != nil { + t.Fatal(err) + } + pending, err = r.ListPendingFederationTunnelBindings() + if err != nil || len(pending) != 0 { + t.Fatalf("completed cleanup remains: %v %v", pending, err) + } + active, err := r.ListActiveFederationTunnelBindingsByTunnel(43) + if err != nil || len(active) != 1 || releases.Load() != 3 { + t.Fatalf("retry touched active binding: %v %v calls=%d", active, err, releases.Load()) + } +} + +func TestTunnelDeleteReportsRemoteFailureAndKeepsTunnel(t *testing.T) { + remote := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + http.Error(w, "peer unavailable", http.StatusServiceUnavailable) + })) + defer remote.Close() + r, h := consumerCleanupFixture(t, remote.URL) + rec := httptest.NewRecorder() + h.tunnelDelete(rec, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/delete", strings.NewReader(`{"id":42}`))) + var payload response.R + if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil { + t.Fatal(err) + } + if payload.Code == 0 { + t.Fatal("delete reported success before peer cleanup") + } + if name, err := r.GetTunnelName(42); err != nil || name != "live tunnel" { + t.Fatalf("lost tunnel: %s %v", name, err) + } + pending, err := r.ListPendingFederationTunnelBindings() + if err != nil || len(pending) != 1 { + t.Fatalf("lost binding: %v %v", pending, err) + } +} + +func TestFederationRollbackReleaseIsDurableAndRetryable(t *testing.T) { + var unavailable atomic.Bool + unavailable.Store(true) + remote := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + if unavailable.Load() { + http.Error(w, "peer unavailable", http.StatusServiceUnavailable) + return + } + response.WriteJSON(w, response.OKEmpty()) + })) + defer remote.Close() + r, h := consumerCleanupFixture(t, remote.URL) + refs := []federationRuntimeReleaseRef{{RemoteURL: remote.URL, RemoteToken: "token", BindingID: "new-binding", ReservationID: "new-reservation", ResourceKey: "new-key"}} + if err := h.releaseFederationRuntimeRefs(refs); err == nil { + t.Fatal("expected release error") + } + if err := h.releaseFederationRuntimeRefs(refs); err == nil { + t.Fatal("expected repeat error") + } + pending, err := r.ListPendingFederationReleases() + if err != nil || len(pending) != 1 { + t.Fatalf("rollback queue lost or duplicated release: %v %v", pending, err) + } + unavailable.Store(false) + // A new Handler has no in-memory state from the failed rollback. + restarted := &Handler{repo: r} + if err := restarted.retryPendingFederationRuntimeCleanup(); err != nil { + t.Fatal(err) + } + pending, err = r.ListPendingFederationReleases() + if err != nil || len(pending) != 0 { + t.Fatalf("completed rollback remains queued: %v %v", pending, err) + } + active, err := r.ListActiveFederationTunnelBindingsByTunnel(42) + if err != nil || len(active) != 1 { + t.Fatalf("rollback retry removed active binding: %v %v", active, err) + } +} + +func TestFederationApplyFailureReturnsReservationForRollback(t *testing.T) { + remote := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + if strings.HasSuffix(req.URL.Path, "reserve-port") { + response.WriteJSON(w, response.OK(map[string]interface{}{"reservationId": "reserved", "bindingId": "binding", "allocatedPort": 31001})) + return + } + http.Error(w, "apply failed", http.StatusServiceUnavailable) + })) + defer remote.Close() + _, h := consumerCleanupFixture(t, remote.URL) + state := &tunnelCreateState{TunnelID: 43, Type: 2, OutNodes: []tunnelRuntimeNode{{NodeID: 1, Port: 31001}}, Nodes: map[int64]*nodeRecord{1: {ID: 1, Name: "remote", IsRemote: 1, RemoteURL: remote.URL, RemoteToken: "token"}}} + _, refs, err := h.applyFederationRuntime(state, "") + if err == nil || len(refs) != 1 || refs[0].ReservationID != "reserved" { + t.Fatalf("failed apply lost reservation: refs=%+v err=%v", refs, err) + } +} + +func TestTunnelUpdateDatabaseValidationPreservesSharedRuntime(t *testing.T) { + var releases atomic.Int64 + remote := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + releases.Add(1) + response.WriteJSON(w, response.OKEmpty()) + })) + defer remote.Close() + r, h := consumerCleanupFixture(t, remote.URL) + if err := r.DB().Create(&model.Node{ID: 2, Name: "entry", Secret: "entry", ServerIP: "127.0.0.2", Port: "31000-31010", Status: 1}).Error; err != nil { + t.Fatal(err) + } + if err := r.DB().Exec(`CREATE TRIGGER reject_tunnel_update BEFORE UPDATE ON tunnel BEGIN SELECT RAISE(ABORT, 'test constraint rejected'); END`).Error; err != nil { + t.Fatal(err) + } + rec := httptest.NewRecorder() + h.tunnelUpdate(rec, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", strings.NewReader(`{"id":42,"type":2,"name":"new name","inNodeId":[{"nodeId":2}],"outNodeId":[{"nodeId":1,"port":31001}]}`))) + var payload response.R + if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil { + t.Fatal(err) + } + active, err := r.ListActiveFederationTunnelBindingsByTunnel(42) + if err != nil || payload.Code == 0 || !strings.Contains(payload.Msg, "test constraint rejected") || releases.Load() != 0 || len(active) != 1 { + t.Fatalf("DB validation touched runtime: payload=%+v bindings=%v calls=%d err=%v", payload, active, releases.Load(), err) + } +} + +func TestTunnelUpdateRemoteReleaseFailurePreservesMetadata(t *testing.T) { + var requests atomic.Int64 + remote := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + requests.Add(1) + http.Error(w, "peer unavailable", http.StatusServiceUnavailable) + })) + defer remote.Close() + r, h := consumerCleanupFixture(t, remote.URL) + if err := r.DB().Create(&model.Node{ID: 2, Name: "entry", Secret: "entry", ServerIP: "127.0.0.2", Port: "31000-31010", Status: 1}).Error; err != nil { + t.Fatal(err) + } + rec := httptest.NewRecorder() + h.tunnelUpdate(rec, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", strings.NewReader(`{"id":42,"type":2,"name":"new name","inNodeId":[{"nodeId":2}],"outNodeId":[{"nodeId":1,"port":31001}]}`))) + var payload response.R + if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil { + t.Fatal(err) + } + name, err := r.GetTunnelName(42) + pending, pendingErr := r.ListPendingFederationTunnelBindings() + if err != nil || pendingErr != nil || payload.Code == 0 || name != "live tunnel" || requests.Load() != 1 || len(pending) != 1 { + t.Fatalf("failed cleanup continued update: payload=%+v name=%s calls=%d pending=%v err=%v/%v", payload, name, requests.Load(), pending, err, pendingErr) + } +} + +func TestTunnelUpdateChainWriteValidationPreservesSharedRuntime(t *testing.T) { + var calls atomic.Int64 + remote := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + calls.Add(1) + response.WriteJSON(w, response.OKEmpty()) + })) + defer remote.Close() + r, h := consumerCleanupFixture(t, remote.URL) + if err := r.DB().Create(&model.Node{ID: 2, Name: "entry", Secret: "entry", ServerIP: "127.0.0.2", Port: "31000-31010", Status: 1}).Error; err != nil { + t.Fatal(err) + } + if err := r.DB().Exec(`CREATE TRIGGER reject_chain_insert BEFORE INSERT ON chain_tunnel BEGIN SELECT RAISE(ABORT, 'chain write rejected'); END`).Error; err != nil { + t.Fatal(err) + } + rec := httptest.NewRecorder() + h.tunnelUpdate(rec, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", strings.NewReader(`{"id":42,"type":2,"name":"new name","inNodeId":[{"nodeId":2}],"outNodeId":[{"nodeId":1,"port":31001}]}`))) + var payload response.R + if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil { + t.Fatal(err) + } + active, err := r.ListActiveFederationTunnelBindingsByTunnel(42) + if err != nil || payload.Code == 0 || !strings.Contains(payload.Msg, "chain write rejected") || calls.Load() != 0 || len(active) != 1 { + t.Fatalf("chain validation touched runtime: payload=%+v active=%v calls=%d err=%v", payload, active, calls.Load(), err) + } + if name, err := r.GetTunnelName(42); err != nil || name != "live tunnel" { + t.Fatalf("preflight metadata update was committed: %s %v", name, err) + } +} diff --git a/go-backend/internal/http/handler/federation_resources.go b/go-backend/internal/http/handler/federation_resources.go new file mode 100644 index 0000000..8913471 --- /dev/null +++ b/go-backend/internal/http/handler/federation_resources.go @@ -0,0 +1,703 @@ +package handler + +import ( + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "sort" + "strconv" + "strings" + "time" + + "go-backend/internal/store/repo" + "go-backend/internal/ws" +) + +const peerShareResourcePrefix = "peer-share-" + +func peerShareResourceName(shareID int64, kind, original string) string { + return fmt.Sprintf("%s%d-%s-%s", peerShareResourcePrefix, shareID, kind, base64.RawURLEncoding.EncodeToString([]byte(original))) +} + +func parsePeerShareServiceName(name string) (shareID int64, originalName string, ok bool) { + if !strings.HasPrefix(name, peerShareResourcePrefix) { + return + } + parts := strings.SplitN(strings.TrimPrefix(name, peerShareResourcePrefix), "-", 3) + if len(parts) != 3 || parts[1] != "service" { + return + } + shareID, err := strconv.ParseInt(parts[0], 10, 64) + if err != nil || shareID <= 0 { + return 0, "", false + } + decoded, err := base64.RawURLEncoding.DecodeString(parts[2]) + if err != nil || len(decoded) == 0 { + return 0, "", false + } + return shareID, string(decoded), true +} + +func federationResourceCommandKind(cmd string) (kind, action string) { + lower := strings.ToLower(cmd) + for _, entry := range []struct{ suffix, kind string }{{"service", "service"}, {"chains", "chain"}, {"climiters", "climiter"}, {"limiters", "limiter"}} { + if strings.HasSuffix(lower, entry.suffix) { + return entry.kind, strings.TrimSuffix(lower, entry.suffix) + } + } + return "", "" +} + +// Every peer-supplied reference is resolved within the same share namespace. +// Composite traffic limiters use a comma-separated list in GOST. +func scopePeerResourceReferences(value interface{}, shareID int64) { + switch v := value.(type) { + case map[string]interface{}: + for key, child := range v { + if strings.EqualFold(key, "chains") { + if names, ok := child.([]interface{}); ok { + for i, name := range names { + if raw, ok := name.(string); ok { + names[i] = peerShareResourceName(shareID, "chain", raw) + } + } + continue + } + } + kind := "" + switch strings.ToLower(key) { + case "chain": + kind = "chain" + case "limiter": + kind = "limiter" + case "climiter": + kind = "climiter" + } + if raw, ok := child.(string); ok && kind != "" && strings.TrimSpace(raw) != "" { + names := strings.Split(raw, ",") + for i, name := range names { + names[i] = peerShareResourceName(shareID, kind, strings.TrimSpace(name)) + } + v[key] = strings.Join(names, ",") + } else { + scopePeerResourceReferences(child, shareID) + } + } + case []interface{}: + for _, child := range v { + scopePeerResourceReferences(child, shareID) + } + } +} + +func peerResourceDeletePayload(item repo.PeerShareResource) interface{} { + if item.Kind == "service" { + return map[string]interface{}{"services": []string{item.RuntimeName}} + } + key := item.Kind + if key == "climiter" { + key = "limiter" + } + return map[string]interface{}{key: item.RuntimeName} +} +func peerResourceDeleteCommand(kind string) string { + switch kind { + case "service": + return "DeleteService" + case "chain": + return "DeleteChains" + case "climiter": + return "DeleteCLimiters" + default: + return "DeleteLimiters" + } +} + +// preparePeerResourceCommand validates and writes desired state before the node +// sees any command, closing the flow-report race even without a reservation. +func (h *Handler) preparePeerResourceCommand(share *repo.PeerShare, cmd string, data interface{}) ([]repo.PeerShareResource, error) { + kind, action := federationResourceCommandKind(cmd) + if kind == "" { + return nil, fmt.Errorf("command not allowed") + } + raw, err := json.Marshal(data) + if err != nil { + return nil, err + } + var decoded interface{} + if err = json.Unmarshal(raw, &decoded); err != nil { + return nil, err + } + var configs []map[string]interface{} + var names []string + setting := action == "add" || action == "update" + if setting { + if kind == "service" { + configs = extractFederationServiceEntries(decoded) + } else if m, ok := decoded.(map[string]interface{}); ok { + if nested, ok := m["data"].(map[string]interface{}); ok { + configs = []map[string]interface{}{nested} + } else { + configs = []map[string]interface{}{m} + } + } + if len(configs) == 0 { + return nil, fmt.Errorf("resource configuration is required") + } + for _, c := range configs { + names = append(names, strings.TrimSpace(asString(c["name"]))) + } + } else { + if kind == "service" { + if m, ok := decoded.(map[string]interface{}); ok { + for _, v := range asAnySlice(m["services"]) { + names = append(names, strings.TrimSpace(asString(v))) + } + } + } else { + if s, ok := decoded.(string); ok { + names = []string{strings.TrimSpace(s)} + } else if m, ok := decoded.(map[string]interface{}); ok { + key := kind + if key == "climiter" { + key = "limiter" + } + names = []string{strings.TrimSpace(asString(m[key]))} + } + } + if len(names) == 0 { + return nil, fmt.Errorf("resource name is required") + } + } + requestedNames := make(map[string]bool, len(names)) + for _, name := range names { + requestedNames[name] = true + } + items := make([]repo.PeerShareResource, 0, len(names)) + seen := map[string]bool{} + for i, name := range names { + if name == "" || strings.HasPrefix(name, peerShareResourcePrefix) { + return nil, fmt.Errorf("invalid peer resource name") + } + if seen[name] { + return nil, fmt.Errorf("duplicate resource name") + } + seen[name] = true + old, err := h.repo.GetPeerShareResource(share.ID, kind, name) + if err != nil { + return nil, err + } + item := repo.PeerShareResource{ShareID: share.ID, NodeID: share.NodeID, Kind: kind, OriginalName: name, RuntimeName: peerShareResourceName(share.ID, kind, name), DesiredState: "active", UpdatedTime: time.Now().UnixMilli()} + if old != nil { + item = *old + item.Applied = 0 + item.UpdatedTime = time.Now().UnixMilli() + } + if setting { + if old == nil && kind == "service" { + legacy, err := h.peerShareLegacyServiceNames(share, name) + if err != nil { + return nil, err + } + if len(legacy) > 0 { + encoded, _ := json.Marshal(legacy) + item.LegacyNames = string(encoded) + item.LegacyServiceBase = normalizeForwardRuntimeServiceName(name) + } + } + config := configs[i] + if err := validatePeerResourceReferences(config, kind); err != nil { + return nil, err + } + scopePeerResourceReferences(config, share.ID) + config["name"] = item.RuntimeName + encoded, err := json.Marshal(config) + if err != nil { + return nil, err + } + item.Config = string(encoded) + item.DesiredState = "active" + item.ReleaseLegacyFamily = false + } else { + if old == nil { + // Delete callers send candidate base/TCP/UDP names. Missing names + // are safe no-ops; never forward a raw legacy name to the node. + if action == "delete" { + if kind != "service" { + continue + } + legacy, err := h.peerShareLegacyServiceNames(share, name) + if err != nil { + return nil, err + } + if len(legacy) == 0 { + continue + } + encoded, _ := json.Marshal(legacy) + item.LegacyNames = string(encoded) + item.LegacyServiceBase = normalizeForwardRuntimeServiceName(name) + } else { + return nil, fmt.Errorf("service %q not found", name) + } + } + if old != nil && old.DesiredState == "deleted" && action != "delete" { + return nil, fmt.Errorf("service %q not found", name) + } + switch action { + case "delete": + item.DesiredState = "deleted" + base := normalizeForwardRuntimeServiceName(name) + if requestedNames[base] && requestedNames[base+"_tcp"] && requestedNames[base+"_udp"] { + item.ReleaseLegacyFamily = true + } + case "pause": + item.DesiredState = "paused" + case "resume": + if item.Config == "" { + return nil, fmt.Errorf("resource has no saved configuration") + } + item.DesiredState = "active" + default: + return nil, fmt.Errorf("command not allowed") + } + } + items = append(items, item) + } + // Ownership and desired config commit atomically. A failed resource write + // must not rename a legacy binding and expose its remaining transports. + err = h.repo.WithPeerShareResourceTransaction(func(tx *repo.Repository) error { + if setting && kind == "service" { + scoped := make([]interface{}, 0, len(items)) + for _, item := range items { + var c map[string]interface{} + _ = json.Unmarshal([]byte(item.Config), &c) + scoped = append(scoped, c) + } + binder := &Handler{repo: tx} + if err := binder.bindPeerShareForwardRuntimeServices(share, scoped); err != nil { + return err + } + } + return tx.SavePeerShareResources(items) + }) + if err != nil { + return nil, err + } + return items, nil +} + +func (h *Handler) applyPeerShareResource(item repo.PeerShareResource) (ws.CommandResult, error) { + var result ws.CommandResult + var err error + if item.LegacyNames != "" { + var names []string + if err := json.Unmarshal([]byte(item.LegacyNames), &names); err != nil { + return result, err + } + share, loadErr := h.repo.GetPeerShare(item.ShareID) + if loadErr != nil { + return result, loadErr + } + if share == nil { + return result, fmt.Errorf("legacy resource ownership missing") + } + for _, name := range names { + owned, checkErr := h.peerShareLegacyServiceNames(share, name) + if checkErr != nil { + return result, checkErr + } + if len(owned) == 0 { + return result, fmt.Errorf("legacy resource ownership missing") + } + } + if _, err = h.sendNodeCommand(item.NodeID, "DeleteService", map[string]interface{}{"services": names}, false, true); err != nil { + return result, err + } + if err = h.repo.ClearPeerShareResourceLegacyNames(item.ShareID, item.Kind, item.OriginalName); err != nil { + return result, err + } + } + if item.DesiredState == "deleted" || item.DesiredState == "paused" { + // The provider retains the paused configuration durably. Keep the + // listener absent on reconnect instead of briefly starting it before + // sending a second pause command. Resume reapplies the saved config. + result, err = h.sendNodeCommand(item.NodeID, peerResourceDeleteCommand(item.Kind), peerResourceDeletePayload(item), false, true) + } else { + var config map[string]interface{} + if err = json.Unmarshal([]byte(item.Config), &config); err != nil { + return result, err + } + if item.Kind == "service" { + result, err = h.sendNodeCommand(item.NodeID, "UpdateService", []interface{}{config}, false, false) + } else { + suffix := "Limiters" + key := "limiter" + if item.Kind == "chain" { + suffix = "Chains" + key = "chain" + } + if item.Kind == "climiter" { + suffix = "CLimiters" + } + result, err = h.sendNodeCommand(item.NodeID, "Add"+suffix, config, false, false) + if err != nil && isAlreadyExistsMessage(err.Error()) { + result, err = h.sendNodeCommand(item.NodeID, "Update"+suffix, map[string]interface{}{key: item.RuntimeName, "data": config}, false, false) + } + } + } + if err != nil { + return result, err + } + if item.Kind == "service" && item.DesiredState == "deleted" { + // TCP and UDP may share one reservation. Release only after all variants + // have an acknowledged tombstone. + items, listErr := h.repo.ListPeerShareResourcesByNode(item.NodeID) + if listErr != nil { + return result, listErr + } + base := normalizeForwardRuntimeServiceName(item.OriginalName) + for _, other := range items { + if other.ShareID == item.ShareID && other.Kind == "service" && other.OriginalName != item.OriginalName && normalizeForwardRuntimeServiceName(other.OriginalName) == base && (other.DesiredState != "deleted" || other.Applied == 0) { + return result, h.repo.MarkPeerShareResourceApplied(item.ShareID, item.Kind, item.OriginalName) + } + } + legacyBase := item.LegacyServiceBase + if legacyBase == "" { + for _, other := range items { + if other.ShareID == item.ShareID && normalizeForwardRuntimeServiceName(other.OriginalName) == base && other.LegacyServiceBase != "" { + legacyBase = other.LegacyServiceBase + break + } + } + } + if legacyBase != "" && item.ReleaseLegacyFamily { + share, loadErr := h.repo.GetPeerShare(item.ShareID) + if loadErr != nil { + return result, loadErr + } + if share == nil { + return result, fmt.Errorf("share ownership missing") + } + if _, checkErr := h.peerShareLegacyServiceNames(share, legacyBase); checkErr != nil { + return result, checkErr + } + if _, deleteErr := h.sendNodeCommand(item.NodeID, "DeleteService", map[string]interface{}{"services": buildForwardServiceDeleteNames([]string{legacyBase})}, false, true); deleteErr != nil { + return result, deleteErr + } + } + if legacyBase != "" && !item.ReleaseLegacyFamily { + return result, h.repo.MarkPeerShareResourceApplied(item.ShareID, item.Kind, item.OriginalName) + } + runtimes, listErr := h.repo.ListActivePeerShareRuntimesByShareID(item.ShareID) + if listErr != nil { + return result, listErr + } + return result, h.repo.WithPeerShareResourceTransaction(func(tx *repo.Repository) error { + for _, runtime := range runtimes { + sid, original, scoped := parsePeerShareServiceName(runtime.ServiceName) + ownedScoped := scoped && sid == item.ShareID && normalizeForwardRuntimeServiceName(original) == base + ownedLegacy := !scoped && runtime.Role == "forward" && item.ReleaseLegacyFamily && legacyBase != "" && normalizeForwardRuntimeServiceName(runtime.ServiceName) == legacyBase + if ownedScoped || ownedLegacy { + if err = tx.CompletePeerShareRuntimeRelease(runtime.ID); err != nil { + return err + } + } + } + if legacyBase != "" && item.ReleaseLegacyFamily { + if err := tx.ClearPeerShareResourceLegacyFamily(item.ShareID, legacyBase); err != nil { + return err + } + } + return tx.MarkPeerShareResourceApplied(item.ShareID, item.Kind, item.OriginalName) + }) + } + return result, h.repo.MarkPeerShareResourceApplied(item.ShareID, item.Kind, item.OriginalName) +} + +func orderPeerShareResources(items []repo.PeerShareResource) { + rank := func(item repo.PeerShareResource) int { + if item.DesiredState == "deleted" { + if item.Kind == "service" { + return 0 + } + return 1 + } + if item.Kind == "service" { + return 4 + } + if item.Kind == "chain" { + return 3 + } + return 2 + } + sort.SliceStable(items, func(i, j int) bool { return rank(items[i]) < rank(items[j]) }) +} + +func (h *Handler) reconcilePeerShareResourcesOnNode(nodeID int64) error { + return h.reconcilePeerShareResources(nodeID, false) +} +func (h *Handler) retryPendingPeerShareResourcesOnNode(nodeID int64) error { + return h.reconcilePeerShareResources(nodeID, true) +} +func (h *Handler) reconcilePeerShareResources(nodeID int64, pendingOnly bool) error { + h.peerResourceMu.Lock() + defer h.peerResourceMu.Unlock() + items, err := h.repo.ListPeerShareResourcesByNode(nodeID) + if err != nil { + return err + } + runtimes, err := h.repo.ListActivePeerShareRuntimesByNode(nodeID) + if err != nil { + return err + } + pending := map[string]bool{} + for _, runtime := range runtimes { + if runtime.ReleasePending != 0 { + sid, name, ok := parsePeerShareServiceName(runtime.ServiceName) + if ok { + pending[fmt.Sprintf("%d:%s", sid, normalizeForwardRuntimeServiceName(name))] = true + } + } + } + groups := map[int64][]repo.PeerShareResource{} + var ids []int64 + for _, item := range items { + if _, ok := groups[item.ShareID]; !ok { + ids = append(ids, item.ShareID) + } + groups[item.ShareID] = append(groups[item.ShareID], item) + } + var failures []error + for _, shareID := range ids { + share, err := h.repo.GetPeerShare(shareID) + if err != nil { + failures = append(failures, err) + continue + } + expired := share == nil || share.IsActive != 1 || (share.ExpiryTime > 0 && share.ExpiryTime <= time.Now().UnixMilli()) || isPeerShareFlowExceeded(share) + group := groups[shareID] + var changed []repo.PeerShareResource + for i := range group { + item := &group[i] + if (expired || (item.Kind == "service" && pending[fmt.Sprintf("%d:%s", item.ShareID, normalizeForwardRuntimeServiceName(item.OriginalName))])) && (item.DesiredState != "deleted" || item.Applied == 0 || item.LegacyServiceBase != "") { + item.DesiredState = "deleted" + item.ReleaseLegacyFamily = true + item.Applied = 0 + changed = append(changed, *item) + } + } + if err := h.repo.SavePeerShareResources(changed); err != nil { + failures = append(failures, err) + continue + } + orderPeerShareResources(group) + dependencyFailed := false + deletionFailed := false + for _, item := range group { + if item.Applied == 1 && (pendingOnly || item.DesiredState == "deleted") { + continue + } + if dependencyFailed && item.Kind == "service" && item.DesiredState != "deleted" { + continue + } + if deletionFailed && item.Kind != "service" && item.DesiredState == "deleted" { + continue + } + if _, err := h.applyPeerShareResource(item); err != nil { + failures = append(failures, fmt.Errorf("share %d %s %s: %w", item.ShareID, item.Kind, item.OriginalName, err)) + if item.Kind != "service" { + dependencyFailed = true + } + if item.DesiredState == "deleted" { + deletionFailed = true + } + } + } + } + return errors.Join(failures...) +} + +func (h *Handler) releasePeerShareResources(shareID int64) error { + h.peerResourceMu.Lock() + defer h.peerResourceMu.Unlock() + share, err := h.repo.GetPeerShare(shareID) + if err != nil { + return err + } + if share == nil { + return nil + } + all, err := h.repo.ListPeerShareResourcesByNode(share.NodeID) + if err != nil { + return err + } + var items []repo.PeerShareResource + for _, item := range all { + if item.ShareID == shareID && !(item.DesiredState == "deleted" && item.Applied == 1 && item.LegacyServiceBase == "") { + item.DesiredState = "deleted" + item.ReleaseLegacyFamily = true + item.Applied = 0 + items = append(items, item) + } + } + if err = h.repo.SavePeerShareResources(items); err != nil { + return err + } + orderPeerShareResources(items) + for _, item := range items { + if _, err = h.applyPeerShareResource(item); err != nil { + return err + } + } + return nil +} + +// A legacy name is eligible for migration only when its runtime registration +// proves unique ownership. An identically numbered local forward is ambiguous. +func (h *Handler) peerShareLegacyServiceNames(share *repo.PeerShare, original string) ([]string, error) { + base := normalizeForwardRuntimeServiceName(original) + runtimes, err := h.repo.ListActiveForwardPeerShareRuntimesByNodeAndServiceName(share.NodeID, base) + if err != nil { + return nil, err + } + owned := false + for _, runtime := range runtimes { + if runtime.ShareID != share.ID { + return nil, fmt.Errorf("legacy resource %q has ambiguous ownership", original) + } + if runtime.ReleasePending != 0 { + return nil, fmt.Errorf("runtime release is pending") + } + owned = true + } + if !owned { + resources, err := h.repo.ListPeerShareResourcesByNode(share.NodeID) + if err != nil { + return nil, err + } + for _, resource := range resources { + if resource.ShareID == share.ID && resource.LegacyServiceBase == base { + owned = true + break + } + } + } + if !owned { + return nil, nil + } + if id, _, _, ok := parseFlowServiceIDs(base); ok { + local, err := h.repo.GetForwardRecord(id) + if err != nil { + return nil, err + } + if local != nil { + return nil, fmt.Errorf("legacy resource %q collides with a local forward", original) + } + } + // Only delete the requested transport. Other variants remain registered to + // their legacy runtime until they are independently migrated. + return []string{original}, nil +} + +// releasePeerShareForwardRuntimeResources removes exact persisted transport +// names belonging to a single forward reservation, retaining failed tombstones. +func (h *Handler) releasePeerShareForwardRuntimeResources(runtime *repo.PeerShareRuntime) error { + h.peerResourceMu.Lock() + defer h.peerResourceMu.Unlock() + shareID, original, ok := parsePeerShareServiceName(runtime.ServiceName) + if !ok || shareID != runtime.ShareID { + return fmt.Errorf("runtime is not a scoped shared forward") + } + all, err := h.repo.ListPeerShareResourcesByNode(runtime.NodeID) + if err != nil { + return err + } + var items []repo.PeerShareResource + for _, item := range all { + if item.ShareID == runtime.ShareID && item.Kind == "service" && normalizeForwardRuntimeServiceName(item.OriginalName) == normalizeForwardRuntimeServiceName(original) { + item.DesiredState = "deleted" + item.ReleaseLegacyFamily = true + item.Applied = 0 + items = append(items, item) + } + } + if len(items) == 0 { + return fmt.Errorf("shared forward resource ownership is missing") + } + if err = h.repo.SavePeerShareResources(items); err != nil { + return err + } + for _, item := range items { + if _, err = h.applyPeerShareResource(item); err != nil { + return err + } + } + return nil +} + +// Registry kinds without an owned command API cannot be referenced by peers. +func validatePeerResourceReferences(value interface{}, kind string) error { + var walk func(interface{}) error + walk = func(value interface{}) error { + switch v := value.(type) { + case map[string]interface{}: + for key, child := range v { + switch strings.ToLower(key) { + case "auther", "authers", "admission", "admissions", "bypass", "bypasses", "resolver", "hosts", "rlimiter", "logger", "loggers", "observer", "recorders", "hop", "sd": + if child != nil { + empty := false + switch x := child.(type) { + case string: + empty = strings.TrimSpace(x) == "" + case []interface{}: + empty = len(x) == 0 + } + if !empty { + return fmt.Errorf("unsupported shared registry reference: %s", key) + } + } + case "forwarder": + if f, ok := child.(map[string]interface{}); ok { + for field, value := range f { + if strings.EqualFold(field, "name") && strings.TrimSpace(asString(value)) != "" { + return fmt.Errorf("named shared forwarder references are unsupported") + } + } + } + case "hops": + for _, hop := range asMapSlice(child) { + if _, ok := hop["nodes"]; !ok { + return fmt.Errorf("shared chains require inline hop nodes") + } + for _, loader := range []string{"file", "redis", "http", "plugin"} { + if hop[loader] != nil { + return fmt.Errorf("shared hop loaders are unsupported") + } + } + } + } + if err := walk(child); err != nil { + return err + } + } + case []interface{}: + for _, child := range v { + if err := walk(child); err != nil { + return err + } + } + } + return nil + } + if kind == "limiter" || kind == "climiter" { + if config, ok := value.(map[string]interface{}); ok { + for _, loader := range []string{"file", "redis", "http", "plugin"} { + if config[loader] != nil { + return fmt.Errorf("shared limiter loaders are unsupported") + } + } + } + } + return walk(value) +} diff --git a/go-backend/internal/http/handler/federation_resources_test.go b/go-backend/internal/http/handler/federation_resources_test.go new file mode 100644 index 0000000..9c7a993 --- /dev/null +++ b/go-backend/internal/http/handler/federation_resources_test.go @@ -0,0 +1,538 @@ +package handler + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "go-backend/internal/http/response" + "go-backend/internal/store/repo" + "gorm.io/gorm" +) + +func resourceTestShare(t *testing.T, h *Handler, token string) *repo.PeerShare { + t.Helper() + now := time.Now().UnixMilli() + share := &repo.PeerShare{Name: token, NodeID: 1, Token: token, IsActive: 1, PortRangeStart: 31000, PortRangeEnd: 32000, CreatedTime: now, UpdatedTime: now} + if err := h.repo.CreatePeerShare(share); err != nil { + t.Fatal(err) + } + return share +} +func resourceTestCommand(t *testing.T, h *Handler, share *repo.PeerShare, cmd string, data interface{}) response.R { + t.Helper() + body, err := json.Marshal(federationRuntimeCommandRequest{CommandType: cmd, Data: data}) + if err != nil { + t.Fatal(err) + } + req := httptest.NewRequest(http.MethodPost, "/api/v1/federation/runtime/command", bytes.NewReader(body)) + req.Header.Set("Authorization", "Bearer "+share.Token) + rec := httptest.NewRecorder() + h.federationRuntimeCommand(rec, req) + var result response.R + if err = json.Unmarshal(rec.Body.Bytes(), &result); err != nil { + t.Fatal(err) + } + return result +} +func resourceTestService(name string, port int) []interface{} { + return []interface{}{map[string]interface{}{"name": name, "addr": fmt.Sprintf(":%d", port), "handler": map[string]interface{}{"type": "tcp", "chain": "70"}, "listener": map[string]interface{}{"type": "tcp"}, "limiter": "10,20", "climiter": "30"}} +} + +func TestPeerResourceCommandIsolationAndDurableRestore(t *testing.T) { + a := newCleanupAgent(t) + first := resourceTestShare(t, a.h, "resource-first") + second := resourceTestShare(t, a.h, "resource-second") + for i, share := range []*repo.PeerShare{first, second} { + for _, entry := range []struct{ cmd, name string }{{"AddLimiters", "10"}, {"AddCLimiters", "30"}, {"AddChains", "70"}} { + result := resourceTestCommand(t, a.h, share, entry.cmd, map[string]interface{}{"name": entry.name}) + if result.Code != 0 { + t.Fatal(result.Msg) + } + } + result := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("70_1_0_tcp", 31001+i)) + if result.Code != 0 { + t.Fatal(result.Msg) + } + } + commands := a.commandsOfType("UpdateService") + if len(commands) != 2 { + t.Fatalf("commands: %+v", commands) + } + for i, share := range []*repo.PeerShare{first, second} { + var configs []map[string]interface{} + if err := json.Unmarshal(commands[i].Data, &configs); err != nil { + t.Fatal(err) + } + c := configs[0] + if c["name"] != peerShareResourceName(share.ID, "service", "70_1_0_tcp") { + t.Fatalf("unscoped service: %v", c) + } + if c["handler"].(map[string]interface{})["chain"] != peerShareResourceName(share.ID, "chain", "70") { + t.Fatalf("unscoped chain: %v", c) + } + want := peerShareResourceName(share.ID, "limiter", "10") + "," + peerShareResourceName(share.ID, "limiter", "20") + if c["limiter"] != want || c["climiter"] != peerShareResourceName(share.ID, "climiter", "30") { + t.Fatalf("unscoped limiter: %v", c) + } + } + result := resourceTestCommand(t, a.h, first, "DeleteService", map[string]interface{}{"services": []string{"70_1_0_tcp"}}) + if result.Code != 0 { + t.Fatal(result.Msg) + } + deleted := a.commandsOfType("DeleteService") + if len(deleted) != 1 || !strings.Contains(string(deleted[0].Data), peerShareResourceName(first.ID, "service", "70_1_0_tcp")) { + t.Fatalf("wrong deletion: %+v", deleted) + } + // A new Handler simulates restart with only durable desired state retained. + restarted := &Handler{repo: a.h.repo, wsServer: a.h.wsServer} + if err := restarted.reconcilePeerShareResourcesOnNode(1); err != nil { + t.Fatal(err) + } + commands = a.commandsOfType("UpdateService") + if len(commands) != 3 || !strings.Contains(string(commands[2].Data), peerShareResourceName(second.ID, "service", "70_1_0_tcp")) { + t.Fatalf("wrong recovery: %+v", commands) + } + if got := resourceTestCommand(t, a.h, first, "DeleteService", map[string]interface{}{"services": []string{peerShareResourceName(second.ID, "service", "70_1_0_tcp")}}); got.Code == 0 { + t.Fatal("accepted foreign scoped name") + } + if got := resourceTestCommand(t, a.h, first, "Reload", nil); got.Code == 0 { + t.Fatal("accepted global reload") + } +} + +func TestPeerResourcePreRegistrationAndDatabaseFailure(t *testing.T) { + a := newCleanupAgent(t) + share := resourceTestShare(t, a.h, "resource-preregister") + items, err := a.h.preparePeerResourceCommand(share, "AddService", resourceTestService("70_1_0_tcp", 31001)) + if err != nil { + t.Fatal(err) + } + if len(a.commandsOfType("UpdateService")) != 0 { + t.Fatal("prepare sent service before persistence") + } + stored, err := a.h.repo.GetPeerShareResource(share.ID, "service", "70_1_0_tcp") + if err != nil || stored == nil || stored.Applied != 0 { + t.Fatalf("missing pending ownership: %+v %v", stored, err) + } + runCleanupPath(a.h, "single", []string{items[0].RuntimeName}) + a.probe(t) + if len(a.commandsOfType("DeleteService")) != 0 { + t.Fatal("flow report removed pending service") + } + if err := a.h.repo.DB().Callback().Create().Before("gorm:create").Register("fail-resource-registration", func(tx *gorm.DB) { + if tx.Statement.Table == "peer_share_resource" { + tx.AddError(errors.New("simulated resource write failure")) + } + }); err != nil { + t.Fatal(err) + } + defer a.h.repo.DB().Callback().Create().Remove("fail-resource-registration") + got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("71_1_0", 31002)) + if got.Code == 0 { + t.Fatal("database failure returned success") + } + if len(a.commandsOfType("UpdateService")) != 0 { + t.Fatal("node received command after persistence failure") + } +} + +func TestPeerResourceBindingDatabaseFailure(t *testing.T) { + a := newCleanupAgent(t) + share := resourceTestShare(t, a.h, "resource-bind-failure") + if err := a.h.repo.DB().Callback().Create().Before("gorm:create").Register("fail-runtime-registration", func(tx *gorm.DB) { + if tx.Statement.Table == "peer_share_runtime" { + tx.AddError(errors.New("simulated binding failure")) + } + }); err != nil { + t.Fatal(err) + } + defer a.h.repo.DB().Callback().Create().Remove("fail-runtime-registration") + got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("70_1_0", 31001)) + if got.Code == 0 { + t.Fatal("binding failure returned success") + } + if len(a.commandsOfType("UpdateService")) != 0 { + t.Fatal("service sent before successful binding") + } +} + +func TestPeerResourceScopedNameRoundTrip(t *testing.T) { + for _, name := range []string{"70_1_0", "70_1_0_tcp", "70_1_0_udp", "a-b_c"} { + scoped := peerShareResourceName(12, "service", name) + id, got, ok := parsePeerShareServiceName(scoped) + if !ok || id != 12 || got != name { + t.Fatalf("round trip failed: %q", scoped) + } + } +} + +func TestPeerResourceDeletesCandidateNamesAndPreservesFailedTombstone(t *testing.T) { + a := newCleanupAgent(t) + share := resourceTestShare(t, a.h, "resource-delete-candidates") + got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("70_1_0_tcp", 31001)) + if got.Code != 0 { + t.Fatal(got.Msg) + } + got = resourceTestCommand(t, a.h, share, "DeleteService", map[string]interface{}{"services": []string{"70_1_0_tcp", "70_1_0_udp", "70_1_0"}}) + if got.Code != 0 { + t.Fatal(got.Msg) + } + if len(a.commandsOfType("DeleteService")) != 1 { + t.Fatal("unregistered names were sent to the node") + } + // A disconnected node cannot acknowledge deletion: keep the pending row. + if got = resourceTestCommand(t, a.h, share, "AddService", resourceTestService("71_1_0", 31002)); got.Code != 0 { + t.Fatal(got.Msg) + } + offline := &Handler{repo: a.h.repo} + got = resourceTestCommand(t, offline, share, "DeleteService", map[string]interface{}{"services": []string{"71_1_0"}}) + if got.Code == 0 { + t.Fatal("offline deletion reported success") + } + pending, err := a.h.repo.GetPeerShareResource(share.ID, "service", "71_1_0") + if err != nil || pending.DesiredState != "deleted" || pending.Applied != 0 { + t.Fatalf("lost pending deletion: %+v %v", pending, err) + } + if err := a.h.reconcilePeerShareResourcesOnNode(1); err != nil { + t.Fatal(err) + } + pending, err = a.h.repo.GetPeerShareResource(share.ID, "service", "71_1_0") + if err != nil || pending.Applied != 1 { + t.Fatalf("deletion not retried: %+v %v", pending, err) + } + if got = resourceTestCommand(t, a.h, share, "ResumeService", map[string]interface{}{"services": []string{"71_1_0"}}); got.Code == 0 { + t.Fatal("resume resurrected a deleted resource") + } +} + +func TestPeerResourceLegacyMigrationRequiresUnambiguousOwner(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("70_1_0_tcp", 31001)) + if got.Code != 0 { + t.Fatal(got.Msg) + } + deleted := a.commandsOfType("DeleteService") + if len(deleted) != 1 || string(deleted[0].Data) != `{"services":["70_1_0_tcp"]}` { + t.Fatalf("legacy deletion was not exact: %+v", deleted) + } + item, err := a.h.repo.GetPeerShareResource(share.ID, "service", "70_1_0_tcp") + if err != nil || item.LegacyNames != "" { + t.Fatalf("legacy migration acknowledgment not persisted: %+v %v", item, err) + } + // Two shares with the same legacy name are never resolved by guessing. + other := resourceTestShare(t, a.h, "resource-ambiguous") + now := time.Now().UnixMilli() + for i, sid := range []int64{share.ID, other.ID} { + if err := a.h.repo.CreatePeerShareRuntime(&repo.PeerShareRuntime{ShareID: sid, NodeID: 1, ReservationID: fmt.Sprintf("ambiguous-%d", i), ResourceKey: fmt.Sprintf("ambiguous-%d", i), Role: "forward", ServiceName: "71_1_0", Port: 31003 + i, Applied: 1, Status: 1, CreatedTime: now, UpdatedTime: now}); err != nil { + t.Fatal(err) + } + } + got = resourceTestCommand(t, a.h, share, "AddService", resourceTestService("71_1_0_tcp", 31003)) + if got.Code == 0 { + t.Fatal("ambiguous legacy owner accepted") + } + if len(a.commandsOfType("DeleteService")) != 1 { + t.Fatal("ambiguous legacy service deleted") + } +} + +func TestPeerResourceLegacyFamilyPersistsAcrossPartialMigration(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + if got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("70_1_0_tcp", 31001)); got.Code != 0 { + t.Fatal(got.Msg) + } + item, err := a.h.repo.GetPeerShareResource(share.ID, "service", "70_1_0_tcp") + if err != nil || item.LegacyServiceBase != "70_1_0" { + t.Fatalf("lost legacy family ownership: %+v %v", item, err) + } + // A later request can still prove ownership of the old UDP transport. + if got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("70_1_0_udp", 31001)); got.Code != 0 { + t.Fatal(got.Msg) + } + deletes := a.commandsOfType("DeleteService") + if len(deletes) != 2 || string(deletes[1].Data) != `{"services":["70_1_0_udp"]}` { + t.Fatalf("lost UDP migration ownership: %+v", deletes) + } + if got := resourceTestCommand(t, a.h, share, "DeleteService", map[string]interface{}{"services": []string{"70_1_0_tcp", "70_1_0_udp", "70_1_0"}}); got.Code != 0 { + t.Fatal(got.Msg) + } + item, err = a.h.repo.GetPeerShareResource(share.ID, "service", "70_1_0_tcp") + if err != nil || item.LegacyServiceBase != "" { + t.Fatalf("legacy family not released: %+v %v", item, err) + } +} + +func TestPeerResourceFailedRegistrationDoesNotRenameLegacyRuntime(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + if err := a.h.repo.DB().Callback().Create().Before("gorm:create").Register("fail-atomic-resource", func(tx *gorm.DB) { + if tx.Statement.Table == "peer_share_resource" { + tx.AddError(errors.New("simulated desired-state failure")) + } + }); err != nil { + t.Fatal(err) + } + defer a.h.repo.DB().Callback().Create().Remove("fail-atomic-resource") + if got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("70_1_0_tcp", 31001)); got.Code == 0 { + t.Fatal("registration failure returned success") + } + runtimes, err := a.h.repo.ListActivePeerShareRuntimesByShareID(share.ID) + if err != nil || len(runtimes) != 1 || runtimes[0].ServiceName != "70_1_0" { + t.Fatalf("legacy binding changed after rollback: %+v %v", runtimes, err) + } + if len(a.commandsOfType("DeleteService"))+len(a.commandsOfType("UpdateService")) != 0 { + t.Fatal("node mutated despite transaction rollback") + } +} + +func TestPeerResourcePausedReconcileNeverStartsListener(t *testing.T) { + a := newCleanupAgent(t) + share := resourceTestShare(t, a.h, "resource-pause-recovery") + for _, cmd := range []string{"AddService", "PauseService"} { + var data interface{} = resourceTestService("70_1_0_tcp", 31001) + if cmd == "PauseService" { + data = map[string]interface{}{"services": []string{"70_1_0_tcp"}} + } + if got := resourceTestCommand(t, a.h, share, cmd, data); got.Code != 0 { + t.Fatal(got.Msg) + } + } + if err := a.h.reconcilePeerShareResourcesOnNode(1); err != nil { + t.Fatal(err) + } + if len(a.commandsOfType("UpdateService")) != 1 { + t.Fatal("paused listener started during reconciliation") + } + if got := resourceTestCommand(t, a.h, share, "ResumeService", map[string]interface{}{"services": []string{"70_1_0_tcp"}}); got.Code != 0 { + t.Fatal(got.Msg) + } + if len(a.commandsOfType("UpdateService")) != 2 { + t.Fatal("resume did not restore saved service") + } +} + +func TestPeerResourceDeleteUnmigratedLegacyOwner(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + got := resourceTestCommand(t, a.h, share, "DeleteService", map[string]interface{}{"services": []string{"70_1_0_tcp", "70_1_0_udp", "70_1_0"}}) + if got.Code != 0 { + t.Fatal(got.Msg) + } + removed := false + for _, cmd := range a.commandsOfType("DeleteService") { + var body struct { + Services []string `json:"services"` + } + if err = json.Unmarshal(cmd.Data, &body); err != nil { + t.Fatal(err) + } + for _, name := range body.Services { + if name == "70_1_0" { + removed = true + } + } + } + if !removed { + t.Fatal("legacy delete reported success without sending old service deletion") + } +} + +func TestPeerResourcePartialDeletePreservesLegacyFamily(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + if got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("70_1_0_tcp", 31001)); got.Code != 0 { + t.Fatal(got.Msg) + } + if got := resourceTestCommand(t, a.h, share, "DeleteService", map[string]interface{}{"services": []string{"70_1_0_tcp"}}); got.Code != 0 { + t.Fatal(got.Msg) + } + item, err := a.h.repo.GetPeerShareResource(share.ID, "service", "70_1_0_tcp") + if err != nil || item.LegacyServiceBase != "70_1_0" { + t.Fatalf("partial delete lost remaining legacy family: %+v %v", item, err) + } + for _, cmd := range a.commandsOfType("DeleteService") { + if strings.Contains(string(cmd.Data), `"70_1_0_udp"`) { + t.Fatal("partial TCP delete removed legacy UDP") + } + } + if err = a.h.releasePeerShareResources(share.ID); err != nil { + t.Fatal(err) + } + item, err = a.h.repo.GetPeerShareResource(share.ID, "service", "70_1_0_tcp") + if err != nil || item.LegacyServiceBase != "" { + t.Fatalf("full release lost legacy cleanup: %+v %v", item, err) + } +} + +func TestPeerResourceChainGroupsAreScopedAndOtherRegistriesRejected(t *testing.T) { + a := newCleanupAgent(t) + share := resourceTestShare(t, a.h, "resource-chain-groups") + data := resourceTestService("70_1_0", 31001) + data[0].(map[string]interface{})["handler"].(map[string]interface{})["chainGroup"] = map[string]interface{}{"chains": []string{"one", "two"}} + if got := resourceTestCommand(t, a.h, share, "AddService", data); got.Code != 0 { + t.Fatal(got.Msg) + } + body := string(a.commandsOfType("UpdateService")[0].Data) + for _, name := range []string{"one", "two"} { + if !strings.Contains(body, peerShareResourceName(share.ID, "chain", name)) { + t.Fatalf("chainGroup reference was not scoped: %s", body) + } + } + for _, reference := range []string{"resolver", "auther", "observer", "hop"} { + data := resourceTestService("71_1_0", 31002) + data[0].(map[string]interface{})[reference] = "global-resource" + if got := resourceTestCommand(t, a.h, share, "AddService", data); got.Code == 0 { + t.Fatalf("accepted global %s reference", reference) + } + } +} + +func TestPeerResourceReconcileContinuesAfterAnotherShareFails(t *testing.T) { + a := newCleanupAgent(t) + first := resourceTestShare(t, a.h, "resource-failed-share") + second := resourceTestShare(t, a.h, "resource-good-share") + for i, share := range []*repo.PeerShare{first, second} { + if got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("70_1_0", 31001+i)); got.Code != 0 { + t.Fatal(got.Msg) + } + } + if err := a.h.repo.SavePeerShareResources([]repo.PeerShareResource{{ShareID: first.ID, NodeID: 1, Kind: "limiter", OriginalName: "broken", RuntimeName: peerShareResourceName(first.ID, "limiter", "broken"), Config: "{", DesiredState: "active", UpdatedTime: time.Now().UnixMilli()}}); err != nil { + t.Fatal(err) + } + if err := a.h.reconcilePeerShareResourcesOnNode(1); err == nil { + t.Fatal("invalid dependency was not reported") + } + commands := a.commandsOfType("UpdateService") + if len(commands) != 3 || !strings.Contains(string(commands[2].Data), peerShareResourceName(second.ID, "service", "70_1_0")) { + t.Fatalf("failed share blocked healthy share, or failed dependency service was started: %+v", commands) + } +} + +func TestPeerResourceRechecksShareBeforeRecreation(t *testing.T) { + for _, state := range []string{"inactive", "expired", "exceeded"} { + t.Run(state, func(t *testing.T) { + a := newCleanupAgent(t) + share := resourceTestShare(t, a.h, "resource-state-"+state) + switch state { + case "inactive": + share.IsActive = 0 + case "expired": + share.ExpiryTime = time.Now().Add(-time.Hour).UnixMilli() + case "exceeded": + share.MaxBandwidth = 1 + share.CurrentFlow = 1 << 40 + } + if err := a.h.repo.UpdatePeerShare(share); err != nil { + t.Fatal(err) + } + if state == "exceeded" { + if err := a.h.repo.AddPeerShareCurrentFlow(share.ID, share.CurrentFlow); err != nil { + t.Fatal(err) + } + } + if got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("70_1_0", 31001)); got.Code == 0 { + t.Fatal("invalid share recreated resources") + } + if len(a.commandsOfType("UpdateService")) != 0 { + t.Fatal("invalid share reached the node") + } + }) + } +} + +func TestPeerResourcePendingRetryLeavesAppliedSiblingAlone(t *testing.T) { + a := newCleanupAgent(t) + share := resourceTestShare(t, a.h, "resource-pending-only") + offline := &Handler{repo: a.h.repo} + if got := resourceTestCommand(t, offline, share, "AddService", resourceTestService("70_1_0", 31001)); got.Code == 0 { + t.Fatal("offline apply reported success") + } + if got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("71_1_0", 31002)); got.Code != 0 { + t.Fatal(got.Msg) + } + if err := a.h.retryPendingPeerShareResourcesOnNode(1); err != nil { + t.Fatal(err) + } + commands := a.commandsOfType("UpdateService") + if len(commands) != 2 || !strings.Contains(string(commands[1].Data), peerShareResourceName(share.ID, "service", "70_1_0")) { + t.Fatalf("pending retry restarted applied sibling: %+v", commands) + } + if err := a.h.retryPendingPeerShareResourcesOnNode(1); err != nil { + t.Fatal(err) + } + if len(a.commandsOfType("UpdateService")) != 2 { + t.Fatal("no-op pending retry restarted applied services") + } +} + +func TestPeerResourceLegacyReleaseAcknowledgmentIsAtomic(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + if err := a.h.repo.DB().Callback().Update().Before("gorm:update").Register("fail-runtime-release", func(tx *gorm.DB) { + if tx.Statement.Table == "peer_share_runtime" { + tx.AddError(errors.New("simulated completion failure")) + } + }); err != nil { + t.Fatal(err) + } + got := resourceTestCommand(t, a.h, share, "DeleteService", map[string]interface{}{"services": []string{"70_1_0", "70_1_0_tcp", "70_1_0_udp"}}) + if got.Code == 0 { + t.Fatal("completion database failure reported success") + } + items, err := a.h.repo.ListPeerShareResourcesByNode(1) + if err != nil { + t.Fatal(err) + } + proof := false + pending := false + for _, item := range items { + proof = proof || item.LegacyServiceBase == "70_1_0" + pending = pending || item.Applied == 0 + } + if !proof || !pending { + t.Fatalf("failure lost ownership or retry marker: %+v", items) + } + if err := a.h.repo.DB().Callback().Update().Remove("fail-runtime-release"); err != nil { + t.Fatal(err) + } + if err := a.h.retryPendingPeerShareResourcesOnNode(1); err != nil { + t.Fatal(err) + } + runtimes, err := a.h.repo.ListActivePeerShareRuntimesByShareID(share.ID) + if err != nil || len(runtimes) != 0 { + t.Fatalf("retry leaked legacy reservation: %+v %v", runtimes, err) + } +} diff --git a/go-backend/internal/http/handler/federation_share_test.go b/go-backend/internal/http/handler/federation_share_test.go index c0a054a..586fdc2 100644 --- a/go-backend/internal/http/handler/federation_share_test.go +++ b/go-backend/internal/http/handler/federation_share_test.go @@ -224,18 +224,14 @@ func TestFederationShareListIncludesRemoteUsedPorts(t *testing.T) { } func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) { - r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) - if err != nil { - t.Fatalf("open sqlite: %v", err) - } - t.Cleanup(func() { _ = r.Close() }) - - h := New(r, "test-jwt-secret") + agent := newCleanupAgent(t) + r := agent.h.repo + h := agent.h now := time.Now().UnixMilli() if err := r.CreatePeerShare(&repo.PeerShare{ Name: "delete-cleanup-share", - NodeID: 99, + NodeID: 1, Token: "delete-cleanup-token", MaxBandwidth: 4096, PortRangeStart: 40000, @@ -257,8 +253,8 @@ func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) { VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?), (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) `, - share.ID, 99, "dc-r1", "dc-rk1", "dc-b1", "exit", "", "fed_svc_dc1", "tls", "round", 40001, "", 1, 1, now, now, - share.ID, 99, "dc-r2", "dc-rk2", "dc-b2", "middle", "fed_chain_dc2", "fed_svc_dc2", "tls", "round", 40002, "", 1, 1, now, now, + share.ID, 1, "dc-r1", "dc-rk1", "dc-b1", "exit", "", "fed_svc_dc1", "tls", "round", 40001, "", 1, 1, now, now, + share.ID, 1, "dc-r2", "dc-rk2", "dc-b2", "middle", "fed_chain_dc2", "fed_svc_dc2", "tls", "round", 40002, "", 1, 1, now, now, ).Error; err != nil { t.Fatalf("insert peer_share_runtime rows: %v", err) } diff --git a/go-backend/internal/http/handler/flow_cleanup_regression_test.go b/go-backend/internal/http/handler/flow_cleanup_regression_test.go new file mode 100644 index 0000000..5513100 --- /dev/null +++ b/go-backend/internal/http/handler/flow_cleanup_regression_test.go @@ -0,0 +1,300 @@ +package handler + +import ( + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "path/filepath" + "reflect" + "sort" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/gorilla/websocket" + "gorm.io/gorm" + + "go-backend/internal/security" + "go-backend/internal/store/model" + "go-backend/internal/store/repo" + "go-backend/internal/ws" +) + +type cleanupAgentCommand struct { + Type string `json:"type"` + Data json.RawMessage `json:"data"` + RequestID string `json:"requestId"` +} + +// cleanupAgent exercises the actual command transport. Each command is recorded +// before its ACK, so synchronous handler calls need no sleeps to inspect it. +type cleanupAgent struct { + h *Handler + mu sync.Mutex + commands []cleanupAgentCommand +} + +func newCleanupAgent(t *testing.T) *cleanupAgent { + t.Helper() + r, err := repo.Open(filepath.Join(t.TempDir(), "cleanup.db")) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = r.Close() }) + const secret = "cleanup-regression-node" + if err := r.DB().Create(&model.Node{ID: 1, Name: "relay", Secret: secret, ServerIP: "127.0.0.1", Port: "30000", Status: 1}).Error; err != nil { + t.Fatal(err) + } + server := ws.NewServer(r, "test-jwt-secret") + online := make(chan struct{}) + server.SetNodeOnlineHook(func(int64) { close(online) }) + serverDone := make(chan struct{}) + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + defer close(serverDone) + server.ServeHTTP(w, req) + })) + t.Cleanup(ts.Close) + conn, _, err := websocket.DefaultDialer.Dial("ws"+strings.TrimPrefix(ts.URL, "http")+"?type=1&secret="+secret, nil) + if err != nil { + t.Fatal(err) + } + crypto, err := security.NewAESCrypto(secret) + if err != nil { + t.Fatal(err) + } + a := &cleanupAgent{h: &Handler{repo: r, wsServer: server}} + readerDone := make(chan struct{}) + go func() { + defer close(readerDone) + for { + _, payload, err := conn.ReadMessage() + if err != nil { + return + } + var envelope struct { + Encrypted bool `json:"encrypted"` + Data string `json:"data"` + } + if err := json.Unmarshal(payload, &envelope); err != nil { + t.Errorf("decode command envelope: %v", err) + return + } + if envelope.Encrypted { + payload, err = crypto.Decrypt(envelope.Data) + if err != nil { + t.Errorf("decrypt command: %v", err) + return + } + } + var command cleanupAgentCommand + if err := json.Unmarshal(payload, &command); err != nil { + t.Errorf("decode command: %v", err) + return + } + a.mu.Lock() + a.commands = append(a.commands, command) + a.mu.Unlock() + if err := conn.WriteJSON(map[string]interface{}{"type": command.Type, "requestId": command.RequestID, "success": true}); err != nil { + return + } + } + }() + t.Cleanup(func() { + _ = conn.Close() + <-readerDone + <-serverDone + }) + select { + case <-online: + case <-time.After(3 * time.Second): + t.Fatal("mock node did not come online") + } + a.probe(t) + return a +} + +func (a *cleanupAgent) probe(t *testing.T) { + t.Helper() + if _, err := a.h.wsServer.SendCommand(1, "CleanupTestProbe", nil, time.Second); err != nil { + t.Fatalf("mock agent transport unavailable: %v", err) + } +} + +func (a *cleanupAgent) commandsOfType(commandType string) []cleanupAgentCommand { + a.mu.Lock() + defer a.mu.Unlock() + var commands []cleanupAgentCommand + for _, command := range a.commands { + if command.Type == commandType { + commands = append(commands, command) + } + } + return commands +} + +func (a *cleanupAgent) addRuntime(t *testing.T, name string, nodeID int64, status, applied int, updated time.Time) int64 { + t.Helper() + now := time.Now().UnixMilli() + share := &repo.PeerShare{Name: "shared", NodeID: nodeID, Token: "cleanup-share-token", PortRangeStart: 31000, PortRangeEnd: 31010, IsActive: 1, CreatedTime: now, UpdatedTime: now} + if err := a.h.repo.CreatePeerShare(share); err != nil { + t.Fatal(err) + } + stored, err := a.h.repo.GetPeerShareByToken(share.Token) + if err != nil || stored == nil { + t.Fatalf("load share: %v", err) + } + // SQL preserves status=0; GORM's default tag would replace that with 1. + if err := a.h.repo.DB().Exec(`INSERT INTO peer_share_runtime + (share_id, node_id, reservation_id, resource_key, role, service_name, applied, status, created_time, updated_time) + VALUES (?, ?, 'cleanup-reservation', 'cleanup-resource', 'forward', ?, ?, ?, ?, ?)`, + stored.ID, nodeID, name, applied, status, now, updated.UnixMilli()).Error; err != nil { + t.Fatal(err) + } + return stored.ID +} + +func TestSharedForwardCleanupRegression(t *testing.T) { + for _, mode := range []string{"single", "batch", "config"} { + for _, runtimeName := range []string{"70_1_0", "70_1_0_tcp", "70_1_0_udp"} { + t.Run(mode+"/active/"+runtimeName, func(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, runtimeName, 1, 1, 1, time.Now()) + runCleanupPath(a.h, mode, []string{"70_1_0", "70_1_0_tcp", "70_1_0_udp"}) + a.probe(t) + if commands := a.commandsOfType("DeleteService"); len(commands) != 0 { + t.Fatalf("active shared service family was deleted: %+v", commands) + } + if mode != "config" && runtimeName == "70_1_0" { + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil || share == nil || share.CurrentFlow != 600 { + t.Fatalf("shared traffic should accumulate 600 bytes: share=%+v err=%v", share, err) + } + } + }) + } + for _, tc := range []struct { + name string + serviceName string + nodeID int64 + status int + applied int + age time.Duration + wantDelete bool + }{ + {name: "recent-unbound", nodeID: 1, status: 1, age: time.Minute}, + {name: "stale-unbound", nodeID: 1, status: 1, age: 11 * time.Minute, wantDelete: true}, + {name: "released", serviceName: "70_1_0", nodeID: 1, applied: 1, wantDelete: true}, + {name: "other-node-bound", serviceName: "70_1_0", nodeID: 2, status: 1, applied: 1, wantDelete: true}, + {name: "other-node-unbound", nodeID: 2, status: 1, wantDelete: true}, + {name: "unrelated-active-family", serviceName: "71_1_0", nodeID: 1, status: 1, applied: 1, wantDelete: true}, + } { + t.Run(mode+"/"+tc.name, func(t *testing.T) { + a := newCleanupAgent(t) + a.addRuntime(t, tc.serviceName, tc.nodeID, tc.status, tc.applied, time.Now().Add(-tc.age)) + runCleanupPath(a.h, mode, []string{"70_1_0_tcp"}) + a.probe(t) + commands := a.commandsOfType("DeleteService") + if !tc.wantDelete { + if len(commands) != 0 { + t.Fatalf("recent unbound shared runtime was deleted: %+v", commands) + } + return + } + if len(commands) != 1 { + t.Fatalf("expected one orphan cleanup command, got %+v", commands) + } + var data struct { + Services []string `json:"services"` + } + if err := json.Unmarshal(commands[0].Data, &data); err != nil { + t.Fatal(err) + } + sort.Strings(data.Services) + if want := []string{"70_1_0", "70_1_0_tcp", "70_1_0_udp"}; !reflect.DeepEqual(data.Services, want) { + t.Fatalf("orphan cleanup must delete complete family: got %v, want %v", data.Services, want) + } + }) + } + } +} + +func runCleanupPath(h *Handler, mode string, names []string) { + items := make([]flowItem, 0, len(names)) + services := make([]namedConfigItem, 0, len(names)) + for _, name := range names { + items = append(items, flowItem{N: name, U: 120, D: 80}) + services = append(services, namedConfigItem{Name: name}) + } + switch mode { + case "single": + for _, item := range items { + h.processFlowItem(1, item) + } + case "batch": + h.applyFlowUploadBatch(1, h.buildNodeFlowUploadBatch(1, items, nil), time.Now()) + case "config": + h.cleanOrphanedServices(1, services) + } +} + +func TestSharedForwardCleanupMissingMetadataPreservesLocalForward(t *testing.T) { + a := newCleanupAgent(t) + if err := a.h.repo.DB().Create(&model.Forward{ID: 70, UserID: 1, UserName: "local", Name: "local", TunnelID: 1, RemoteAddr: "127.0.0.1:80", Status: 1}).Error; err != nil { + t.Fatal(err) + } + runCleanupPath(a.h, "batch", []string{"70_1_0", "70_1_0_tcp", "70_1_0_udp"}) + a.probe(t) + if commands := a.commandsOfType("DeleteService"); len(commands) != 0 { + t.Fatalf("missing batch metadata must not delete a local forward: %+v", commands) + } +} + +func TestSharedForwardCleanupChecksActualDeleteFamily(t *testing.T) { + for _, mode := range []string{"single", "batch", "config"} { + t.Run(mode, func(t *testing.T) { + a := newCleanupAgent(t) + a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + // Legacy parsing accepts additional suffixes. Cleanup must check + // the family it would delete, not just the reported service name. + runCleanupPath(a.h, mode, []string{"70_1_0_old", "70_1_0_old_tcp"}) + a.probe(t) + if commands := a.commandsOfType("DeleteService"); len(commands) != 0 { + t.Fatalf("suffix variation must not delete a protected shared family: %+v", commands) + } + }) + } +} + +func TestSharedForwardCleanupQueryFailurePreservesServices(t *testing.T) { + for _, mode := range []string{"single", "batch", "config"} { + for _, failure := range []string{"forward", "peer_share_runtime"} { + t.Run(mode+"/"+failure, func(t *testing.T) { + a := newCleanupAgent(t) + // Fail only the ownership lookup; leave node lookup and WebSocket + // delivery functional so an erroneous delete remains observable. + callback := "test:cleanup-query-failure" + var injected atomic.Int32 + if err := a.h.repo.DB().Callback().Query().Before("gorm:query").Register(callback, func(tx *gorm.DB) { + if tx.Statement.Table == failure { + injected.Add(1) + tx.AddError(errors.New("injected ownership query failure")) + } + }); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = a.h.repo.DB().Callback().Query().Remove(callback) }) + runCleanupPath(a.h, mode, []string{"70_1_0_tcp"}) + a.probe(t) + if injected.Load() == 0 { + t.Fatal("test did not inject the expected query failure") + } + if commands := a.commandsOfType("DeleteService"); len(commands) != 0 { + t.Fatalf("ownership query failure must preserve services: %+v", commands) + } + }) + } + } +} diff --git a/go-backend/internal/http/handler/flow_config_cleanup_test.go b/go-backend/internal/http/handler/flow_config_cleanup_test.go new file mode 100644 index 0000000..25bccdc --- /dev/null +++ b/go-backend/internal/http/handler/flow_config_cleanup_test.go @@ -0,0 +1,145 @@ +package handler + +import ( + "encoding/json" + "errors" + "testing" + "time" + + "gorm.io/gorm" +) + +func TestConfigCleanupPreservesSharedDependencies(t *testing.T) { + a := newCleanupAgent(t) + a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + a.h.cleanNodeConfigs(1, `{ + "services": [{"name":"70_1_0_tcp", "handler":{"chain":"chains_88"}, "limiter":"13, rule_traffic_limit_70"}], + "chains": [{"name":"fed_chain_17"}, {"name":"chains_88"}, {"name":"chains_999"}], + "limiters": [{"name":"13"}, {"name":"rule_traffic_limit_70"}, {"name":"99"}] + }`) + a.probe(t) + if commands := a.commandsOfType("DeleteService"); len(commands) != 0 { + t.Fatalf("shared service was deleted: %+v", commands) + } + assertCleanupDependency(t, a, "DeleteChains", "chain", "chains_999") + assertCleanupDependency(t, a, "DeleteLimiters", "limiter", "99") +} + +func TestConfigCleanupRemovesOrphanedForwardLimiter(t *testing.T) { + a := newCleanupAgent(t) + a.h.cleanNodeConfigs(1, `{"limiters":[{"name":"rule_traffic_limit_70"}]}`) + a.probe(t) + assertCleanupDependency(t, a, "DeleteLimiters", "limiter", "rule_traffic_limit_70") +} + +func TestConfigCleanupProtectsPendingSharedDependencies(t *testing.T) { + for _, tc := range []struct { + name string + age time.Duration + keep bool + }{ + {name: "pending", age: time.Minute, keep: true}, + {name: "expired-reservation", age: 11 * time.Minute}, + } { + t.Run(tc.name, func(t *testing.T) { + a := newCleanupAgent(t) + a.addRuntime(t, "", 1, 1, 0, time.Now().Add(-tc.age)) + // Dependencies may arrive before the service and its runtime binding. + a.h.cleanNodeConfigs(1, `{"chains":[{"name":"chains_88"}],"limiters":[{"name":"13"}]}`) + a.probe(t) + if tc.keep { + for _, commandType := range []string{"DeleteChains", "DeleteLimiters"} { + if commands := a.commandsOfType(commandType); len(commands) != 0 { + t.Fatalf("pending shared dependencies deleted: %+v", commands) + } + } + return + } + assertCleanupDependency(t, a, "DeleteChains", "chain", "chains_88") + assertCleanupDependency(t, a, "DeleteLimiters", "limiter", "13") + }) + } +} + +func TestConfigCleanupPreservesDependenciesOnLookupFailure(t *testing.T) { + for _, table := range []string{"peer_share_runtime", "tunnel", "speed_limit", "forward"} { + t.Run(table, func(t *testing.T) { + a := newCleanupAgent(t) + callback := "test:config-cleanup-query-failure" + injected := false + if err := a.h.repo.DB().Callback().Query().Before("gorm:query").Register(callback, func(tx *gorm.DB) { + if tx.Statement.Table == table { + injected = true + tx.AddError(errors.New("injected dependency lookup failure")) + } + }); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = a.h.repo.DB().Callback().Query().Remove(callback) }) + configs := map[string]string{ + "peer_share_runtime": `{"chains":[{"name":"chains_88"}],"limiters":[{"name":"13"}]}`, + "tunnel": `{"chains":[{"name":"chains_88"}]}`, + "speed_limit": `{"limiters":[{"name":"13"}]}`, + "forward": `{"limiters":[{"name":"rule_traffic_limit_70"}]}`, + } + a.h.cleanNodeConfigs(1, configs[table]) + a.probe(t) + if !injected { + t.Fatal("expected dependency lookup failure to be injected") + } + for _, commandType := range []string{"DeleteChains", "DeleteLimiters"} { + if commands := a.commandsOfType(commandType); len(commands) != 0 { + t.Fatalf("lookup failure must not authorize cleanup: %+v", commands) + } + } + }) + } +} + +func TestSharedForwardCleanupDuringRuntimeBinding(t *testing.T) { + for _, mode := range []string{"single", "batch", "config"} { + t.Run(mode, func(t *testing.T) { + a := newCleanupAgent(t) + a.addRuntime(t, "", 1, 1, 0, time.Now()) + callback := "test:bind-during-cleanup" + bound := false + if err := a.h.repo.DB().Callback().Query().After("gorm:query").Register(callback, func(tx *gorm.DB) { + if tx.Statement.Table != "peer_share_runtime" || bound { + return + } + bound = true + // Reproduce binding immediately after the first ownership read. + // Separate name/unbound queries would both miss this runtime. + if err := a.h.repo.DB().Exec("UPDATE peer_share_runtime SET service_name = ?, applied = 1", "70_1_0").Error; err != nil { + t.Errorf("bind shared runtime: %v", err) + } + }); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = a.h.repo.DB().Callback().Query().Remove(callback) }) + runCleanupPath(a.h, mode, []string{"70_1_0_tcp"}) + a.probe(t) + if !bound { + t.Fatal("binding transition did not run") + } + if commands := a.commandsOfType("DeleteService"); len(commands) != 0 { + t.Fatalf("service was deleted during binding: %+v", commands) + } + }) + } +} + +func assertCleanupDependency(t *testing.T, a *cleanupAgent, commandType, key, want string) { + t.Helper() + commands := a.commandsOfType(commandType) + if len(commands) != 1 { + t.Fatalf("expected one %s for %s, got %+v", commandType, want, commands) + } + var data map[string]string + if err := json.Unmarshal(commands[0].Data, &data); err != nil { + t.Fatal(err) + } + if data[key] != want { + t.Fatalf("unexpected %s target: got %q, want %q", commandType, data[key], want) + } +} diff --git a/go-backend/internal/http/handler/flow_ownership.go b/go-backend/internal/http/handler/flow_ownership.go new file mode 100644 index 0000000..d494376 --- /dev/null +++ b/go-backend/internal/http/handler/flow_ownership.go @@ -0,0 +1,118 @@ +package handler + +import ( + "log" + "strings" + + "go-backend/internal/store/repo" +) + +// Classify ownership before building local counters or enforcing local quotas. +// Numeric IDs from a consuming panel are not IDs in the provider's database. +func (h *Handler) buildNodeFlowUploadBatch(nodeID int64, items []flowItem, metas map[int64]repo.FlowUploadForwardMeta) flowUploadBatch { + if h == nil || h.repo == nil || nodeID <= 0 { + return flowUploadBatch{} + } + runtimes, err := h.repo.ListActiveForwardPeerShareRuntimesByNode(nodeID) + if err != nil { + log.Printf("flow ownership lookup failed node_id=%d err=%v", nodeID, err) + return flowUploadBatch{} + } + resources, err := h.repo.ListPeerShareResourcesByNode(nodeID) + if err != nil { + log.Printf("flow resource lookup failed node_id=%d err=%v", nodeID, err) + return flowUploadBatch{} + } + forwardIDs, err := h.repo.ListForwardIDsByNode(nodeID) + if err != nil { + log.Printf("flow node ownership lookup failed node_id=%d err=%v", nodeID, err) + return flowUploadBatch{} + } + nodeForwards := make(map[int64]struct{}, len(forwardIDs)) + for _, forwardID := range forwardIDs { + nodeForwards[forwardID] = struct{}{} + } + legacyOwners := make(map[string]map[int64]struct{}) + addLegacyOwner := func(name string, shareID int64) { + name = normalizeForwardRuntimeServiceName(name) + if name == "" { + return + } + if legacyOwners[name] == nil { + legacyOwners[name] = make(map[int64]struct{}) + } + legacyOwners[name][shareID] = struct{}{} + } + for _, runtime := range runtimes { + addLegacyOwner(runtime.ServiceName, runtime.ShareID) + } + resourceOwners := make(map[string]int64) + for _, resource := range resources { + // A migrated TCP service may still own a legacy UDP sibling. This alias + // remains authoritative until the whole legacy family is acknowledged gone. + addLegacyOwner(resource.LegacyServiceBase, resource.ShareID) + // A deletion intent does not prove the listener is gone. Failed or + // timed-out commands retain ownership until the node acknowledges it. + if resource.Kind == "service" && (resource.DesiredState != "deleted" || resource.Applied == 0) { + resourceOwners[resource.RuntimeName] = resource.ShareID + } + } + localItems := make([]flowItem, 0, len(items)) + sharedUsage := make(map[int64]int64) + for _, item := range items { + name := strings.TrimSpace(item.N) + if strings.HasPrefix(name, "peer-share-") { + shareID, _, ok := parsePeerShareServiceName(name) + if ok && resourceOwners[name] == shareID && item.U >= 0 && item.D >= 0 { + sharedUsage[shareID] += item.U + item.D + } + // Unknown or confirmed-deleted scoped names must never become local IDs. + continue + } + if owners := legacyOwners[normalizeForwardRuntimeServiceName(name)]; len(owners) > 0 { + forwardID, userID, userTunnelID, parsed := parseFlowServiceIDs(name) + meta, local := metas[forwardID] + _, onNode := nodeForwards[forwardID] + if parsed && local && onNode && meta.UserID == userID && meta.UserTunnelID == userTunnelID { + // Legacy names can be genuinely ambiguous. Preserve the listener + // but do not debit either owner based on an ID guess. + log.Printf("ambiguous legacy flow ownership node_id=%d service=%s", nodeID, name) + continue + } + if len(owners) == 1 && item.U >= 0 && item.D >= 0 { + for shareID := range owners { + sharedUsage[shareID] += item.U + item.D + } + } + continue + } + if forwardID, userID, _, ok := parseFlowServiceIDs(name); ok { + if meta, exists := metas[forwardID]; exists { + if _, onNode := nodeForwards[forwardID]; !onNode || meta.UserID != userID { + continue + } + } + } + localItems = append(localItems, item) + } + batch := h.buildFlowUploadBatch(localItems, metas) + batch.peerShareUsage = sharedUsage + return batch +} + +func (h *Handler) addPeerShareFlow(nodeID, shareID, delta int64) { + if nodeID <= 0 || shareID <= 0 || delta <= 0 { + return + } + share, err := h.repo.GetPeerShare(shareID) + if err != nil || share == nil || share.NodeID != nodeID { + return + } + if err := h.repo.AddPeerShareCurrentFlow(shareID, delta); err != nil { + return + } + share, err = h.repo.GetPeerShare(shareID) + if err == nil && isPeerShareFlowExceeded(share) { + h.enforcePeerShareFlowLimit(shareID) + } +} diff --git a/go-backend/internal/http/handler/flow_ownership_test.go b/go-backend/internal/http/handler/flow_ownership_test.go new file mode 100644 index 0000000..92db232 --- /dev/null +++ b/go-backend/internal/http/handler/flow_ownership_test.go @@ -0,0 +1,235 @@ +package handler + +import ( + "go-backend/internal/store/model" + "go-backend/internal/store/repo" + "testing" + "time" +) + +func TestFlowOwnershipSharedFlowDoesNotChargeLocalCollision(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + if err := a.h.repo.DB().Create(&model.Tunnel{ID: 9, Name: "local", TrafficRatio: 1, Flow: 1, Type: 1, Status: 1}).Error; err != nil { + t.Fatal(err) + } + if err := a.h.repo.DB().Create(&model.Forward{ID: 70, UserID: 2, UserName: "local", Name: "local", TunnelID: 9, RemoteAddr: "127.0.0.1:80", Status: 1}).Error; err != nil { + t.Fatal(err) + } + metas, err := a.h.repo.GetFlowUploadForwardMetas([]int64{70}) + if err != nil { + t.Fatal(err) + } + a.h.applyFlowUploadBatch(1, a.h.buildNodeFlowUploadBatch(1, []flowItem{{N: "70_1_0_tcp", U: 120, D: 80}}, metas), time.Now()) + var local model.Forward + if err := a.h.repo.DB().First(&local, 70).Error; err != nil { + t.Fatal(err) + } + shared, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + t.Logf("local user=%d local bytes=%d shared bytes=%d", local.UserID, local.InFlow+local.OutFlow, shared.CurrentFlow) + if local.InFlow+local.OutFlow != 0 { + t.Errorf("shared flow charged colliding local forward: %d", local.InFlow+local.OutFlow) + } + if shared.CurrentFlow != 200 { + t.Errorf("shared flow missing: %d", shared.CurrentFlow) + } +} + +func TestFlowOwnershipScopedServiceRequiresStoredNodeOwnership(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + name := peerShareResourceName(shareID, "service", "70_1_0_tcp") + if err := a.h.repo.SavePeerShareResources([]repo.PeerShareResource{{ + ShareID: shareID, NodeID: 1, Kind: "service", OriginalName: "70_1_0_tcp", + RuntimeName: name, DesiredState: "active", Applied: 1, + }}); err != nil { + t.Fatal(err) + } + for _, nodeID := range []int64{2, 1} { + batch := a.h.buildNodeFlowUploadBatch(nodeID, []flowItem{{N: name, U: 120, D: 80}}, nil) + a.h.applyFlowUploadBatch(nodeID, batch, time.Now()) + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + want := int64(0) + if nodeID == 1 { + want = 200 + } + if share.CurrentFlow != want { + t.Fatalf("node=%d: got flow=%d want=%d", nodeID, share.CurrentFlow, want) + } + if len(batch.flowDeltas) != 0 || len(batch.quotaUsage) != 0 || len(batch.orphanServices) != 0 { + t.Fatalf("scoped traffic entered local accounting: %+v", batch) + } + } + a.probe(t) + if cmds := a.commandsOfType("DeleteService"); len(cmds) != 0 { + t.Fatalf("scoped service deleted: %+v", cmds) + } +} + +func TestFlowOwnershipRoleRuntimeRequiresReportingNode(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "fed_svc_17", 1, 1, 1, time.Now()) + if err := a.h.repo.DB().Exec("UPDATE peer_share_runtime SET id = 17, role = 'exit'").Error; err != nil { + t.Fatal(err) + } + a.h.processFlowItem(2, flowItem{N: "fed_svc_17", U: 120, D: 80}) + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + if share.CurrentFlow != 0 { + t.Fatalf("wrong node charged role runtime: %d", share.CurrentFlow) + } + a.h.processFlowItem(1, flowItem{N: "fed_svc_17", U: 120, D: 80}) + share, err = a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + if share.CurrentFlow != 200 { + t.Fatalf("owner node flow missing: %d", share.CurrentFlow) + } +} + +func TestFlowOwnershipAmbiguousLegacyNameDoesNotDebitEitherOwner(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + if err := a.h.repo.DB().Create(&model.Forward{ID: 70, UserID: 1, UserName: "local", Name: "local", TunnelID: 9, RemoteAddr: "127.0.0.1:80", Status: 1}).Error; err != nil { + t.Fatal(err) + } + if err := a.h.repo.DB().Create(&model.ForwardPort{ForwardID: 70, NodeID: 1, Port: 32000}).Error; err != nil { + t.Fatal(err) + } + metas, err := a.h.repo.GetFlowUploadForwardMetas([]int64{70}) + if err != nil { + t.Fatal(err) + } + batch := a.h.buildNodeFlowUploadBatch(1, []flowItem{{N: "70_1_0_tcp", U: 120, D: 80}}, metas) + a.h.applyFlowUploadBatch(1, batch, time.Now()) + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + if share.CurrentFlow != 0 || len(batch.flowDeltas) != 0 { + t.Fatalf("ambiguous flow charged an owner: share=%d local=%+v", share.CurrentFlow, batch.flowDeltas) + } + a.probe(t) + if commands := a.commandsOfType("DeleteService"); len(commands) != 0 { + t.Fatalf("ambiguous legacy service deleted: %+v", commands) + } +} + +func TestFlowOwnershipPreservesUnmigratedLegacyTransport(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, peerShareResourceName(1, "service", "70_1_0"), 1, 1, 1, time.Now()) + if err := a.h.repo.SavePeerShareResources([]repo.PeerShareResource{{ + ShareID: shareID, NodeID: 1, Kind: "service", OriginalName: "70_1_0_tcp", + RuntimeName: peerShareResourceName(shareID, "service", "70_1_0_tcp"), + LegacyServiceBase: "70_1_0", DesiredState: "deleted", Applied: 1, + }}); err != nil { + t.Fatal(err) + } + a.h.processFlowItem(1, flowItem{N: "70_1_0_udp", U: 120, D: 80}) + a.h.cleanNodeConfigs(1, `{"services":[{"name":"70_1_0_udp"}]}`) + a.probe(t) + if commands := a.commandsOfType("DeleteService"); len(commands) != 0 { + t.Fatalf("unmigrated transport deleted: %+v", commands) + } + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + if share.CurrentFlow != 200 { + t.Fatalf("legacy transport accounting lost: %d", share.CurrentFlow) + } +} + +func TestFlowOwnershipLocalForwardRequiresReportingNode(t *testing.T) { + a := newCleanupAgent(t) + if err := a.h.repo.DB().Create(&model.Forward{ID: 70, UserID: 1, UserName: "local", Name: "local", TunnelID: 9, RemoteAddr: "127.0.0.1:80", Status: 1}).Error; err != nil { + t.Fatal(err) + } + if err := a.h.repo.DB().Create(&model.ForwardPort{ForwardID: 70, NodeID: 2, Port: 32000}).Error; err != nil { + t.Fatal(err) + } + metas, err := a.h.repo.GetFlowUploadForwardMetas([]int64{70}) + if err != nil { + t.Fatal(err) + } + for _, nodeID := range []int64{1, 2} { + batch := a.h.buildNodeFlowUploadBatch(nodeID, []flowItem{{N: "70_1_0_tcp", U: 120, D: 80}}, metas) + a.h.applyFlowUploadBatch(nodeID, batch, time.Now()) + var forward model.Forward + if err := a.h.repo.DB().First(&forward, 70).Error; err != nil { + t.Fatal(err) + } + want := int64(0) + if nodeID == 2 { + want = 200 + } + if forward.InFlow+forward.OutFlow != want { + t.Fatalf("node %d local traffic=%d want=%d", nodeID, forward.InFlow+forward.OutFlow, want) + } + } +} + +func TestPeerShareMaintenanceRetriesOnlyUnfinishedOperations(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + name := peerShareResourceName(shareID, "service", "80_1_0_tcp") + if err := a.h.repo.SavePeerShareResources([]repo.PeerShareResource{{ShareID: shareID, NodeID: 1, + Kind: "service", OriginalName: "80_1_0_tcp", RuntimeName: name, DesiredState: "deleted", Applied: 0, + }}); err != nil { + t.Fatal(err) + } + a.h.retryPendingPeerShareOperations() + a.probe(t) + if commands := a.commandsOfType("DeleteService"); len(commands) != 1 { + t.Fatalf("pending delete was not retried: %+v", commands) + } + a.h.retryPendingPeerShareOperations() + a.probe(t) + if commands := a.commandsOfType("DeleteService"); len(commands) != 1 { + t.Fatalf("completed delete was replayed: %+v", commands) + } + ids, err := a.h.repo.ListPendingPeerShareNodeIDs() + if err != nil { + t.Fatal(err) + } + if len(ids) != 0 { + t.Fatalf("completed operation still pending: %v", ids) + } +} + +func TestFlowOwnershipDoesNotCrossNodes(t *testing.T) { + a := newCleanupAgent(t) + shareID := a.addRuntime(t, "70_1_0", 1, 1, 1, time.Now()) + share, err := a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + share.MaxBandwidth = 100 + if err := a.h.repo.UpdatePeerShare(share); err != nil { + t.Fatal(err) + } + // Node 2 has no shared runtime. It reports the same local numeric service name. + a.h.applyFlowUploadBatch(2, a.h.buildNodeFlowUploadBatch(2, []flowItem{{N: "70_1_0_tcp", U: 120, D: 80}}, nil), time.Now()) + a.probe(t) + share, err = a.h.repo.GetPeerShare(shareID) + if err != nil { + t.Fatal(err) + } + commands := a.commandsOfType("DeleteService") + t.Logf("node1 share charged from node2: %d; node1 delete commands: %+v", share.CurrentFlow, commands) + if share.CurrentFlow != 0 { + t.Errorf("flow crossed node boundary: %d", share.CurrentFlow) + } + if len(commands) != 0 { + t.Errorf("node2 flow deleted node1 shared service: %+v", commands) + } +} diff --git a/go-backend/internal/http/handler/flow_pending_deletion_test.go b/go-backend/internal/http/handler/flow_pending_deletion_test.go new file mode 100644 index 0000000..d578ba7 --- /dev/null +++ b/go-backend/internal/http/handler/flow_pending_deletion_test.go @@ -0,0 +1,57 @@ +package handler + +import ( + "testing" + "time" +) + +func TestFlowOwnershipPendingDeletionRemainsBillableUntilAcknowledged(t *testing.T) { + for _, mode := range []string{"single", "batch"} { + t.Run(mode, func(t *testing.T) { + a := newCleanupAgent(t) + share := resourceTestShare(t, a.h, "pending-deletion-"+mode) + if got := resourceTestCommand(t, a.h, share, "AddService", resourceTestService("70_1_0_tcp", 31001)); got.Code != 0 { + t.Fatal(got.Msg) + } + name := peerShareResourceName(share.ID, "service", "70_1_0_tcp") + // Simulate a failed delivery after the deletion intent is committed. The + // mock node's previously installed listener has received no delete command. + disconnected := &Handler{repo: a.h.repo} + if got := resourceTestCommand(t, disconnected, share, "DeleteService", map[string]interface{}{"services": []string{"70_1_0_tcp"}}); got.Code == 0 { + t.Fatal("failed deletion reported success") + } + resource, err := a.h.repo.GetPeerShareResource(share.ID, "service", "70_1_0_tcp") + if err != nil || resource == nil || resource.DesiredState != "deleted" || resource.Applied != 0 { + t.Fatalf("missing pending tombstone: %+v %v", resource, err) + } + report := func() { + item := flowItem{N: name, U: 120, D: 80} + if mode == "single" { + a.h.processFlowItem(1, item) + } else { + a.h.applyFlowUploadBatch(1, a.h.buildNodeFlowUploadBatch(1, []flowItem{item}, nil), time.Now()) + } + } + report() + stored, err := a.h.repo.GetPeerShare(share.ID) + if err != nil || stored.CurrentFlow != 200 { + t.Fatalf("unconfirmed deletion stopped billing: share=%+v err=%v", stored, err) + } + if len(a.commandsOfType("DeleteService")) != 0 { + t.Fatal("flow handling deleted a pending shared listener") + } + if err := a.h.retryPendingPeerShareResourcesOnNode(1); err != nil { + t.Fatal(err) + } + resource, err = a.h.repo.GetPeerShareResource(share.ID, "service", "70_1_0_tcp") + if err != nil || resource.Applied != 1 { + t.Fatalf("deletion acknowledgment not recorded: %+v %v", resource, err) + } + report() + stored, err = a.h.repo.GetPeerShare(share.ID) + if err != nil || stored.CurrentFlow != 200 { + t.Fatalf("confirmed-deleted listener was billed again: share=%+v err=%v", stored, err) + } + }) + } +} diff --git a/go-backend/internal/http/handler/flow_policy.go b/go-backend/internal/http/handler/flow_policy.go index cca261f..d85b67f 100644 --- a/go-backend/internal/http/handler/flow_policy.go +++ b/go-backend/internal/http/handler/flow_policy.go @@ -8,8 +8,6 @@ import ( "strconv" "strings" "time" - - "go-backend/internal/store/model" ) const bytesPerGB int64 = 1024 * 1024 * 1024 @@ -48,38 +46,22 @@ type gostConfigSnapshot struct { } type namedConfigItem struct { - Name string `json:"name"` + Name string `json:"name"` + Limiter string `json:"limiter,omitempty"` + Handler *struct { + Chain string `json:"chain"` + } `json:"handler,omitempty"` } func (h *Handler) processFlowItem(nodeID int64, item flowItem) { - serviceName := strings.TrimSpace(item.N) - if serviceName == "" || serviceName == "web_api" { + if h == nil || h.repo == nil || nodeID <= 0 { return } - - forwardID, userID, userTunnelID, ok := parseFlowServiceIDs(serviceName) - if ok { - if h.forwardExists(forwardID) { - inFlow, outFlow := h.scaleFlowByTunnel(forwardID, item.D, item.U) - _ = h.repo.AddFlow(forwardID, userID, userTunnelID, inFlow, outFlow) - if quota, quotaErr := h.repo.AddUserQuotaUsage(userID, inFlow+outFlow, time.Now()); quotaErr == nil { - h.enforceUserQuotaIfNeeded(userID, quota) - } - if userTunnelID > 0 { - h.enforceFlowPolicies(userID, userTunnelID) - } - } else if nodeID > 0 { - h.sendDeleteOrphanedForwardService(nodeID, serviceName) - } - h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item) - return + metas, err := h.repo.GetFlowUploadForwardMetas(collectFlowUploadForwardIDs([]flowItem{item})) + if err != nil { + metas = nil } - - runtimeID, ok := parsePeerShareRuntimeServiceID(serviceName) - if !ok { - return - } - h.processPeerShareFlow(runtimeID, item) + h.applyFlowUploadBatch(nodeID, h.buildNodeFlowUploadBatch(nodeID, []flowItem{item}, metas), time.Now()) } func parseFlowServiceIDs(serviceName string) (int64, int64, int64, bool) { @@ -154,73 +136,49 @@ func parsePeerShareIDFromFederationTunnelName(tunnelName string) (int64, bool) { return shareID, true } -func (h *Handler) processPeerShareFlow(runtimeID int64, item flowItem) { - if h == nil || h.repo == nil || runtimeID <= 0 { +func (h *Handler) processPeerShareFlow(nodeID, runtimeID int64, item flowItem) { + if h == nil || h.repo == nil || nodeID <= 0 || runtimeID <= 0 { return } runtime, err := h.repo.GetPeerShareRuntimeByID(runtimeID) - if err != nil || runtime == nil || runtime.ShareID <= 0 || runtime.Status != 1 { + if err != nil || runtime == nil || runtime.NodeID != nodeID || runtime.Status != 1 { return } - - delta := item.D + item.U - if delta <= 0 { - return - } - - _ = h.repo.AddPeerShareCurrentFlow(runtime.ShareID, delta) - - share, err := h.repo.GetPeerShare(runtime.ShareID) - if err != nil || share == nil { - return - } - if !isPeerShareFlowExceeded(share) { - return - } - h.enforcePeerShareFlowLimit(share.ID) + h.addPeerShareFlow(nodeID, runtime.ShareID, item.D+item.U) } func (h *Handler) processPeerShareFlowFromForward(forwardID int64, nodeID int64, serviceName string, item flowItem) { - if h == nil || h.repo == nil || forwardID <= 0 { + if h == nil || h.repo == nil || forwardID <= 0 || nodeID <= 0 { return } - - delta := item.D + item.U - if delta <= 0 { + // Prefer the reporting node's explicit shared ownership over a coincidentally + // equal local forward ID. Never fall back to a service on another node. + runtimes, err := h.repo.ListActiveForwardPeerShareRuntimesByNode(nodeID) + if err != nil { return } - + for _, runtime := range runtimes { + if normalizeForwardRuntimeServiceName(runtime.ServiceName) == normalizeForwardRuntimeServiceName(serviceName) { + h.processPeerShareFlowByServiceName(nodeID, serviceName, item) + return + } + } forward, err := h.getForwardRecord(forwardID) if err != nil || forward == nil { - // Forward not found in local database - might be a federation port-forward - // Try to find by service name in peer_share_runtime - h.processPeerShareFlowByServiceName(nodeID, serviceName, item) + return + } + _, userID, _, ok := parseFlowServiceIDs(serviceName) + if !ok || userID != forward.UserID { return } tunnelName, err := h.repo.GetTunnelName(forward.TunnelID) if err != nil { - h.processPeerShareFlowByServiceName(nodeID, serviceName, item) return } shareID, ok := parsePeerShareIDFromFederationTunnelName(tunnelName) - if !ok { - h.processPeerShareFlowByServiceName(nodeID, serviceName, item) - return + if ok { + h.addPeerShareFlow(nodeID, shareID, item.D+item.U) } - - if err := h.repo.AddPeerShareCurrentFlow(shareID, delta); err != nil { - h.processPeerShareFlowByServiceName(nodeID, serviceName, item) - return - } - - share, err := h.repo.GetPeerShare(shareID) - if err != nil || share == nil { - return - } - if !isPeerShareFlowExceeded(share) { - return - } - h.enforcePeerShareFlowLimit(share.ID) } func normalizeForwardRuntimeServiceName(serviceName string) string { @@ -235,63 +193,26 @@ func normalizeForwardRuntimeServiceName(serviceName string) string { } func (h *Handler) processPeerShareFlowByServiceName(nodeID int64, serviceName string, item flowItem) { - if h == nil || h.repo == nil || strings.TrimSpace(serviceName) == "" { + if h == nil || h.repo == nil || nodeID <= 0 || strings.TrimSpace(serviceName) == "" { return } - - delta := item.D + item.U - if delta <= 0 { + runtimes, err := h.repo.ListActiveForwardPeerShareRuntimesByNode(nodeID) + if err != nil { return } - - normalized := normalizeForwardRuntimeServiceName(serviceName) - var runtimes []model.PeerShareRuntime - var err error - - // Try node-scoped query first if nodeID is valid - if nodeID > 0 { - runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByNodeAndServiceName(nodeID, normalized) - if err != nil { + var shareID int64 + for _, runtime := range runtimes { + if normalizeForwardRuntimeServiceName(runtime.ServiceName) != normalizeForwardRuntimeServiceName(serviceName) { + continue + } + if shareID != 0 { + log.Printf("ambiguous peer share runtime service=%s node_id=%d", serviceName, nodeID) return } - if len(runtimes) == 0 && normalized != serviceName { - runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByNodeAndServiceName(nodeID, serviceName) - if err != nil { - return - } - } + shareID = runtime.ShareID } - - // Fallback to global query if node-scoped query returned nothing or nodeID is invalid - if len(runtimes) == 0 { - runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByServiceName(normalized) - if err != nil { - return - } - if len(runtimes) == 0 && normalized != serviceName { - runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByServiceName(serviceName) - if err != nil { - return - } - } - } - - if len(runtimes) != 1 { - if len(runtimes) > 1 { - log.Printf("WARN: ambiguous peer share runtime match for service=%s nodeID=%d count=%d", serviceName, nodeID, len(runtimes)) - } - return - } - runtime := runtimes[0] - - _ = h.repo.AddPeerShareCurrentFlow(runtime.ShareID, delta) - - matchedShare, err := h.repo.GetPeerShare(runtime.ShareID) - if err != nil || matchedShare == nil { - return - } - if isPeerShareFlowExceeded(matchedShare) { - h.enforcePeerShareFlowLimit(matchedShare.ID) + if shareID > 0 { + h.addPeerShareFlow(nodeID, shareID, item.D+item.U) } } @@ -299,22 +220,8 @@ func (h *Handler) enforcePeerShareFlowLimit(shareID int64) { if h == nil || h.repo == nil || shareID <= 0 { return } - runtimes, err := h.repo.ListActivePeerShareRuntimesByShareID(shareID) - if err != nil || len(runtimes) == 0 { - return - } - - now := time.Now().UnixMilli() - for _, runtime := range runtimes { - if h.wsServer != nil && runtime.Applied == 1 { - if strings.TrimSpace(runtime.ServiceName) != "" { - _, _ = h.sendNodeCommand(runtime.NodeID, "DeleteService", map[string]interface{}{"services": []string{runtime.ServiceName}}, false, true) - } - if strings.TrimSpace(runtime.Role) == "middle" && strings.TrimSpace(runtime.ChainName) != "" { - _, _ = h.sendNodeCommand(runtime.NodeID, "DeleteChains", map[string]interface{}{"chain": runtime.ChainName}, false, true) - } - } - _ = h.repo.MarkPeerShareRuntimeReleased(runtime.ID, now) + if err := h.cleanupPeerShareRuntimes(shareID); err != nil { + log.Printf("peer share quota cleanup pending share_id=%d err=%v", shareID, err) } } @@ -531,43 +438,102 @@ func (h *Handler) cleanNodeConfigs(nodeID int64, rawConfig string) { return } - h.cleanOrphanedServices(nodeID, snapshot.Services) - h.cleanOrphanedChains(nodeID, snapshot.Chains) - h.cleanOrphanedLimiters(nodeID, snapshot.Limiters) -} - -func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem) { - runtimeServiceNames, err := h.repo.ListActiveForwardPeerShareRuntimeServiceNamesByNode(nodeID) + protection, err := h.loadForwardServiceProtection(nodeID) if err != nil { return } - minUpdatedTime := time.Now().Add(-10 * time.Minute).UnixMilli() - hasUnboundForwardPeerRuntime, err := h.repo.HasRecentUnboundForwardPeerShareRuntimeOnNode(nodeID, minUpdatedTime) - if err != nil { - hasUnboundForwardPeerRuntime = false + h.cleanOrphanedServicesWithProtection(nodeID, snapshot.Services, protection) + // Dependencies are sent before services. A pending shared reservation may + // therefore have chains/limiters that are not referenced in this snapshot yet. + if protection.unbound { + return } - runtimeServiceSet := make(map[string]struct{}, len(runtimeServiceNames)) - for _, serviceName := range runtimeServiceNames { - serviceName = strings.TrimSpace(serviceName) + chainsInUse := make(map[string]struct{}) + limitersInUse := make(map[string]struct{}) + for _, service := range snapshot.Services { + if service.Handler != nil { + if chain := strings.TrimSpace(service.Handler.Chain); chain != "" { + chainsInUse[chain] = struct{}{} + } + } + for _, limiter := range strings.Split(service.Limiter, ",") { + if limiter = strings.TrimSpace(limiter); limiter != "" { + limitersInUse[limiter] = struct{}{} + } + } + } + // Keep dependencies referenced by the reported services, even when their + // IDs belong to a different panel. Orphan dependencies can be collected on + // the next report after their services have actually disappeared. + h.cleanOrphanedChains(nodeID, snapshot.Chains, chainsInUse) + h.cleanOrphanedLimiters(nodeID, snapshot.Limiters, limitersInUse) +} + +type forwardServiceProtection struct { + sharedNames map[string]struct{} + unbound bool +} + +// Shared forward IDs belong to another panel and need not exist in our forward +// table. Use the same node-scoped ownership check for config and flow reports. +func (h *Handler) loadForwardServiceProtection(nodeID int64) (forwardServiceProtection, error) { + protection := forwardServiceProtection{sharedNames: make(map[string]struct{})} + // Read names and pending bindings in one snapshot, so a concurrent bind + // cannot fall between two queries and disappear from both protections. + runtimes, err := h.repo.ListActiveForwardPeerShareRuntimesByNode(nodeID) + if err != nil { + return protection, err + } + minUpdatedTime := time.Now().Add(-10 * time.Minute).UnixMilli() + for _, runtime := range runtimes { + serviceName := normalizeForwardRuntimeServiceName(runtime.ServiceName) if serviceName == "" { + if runtime.Applied == 0 && runtime.UpdatedTime >= minUpdatedTime { + protection.unbound = true + } continue } - runtimeServiceSet[serviceName] = struct{}{} + protection.sharedNames[serviceName] = struct{}{} } + resources, err := h.repo.ListPeerShareResourcesByNode(nodeID) + if err != nil { + return protection, err + } + for _, resource := range resources { + if base := normalizeForwardRuntimeServiceName(resource.LegacyServiceBase); base != "" { + protection.sharedNames[base] = struct{}{} + } + } + return protection, nil +} +func (p forwardServiceProtection) preserves(serviceName string) bool { + _, shared := p.sharedNames[normalizeForwardRuntimeServiceName(serviceName)] + return shared || p.unbound +} + +func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem) { + if h == nil || h.repo == nil || nodeID <= 0 { + return + } + protection, err := h.loadForwardServiceProtection(nodeID) + if err != nil { + // A failed ownership lookup must never authorize deletion. + return + } + h.cleanOrphanedServicesWithProtection(nodeID, services, protection) +} + +func (h *Handler) cleanOrphanedServicesWithProtection(nodeID int64, services []namedConfigItem, protection forwardServiceProtection) { for _, item := range services { name := strings.TrimSpace(item.Name) if name == "" || name == "web_api" { continue } - if strings.HasPrefix(name, "fed_svc_") { + if strings.HasPrefix(name, "fed_svc_") || strings.HasPrefix(name, "peer-share-") { continue } - normalizedName := normalizeForwardRuntimeServiceName(name) - if _, ok := runtimeServiceSet[normalizedName]; ok { - continue - } - if _, ok := runtimeServiceSet[name]; ok { + if _, ok := protection.sharedNames[normalizeForwardRuntimeServiceName(name)]; ok { continue } @@ -580,15 +546,9 @@ func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem continue } - if len(parts) >= 3 { - forwardID, err := strconv.ParseInt(parts[0], 10, 64) - if err == nil && forwardID > 0 && hasUnboundForwardPeerRuntime { - continue - } - if err == nil && forwardID > 0 && !h.forwardExists(forwardID) { - _, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name, parts[0] + "_" + parts[1] + "_" + parts[2], parts[0] + "_" + parts[1] + "_" + parts[2] + "_tcp", parts[0] + "_" + parts[1] + "_" + parts[2] + "_udp"}}, false, true) - continue - } + if _, _, _, ok := parseFlowServiceIDs(name); ok { + h.deleteOrphanedForwardService(nodeID, name, protection) + continue } suffix := parts[len(parts)-1] @@ -607,23 +567,17 @@ func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem } continue } - forwardID, err := strconv.ParseInt(parts[0], 10, 64) - if err == nil && forwardID > 0 && hasUnboundForwardPeerRuntime { - continue - } - if err != nil || forwardID <= 0 || h.forwardExists(forwardID) { - continue - } - base := strings.TrimSuffix(name, "_tcp") - _, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{base + "_tcp", base + "_udp"}}, false, true) } } } -func (h *Handler) cleanOrphanedChains(nodeID int64, chains []namedConfigItem) { +func (h *Handler) cleanOrphanedChains(nodeID int64, chains []namedConfigItem, inUse map[string]struct{}) { for _, item := range chains { name := strings.TrimSpace(item.Name) - if name == "" { + if name == "" || strings.HasPrefix(name, "fed_chain_") || strings.HasPrefix(name, "peer-share-") { + continue + } + if _, ok := inUse[name]; ok { continue } @@ -632,17 +586,28 @@ func (h *Handler) cleanOrphanedChains(nodeID int64, chains []namedConfigItem) { continue } tunnelID, err := strconv.ParseInt(name[idx+1:], 10, 64) - if err != nil || tunnelID <= 0 || h.tunnelExists(tunnelID) { + if err != nil || tunnelID <= 0 { + continue + } + exists, err := h.repo.TunnelExists(tunnelID) + if err != nil || exists { continue } _, _ = h.sendNodeCommand(nodeID, "DeleteChains", map[string]interface{}{"chain": name}, false, true) } } -func (h *Handler) cleanOrphanedLimiters(nodeID int64, limiters []namedConfigItem) { +func (h *Handler) cleanOrphanedLimiters(nodeID int64, limiters []namedConfigItem, inUse map[string]struct{}) { for _, item := range limiters { name := strings.TrimSpace(item.Name) - if name == "" || h.speedLimiterExists(name) { + if name == "" || strings.HasPrefix(name, "peer-share-") { + continue + } + if _, ok := inUse[name]; ok { + continue + } + exists, err := h.lookupSpeedLimiter(name) + if err != nil || exists { continue } _, _ = h.sendNodeCommand(nodeID, "DeleteLimiters", map[string]interface{}{"limiter": name}, false, true) @@ -660,40 +625,78 @@ func (h *Handler) forwardExists(forwardID int64) bool { } func (h *Handler) sendDeleteOrphanedForwardService(nodeID int64, serviceName string) { + h.sendDeleteOrphanedForwardServices(nodeID, []string{serviceName}) +} + +func (h *Handler) sendDeleteOrphanedForwardServices(nodeID int64, serviceNames []string) { + if h == nil || h.repo == nil || nodeID <= 0 || len(serviceNames) == 0 { + return + } + protection, err := h.loadForwardServiceProtection(nodeID) + if err != nil { + return + } + seen := make(map[string]struct{}, len(serviceNames)) + for _, serviceName := range serviceNames { + serviceName = normalizeForwardRuntimeServiceName(serviceName) + if _, ok := seen[serviceName]; ok { + continue + } + seen[serviceName] = struct{}{} + h.deleteOrphanedForwardService(nodeID, serviceName, protection) + } +} + +func (h *Handler) deleteOrphanedForwardService(nodeID int64, serviceName string, protection forwardServiceProtection) { + forwardID, _, _, ok := parseFlowServiceIDs(serviceName) + if !ok || protection.preserves(serviceName) { + return + } parts := strings.Split(serviceName, "_") - if len(parts) < 3 { - return - } - forwardID, err := strconv.ParseInt(parts[0], 10, 64) - if err != nil || forwardID <= 0 { - return - } base := parts[0] + "_" + parts[1] + "_" + parts[2] + // Parsing accepts legacy suffixes, while deletion targets the entire base + // family. Verify the actual targets cannot include a protected share. + if protection.preserves(base) { + return + } + // Batch metadata can be missing after a read failure, or stale by the time + // cleanup runs. Confirm absence before issuing a destructive command. + exists, err := h.repo.ForwardExists(forwardID) + if err != nil || exists { + return + } _, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{ - "services": []string{base + "_tcp", base + "_udp"}, + "services": buildForwardServiceDeleteNames([]string{base}), }, false, true) } func (h *Handler) speedLimiterExists(name string) bool { + exists, _ := h.lookupSpeedLimiter(name) + return exists +} + +func (h *Handler) lookupSpeedLimiter(name string) (bool, error) { name = strings.TrimSpace(name) if name == "" { - return false + return false, nil } const forwardRulePrefix = "rule_traffic_limit_" if strings.HasPrefix(name, forwardRulePrefix) { forwardID, err := strconv.ParseInt(strings.TrimPrefix(name, forwardRulePrefix), 10, 64) if err != nil || forwardID <= 0 { - return false + return false, nil } forward, err := h.getForwardRecord(forwardID) - return err == nil && forward != nil && forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0 + if errors.Is(err, errForwardNotFound) { + return false, nil + } + return forward != nil && forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0, err } id, err := strconv.ParseInt(name, 10, 64) if err != nil || id <= 0 { - return false + return false, nil } - ok, _ := h.repo.SpeedLimitExists(id) - return ok + return h.repo.SpeedLimitExists(id) } diff --git a/go-backend/internal/http/handler/flow_policy_federation_test.go b/go-backend/internal/http/handler/flow_policy_federation_test.go index c2a3bfc..61f4bcb 100644 --- a/go-backend/internal/http/handler/flow_policy_federation_test.go +++ b/go-backend/internal/http/handler/flow_policy_federation_test.go @@ -10,11 +10,8 @@ import ( ) func TestProcessFlowItemTracksPeerShareFlowAndEnforcesLimit(t *testing.T) { - r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) - if err != nil { - t.Fatalf("open repo: %v", err) - } - defer r.Close() + a := newCleanupAgent(t) + r := a.h.repo now := time.Now().UnixMilli() if err := r.CreatePeerShare(&repo.PeerShare{ @@ -43,7 +40,7 @@ func TestProcessFlowItemTracksPeerShareFlowAndEnforcesLimit(t *testing.T) { t.Fatalf("insert peer_share_runtime: %v", err) } - h := &Handler{repo: r} + h := a.h h.processFlowItem(1, flowItem{N: "fed_svc_17", U: 1200, D: 900}) updatedShare, err := r.GetPeerShare(share.ID) @@ -120,6 +117,9 @@ func TestProcessFlowItemTracksPeerShareFlowForFederationPortForward(t *testing.T } h := &Handler{repo: r} + if err := r.DB().Exec("INSERT INTO forward_port(forward_id, node_id, port) VALUES(20, 1, 30001)").Error; err != nil { + t.Fatal(err) + } h.processFlowItem(1, flowItem{N: "20_2_10", U: 120, D: 80}) updatedShare, err := r.GetPeerShare(share.ID) diff --git a/go-backend/internal/http/handler/flow_upload_batch.go b/go-backend/internal/http/handler/flow_upload_batch.go index 4cb8cfc..c5b19ce 100644 --- a/go-backend/internal/http/handler/flow_upload_batch.go +++ b/go-backend/internal/http/handler/flow_upload_batch.go @@ -23,6 +23,7 @@ type flowUploadBatch struct { orphanServices map[string]struct{} peerShareForwardItems map[string]flowItem peerShareRuntimeItems map[int64]flowItem + peerShareUsage map[int64]int64 } func (h *Handler) buildFlowUploadBatch(items []flowItem, metas map[int64]repo.FlowUploadForwardMeta) flowUploadBatch { @@ -65,6 +66,9 @@ func (h *Handler) buildFlowUploadBatch(items []flowItem, metas map[int64]repo.Fl batch.orphanServices[serviceName] = struct{}{} continue } + // Local accounting uses database ownership, not foreign IDs embedded in + // a service name (including stale user-tunnel IDs after reassignment). + userID, userTunnelID = meta.UserID, meta.UserTunnelID raw := batch.forwardTraffic[forwardID] raw.bytesIn += item.D @@ -120,8 +124,12 @@ func (h *Handler) applyFlowUploadBatch(nodeID int64, batch flowUploadBatch, now } h.enforceFlowPolicies(target.UserID, target.UserTunnelID) } - for serviceName := range batch.orphanServices { - h.sendDeleteOrphanedForwardService(nodeID, serviceName) + if len(batch.orphanServices) > 0 { + serviceNames := make([]string, 0, len(batch.orphanServices)) + for serviceName := range batch.orphanServices { + serviceNames = append(serviceNames, serviceName) + } + h.sendDeleteOrphanedForwardServices(nodeID, serviceNames) } for serviceName, item := range batch.peerShareForwardItems { forwardID, _, _, ok := parseFlowServiceIDs(serviceName) @@ -130,7 +138,10 @@ func (h *Handler) applyFlowUploadBatch(nodeID int64, batch flowUploadBatch, now } } for runtimeID, item := range batch.peerShareRuntimeItems { - h.processPeerShareFlow(runtimeID, item) + h.processPeerShareFlow(nodeID, runtimeID, item) + } + for shareID, delta := range batch.peerShareUsage { + h.addPeerShareFlow(nodeID, shareID, delta) } } diff --git a/go-backend/internal/http/handler/flow_upload_batch_test.go b/go-backend/internal/http/handler/flow_upload_batch_test.go index 2fb20f5..861857f 100644 --- a/go-backend/internal/http/handler/flow_upload_batch_test.go +++ b/go-backend/internal/http/handler/flow_upload_batch_test.go @@ -14,6 +14,8 @@ func TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets(t metas := map[int64]repo.FlowUploadForwardMeta{ 20: { ForwardID: 20, + UserID: 2, + UserTunnelID: 10, TunnelID: 1, TrafficRatio: 2, TunnelFlow: 3, diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 4d67883..4cfcfe9 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -39,8 +39,9 @@ type Handler struct { healthCheck *health.Checker nftablesManager nftablesRuntimeManager - captchaMu sync.Mutex - captchaTokens map[string]int64 + peerResourceMu sync.Mutex + captchaMu sync.Mutex + captchaTokens map[string]int64 jobsMu sync.Mutex jobsCancel context.CancelFunc @@ -49,12 +50,13 @@ type Handler struct { fingerprintMu sync.Mutex licenseValidationMu sync.Mutex - upgradeMu sync.Mutex - systemUpgradeMu sync.Mutex - pendingUpgradeRedeploy map[int64]struct{} - nodeOnlineRedeployAt map[int64]time.Time - nodeOnlineRedeployQueued map[int64]struct{} - nodeOnlineRedeploying map[int64]struct{} + upgradeMu sync.Mutex + systemUpgradeMu sync.Mutex + pendingUpgradeRedeploy map[int64]struct{} + nodeOnlineRedeployAt map[int64]time.Time + nodeOnlineRedeployQueued map[int64]struct{} + nodeOnlineRedeploying map[int64]struct{} + nodeLocalRuntimeRetryQueued map[int64]struct{} qualityProber *tunnelQualityProber bestExit *bestExitManager @@ -876,7 +878,7 @@ func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) { log.Printf("flow upload metadata lookup failed node_id=%d err=%v", node.ID, metaErr) metas = map[int64]repo.FlowUploadForwardMeta{} } - batch := h.buildFlowUploadBatch(items, metas) + batch := h.buildNodeFlowUploadBatch(node.ID, items, metas) h.recordTunnelMetricsFromForwardBatch(node.ID, batch.forwardTraffic, metas, now.UnixMilli()) h.applyFlowUploadBatch(node.ID, batch, now) } diff --git a/go-backend/internal/http/handler/jobs.go b/go-backend/internal/http/handler/jobs.go index cebca8c..72df747 100644 --- a/go-backend/internal/http/handler/jobs.go +++ b/go-backend/internal/http/handler/jobs.go @@ -21,7 +21,7 @@ func (h *Handler) StartBackgroundJobs() { ctx, cancel := context.WithCancel(context.Background()) h.jobsCancel = cancel h.jobsStarted = true - h.jobsWG.Add(8) + h.jobsWG.Add(9) h.jobsMu.Unlock() go h.runHourlyStatsLoop(ctx) @@ -32,6 +32,49 @@ func (h *Handler) StartBackgroundJobs() { go h.runTunnelQualityProber(ctx) go h.runValidateLicenseJob(ctx) go h.runNftablesTrafficCollectLoop(ctx) + go h.runFederationCleanupRetryLoop(ctx) +} + +func (h *Handler) runFederationCleanupRetryLoop(ctx context.Context) { + defer h.jobsWG.Done() + ticker := time.NewTicker(30 * time.Second) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + default: + } + if err := h.retryPendingFederationRuntimeCleanup(); err != nil { + log.Printf("federation cleanup remains pending: %v", err) + } + h.retryPendingPeerShareOperations() + select { + case <-ctx.Done(): + return + case <-ticker.C: + } + } +} + +func (h *Handler) retryPendingPeerShareOperations() { + nodeIDs, err := h.repo.ListPendingPeerShareNodeIDs() + if err != nil { + log.Printf("peer share pending operation lookup failed: %v", err) + return + } + for _, nodeID := range nodeIDs { + node, err := h.repo.GetNodeByID(nodeID) + if err != nil || node == nil || node.Status != 1 { + continue + } + if err := h.retryPendingPeerShareResourcesOnNode(nodeID); err != nil { + log.Printf("peer share resource retry failed node_id=%d err=%v", nodeID, err) + } + if err := h.retryPendingPeerShareRoleRuntimesOnNode(nodeID); err != nil { + log.Printf("peer share role retry failed node_id=%d err=%v", nodeID, err) + } + } } func (h *Handler) runValidateLicenseJob(ctx context.Context) { diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 1a11ab1..4e57ea8 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -827,6 +827,10 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault(err.Error())) return } + if err := validateTunnelRuntimeBeforeApply(runtimeState); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } if len(runtimeState.InNodes) > 0 { firstNodeID := runtimeState.InNodes[0].NodeID @@ -909,22 +913,27 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { var federationReleaseRefs []federationRuntimeReleaseRef federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState, localDomain) if err != nil { + tx.Rollback() + err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs)) response.WriteJSON(w, response.ErrDefault(err.Error())) return } applyTunnelPortsToRequest(req, runtimeState) if err := h.replaceTunnelChainsTx(tx, tunnelID, req); err != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) + tx.Rollback() + err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs)) response.WriteJSON(w, response.Err(-2, err.Error())) return } if err := h.repo.ReplaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); err != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) + tx.Rollback() + err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs)) response.WriteJSON(w, response.Err(-2, err.Error())) return } if err := tx.Commit().Error; err != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) + tx.Rollback() + err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs)) response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -932,8 +941,11 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState) if applyErr != nil { h.rollbackTunnelRuntime(createdChains, createdServices, tunnelID, tunnelProtocol) - h.releaseFederationRuntimeRefs(federationReleaseRefs) - _ = h.deleteTunnelByID(tunnelID) + if cleanupErr := h.cleanupFederationRuntime(tunnelID); cleanupErr != nil { + applyErr = errors.Join(applyErr, cleanupErr) + } else { + _ = h.deleteTunnelByID(tunnelID) + } response.WriteJSON(w, response.ErrDefault(applyErr.Error())) return } @@ -1091,22 +1103,38 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) { probeTargetPort = probeTarget.Port } } - oldEntryNodeIDs, _ := h.tunnelEntryNodeIDs(id) - oldTunnel, _ := h.getTunnelRecord(id) + oldEntryNodeIDs, err := h.tunnelEntryNodeIDs(id) + if err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + oldTunnel, err := h.getTunnelRecord(id) + if err != nil || oldTunnel == nil { + response.WriteJSON(w, response.ErrDefault("隧道不存在")) + return + } if !probeTargetFieldsPresent && oldTunnel != nil { probeTargetHost = oldTunnel.ProbeTargetHost probeTargetPort = oldTunnel.ProbeTargetPort } - oldChainRows, _ := h.listChainNodesForTunnel(id) - if oldTunnel != nil && oldTunnel.Type == 2 && typeVal != 2 { - h.cleanupTunnelRuntime(id) + oldChainRows, err := h.listChainNodesForTunnel(id) + if err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return } - h.cleanupFederationRuntime(id) now := time.Now().UnixMilli() localDomain := h.federationLocalDomain() - runtimeState, err := h.prepareTunnelCreateState(h.repo.DB(), req, typeVal, id) + // Check the complete request and known database constraints before changing + // any existing listener or remote binding. + validationTx := h.repo.BeginTx() + if validationTx.Error != nil { + response.WriteJSON(w, response.Err(-2, validationTx.Error.Error())) + return + } + defer validationTx.Rollback() + runtimeState, err := h.prepareTunnelCreateState(validationTx, req, typeVal, id) if err != nil { response.WriteJSON(w, response.ErrDefault(err.Error())) return @@ -1119,36 +1147,69 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) { entryNodeIDs = append(entryNodeIDs, inNode.NodeID) } } - if err := h.validateNftablesTunnelState(entryNodeIDs); err != nil { + if err := h.validateNftablesTunnelStateTx(validationTx, entryNodeIDs); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + + if err := h.validateTunnelEntryPortConflictsForNewEntriesTx(validationTx, id, oldEntryNodeIDs, entryNodeIDs); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + if err := validateTunnelRuntimeBeforeApply(runtimeState); err != nil { response.WriteJSON(w, response.ErrDefault(err.Error())) return } inIp := buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes, ipPreference) - var federationBindings []repo.FederationTunnelBinding - var federationReleaseRefs []federationRuntimeReleaseRef - federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState, localDomain) - if err != nil { - response.WriteJSON(w, response.ErrDefault(err.Error())) - return - } - applyTunnelPortsToRequest(req, runtimeState) - - tx := h.repo.BeginTx() - if tx.Error != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) - response.WriteJSON(w, response.Err(-2, tx.Error.Error())) - return - } - defer func() { tx.Rollback() }() - updateProtocol := "tls" if len(runtimeState.OutNodes) > 0 && strings.TrimSpace(runtimeState.OutNodes[0].Protocol) != "" { updateProtocol = strings.TrimSpace(runtimeState.OutNodes[0].Protocol) } else if len(runtimeState.InNodes) > 0 && strings.TrimSpace(runtimeState.InNodes[0].Protocol) != "" { updateProtocol = strings.TrimSpace(runtimeState.InNodes[0].Protocol) } + if err := h.repo.UpdateTunnelTx(validationTx, id, asString(req["name"]), typeVal, + asInt64(req["flow"], 1), asFloat(req["trafficRatio"], 1.0), asInt(req["status"], 1), + inIp, ipPreference, updateProtocol, probeTargetHost, probeTargetPort, now); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err := h.validateTunnelChainReplacementTx(validationTx, runtimeState); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err := validationTx.Rollback().Error; err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err := h.cleanupFederationRuntime(id); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + if oldTunnel.Type == 2 && typeVal != 2 { + h.cleanupTunnelRuntime(id) + } + + var federationBindings []repo.FederationTunnelBinding + var federationReleaseRefs []federationRuntimeReleaseRef + federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState, localDomain) + if err != nil { + err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs)) + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + applyTunnelPortsToRequest(req, runtimeState) + + tx := h.repo.BeginTx() + if tx.Error != nil { + tx.Rollback() + err := errors.Join(tx.Error, h.releaseFederationRuntimeRefs(federationReleaseRefs)) + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + defer func() { tx.Rollback() }() + if err := h.repo.UpdateTunnelTx( tx, id, @@ -1164,21 +1225,27 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) { probeTargetPort, now, ); err != nil { + tx.Rollback() + err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs)) response.WriteJSON(w, response.Err(-2, err.Error())) return } if err := h.repo.DeleteChainTunnelsByTunnelTx(tx, id); err != nil { + tx.Rollback() + err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs)) response.WriteJSON(w, response.Err(-2, err.Error())) return } if err := h.replaceTunnelChainsTx(tx, id, req); err != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) + tx.Rollback() + err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs)) response.WriteJSON(w, response.Err(-2, err.Error())) return } if err := h.repo.ReplaceFederationTunnelBindingsTx(tx, id, federationBindings); err != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) + tx.Rollback() + err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs)) response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1190,12 +1257,15 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) { } } if err := h.validateTunnelEntryPortConflictsForNewEntriesTx(tx, id, oldEntryNodeIDs, newEntryNodeIDs); err != nil { + tx.Rollback() + err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs)) response.WriteJSON(w, response.ErrDefault(err.Error())) return } if err := tx.Commit().Error; err != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) + tx.Rollback() + err = errors.Join(err, h.releaseFederationRuntimeRefs(federationReleaseRefs)) response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1220,8 +1290,7 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) { if oldTunnel == nil || oldTunnel.Type != 2 { h.rollbackTunnelRuntime(createdChains, createdServices, id, updateProtocol) } - h.releaseFederationRuntimeRefs(federationReleaseRefs) - _ = h.repo.DeleteFederationTunnelBindingsByTunnel(id) + applyErr = errors.Join(applyErr, h.cleanupFederationRuntime(id)) if len(federationReleaseRefs) == 0 && shouldDeferTunnelRuntimeApplyError(applyErr) { response.WriteJSON(w, response.OKEmpty()) return @@ -1430,7 +1499,10 @@ func (h *Handler) validateTunnelEntryPortConflictsForNewEntriesTx(tx *gorm.DB, t } forwards, err := h.repo.ListForwardsByTunnelTx(tx, tunnelID) - if err != nil || len(forwards) == 0 { + if err != nil { + return err + } + if len(forwards) == 0 { return nil } @@ -1441,7 +1513,7 @@ func (h *Handler) validateTunnelEntryPortConflictsForNewEntriesTx(tx *gorm.DB, t } oldPorts, portsErr := h.repo.ListForwardPortsTx(tx, f.ID) if portsErr != nil { - continue + return portsErr } port := pickForwardPortFromRecords(oldPorts) if port <= 0 { @@ -1451,7 +1523,7 @@ func (h *Handler) validateTunnelEntryPortConflictsForNewEntriesTx(tx *gorm.DB, t for _, nodeID := range addedNodeIDs { node, nodeErr := h.repo.GetNodeRecordTx(tx, nodeID) if nodeErr != nil { - continue + return nodeErr } if err := h.validateForwardPortAvailabilityTx(tx, node, port, f.ID); err != nil { @@ -1603,8 +1675,11 @@ func (h *Handler) tunnelDelete(w http.ResponseWriter, r *http.Request) { if id <= 0 { return } + if err := h.cleanupFederationRuntime(id); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } h.cleanupTunnelRuntime(id) - h.cleanupFederationRuntime(id) if err := h.deleteTunnelByID(id); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -1671,8 +1746,12 @@ func (h *Handler) tunnelBatchDelete(w http.ResponseWriter, r *http.Request) { failures = appendBatchFailure(failures, id, tunnelName, err) continue } + if err := h.cleanupFederationRuntime(id); err != nil { + fail++ + failures = appendBatchFailure(failures, id, tunnelName, err) + continue + } h.cleanupTunnelRuntime(id) - h.cleanupFederationRuntime(id) if err := h.deleteTunnelByID(id); err != nil { fail++ failures = appendBatchFailure(failures, id, tunnelName, err) @@ -1771,34 +1850,35 @@ func (h *Handler) redeployTunnelAndForwards(tunnelID int64) error { } if tunnel.Type == 2 { - h.cleanupTunnelRuntime(tunnelID) - h.cleanupFederationRuntime(tunnelID) state, err := h.reconstructTunnelState(tunnelID) if err != nil { return err } + if err := validateTunnelRuntimeBeforeApply(state); err != nil { + return err + } + if err := h.cleanupFederationRuntime(tunnelID); err != nil { + return err + } + h.cleanupTunnelRuntime(tunnelID) federationBindings, federationReleaseRefs, fedErr := h.applyFederationRuntime(state, h.federationLocalDomain()) if fedErr != nil { - return fedErr + return errors.Join(fedErr, h.releaseFederationRuntimeRefs(federationReleaseRefs)) } tx := h.repo.BeginTx() if tx.Error != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) - return tx.Error + return errors.Join(tx.Error, h.releaseFederationRuntimeRefs(federationReleaseRefs)) } if replaceErr := h.repo.ReplaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); replaceErr != nil { tx.Rollback() - h.releaseFederationRuntimeRefs(federationReleaseRefs) - return replaceErr + return errors.Join(replaceErr, h.releaseFederationRuntimeRefs(federationReleaseRefs)) } if commitErr := tx.Commit().Error; commitErr != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) - return commitErr + return errors.Join(commitErr, h.releaseFederationRuntimeRefs(federationReleaseRefs)) } _, _, applyErr := h.applyTunnelRuntime(state) if applyErr != nil { - h.releaseFederationRuntimeRefs(federationReleaseRefs) - _ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID) + applyErr = errors.Join(applyErr, h.cleanupFederationRuntime(tunnelID)) return applyErr } } @@ -3393,12 +3473,13 @@ func (h *Handler) prepareTunnelCreateState(tx *gorm.DB, req map[string]interface existingNodeIDs := make(map[int64]struct{}) if excludeTunnelID > 0 { var existIDs []int64 - if err := tx.Model(&model.ChainTunnel{}). - Where("tunnel_id = ?", excludeTunnelID). - Pluck("node_id", &existIDs).Error; err == nil { - for _, eid := range existIDs { - existingNodeIDs[eid] = struct{}{} - } + var err error + existIDs, err = h.repo.ListTunnelChainNodeIDsTx(tx, excludeTunnelID) + if err != nil { + return nil, err + } + for _, eid := range existIDs { + existingNodeIDs[eid] = struct{}{} } } @@ -3446,6 +3527,59 @@ func (h *Handler) prepareTunnelCreateState(tx *gorm.DB, req map[string]interface return state, nil } +func validateTunnelRuntimeBeforeApply(state *tunnelCreateState) error { + if state.Type != 2 { + return nil + } + for _, node := range state.Nodes { + if node.IsRemote == 1 && (strings.TrimSpace(node.RemoteURL) == "" || strings.TrimSpace(node.RemoteToken) == "") { + return fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node)) + } + } + groups := append([][]tunnelRuntimeNode{state.InNodes}, state.ChainHops...) + groups = append(groups, state.OutNodes) + for i := 0; i+1 < len(groups); i++ { + for _, source := range groups[i] { + for _, target := range groups[i+1] { + node := state.Nodes[target.NodeID] + if node == nil { + return errors.New("节点不存在") + } + if _, err := selectTunnelDialHost(state.Nodes[source.NodeID], node, state.IPPreference, target.ConnectIP); err != nil { + return err + } + if node.IsRemote != 1 && target.Port <= 0 { + return errors.New("节点端口不能为空") + } + } + } + } + return nil +} + +// Exercise known chain write constraints inside the validation transaction. +// Remote ports that have not been reserved yet remain NULL until actual apply. +func (h *Handler) validateTunnelChainReplacementTx(tx *gorm.DB, state *tunnelCreateState) error { + if err := h.repo.DeleteChainTunnelsByTunnelTx(tx, state.TunnelID); err != nil { + return err + } + groups := append([][]tunnelRuntimeNode{state.InNodes, state.OutNodes}, state.ChainHops...) + for groupIndex, group := range groups { + for nodeIndex, node := range group { + inx := nodeIndex + 1 + if node.ChainType == 2 { + inx = groupIndex - 1 + } + port := sql.NullInt64{Int64: int64(node.Port), Valid: node.Port > 0} + if err := h.repo.CreateChainTunnelTx(tx, state.TunnelID, strconv.Itoa(node.ChainType), + node.NodeID, port, node.Strategy, inx, node.Protocol, node.ConnectIP); err != nil { + return err + } + } + } + return nil +} + func buildTunnelInIP(inNodes []tunnelRuntimeNode, nodes map[int64]*nodeRecord, ipPreference string) string { set := make(map[string]struct{}) ordered := make([]string, 0) @@ -3595,8 +3729,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s remoteURL := strings.TrimSpace(node.RemoteURL) remoteToken := strings.TrimSpace(node.RemoteToken) if remoteURL == "" || remoteToken == "" { - h.releaseFederationRuntimeRefs(releaseRefs) - return nil, nil, fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node)) + return bindings, releaseRefs, fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node)) } resourceKey := federationRuntimeResourceKey(state.TunnelID, outNode.NodeID, 3, 0) @@ -3611,10 +3744,14 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s reserveRes, err = fc.ReservePort(remoteURL, remoteToken, localDomain, reserveReq) } if err != nil { - h.releaseFederationRuntimeRefs(releaseRefs) - return nil, nil, fmt.Errorf("远程节点 %s 端口分配失败: %w", nodeDisplayName(node), err) + return bindings, releaseRefs, fmt.Errorf("远程节点 %s 端口分配失败: %w", nodeDisplayName(node), err) } + releaseRefs = append(releaseRefs, federationRuntimeReleaseRef{ + RemoteURL: remoteURL, RemoteToken: remoteToken, BindingID: reserveRes.BindingID, + ReservationID: reserveRes.ReservationID, ResourceKey: resourceKey, + }) + state.OutNodes[outIdx].Port = reserveRes.AllocatedPort outNode = state.OutNodes[outIdx] @@ -3627,8 +3764,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s } applyRes, err := fc.ApplyRole(remoteURL, remoteToken, localDomain, applyReq) if err != nil { - h.releaseFederationRuntimeRefs(releaseRefs) - return nil, nil, fmt.Errorf("远程节点 %s 运行时下发失败: %w", nodeDisplayName(node), err) + return bindings, releaseRefs, fmt.Errorf("远程节点 %s 运行时下发失败: %w", nodeDisplayName(node), err) } if applyRes.AllocatedPort > 0 { state.OutNodes[outIdx].Port = applyRes.AllocatedPort @@ -3648,13 +3784,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s CreatedTime: now, UpdatedTime: now, }) - releaseRefs = append(releaseRefs, federationRuntimeReleaseRef{ - RemoteURL: remoteURL, - RemoteToken: remoteToken, - BindingID: applyRes.BindingID, - ReservationID: reserveRes.ReservationID, - ResourceKey: resourceKey, - }) + releaseRefs[len(releaseRefs)-1].BindingID = defaultString(applyRes.BindingID, reserveRes.BindingID) } for hopIdx := len(state.ChainHops) - 1; hopIdx >= 0; hopIdx-- { @@ -3667,8 +3797,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s remoteURL := strings.TrimSpace(node.RemoteURL) remoteToken := strings.TrimSpace(node.RemoteToken) if remoteURL == "" || remoteToken == "" { - h.releaseFederationRuntimeRefs(releaseRefs) - return nil, nil, fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node)) + return bindings, releaseRefs, fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node)) } resourceKey := federationRuntimeResourceKey(state.TunnelID, chainNode.NodeID, 2, hopIdx+1) @@ -3683,10 +3812,14 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s reserveRes, err = fc.ReservePort(remoteURL, remoteToken, localDomain, reserveReq) } if err != nil { - h.releaseFederationRuntimeRefs(releaseRefs) - return nil, nil, fmt.Errorf("远程节点 %s 端口分配失败: %w", nodeDisplayName(node), err) + return bindings, releaseRefs, fmt.Errorf("远程节点 %s 端口分配失败: %w", nodeDisplayName(node), err) } + releaseRefs = append(releaseRefs, federationRuntimeReleaseRef{ + RemoteURL: remoteURL, RemoteToken: remoteToken, BindingID: reserveRes.BindingID, + ReservationID: reserveRes.ReservationID, ResourceKey: resourceKey, + }) + state.ChainHops[hopIdx][nodeIdx].Port = reserveRes.AllocatedPort chainNode = state.ChainHops[hopIdx][nodeIdx] @@ -3700,17 +3833,14 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s for _, target := range nextTargets { targetNode := state.Nodes[target.NodeID] if targetNode == nil { - h.releaseFederationRuntimeRefs(releaseRefs) - return nil, nil, errors.New("节点不存在") + return bindings, releaseRefs, errors.New("节点不存在") } host, hostErr := selectTunnelDialHost(node, targetNode, state.IPPreference, target.ConnectIP) if hostErr != nil { - h.releaseFederationRuntimeRefs(releaseRefs) - return nil, nil, hostErr + return bindings, releaseRefs, hostErr } if target.Port <= 0 { - h.releaseFederationRuntimeRefs(releaseRefs) - return nil, nil, errors.New("节点端口不能为空") + return bindings, releaseRefs, errors.New("节点端口不能为空") } applyTargets = append(applyTargets, client.RuntimeTarget{ Host: host, @@ -3729,8 +3859,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s } applyRes, err := fc.ApplyRole(remoteURL, remoteToken, localDomain, applyReq) if err != nil { - h.releaseFederationRuntimeRefs(releaseRefs) - return nil, nil, fmt.Errorf("远程节点 %s 运行时下发失败: %w", nodeDisplayName(node), err) + return bindings, releaseRefs, fmt.Errorf("远程节点 %s 运行时下发失败: %w", nodeDisplayName(node), err) } if applyRes.AllocatedPort > 0 { state.ChainHops[hopIdx][nodeIdx].Port = applyRes.AllocatedPort @@ -3750,70 +3879,124 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s CreatedTime: now, UpdatedTime: now, }) - releaseRefs = append(releaseRefs, federationRuntimeReleaseRef{ - RemoteURL: remoteURL, - RemoteToken: remoteToken, - BindingID: applyRes.BindingID, - ReservationID: reserveRes.ReservationID, - ResourceKey: resourceKey, - }) + releaseRefs[len(releaseRefs)-1].BindingID = defaultString(applyRes.BindingID, reserveRes.BindingID) } } return bindings, releaseRefs, nil } -func (h *Handler) releaseFederationRuntimeRefs(refs []federationRuntimeReleaseRef) { +// releaseFederationRuntimeRefs is used after the caller has rolled back its +// transaction. Failed releases survive process restarts in a separate queue. +func (h *Handler) releaseFederationRuntimeRefs(refs []federationRuntimeReleaseRef) error { if h == nil || len(refs) == 0 { - return + return nil } fc := client.NewFederationClient() localDomain := h.federationLocalDomain() + var failures []error for i := len(refs) - 1; i >= 0; i-- { ref := refs[i] - if strings.TrimSpace(ref.RemoteURL) == "" || strings.TrimSpace(ref.RemoteToken) == "" { + pending := &model.FederationPendingRelease{ + RemoteURL: ref.RemoteURL, RemoteToken: ref.RemoteToken, + BindingID: ref.BindingID, ReservationID: ref.ReservationID, ResourceKey: ref.ResourceKey, + } + if err := h.repo.SavePendingFederationRelease(pending); err != nil { + failures = append(failures, fmt.Errorf("保存共享运行时清理任务失败: %w", err)) continue } req := client.RuntimeReleaseRoleRequest{ - BindingID: ref.BindingID, - ReservationID: ref.ReservationID, - ResourceKey: ref.ResourceKey, + BindingID: ref.BindingID, ReservationID: ref.ReservationID, ResourceKey: ref.ResourceKey, + } + if err := fc.ReleaseRole(ref.RemoteURL, ref.RemoteToken, localDomain, req); err != nil { + failures = append(failures, err) + continue + } + if err := h.repo.DeletePendingFederationRelease(pending.ID); err != nil { + failures = append(failures, err) } - _ = fc.ReleaseRole(ref.RemoteURL, ref.RemoteToken, localDomain, req) } + return errors.Join(failures...) } -func (h *Handler) cleanupFederationRuntime(tunnelID int64) { +func (h *Handler) cleanupFederationRuntime(tunnelID int64) error { if h == nil || tunnelID <= 0 { - return + return nil } - bindings, err := h.repo.ListActiveFederationTunnelBindingsByTunnel(tunnelID) - if err != nil || len(bindings) == 0 { - return + bindings, err := h.repo.ListFederationTunnelBindingsForCleanup(tunnelID) + if err != nil { + return err } - - fc := client.NewFederationClient() - localDomain := h.federationLocalDomain() + var failures []error for _, b := range bindings { - node, nodeErr := h.repo.GetNodeByID(b.NodeID) - if nodeErr != nil || node == nil { + // Persist intent before sending the request: timeouts can occur after + // the peer has already acted, and must remain safely retryable. + if err := h.repo.MarkFederationTunnelBindingPendingRelease(b.ID); err != nil { + failures = append(failures, err) continue } - remoteURL := strings.TrimSpace(node.RemoteURL.String) - if remoteURL == "" { - remoteURL = strings.TrimSpace(b.RemoteURL) + if err := h.releaseFederationTunnelBinding(b); err != nil { + failures = append(failures, err) } - remoteToken := strings.TrimSpace(node.RemoteToken.String) - if remoteURL == "" || remoteToken == "" { - continue - } - req := client.RuntimeReleaseRoleRequest{ - BindingID: strings.TrimSpace(b.RemoteBindingID), - ResourceKey: strings.TrimSpace(b.ResourceKey), - } - _ = fc.ReleaseRole(remoteURL, remoteToken, localDomain, req) } - _ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID) + return errors.Join(failures...) +} + +func (h *Handler) releaseFederationTunnelBinding(b repo.FederationTunnelBinding) error { + node, err := h.repo.GetNodeByID(b.NodeID) + if err != nil { + return err + } + if node == nil { + return fmt.Errorf("共享节点 %d 不存在,保留待清理绑定", b.NodeID) + } + remoteURL := strings.TrimSpace(b.RemoteURL) + if remoteURL == "" { + remoteURL = strings.TrimSpace(node.RemoteURL.String) + } + remoteToken := strings.TrimSpace(node.RemoteToken.String) + if remoteURL == "" || remoteToken == "" { + return fmt.Errorf("共享节点 %d 缺少连接配置,保留待清理绑定", b.NodeID) + } + req := client.RuntimeReleaseRoleRequest{ + BindingID: strings.TrimSpace(b.RemoteBindingID), ResourceKey: strings.TrimSpace(b.ResourceKey), + } + if err := client.NewFederationClient().ReleaseRole(remoteURL, remoteToken, h.federationLocalDomain(), req); err != nil { + return fmt.Errorf("共享节点 %d 清理失败: %w", b.NodeID, err) + } + return h.repo.DeleteFederationTunnelBinding(b.ID) +} + +// retryPendingFederationRuntimeCleanup never touches active bindings. +func (h *Handler) retryPendingFederationRuntimeCleanup() error { + bindings, err := h.repo.ListPendingFederationTunnelBindings() + if err != nil { + return err + } + var failures []error + for _, binding := range bindings { + if err := h.releaseFederationTunnelBinding(binding); err != nil { + failures = append(failures, err) + } + } + pending, err := h.repo.ListPendingFederationReleases() + if err != nil { + return errors.Join(append(failures, err)...) + } + fc := client.NewFederationClient() + for _, item := range pending { + req := client.RuntimeReleaseRoleRequest{ + BindingID: item.BindingID, ReservationID: item.ReservationID, ResourceKey: item.ResourceKey, + } + if err := fc.ReleaseRole(item.RemoteURL, item.RemoteToken, h.federationLocalDomain(), req); err != nil { + failures = append(failures, err) + continue + } + if err := h.repo.DeletePendingFederationRelease(item.ID); err != nil { + failures = append(failures, err) + } + } + return errors.Join(failures...) } func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64, error) { diff --git a/go-backend/internal/http/handler/peer_share_auth_cleanup_test.go b/go-backend/internal/http/handler/peer_share_auth_cleanup_test.go new file mode 100644 index 0000000..6c0b9cd --- /dev/null +++ b/go-backend/internal/http/handler/peer_share_auth_cleanup_test.go @@ -0,0 +1,160 @@ +package handler + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "go-backend/internal/http/response" + "go-backend/internal/store/repo" +) + +func TestPeerShareRestrictedShareAllowsOnlyAuthenticatedCleanup(t *testing.T) { + for _, state := range []string{"disabled", "expired", "over-quota"} { + t.Run(state, func(t *testing.T) { + agent := newCleanupAgent(t) + runtime := roleRuntimeFixture(t, agent, "exit") + changes := map[string]interface{}{"allowed_ips": "203.0.113.10", "allowed_domains": "owner.example"} + switch state { + case "disabled": + changes["is_active"] = 0 + case "expired": + changes["expiry_time"] = time.Now().Add(-time.Hour).UnixMilli() + case "over-quota": + changes["max_bandwidth"] = 1 + changes["current_flow"] = 2 + } + if err := agent.h.repo.DB().Model(&repo.PeerShare{}).Where("id = ?", runtime.ShareID).Updates(changes).Error; err != nil { + t.Fatal(err) + } + tests := []struct { + command string + allowed bool + }{ + {"release-role", true}, {"DeleteService", true}, {"DeleteChains", true}, {"DeleteLimiters", true}, {"DeleteCLimiters", true}, + {"AddService", false}, {"UpdateService", false}, {"ResumeService", false}, {"PauseService", false}, {"DeleteEverything", false}, + } + for _, test := range tests { + t.Run(test.command, func(t *testing.T) { + path := "/api/v1/federation/runtime/command" + body := fmt.Sprintf(`{"commandType":%q,"data":{"services":["70_1_0_tcp"]}}`, test.command) + if test.command == "release-role" { + path = "/api/v1/federation/runtime/release-role" + body = `{"reservationId":"role-reservation"}` + } + req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(body)) + req.Header.Set("Authorization", "Bearer role-recovery-token") + req.Header.Set("X-Panel-Domain", "owner.example") + req.RemoteAddr = "203.0.113.10:12345" + reached := false + next := func(w http.ResponseWriter, r *http.Request) { + reached = true + received, err := io.ReadAll(r.Body) + if err != nil || string(received) != body { + t.Errorf("auth consumed request body: %s %v", received, err) + } + response.WriteJSON(w, response.OKEmpty()) + } + agent.h.authPeer(next)(httptest.NewRecorder(), req) + if reached != test.allowed { + t.Fatalf("restricted %s %s allowed=%t want=%t", state, test.command, reached, test.allowed) + } + }) + } + for _, invalid := range []string{"token", "domain", "ip"} { + t.Run("reject-"+invalid, func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/federation/runtime/release-role", strings.NewReader(`{"reservationId":"role-reservation"}`)) + req.Header.Set("Authorization", "Bearer role-recovery-token") + req.Header.Set("X-Panel-Domain", "owner.example") + req.RemoteAddr = "203.0.113.10:12345" + switch invalid { + case "token": + req.Header.Set("Authorization", "Bearer wrong-token") + case "domain": + req.Header.Set("X-Panel-Domain", "other.example") + case "ip": + req.RemoteAddr = "198.51.100.1:12345" + } + reached := false + agent.h.authPeer(func(http.ResponseWriter, *http.Request) { reached = true })(httptest.NewRecorder(), req) + if reached { + t.Fatalf("cleanup bypassed %s authentication", invalid) + } + }) + } + // Exercise the actual release handler through auth, not only the gate. + req := httptest.NewRequest(http.MethodPost, "/api/v1/federation/runtime/release-role", strings.NewReader(`{"reservationId":"role-reservation"}`)) + req.Header.Set("Authorization", "Bearer role-recovery-token") + req.Header.Set("X-Panel-Domain", "owner.example") + req.RemoteAddr = "203.0.113.10:12345" + res := httptest.NewRecorder() + agent.h.authPeer(agent.h.federationRuntimeReleaseRole)(res, req) + var result response.R + if err := json.Unmarshal(res.Body.Bytes(), &result); err != nil || result.Code != 0 { + t.Fatalf("authenticated cleanup did not complete: %s %v", res.Body.String(), err) + } + stored, err := agent.h.repo.GetPeerShareRuntimeByID(runtime.ID) + if err != nil || stored.Status != 0 { + t.Fatalf("cleanup did not release runtime: %+v %v", stored, err) + } + }) + } +} + +func TestPeerShareOldReleaseIdentityCannotDeleteReusedReservation(t *testing.T) { + agent := newCleanupAgent(t) + old := roleRuntimeFixture(t, agent, "exit") + if err := agent.h.releasePeerShareRuntime(old); err != nil { + t.Fatal(err) + } + req := httptest.NewRequest(http.MethodPost, "/api/v1/federation/runtime/reserve-port", strings.NewReader(`{"resourceKey":"role-resource","requestedPort":31000,"protocol":"tls"}`)) + req.Header.Set("Authorization", "Bearer role-recovery-token") + res := httptest.NewRecorder() + agent.h.federationRuntimeReservePort(res, req) + var result response.R + if err := json.Unmarshal(res.Body.Bytes(), &result); err != nil || result.Code != 0 { + t.Fatalf("new generation reserve failed: %s %v", res.Body.String(), err) + } + fresh, err := agent.h.repo.GetPeerShareRuntimeByID(old.ID) + if err != nil { + t.Fatal(err) + } + if fresh.ReservationID == old.ReservationID { + t.Fatal("reservation identity was reused") + } + if code := roleRuntimeRequest(t, agent.h, fmt.Sprintf(`{"reservationId":%q,"role":"exit"}`, fresh.ReservationID), false); code != 0 { + t.Fatal("new generation apply failed") + } + fresh, err = agent.h.repo.GetPeerShareRuntimeByID(old.ID) + if err != nil { + t.Fatal(err) + } + if fresh.BindingID == old.BindingID || fresh.BindingID == "" { + t.Fatal("binding identity was reused") + } + deletesBefore := len(agent.commandsOfType("DeleteService")) + for _, body := range []string{ + fmt.Sprintf(`{"bindingId":%q,"reservationId":%q,"resourceKey":"role-resource"}`, old.BindingID, old.ReservationID), + fmt.Sprintf(`{"reservationId":%q,"resourceKey":"role-resource"}`, old.ReservationID), + } { + if code := roleRuntimeRequest(t, agent.h, body, true); code != 0 { + t.Fatal("obsolete release should acknowledge completion") + } + } + // Also cover an old lookup snapshot waiting behind reservation renewal. + if err := agent.h.releasePeerShareRuntime(old); err != nil { + t.Fatal(err) + } + after, err := agent.h.repo.GetPeerShareRuntimeByID(old.ID) + if err != nil || after.Status != 1 || after.BindingID != fresh.BindingID { + t.Fatalf("old cleanup released new generation: %+v %v", after, err) + } + if len(agent.commandsOfType("DeleteService")) != deletesBefore { + t.Fatal("old cleanup sent deletion for new generation") + } +} diff --git a/go-backend/internal/http/handler/peer_share_runtime_lifecycle.go b/go-backend/internal/http/handler/peer_share_runtime_lifecycle.go new file mode 100644 index 0000000..6421a9c --- /dev/null +++ b/go-backend/internal/http/handler/peer_share_runtime_lifecycle.go @@ -0,0 +1,190 @@ +package handler + +import ( + "encoding/json" + "errors" + "fmt" + "strings" + "sync" + "time" + + "go-backend/internal/store/repo" +) + +// Serialize desired-state changes with reconnect reconciliation. In particular, +// a stale reconnect snapshot must never revive an acknowledged release. +var peerRoleRuntimeMu sync.Mutex + +// applyPeerShareRoleRuntime requires peerRoleRuntimeMu. Persist the validated +// desired configuration before creating resources, so failed or interrupted +// commands remain recoverable rather than producing untracked listeners. +func (h *Handler) applyPeerShareRoleRuntime(runtime *repo.PeerShareRuntime) error { + if runtime.ReleasePending != 0 || runtime.Status != 1 { + return errors.New("runtime is being released") + } + if runtime.Role != "middle" && runtime.Role != "exit" { + return errors.New("invalid runtime role") + } + node, err := h.getNodeRecord(runtime.NodeID) + if err != nil { + return err + } + var targets []federationRuntimeTarget + if strings.TrimSpace(runtime.Target) != "" { + if err := json.Unmarshal([]byte(runtime.Target), &targets); err != nil { + return err + } + } + var chain map[string]interface{} + if runtime.Role == "middle" { + chain, err = buildFederationMiddleChainConfig(runtime.ChainName, runtime.ID, runtime.Protocol, runtime.Strategy, targets, node.InterfaceName) + if err != nil { + return err + } + } + service := buildFederationServiceConfig(runtime.ServiceName, fmt.Sprintf("%s:%d", node.TCPListenAddr, runtime.Port), runtime.Protocol, runtime.Role, runtime.ChainName, len(targets), node.InterfaceName) + runtime.Applied = 0 + runtime.UpdatedTime = time.Now().UnixMilli() + if err := h.repo.UpdatePeerShareRuntime(runtime); err != nil { + return err + } + if h.wsServer == nil { + return errors.New("node command transport unavailable") + } + if chain != nil { + if _, err := h.sendNodeCommand(runtime.NodeID, "UpdateChains", updateChainPayload(runtime.ChainName, chain), false, false); err != nil { + return err + } + } + // UpdateService and UpdateChains are upserts, including on an empty agent. + if _, err := h.sendNodeCommand(runtime.NodeID, "UpdateService", []map[string]interface{}{service}, false, false); err != nil { + return err + } + runtime.Applied = 1 + runtime.UpdatedTime = time.Now().UnixMilli() + return h.repo.UpdatePeerShareRuntime(runtime) +} + +func (h *Handler) releasePeerShareRuntime(runtime *repo.PeerShareRuntime) error { + peerRoleRuntimeMu.Lock() + defer peerRoleRuntimeMu.Unlock() + if runtime == nil { + return nil + } + current, err := h.repo.GetPeerShareRuntimeByID(runtime.ID) + if err != nil { + return err + } + if current != nil && (current.ReservationID != runtime.ReservationID || current.BindingID != runtime.BindingID) { + // A completed reservation may have been reused while this release + // waited for the mutation lock. Never release the new generation. + return nil + } + return h.releasePeerShareRuntimeLocked(current) +} + +func (h *Handler) releasePeerShareRuntimeLocked(runtime *repo.PeerShareRuntime) error { + if runtime == nil || runtime.Status == 0 { + return nil + } + if err := h.repo.SetPeerShareRuntimeReleasePending(runtime.ID); err != nil { + return err + } + runtime.ReleasePending = 1 + if runtime.Role == "forward" { + if _, _, scoped := parsePeerShareServiceName(runtime.ServiceName); scoped { + if err := h.releasePeerShareForwardRuntimeResources(runtime); err != nil { + return err + } + return h.repo.CompletePeerShareRuntimeRelease(runtime.ID) + } + // Older agents used unscoped names. Do not delete a name shared by + // another reservation or by a local forward during migration. + owners, err := h.repo.ListActiveForwardPeerShareRuntimesByNodeAndServiceName(runtime.NodeID, runtime.ServiceName) + if err != nil { + return err + } + for _, owner := range owners { + if owner.ShareID != runtime.ShareID { + return fmt.Errorf("legacy runtime %q has ambiguous ownership", runtime.ServiceName) + } + } + if forwardID, _, _, ok := parseFlowServiceIDs(normalizeForwardRuntimeServiceName(runtime.ServiceName)); ok { + local, err := h.repo.GetForwardRecord(forwardID) + if err != nil { + return err + } + if local != nil { + return fmt.Errorf("legacy runtime %q collides with a local forward", runtime.ServiceName) + } + } + } + // ServiceName is saved before apply; Applied=0 can therefore mean an + // unacknowledged command, and must not bypass deletion. + if strings.TrimSpace(runtime.ServiceName) != "" || strings.TrimSpace(runtime.ChainName) != "" { + if h.wsServer == nil { + return errors.New("node command transport unavailable") + } + if strings.TrimSpace(runtime.ServiceName) != "" { + names := []string{runtime.ServiceName} + if runtime.Role == "forward" { + names = append(names, runtime.ServiceName+"_tcp", runtime.ServiceName+"_udp") + } + if _, err := h.sendNodeCommand(runtime.NodeID, "DeleteService", map[string]interface{}{"services": names}, false, true); err != nil { + return err + } + } + if strings.TrimSpace(runtime.ChainName) != "" { + if _, err := h.sendNodeCommand(runtime.NodeID, "DeleteChains", map[string]interface{}{"chain": runtime.ChainName}, false, true); err != nil { + return err + } + } + } + return h.repo.CompletePeerShareRuntimeRelease(runtime.ID) +} + +func (h *Handler) reconcilePeerShareRoleRuntimesOnNode(nodeID int64) error { + return h.reconcilePeerShareRoleRuntimes(nodeID, false) +} + +// Maintenance retries only unfinished desired state. Replaying acknowledged +// services while an unrelated operation is failing can interrupt live traffic. +func (h *Handler) retryPendingPeerShareRoleRuntimesOnNode(nodeID int64) error { + return h.reconcilePeerShareRoleRuntimes(nodeID, true) +} + +func (h *Handler) reconcilePeerShareRoleRuntimes(nodeID int64, pendingOnly bool) error { + peerRoleRuntimeMu.Lock() + defer peerRoleRuntimeMu.Unlock() + runtimes, err := h.repo.ListActivePeerShareRuntimesByNode(nodeID) + if err != nil { + return err + } + var reconcileErr error + for i := range runtimes { + runtime := &runtimes[i] + if pendingOnly && runtime.Applied == 1 && runtime.ReleasePending == 0 { + continue + } + if runtime.ReleasePending != 0 { + reconcileErr = errors.Join(reconcileErr, h.releasePeerShareRuntimeLocked(runtime)) + continue + } + if runtime.Role != "middle" && runtime.Role != "exit" { + continue + } + share, err := h.repo.GetPeerShare(runtime.ShareID) + if err != nil { + reconcileErr = errors.Join(reconcileErr, err) + continue + } + if share == nil || share.IsActive != 1 || (share.ExpiryTime > 0 && share.ExpiryTime <= time.Now().UnixMilli()) || isPeerShareFlowExceeded(share) { + reconcileErr = errors.Join(reconcileErr, h.releasePeerShareRuntimeLocked(runtime)) + continue + } + if err := h.applyPeerShareRoleRuntime(runtime); err != nil { + reconcileErr = errors.Join(reconcileErr, fmt.Errorf("shared runtime %d: %w", runtime.ID, err)) + } + } + return reconcileErr +} diff --git a/go-backend/internal/http/handler/peer_share_runtime_lifecycle_test.go b/go-backend/internal/http/handler/peer_share_runtime_lifecycle_test.go new file mode 100644 index 0000000..f3ee5d3 --- /dev/null +++ b/go-backend/internal/http/handler/peer_share_runtime_lifecycle_test.go @@ -0,0 +1,356 @@ +package handler + +import ( + "encoding/json" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + "time" + + "go-backend/internal/store/repo" + "go-backend/internal/ws" + "gorm.io/gorm" +) + +func roleRuntimeFixture(t *testing.T, agent *cleanupAgent, role string) *repo.PeerShareRuntime { + t.Helper() + now := time.Now().UnixMilli() + share := &repo.PeerShare{Name: "role-recovery", NodeID: 1, Token: "role-recovery-token", PortRangeStart: 31000, PortRangeEnd: 31010, IsActive: 1, CreatedTime: now, UpdatedTime: now} + if err := agent.h.repo.CreatePeerShare(share); err != nil { + t.Fatal(err) + } + runtime := &repo.PeerShareRuntime{ShareID: share.ID, NodeID: 1, ReservationID: "role-reservation", ResourceKey: "role-resource", BindingID: "333", Role: role, ServiceName: "fed_svc_333", Protocol: "tls", Strategy: "round", Port: 31000, Applied: 1, Status: 1, CreatedTime: now, UpdatedTime: now} + if role == "middle" { + runtime.ChainName = "fed_chain_333" + runtime.Target = `[{"host":"127.0.0.1","port":32000,"protocol":"tls"}]` + } + if err := agent.h.repo.CreatePeerShareRuntime(runtime); err != nil { + t.Fatal(err) + } + // Runtime-generated names use the provider's globally unique runtime ID. + runtime.BindingID = fmt.Sprint(runtime.ID) + runtime.ServiceName = fmt.Sprintf("fed_svc_%d", runtime.ID) + if role == "middle" { + runtime.ChainName = federationRuntimeChainName(runtime.BindingID) + } + if err := agent.h.repo.UpdatePeerShareRuntime(runtime); err != nil { + t.Fatal(err) + } + return runtime +} + +func roleRuntimeRequest(t *testing.T, h *Handler, body string, release bool) int { + t.Helper() + req := httptest.NewRequest(http.MethodPost, "/runtime", strings.NewReader(body)) + req.Header.Set("Authorization", "Bearer role-recovery-token") + res := httptest.NewRecorder() + if release { + h.federationRuntimeReleaseRole(res, req) + } else { + h.federationRuntimeApplyRole(res, req) + } + var result struct { + Code int `json:"code"` + Msg string `json:"msg"` + } + if err := json.Unmarshal(res.Body.Bytes(), &result); err != nil { + t.Fatal(err) + } + t.Logf("runtime response: %s", res.Body.String()) + return result.Code +} + +func TestPeerShareAppliedRuntimeRepairsEmptyAgent(t *testing.T) { + for _, role := range []string{"exit", "middle"} { + t.Run(role, func(t *testing.T) { + agent := newCleanupAgent(t) + roleRuntimeFixture(t, agent, role) + if !agent.h.redeployNodeRuntimeAfterUpgrade(1) { + t.Fatal("reconnect reconciliation failed") + } + if got := len(agent.commandsOfType("UpdateService")); got != 1 { + t.Fatalf("reconnect did not restore service: %d", got) + } + targets := "" + if role == "middle" { + targets = `,"targets":[{"host":"127.0.0.1","port":32000,"protocol":"tls"}]` + } + if code := roleRuntimeRequest(t, agent.h, `{"reservationId":"role-reservation","role":"`+role+`","protocol":"tls"`+targets+`}`, false); code != 0 { + t.Fatalf("apply failed: %d", code) + } + if got := len(agent.commandsOfType("UpdateService")); got != 2 { + t.Fatalf("Applied=1 skipped service repair: %d", got) + } + if role == "middle" { + if got := len(agent.commandsOfType("UpdateChains")); got != 2 { + t.Fatalf("missing middle chains: %d", got) + } + agent.mu.Lock() + defer agent.mu.Unlock() + chainSeen := false + for _, cmd := range agent.commands { + if cmd.Type == "UpdateChains" { + chainSeen = true + } + if cmd.Type == "UpdateService" && !chainSeen { + t.Fatal("service started before chain") + } + } + } + }) + } +} + +func TestPeerShareReleaseOfflineKeepsPortAndReconnectDeletes(t *testing.T) { + agent := newCleanupAgent(t) + runtime := roleRuntimeFixture(t, agent, "middle") + liveServer := agent.h.wsServer + agent.h.wsServer = ws.NewServer(agent.h.repo, "offline-test") + if code := roleRuntimeRequest(t, agent.h, `{"reservationId":"role-reservation"}`, true); code == 0 { + t.Fatal("offline release reported success") + } + stored, err := agent.h.repo.GetPeerShareRuntimeByID(runtime.ID) + if err != nil || stored.Status != 1 || stored.ReleasePending != 1 { + t.Fatalf("pending release lost: %+v err=%v", stored, err) + } + occupied, err := agent.h.repo.ExistsActivePeerShareRuntimeOnNodePort(1, 31000) + if err != nil || !occupied { + t.Fatalf("pending release freed occupied port: %t %v", occupied, err) + } + if code := roleRuntimeRequest(t, agent.h, `{"reservationId":"role-reservation","role":"middle","targets":[{"host":"127.0.0.1","port":32000}]}`, false); code == 0 { + t.Fatal("pending release was revived by apply") + } + agent.h.wsServer = liveServer + if !agent.h.redeployNodeRuntimeAfterUpgrade(1) { + t.Fatal("reconnect cleanup failed") + } + stored, err = agent.h.repo.GetPeerShareRuntimeByID(runtime.ID) + if err != nil || stored.Status != 0 || stored.Applied != 0 || stored.ReleasePending != 0 { + t.Fatalf("release not completed: %+v %v", stored, err) + } + if len(agent.commandsOfType("DeleteService")) != 1 || len(agent.commandsOfType("DeleteChains")) != 1 { + t.Fatal("reconnect did not delete both service and chain") + } + if len(agent.commandsOfType("UpdateService")) != 0 || len(agent.commandsOfType("UpdateChains")) != 0 { + t.Fatal("reconnect revived pending release") + } + occupied, err = agent.h.repo.ExistsActivePeerShareRuntimeOnNodePort(1, 31000) + if err != nil || occupied { + t.Fatalf("acknowledged release still occupies port: %t %v", occupied, err) + } +} + +func TestPeerShareApplyPersistsBeforeCommand(t *testing.T) { + agent := newCleanupAgent(t) + runtime := roleRuntimeFixture(t, agent, "exit") + callback := "test:fail-role-desired-write" + if err := agent.h.repo.DB().Callback().Update().Before("gorm:update").Register(callback, func(tx *gorm.DB) { + if tx.Statement.Table == "peer_share_runtime" { + tx.AddError(errors.New("desired-state write failed")) + } + }); err != nil { + t.Fatal(err) + } + // This database is test-local; keep the callback registered through the + // websocket teardown so callback mutation cannot race node status writes. + if code := roleRuntimeRequest(t, agent.h, `{"reservationId":"role-reservation","role":"exit"}`, false); code == 0 { + t.Fatal("failed ownership persistence reported success") + } + if len(agent.commandsOfType("UpdateService")) != 0 { + t.Fatal("service started before ownership persisted") + } + stored, err := agent.h.repo.GetPeerShareRuntimeByID(runtime.ID) + if err != nil || stored.Applied != 1 { + t.Fatalf("old runtime changed on failed persistence: %+v %v", stored, err) + } +} + +func TestPeerShareApplyInvalidMiddleKeepsDesiredConfig(t *testing.T) { + agent := newCleanupAgent(t) + runtime := roleRuntimeFixture(t, agent, "middle") + if code := roleRuntimeRequest(t, agent.h, `{"reservationId":"role-reservation","role":"middle","targets":[]}`, false); code == 0 { + t.Fatal("invalid middle targets accepted") + } + stored, err := agent.h.repo.GetPeerShareRuntimeByID(runtime.ID) + if err != nil || stored.Target != runtime.Target || stored.Applied != 1 { + t.Fatalf("invalid update changed desired runtime: %+v %v", stored, err) + } + if len(agent.commandsOfType("UpdateService"))+len(agent.commandsOfType("UpdateChains")) != 0 { + t.Fatal("invalid config sent to agent") + } +} + +func TestPeerShareUnacknowledgedApplyRecoversOrReleases(t *testing.T) { + for _, release := range []bool{false, true} { + t.Run(fmt.Sprintf("release=%t", release), func(t *testing.T) { + agent := newCleanupAgent(t) + runtime := roleRuntimeFixture(t, agent, "middle") + liveServer := agent.h.wsServer + agent.h.wsServer = ws.NewServer(agent.h.repo, "offline-test") + if code := roleRuntimeRequest(t, agent.h, `{"reservationId":"role-reservation","role":"middle","targets":[{"host":"127.0.0.1","port":32001,"protocol":"tls"}]}`, false); code == 0 { + t.Fatal("offline apply reported success") + } + stored, err := agent.h.repo.GetPeerShareRuntimeByID(runtime.ID) + if err != nil || stored.Applied != 0 || !strings.Contains(stored.Target, "32001") || stored.ServiceName == "" { + t.Fatalf("unacknowledged desired state lost: %+v %v", stored, err) + } + if release { + if code := roleRuntimeRequest(t, agent.h, `{"reservationId":"role-reservation"}`, true); code == 0 { + t.Fatal("unacknowledged listener was freed offline") + } + } + agent.h.wsServer = liveServer + if !agent.h.redeployNodeRuntimeAfterUpgrade(1) { + t.Fatal("reconnect failed") + } + stored, err = agent.h.repo.GetPeerShareRuntimeByID(runtime.ID) + if err != nil { + t.Fatal(err) + } + if release { + if stored.Status != 0 || len(agent.commandsOfType("DeleteService")) != 1 || len(agent.commandsOfType("UpdateService")) != 0 { + t.Fatalf("unacknowledged apply revived after release: %+v", stored) + } + } else { + if stored.Status != 1 || stored.Applied != 1 || len(agent.commandsOfType("UpdateService")) != 1 { + t.Fatalf("unacknowledged desired state not recovered: %+v", stored) + } + } + }) + } +} + +func TestPeerShareDeleteOfflineRetainsDisabledShare(t *testing.T) { + agent := newCleanupAgent(t) + runtime := roleRuntimeFixture(t, agent, "exit") + liveServer := agent.h.wsServer + agent.h.wsServer = ws.NewServer(agent.h.repo, "offline-test") + req := httptest.NewRequest(http.MethodPost, "/share/delete", strings.NewReader(fmt.Sprintf(`{"id":%d}`, runtime.ShareID))) + res := httptest.NewRecorder() + agent.h.federationShareDelete(res, req) + var result struct { + Code int `json:"code"` + } + if err := json.Unmarshal(res.Body.Bytes(), &result); err != nil { + t.Fatal(err) + } + if result.Code == 0 { + t.Fatal("offline share delete reported success") + } + share, err := agent.h.repo.GetPeerShare(runtime.ShareID) + if err != nil || share == nil || share.IsActive != 0 { + t.Fatalf("cleanup retry record lost or still accepts allocations: %+v %v", share, err) + } + occupied, err := agent.h.repo.ExistsActivePeerShareRuntimeOnNodePort(1, 31000) + if err != nil || !occupied { + t.Fatal("offline share deletion released its port") + } + agent.h.wsServer = liveServer + if !agent.h.redeployNodeRuntimeAfterUpgrade(1) { + t.Fatal("disabled share cleanup did not retry") + } + if len(agent.commandsOfType("DeleteService")) != 1 || len(agent.commandsOfType("UpdateService")) != 0 { + t.Fatal("disabled share was recreated") + } +} + +func TestPeerShareLegacyReleaseWithoutLocalCollisionDeletesFamily(t *testing.T) { + agent := newCleanupAgent(t) + runtime := roleRuntimeFixture(t, agent, "forward") + runtime.ServiceName = "70_1_0" + if err := agent.h.repo.UpdatePeerShareRuntime(runtime); err != nil { + t.Fatal(err) + } + if err := agent.h.releasePeerShareRuntime(runtime); err != nil { + t.Fatal(err) + } + commands := agent.commandsOfType("DeleteService") + if len(commands) != 1 { + t.Fatalf("expected one family cleanup, got %d", len(commands)) + } + var payload struct { + Services []string `json:"services"` + } + if err := json.Unmarshal(commands[0].Data, &payload); err != nil { + t.Fatal(err) + } + got := strings.Join(payload.Services, ",") + if got != "70_1_0,70_1_0_tcp,70_1_0_udp" { + t.Fatalf("incomplete legacy cleanup: %s", got) + } + stored, err := agent.h.repo.GetPeerShareRuntimeByID(runtime.ID) + if err != nil || stored.Status != 0 || stored.ReleasePending != 0 { + t.Fatalf("legacy release incomplete: %+v %v", stored, err) + } +} + +func TestPeerSharePendingRoleRetryDoesNotReplaySuccessfulServices(t *testing.T) { + agent := newCleanupAgent(t) + runtime := roleRuntimeFixture(t, agent, "exit") + if err := agent.h.retryPendingPeerShareRoleRuntimesOnNode(1); err != nil { + t.Fatal(err) + } + if len(agent.commandsOfType("UpdateService")) != 0 { + t.Fatal("maintenance replayed an acknowledged service") + } + runtime.Applied = 0 + if err := agent.h.repo.UpdatePeerShareRuntime(runtime); err != nil { + t.Fatal(err) + } + if err := agent.h.retryPendingPeerShareRoleRuntimesOnNode(1); err != nil { + t.Fatal(err) + } + if len(agent.commandsOfType("UpdateService")) != 1 { + t.Fatal("maintenance did not retry unfinished service") + } + if err := agent.h.retryPendingPeerShareRoleRuntimesOnNode(1); err != nil { + t.Fatal(err) + } + if len(agent.commandsOfType("UpdateService")) != 1 { + t.Fatal("maintenance repeated a successful retry") + } +} + +func TestPeerShareResourceFailureDoesNotBlockRoleOrLocalRecovery(t *testing.T) { + agent := newCleanupAgent(t) + runtime := roleRuntimeFixture(t, agent, "exit") + if err := agent.h.repo.SavePeerShareResources([]repo.PeerShareResource{{ShareID: runtime.ShareID, NodeID: 1, Kind: "chain", OriginalName: "broken", RuntimeName: "broken", Config: "invalid-json", DesiredState: "active", Applied: 0, UpdatedTime: time.Now().UnixMilli()}}); err != nil { + t.Fatal(err) + } + var localQueries atomic.Int32 + if err := agent.h.repo.DB().Callback().Query().Before("gorm:query").Register("test:observe-local-recovery", func(tx *gorm.DB) { + if tx.Statement.Table == "chain_tunnel" || tx.Statement.Table == "forward_port" { + localQueries.Add(1) + } + }); err != nil { + t.Fatal(err) + } + if agent.h.redeployNodeRuntimeAfterUpgrade(1) { + t.Fatal("invalid shared resource was reported recovered") + } + if len(agent.commandsOfType("UpdateService")) != 1 { + t.Fatal("resource failure blocked independent role service recovery") + } + if localQueries.Load() < 2 { + t.Fatal("shared failure blocked independent local recovery") + } + // A shared failure is owned by pending maintenance, without full-redeploy + // timers that would periodically restart healthy listeners. + agent.h.onNodeOnline(1) + agent.h.upgradeMu.Lock() + _, fullQueued := agent.h.nodeOnlineRedeployQueued[1] + _, localQueued := agent.h.nodeLocalRuntimeRetryQueued[1] + agent.h.upgradeMu.Unlock() + if fullQueued || localQueued { + t.Fatal("shared-only failure queued a full or local redeploy") + } + before := len(agent.commandsOfType("UpdateService")) + agent.h.retryNodeLocalRuntime(1) + if len(agent.commandsOfType("UpdateService")) != before { + t.Fatal("local retry replayed shared listeners") + } +} diff --git a/go-backend/internal/http/handler/upgrade.go b/go-backend/internal/http/handler/upgrade.go index f7164b0..f4d8da3 100644 --- a/go-backend/internal/http/handler/upgrade.go +++ b/go-backend/internal/http/handler/upgrade.go @@ -398,15 +398,80 @@ func (h *Handler) consumeNodePendingUpgradeRedeploy(nodeID int64) bool { } func (h *Handler) onNodeOnline(nodeID int64) { + if h == nil || h.repo == nil || h.wsServer == nil { + return + } + if node, err := h.getNodeRecord(nodeID); err == nil && (node == nil || node.Status != 1) { + return // The next connection will resume any pending reconciliation. + } if !h.startNodeOnlineRedeploy(nodeID, time.Now()) { return } defer h.finishNodeOnlineRedeploy(nodeID) - // Reconcile node runtime on the first reconnect, but suppress rapid flapping - // so websocket churn does not trigger repeated full redeploy storms. - if !h.redeployNodeRuntimeAfterUpgrade(nodeID) { - h.markNodePendingUpgradeRedeploy(nodeID) + // A fresh agent needs a full restore. Shared failures remain persisted and + // are retried by maintenance without restarting acknowledged services. + h.reconcileSharedNodeRuntime(nodeID) + if !h.redeployLocalNodeRuntime(nodeID) { + h.scheduleNodeLocalRuntimeRetry(nodeID) + } else { + h.clearNodeLocalRuntimeRetry(nodeID) + } +} + +// Local retries are separate from reconnect reconciliation: a failed local +// forward must not cause every shared listener to be replayed every 30 seconds. +func (h *Handler) scheduleNodeLocalRuntimeRetry(nodeID int64) { + h.upgradeMu.Lock() + defer h.upgradeMu.Unlock() + if h.nodeLocalRuntimeRetryQueued == nil { + h.nodeLocalRuntimeRetryQueued = make(map[int64]struct{}) + } + if _, queued := h.nodeLocalRuntimeRetryQueued[nodeID]; queued { + return + } + h.nodeLocalRuntimeRetryQueued[nodeID] = struct{}{} + time.AfterFunc(nodeOnlineRedeployCooldown, func() { + h.upgradeMu.Lock() + _, queued := h.nodeLocalRuntimeRetryQueued[nodeID] + delete(h.nodeLocalRuntimeRetryQueued, nodeID) + h.upgradeMu.Unlock() + if !queued { + return + } + h.retryNodeLocalRuntime(nodeID) + }) +} + +func (h *Handler) clearNodeLocalRuntimeRetry(nodeID int64) { + h.upgradeMu.Lock() + delete(h.nodeLocalRuntimeRetryQueued, nodeID) + h.upgradeMu.Unlock() +} + +func (h *Handler) retryNodeLocalRuntime(nodeID int64) { + if h == nil || h.repo == nil || h.wsServer == nil { + return + } + if node, err := h.getNodeRecord(nodeID); err == nil && (node == nil || node.Status != 1) { + return + } + h.upgradeMu.Lock() + if _, inFlight := h.nodeOnlineRedeploying[nodeID]; inFlight { + h.upgradeMu.Unlock() + h.scheduleNodeLocalRuntimeRetry(nodeID) + return + } + if h.nodeOnlineRedeploying == nil { + h.nodeOnlineRedeploying = make(map[int64]struct{}) + } + h.nodeOnlineRedeploying[nodeID] = struct{}{} + h.upgradeMu.Unlock() + defer h.finishNodeOnlineRedeploy(nodeID) + if !h.redeployLocalNodeRuntime(nodeID) { + h.scheduleNodeLocalRuntimeRetry(nodeID) + } else { + h.clearNodeLocalRuntimeRetry(nodeID) } } @@ -503,6 +568,25 @@ func (h *Handler) finishNodeOnlineRedeploy(nodeID int64) { } func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) bool { + sharedOK := h.reconcileSharedNodeRuntime(nodeID) + localOK := h.redeployLocalNodeRuntime(nodeID) + return sharedOK && localOK +} + +func (h *Handler) reconcileSharedNodeRuntime(nodeID int64) bool { + succeeded := true + if err := h.reconcilePeerShareResourcesOnNode(nodeID); err != nil { + fmt.Printf("reconnect shared resource reconciliation failed on node %d: %v\n", nodeID, err) + succeeded = false + } + if err := h.reconcilePeerShareRoleRuntimesOnNode(nodeID); err != nil { + fmt.Printf("reconnect shared role reconciliation failed on node %d: %v\n", nodeID, err) + succeeded = false + } + return succeeded +} + +func (h *Handler) redeployLocalNodeRuntime(nodeID int64) bool { tunnelIDs, err := h.repo.ListActiveTunnelIDsByNode(nodeID) if err != nil { fmt.Printf("post-upgrade redeploy: list tunnels for node %d failed: %v\n", nodeID, err) @@ -567,6 +651,7 @@ func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]stru return true } + permanentFailure := false const maxRetries = 3 baseDelay := time.Second @@ -580,7 +665,8 @@ func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]stru delete(tunnelFailed, tunnelID) fmt.Printf("post-upgrade redeploy retry: tunnel %d succeeded on node %d (attempt %d)\n", tunnelID, nodeID, attempt) } else if !isRetryableError(err) { - delete(tunnelFailed, tunnelID) // Non-retryable, don't retry again + permanentFailure = true + delete(tunnelFailed, tunnelID) // Preserve failure while avoiding immediate retries. } else { fmt.Printf("post-upgrade redeploy retry: tunnel %d still failing on node %d (attempt %d): %v\n", tunnelID, nodeID, attempt, err) } @@ -596,7 +682,7 @@ func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]stru if err := h.syncForwardServices(ff.forward, "UpdateService", true); err == nil { fmt.Printf("post-upgrade redeploy retry: forward %d succeeded on node %d (attempt %d)\n", ff.id, nodeID, attempt) } else if !isRetryableError(err) { - // Non-retryable, drop it + permanentFailure = true // Keep reconciliation pending for a later reconnect. } else { stillFailed = append(stillFailed, ff) fmt.Printf("post-upgrade redeploy retry: forward %d still failing on node %d (attempt %d): %v\n", ff.id, nodeID, attempt, err) @@ -606,7 +692,7 @@ func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]stru if len(tunnelFailed) == 0 && len(failedForwards) == 0 { fmt.Printf("post-upgrade redeploy retry: all items recovered on node %d\n", nodeID) - return true + return !permanentFailure } } diff --git a/go-backend/internal/http/handler/upgrade_test.go b/go-backend/internal/http/handler/upgrade_test.go index abe8d33..0482e25 100644 --- a/go-backend/internal/http/handler/upgrade_test.go +++ b/go-backend/internal/http/handler/upgrade_test.go @@ -12,7 +12,7 @@ func TestStartNodeOnlineRedeploySkipsRecentReconnects(t *testing.T) { nodeOnlineRedeployQueued: map[int64]struct{}{}, nodeOnlineRedeploying: map[int64]struct{}{}, } - now := time.Unix(1_777_176_720, 0) + now := time.Now() if !h.startNodeOnlineRedeploy(54, now) { t.Fatalf("expected first reconnect to redeploy") @@ -34,7 +34,7 @@ func TestStartNodeOnlineRedeployAllowsPendingUpgradeDuringCooldown(t *testing.T) nodeOnlineRedeployQueued: map[int64]struct{}{}, nodeOnlineRedeploying: map[int64]struct{}{}, } - now := time.Unix(1_777_176_720, 0) + now := time.Now() if !h.startNodeOnlineRedeploy(54, now) { t.Fatalf("expected first reconnect to redeploy") @@ -57,7 +57,7 @@ func TestStartNodeOnlineRedeployQueuesCooldownReconnect(t *testing.T) { nodeOnlineRedeployQueued: map[int64]struct{}{}, nodeOnlineRedeploying: map[int64]struct{}{}, } - now := time.Unix(1_777_176_720, 0) + now := time.Now() if !h.startNodeOnlineRedeploy(54, now) { t.Fatalf("expected first reconnect to redeploy") @@ -67,7 +67,10 @@ func TestStartNodeOnlineRedeployQueuesCooldownReconnect(t *testing.T) { if h.startNodeOnlineRedeploy(54, now.Add(5*time.Second)) { t.Fatalf("expected cooldown reconnect to skip immediate redeploy") } - if _, queued := h.nodeOnlineRedeployQueued[54]; !queued { + h.upgradeMu.Lock() + _, queued := h.nodeOnlineRedeployQueued[54] + h.upgradeMu.Unlock() + if !queued { t.Fatalf("expected cooldown reconnect to queue a follow-up redeploy") } } @@ -79,7 +82,7 @@ func TestStartNodeOnlineRedeployKeepsPendingUpgradeWhileInFlight(t *testing.T) { nodeOnlineRedeployQueued: map[int64]struct{}{}, nodeOnlineRedeploying: map[int64]struct{}{}, } - now := time.Unix(1_777_176_720, 0) + now := time.Now() if !h.startNodeOnlineRedeploy(54, now) { t.Fatalf("expected first reconnect to redeploy") @@ -96,7 +99,7 @@ func TestStartNodeOnlineRedeployKeepsPendingUpgradeWhileInFlight(t *testing.T) { } func TestNextNodeOnlineRedeployFireAtDefersExpiredInFlightReconnect(t *testing.T) { - now := time.Unix(1_777_176_720, 0) + now := time.Now() last := now.Add(-nodeOnlineRedeployCooldown - 5*time.Second) fireAt, start := nextNodeOnlineRedeployFireAt(last, now, false, true) diff --git a/go-backend/internal/store/model/federation_release.go b/go-backend/internal/store/model/federation_release.go new file mode 100644 index 0000000..eafbff5 --- /dev/null +++ b/go-backend/internal/store/model/federation_release.go @@ -0,0 +1,15 @@ +package model + +// FederationPendingRelease keeps rollback work after a failed remote release. +// It is independent of tunnels, which may never have committed or be deleted. +type FederationPendingRelease struct { + ID string `gorm:"primaryKey;size:64"` + RemoteURL string `gorm:"not null"` + RemoteToken string `gorm:"not null"` + BindingID string + ReservationID string + ResourceKey string + CreatedTime int64 +} + +func (FederationPendingRelease) TableName() string { return "federation_pending_release" } diff --git a/go-backend/internal/store/model/model.go b/go-backend/internal/store/model/model.go index afca82b..e46154d 100644 --- a/go-backend/internal/store/model/model.go +++ b/go-backend/internal/store/model/model.go @@ -354,23 +354,24 @@ type PeerShare struct { func (PeerShare) TableName() string { return "peer_share" } type PeerShareRuntime struct { - ID int64 `gorm:"primaryKey;autoIncrement"` - ShareID int64 `gorm:"column:share_id;not null;index:idx_peer_share_runtime_share_node_status"` - NodeID int64 `gorm:"column:node_id;not null;index:idx_peer_share_runtime_share_node_status"` - ReservationID string `gorm:"column:reservation_id;type:text;not null;uniqueIndex"` - ResourceKey string `gorm:"column:resource_key;type:text;not null;uniqueIndex"` - BindingID string `gorm:"column:binding_id;type:text;not null;default:'';index:idx_peer_share_runtime_binding_id"` - Role string `gorm:"type:text;not null;default:''"` - ChainName string `gorm:"column:chain_name;type:text;not null;default:''"` - ServiceName string `gorm:"column:service_name;type:text;not null;default:''"` - Protocol string `gorm:"type:text;not null;default:'tls'"` - Strategy string `gorm:"type:text;not null;default:'round'"` - Port int `gorm:"not null;default:0"` - Target string `gorm:"type:text;not null;default:''"` - Applied int `gorm:"not null;default:0"` - Status int `gorm:"not null;default:1;index:idx_peer_share_runtime_share_node_status"` - CreatedTime int64 `gorm:"column:created_time;not null"` - UpdatedTime int64 `gorm:"column:updated_time;not null"` + ID int64 `gorm:"primaryKey;autoIncrement"` + ShareID int64 `gorm:"column:share_id;not null;index:idx_peer_share_runtime_share_node_status"` + NodeID int64 `gorm:"column:node_id;not null;index:idx_peer_share_runtime_share_node_status"` + ReservationID string `gorm:"column:reservation_id;type:text;not null;uniqueIndex"` + ResourceKey string `gorm:"column:resource_key;type:text;not null;uniqueIndex"` + BindingID string `gorm:"column:binding_id;type:text;not null;default:'';index:idx_peer_share_runtime_binding_id"` + Role string `gorm:"type:text;not null;default:''"` + ChainName string `gorm:"column:chain_name;type:text;not null;default:''"` + ServiceName string `gorm:"column:service_name;type:text;not null;default:''"` + Protocol string `gorm:"type:text;not null;default:'tls'"` + Strategy string `gorm:"type:text;not null;default:'round'"` + Port int `gorm:"not null;default:0"` + Target string `gorm:"type:text;not null;default:''"` + Applied int `gorm:"not null;default:0"` + ReleasePending int `gorm:"column:release_pending;not null;default:0"` + Status int `gorm:"not null;default:1;index:idx_peer_share_runtime_share_node_status"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime int64 `gorm:"column:updated_time;not null"` } func (PeerShareRuntime) TableName() string { return "peer_share_runtime" } @@ -816,3 +817,23 @@ type TunnelQuality struct { } func (TunnelQuality) TableName() string { return "tunnel_quality" } + +// PeerShareResource is the durable desired state for a namespaced peer command. +// Rows are retained as tombstones until deletion has been acknowledged. +type PeerShareResource struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + ShareID int64 `gorm:"column:share_id;not null;uniqueIndex:idx_peer_share_resource_key"` + NodeID int64 `gorm:"column:node_id;not null;index"` + Kind string `gorm:"type:text;not null;uniqueIndex:idx_peer_share_resource_key"` + OriginalName string `gorm:"column:original_name;type:text;not null;uniqueIndex:idx_peer_share_resource_key"` + RuntimeName string `gorm:"column:runtime_name;type:text;not null;index"` + LegacyNames string `gorm:"column:legacy_names;type:text;not null;default:''"` + LegacyServiceBase string `gorm:"column:legacy_service_base;type:text;not null;default:''"` + ReleaseLegacyFamily bool `gorm:"column:release_legacy_family;not null;default:false"` + Config string `gorm:"type:text;not null;default:''"` + DesiredState string `gorm:"column:desired_state;type:text;not null;default:'active'"` + Applied int `gorm:"not null;default:0"` + UpdatedTime int64 `gorm:"column:updated_time;not null"` +} + +func (PeerShareResource) TableName() string { return "peer_share_resource" } diff --git a/go-backend/internal/store/repo/peer_share_resources.go b/go-backend/internal/store/repo/peer_share_resources.go new file mode 100644 index 0000000..08edf1f --- /dev/null +++ b/go-backend/internal/store/repo/peer_share_resources.go @@ -0,0 +1,55 @@ +package repo + +import ( + "errors" + "go-backend/internal/store/model" + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +func (r *Repository) SavePeerShareResources(items []PeerShareResource) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + if len(items) == 0 { + return nil + } + return r.db.Transaction(func(tx *gorm.DB) error { + for i := range items { + if err := tx.Clauses(clause.OnConflict{Columns: []clause.Column{{Name: "share_id"}, {Name: "kind"}, {Name: "original_name"}}, DoUpdates: clause.AssignmentColumns([]string{"node_id", "runtime_name", "legacy_names", "legacy_service_base", "release_legacy_family", "config", "desired_state", "applied", "updated_time"})}).Create(&items[i]).Error; err != nil { + return err + } + } + return nil + }) +} + +func (r *Repository) GetPeerShareResource(shareID int64, kind, originalName string) (*PeerShareResource, error) { + var item model.PeerShareResource + err := r.db.Where("share_id = ? AND kind = ? AND original_name = ?", shareID, kind, originalName).First(&item).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + return &item, err +} + +func (r *Repository) ListPeerShareResourcesByNode(nodeID int64) ([]PeerShareResource, error) { + var items []PeerShareResource + err := r.db.Where("node_id = ?", nodeID).Order("id").Find(&items).Error + return items, err +} + +func (r *Repository) MarkPeerShareResourceApplied(shareID int64, kind, name string) error { + return r.db.Model(&model.PeerShareResource{}).Where("share_id = ? AND kind = ? AND original_name = ?", shareID, kind, name).Update("applied", 1).Error +} + +func (r *Repository) ClearPeerShareResourceLegacyNames(shareID int64, kind, name string) error { + return r.db.Model(&model.PeerShareResource{}).Where("share_id = ? AND kind = ? AND original_name = ?", shareID, kind, name).Update("legacy_names", "").Error +} + +func (r *Repository) WithPeerShareResourceTransaction(fn func(*Repository) error) error { + return r.db.Transaction(func(tx *gorm.DB) error { return fn(&Repository{db: tx, dbPath: r.dbPath}) }) +} +func (r *Repository) ClearPeerShareResourceLegacyFamily(shareID int64, base string) error { + return r.db.Model(&model.PeerShareResource{}).Where("share_id = ? AND legacy_service_base = ?", shareID, base).Update("legacy_service_base", "").Error +} diff --git a/go-backend/internal/store/repo/peer_share_runtime_lifecycle.go b/go-backend/internal/store/repo/peer_share_runtime_lifecycle.go new file mode 100644 index 0000000..2eda5be --- /dev/null +++ b/go-backend/internal/store/repo/peer_share_runtime_lifecycle.go @@ -0,0 +1,32 @@ +package repo + +import ( + "errors" + "go-backend/internal/store/model" + "time" +) + +// A pending release remains active until the agent acknowledges deletion. This +// keeps its port reserved even if the control connection is unavailable. +func (r *Repository) SetPeerShareRuntimeReleasePending(id int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.PeerShareRuntime{}).Where("id = ? AND status = 1", id).Updates(map[string]interface{}{"release_pending": 1, "updated_time": time.Now().UnixMilli()}).Error +} + +func (r *Repository) CompletePeerShareRuntimeRelease(id int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.PeerShareRuntime{}).Where("id = ?", id).Updates(map[string]interface{}{"status": 0, "applied": 0, "release_pending": 0, "updated_time": time.Now().UnixMilli()}).Error +} + +func (r *Repository) ListActivePeerShareRuntimesByNode(nodeID int64) ([]model.PeerShareRuntime, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var items []model.PeerShareRuntime + err := r.db.Where("node_id = ? AND status = 1", nodeID).Order("release_pending DESC, id ASC").Find(&items).Error + return items, err +} diff --git a/go-backend/internal/store/repo/repository.go b/go-backend/internal/store/repo/repository.go index a238d63..ef86b29 100644 --- a/go-backend/internal/store/repo/repository.go +++ b/go-backend/internal/store/repo/repository.go @@ -43,6 +43,7 @@ type UserForwardDetail = model.UserForwardDetail type StatisticsFlow = model.StatisticsFlow type Node = model.Node type PeerShare = model.PeerShare +type PeerShareResource = model.PeerShareResource type PeerShareRuntime = model.PeerShareRuntime type FederationTunnelBinding = model.FederationTunnelBinding type BackupData = model.BackupData @@ -309,7 +310,9 @@ func autoMigrateAll(db *gorm.DB) error { &model.ViteConfig{}, &model.PeerShare{}, &model.PeerShareRuntime{}, + &model.PeerShareResource{}, &model.FederationTunnelBinding{}, + &model.FederationPendingRelease{}, &model.Announcement{}, &model.SchemaVersion{}, &model.NodeMetric{}, @@ -1514,7 +1517,12 @@ func (r *Repository) DeletePeerShare(id int64) error { return errors.New("repository not initialized") } return r.db.Transaction(func(tx *gorm.DB) error { - tx.Where("share_id = ?", id).Delete(&model.PeerShareRuntime{}) + if err := tx.Where("share_id = ?", id).Delete(&model.PeerShareRuntime{}).Error; err != nil { + return err + } + if err := tx.Where("share_id = ?", id).Delete(&model.PeerShareResource{}).Error; err != nil { + return err + } return tx.Where("id = ?", id).Delete(&model.PeerShare{}).Error }) } @@ -1607,7 +1615,8 @@ func (r *Repository) UpdatePeerShareRuntime(item *model.PeerShareRuntime) error return errors.New("runtime item is nil") } return r.db.Model(&model.PeerShareRuntime{}).Where("id = ?", item.ID).Updates(map[string]interface{}{ - "binding_id": item.BindingID, "role": item.Role, + "reservation_id": item.ReservationID, + "binding_id": item.BindingID, "role": item.Role, "chain_name": item.ChainName, "service_name": item.ServiceName, "protocol": item.Protocol, "strategy": item.Strategy, "port": item.Port, "target": item.Target, @@ -1755,35 +1764,19 @@ func (r *Repository) ListActiveForwardPeerShareRuntimesByNodeAndServiceName(node return items, nil } -func (r *Repository) ListActiveForwardPeerShareRuntimeServiceNamesByNode(nodeID int64) ([]string, error) { +func (r *Repository) ListActiveForwardPeerShareRuntimesByNode(nodeID int64) ([]model.PeerShareRuntime, error) { if r == nil || r.db == nil { return nil, errors.New("repository not initialized") } - var names []string - err := r.db.Model(&model.PeerShareRuntime{}). - Where("node_id = ? AND status = 1 AND role = ? AND service_name <> ''", nodeID, "forward"). - Pluck("service_name", &names).Error + var items []model.PeerShareRuntime + err := r.db.Where("node_id = ? AND status = 1 AND role = ?", nodeID, "forward").Find(&items).Error if err != nil { return nil, err } - if names == nil { - names = make([]string, 0) + if items == nil { + items = make([]model.PeerShareRuntime, 0) } - return names, nil -} - -func (r *Repository) HasRecentUnboundForwardPeerShareRuntimeOnNode(nodeID int64, minUpdatedTime int64) (bool, error) { - if r == nil || r.db == nil { - return false, errors.New("repository not initialized") - } - var count int64 - err := r.db.Model(&model.PeerShareRuntime{}). - Where("node_id = ? AND status = 1 AND role = ? AND applied = 0 AND updated_time >= ? AND (service_name = '' OR service_name IS NULL)", nodeID, "forward", minUpdatedTime). - Count(&count).Error - if err != nil { - return false, err - } - return count > 0, nil + return items, nil } func (r *Repository) GetActiveForwardPeerShareRuntimeByPort(shareID int64, port int) (*model.PeerShareRuntime, error) { diff --git a/go-backend/internal/store/repo/repository_federation_cleanup.go b/go-backend/internal/store/repo/repository_federation_cleanup.go new file mode 100644 index 0000000..554b3f3 --- /dev/null +++ b/go-backend/internal/store/repo/repository_federation_cleanup.go @@ -0,0 +1,55 @@ +package repo + +import ( + "crypto/sha256" + "fmt" + + "go-backend/internal/store/model" + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +const FederationBindingPendingRelease = 2 + +func (r *Repository) ListTunnelChainNodeIDsTx(tx *gorm.DB, tunnelID int64) ([]int64, error) { + var ids []int64 + err := tx.Model(&model.ChainTunnel{}).Where("tunnel_id = ?", tunnelID).Pluck("node_id", &ids).Error + return ids, err +} + +func (r *Repository) ListFederationTunnelBindingsForCleanup(tunnelID int64) ([]model.FederationTunnelBinding, error) { + var rows []model.FederationTunnelBinding + err := r.db.Where("tunnel_id = ? AND status IN ?", tunnelID, []int{1, FederationBindingPendingRelease}).Order("id").Find(&rows).Error + return rows, err +} + +func (r *Repository) ListPendingFederationTunnelBindings() ([]model.FederationTunnelBinding, error) { + var rows []model.FederationTunnelBinding + err := r.db.Where("status = ?", FederationBindingPendingRelease).Order("id").Find(&rows).Error + return rows, err +} + +func (r *Repository) MarkFederationTunnelBindingPendingRelease(id int64) error { + return r.db.Model(&model.FederationTunnelBinding{}).Where("id = ?", id). + Updates(map[string]interface{}{"status": FederationBindingPendingRelease, "updated_time": unixMilliNow()}).Error +} + +func (r *Repository) DeleteFederationTunnelBinding(id int64) error { + return r.db.Where("id = ?", id).Delete(&model.FederationTunnelBinding{}).Error +} + +func (r *Repository) SavePendingFederationRelease(item *model.FederationPendingRelease) error { + item.ID = fmt.Sprintf("%x", sha256.Sum256([]byte(item.RemoteURL+"\n"+item.BindingID+"\n"+item.ReservationID+"\n"+item.ResourceKey))) + item.CreatedTime = unixMilliNow() + return r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(item).Error +} + +func (r *Repository) ListPendingFederationReleases() ([]model.FederationPendingRelease, error) { + var rows []model.FederationPendingRelease + err := r.db.Order("created_time, id").Find(&rows).Error + return rows, err +} + +func (r *Repository) DeletePendingFederationRelease(id string) error { + return r.db.Where("id = ?", id).Delete(&model.FederationPendingRelease{}).Error +} diff --git a/go-backend/internal/store/repo/repository_federation_reconcile.go b/go-backend/internal/store/repo/repository_federation_reconcile.go new file mode 100644 index 0000000..10f4d55 --- /dev/null +++ b/go-backend/internal/store/repo/repository_federation_reconcile.go @@ -0,0 +1,35 @@ +package repo + +import ( + "sort" + + "go-backend/internal/store/model" +) + +// Only retry unfinished operations. Successful desired state must be replayed +// on a real reconnect, not every maintenance tick while the agent stays online. +func (r *Repository) ListPendingPeerShareNodeIDs() ([]int64, error) { + var resourceNodes, runtimeNodes []int64 + if err := r.db.Model(&model.PeerShareResource{}).Where("applied = 0").Distinct("node_id").Pluck("node_id", &resourceNodes).Error; err != nil { + return nil, err + } + if err := r.db.Model(&model.PeerShareRuntime{}). + Where("status = 1 AND (release_pending <> 0 OR (applied = 0 AND service_name <> '' AND role IN ?))", []string{"middle", "exit"}). + Distinct("node_id").Pluck("node_id", &runtimeNodes).Error; err != nil { + return nil, err + } + seen := make(map[int64]struct{}) + for _, ids := range [][]int64{resourceNodes, runtimeNodes} { + for _, id := range ids { + if id > 0 { + seen[id] = struct{}{} + } + } + } + out := make([]int64, 0, len(seen)) + for id := range seen { + out = append(out, id) + } + sort.Slice(out, func(i, j int) bool { return out[i] < out[j] }) + return out, nil +} diff --git a/go-backend/tests/contract/federation_dual_panel_contract_test.go b/go-backend/tests/contract/federation_dual_panel_contract_test.go index 7c1e40c..1e6530f 100644 --- a/go-backend/tests/contract/federation_dual_panel_contract_test.go +++ b/go-backend/tests/contract/federation_dual_panel_contract_test.go @@ -743,7 +743,7 @@ func TestFederationRuntimeCommandPortRangeEnforcement(t *testing.T) { if err := json.NewDecoder(res.Body).Decode(&out); err != nil { t.Fatalf("decode response: %v", err) } - if out.Code != 0 { - t.Fatalf("expected code 0 for reload command, got %d (msg: %s)", out.Code, out.Msg) + if out.Code == 0 { + t.Fatal("a shared-node token must not reload the entire provider node") } } diff --git a/go-backend/tests/contract/flow_upload_batch_contract_test.go b/go-backend/tests/contract/flow_upload_batch_contract_test.go index 435df2f..ce7488b 100644 --- a/go-backend/tests/contract/flow_upload_batch_contract_test.go +++ b/go-backend/tests/contract/flow_upload_batch_contract_test.go @@ -38,6 +38,9 @@ func TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately(t *testing if err := repo.DB().Create(forward).Error; err != nil { t.Fatalf("seed forward: %v", err) } + if err := repo.DB().Create(&model.ForwardPort{ForwardID: forward.ID, NodeID: node.ID, Port: 10000}).Error; err != nil { + t.Fatalf("seed forward node ownership: %v", err) + } if err := repo.DB().Exec(`INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time) VALUES(2, 1, 0, ?, ?, ?, ?, 0, 0, '', ?, ?)`, bytesPerGB-100, bytesPerGB-100, dayKey, monthKey, nowMs, nowMs).Error; err != nil { t.Fatalf("insert user_quota: %v", err) } diff --git a/go-backend/tests/contract/tunnel_metrics_ingestion_contract_test.go b/go-backend/tests/contract/tunnel_metrics_ingestion_contract_test.go index b53136f..0c2209c 100644 --- a/go-backend/tests/contract/tunnel_metrics_ingestion_contract_test.go +++ b/go-backend/tests/contract/tunnel_metrics_ingestion_contract_test.go @@ -58,6 +58,9 @@ func TestFlowUploadInsertsTunnelMetrics(t *testing.T) { if err := repo.DB().Create(forward).Error; err != nil { t.Fatalf("seed forward: %v", err) } + if err := repo.DB().Create(&model.ForwardPort{ForwardID: forward.ID, NodeID: node.ID, Port: 10000}).Error; err != nil { + t.Fatalf("seed forward node ownership: %v", err) + } serviceName := jsonNumber(forward.ID) + "_123_0" body, _ := json.Marshal([]map[string]interface{}{{ diff --git a/go-gost/main.go b/go-gost/main.go index d026b5f..4e79057 100644 --- a/go-gost/main.go +++ b/go-gost/main.go @@ -125,11 +125,13 @@ func main() { distro := socket.DetectDistro() fullVersion := fmt.Sprintf("%s (%s/%s)", version, distro, runtime.GOARCH) - wsReporter := socket.StartWebSocketReporterWithConfig(config.Addr, config.Secret, config.Http, config.Tls, config.Socks, fullVersion) - defer wsReporter.Stop() service.SetHTTPReportURL(config.Addr, config.Secret) - p := &program{} + p := &program{ + startReporter: func() reporter { + return socket.StartWebSocketReporterWithConfig(config.Addr, config.Secret, config.Http, config.Tls, config.Socks, fullVersion) + }, + } if err := svc.Run(p); err != nil { logger.Default().Fatal(err) } diff --git a/go-gost/program.go b/go-gost/program.go index eae3d1b..4b425c3 100644 --- a/go-gost/program.go +++ b/go-gost/program.go @@ -3,6 +3,14 @@ package main import ( "context" "errors" + "net" + "net/http" + "os" + "os/signal" + "strings" + "syscall" + "time" + "github.com/go-gost/core/auth" "github.com/go-gost/core/logger" "github.com/go-gost/core/service" @@ -18,20 +26,22 @@ import ( xservice "github.com/go-gost/x/service" "github.com/go-gost/x/socket" "github.com/judwhite/go-svc" - "net/http" - "os" - "os/signal" - "strings" - "syscall" - "time" ) -type program struct { - srvApi service.Service - srvMetrics service.Service - srvProfiling *http.Server +type reporter interface { + Stop() +} - cancel context.CancelFunc +type program struct { + startReporter func() reporter + reporter reporter + srvApi service.Service + srvMetrics service.Service + srvProfiling *http.Server + profilingListener net.Listener + + cancel context.CancelFunc + stopped bool } func (p *program) Init(env svc.Environment) error { @@ -48,7 +58,15 @@ func (p *program) Init(env svc.Environment) error { return nil } -func (p *program) Start() error { +func (p *program) Start() (err error) { + unlock := config.LockMutation() + defer unlock() + p.stopped = false + defer func() { + if err != nil { + p.stopRuntime() + } + }() cfg, err := parser.Parse() if err != nil { return err @@ -61,23 +79,28 @@ func (p *program) Start() error { os.Exit(0) } - config.Set(cfg) - if err := loader.Load(cfg); err != nil { return err } - // Enable config persistence after initial load so runtime mutations - // (AddService, UpdateService, DeleteService, etc.) are saved to disk. - socket.EnableConfigPersist() - if err := p.run(cfg); err != nil { return err } + config.Set(cfg) + socket.EnableConfigPersist() + ctx, cancel := context.WithCancel(context.Background()) p.cancel = cancel - go p.reload(ctx) + c := make(chan os.Signal, 1) + signal.Notify(c, syscall.SIGHUP) + go p.reload(ctx, c) + + // A connected panel may immediately send commands. Only expose the agent + // after initial config loading, runtime startup and persistence are ready. + if p.startReporter != nil { + p.reporter = p.startReporter() + } go func() { select { @@ -91,7 +114,14 @@ func (p *program) Start() error { return nil } -func (p *program) run(cfg *config.Config) error { +func (p *program) run(cfg *config.Config) (err error) { + defer func() { + if err != nil { + // Auxiliary listeners may occupy ports required by the rollback config. + // Release all resources opened by this attempt before rebuilding it. + p.stopRuntime() + } + }() for _, svc := range registry.ServiceRegistry().GetAll() { svc := svc go func() { @@ -152,6 +182,10 @@ func (p *program) run(cfg *config.Config) error { if p.srvProfiling != nil { p.srvProfiling.Close() + if p.profilingListener != nil { + p.profilingListener.Close() + p.profilingListener = nil + } p.srvProfiling = nil } if cfg.Profiling != nil { @@ -162,7 +196,12 @@ func (p *program) run(cfg *config.Config) error { s := &http.Server{ Addr: addr, } + ln, err := net.Listen("tcp", addr) + if err != nil { + return err + } p.srvProfiling = s + p.profilingListener = ln go func() { defer s.Close() @@ -170,7 +209,7 @@ func (p *program) run(cfg *config.Config) error { log := logger.Default().WithFields(map[string]any{"kind": "service", "service": "@profiling"}) log.Info("listening on ", addr) - if err := s.ListenAndServe(); !errors.Is(err, http.ErrServerClosed) { + if err := s.Serve(ln); !errors.Is(err, http.ErrServerClosed) { log.Error(err) } }() @@ -184,30 +223,45 @@ func (p *program) Stop() error { p.cancel() } - for name, srv := range registry.ServiceRegistry().GetAll() { - srv.Close() + if p.reporter != nil { + p.reporter.Stop() + } + unlock := config.LockMutation() + defer unlock() + p.stopped = true + p.stopRuntime() + return nil +} + +func (p *program) stopRuntime() { + for name := range registry.ServiceRegistry().GetAll() { + registry.ServiceRegistry().Unregister(name) logger.Default().Debugf("service %s shutdown", name) } if p.srvApi != nil { p.srvApi.Close() + p.srvApi = nil logger.Default().Debug("service @api shutdown") } if p.srvMetrics != nil { p.srvMetrics.Close() + p.srvMetrics = nil logger.Default().Debug("service @metrics shutdown") } if p.srvProfiling != nil { p.srvProfiling.Close() + if p.profilingListener != nil { + p.profilingListener.Close() + p.profilingListener = nil + } + p.srvProfiling = nil logger.Default().Debug("service @profiling shutdown") } - - return nil } -func (p *program) reload(ctx context.Context) { - c := make(chan os.Signal, 1) - signal.Notify(c, syscall.SIGHUP) +func (p *program) reload(ctx context.Context, c chan os.Signal) { + defer signal.Stop(c) for { select { @@ -225,13 +279,16 @@ func (p *program) reload(ctx context.Context) { } func (p *program) reloadConfig() error { + unlock := config.LockMutation() + defer unlock() + if p.stopped { + return errors.New("agent is shutting down") + } cfg, err := parser.Parse() if err != nil { return err } - config.Set(cfg) - - if err := loader.Load(cfg); err != nil { + if err := loader.Reload(cfg, p.run); err != nil { return err } activeServices := make(map[string]struct{}, len(cfg.Services)) @@ -242,10 +299,6 @@ func (p *program) reloadConfig() error { } xservice.GetGlobalTrafficManager().RetainServices(activeServices) - if err := p.run(cfg); err != nil { - return err - } - return nil } diff --git a/go-gost/tests/lifecycle/lifecycle_test.go b/go-gost/tests/lifecycle/lifecycle_test.go new file mode 100644 index 0000000..07bb10e --- /dev/null +++ b/go-gost/tests/lifecycle/lifecycle_test.go @@ -0,0 +1,313 @@ +//go:build linux || darwin + +package lifecycle_test + +import ( + "context" + "crypto/aes" + "crypto/cipher" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "fmt" + "net" + "net/http" + "net/http/httptest" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "syscall" + "testing" + "time" + + "github.com/gorilla/websocket" +) + +// Exercise the real binary: initial parsing is held at a FIFO while the panel +// tries to send a rule as soon as the WebSocket connects. This reproduced the +// old startup overwrite reliably without timing a large config load. +func TestAgentStartupAndFailedReloadPreservePanelRules(t *testing.T) { + if testing.Short() { + t.Skip("builds and runs the real agent") + } + binary := filepath.Join(t.TempDir(), "gost") + build := exec.Command("go", "build", "-o", binary, "../..") + if output, err := build.CombinedOutput(); err != nil { + t.Fatalf("build: %v\n%s", err, output) + } + dir := t.TempDir() + early, baseline := address(t), address(t) + apiAddr := address(t) + panelConnections := make(chan *websocket.Conn, 2) + wsReady := make(chan struct{}) + responses := make(chan map[string]any, 8) + var ready sync.Once + upgrader := websocket.Upgrader{} + panel := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/system-info" { + w.Write([]byte("ok")) + return + } + c, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer c.Close() + ready.Do(func() { close(wsReady) }) + panelConnections <- c + if err := c.WriteJSON(map[string]any{"type": "AddService", "requestId": "early-rule", "data": []any{service("70_1_0", early)}}); err != nil { + return + } + for { + _, payload, err := c.ReadMessage() + if err != nil { + return + } + env := decodeResponse(t, payload) + if env["requestId"] == "early-rule" { + responses <- env + } + } + })) + defer panel.Close() + writeJSON(t, filepath.Join(dir, "config.json"), map[string]any{"addr": panel.URL, "secret": "audit-secret", "http": 1, "tls": 1, "socks": 1}) + fifo := filepath.Join(dir, "delayed.json") + if err := syscall.Mkfifo(fifo, 0600); err != nil { + t.Fatal(err) + } + logPath := filepath.Join(dir, "agent.log") + logfile, err := os.Create(logPath) + if err != nil { + t.Fatal(err) + } + defer logfile.Close() + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + agent := exec.CommandContext(ctx, binary, "-C", fifo) + agent.Dir, agent.Stdout, agent.Stderr = dir, logfile, logfile + if err := agent.Start(); err != nil { + t.Fatal(err) + } + exited := make(chan error, 1) + go func() { exited <- agent.Wait() }() + defer func() { + agent.Process.Kill() + if t.Failed() { + b, _ := os.ReadFile(logPath) + t.Logf("agent log:\n%s", b) + } + }() + select { + case <-wsReady: + t.Fatal("panel connected before initial config loaded") + case err := <-exited: + t.Fatalf("agent exited during startup: %v", err) + case <-time.After(500 * time.Millisecond): + } + boot := map[string]any{"services": []any{service("71_1_0", baseline)}, "api": map[string]any{"addr": apiAddr}} + writeJSON(t, filepath.Join(dir, "gost.json"), boot) + writeFIFO(t, fifo, boot) + select { + case <-wsReady: + case <-time.After(10 * time.Second): + t.Fatal("panel did not connect after startup") + } + select { + case response := <-responses: + if response["success"] != true { + t.Fatalf("AddService failed: %v", response) + } + case <-time.After(5 * time.Second): + t.Fatal("missing AddService response") + } + await(t, "both startup and panel listeners", func() bool { return listening(early) && listening(baseline) }) + saved, err := os.ReadFile(fifo) + if err != nil || !strings.Contains(string(saved), "70_1_0") { + t.Fatalf("acknowledged rule not persisted: %v", err) + } + + // A valid listener plus an invalid handler also exercises rollback of a + // partially initialized service which already bound the original port. + invalid := service("71_1_0", baseline) + invalid["handler"] = map[string]any{"type": "handler-does-not-exist"} + // The first persisted mutation atomically replaced the FIFO with a regular + // config file, so subsequent reloads use the same real persistence path. + writeJSON(t, fifo, map[string]any{"services": []any{invalid}}) + if err := agent.Process.Signal(syscall.SIGHUP); err != nil { + t.Fatal(err) + } + await(t, "failed reload rollback", func() bool { + b, _ := os.ReadFile(logPath) + return strings.Contains(string(b), "previous config restored") + }) + if !listening(early) || !listening(baseline) { + t.Fatal("failed reload lost a previously acknowledged listener") + } + + // Starting the candidate API on the old service port must not prevent + // rollback when a later auxiliary listener fails to initialize. + logBefore, _ := os.ReadFile(logPath) + writeJSON(t, fifo, map[string]any{ + "services": []any{service("candidate", address(t))}, + "api": map[string]any{"addr": baseline}, + "metrics": map[string]any{"addr": "127.0.0.1:not-a-port"}, + }) + if err := agent.Process.Signal(syscall.SIGHUP); err != nil { + t.Fatal(err) + } + await(t, "auxiliary listener rollback", func() bool { + b, _ := os.ReadFile(logPath) + return strings.Count(string(b), "previous config restored") > strings.Count(string(logBefore), "previous config restored") + }) + if !listening(early) || !listening(baseline) || !listening(apiAddr) { + t.Fatal("candidate auxiliary listener prevented rollback") + } + + // An authenticated API request that never finishes its body must not hold + // the runtime transaction lock against WS mutations or process shutdown. + slow, err := net.Dial("tcp", apiAddr) + if err != nil { + t.Fatal(err) + } + defer slow.Close() + if _, err = fmt.Fprintf(slow, "POST /config/services HTTP/1.1\r\nHost: localhost\r\nAuthorization: Basic dGVzdDp0ZXN0\r\nContent-Type: application/json\r\nContent-Length: 100000\r\n\r\n{"); err != nil { + t.Fatal(err) + } + time.Sleep(50 * time.Millisecond) + panelConn := <-panelConnections + if err := panelConn.WriteJSON(map[string]any{"type": "DeleteService", "requestId": "slow-body-check", "data": map[string]any{"services": []string{"70_1_0"}}}); err != nil { + t.Fatal(err) + } + await(t, "WS mutation despite slow API upload", func() bool { return !listening(early) }) + if err := agent.Process.Signal(syscall.SIGTERM); err != nil { + t.Fatal(err) + } + select { + case err := <-exited: + if err != nil { + t.Fatalf("shutdown: %v", err) + } + case <-time.After(5 * time.Second): + t.Fatal("agent did not stop") + } + if listening(early) || listening(baseline) { + t.Fatal("shutdown left listeners open") + } + + // Startup errors must exit without advertising a node ready to accept rules. + bad := filepath.Join(dir, "invalid.json") + writeJSON(t, bad, map[string]any{"services": []any{invalid}}) + failed := exec.CommandContext(ctx, binary, "-C", bad) + failed.Dir = dir + if output, err := failed.CombinedOutput(); err == nil { + t.Fatalf("invalid startup succeeded: %s", output) + } + select { + case response := <-responses: + t.Fatalf("failed startup accepted panel command: %v", response) + default: + } +} + +func service(name, addr string) map[string]any { + return map[string]any{"name": name, "addr": addr, "listener": map[string]any{"type": "tcp"}, "handler": map[string]any{"type": "auto"}} +} +func address(t *testing.T) string { + t.Helper() + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer l.Close() + return l.Addr().String() +} +func listening(addr string) bool { + c, err := net.DialTimeout("tcp", addr, 50*time.Millisecond) + if err != nil { + return false + } + c.Close() + return true +} +func await(t *testing.T, label string, pred func() bool) { + t.Helper() + for until := time.Now().Add(5 * time.Second); time.Now().Before(until); time.Sleep(10 * time.Millisecond) { + if pred() { + return + } + } + t.Fatal("timeout: " + label) +} +func writeJSON(t *testing.T, path string, value any) { + t.Helper() + b, err := json.Marshal(value) + if err != nil { + t.Fatal(err) + } + if err := os.WriteFile(path, b, 0600); err != nil { + t.Fatal(err) + } +} +func writeFIFO(t *testing.T, path string, value any) { + t.Helper() + b, err := json.Marshal(value) + if err != nil { + t.Fatal(err) + } + done := make(chan error, 1) + go func() { + f, err := os.OpenFile(path, os.O_WRONLY, 0) + if err != nil { + done <- err + return + } + _, err = f.Write(b) + f.Close() + done <- err + }() + select { + case err := <-done: + if err != nil { + t.Fatal(err) + } + case <-time.After(5 * time.Second): + t.Fatal("agent did not read config FIFO") + } +} +func decodeResponse(t *testing.T, payload []byte) map[string]any { + t.Helper() + env := map[string]any{} + if err := json.Unmarshal(payload, &env); err != nil { + t.Error(err) + return nil + } + if encrypted, _ := env["encrypted"].(bool); !encrypted { + return env + } + encoded, _ := env["data"].(string) + raw, err := base64.StdEncoding.DecodeString(encoded) + if err != nil { + t.Error(err) + return nil + } + hash := sha256.Sum256([]byte("audit-secret")) + block, _ := aes.NewCipher(hash[:]) + gcm, _ := cipher.NewGCM(block) + if len(raw) < gcm.NonceSize() { + t.Error("short encrypted response") + return nil + } + payload, err = gcm.Open(nil, raw[:gcm.NonceSize()], raw[gcm.NonceSize():], nil) + if err != nil { + t.Error(err) + return nil + } + env = map[string]any{} + if err := json.Unmarshal(payload, &env); err != nil { + t.Error(err) + return nil + } + return env +} diff --git a/go-gost/x/api/api.go b/go-gost/x/api/api.go index 493f3c6..df1ae57 100644 --- a/go-gost/x/api/api.go +++ b/go-gost/x/api/api.go @@ -52,7 +52,7 @@ func Register(r *gin.Engine, opts *Options) { router.StaticFS("/docs", http.FS(swaggerDoc)) config := router.Group("/config") - config.Use(mwBasicAuth(opts.Auther)) + config.Use(mwBasicAuth(opts.Auther), configTransaction()) config.GET("", getConfig) config.POST("", saveConfig) diff --git a/go-gost/x/api/config_reload.go b/go-gost/x/api/config_reload.go index dbfb003..a971813 100644 --- a/go-gost/x/api/config_reload.go +++ b/go-gost/x/api/config_reload.go @@ -37,9 +37,12 @@ func reloadConfig(ctx *gin.Context) { return } - config.Set(cfg) - - if err := loader.Load(cfg); err != nil { + if err := loader.Reload(cfg, func(*config.Config) error { + for _, svc := range registry.ServiceRegistry().GetAll() { + go svc.Serve() + } + return nil + }); err != nil { writeError(ctx, NewError(http.StatusBadRequest, ErrCodeInvalid, err.Error())) return } @@ -51,13 +54,6 @@ func reloadConfig(ctx *gin.Context) { } xservice.GetGlobalTrafficManager().RetainServices(activeServices) - for _, svc := range registry.ServiceRegistry().GetAll() { - svc := svc - go func() { - svc.Serve() - }() - } - ctx.JSON(http.StatusOK, Response{ Msg: "OK", }) diff --git a/go-gost/x/api/config_transaction.go b/go-gost/x/api/config_transaction.go new file mode 100644 index 0000000..59575fc --- /dev/null +++ b/go-gost/x/api/config_transaction.go @@ -0,0 +1,90 @@ +package api + +import ( + "bytes" + "errors" + "io" + "net/http" + + "github.com/gin-gonic/gin" + "github.com/go-gost/x/config" +) + +const maxConfigRequestBody = 16 << 20 + +// Read the complete bounded request before acquiring the runtime transaction +// lock. Buffer the response until after it is released: neither a slow upload +// nor a client that stops reading may block panel commands, reload or shutdown. +func configTransaction() gin.HandlerFunc { + return func(ctx *gin.Context) { + if ctx.Request.Body != nil { + body := http.MaxBytesReader(ctx.Writer, ctx.Request.Body, maxConfigRequestBody) + data, err := io.ReadAll(body) + body.Close() + if err != nil { + status := http.StatusBadRequest + var tooLarge *http.MaxBytesError + if errors.As(err, &tooLarge) { + status = http.StatusRequestEntityTooLarge + } + ctx.AbortWithStatusJSON(status, Response{Code: status, Msg: "Unable to read configuration request"}) + return + } + ctx.Request.Body = io.NopCloser(bytes.NewReader(data)) + } + + writer := ctx.Writer + buffered := &configResponseWriter{ResponseWriter: writer, header: writer.Header().Clone(), status: http.StatusOK, size: -1} + ctx.Writer = buffered + defer func() { ctx.Writer = writer }() + func() { + unlock := config.LockMutation() + defer unlock() + // A request waiting behind reload may have been closed during shutdown. + if ctx.Request.Context().Err() != nil { + ctx.Abort() + return + } + ctx.Next() + }() + ctx.Writer = writer + for key, values := range buffered.header { + writer.Header()[key] = values + } + writer.WriteHeader(buffered.status) + writer.Write(buffered.body.Bytes()) + } +} + +// Config endpoints return JSON rather than streaming. Preserve Gin's response +// bookkeeping while delaying all network writes until the transaction ends. +type configResponseWriter struct { + gin.ResponseWriter + header http.Header + body bytes.Buffer + status int + size int +} + +func (w *configResponseWriter) Header() http.Header { return w.header } +func (w *configResponseWriter) WriteHeader(status int) { + if !w.Written() && status > 0 { + w.status = status + } +} +func (w *configResponseWriter) WriteHeaderNow() { + if !w.Written() { + w.size = 0 + } +} +func (w *configResponseWriter) Write(p []byte) (int, error) { + w.WriteHeaderNow() + n, err := w.body.Write(p) + w.size += n + return n, err +} +func (w *configResponseWriter) WriteString(s string) (int, error) { return w.Write([]byte(s)) } +func (w *configResponseWriter) Status() int { return w.status } +func (w *configResponseWriter) Size() int { return w.size } +func (w *configResponseWriter) Written() bool { return w.size >= 0 } +func (w *configResponseWriter) Flush() { w.WriteHeaderNow() } diff --git a/go-gost/x/api/config_transaction_test.go b/go-gost/x/api/config_transaction_test.go new file mode 100644 index 0000000..9e2b311 --- /dev/null +++ b/go-gost/x/api/config_transaction_test.go @@ -0,0 +1,76 @@ +package api + +import ( + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/go-gost/x/config" +) + +func assertMutationAvailable(t *testing.T) { + t.Helper() + done := make(chan struct{}) + go func() { unlock := config.LockMutation(); unlock(); close(done) }() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("network I/O holds runtime mutation lock") + } +} + +func TestConfigTransactionDoesNotLockWhileReadingBody(t *testing.T) { + router := gin.New() + router.Use(configTransaction()) + router.POST("/config", func(c *gin.Context) { c.JSON(http.StatusOK, Response{Msg: "OK"}) }) + reader, writer := io.Pipe() + defer reader.Close() + defer writer.Close() + request := httptest.NewRequest(http.MethodPost, "/config", reader) + done := make(chan struct{}) + go func() { defer close(done); router.ServeHTTP(httptest.NewRecorder(), request) }() + if _, err := writer.Write([]byte("{")); err != nil { + t.Fatal(err) + } + assertMutationAvailable(t) + writer.Close() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("handler did not finish") + } +} + +type blockedResponse struct { + header http.Header + started chan struct{} + release chan struct{} +} + +func (w *blockedResponse) Header() http.Header { return w.header } +func (w *blockedResponse) WriteHeader(int) {} +func (w *blockedResponse) Write(p []byte) (int, error) { + close(w.started) + <-w.release + return len(p), nil +} + +func TestConfigTransactionReleasesLockBeforeSendingResponse(t *testing.T) { + router := gin.New() + router.Use(configTransaction()) + router.POST("/config", func(c *gin.Context) { c.JSON(http.StatusOK, Response{Msg: "OK"}) }) + writer := &blockedResponse{header: make(http.Header), started: make(chan struct{}), release: make(chan struct{})} + defer close(writer.release) + request := httptest.NewRequest(http.MethodPost, "/config", strings.NewReader("{}")) + go router.ServeHTTP(writer, request) + select { + case <-writer.started: + case <-time.After(time.Second): + t.Fatal("response did not start") + } + assertMutationAvailable(t) +} diff --git a/go-gost/x/api/service/service.go b/go-gost/x/api/service/service.go index 914e6ee..2089d2c 100644 --- a/go-gost/x/api/service/service.go +++ b/go-gost/x/api/service/service.go @@ -3,6 +3,7 @@ package service import ( "net" "net/http" + "time" "github.com/gin-gonic/gin" "github.com/go-gost/core/auth" @@ -67,7 +68,11 @@ func NewService(network, addr string, opts ...Option) (service.Service, error) { return &server{ s: &http.Server{ - Handler: r, + Handler: r, + ReadHeaderTimeout: 5 * time.Second, + ReadTimeout: 15 * time.Second, + WriteTimeout: 30 * time.Second, + IdleTimeout: 60 * time.Second, }, ln: ln, cclose: make(chan struct{}), @@ -83,7 +88,11 @@ func (s *server) Addr() net.Addr { } func (s *server) Close() error { - return s.s.Close() + // Close can race the goroutine entering Serve during a failed startup. + // http.Server.Close alone does not own the listener until Serve starts. + err := s.s.Close() + s.ln.Close() + return err } func (s *server) IsClosed() bool { diff --git a/go-gost/x/config/config.go b/go-gost/x/config/config.go index 5ca370a..6e8ab57 100644 --- a/go-gost/x/config/config.go +++ b/go-gost/x/config/config.go @@ -23,10 +23,20 @@ func init() { } var ( - global = &Config{} - globalMux sync.RWMutex + global = &Config{} + globalMux sync.RWMutex + mutationMux sync.Mutex ) +// LockMutation serializes complete runtime/config transactions across panel +// commands, management API requests, startup and reload. It is separate from +// globalMux so callers may safely use Global, Set and OnUpdate while holding it. +// Lock before reading/parsing the config and hold through persistence/rollback. +func LockMutation() func() { + mutationMux.Lock() + return mutationMux.Unlock +} + func Global() *Config { globalMux.RLock() defer globalMux.RUnlock() diff --git a/go-gost/x/config/loader/loader.go b/go-gost/x/config/loader/loader.go index 7273271..10b0cc7 100644 --- a/go-gost/x/config/loader/loader.go +++ b/go-gost/x/config/loader/loader.go @@ -1,6 +1,8 @@ package loader import ( + "fmt" + "github.com/go-gost/core/logger" "github.com/go-gost/x/config" "github.com/go-gost/x/config/parsing" @@ -30,6 +32,34 @@ func Load(cfg *config.Config) error { return defaultLoader.Load(cfg) } +// Reload replaces the runtime and commits the config only after it starts. +// The caller must hold config.LockMutation, including while parsing cfg, so the +// rollback snapshot includes every previously acknowledged runtime command. +// Failed loads can partially replace registries and bind listeners; always +// rebuild the last successful snapshot before returning an error. +func Reload(cfg *config.Config, run func(*config.Config) error) error { + previous := config.Global() + if err := apply(cfg, run); err != nil { + // A failed partial load may have left both old and new listeners behind. + for name := range registry.ServiceRegistry().GetAll() { + registry.ServiceRegistry().Unregister(name) + } + if rollbackErr := apply(previous, run); rollbackErr != nil { + return fmt.Errorf("reload failed: %w; restore previous config failed: %v", err, rollbackErr) + } + return fmt.Errorf("reload failed (previous config restored): %w", err) + } + config.Set(cfg) + return nil +} + +func apply(cfg *config.Config, run func(*config.Config) error) error { + if err := Load(cfg); err != nil { + return err + } + return run(cfg) +} + type loader struct{} func (l *loader) Load(cfg *config.Config) error { @@ -217,12 +247,19 @@ func register(cfg *config.Config) error { registry.ServiceRegistry().Unregister(name) } for _, svcCfg := range cfg.Services { + if svcCfg == nil { + return fmt.Errorf("service config is nil") + } + if paused, _ := svcCfg.Metadata["paused"].(bool); paused { + continue + } svc, err := service_parser.ParseService(svcCfg) if err != nil { return err } if svc != nil { if err := registry.ServiceRegistry().Register(svcCfg.Name, svc); err != nil { + svc.Close() return err } } diff --git a/go-gost/x/config/loader/reload_test.go b/go-gost/x/config/loader/reload_test.go new file mode 100644 index 0000000..a31977c --- /dev/null +++ b/go-gost/x/config/loader/reload_test.go @@ -0,0 +1,84 @@ +package loader_test + +import ( + "net" + "testing" + "time" + + "github.com/go-gost/x/config" + "github.com/go-gost/x/config/loader" + _ "github.com/go-gost/x/handler/auto" + _ "github.com/go-gost/x/listener/tcp" + "github.com/go-gost/x/registry" +) + +func TestReloadRestoresPreviousRuntime(t *testing.T) { + for _, failure := range []string{"listener", "handler", "partial", "run"} { + t.Run(failure, func(t *testing.T) { + unlock := config.LockMutation() + defer unlock() + original := config.Global() + defer config.Set(original) + defer func() { + for name := range registry.ServiceRegistry().GetAll() { + registry.ServiceRegistry().Unregister(name) + } + }() + old := &config.Config{Services: []*config.ServiceConfig{ + {Name: "shared-rule", Addr: "127.0.0.1:0", Handler: &config.HandlerConfig{Type: "auto"}, Listener: &config.ListenerConfig{Type: "tcp"}}, + {Name: "paused-rule", Addr: "127.0.0.1:0", Metadata: map[string]any{"paused": true}}, + }} + if err := loader.Load(old); err != nil { + t.Fatal(err) + } + addr := registry.ServiceRegistry().Get("shared-rule").Addr().String() + old.Services[0].Addr = addr + config.Set(old) + serve := func(*config.Config) error { + for _, svc := range registry.ServiceRegistry().GetAll() { + go svc.Serve() + } + return nil + } + serve(old) + replacement := config.Global() + replacement.Services = replacement.Services[:1] + switch failure { + case "listener": + replacement.Services[0].Listener.Type = "invalid-listener" + case "handler": + replacement.Services[0].Handler.Type = "invalid-handler" + case "partial": + replacement.Services = append(replacement.Services, &config.ServiceConfig{Name: "broken", Listener: &config.ListenerConfig{Type: "invalid-listener"}}) + } + run := serve + if failure == "run" { + run = func(cfg *config.Config) error { + if cfg == replacement { + return &net.AddrError{Err: "auxiliary listener failed", Addr: "test"} + } + return serve(cfg) + } + } + if err := loader.Reload(replacement, run); err == nil { + t.Fatal("expected reload failure") + } + restored := registry.ServiceRegistry().Get("shared-rule") + if restored == nil || restored.Addr().String() != addr { + t.Fatal("previous service was not restored on its original port") + } + if registry.ServiceRegistry().Get("paused-rule") != nil { + t.Fatal("rollback resumed a paused service") + } + conn, err := net.DialTimeout("tcp", addr, time.Second) + if err != nil { + t.Fatalf("restored listener unavailable: %v", err) + } + conn.Close() + got := config.Global() + if len(got.Services) != 2 || got.Services[0].Listener.Type != "tcp" || got.Services[0].Handler.Type != "auto" { + t.Fatalf("failed config was committed: %+v", got.Services) + } + }) + } +} diff --git a/go-gost/x/config/parsing/service/parse.go b/go-gost/x/config/parsing/service/parse.go index 584398a..ec2df4c 100644 --- a/go-gost/x/config/parsing/service/parse.go +++ b/go-gost/x/config/parsing/service/parse.go @@ -255,6 +255,15 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) { return nil, err } + // Listener initialization binds the port. If handler/TLS/forwarder parsing + // fails, release it so a reload rollback can restore the previous listener. + configured := false + defer func() { + if !configured { + ln.Close() + } + }() + handlerLogger := serviceLogger.WithFields(map[string]any{ "kind": "handler", }) @@ -379,6 +388,7 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) { ) serviceLogger.Infof("listening on %s/%s", s.Addr().String(), s.Addr().Network()) + configured = true return s, nil } diff --git a/go-gost/x/metrics/service/service.go b/go-gost/x/metrics/service/service.go index 3e585c9..0871cad 100644 --- a/go-gost/x/metrics/service/service.go +++ b/go-gost/x/metrics/service/service.go @@ -84,7 +84,9 @@ func (s *metricService) Addr() net.Addr { } func (s *metricService) Close() error { - return s.s.Close() + err := s.s.Close() + s.ln.Close() + return err } func (s *metricService) IsClosed() bool { diff --git a/go-gost/x/service/traffic_reporter.go b/go-gost/x/service/traffic_reporter.go index 520f560..fee6065 100644 --- a/go-gost/x/service/traffic_reporter.go +++ b/go-gost/x/service/traffic_reporter.go @@ -370,43 +370,43 @@ type getConfigResponse struct { // getConfigData 获取配置数据(避免循环依赖) func getConfigData() ([]byte, error) { - config.OnUpdate(func(c *config.Config) error { - for _, svc := range c.Services { - if svc == nil { - continue + // Reporting is read-only: enriching a detached snapshot must not persist + // stale state over a config currently being reloaded or acknowledged. + cfg := config.Global() + for _, svc := range cfg.Services { + if svc == nil { + continue + } + s := registry.ServiceRegistry().Get(svc.Name) + ss, ok := s.(serviceStatus) + if ok && ss != nil { + status := ss.Status() + svc.Status = &config.ServiceStatus{ + CreateTime: status.CreateTime().Unix(), + State: string(status.State()), } - s := registry.ServiceRegistry().Get(svc.Name) - ss, ok := s.(serviceStatus) - if ok && ss != nil { - status := ss.Status() - svc.Status = &config.ServiceStatus{ - CreateTime: status.CreateTime().Unix(), - State: string(status.State()), + if st := status.Stats(); st != nil { + svc.Status.Stats = &config.ServiceStats{ + TotalConns: st.Get(stats.KindTotalConns), + CurrentConns: st.Get(stats.KindCurrentConns), + TotalErrs: st.Get(stats.KindTotalErrs), + InputBytes: st.Get(stats.KindInputBytes), + OutputBytes: st.Get(stats.KindOutputBytes), } - if st := status.Stats(); st != nil { - svc.Status.Stats = &config.ServiceStats{ - TotalConns: st.Get(stats.KindTotalConns), - CurrentConns: st.Get(stats.KindCurrentConns), - TotalErrs: st.Get(stats.KindTotalErrs), - InputBytes: st.Get(stats.KindInputBytes), - OutputBytes: st.Get(stats.KindOutputBytes), - } - } - for _, ev := range status.Events() { - if !ev.Time.IsZero() { - svc.Status.Events = append(svc.Status.Events, config.ServiceEvent{ - Time: ev.Time.Unix(), - Msg: ev.Message, - }) - } + } + for _, ev := range status.Events() { + if !ev.Time.IsZero() { + svc.Status.Events = append(svc.Status.Events, config.ServiceEvent{ + Time: ev.Time.Unix(), + Msg: ev.Message, + }) } } } - return nil - }) + } var resp getConfigResponse - resp.Config = config.Global() + resp.Config = cfg buf := &bytes.Buffer{} resp.Config.Write(buf, "json") diff --git a/go-gost/x/service/traffic_reporter_test.go b/go-gost/x/service/traffic_reporter_test.go index 0fdc651..ac727cb 100644 --- a/go-gost/x/service/traffic_reporter_test.go +++ b/go-gost/x/service/traffic_reporter_test.go @@ -1,10 +1,14 @@ package service import ( + "bytes" "context" "errors" + "github.com/go-gost/x/config" "io" "net/http" + "os" + "path/filepath" "strings" "testing" "time" @@ -148,3 +152,31 @@ func TestPostJSONWithFallbackRemembersDetectedURL(t *testing.T) { t.Fatalf("expected remembered http url first, got %s", calls[0]) } } + +func TestConfigReportDoesNotPersistStaleRuntime(t *testing.T) { + previous, path := config.Global(), config.PersistPath() + defer config.Set(previous) + defer config.SetPersistPath(path) + config.Set(&config.Config{Services: []*config.ServiceConfig{{Name: "old-runtime"}}}) + filename := filepath.Join(t.TempDir(), "gost.json") + candidate := []byte(`{"services":[{"name":"new-on-disk"}]}`) + if err := os.WriteFile(filename, candidate, 0600); err != nil { + t.Fatal(err) + } + config.SetPersistPath(filename) + config.EnablePersist() + report, err := getConfigData() + if err != nil { + t.Fatal(err) + } + if !bytes.Contains(report, []byte("old-runtime")) { + t.Fatalf("unexpected report: %s", report) + } + saved, err := os.ReadFile(filename) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(saved, candidate) { + t.Fatalf("report overwrote candidate config: %s", saved) + } +} diff --git a/go-gost/x/socket/command_dispatch_test.go b/go-gost/x/socket/command_dispatch_test.go index 79a8e91..f063d29 100644 --- a/go-gost/x/socket/command_dispatch_test.go +++ b/go-gost/x/socket/command_dispatch_test.go @@ -115,3 +115,39 @@ func TestMutationQueueExecutesCommandsInArrivalOrder(t *testing.T) { } close(second.release) } + +func TestMutationCommandWaitsForRuntimeTransaction(t *testing.T) { + original := config.Global() + defer config.Set(original) + name := "reload_transaction_service" + svc := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})} + close(svc.release) + if err := registry.ServiceRegistry().Register(name, svc); err != nil { + t.Fatal(err) + } + defer registry.ServiceRegistry().Unregister(name) + config.Set(&config.Config{Services: []*config.ServiceConfig{{Name: name}}}) + reporter := NewWebSocketReporter("", "transaction-test-secret") + defer reporter.Stop() + unlock := config.LockMutation() + done := make(chan struct{}) + go func() { + defer close(done) + reporter.routeCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{name}}}) + }() + select { + case <-svc.started: + unlock() + t.Fatal("command interleaved with runtime transaction") + case <-time.After(50 * time.Millisecond): + } + unlock() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("command did not resume after transaction") + } + if registry.ServiceRegistry().Get(name) != nil { + t.Fatal("service was not removed") + } +} diff --git a/go-gost/x/socket/websocket_reporter.go b/go-gost/x/socket/websocket_reporter.go index e84ddca..48d5338 100644 --- a/go-gost/x/socket/websocket_reporter.go +++ b/go-gost/x/socket/websocket_reporter.go @@ -179,6 +179,7 @@ type WebSocketReporter struct { tcpPingSem chan struct{} // 限制诊断探测并发,避免离线目标耗尽连接 readCommandSem chan struct{} // 限制只读命令并发,避免诊断请求耗尽资源 mutationQueue chan CommandMessage + workers sync.WaitGroup } var wsDial = func(dialer *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) { @@ -238,8 +239,15 @@ func (w *WebSocketReporter) releaseTCPPingSlot() { // Start 启动WebSocket报告器 func (w *WebSocketReporter) Start() { - go w.runMutationCommands() - go w.run() + w.workers.Add(2) + go func() { + defer w.workers.Done() + w.runMutationCommands() + }() + go func() { + defer w.workers.Done() + w.run() + }() } // Stop 停止WebSocket报告器 @@ -250,6 +258,7 @@ func (w *WebSocketReporter) Stop() { w.conn.Close() } w.connMutex.Unlock() + w.workers.Wait() } // backoffWithJitter 返回带随机抖动的退避时间(±25%) @@ -339,10 +348,10 @@ func (w *WebSocketReporter) connect() error { candidates := buildWebSocketCandidates(w.addr, w.secret, w.version, cfg.Http, cfg.Tls, cfg.Socks, w.preferredWSScheme) - dialer := websocket.DefaultDialer + dialer := *websocket.DefaultDialer dialer.HandshakeTimeout = 10 * time.Second - conn, usedURL, err := dialWebSocketWithFallback(dialer, candidates) + conn, usedURL, err := dialWebSocketWithFallback(&dialer, candidates) if err != nil { return err } @@ -528,7 +537,11 @@ func (w *WebSocketReporter) handleConnection() { }() // 启动消息接收goroutine - go w.receiveMessages() + w.workers.Add(1) + go func() { + defer w.workers.Done() + w.receiveMessages() + }() // 指标上报 ticker metricTicker := time.NewTicker(w.pingInterval) @@ -887,6 +900,14 @@ func isMutationCommand(commandType string) bool { // routeCommand 路由命令到对应的处理函数 func (w *WebSocketReporter) routeCommand(cmd CommandMessage) { + if isMutationCommand(cmd.Type) { + unlock := config.LockMutation() + defer unlock() + if w.ctx.Err() != nil { + w.sendCommandFailure(cmd, "Agent is shutting down") + return + } + } jsonBytes, errs := json.Marshal(cmd) if errs != nil { fmt.Println("Error marshaling JSON:", errs)