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
+200 -197
View File
@@ -8,8 +8,6 @@ import (
"strconv"
"strings"
"time"
"go-backend/internal/store/model"
)
const bytesPerGB int64 = 1024 * 1024 * 1024
@@ -48,38 +46,22 @@ type gostConfigSnapshot struct {
}
type namedConfigItem struct {
Name string `json:"name"`
Name string `json:"name"`
Limiter string `json:"limiter,omitempty"`
Handler *struct {
Chain string `json:"chain"`
} `json:"handler,omitempty"`
}
func (h *Handler) processFlowItem(nodeID int64, item flowItem) {
serviceName := strings.TrimSpace(item.N)
if serviceName == "" || serviceName == "web_api" {
if h == nil || h.repo == nil || nodeID <= 0 {
return
}
forwardID, userID, userTunnelID, ok := parseFlowServiceIDs(serviceName)
if ok {
if h.forwardExists(forwardID) {
inFlow, outFlow := h.scaleFlowByTunnel(forwardID, item.D, item.U)
_ = h.repo.AddFlow(forwardID, userID, userTunnelID, inFlow, outFlow)
if quota, quotaErr := h.repo.AddUserQuotaUsage(userID, inFlow+outFlow, time.Now()); quotaErr == nil {
h.enforceUserQuotaIfNeeded(userID, quota)
}
if userTunnelID > 0 {
h.enforceFlowPolicies(userID, userTunnelID)
}
} else if nodeID > 0 {
h.sendDeleteOrphanedForwardService(nodeID, serviceName)
}
h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item)
return
metas, err := h.repo.GetFlowUploadForwardMetas(collectFlowUploadForwardIDs([]flowItem{item}))
if err != nil {
metas = nil
}
runtimeID, ok := parsePeerShareRuntimeServiceID(serviceName)
if !ok {
return
}
h.processPeerShareFlow(runtimeID, item)
h.applyFlowUploadBatch(nodeID, h.buildNodeFlowUploadBatch(nodeID, []flowItem{item}, metas), time.Now())
}
func parseFlowServiceIDs(serviceName string) (int64, int64, int64, bool) {
@@ -154,73 +136,49 @@ func parsePeerShareIDFromFederationTunnelName(tunnelName string) (int64, bool) {
return shareID, true
}
func (h *Handler) processPeerShareFlow(runtimeID int64, item flowItem) {
if h == nil || h.repo == nil || runtimeID <= 0 {
func (h *Handler) processPeerShareFlow(nodeID, runtimeID int64, item flowItem) {
if h == nil || h.repo == nil || nodeID <= 0 || runtimeID <= 0 {
return
}
runtime, err := h.repo.GetPeerShareRuntimeByID(runtimeID)
if err != nil || runtime == nil || runtime.ShareID <= 0 || runtime.Status != 1 {
if err != nil || runtime == nil || runtime.NodeID != nodeID || runtime.Status != 1 {
return
}
delta := item.D + item.U
if delta <= 0 {
return
}
_ = h.repo.AddPeerShareCurrentFlow(runtime.ShareID, delta)
share, err := h.repo.GetPeerShare(runtime.ShareID)
if err != nil || share == nil {
return
}
if !isPeerShareFlowExceeded(share) {
return
}
h.enforcePeerShareFlowLimit(share.ID)
h.addPeerShareFlow(nodeID, runtime.ShareID, item.D+item.U)
}
func (h *Handler) processPeerShareFlowFromForward(forwardID int64, nodeID int64, serviceName string, item flowItem) {
if h == nil || h.repo == nil || forwardID <= 0 {
if h == nil || h.repo == nil || forwardID <= 0 || nodeID <= 0 {
return
}
delta := item.D + item.U
if delta <= 0 {
// Prefer the reporting node's explicit shared ownership over a coincidentally
// equal local forward ID. Never fall back to a service on another node.
runtimes, err := h.repo.ListActiveForwardPeerShareRuntimesByNode(nodeID)
if err != nil {
return
}
for _, runtime := range runtimes {
if normalizeForwardRuntimeServiceName(runtime.ServiceName) == normalizeForwardRuntimeServiceName(serviceName) {
h.processPeerShareFlowByServiceName(nodeID, serviceName, item)
return
}
}
forward, err := h.getForwardRecord(forwardID)
if err != nil || forward == nil {
// Forward not found in local database - might be a federation port-forward
// Try to find by service name in peer_share_runtime
h.processPeerShareFlowByServiceName(nodeID, serviceName, item)
return
}
_, userID, _, ok := parseFlowServiceIDs(serviceName)
if !ok || userID != forward.UserID {
return
}
tunnelName, err := h.repo.GetTunnelName(forward.TunnelID)
if err != nil {
h.processPeerShareFlowByServiceName(nodeID, serviceName, item)
return
}
shareID, ok := parsePeerShareIDFromFederationTunnelName(tunnelName)
if !ok {
h.processPeerShareFlowByServiceName(nodeID, serviceName, item)
return
if ok {
h.addPeerShareFlow(nodeID, shareID, item.D+item.U)
}
if err := h.repo.AddPeerShareCurrentFlow(shareID, delta); err != nil {
h.processPeerShareFlowByServiceName(nodeID, serviceName, item)
return
}
share, err := h.repo.GetPeerShare(shareID)
if err != nil || share == nil {
return
}
if !isPeerShareFlowExceeded(share) {
return
}
h.enforcePeerShareFlowLimit(share.ID)
}
func normalizeForwardRuntimeServiceName(serviceName string) string {
@@ -235,63 +193,26 @@ func normalizeForwardRuntimeServiceName(serviceName string) string {
}
func (h *Handler) processPeerShareFlowByServiceName(nodeID int64, serviceName string, item flowItem) {
if h == nil || h.repo == nil || strings.TrimSpace(serviceName) == "" {
if h == nil || h.repo == nil || nodeID <= 0 || strings.TrimSpace(serviceName) == "" {
return
}
delta := item.D + item.U
if delta <= 0 {
runtimes, err := h.repo.ListActiveForwardPeerShareRuntimesByNode(nodeID)
if err != nil {
return
}
normalized := normalizeForwardRuntimeServiceName(serviceName)
var runtimes []model.PeerShareRuntime
var err error
// Try node-scoped query first if nodeID is valid
if nodeID > 0 {
runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByNodeAndServiceName(nodeID, normalized)
if err != nil {
var shareID int64
for _, runtime := range runtimes {
if normalizeForwardRuntimeServiceName(runtime.ServiceName) != normalizeForwardRuntimeServiceName(serviceName) {
continue
}
if shareID != 0 {
log.Printf("ambiguous peer share runtime service=%s node_id=%d", serviceName, nodeID)
return
}
if len(runtimes) == 0 && normalized != serviceName {
runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByNodeAndServiceName(nodeID, serviceName)
if err != nil {
return
}
}
shareID = runtime.ShareID
}
// Fallback to global query if node-scoped query returned nothing or nodeID is invalid
if len(runtimes) == 0 {
runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByServiceName(normalized)
if err != nil {
return
}
if len(runtimes) == 0 && normalized != serviceName {
runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByServiceName(serviceName)
if err != nil {
return
}
}
}
if len(runtimes) != 1 {
if len(runtimes) > 1 {
log.Printf("WARN: ambiguous peer share runtime match for service=%s nodeID=%d count=%d", serviceName, nodeID, len(runtimes))
}
return
}
runtime := runtimes[0]
_ = h.repo.AddPeerShareCurrentFlow(runtime.ShareID, delta)
matchedShare, err := h.repo.GetPeerShare(runtime.ShareID)
if err != nil || matchedShare == nil {
return
}
if isPeerShareFlowExceeded(matchedShare) {
h.enforcePeerShareFlowLimit(matchedShare.ID)
if shareID > 0 {
h.addPeerShareFlow(nodeID, shareID, item.D+item.U)
}
}
@@ -299,22 +220,8 @@ func (h *Handler) enforcePeerShareFlowLimit(shareID int64) {
if h == nil || h.repo == nil || shareID <= 0 {
return
}
runtimes, err := h.repo.ListActivePeerShareRuntimesByShareID(shareID)
if err != nil || len(runtimes) == 0 {
return
}
now := time.Now().UnixMilli()
for _, runtime := range runtimes {
if h.wsServer != nil && runtime.Applied == 1 {
if strings.TrimSpace(runtime.ServiceName) != "" {
_, _ = h.sendNodeCommand(runtime.NodeID, "DeleteService", map[string]interface{}{"services": []string{runtime.ServiceName}}, false, true)
}
if strings.TrimSpace(runtime.Role) == "middle" && strings.TrimSpace(runtime.ChainName) != "" {
_, _ = h.sendNodeCommand(runtime.NodeID, "DeleteChains", map[string]interface{}{"chain": runtime.ChainName}, false, true)
}
}
_ = h.repo.MarkPeerShareRuntimeReleased(runtime.ID, now)
if err := h.cleanupPeerShareRuntimes(shareID); err != nil {
log.Printf("peer share quota cleanup pending share_id=%d err=%v", shareID, err)
}
}
@@ -531,43 +438,102 @@ func (h *Handler) cleanNodeConfigs(nodeID int64, rawConfig string) {
return
}
h.cleanOrphanedServices(nodeID, snapshot.Services)
h.cleanOrphanedChains(nodeID, snapshot.Chains)
h.cleanOrphanedLimiters(nodeID, snapshot.Limiters)
}
func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem) {
runtimeServiceNames, err := h.repo.ListActiveForwardPeerShareRuntimeServiceNamesByNode(nodeID)
protection, err := h.loadForwardServiceProtection(nodeID)
if err != nil {
return
}
minUpdatedTime := time.Now().Add(-10 * time.Minute).UnixMilli()
hasUnboundForwardPeerRuntime, err := h.repo.HasRecentUnboundForwardPeerShareRuntimeOnNode(nodeID, minUpdatedTime)
if err != nil {
hasUnboundForwardPeerRuntime = false
h.cleanOrphanedServicesWithProtection(nodeID, snapshot.Services, protection)
// Dependencies are sent before services. A pending shared reservation may
// therefore have chains/limiters that are not referenced in this snapshot yet.
if protection.unbound {
return
}
runtimeServiceSet := make(map[string]struct{}, len(runtimeServiceNames))
for _, serviceName := range runtimeServiceNames {
serviceName = strings.TrimSpace(serviceName)
chainsInUse := make(map[string]struct{})
limitersInUse := make(map[string]struct{})
for _, service := range snapshot.Services {
if service.Handler != nil {
if chain := strings.TrimSpace(service.Handler.Chain); chain != "" {
chainsInUse[chain] = struct{}{}
}
}
for _, limiter := range strings.Split(service.Limiter, ",") {
if limiter = strings.TrimSpace(limiter); limiter != "" {
limitersInUse[limiter] = struct{}{}
}
}
}
// Keep dependencies referenced by the reported services, even when their
// IDs belong to a different panel. Orphan dependencies can be collected on
// the next report after their services have actually disappeared.
h.cleanOrphanedChains(nodeID, snapshot.Chains, chainsInUse)
h.cleanOrphanedLimiters(nodeID, snapshot.Limiters, limitersInUse)
}
type forwardServiceProtection struct {
sharedNames map[string]struct{}
unbound bool
}
// Shared forward IDs belong to another panel and need not exist in our forward
// table. Use the same node-scoped ownership check for config and flow reports.
func (h *Handler) loadForwardServiceProtection(nodeID int64) (forwardServiceProtection, error) {
protection := forwardServiceProtection{sharedNames: make(map[string]struct{})}
// Read names and pending bindings in one snapshot, so a concurrent bind
// cannot fall between two queries and disappear from both protections.
runtimes, err := h.repo.ListActiveForwardPeerShareRuntimesByNode(nodeID)
if err != nil {
return protection, err
}
minUpdatedTime := time.Now().Add(-10 * time.Minute).UnixMilli()
for _, runtime := range runtimes {
serviceName := normalizeForwardRuntimeServiceName(runtime.ServiceName)
if serviceName == "" {
if runtime.Applied == 0 && runtime.UpdatedTime >= minUpdatedTime {
protection.unbound = true
}
continue
}
runtimeServiceSet[serviceName] = struct{}{}
protection.sharedNames[serviceName] = struct{}{}
}
resources, err := h.repo.ListPeerShareResourcesByNode(nodeID)
if err != nil {
return protection, err
}
for _, resource := range resources {
if base := normalizeForwardRuntimeServiceName(resource.LegacyServiceBase); base != "" {
protection.sharedNames[base] = struct{}{}
}
}
return protection, nil
}
func (p forwardServiceProtection) preserves(serviceName string) bool {
_, shared := p.sharedNames[normalizeForwardRuntimeServiceName(serviceName)]
return shared || p.unbound
}
func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem) {
if h == nil || h.repo == nil || nodeID <= 0 {
return
}
protection, err := h.loadForwardServiceProtection(nodeID)
if err != nil {
// A failed ownership lookup must never authorize deletion.
return
}
h.cleanOrphanedServicesWithProtection(nodeID, services, protection)
}
func (h *Handler) cleanOrphanedServicesWithProtection(nodeID int64, services []namedConfigItem, protection forwardServiceProtection) {
for _, item := range services {
name := strings.TrimSpace(item.Name)
if name == "" || name == "web_api" {
continue
}
if strings.HasPrefix(name, "fed_svc_") {
if strings.HasPrefix(name, "fed_svc_") || strings.HasPrefix(name, "peer-share-") {
continue
}
normalizedName := normalizeForwardRuntimeServiceName(name)
if _, ok := runtimeServiceSet[normalizedName]; ok {
continue
}
if _, ok := runtimeServiceSet[name]; ok {
if _, ok := protection.sharedNames[normalizeForwardRuntimeServiceName(name)]; ok {
continue
}
@@ -580,15 +546,9 @@ func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem
continue
}
if len(parts) >= 3 {
forwardID, err := strconv.ParseInt(parts[0], 10, 64)
if err == nil && forwardID > 0 && hasUnboundForwardPeerRuntime {
continue
}
if err == nil && forwardID > 0 && !h.forwardExists(forwardID) {
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name, parts[0] + "_" + parts[1] + "_" + parts[2], parts[0] + "_" + parts[1] + "_" + parts[2] + "_tcp", parts[0] + "_" + parts[1] + "_" + parts[2] + "_udp"}}, false, true)
continue
}
if _, _, _, ok := parseFlowServiceIDs(name); ok {
h.deleteOrphanedForwardService(nodeID, name, protection)
continue
}
suffix := parts[len(parts)-1]
@@ -607,23 +567,17 @@ func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem
}
continue
}
forwardID, err := strconv.ParseInt(parts[0], 10, 64)
if err == nil && forwardID > 0 && hasUnboundForwardPeerRuntime {
continue
}
if err != nil || forwardID <= 0 || h.forwardExists(forwardID) {
continue
}
base := strings.TrimSuffix(name, "_tcp")
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{base + "_tcp", base + "_udp"}}, false, true)
}
}
}
func (h *Handler) cleanOrphanedChains(nodeID int64, chains []namedConfigItem) {
func (h *Handler) cleanOrphanedChains(nodeID int64, chains []namedConfigItem, inUse map[string]struct{}) {
for _, item := range chains {
name := strings.TrimSpace(item.Name)
if name == "" {
if name == "" || strings.HasPrefix(name, "fed_chain_") || strings.HasPrefix(name, "peer-share-") {
continue
}
if _, ok := inUse[name]; ok {
continue
}
@@ -632,17 +586,28 @@ func (h *Handler) cleanOrphanedChains(nodeID int64, chains []namedConfigItem) {
continue
}
tunnelID, err := strconv.ParseInt(name[idx+1:], 10, 64)
if err != nil || tunnelID <= 0 || h.tunnelExists(tunnelID) {
if err != nil || tunnelID <= 0 {
continue
}
exists, err := h.repo.TunnelExists(tunnelID)
if err != nil || exists {
continue
}
_, _ = h.sendNodeCommand(nodeID, "DeleteChains", map[string]interface{}{"chain": name}, false, true)
}
}
func (h *Handler) cleanOrphanedLimiters(nodeID int64, limiters []namedConfigItem) {
func (h *Handler) cleanOrphanedLimiters(nodeID int64, limiters []namedConfigItem, inUse map[string]struct{}) {
for _, item := range limiters {
name := strings.TrimSpace(item.Name)
if name == "" || h.speedLimiterExists(name) {
if name == "" || strings.HasPrefix(name, "peer-share-") {
continue
}
if _, ok := inUse[name]; ok {
continue
}
exists, err := h.lookupSpeedLimiter(name)
if err != nil || exists {
continue
}
_, _ = h.sendNodeCommand(nodeID, "DeleteLimiters", map[string]interface{}{"limiter": name}, false, true)
@@ -660,40 +625,78 @@ func (h *Handler) forwardExists(forwardID int64) bool {
}
func (h *Handler) sendDeleteOrphanedForwardService(nodeID int64, serviceName string) {
h.sendDeleteOrphanedForwardServices(nodeID, []string{serviceName})
}
func (h *Handler) sendDeleteOrphanedForwardServices(nodeID int64, serviceNames []string) {
if h == nil || h.repo == nil || nodeID <= 0 || len(serviceNames) == 0 {
return
}
protection, err := h.loadForwardServiceProtection(nodeID)
if err != nil {
return
}
seen := make(map[string]struct{}, len(serviceNames))
for _, serviceName := range serviceNames {
serviceName = normalizeForwardRuntimeServiceName(serviceName)
if _, ok := seen[serviceName]; ok {
continue
}
seen[serviceName] = struct{}{}
h.deleteOrphanedForwardService(nodeID, serviceName, protection)
}
}
func (h *Handler) deleteOrphanedForwardService(nodeID int64, serviceName string, protection forwardServiceProtection) {
forwardID, _, _, ok := parseFlowServiceIDs(serviceName)
if !ok || protection.preserves(serviceName) {
return
}
parts := strings.Split(serviceName, "_")
if len(parts) < 3 {
return
}
forwardID, err := strconv.ParseInt(parts[0], 10, 64)
if err != nil || forwardID <= 0 {
return
}
base := parts[0] + "_" + parts[1] + "_" + parts[2]
// Parsing accepts legacy suffixes, while deletion targets the entire base
// family. Verify the actual targets cannot include a protected share.
if protection.preserves(base) {
return
}
// Batch metadata can be missing after a read failure, or stale by the time
// cleanup runs. Confirm absence before issuing a destructive command.
exists, err := h.repo.ForwardExists(forwardID)
if err != nil || exists {
return
}
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{
"services": []string{base + "_tcp", base + "_udp"},
"services": buildForwardServiceDeleteNames([]string{base}),
}, false, true)
}
func (h *Handler) speedLimiterExists(name string) bool {
exists, _ := h.lookupSpeedLimiter(name)
return exists
}
func (h *Handler) lookupSpeedLimiter(name string) (bool, error) {
name = strings.TrimSpace(name)
if name == "" {
return false
return false, nil
}
const forwardRulePrefix = "rule_traffic_limit_"
if strings.HasPrefix(name, forwardRulePrefix) {
forwardID, err := strconv.ParseInt(strings.TrimPrefix(name, forwardRulePrefix), 10, 64)
if err != nil || forwardID <= 0 {
return false
return false, nil
}
forward, err := h.getForwardRecord(forwardID)
return err == nil && forward != nil && forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0
if errors.Is(err, errForwardNotFound) {
return false, nil
}
return forward != nil && forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0, err
}
id, err := strconv.ParseInt(name, 10, 64)
if err != nil || id <= 0 {
return false
return false, nil
}
ok, _ := h.repo.SpeedLimitExists(id)
return ok
return h.repo.SpeedLimitExists(id)
}