From 777db8767f8136608b5fb306d8b543e8e4b00363 Mon Sep 17 00:00:00 2001 From: sagitchu Date: Tue, 30 Jun 2026 17:07:50 +0800 Subject: [PATCH] fix nftables recovery and diagnostics --- .../internal/http/handler/control_plane.go | 109 ++++++++++++++++ go-backend/internal/http/handler/jobs.go | 38 +++++- .../http/handler/nftables_runtime_test.go | 122 +++++++++++++++++- 3 files changed, 267 insertions(+), 2 deletions(-) diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 4025b91..7b5ccb2 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -14,6 +14,7 @@ import ( "time" "go-backend/internal/http/client" + runtimenft "go-backend/internal/runtime/nftables" "go-backend/internal/store/model" "go-backend/internal/ws" ) @@ -772,6 +773,9 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe if ctx == nil { ctx = context.Background() } + if payload, handled, err := h.diagnoseNftablesForwardRuntime(forward); handled || err != nil { + return payload, err + } forwardName, workItems, err := h.prepareForwardDiagnosis(forward) if err != nil { return nil, err @@ -787,6 +791,111 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe return payload, nil } +func (h *Handler) diagnoseNftablesForwardRuntime(forward *forwardRecord) (map[string]interface{}, bool, error) { + if forward == nil { + return nil, false, errForwardNotFound + } + nftMode, entryNodeIDs, err := h.tunnelUsesNftables(forward.TunnelID) + if err != nil { + return nil, false, err + } + if !nftMode { + return nil, false, nil + } + if len(entryNodeIDs) == 0 { + return nil, true, errors.New("nftables 转发缺少入口节点") + } + targets, err := resolveDiagnosisTargets(forward.RemoteAddr) + if err != nil { + return nil, true, err + } + results, err := h.buildNftablesForwardDiagnosisResults(forward, entryNodeIDs[0], targets) + if err != nil { + return nil, true, err + } + payload := map[string]interface{}{ + "forwardName": forward.Name, + "timestamp": time.Now().UnixMilli(), + "results": results, + } + return payload, true, nil +} + +func (h *Handler) buildNftablesForwardDiagnosisResults(forward *forwardRecord, nodeID int64, targets []diagnosisTarget) ([]map[string]interface{}, error) { + if h == nil || h.repo == nil { + return nil, errors.New("handler not initialized") + } + node, err := h.getNodeRecord(nodeID) + if err != nil { + return nil, err + } + bindings, err := h.repo.ListNftRuleBindingsByNode(nodeID) + if err != nil { + return nil, err + } + var binding *model.NftRuleBinding + for i := range bindings { + if bindings[i].ForwardID == forward.ID { + binding = &bindings[i] + break + } + } + target := diagnosisTarget{} + if len(targets) > 0 { + target = targets[0] + } + + status := "missing" + message := "nftables 规则未下发" + success := false + inPort := 0 + protocols := "" + targetAddr := strings.TrimSpace(forward.RemoteAddr) + ruleHash := "" + if binding != nil { + status = strings.ToLower(strings.TrimSpace(binding.Status)) + inPort = binding.InPort + protocols = strings.TrimSpace(binding.Protocols) + targetAddr = strings.TrimSpace(binding.TargetAddr) + ruleHash = strings.TrimSpace(binding.RuleHash) + if status == "" { + status = "pending" + } + if status == runtimenft.StatusApplied { + success = true + message = "nftables 规则已下发" + } else if strings.TrimSpace(binding.LastError) != "" { + message = binding.LastError + } else { + message = "nftables 规则未完成下发" + } + } + + packetLoss := 100 + if success { + packetLoss = 0 + } + result := map[string]interface{}{ + "success": success, + "nodeName": node.Name, + "nodeId": strconv.FormatInt(nodeID, 10), + "targetIp": target.IP, + "targetPort": target.Port, + "description": fmt.Sprintf("nftables规则(%s)->目标(%s)", node.Name, defaultString(target.Address, targetAddr)), + "averageTime": 0, + "packetLoss": packetLoss, + "message": message, + "fromChainType": 1, + "forwardMode": "nftables", + "nftRuleStatus": status, + "nftRuleHash": ruleHash, + "inPort": inPort, + "protocols": protocols, + "targetAddr": targetAddr, + } + return []map[string]interface{}{result}, nil +} + func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []diagnosisWorkItem, error) { if forward == nil { return "", nil, errForwardNotFound diff --git a/go-backend/internal/http/handler/jobs.go b/go-backend/internal/http/handler/jobs.go index c89f9e8..9fbddc0 100644 --- a/go-backend/internal/http/handler/jobs.go +++ b/go-backend/internal/http/handler/jobs.go @@ -2,11 +2,14 @@ package handler import ( "context" + "log" "time" "go-backend/internal/license" ) +var nftablesTrafficCollectInterval = 30 * time.Second + func (h *Handler) StartBackgroundJobs() { if h == nil || h.repo == nil { return @@ -131,7 +134,19 @@ func (h *Handler) runTunnelQualityProber(ctx context.Context) { func (h *Handler) runNftablesTrafficCollectLoop(ctx context.Context) { defer h.jobsWG.Done() - ticker := time.NewTicker(time.Minute) + h.runNftablesStartupReconcile(ctx) + select { + case <-ctx.Done(): + return + default: + h.runNftablesTrafficCollectJob(time.Now()) + } + + interval := nftablesTrafficCollectInterval + if interval <= 0 { + interval = 30 * time.Second + } + ticker := time.NewTicker(interval) defer ticker.Stop() for { @@ -144,6 +159,27 @@ func (h *Handler) runNftablesTrafficCollectLoop(ctx context.Context) { } } +func (h *Handler) runNftablesStartupReconcile(ctx context.Context) { + if h == nil || h.repo == nil { + return + } + nodes, err := h.repo.ListNftablesNodesForCollection() + if err != nil { + log.Printf("nftables startup reconcile failed op=list_nodes err=%v", err) + return + } + for _, node := range nodes { + select { + case <-ctx.Done(): + return + default: + } + if err := h.syncNftablesNode(node.NodeID); err != nil { + log.Printf("nftables startup reconcile failed node_id=%d err=%v", node.NodeID, err) + } + } +} + func (h *Handler) runHourlyStatsLoop(ctx context.Context) { defer h.jobsWG.Done() diff --git a/go-backend/internal/http/handler/nftables_runtime_test.go b/go-backend/internal/http/handler/nftables_runtime_test.go index e2bd313..b6b64a0 100644 --- a/go-backend/internal/http/handler/nftables_runtime_test.go +++ b/go-backend/internal/http/handler/nftables_runtime_test.go @@ -10,6 +10,7 @@ import ( "net/http/httptest" "path/filepath" "strings" + "sync" "testing" "time" @@ -20,6 +21,7 @@ import ( ) type fakeNftablesManager struct { + mu sync.Mutex testErr error reconcileErr error reconcileHit int @@ -33,11 +35,15 @@ type fakeNftablesManager struct { } func (f *fakeNftablesManager) Test(_ context.Context, cfg runtimenft.SSHConfig) error { + f.mu.Lock() + defer f.mu.Unlock() f.lastConfig = cfg return f.testErr } func (f *fakeNftablesManager) Reconcile(_ context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error) { + f.mu.Lock() + defer f.mu.Unlock() f.reconcileHit++ f.lastConfig = cfg f.lastPlan = plan @@ -47,16 +53,20 @@ func (f *fakeNftablesManager) Reconcile(_ context.Context, cfg runtimenft.SSHCon return runtimenft.ApplyResult{ NodeID: plan.NodeID, Script: "table inet flvx {}", - Hashes: map[int64]string{plan.NodeID: "hash"}, + Hashes: runtimenft.PlanHashes(plan), }, nil } func (f *fakeNftablesManager) Clear(context.Context, runtimenft.SSHConfig) error { + f.mu.Lock() + defer f.mu.Unlock() f.clearHit++ return f.clearErr } func (f *fakeNftablesManager) CollectCounters(_ context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error) { + f.mu.Lock() + defer f.mu.Unlock() f.collectHit++ f.lastConfig = cfg if f.collectErr != nil { @@ -65,6 +75,18 @@ func (f *fakeNftablesManager) CollectCounters(_ context.Context, cfg runtimenft. return f.counterSamples, nil } +func (f *fakeNftablesManager) reconcileCount() int { + f.mu.Lock() + defer f.mu.Unlock() + return f.reconcileHit +} + +func (f *fakeNftablesManager) collectCount() int { + f.mu.Lock() + defer f.mu.Unlock() + return f.collectHit +} + type nftablesTestFixture struct { handler *Handler nodeID int64 @@ -158,6 +180,49 @@ func TestNodeNftablesReconcileEndpointPersistsBindings(t *testing.T) { } } +func TestStartBackgroundJobsReconcilesNftablesRulesAtStartup(t *testing.T) { + fixture := setupNftablesHandler(t) + h := fixture.handler + seedNftablesSSHConfig(t, h, fixture.nodeID) + tunnelID := seedTunnelForNftables(t, h, "nft-startup-tunnel", fixture.nodeID) + seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080") + manager := &fakeNftablesManager{} + h.nftablesManager = manager + + h.StartBackgroundJobs() + t.Cleanup(h.StopBackgroundJobs) + + waitForCondition(t, time.Second, func() bool { + return manager.reconcileCount() > 0 + }, "nftables startup reconcile") +} + +func TestStartBackgroundJobsCollectsNftablesTrafficImmediatelyAndUsesFastInterval(t *testing.T) { + fixture := setupNftablesCollectionFixture(t) + h := fixture.handler + manager := &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{ + {ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10}, + }} + h.nftablesManager = manager + + oldInterval := nftablesTrafficCollectInterval + nftablesTrafficCollectInterval = 20 * time.Millisecond + t.Cleanup(func() { nftablesTrafficCollectInterval = oldInterval }) + + h.StartBackgroundJobs() + t.Cleanup(h.StopBackgroundJobs) + + waitForCondition(t, time.Second, func() bool { + return manager.collectCount() >= 2 + }, "immediate and repeated nftables traffic collection") +} + +func TestNftablesTrafficCollectIntervalDefaultsToThirtySeconds(t *testing.T) { + if nftablesTrafficCollectInterval != 30*time.Second { + t.Fatalf("expected default nftables traffic collection interval 30s, got %s", nftablesTrafficCollectInterval) + } +} + func TestNodeNftablesClearEndpointClearsBindings(t *testing.T) { fixture := setupNftablesHandler(t) seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID) @@ -421,6 +486,61 @@ func TestTunnelBatchRedeployUsesNftablesReconcile(t *testing.T) { } } +func TestDiagnoseForwardRuntimeReturnsNftablesRuleStatus(t *testing.T) { + fixture := setupNftablesHandler(t) + h := fixture.handler + tunnelID := seedTunnelForNftables(t, h, "nft-diagnose-tunnel", fixture.nodeID) + forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080") + now := time.Now().UnixMilli() + if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{ + ForwardID: forward.ID, + NodeID: fixture.nodeID, + InPort: 20000, + Protocols: "tcp,udp", + TargetAddr: "203.0.113.9:8080", + RuleHash: "hash-a", + Status: runtimenft.StatusApplied, + }, now); err != nil { + t.Fatalf("seed nft binding: %v", err) + } + + payload, err := h.diagnoseForwardRuntime(context.Background(), &forwardRecord{ + ID: forward.ID, + Name: forward.Name, + TunnelID: tunnelID, + RemoteAddr: "203.0.113.9:8080", + }) + if err != nil { + t.Fatalf("diagnose forward: %v", err) + } + results, ok := payload["results"].([]map[string]interface{}) + if !ok || len(results) != 1 { + t.Fatalf("expected one nftables diagnosis result, got %#v", payload["results"]) + } + result := results[0] + if result["forwardMode"] != "nftables" || result["nftRuleStatus"] != runtimenft.StatusApplied { + t.Fatalf("expected nftables applied result, got %#v", result) + } + if result["success"] != true { + t.Fatalf("expected nftables diagnosis success, got %#v", result) + } + if !strings.Contains(asString(result["message"]), "已下发") { + t.Fatalf("expected applied message, got %#v", result["message"]) + } +} + +func waitForCondition(t *testing.T, timeout time.Duration, condition func() bool, description string) { + t.Helper() + deadline := time.Now().Add(timeout) + for time.Now().Before(deadline) { + if condition() { + return + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("timed out waiting for %s", description) +} + func setupNftablesHandler(t *testing.T) nftablesTestFixture { t.Helper()