mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-02 08:56:38 +08:00
fix: protect shared rules across delivery and recovery (#560)
This commit is contained in:
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
package model
|
||||
|
||||
// FederationPendingRelease keeps rollback work after a failed remote release.
|
||||
// It is independent of tunnels, which may never have committed or be deleted.
|
||||
type FederationPendingRelease struct {
|
||||
ID string `gorm:"primaryKey;size:64"`
|
||||
RemoteURL string `gorm:"not null"`
|
||||
RemoteToken string `gorm:"not null"`
|
||||
BindingID string
|
||||
ReservationID string
|
||||
ResourceKey string
|
||||
CreatedTime int64
|
||||
}
|
||||
|
||||
func (FederationPendingRelease) TableName() string { return "federation_pending_release" }
|
||||
@@ -354,23 +354,24 @@ type PeerShare struct {
|
||||
func (PeerShare) TableName() string { return "peer_share" }
|
||||
|
||||
type PeerShareRuntime struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
ShareID int64 `gorm:"column:share_id;not null;index:idx_peer_share_runtime_share_node_status"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;index:idx_peer_share_runtime_share_node_status"`
|
||||
ReservationID string `gorm:"column:reservation_id;type:text;not null;uniqueIndex"`
|
||||
ResourceKey string `gorm:"column:resource_key;type:text;not null;uniqueIndex"`
|
||||
BindingID string `gorm:"column:binding_id;type:text;not null;default:'';index:idx_peer_share_runtime_binding_id"`
|
||||
Role string `gorm:"type:text;not null;default:''"`
|
||||
ChainName string `gorm:"column:chain_name;type:text;not null;default:''"`
|
||||
ServiceName string `gorm:"column:service_name;type:text;not null;default:''"`
|
||||
Protocol string `gorm:"type:text;not null;default:'tls'"`
|
||||
Strategy string `gorm:"type:text;not null;default:'round'"`
|
||||
Port int `gorm:"not null;default:0"`
|
||||
Target string `gorm:"type:text;not null;default:''"`
|
||||
Applied int `gorm:"not null;default:0"`
|
||||
Status int `gorm:"not null;default:1;index:idx_peer_share_runtime_share_node_status"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
ShareID int64 `gorm:"column:share_id;not null;index:idx_peer_share_runtime_share_node_status"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;index:idx_peer_share_runtime_share_node_status"`
|
||||
ReservationID string `gorm:"column:reservation_id;type:text;not null;uniqueIndex"`
|
||||
ResourceKey string `gorm:"column:resource_key;type:text;not null;uniqueIndex"`
|
||||
BindingID string `gorm:"column:binding_id;type:text;not null;default:'';index:idx_peer_share_runtime_binding_id"`
|
||||
Role string `gorm:"type:text;not null;default:''"`
|
||||
ChainName string `gorm:"column:chain_name;type:text;not null;default:''"`
|
||||
ServiceName string `gorm:"column:service_name;type:text;not null;default:''"`
|
||||
Protocol string `gorm:"type:text;not null;default:'tls'"`
|
||||
Strategy string `gorm:"type:text;not null;default:'round'"`
|
||||
Port int `gorm:"not null;default:0"`
|
||||
Target string `gorm:"type:text;not null;default:''"`
|
||||
Applied int `gorm:"not null;default:0"`
|
||||
ReleasePending int `gorm:"column:release_pending;not null;default:0"`
|
||||
Status int `gorm:"not null;default:1;index:idx_peer_share_runtime_share_node_status"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (PeerShareRuntime) TableName() string { return "peer_share_runtime" }
|
||||
@@ -816,3 +817,23 @@ type TunnelQuality struct {
|
||||
}
|
||||
|
||||
func (TunnelQuality) TableName() string { return "tunnel_quality" }
|
||||
|
||||
// PeerShareResource is the durable desired state for a namespaced peer command.
|
||||
// Rows are retained as tombstones until deletion has been acknowledged.
|
||||
type PeerShareResource struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
ShareID int64 `gorm:"column:share_id;not null;uniqueIndex:idx_peer_share_resource_key"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;index"`
|
||||
Kind string `gorm:"type:text;not null;uniqueIndex:idx_peer_share_resource_key"`
|
||||
OriginalName string `gorm:"column:original_name;type:text;not null;uniqueIndex:idx_peer_share_resource_key"`
|
||||
RuntimeName string `gorm:"column:runtime_name;type:text;not null;index"`
|
||||
LegacyNames string `gorm:"column:legacy_names;type:text;not null;default:''"`
|
||||
LegacyServiceBase string `gorm:"column:legacy_service_base;type:text;not null;default:''"`
|
||||
ReleaseLegacyFamily bool `gorm:"column:release_legacy_family;not null;default:false"`
|
||||
Config string `gorm:"type:text;not null;default:''"`
|
||||
DesiredState string `gorm:"column:desired_state;type:text;not null;default:'active'"`
|
||||
Applied int `gorm:"not null;default:0"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (PeerShareResource) TableName() string { return "peer_share_resource" }
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"go-backend/internal/store/model"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
func (r *Repository) SavePeerShareResources(items []PeerShareResource) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
for i := range items {
|
||||
if err := tx.Clauses(clause.OnConflict{Columns: []clause.Column{{Name: "share_id"}, {Name: "kind"}, {Name: "original_name"}}, DoUpdates: clause.AssignmentColumns([]string{"node_id", "runtime_name", "legacy_names", "legacy_service_base", "release_legacy_family", "config", "desired_state", "applied", "updated_time"})}).Create(&items[i]).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func (r *Repository) GetPeerShareResource(shareID int64, kind, originalName string) (*PeerShareResource, error) {
|
||||
var item model.PeerShareResource
|
||||
err := r.db.Where("share_id = ? AND kind = ? AND original_name = ?", shareID, kind, originalName).First(&item).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
return &item, err
|
||||
}
|
||||
|
||||
func (r *Repository) ListPeerShareResourcesByNode(nodeID int64) ([]PeerShareResource, error) {
|
||||
var items []PeerShareResource
|
||||
err := r.db.Where("node_id = ?", nodeID).Order("id").Find(&items).Error
|
||||
return items, err
|
||||
}
|
||||
|
||||
func (r *Repository) MarkPeerShareResourceApplied(shareID int64, kind, name string) error {
|
||||
return r.db.Model(&model.PeerShareResource{}).Where("share_id = ? AND kind = ? AND original_name = ?", shareID, kind, name).Update("applied", 1).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ClearPeerShareResourceLegacyNames(shareID int64, kind, name string) error {
|
||||
return r.db.Model(&model.PeerShareResource{}).Where("share_id = ? AND kind = ? AND original_name = ?", shareID, kind, name).Update("legacy_names", "").Error
|
||||
}
|
||||
|
||||
func (r *Repository) WithPeerShareResourceTransaction(fn func(*Repository) error) error {
|
||||
return r.db.Transaction(func(tx *gorm.DB) error { return fn(&Repository{db: tx, dbPath: r.dbPath}) })
|
||||
}
|
||||
func (r *Repository) ClearPeerShareResourceLegacyFamily(shareID int64, base string) error {
|
||||
return r.db.Model(&model.PeerShareResource{}).Where("share_id = ? AND legacy_service_base = ?", shareID, base).Update("legacy_service_base", "").Error
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"go-backend/internal/store/model"
|
||||
"time"
|
||||
)
|
||||
|
||||
// A pending release remains active until the agent acknowledges deletion. This
|
||||
// keeps its port reserved even if the control connection is unavailable.
|
||||
func (r *Repository) SetPeerShareRuntimeReleasePending(id int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.PeerShareRuntime{}).Where("id = ? AND status = 1", id).Updates(map[string]interface{}{"release_pending": 1, "updated_time": time.Now().UnixMilli()}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) CompletePeerShareRuntimeRelease(id int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.PeerShareRuntime{}).Where("id = ?", id).Updates(map[string]interface{}{"status": 0, "applied": 0, "release_pending": 0, "updated_time": time.Now().UnixMilli()}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ListActivePeerShareRuntimesByNode(nodeID int64) ([]model.PeerShareRuntime, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var items []model.PeerShareRuntime
|
||||
err := r.db.Where("node_id = ? AND status = 1", nodeID).Order("release_pending DESC, id ASC").Find(&items).Error
|
||||
return items, err
|
||||
}
|
||||
@@ -43,6 +43,7 @@ type UserForwardDetail = model.UserForwardDetail
|
||||
type StatisticsFlow = model.StatisticsFlow
|
||||
type Node = model.Node
|
||||
type PeerShare = model.PeerShare
|
||||
type PeerShareResource = model.PeerShareResource
|
||||
type PeerShareRuntime = model.PeerShareRuntime
|
||||
type FederationTunnelBinding = model.FederationTunnelBinding
|
||||
type BackupData = model.BackupData
|
||||
@@ -309,7 +310,9 @@ func autoMigrateAll(db *gorm.DB) error {
|
||||
&model.ViteConfig{},
|
||||
&model.PeerShare{},
|
||||
&model.PeerShareRuntime{},
|
||||
&model.PeerShareResource{},
|
||||
&model.FederationTunnelBinding{},
|
||||
&model.FederationPendingRelease{},
|
||||
&model.Announcement{},
|
||||
&model.SchemaVersion{},
|
||||
&model.NodeMetric{},
|
||||
@@ -1514,7 +1517,12 @@ func (r *Repository) DeletePeerShare(id int64) error {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
tx.Where("share_id = ?", id).Delete(&model.PeerShareRuntime{})
|
||||
if err := tx.Where("share_id = ?", id).Delete(&model.PeerShareRuntime{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("share_id = ?", id).Delete(&model.PeerShareResource{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Where("id = ?", id).Delete(&model.PeerShare{}).Error
|
||||
})
|
||||
}
|
||||
@@ -1607,7 +1615,8 @@ func (r *Repository) UpdatePeerShareRuntime(item *model.PeerShareRuntime) error
|
||||
return errors.New("runtime item is nil")
|
||||
}
|
||||
return r.db.Model(&model.PeerShareRuntime{}).Where("id = ?", item.ID).Updates(map[string]interface{}{
|
||||
"binding_id": item.BindingID, "role": item.Role,
|
||||
"reservation_id": item.ReservationID,
|
||||
"binding_id": item.BindingID, "role": item.Role,
|
||||
"chain_name": item.ChainName, "service_name": item.ServiceName,
|
||||
"protocol": item.Protocol, "strategy": item.Strategy,
|
||||
"port": item.Port, "target": item.Target,
|
||||
@@ -1755,35 +1764,19 @@ func (r *Repository) ListActiveForwardPeerShareRuntimesByNodeAndServiceName(node
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func (r *Repository) ListActiveForwardPeerShareRuntimeServiceNamesByNode(nodeID int64) ([]string, error) {
|
||||
func (r *Repository) ListActiveForwardPeerShareRuntimesByNode(nodeID int64) ([]model.PeerShareRuntime, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var names []string
|
||||
err := r.db.Model(&model.PeerShareRuntime{}).
|
||||
Where("node_id = ? AND status = 1 AND role = ? AND service_name <> ''", nodeID, "forward").
|
||||
Pluck("service_name", &names).Error
|
||||
var items []model.PeerShareRuntime
|
||||
err := r.db.Where("node_id = ? AND status = 1 AND role = ?", nodeID, "forward").Find(&items).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if names == nil {
|
||||
names = make([]string, 0)
|
||||
if items == nil {
|
||||
items = make([]model.PeerShareRuntime, 0)
|
||||
}
|
||||
return names, nil
|
||||
}
|
||||
|
||||
func (r *Repository) HasRecentUnboundForwardPeerShareRuntimeOnNode(nodeID int64, minUpdatedTime int64) (bool, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return false, errors.New("repository not initialized")
|
||||
}
|
||||
var count int64
|
||||
err := r.db.Model(&model.PeerShareRuntime{}).
|
||||
Where("node_id = ? AND status = 1 AND role = ? AND applied = 0 AND updated_time >= ? AND (service_name = '' OR service_name IS NULL)", nodeID, "forward", minUpdatedTime).
|
||||
Count(&count).Error
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func (r *Repository) GetActiveForwardPeerShareRuntimeByPort(shareID int64, port int) (*model.PeerShareRuntime, error) {
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"fmt"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
const FederationBindingPendingRelease = 2
|
||||
|
||||
func (r *Repository) ListTunnelChainNodeIDsTx(tx *gorm.DB, tunnelID int64) ([]int64, error) {
|
||||
var ids []int64
|
||||
err := tx.Model(&model.ChainTunnel{}).Where("tunnel_id = ?", tunnelID).Pluck("node_id", &ids).Error
|
||||
return ids, err
|
||||
}
|
||||
|
||||
func (r *Repository) ListFederationTunnelBindingsForCleanup(tunnelID int64) ([]model.FederationTunnelBinding, error) {
|
||||
var rows []model.FederationTunnelBinding
|
||||
err := r.db.Where("tunnel_id = ? AND status IN ?", tunnelID, []int{1, FederationBindingPendingRelease}).Order("id").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
func (r *Repository) ListPendingFederationTunnelBindings() ([]model.FederationTunnelBinding, error) {
|
||||
var rows []model.FederationTunnelBinding
|
||||
err := r.db.Where("status = ?", FederationBindingPendingRelease).Order("id").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
func (r *Repository) MarkFederationTunnelBindingPendingRelease(id int64) error {
|
||||
return r.db.Model(&model.FederationTunnelBinding{}).Where("id = ?", id).
|
||||
Updates(map[string]interface{}{"status": FederationBindingPendingRelease, "updated_time": unixMilliNow()}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) DeleteFederationTunnelBinding(id int64) error {
|
||||
return r.db.Where("id = ?", id).Delete(&model.FederationTunnelBinding{}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) SavePendingFederationRelease(item *model.FederationPendingRelease) error {
|
||||
item.ID = fmt.Sprintf("%x", sha256.Sum256([]byte(item.RemoteURL+"\n"+item.BindingID+"\n"+item.ReservationID+"\n"+item.ResourceKey)))
|
||||
item.CreatedTime = unixMilliNow()
|
||||
return r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(item).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ListPendingFederationReleases() ([]model.FederationPendingRelease, error) {
|
||||
var rows []model.FederationPendingRelease
|
||||
err := r.db.Order("created_time, id").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
func (r *Repository) DeletePendingFederationRelease(id string) error {
|
||||
return r.db.Where("id = ?", id).Delete(&model.FederationPendingRelease{}).Error
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"sort"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
// Only retry unfinished operations. Successful desired state must be replayed
|
||||
// on a real reconnect, not every maintenance tick while the agent stays online.
|
||||
func (r *Repository) ListPendingPeerShareNodeIDs() ([]int64, error) {
|
||||
var resourceNodes, runtimeNodes []int64
|
||||
if err := r.db.Model(&model.PeerShareResource{}).Where("applied = 0").Distinct("node_id").Pluck("node_id", &resourceNodes).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := r.db.Model(&model.PeerShareRuntime{}).
|
||||
Where("status = 1 AND (release_pending <> 0 OR (applied = 0 AND service_name <> '' AND role IN ?))", []string{"middle", "exit"}).
|
||||
Distinct("node_id").Pluck("node_id", &runtimeNodes).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
seen := make(map[int64]struct{})
|
||||
for _, ids := range [][]int64{resourceNodes, runtimeNodes} {
|
||||
for _, id := range ids {
|
||||
if id > 0 {
|
||||
seen[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
out := make([]int64, 0, len(seen))
|
||||
for id := range seen {
|
||||
out = append(out, id)
|
||||
}
|
||||
sort.Slice(out, func(i, j int) bool { return out[i] < out[j] })
|
||||
return out, nil
|
||||
}
|
||||
@@ -743,7 +743,7 @@ func TestFederationRuntimeCommandPortRangeEnforcement(t *testing.T) {
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0 for reload command, got %d (msg: %s)", out.Code, out.Msg)
|
||||
if out.Code == 0 {
|
||||
t.Fatal("a shared-node token must not reload the entire provider node")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -38,6 +38,9 @@ func TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately(t *testing
|
||||
if err := repo.DB().Create(forward).Error; err != nil {
|
||||
t.Fatalf("seed forward: %v", err)
|
||||
}
|
||||
if err := repo.DB().Create(&model.ForwardPort{ForwardID: forward.ID, NodeID: node.ID, Port: 10000}).Error; err != nil {
|
||||
t.Fatalf("seed forward node ownership: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time) VALUES(2, 1, 0, ?, ?, ?, ?, 0, 0, '', ?, ?)`, bytesPerGB-100, bytesPerGB-100, dayKey, monthKey, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user_quota: %v", err)
|
||||
}
|
||||
|
||||
@@ -58,6 +58,9 @@ func TestFlowUploadInsertsTunnelMetrics(t *testing.T) {
|
||||
if err := repo.DB().Create(forward).Error; err != nil {
|
||||
t.Fatalf("seed forward: %v", err)
|
||||
}
|
||||
if err := repo.DB().Create(&model.ForwardPort{ForwardID: forward.ID, NodeID: node.ID, Port: 10000}).Error; err != nil {
|
||||
t.Fatalf("seed forward node ownership: %v", err)
|
||||
}
|
||||
|
||||
serviceName := jsonNumber(forward.ID) + "_123_0"
|
||||
body, _ := json.Marshal([]map[string]interface{}{{
|
||||
|
||||
+5
-3
@@ -125,11 +125,13 @@ func main() {
|
||||
|
||||
distro := socket.DetectDistro()
|
||||
fullVersion := fmt.Sprintf("%s (%s/%s)", version, distro, runtime.GOARCH)
|
||||
wsReporter := socket.StartWebSocketReporterWithConfig(config.Addr, config.Secret, config.Http, config.Tls, config.Socks, fullVersion)
|
||||
defer wsReporter.Stop()
|
||||
service.SetHTTPReportURL(config.Addr, config.Secret)
|
||||
|
||||
p := &program{}
|
||||
p := &program{
|
||||
startReporter: func() reporter {
|
||||
return socket.StartWebSocketReporterWithConfig(config.Addr, config.Secret, config.Http, config.Tls, config.Socks, fullVersion)
|
||||
},
|
||||
}
|
||||
if err := svc.Run(p); err != nil {
|
||||
logger.Default().Fatal(err)
|
||||
}
|
||||
|
||||
+88
-35
@@ -3,6 +3,14 @@ package main
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/go-gost/core/auth"
|
||||
"github.com/go-gost/core/logger"
|
||||
"github.com/go-gost/core/service"
|
||||
@@ -18,20 +26,22 @@ import (
|
||||
xservice "github.com/go-gost/x/service"
|
||||
"github.com/go-gost/x/socket"
|
||||
"github.com/judwhite/go-svc"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
)
|
||||
|
||||
type program struct {
|
||||
srvApi service.Service
|
||||
srvMetrics service.Service
|
||||
srvProfiling *http.Server
|
||||
type reporter interface {
|
||||
Stop()
|
||||
}
|
||||
|
||||
cancel context.CancelFunc
|
||||
type program struct {
|
||||
startReporter func() reporter
|
||||
reporter reporter
|
||||
srvApi service.Service
|
||||
srvMetrics service.Service
|
||||
srvProfiling *http.Server
|
||||
profilingListener net.Listener
|
||||
|
||||
cancel context.CancelFunc
|
||||
stopped bool
|
||||
}
|
||||
|
||||
func (p *program) Init(env svc.Environment) error {
|
||||
@@ -48,7 +58,15 @@ func (p *program) Init(env svc.Environment) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *program) Start() error {
|
||||
func (p *program) Start() (err error) {
|
||||
unlock := config.LockMutation()
|
||||
defer unlock()
|
||||
p.stopped = false
|
||||
defer func() {
|
||||
if err != nil {
|
||||
p.stopRuntime()
|
||||
}
|
||||
}()
|
||||
cfg, err := parser.Parse()
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -61,23 +79,28 @@ func (p *program) Start() error {
|
||||
os.Exit(0)
|
||||
}
|
||||
|
||||
config.Set(cfg)
|
||||
|
||||
if err := loader.Load(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Enable config persistence after initial load so runtime mutations
|
||||
// (AddService, UpdateService, DeleteService, etc.) are saved to disk.
|
||||
socket.EnableConfigPersist()
|
||||
|
||||
if err := p.run(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
config.Set(cfg)
|
||||
socket.EnableConfigPersist()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
p.cancel = cancel
|
||||
go p.reload(ctx)
|
||||
c := make(chan os.Signal, 1)
|
||||
signal.Notify(c, syscall.SIGHUP)
|
||||
go p.reload(ctx, c)
|
||||
|
||||
// A connected panel may immediately send commands. Only expose the agent
|
||||
// after initial config loading, runtime startup and persistence are ready.
|
||||
if p.startReporter != nil {
|
||||
p.reporter = p.startReporter()
|
||||
}
|
||||
|
||||
go func() {
|
||||
select {
|
||||
@@ -91,7 +114,14 @@ func (p *program) Start() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *program) run(cfg *config.Config) error {
|
||||
func (p *program) run(cfg *config.Config) (err error) {
|
||||
defer func() {
|
||||
if err != nil {
|
||||
// Auxiliary listeners may occupy ports required by the rollback config.
|
||||
// Release all resources opened by this attempt before rebuilding it.
|
||||
p.stopRuntime()
|
||||
}
|
||||
}()
|
||||
for _, svc := range registry.ServiceRegistry().GetAll() {
|
||||
svc := svc
|
||||
go func() {
|
||||
@@ -152,6 +182,10 @@ func (p *program) run(cfg *config.Config) error {
|
||||
|
||||
if p.srvProfiling != nil {
|
||||
p.srvProfiling.Close()
|
||||
if p.profilingListener != nil {
|
||||
p.profilingListener.Close()
|
||||
p.profilingListener = nil
|
||||
}
|
||||
p.srvProfiling = nil
|
||||
}
|
||||
if cfg.Profiling != nil {
|
||||
@@ -162,7 +196,12 @@ func (p *program) run(cfg *config.Config) error {
|
||||
s := &http.Server{
|
||||
Addr: addr,
|
||||
}
|
||||
ln, err := net.Listen("tcp", addr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.srvProfiling = s
|
||||
p.profilingListener = ln
|
||||
|
||||
go func() {
|
||||
defer s.Close()
|
||||
@@ -170,7 +209,7 @@ func (p *program) run(cfg *config.Config) error {
|
||||
log := logger.Default().WithFields(map[string]any{"kind": "service", "service": "@profiling"})
|
||||
|
||||
log.Info("listening on ", addr)
|
||||
if err := s.ListenAndServe(); !errors.Is(err, http.ErrServerClosed) {
|
||||
if err := s.Serve(ln); !errors.Is(err, http.ErrServerClosed) {
|
||||
log.Error(err)
|
||||
}
|
||||
}()
|
||||
@@ -184,30 +223,45 @@ func (p *program) Stop() error {
|
||||
p.cancel()
|
||||
}
|
||||
|
||||
for name, srv := range registry.ServiceRegistry().GetAll() {
|
||||
srv.Close()
|
||||
if p.reporter != nil {
|
||||
p.reporter.Stop()
|
||||
}
|
||||
unlock := config.LockMutation()
|
||||
defer unlock()
|
||||
p.stopped = true
|
||||
p.stopRuntime()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *program) stopRuntime() {
|
||||
for name := range registry.ServiceRegistry().GetAll() {
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
logger.Default().Debugf("service %s shutdown", name)
|
||||
}
|
||||
|
||||
if p.srvApi != nil {
|
||||
p.srvApi.Close()
|
||||
p.srvApi = nil
|
||||
logger.Default().Debug("service @api shutdown")
|
||||
}
|
||||
if p.srvMetrics != nil {
|
||||
p.srvMetrics.Close()
|
||||
p.srvMetrics = nil
|
||||
logger.Default().Debug("service @metrics shutdown")
|
||||
}
|
||||
if p.srvProfiling != nil {
|
||||
p.srvProfiling.Close()
|
||||
if p.profilingListener != nil {
|
||||
p.profilingListener.Close()
|
||||
p.profilingListener = nil
|
||||
}
|
||||
p.srvProfiling = nil
|
||||
logger.Default().Debug("service @profiling shutdown")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *program) reload(ctx context.Context) {
|
||||
c := make(chan os.Signal, 1)
|
||||
signal.Notify(c, syscall.SIGHUP)
|
||||
func (p *program) reload(ctx context.Context, c chan os.Signal) {
|
||||
defer signal.Stop(c)
|
||||
|
||||
for {
|
||||
select {
|
||||
@@ -225,13 +279,16 @@ func (p *program) reload(ctx context.Context) {
|
||||
}
|
||||
|
||||
func (p *program) reloadConfig() error {
|
||||
unlock := config.LockMutation()
|
||||
defer unlock()
|
||||
if p.stopped {
|
||||
return errors.New("agent is shutting down")
|
||||
}
|
||||
cfg, err := parser.Parse()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
config.Set(cfg)
|
||||
|
||||
if err := loader.Load(cfg); err != nil {
|
||||
if err := loader.Reload(cfg, p.run); err != nil {
|
||||
return err
|
||||
}
|
||||
activeServices := make(map[string]struct{}, len(cfg.Services))
|
||||
@@ -242,10 +299,6 @@ func (p *program) reloadConfig() error {
|
||||
}
|
||||
xservice.GetGlobalTrafficManager().RetainServices(activeServices)
|
||||
|
||||
if err := p.run(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,313 @@
|
||||
//go:build linux || darwin
|
||||
|
||||
package lifecycle_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
// Exercise the real binary: initial parsing is held at a FIFO while the panel
|
||||
// tries to send a rule as soon as the WebSocket connects. This reproduced the
|
||||
// old startup overwrite reliably without timing a large config load.
|
||||
func TestAgentStartupAndFailedReloadPreservePanelRules(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("builds and runs the real agent")
|
||||
}
|
||||
binary := filepath.Join(t.TempDir(), "gost")
|
||||
build := exec.Command("go", "build", "-o", binary, "../..")
|
||||
if output, err := build.CombinedOutput(); err != nil {
|
||||
t.Fatalf("build: %v\n%s", err, output)
|
||||
}
|
||||
dir := t.TempDir()
|
||||
early, baseline := address(t), address(t)
|
||||
apiAddr := address(t)
|
||||
panelConnections := make(chan *websocket.Conn, 2)
|
||||
wsReady := make(chan struct{})
|
||||
responses := make(chan map[string]any, 8)
|
||||
var ready sync.Once
|
||||
upgrader := websocket.Upgrader{}
|
||||
panel := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/system-info" {
|
||||
w.Write([]byte("ok"))
|
||||
return
|
||||
}
|
||||
c, err := upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer c.Close()
|
||||
ready.Do(func() { close(wsReady) })
|
||||
panelConnections <- c
|
||||
if err := c.WriteJSON(map[string]any{"type": "AddService", "requestId": "early-rule", "data": []any{service("70_1_0", early)}}); err != nil {
|
||||
return
|
||||
}
|
||||
for {
|
||||
_, payload, err := c.ReadMessage()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
env := decodeResponse(t, payload)
|
||||
if env["requestId"] == "early-rule" {
|
||||
responses <- env
|
||||
}
|
||||
}
|
||||
}))
|
||||
defer panel.Close()
|
||||
writeJSON(t, filepath.Join(dir, "config.json"), map[string]any{"addr": panel.URL, "secret": "audit-secret", "http": 1, "tls": 1, "socks": 1})
|
||||
fifo := filepath.Join(dir, "delayed.json")
|
||||
if err := syscall.Mkfifo(fifo, 0600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
logPath := filepath.Join(dir, "agent.log")
|
||||
logfile, err := os.Create(logPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer logfile.Close()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
agent := exec.CommandContext(ctx, binary, "-C", fifo)
|
||||
agent.Dir, agent.Stdout, agent.Stderr = dir, logfile, logfile
|
||||
if err := agent.Start(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
exited := make(chan error, 1)
|
||||
go func() { exited <- agent.Wait() }()
|
||||
defer func() {
|
||||
agent.Process.Kill()
|
||||
if t.Failed() {
|
||||
b, _ := os.ReadFile(logPath)
|
||||
t.Logf("agent log:\n%s", b)
|
||||
}
|
||||
}()
|
||||
select {
|
||||
case <-wsReady:
|
||||
t.Fatal("panel connected before initial config loaded")
|
||||
case err := <-exited:
|
||||
t.Fatalf("agent exited during startup: %v", err)
|
||||
case <-time.After(500 * time.Millisecond):
|
||||
}
|
||||
boot := map[string]any{"services": []any{service("71_1_0", baseline)}, "api": map[string]any{"addr": apiAddr}}
|
||||
writeJSON(t, filepath.Join(dir, "gost.json"), boot)
|
||||
writeFIFO(t, fifo, boot)
|
||||
select {
|
||||
case <-wsReady:
|
||||
case <-time.After(10 * time.Second):
|
||||
t.Fatal("panel did not connect after startup")
|
||||
}
|
||||
select {
|
||||
case response := <-responses:
|
||||
if response["success"] != true {
|
||||
t.Fatalf("AddService failed: %v", response)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("missing AddService response")
|
||||
}
|
||||
await(t, "both startup and panel listeners", func() bool { return listening(early) && listening(baseline) })
|
||||
saved, err := os.ReadFile(fifo)
|
||||
if err != nil || !strings.Contains(string(saved), "70_1_0") {
|
||||
t.Fatalf("acknowledged rule not persisted: %v", err)
|
||||
}
|
||||
|
||||
// A valid listener plus an invalid handler also exercises rollback of a
|
||||
// partially initialized service which already bound the original port.
|
||||
invalid := service("71_1_0", baseline)
|
||||
invalid["handler"] = map[string]any{"type": "handler-does-not-exist"}
|
||||
// The first persisted mutation atomically replaced the FIFO with a regular
|
||||
// config file, so subsequent reloads use the same real persistence path.
|
||||
writeJSON(t, fifo, map[string]any{"services": []any{invalid}})
|
||||
if err := agent.Process.Signal(syscall.SIGHUP); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
await(t, "failed reload rollback", func() bool {
|
||||
b, _ := os.ReadFile(logPath)
|
||||
return strings.Contains(string(b), "previous config restored")
|
||||
})
|
||||
if !listening(early) || !listening(baseline) {
|
||||
t.Fatal("failed reload lost a previously acknowledged listener")
|
||||
}
|
||||
|
||||
// Starting the candidate API on the old service port must not prevent
|
||||
// rollback when a later auxiliary listener fails to initialize.
|
||||
logBefore, _ := os.ReadFile(logPath)
|
||||
writeJSON(t, fifo, map[string]any{
|
||||
"services": []any{service("candidate", address(t))},
|
||||
"api": map[string]any{"addr": baseline},
|
||||
"metrics": map[string]any{"addr": "127.0.0.1:not-a-port"},
|
||||
})
|
||||
if err := agent.Process.Signal(syscall.SIGHUP); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
await(t, "auxiliary listener rollback", func() bool {
|
||||
b, _ := os.ReadFile(logPath)
|
||||
return strings.Count(string(b), "previous config restored") > strings.Count(string(logBefore), "previous config restored")
|
||||
})
|
||||
if !listening(early) || !listening(baseline) || !listening(apiAddr) {
|
||||
t.Fatal("candidate auxiliary listener prevented rollback")
|
||||
}
|
||||
|
||||
// An authenticated API request that never finishes its body must not hold
|
||||
// the runtime transaction lock against WS mutations or process shutdown.
|
||||
slow, err := net.Dial("tcp", apiAddr)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer slow.Close()
|
||||
if _, err = fmt.Fprintf(slow, "POST /config/services HTTP/1.1\r\nHost: localhost\r\nAuthorization: Basic dGVzdDp0ZXN0\r\nContent-Type: application/json\r\nContent-Length: 100000\r\n\r\n{"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
panelConn := <-panelConnections
|
||||
if err := panelConn.WriteJSON(map[string]any{"type": "DeleteService", "requestId": "slow-body-check", "data": map[string]any{"services": []string{"70_1_0"}}}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
await(t, "WS mutation despite slow API upload", func() bool { return !listening(early) })
|
||||
if err := agent.Process.Signal(syscall.SIGTERM); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
select {
|
||||
case err := <-exited:
|
||||
if err != nil {
|
||||
t.Fatalf("shutdown: %v", err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("agent did not stop")
|
||||
}
|
||||
if listening(early) || listening(baseline) {
|
||||
t.Fatal("shutdown left listeners open")
|
||||
}
|
||||
|
||||
// Startup errors must exit without advertising a node ready to accept rules.
|
||||
bad := filepath.Join(dir, "invalid.json")
|
||||
writeJSON(t, bad, map[string]any{"services": []any{invalid}})
|
||||
failed := exec.CommandContext(ctx, binary, "-C", bad)
|
||||
failed.Dir = dir
|
||||
if output, err := failed.CombinedOutput(); err == nil {
|
||||
t.Fatalf("invalid startup succeeded: %s", output)
|
||||
}
|
||||
select {
|
||||
case response := <-responses:
|
||||
t.Fatalf("failed startup accepted panel command: %v", response)
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func service(name, addr string) map[string]any {
|
||||
return map[string]any{"name": name, "addr": addr, "listener": map[string]any{"type": "tcp"}, "handler": map[string]any{"type": "auto"}}
|
||||
}
|
||||
func address(t *testing.T) string {
|
||||
t.Helper()
|
||||
l, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer l.Close()
|
||||
return l.Addr().String()
|
||||
}
|
||||
func listening(addr string) bool {
|
||||
c, err := net.DialTimeout("tcp", addr, 50*time.Millisecond)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
c.Close()
|
||||
return true
|
||||
}
|
||||
func await(t *testing.T, label string, pred func() bool) {
|
||||
t.Helper()
|
||||
for until := time.Now().Add(5 * time.Second); time.Now().Before(until); time.Sleep(10 * time.Millisecond) {
|
||||
if pred() {
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatal("timeout: " + label)
|
||||
}
|
||||
func writeJSON(t *testing.T, path string, value any) {
|
||||
t.Helper()
|
||||
b, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(path, b, 0600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
func writeFIFO(t *testing.T, path string, value any) {
|
||||
t.Helper()
|
||||
b, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
f, err := os.OpenFile(path, os.O_WRONLY, 0)
|
||||
if err != nil {
|
||||
done <- err
|
||||
return
|
||||
}
|
||||
_, err = f.Write(b)
|
||||
f.Close()
|
||||
done <- err
|
||||
}()
|
||||
select {
|
||||
case err := <-done:
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("agent did not read config FIFO")
|
||||
}
|
||||
}
|
||||
func decodeResponse(t *testing.T, payload []byte) map[string]any {
|
||||
t.Helper()
|
||||
env := map[string]any{}
|
||||
if err := json.Unmarshal(payload, &env); err != nil {
|
||||
t.Error(err)
|
||||
return nil
|
||||
}
|
||||
if encrypted, _ := env["encrypted"].(bool); !encrypted {
|
||||
return env
|
||||
}
|
||||
encoded, _ := env["data"].(string)
|
||||
raw, err := base64.StdEncoding.DecodeString(encoded)
|
||||
if err != nil {
|
||||
t.Error(err)
|
||||
return nil
|
||||
}
|
||||
hash := sha256.Sum256([]byte("audit-secret"))
|
||||
block, _ := aes.NewCipher(hash[:])
|
||||
gcm, _ := cipher.NewGCM(block)
|
||||
if len(raw) < gcm.NonceSize() {
|
||||
t.Error("short encrypted response")
|
||||
return nil
|
||||
}
|
||||
payload, err = gcm.Open(nil, raw[:gcm.NonceSize()], raw[gcm.NonceSize():], nil)
|
||||
if err != nil {
|
||||
t.Error(err)
|
||||
return nil
|
||||
}
|
||||
env = map[string]any{}
|
||||
if err := json.Unmarshal(payload, &env); err != nil {
|
||||
t.Error(err)
|
||||
return nil
|
||||
}
|
||||
return env
|
||||
}
|
||||
@@ -52,7 +52,7 @@ func Register(r *gin.Engine, opts *Options) {
|
||||
router.StaticFS("/docs", http.FS(swaggerDoc))
|
||||
|
||||
config := router.Group("/config")
|
||||
config.Use(mwBasicAuth(opts.Auther))
|
||||
config.Use(mwBasicAuth(opts.Auther), configTransaction())
|
||||
|
||||
config.GET("", getConfig)
|
||||
config.POST("", saveConfig)
|
||||
|
||||
@@ -37,9 +37,12 @@ func reloadConfig(ctx *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
config.Set(cfg)
|
||||
|
||||
if err := loader.Load(cfg); err != nil {
|
||||
if err := loader.Reload(cfg, func(*config.Config) error {
|
||||
for _, svc := range registry.ServiceRegistry().GetAll() {
|
||||
go svc.Serve()
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
writeError(ctx, NewError(http.StatusBadRequest, ErrCodeInvalid, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -51,13 +54,6 @@ func reloadConfig(ctx *gin.Context) {
|
||||
}
|
||||
xservice.GetGlobalTrafficManager().RetainServices(activeServices)
|
||||
|
||||
for _, svc := range registry.ServiceRegistry().GetAll() {
|
||||
svc := svc
|
||||
go func() {
|
||||
svc.Serve()
|
||||
}()
|
||||
}
|
||||
|
||||
ctx.JSON(http.StatusOK, Response{
|
||||
Msg: "OK",
|
||||
})
|
||||
|
||||
@@ -0,0 +1,90 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/go-gost/x/config"
|
||||
)
|
||||
|
||||
const maxConfigRequestBody = 16 << 20
|
||||
|
||||
// Read the complete bounded request before acquiring the runtime transaction
|
||||
// lock. Buffer the response until after it is released: neither a slow upload
|
||||
// nor a client that stops reading may block panel commands, reload or shutdown.
|
||||
func configTransaction() gin.HandlerFunc {
|
||||
return func(ctx *gin.Context) {
|
||||
if ctx.Request.Body != nil {
|
||||
body := http.MaxBytesReader(ctx.Writer, ctx.Request.Body, maxConfigRequestBody)
|
||||
data, err := io.ReadAll(body)
|
||||
body.Close()
|
||||
if err != nil {
|
||||
status := http.StatusBadRequest
|
||||
var tooLarge *http.MaxBytesError
|
||||
if errors.As(err, &tooLarge) {
|
||||
status = http.StatusRequestEntityTooLarge
|
||||
}
|
||||
ctx.AbortWithStatusJSON(status, Response{Code: status, Msg: "Unable to read configuration request"})
|
||||
return
|
||||
}
|
||||
ctx.Request.Body = io.NopCloser(bytes.NewReader(data))
|
||||
}
|
||||
|
||||
writer := ctx.Writer
|
||||
buffered := &configResponseWriter{ResponseWriter: writer, header: writer.Header().Clone(), status: http.StatusOK, size: -1}
|
||||
ctx.Writer = buffered
|
||||
defer func() { ctx.Writer = writer }()
|
||||
func() {
|
||||
unlock := config.LockMutation()
|
||||
defer unlock()
|
||||
// A request waiting behind reload may have been closed during shutdown.
|
||||
if ctx.Request.Context().Err() != nil {
|
||||
ctx.Abort()
|
||||
return
|
||||
}
|
||||
ctx.Next()
|
||||
}()
|
||||
ctx.Writer = writer
|
||||
for key, values := range buffered.header {
|
||||
writer.Header()[key] = values
|
||||
}
|
||||
writer.WriteHeader(buffered.status)
|
||||
writer.Write(buffered.body.Bytes())
|
||||
}
|
||||
}
|
||||
|
||||
// Config endpoints return JSON rather than streaming. Preserve Gin's response
|
||||
// bookkeeping while delaying all network writes until the transaction ends.
|
||||
type configResponseWriter struct {
|
||||
gin.ResponseWriter
|
||||
header http.Header
|
||||
body bytes.Buffer
|
||||
status int
|
||||
size int
|
||||
}
|
||||
|
||||
func (w *configResponseWriter) Header() http.Header { return w.header }
|
||||
func (w *configResponseWriter) WriteHeader(status int) {
|
||||
if !w.Written() && status > 0 {
|
||||
w.status = status
|
||||
}
|
||||
}
|
||||
func (w *configResponseWriter) WriteHeaderNow() {
|
||||
if !w.Written() {
|
||||
w.size = 0
|
||||
}
|
||||
}
|
||||
func (w *configResponseWriter) Write(p []byte) (int, error) {
|
||||
w.WriteHeaderNow()
|
||||
n, err := w.body.Write(p)
|
||||
w.size += n
|
||||
return n, err
|
||||
}
|
||||
func (w *configResponseWriter) WriteString(s string) (int, error) { return w.Write([]byte(s)) }
|
||||
func (w *configResponseWriter) Status() int { return w.status }
|
||||
func (w *configResponseWriter) Size() int { return w.size }
|
||||
func (w *configResponseWriter) Written() bool { return w.size >= 0 }
|
||||
func (w *configResponseWriter) Flush() { w.WriteHeaderNow() }
|
||||
@@ -0,0 +1,76 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/go-gost/x/config"
|
||||
)
|
||||
|
||||
func assertMutationAvailable(t *testing.T) {
|
||||
t.Helper()
|
||||
done := make(chan struct{})
|
||||
go func() { unlock := config.LockMutation(); unlock(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("network I/O holds runtime mutation lock")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigTransactionDoesNotLockWhileReadingBody(t *testing.T) {
|
||||
router := gin.New()
|
||||
router.Use(configTransaction())
|
||||
router.POST("/config", func(c *gin.Context) { c.JSON(http.StatusOK, Response{Msg: "OK"}) })
|
||||
reader, writer := io.Pipe()
|
||||
defer reader.Close()
|
||||
defer writer.Close()
|
||||
request := httptest.NewRequest(http.MethodPost, "/config", reader)
|
||||
done := make(chan struct{})
|
||||
go func() { defer close(done); router.ServeHTTP(httptest.NewRecorder(), request) }()
|
||||
if _, err := writer.Write([]byte("{")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertMutationAvailable(t)
|
||||
writer.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("handler did not finish")
|
||||
}
|
||||
}
|
||||
|
||||
type blockedResponse struct {
|
||||
header http.Header
|
||||
started chan struct{}
|
||||
release chan struct{}
|
||||
}
|
||||
|
||||
func (w *blockedResponse) Header() http.Header { return w.header }
|
||||
func (w *blockedResponse) WriteHeader(int) {}
|
||||
func (w *blockedResponse) Write(p []byte) (int, error) {
|
||||
close(w.started)
|
||||
<-w.release
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func TestConfigTransactionReleasesLockBeforeSendingResponse(t *testing.T) {
|
||||
router := gin.New()
|
||||
router.Use(configTransaction())
|
||||
router.POST("/config", func(c *gin.Context) { c.JSON(http.StatusOK, Response{Msg: "OK"}) })
|
||||
writer := &blockedResponse{header: make(http.Header), started: make(chan struct{}), release: make(chan struct{})}
|
||||
defer close(writer.release)
|
||||
request := httptest.NewRequest(http.MethodPost, "/config", strings.NewReader("{}"))
|
||||
go router.ServeHTTP(writer, request)
|
||||
select {
|
||||
case <-writer.started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("response did not start")
|
||||
}
|
||||
assertMutationAvailable(t)
|
||||
}
|
||||
@@ -3,6 +3,7 @@ package service
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/go-gost/core/auth"
|
||||
@@ -67,7 +68,11 @@ func NewService(network, addr string, opts ...Option) (service.Service, error) {
|
||||
|
||||
return &server{
|
||||
s: &http.Server{
|
||||
Handler: r,
|
||||
Handler: r,
|
||||
ReadHeaderTimeout: 5 * time.Second,
|
||||
ReadTimeout: 15 * time.Second,
|
||||
WriteTimeout: 30 * time.Second,
|
||||
IdleTimeout: 60 * time.Second,
|
||||
},
|
||||
ln: ln,
|
||||
cclose: make(chan struct{}),
|
||||
@@ -83,7 +88,11 @@ func (s *server) Addr() net.Addr {
|
||||
}
|
||||
|
||||
func (s *server) Close() error {
|
||||
return s.s.Close()
|
||||
// Close can race the goroutine entering Serve during a failed startup.
|
||||
// http.Server.Close alone does not own the listener until Serve starts.
|
||||
err := s.s.Close()
|
||||
s.ln.Close()
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *server) IsClosed() bool {
|
||||
|
||||
@@ -23,10 +23,20 @@ func init() {
|
||||
}
|
||||
|
||||
var (
|
||||
global = &Config{}
|
||||
globalMux sync.RWMutex
|
||||
global = &Config{}
|
||||
globalMux sync.RWMutex
|
||||
mutationMux sync.Mutex
|
||||
)
|
||||
|
||||
// LockMutation serializes complete runtime/config transactions across panel
|
||||
// commands, management API requests, startup and reload. It is separate from
|
||||
// globalMux so callers may safely use Global, Set and OnUpdate while holding it.
|
||||
// Lock before reading/parsing the config and hold through persistence/rollback.
|
||||
func LockMutation() func() {
|
||||
mutationMux.Lock()
|
||||
return mutationMux.Unlock
|
||||
}
|
||||
|
||||
func Global() *Config {
|
||||
globalMux.RLock()
|
||||
defer globalMux.RUnlock()
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
package loader
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/go-gost/core/logger"
|
||||
"github.com/go-gost/x/config"
|
||||
"github.com/go-gost/x/config/parsing"
|
||||
@@ -30,6 +32,34 @@ func Load(cfg *config.Config) error {
|
||||
return defaultLoader.Load(cfg)
|
||||
}
|
||||
|
||||
// Reload replaces the runtime and commits the config only after it starts.
|
||||
// The caller must hold config.LockMutation, including while parsing cfg, so the
|
||||
// rollback snapshot includes every previously acknowledged runtime command.
|
||||
// Failed loads can partially replace registries and bind listeners; always
|
||||
// rebuild the last successful snapshot before returning an error.
|
||||
func Reload(cfg *config.Config, run func(*config.Config) error) error {
|
||||
previous := config.Global()
|
||||
if err := apply(cfg, run); err != nil {
|
||||
// A failed partial load may have left both old and new listeners behind.
|
||||
for name := range registry.ServiceRegistry().GetAll() {
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
}
|
||||
if rollbackErr := apply(previous, run); rollbackErr != nil {
|
||||
return fmt.Errorf("reload failed: %w; restore previous config failed: %v", err, rollbackErr)
|
||||
}
|
||||
return fmt.Errorf("reload failed (previous config restored): %w", err)
|
||||
}
|
||||
config.Set(cfg)
|
||||
return nil
|
||||
}
|
||||
|
||||
func apply(cfg *config.Config, run func(*config.Config) error) error {
|
||||
if err := Load(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
return run(cfg)
|
||||
}
|
||||
|
||||
type loader struct{}
|
||||
|
||||
func (l *loader) Load(cfg *config.Config) error {
|
||||
@@ -217,12 +247,19 @@ func register(cfg *config.Config) error {
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
}
|
||||
for _, svcCfg := range cfg.Services {
|
||||
if svcCfg == nil {
|
||||
return fmt.Errorf("service config is nil")
|
||||
}
|
||||
if paused, _ := svcCfg.Metadata["paused"].(bool); paused {
|
||||
continue
|
||||
}
|
||||
svc, err := service_parser.ParseService(svcCfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if svc != nil {
|
||||
if err := registry.ServiceRegistry().Register(svcCfg.Name, svc); err != nil {
|
||||
svc.Close()
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
package loader_test
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/go-gost/x/config"
|
||||
"github.com/go-gost/x/config/loader"
|
||||
_ "github.com/go-gost/x/handler/auto"
|
||||
_ "github.com/go-gost/x/listener/tcp"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
func TestReloadRestoresPreviousRuntime(t *testing.T) {
|
||||
for _, failure := range []string{"listener", "handler", "partial", "run"} {
|
||||
t.Run(failure, func(t *testing.T) {
|
||||
unlock := config.LockMutation()
|
||||
defer unlock()
|
||||
original := config.Global()
|
||||
defer config.Set(original)
|
||||
defer func() {
|
||||
for name := range registry.ServiceRegistry().GetAll() {
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
}
|
||||
}()
|
||||
old := &config.Config{Services: []*config.ServiceConfig{
|
||||
{Name: "shared-rule", Addr: "127.0.0.1:0", Handler: &config.HandlerConfig{Type: "auto"}, Listener: &config.ListenerConfig{Type: "tcp"}},
|
||||
{Name: "paused-rule", Addr: "127.0.0.1:0", Metadata: map[string]any{"paused": true}},
|
||||
}}
|
||||
if err := loader.Load(old); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
addr := registry.ServiceRegistry().Get("shared-rule").Addr().String()
|
||||
old.Services[0].Addr = addr
|
||||
config.Set(old)
|
||||
serve := func(*config.Config) error {
|
||||
for _, svc := range registry.ServiceRegistry().GetAll() {
|
||||
go svc.Serve()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
serve(old)
|
||||
replacement := config.Global()
|
||||
replacement.Services = replacement.Services[:1]
|
||||
switch failure {
|
||||
case "listener":
|
||||
replacement.Services[0].Listener.Type = "invalid-listener"
|
||||
case "handler":
|
||||
replacement.Services[0].Handler.Type = "invalid-handler"
|
||||
case "partial":
|
||||
replacement.Services = append(replacement.Services, &config.ServiceConfig{Name: "broken", Listener: &config.ListenerConfig{Type: "invalid-listener"}})
|
||||
}
|
||||
run := serve
|
||||
if failure == "run" {
|
||||
run = func(cfg *config.Config) error {
|
||||
if cfg == replacement {
|
||||
return &net.AddrError{Err: "auxiliary listener failed", Addr: "test"}
|
||||
}
|
||||
return serve(cfg)
|
||||
}
|
||||
}
|
||||
if err := loader.Reload(replacement, run); err == nil {
|
||||
t.Fatal("expected reload failure")
|
||||
}
|
||||
restored := registry.ServiceRegistry().Get("shared-rule")
|
||||
if restored == nil || restored.Addr().String() != addr {
|
||||
t.Fatal("previous service was not restored on its original port")
|
||||
}
|
||||
if registry.ServiceRegistry().Get("paused-rule") != nil {
|
||||
t.Fatal("rollback resumed a paused service")
|
||||
}
|
||||
conn, err := net.DialTimeout("tcp", addr, time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("restored listener unavailable: %v", err)
|
||||
}
|
||||
conn.Close()
|
||||
got := config.Global()
|
||||
if len(got.Services) != 2 || got.Services[0].Listener.Type != "tcp" || got.Services[0].Handler.Type != "auto" {
|
||||
t.Fatalf("failed config was committed: %+v", got.Services)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -255,6 +255,15 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Listener initialization binds the port. If handler/TLS/forwarder parsing
|
||||
// fails, release it so a reload rollback can restore the previous listener.
|
||||
configured := false
|
||||
defer func() {
|
||||
if !configured {
|
||||
ln.Close()
|
||||
}
|
||||
}()
|
||||
|
||||
handlerLogger := serviceLogger.WithFields(map[string]any{
|
||||
"kind": "handler",
|
||||
})
|
||||
@@ -379,6 +388,7 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) {
|
||||
)
|
||||
|
||||
serviceLogger.Infof("listening on %s/%s", s.Addr().String(), s.Addr().Network())
|
||||
configured = true
|
||||
return s, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -84,7 +84,9 @@ func (s *metricService) Addr() net.Addr {
|
||||
}
|
||||
|
||||
func (s *metricService) Close() error {
|
||||
return s.s.Close()
|
||||
err := s.s.Close()
|
||||
s.ln.Close()
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *metricService) IsClosed() bool {
|
||||
|
||||
@@ -370,43 +370,43 @@ type getConfigResponse struct {
|
||||
|
||||
// getConfigData 获取配置数据(避免循环依赖)
|
||||
func getConfigData() ([]byte, error) {
|
||||
config.OnUpdate(func(c *config.Config) error {
|
||||
for _, svc := range c.Services {
|
||||
if svc == nil {
|
||||
continue
|
||||
// Reporting is read-only: enriching a detached snapshot must not persist
|
||||
// stale state over a config currently being reloaded or acknowledged.
|
||||
cfg := config.Global()
|
||||
for _, svc := range cfg.Services {
|
||||
if svc == nil {
|
||||
continue
|
||||
}
|
||||
s := registry.ServiceRegistry().Get(svc.Name)
|
||||
ss, ok := s.(serviceStatus)
|
||||
if ok && ss != nil {
|
||||
status := ss.Status()
|
||||
svc.Status = &config.ServiceStatus{
|
||||
CreateTime: status.CreateTime().Unix(),
|
||||
State: string(status.State()),
|
||||
}
|
||||
s := registry.ServiceRegistry().Get(svc.Name)
|
||||
ss, ok := s.(serviceStatus)
|
||||
if ok && ss != nil {
|
||||
status := ss.Status()
|
||||
svc.Status = &config.ServiceStatus{
|
||||
CreateTime: status.CreateTime().Unix(),
|
||||
State: string(status.State()),
|
||||
if st := status.Stats(); st != nil {
|
||||
svc.Status.Stats = &config.ServiceStats{
|
||||
TotalConns: st.Get(stats.KindTotalConns),
|
||||
CurrentConns: st.Get(stats.KindCurrentConns),
|
||||
TotalErrs: st.Get(stats.KindTotalErrs),
|
||||
InputBytes: st.Get(stats.KindInputBytes),
|
||||
OutputBytes: st.Get(stats.KindOutputBytes),
|
||||
}
|
||||
if st := status.Stats(); st != nil {
|
||||
svc.Status.Stats = &config.ServiceStats{
|
||||
TotalConns: st.Get(stats.KindTotalConns),
|
||||
CurrentConns: st.Get(stats.KindCurrentConns),
|
||||
TotalErrs: st.Get(stats.KindTotalErrs),
|
||||
InputBytes: st.Get(stats.KindInputBytes),
|
||||
OutputBytes: st.Get(stats.KindOutputBytes),
|
||||
}
|
||||
}
|
||||
for _, ev := range status.Events() {
|
||||
if !ev.Time.IsZero() {
|
||||
svc.Status.Events = append(svc.Status.Events, config.ServiceEvent{
|
||||
Time: ev.Time.Unix(),
|
||||
Msg: ev.Message,
|
||||
})
|
||||
}
|
||||
}
|
||||
for _, ev := range status.Events() {
|
||||
if !ev.Time.IsZero() {
|
||||
svc.Status.Events = append(svc.Status.Events, config.ServiceEvent{
|
||||
Time: ev.Time.Unix(),
|
||||
Msg: ev.Message,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
var resp getConfigResponse
|
||||
resp.Config = config.Global()
|
||||
resp.Config = cfg
|
||||
|
||||
buf := &bytes.Buffer{}
|
||||
resp.Config.Write(buf, "json")
|
||||
|
||||
@@ -1,10 +1,14 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"github.com/go-gost/x/config"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -148,3 +152,31 @@ func TestPostJSONWithFallbackRemembersDetectedURL(t *testing.T) {
|
||||
t.Fatalf("expected remembered http url first, got %s", calls[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigReportDoesNotPersistStaleRuntime(t *testing.T) {
|
||||
previous, path := config.Global(), config.PersistPath()
|
||||
defer config.Set(previous)
|
||||
defer config.SetPersistPath(path)
|
||||
config.Set(&config.Config{Services: []*config.ServiceConfig{{Name: "old-runtime"}}})
|
||||
filename := filepath.Join(t.TempDir(), "gost.json")
|
||||
candidate := []byte(`{"services":[{"name":"new-on-disk"}]}`)
|
||||
if err := os.WriteFile(filename, candidate, 0600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
config.SetPersistPath(filename)
|
||||
config.EnablePersist()
|
||||
report, err := getConfigData()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !bytes.Contains(report, []byte("old-runtime")) {
|
||||
t.Fatalf("unexpected report: %s", report)
|
||||
}
|
||||
saved, err := os.ReadFile(filename)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !bytes.Equal(saved, candidate) {
|
||||
t.Fatalf("report overwrote candidate config: %s", saved)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -115,3 +115,39 @@ func TestMutationQueueExecutesCommandsInArrivalOrder(t *testing.T) {
|
||||
}
|
||||
close(second.release)
|
||||
}
|
||||
|
||||
func TestMutationCommandWaitsForRuntimeTransaction(t *testing.T) {
|
||||
original := config.Global()
|
||||
defer config.Set(original)
|
||||
name := "reload_transaction_service"
|
||||
svc := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})}
|
||||
close(svc.release)
|
||||
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer registry.ServiceRegistry().Unregister(name)
|
||||
config.Set(&config.Config{Services: []*config.ServiceConfig{{Name: name}}})
|
||||
reporter := NewWebSocketReporter("", "transaction-test-secret")
|
||||
defer reporter.Stop()
|
||||
unlock := config.LockMutation()
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
reporter.routeCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{name}}})
|
||||
}()
|
||||
select {
|
||||
case <-svc.started:
|
||||
unlock()
|
||||
t.Fatal("command interleaved with runtime transaction")
|
||||
case <-time.After(50 * time.Millisecond):
|
||||
}
|
||||
unlock()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("command did not resume after transaction")
|
||||
}
|
||||
if registry.ServiceRegistry().Get(name) != nil {
|
||||
t.Fatal("service was not removed")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -179,6 +179,7 @@ type WebSocketReporter struct {
|
||||
tcpPingSem chan struct{} // 限制诊断探测并发,避免离线目标耗尽连接
|
||||
readCommandSem chan struct{} // 限制只读命令并发,避免诊断请求耗尽资源
|
||||
mutationQueue chan CommandMessage
|
||||
workers sync.WaitGroup
|
||||
}
|
||||
|
||||
var wsDial = func(dialer *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
|
||||
@@ -238,8 +239,15 @@ func (w *WebSocketReporter) releaseTCPPingSlot() {
|
||||
|
||||
// Start 启动WebSocket报告器
|
||||
func (w *WebSocketReporter) Start() {
|
||||
go w.runMutationCommands()
|
||||
go w.run()
|
||||
w.workers.Add(2)
|
||||
go func() {
|
||||
defer w.workers.Done()
|
||||
w.runMutationCommands()
|
||||
}()
|
||||
go func() {
|
||||
defer w.workers.Done()
|
||||
w.run()
|
||||
}()
|
||||
}
|
||||
|
||||
// Stop 停止WebSocket报告器
|
||||
@@ -250,6 +258,7 @@ func (w *WebSocketReporter) Stop() {
|
||||
w.conn.Close()
|
||||
}
|
||||
w.connMutex.Unlock()
|
||||
w.workers.Wait()
|
||||
}
|
||||
|
||||
// backoffWithJitter 返回带随机抖动的退避时间(±25%)
|
||||
@@ -339,10 +348,10 @@ func (w *WebSocketReporter) connect() error {
|
||||
|
||||
candidates := buildWebSocketCandidates(w.addr, w.secret, w.version, cfg.Http, cfg.Tls, cfg.Socks, w.preferredWSScheme)
|
||||
|
||||
dialer := websocket.DefaultDialer
|
||||
dialer := *websocket.DefaultDialer
|
||||
dialer.HandshakeTimeout = 10 * time.Second
|
||||
|
||||
conn, usedURL, err := dialWebSocketWithFallback(dialer, candidates)
|
||||
conn, usedURL, err := dialWebSocketWithFallback(&dialer, candidates)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -528,7 +537,11 @@ func (w *WebSocketReporter) handleConnection() {
|
||||
}()
|
||||
|
||||
// 启动消息接收goroutine
|
||||
go w.receiveMessages()
|
||||
w.workers.Add(1)
|
||||
go func() {
|
||||
defer w.workers.Done()
|
||||
w.receiveMessages()
|
||||
}()
|
||||
|
||||
// 指标上报 ticker
|
||||
metricTicker := time.NewTicker(w.pingInterval)
|
||||
@@ -887,6 +900,14 @@ func isMutationCommand(commandType string) bool {
|
||||
|
||||
// routeCommand 路由命令到对应的处理函数
|
||||
func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
|
||||
if isMutationCommand(cmd.Type) {
|
||||
unlock := config.LockMutation()
|
||||
defer unlock()
|
||||
if w.ctx.Err() != nil {
|
||||
w.sendCommandFailure(cmd, "Agent is shutting down")
|
||||
return
|
||||
}
|
||||
}
|
||||
jsonBytes, errs := json.Marshal(cmd)
|
||||
if errs != nil {
|
||||
fmt.Println("Error marshaling JSON:", errs)
|
||||
|
||||
Reference in New Issue
Block a user