mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-04 17:16:37 +08:00
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:
@@ -14,6 +14,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"go-backend/internal/http/client"
|
"go-backend/internal/http/client"
|
||||||
|
runtimenft "go-backend/internal/runtime/nftables"
|
||||||
"go-backend/internal/store/model"
|
"go-backend/internal/store/model"
|
||||||
"go-backend/internal/ws"
|
"go-backend/internal/ws"
|
||||||
)
|
)
|
||||||
@@ -772,6 +773,9 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe
|
|||||||
if ctx == nil {
|
if ctx == nil {
|
||||||
ctx = context.Background()
|
ctx = context.Background()
|
||||||
}
|
}
|
||||||
|
if payload, handled, err := h.diagnoseNftablesForwardRuntime(forward); handled || err != nil {
|
||||||
|
return payload, err
|
||||||
|
}
|
||||||
forwardName, workItems, err := h.prepareForwardDiagnosis(forward)
|
forwardName, workItems, err := h.prepareForwardDiagnosis(forward)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -787,6 +791,111 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe
|
|||||||
return payload, nil
|
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) {
|
func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []diagnosisWorkItem, error) {
|
||||||
if forward == nil {
|
if forward == nil {
|
||||||
return "", nil, errForwardNotFound
|
return "", nil, errForwardNotFound
|
||||||
|
|||||||
@@ -2,11 +2,14 @@ package handler
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"log"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"go-backend/internal/license"
|
"go-backend/internal/license"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var nftablesTrafficCollectInterval = 30 * time.Second
|
||||||
|
|
||||||
func (h *Handler) StartBackgroundJobs() {
|
func (h *Handler) StartBackgroundJobs() {
|
||||||
if h == nil || h.repo == nil {
|
if h == nil || h.repo == nil {
|
||||||
return
|
return
|
||||||
@@ -131,7 +134,19 @@ func (h *Handler) runTunnelQualityProber(ctx context.Context) {
|
|||||||
|
|
||||||
func (h *Handler) runNftablesTrafficCollectLoop(ctx context.Context) {
|
func (h *Handler) runNftablesTrafficCollectLoop(ctx context.Context) {
|
||||||
defer h.jobsWG.Done()
|
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()
|
defer ticker.Stop()
|
||||||
|
|
||||||
for {
|
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) {
|
func (h *Handler) runHourlyStatsLoop(ctx context.Context) {
|
||||||
defer h.jobsWG.Done()
|
defer h.jobsWG.Done()
|
||||||
|
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import (
|
|||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -20,6 +21,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type fakeNftablesManager struct {
|
type fakeNftablesManager struct {
|
||||||
|
mu sync.Mutex
|
||||||
testErr error
|
testErr error
|
||||||
reconcileErr error
|
reconcileErr error
|
||||||
reconcileHit int
|
reconcileHit int
|
||||||
@@ -33,11 +35,15 @@ type fakeNftablesManager struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *fakeNftablesManager) Test(_ context.Context, cfg runtimenft.SSHConfig) error {
|
func (f *fakeNftablesManager) Test(_ context.Context, cfg runtimenft.SSHConfig) error {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
f.lastConfig = cfg
|
f.lastConfig = cfg
|
||||||
return f.testErr
|
return f.testErr
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *fakeNftablesManager) Reconcile(_ context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error) {
|
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.reconcileHit++
|
||||||
f.lastConfig = cfg
|
f.lastConfig = cfg
|
||||||
f.lastPlan = plan
|
f.lastPlan = plan
|
||||||
@@ -47,16 +53,20 @@ func (f *fakeNftablesManager) Reconcile(_ context.Context, cfg runtimenft.SSHCon
|
|||||||
return runtimenft.ApplyResult{
|
return runtimenft.ApplyResult{
|
||||||
NodeID: plan.NodeID,
|
NodeID: plan.NodeID,
|
||||||
Script: "table inet flvx {}",
|
Script: "table inet flvx {}",
|
||||||
Hashes: map[int64]string{plan.NodeID: "hash"},
|
Hashes: runtimenft.PlanHashes(plan),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *fakeNftablesManager) Clear(context.Context, runtimenft.SSHConfig) error {
|
func (f *fakeNftablesManager) Clear(context.Context, runtimenft.SSHConfig) error {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
f.clearHit++
|
f.clearHit++
|
||||||
return f.clearErr
|
return f.clearErr
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *fakeNftablesManager) CollectCounters(_ context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error) {
|
func (f *fakeNftablesManager) CollectCounters(_ context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error) {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
f.collectHit++
|
f.collectHit++
|
||||||
f.lastConfig = cfg
|
f.lastConfig = cfg
|
||||||
if f.collectErr != nil {
|
if f.collectErr != nil {
|
||||||
@@ -65,6 +75,18 @@ func (f *fakeNftablesManager) CollectCounters(_ context.Context, cfg runtimenft.
|
|||||||
return f.counterSamples, nil
|
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 {
|
type nftablesTestFixture struct {
|
||||||
handler *Handler
|
handler *Handler
|
||||||
nodeID int64
|
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) {
|
func TestNodeNftablesClearEndpointClearsBindings(t *testing.T) {
|
||||||
fixture := setupNftablesHandler(t)
|
fixture := setupNftablesHandler(t)
|
||||||
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
|
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 {
|
func setupNftablesHandler(t *testing.T) nftablesTestFixture {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user