Fix nftables recovery and diagnostics (#524)

## Summary
- Reconcile nftables nodes when background jobs start so rules are
restored after server reboot.
- Collect nftables traffic immediately at startup and every 30 seconds
by default.
- Return nftables rule binding status in forward diagnostics and cover
it with regression tests.

## Test Plan
- `cd go-backend && go test ./...`
- `cd go-backend && make build`
This commit is contained in:
sagit
2026-06-30 17:11:10 +08:00
committed by GitHub
3 changed files with 267 additions and 2 deletions
@@ -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
+37 -1
View File
@@ -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()
@@ -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()