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

This commit is contained in:
sagit
2026-09-29 14:05:49 +08:00
committed by GitHub
parent 129fa0aa4c
commit c952d2fb3a
49 changed files with 5053 additions and 643 deletions
+174 -153
View File
@@ -1,8 +1,11 @@
package handler
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/url"
@@ -408,9 +411,26 @@ func (h *Handler) federationShareDelete(w http.ResponseWriter, r *http.Request)
return
}
share, _ := h.repo.GetPeerShare(req.ID)
share, err := h.repo.GetPeerShare(req.ID)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
if share != nil {
// Revoke new allocations before cleanup; failed deletions keep the
// disabled share and its reservations available for retry.
share.IsActive = 0
share.UpdatedTime = time.Now().UnixMilli()
if err := h.repo.UpdatePeerShare(share); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
}
h.cleanupPeerShareRuntimes(req.ID)
if err := h.cleanupPeerShareRuntimes(req.ID); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
h.cleanupFederationTunnels(req.ID)
if err := h.repo.DeletePeerShare(req.ID); err != nil {
@@ -418,10 +438,6 @@ func (h *Handler) federationShareDelete(w http.ResponseWriter, r *http.Request)
return
}
if share != nil && h.wsServer != nil {
h.wsServer.SendCommand(share.NodeID, "reload", nil, time.Second*5)
}
response.WriteJSON(w, response.OKEmpty())
}
@@ -805,16 +821,6 @@ func (h *Handler) authPeer(next http.HandlerFunc) http.HandlerFunc {
return
}
if share.IsActive == 0 {
response.WriteJSON(w, response.Err(403, "Share is disabled"))
return
}
if share.ExpiryTime > 0 && share.ExpiryTime < time.Now().UnixMilli() {
response.WriteJSON(w, response.Err(403, "Share expired"))
return
}
if strings.TrimSpace(share.AllowedIPs) != "" {
clientIP := resolvePeerClientIP(r)
if clientIP == nil {
@@ -847,10 +853,48 @@ func (h *Handler) authPeer(next http.HandlerFunc) http.HandlerFunc {
}
}
if share.IsActive != 1 || (share.ExpiryTime > 0 && share.ExpiryTime <= time.Now().UnixMilli()) || isPeerShareFlowExceeded(share) {
if !isFederationCleanupRequest(r) {
response.WriteJSON(w, response.Err(403, "Share is inactive, expired, or over quota"))
return
}
}
next(w, r)
}
}
// Revoked allocation privileges must not revoke cleanup privileges. Inspect
// only supported destructive commands, preserving the body for the handler.
func isFederationCleanupRequest(r *http.Request) bool {
if r.Method != http.MethodPost {
return false
}
switch r.URL.Path {
case "/api/v1/federation/runtime/release-role":
return true
case "/api/v1/federation/runtime/command":
if r.Body == nil {
return false
}
var copied bytes.Buffer
original := r.Body
var request federationRuntimeCommandRequest
err := json.NewDecoder(io.TeeReader(io.LimitReader(original, 1<<20), &copied)).Decode(&request)
r.Body = struct {
io.Reader
io.Closer
}{Reader: io.MultiReader(&copied, original), Closer: original}
if err != nil || !isFederationRuntimeCommandAllowed(request.CommandType) {
return false
}
_, action := federationResourceCommandKind(request.CommandType)
return action == "delete"
default:
return false
}
}
func (h *Handler) federationConnect(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("Invalid method"))
@@ -982,8 +1026,6 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
return
}
h.wsServer.SendCommand(share.NodeID, "reload", nil, time.Second*5)
response.WriteJSON(w, response.OK(map[string]interface{}{
"tunnelId": tunnelID,
}))
@@ -1013,11 +1055,28 @@ func (h *Handler) federationRuntimeReservePort(w http.ResponseWriter, r *http.Re
return
}
peerRoleRuntimeMu.Lock()
defer peerRoleRuntimeMu.Unlock()
share, err = h.repo.GetPeerShare(share.ID)
if err != nil || share == nil || share.IsActive != 1 || (share.ExpiryTime > 0 && share.ExpiryTime <= time.Now().UnixMilli()) {
response.WriteJSON(w, response.Err(403, "Share is unavailable"))
return
}
if isPeerShareFlowExceeded(share) {
response.WriteJSON(w, response.Err(403, "Share traffic limit exceeded"))
return
}
existing, err := h.repo.GetPeerShareRuntimeByResourceKey(share.ID, req.ResourceKey)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if existing != nil && existing.ReleasePending != 0 {
response.WriteJSON(w, response.ErrDefault("Runtime release is pending"))
return
}
if existing != nil && existing.Status == 1 {
response.WriteJSON(w, response.OK(map[string]interface{}{
"reservationId": existing.ReservationID,
@@ -1026,10 +1085,6 @@ func (h *Handler) federationRuntimeReservePort(w http.ResponseWriter, r *http.Re
}))
return
}
if isPeerShareFlowExceeded(share) {
response.WriteJSON(w, response.Err(403, "Share traffic limit exceeded"))
return
}
allocatedPort, err := h.pickPeerSharePort(share, req.RequestedPort)
if err != nil {
@@ -1039,6 +1094,7 @@ func (h *Handler) federationRuntimeReservePort(w http.ResponseWriter, r *http.Re
now := time.Now().UnixMilli()
if existing != nil {
existing.ReservationID = randomToken(24)
existing.Protocol = defaultString(req.Protocol, "tls")
existing.Port = allocatedPort
existing.BindingID = ""
@@ -1116,6 +1172,14 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ
return
}
peerRoleRuntimeMu.Lock()
defer peerRoleRuntimeMu.Unlock()
share, err = h.repo.GetPeerShare(share.ID)
if err != nil || share == nil || share.IsActive != 1 || (share.ExpiryTime > 0 && share.ExpiryTime <= time.Now().UnixMilli()) {
response.WriteJSON(w, response.ErrDefault("Share is unavailable"))
return
}
var runtime *repo.PeerShareRuntime
if strings.TrimSpace(req.ReservationID) != "" {
runtime, err = h.repo.GetPeerShareRuntimeByReservationID(share.ID, strings.TrimSpace(req.ReservationID))
@@ -1131,108 +1195,37 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ
return
}
node, err := h.getNodeRecord(share.NodeID)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
protocol := defaultString(req.Protocol, runtime.Protocol)
strategy := defaultString(req.Strategy, "round")
chainName := defaultString(runtime.ChainName, federationRuntimeChainName(runtime.BindingID))
if chainName == "" {
chainName = federationRuntimeChainName(fmt.Sprintf("%d", runtime.ID))
}
serviceName := fmt.Sprintf("fed_svc_%d", runtime.ID)
if runtime.Applied == 1 && strings.TrimSpace(runtime.BindingID) != "" {
if req.Role == "middle" && len(req.Targets) > 0 {
chainData, buildErr := buildFederationMiddleChainConfig(chainName, runtime.ID, protocol, strategy, req.Targets, node.InterfaceName)
if buildErr != nil {
response.WriteJSON(w, response.ErrDefault(buildErr.Error()))
return
}
if _, err := h.sendNodeCommand(share.NodeID, "UpdateChains", updateChainPayload(chainName, chainData), false, false); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
targetBytes, _ := json.Marshal(req.Targets)
runtime.Role = req.Role
runtime.ChainName = chainName
runtime.Protocol = protocol
runtime.Strategy = strategy
runtime.Target = string(targetBytes)
runtime.Status = 1
runtime.UpdatedTime = time.Now().UnixMilli()
if err := h.repo.UpdatePeerShareRuntime(runtime); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
}
response.WriteJSON(w, response.OK(map[string]interface{}{
"bindingId": runtime.BindingID,
"allocatedPort": runtime.Port,
"reservationId": runtime.ReservationID,
}))
if runtime.ReleasePending != 0 {
response.WriteJSON(w, response.ErrDefault("Runtime release is pending"))
return
}
if isPeerShareFlowExceeded(share) {
response.WriteJSON(w, response.Err(403, "Share traffic limit exceeded"))
return
}
if share.PortRangeStart > 0 && share.PortRangeEnd > 0 && runtime.Port > 0 {
if runtime.Port < share.PortRangeStart || runtime.Port > share.PortRangeEnd {
response.WriteJSON(w, response.Err(403, fmt.Sprintf("port %d out of allowed range %d-%d", runtime.Port, share.PortRangeStart, share.PortRangeEnd)))
return
}
}
if req.Role == "middle" {
chainData, buildErr := buildFederationMiddleChainConfig(chainName, runtime.ID, protocol, strategy, req.Targets, node.InterfaceName)
if buildErr != nil {
response.WriteJSON(w, response.ErrDefault(buildErr.Error()))
return
}
if _, err := h.sendNodeCommand(share.NodeID, "AddChains", chainData, true, false); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
}
targetCount := len(req.Targets)
service := buildFederationServiceConfig(
serviceName,
fmt.Sprintf("%s:%d", node.TCPListenAddr, runtime.Port),
protocol,
req.Role,
chainName,
targetCount,
node.InterfaceName,
)
if _, err := h.sendNodeCommand(share.NodeID, "AddService", []map[string]interface{}{service}, true, false); err != nil {
if req.Role == "middle" {
_, _ = h.sendNodeCommand(share.NodeID, "DeleteChains", map[string]interface{}{"chain": chainName}, false, true)
}
response.WriteJSON(w, response.ErrDefault(err.Error()))
if runtime.Role != "" && runtime.Role != req.Role {
response.WriteJSON(w, response.ErrDefault("Runtime role cannot change without release"))
return
}
targetBytes, _ := json.Marshal(req.Targets)
runtime.BindingID = fmt.Sprintf("%d", runtime.ID)
if share.PortRangeStart > 0 && share.PortRangeEnd > 0 && (runtime.Port < share.PortRangeStart || runtime.Port > share.PortRangeEnd) {
response.WriteJSON(w, response.ErrDefault("Reserved port is outside share range"))
return
}
if strings.TrimSpace(runtime.BindingID) == "" {
runtime.BindingID = randomToken(24)
}
runtime.Role = req.Role
runtime.ServiceName = fmt.Sprintf("fed_svc_%d", runtime.ID)
runtime.ChainName = ""
if req.Role == "middle" {
runtime.ChainName = chainName
runtime.ChainName = federationRuntimeChainName(runtime.BindingID)
}
runtime.ServiceName = serviceName
runtime.Protocol = protocol
runtime.Strategy = strategy
runtime.Protocol = defaultString(req.Protocol, runtime.Protocol)
runtime.Strategy = defaultString(req.Strategy, "round")
targetBytes, _ := json.Marshal(req.Targets)
runtime.Target = string(targetBytes)
runtime.Applied = 1
runtime.Status = 1
runtime.UpdatedTime = time.Now().UnixMilli()
if err := h.repo.UpdatePeerShareRuntime(runtime); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
if err := h.applyPeerShareRoleRuntime(runtime); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
@@ -1282,17 +1275,8 @@ func (h *Handler) federationRuntimeReleaseRole(w http.ResponseWriter, r *http.Re
return
}
if runtime.Applied == 1 {
if strings.TrimSpace(runtime.ServiceName) != "" {
_, _ = h.sendNodeCommand(share.NodeID, "DeleteService", map[string]interface{}{"services": []string{runtime.ServiceName}}, false, true)
}
if strings.TrimSpace(runtime.Role) == "middle" && strings.TrimSpace(runtime.ChainName) != "" {
_, _ = h.sendNodeCommand(share.NodeID, "DeleteChains", map[string]interface{}{"chain": runtime.ChainName}, false, true)
}
}
if err := h.repo.MarkPeerShareRuntimeReleased(runtime.ID, time.Now().UnixMilli()); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
if err := h.releasePeerShareRuntime(runtime); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
@@ -1385,6 +1369,15 @@ func (h *Handler) federationRuntimeCommand(w http.ResponseWriter, r *http.Reques
return
}
h.peerResourceMu.Lock()
defer h.peerResourceMu.Unlock()
// Recheck after acquiring the mutation lock: an earlier authentication
// decision cannot authorize a recreation after quota/expiry cleanup.
share, err = h.repo.GetPeerShare(share.ID)
if err != nil || share == nil {
response.WriteJSON(w, response.ErrDefault("share ownership unavailable"))
return
}
if isFederationServiceCommand(cmd) {
if err := validateFederationCommandPorts(share, req.Data); err != nil {
response.WriteJSON(w, response.Err(403, err.Error()))
@@ -1392,17 +1385,35 @@ func (h *Handler) federationRuntimeCommand(w http.ResponseWriter, r *http.Reques
}
}
res, err := h.sendNodeCommand(share.NodeID, cmd, req.Data, false, false)
_, action := federationResourceCommandKind(cmd)
if action != "delete" && (share.IsActive != 1 || (share.ExpiryTime > 0 && share.ExpiryTime <= time.Now().UnixMilli()) || isPeerShareFlowExceeded(share)) {
response.WriteJSON(w, response.Err(403, "share is inactive, expired, or over quota"))
return
}
if strings.EqualFold(cmd, "tcpping") {
res, err := h.sendNodeCommand(share.NodeID, "TcpPing", req.Data, false, false)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
response.WriteJSON(w, response.OK(res))
return
}
items, err := h.preparePeerResourceCommand(share, cmd, req.Data)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
if strings.EqualFold(cmd, "addservice") || strings.EqualFold(cmd, "updateservice") {
h.bindPeerShareForwardRuntimeServices(share, req.Data)
} else if strings.EqualFold(cmd, "deleteservice") {
h.releasePeerShareForwardRuntimeServices(share, req.Data)
var result interface{} = map[string]interface{}{"success": true}
for _, item := range items {
res, err := h.applyPeerShareResource(item)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
result = res
}
response.WriteJSON(w, response.OK(res))
response.WriteJSON(w, response.OK(result))
}
type federationForwardServiceBinding struct {
@@ -1436,10 +1447,15 @@ func parseFederationForwardServiceBindings(data interface{}) []federationForward
bindings := make([]federationForwardServiceBinding, 0, len(serviceList))
for _, svcMap := range serviceList {
name := normalizeForwardRuntimeServiceName(asString(svcMap["name"]))
originalName := name
if shareID, original, ok := parsePeerShareServiceName(asString(svcMap["name"])); ok {
originalName = normalizeForwardRuntimeServiceName(original)
name = peerShareResourceName(shareID, "service", originalName)
}
if name == "" {
continue
}
if _, _, _, ok := parseFlowServiceIDs(name); !ok {
if _, _, _, ok := parseFlowServiceIDs(originalName); !ok {
continue
}
addr := strings.TrimSpace(asString(svcMap["addr"]))
@@ -1498,29 +1514,29 @@ func parseFederationForwardServiceNamesForRelease(data interface{}) []string {
return out
}
func (h *Handler) bindPeerShareForwardRuntimeServices(share *repo.PeerShare, data interface{}) {
func (h *Handler) bindPeerShareForwardRuntimeServices(share *repo.PeerShare, data interface{}) error {
if h == nil || h.repo == nil || share == nil {
return
return nil
}
bindings := parseFederationForwardServiceBindings(data)
if len(bindings) == 0 {
return
return nil
}
now := time.Now().UnixMilli()
for _, binding := range bindings {
runtime, err := h.repo.GetActiveForwardPeerShareRuntimeByPort(share.ID, binding.Port)
if err != nil {
continue
return err
}
if runtime == nil {
runtime, err = h.repo.GetActiveForwardPeerShareRuntimeByServiceName(share.ID, binding.Name)
if err != nil {
continue
return err
}
}
if runtime == nil {
_ = h.repo.CreatePeerShareRuntime(&repo.PeerShareRuntime{
if err := h.repo.CreatePeerShareRuntime(&repo.PeerShareRuntime{
ShareID: share.ID,
NodeID: share.NodeID,
ReservationID: randomToken(24),
@@ -1537,9 +1553,14 @@ func (h *Handler) bindPeerShareForwardRuntimeServices(share *repo.PeerShare, dat
Status: 1,
CreatedTime: now,
UpdatedTime: now,
})
}); err != nil {
return err
}
continue
}
if runtime.ReleasePending != 0 {
return fmt.Errorf("runtime release is pending")
}
if runtime.ServiceName == binding.Name && runtime.Applied == 1 && runtime.Port == binding.Port && runtime.Status == 1 {
continue
}
@@ -1554,8 +1575,11 @@ func (h *Handler) bindPeerShareForwardRuntimeServices(share *repo.PeerShare, dat
if strings.TrimSpace(runtime.Strategy) == "" {
runtime.Strategy = "fifo"
}
_ = h.repo.UpdatePeerShareRuntime(runtime)
if err := h.repo.UpdatePeerShareRuntime(runtime); err != nil {
return err
}
}
return nil
}
func (h *Handler) releasePeerShareForwardRuntimeServices(share *repo.PeerShare, data interface{}) {
@@ -1575,7 +1599,7 @@ func (h *Handler) releasePeerShareForwardRuntimeServices(share *repo.PeerShare,
func isFederationRuntimeCommandAllowed(commandType string) bool {
switch strings.ToLower(strings.TrimSpace(commandType)) {
case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice", "addchains", "deletechains", "addlimiters", "updatelimiters", "deletelimiters", "tcpping", "reload":
case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice", "addchains", "updatechains", "deletechains", "addlimiters", "updatelimiters", "deletelimiters", "addclimiters", "updateclimiters", "deleteclimiters", "tcpping":
return true
default:
return false
@@ -1893,27 +1917,24 @@ func (h *Handler) syncRemoteNodeStatuses(items []map[string]interface{}) {
}
}
func (h *Handler) cleanupPeerShareRuntimes(shareID int64) {
func (h *Handler) cleanupPeerShareRuntimes(shareID int64) error {
if h == nil || h.repo == nil || shareID <= 0 {
return
return nil
}
if err := h.releasePeerShareResources(shareID); err != nil {
return err
}
runtimes, err := h.repo.ListActivePeerShareRuntimesByShareID(shareID)
if err != nil || len(runtimes) == 0 {
return
if err != nil {
return err
}
now := time.Now().UnixMilli()
for _, runtime := range runtimes {
if h.wsServer != nil && runtime.Applied == 1 {
if strings.TrimSpace(runtime.ServiceName) != "" {
_, _ = h.sendNodeCommand(runtime.NodeID, "DeleteService", map[string]interface{}{"services": []string{runtime.ServiceName}}, false, true)
}
if strings.TrimSpace(runtime.Role) == "middle" && strings.TrimSpace(runtime.ChainName) != "" {
_, _ = h.sendNodeCommand(runtime.NodeID, "DeleteChains", map[string]interface{}{"chain": runtime.ChainName}, false, true)
}
var cleanupErr error
for i := range runtimes {
if err := h.releasePeerShareRuntime(&runtimes[i]); err != nil {
cleanupErr = errors.Join(cleanupErr, err)
}
_ = h.repo.MarkPeerShareRuntimeReleased(runtime.ID, now)
}
return cleanupErr
}
func (h *Handler) cleanupFederationTunnels(shareID int64) {