fix: protect shared rules across delivery and recovery (#560)

This commit is contained in:
sagit
2026-09-29 14:05:49 +08:00
committed by GitHub
parent 129fa0aa4c
commit c952d2fb3a
49 changed files with 5053 additions and 643 deletions
+174 -153
View File
@@ -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) {
@@ -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)
}
}
@@ -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)
}
@@ -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)
}
}
@@ -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)
}
@@ -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)
}
})
}
}
}
@@ -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)
}
}
@@ -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)
}
}
@@ -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)
}
}
@@ -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)
}
})
}
}
+200 -197
View File
@@ -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)
}
@@ -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)
@@ -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)
}
}
@@ -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,
+11 -9
View File
@@ -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)
}
+44 -1
View File
@@ -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) {
+302 -119
View File
@@ -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) {
@@ -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")
}
}
@@ -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
}
@@ -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")
}
}
+93 -7
View File
@@ -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
}
}
@@ -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)