diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index eea7d7c..3fbc5d3 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -2412,23 +2412,30 @@ func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if err := h.controlForwardServices(forward, "DeleteService", true); err != nil { - response.WriteJSON(w, response.ErrDefault(err.Error())) - return - } + var nftNodeID int64 if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil { response.WriteJSON(w, response.Err(-2, modeErr.Error())) return - } else if nftMode && len(entryNodeIDs) > 0 { - if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil { - response.WriteJSON(w, response.ErrDefault(err.Error())) + } else if nftMode { + if len(entryNodeIDs) == 0 { + response.WriteJSON(w, response.ErrDefault("nftables 转发缺少入口节点")) return } + nftNodeID = entryNodeIDs[0] + } else if err := h.controlForwardServices(forward, "DeleteService", true); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return } if err := h.deleteForwardByID(id); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } + if nftNodeID > 0 { + if err := h.reconcileNftablesNodeByRequest(nftNodeID); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + } response.WriteJSON(w, response.OKEmpty()) } @@ -2577,7 +2584,19 @@ func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) { failures = appendBatchFailure(failures, id, "", accessErr) continue } - if err := h.controlForwardServices(forward, "DeleteService", true); err != nil { + var nftNodeID int64 + if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil { + f++ + failures = appendBatchFailure(failures, id, forward.Name, modeErr) + continue + } else if nftMode { + if len(entryNodeIDs) == 0 { + f++ + failures = appendBatchFailureReason(failures, id, forward.Name, "nftables 转发缺少入口节点") + continue + } + nftNodeID = entryNodeIDs[0] + } else if err := h.controlForwardServices(forward, "DeleteService", true); err != nil { f++ failures = appendBatchFailure(failures, id, forward.Name, err) continue @@ -2585,9 +2604,16 @@ func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) { if err := h.deleteForwardByID(id); err != nil { f++ failures = appendBatchFailure(failures, id, forward.Name, err) - } else { - s++ + continue } + if nftNodeID > 0 { + if err := h.reconcileNftablesNodeByRequest(nftNodeID); err != nil { + f++ + failures = appendBatchFailure(failures, id, forward.Name, err) + continue + } + } + s++ } response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: s, FailCount: f, Failures: failures})) } diff --git a/go-backend/internal/http/handler/nftables_runtime.go b/go-backend/internal/http/handler/nftables_runtime.go index eed00d1..d440d51 100644 --- a/go-backend/internal/http/handler/nftables_runtime.go +++ b/go-backend/internal/http/handler/nftables_runtime.go @@ -5,6 +5,7 @@ import ( "database/sql" "errors" "fmt" + "net" "net/http" "strings" "time" @@ -78,9 +79,13 @@ func (h *Handler) validateNftablesForwardRequest(tunnel *tunnelRecord, remoteAdd if len(entryNodeIDs) != 1 { return errors.New("nftables 节点仅支持单入口隧道") } - if _, err := runtimenft.ParseSingleTarget(remoteAddr); err != nil { + target, err := runtimenft.ParseSingleTarget(remoteAddr) + if err != nil { return err } + if net.ParseIP(strings.Trim(strings.TrimSpace(target.Host), "[]")) == nil { + return errors.New("nftables 节点仅支持 IP 目标地址") + } return nil } diff --git a/go-backend/internal/http/handler/nftables_runtime_test.go b/go-backend/internal/http/handler/nftables_runtime_test.go index 14145f6..6cfacef 100644 --- a/go-backend/internal/http/handler/nftables_runtime_test.go +++ b/go-backend/internal/http/handler/nftables_runtime_test.go @@ -280,6 +280,48 @@ func TestNodeUpdatePreservesExistingNftablesSecretsWhenFieldsOmitted(t *testing. } } +func TestValidateNftablesForwardRequestRejectsHostnameTarget(t *testing.T) { + fixture := setupNftablesHandler(t) + h := fixture.handler + tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID) + tunnel, err := h.getTunnelRecord(tunnelID) + if err != nil { + t.Fatalf("load tunnel: %v", err) + } + + err = h.validateNftablesForwardRequest(tunnel, "example.com:443", []int64{fixture.nodeID}) + if err == nil { + t.Fatalf("expected hostname target to be rejected") + } + if !strings.Contains(err.Error(), "IP") { + t.Fatalf("expected IP literal validation error, got %q", err) + } +} + +func TestForwardDeleteReconcilesNftablesAfterDBDelete(t *testing.T) { + fixture := setupNftablesHandler(t) + h := fixture.handler + seedNftablesSSHConfig(t, h, fixture.nodeID) + tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID) + forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080") + manager := &fakeNftablesManager{} + h.nftablesManager = manager + + req := newAuthenticatedJSONRequest(t, map[string]int64{"id": forward.ID}) + res := httptest.NewRecorder() + h.forwardDelete(res, req) + assertNftablesSuccessWithBody(t, res) + if manager.reconcileHit != 1 { + t.Fatalf("expected reconcile once, got %d", manager.reconcileHit) + } + if len(manager.lastPlan.Rules) != 0 { + t.Fatalf("expected reconcile after DB delete to render no rules, got %+v", manager.lastPlan.Rules) + } + if _, err := h.getForwardRecord(forward.ID); !errors.Is(err, errForwardNotFound) { + t.Fatalf("expected forward to be deleted, got %v", err) + } +} + func TestForwardForceDeleteRemovesNftablesBindingAndReconciles(t *testing.T) { fixture := setupNftablesHandler(t) h := fixture.handler @@ -319,6 +361,30 @@ func TestForwardForceDeleteRemovesNftablesBindingAndReconciles(t *testing.T) { } } +func TestForwardBatchDeleteReconcilesNftablesAfterDBDelete(t *testing.T) { + fixture := setupNftablesHandler(t) + h := fixture.handler + seedNftablesSSHConfig(t, h, fixture.nodeID) + tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID) + forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080") + manager := &fakeNftablesManager{} + h.nftablesManager = manager + + req := newAuthenticatedJSONRequest(t, map[string][]int64{"ids": {forward.ID}}) + res := httptest.NewRecorder() + h.forwardBatchDelete(res, req) + assertNftablesSuccessWithBody(t, res) + if manager.reconcileHit != 1 { + t.Fatalf("expected reconcile once, got %d", manager.reconcileHit) + } + if len(manager.lastPlan.Rules) != 0 { + t.Fatalf("expected reconcile after DB delete to render no rules, got %+v", manager.lastPlan.Rules) + } + if _, err := h.getForwardRecord(forward.ID); !errors.Is(err, errForwardNotFound) { + t.Fatalf("expected forward to be deleted, got %v", err) + } +} + func TestForwardBatchRedeployUsesNftablesReconcile(t *testing.T) { fixture := setupNftablesHandler(t) h := fixture.handler diff --git a/go-backend/internal/runtime/nftables/collector.go b/go-backend/internal/runtime/nftables/collector.go index faeabe9..d487402 100644 --- a/go-backend/internal/runtime/nftables/collector.go +++ b/go-backend/internal/runtime/nftables/collector.go @@ -94,6 +94,7 @@ func parseCounterRule(rule nftCounterRule) (CounterSample, bool, error) { var ( counter nftCounter hasCounter bool + comment = rule.Comment ) for _, expr := range rule.Expr { @@ -104,12 +105,17 @@ func parseCounterRule(rule nftCounterRule) (CounterSample, bool, error) { hasCounter = true continue } + if rawComment, ok := expr["comment"]; ok && strings.TrimSpace(comment) == "" { + if err := json.Unmarshal(rawComment, &comment); err != nil { + return CounterSample{}, false, err + } + } } if !hasCounter { return CounterSample{}, false, nil } - sample, ok := ParseCounterComment(rule.Comment) + sample, ok := ParseCounterComment(comment) if !ok { return CounterSample{}, false, nil } diff --git a/go-backend/internal/runtime/nftables/collector_test.go b/go-backend/internal/runtime/nftables/collector_test.go index 2ba065a..e5ee097 100644 --- a/go-backend/internal/runtime/nftables/collector_test.go +++ b/go-backend/internal/runtime/nftables/collector_test.go @@ -110,6 +110,39 @@ func TestParseCounterSamplesUsesRuleLevelComment(t *testing.T) { } } +func TestParseCounterSamplesUsesExprLevelComment(t *testing.T) { + raw := []byte(`{ + "nftables": [ + {"rule": { + "table": "flvx", + "chain": "forward", + "expr": [ + {"counter": {"packets": 4, "bytes": 3072}}, + {"comment": "flvx forward:78 from-target tcp"} + ] + }} + ] + }`) + + samples, err := ParseCounterSamples(raw) + if err != nil { + t.Fatalf("ParseCounterSamples: %v", err) + } + if len(samples) != 1 { + t.Fatalf("expected 1 sample, got %d: %+v", len(samples), samples) + } + want := CounterSample{ + ForwardID: 78, + Direction: CounterDirectionFromTarget, + Protocol: "tcp", + Bytes: 3072, + Packets: 4, + } + if samples[0] != want { + t.Fatalf("expected %+v, got %+v", want, samples[0]) + } +} + func TestParseCounterSamplesMalformedJSONReturnsError(t *testing.T) { if _, err := ParseCounterSamples([]byte(`{"nftables": [`)); err == nil { t.Fatal("expected malformed JSON error") diff --git a/go-backend/internal/runtime/nftables/renderer.go b/go-backend/internal/runtime/nftables/renderer.go index 34f6998..cb816f9 100644 --- a/go-backend/internal/runtime/nftables/renderer.go +++ b/go-backend/internal/runtime/nftables/renderer.go @@ -46,14 +46,16 @@ func RenderTable(plan NodePlan) string { } targetHost := strings.Trim(strings.TrimSpace(rule.TargetHost), "[]") for _, protocol := range normalizedProtocols(rule.Protocols) { - b.WriteString(fmt.Sprintf(" %s daddr %s %s dport %d counter comment %q\n", + b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n", + rule.InPort, family, targetHost, protocol, rule.TargetPort, counterComment(rule.ForwardID, CounterDirectionToTarget, protocol), )) - b.WriteString(fmt.Sprintf(" %s saddr %s %s sport %d counter comment %q\n", + b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n", + rule.InPort, family, targetHost, protocol, diff --git a/go-backend/internal/runtime/nftables/renderer_test.go b/go-backend/internal/runtime/nftables/renderer_test.go index b566f80..001015f 100644 --- a/go-backend/internal/runtime/nftables/renderer_test.go +++ b/go-backend/internal/runtime/nftables/renderer_test.go @@ -62,10 +62,10 @@ func TestRenderTableIncludesForwardAccountingCounters(t *testing.T) { wantLines := []string{ `tcp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat tcp"`, `udp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat udp"`, - `ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`, - `ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`, - `ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`, - `ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`, + `ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`, + `ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`, + `ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`, + `ct original proto-dst 12345 ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`, } for _, want := range wantLines { if !strings.Contains(got, want) { @@ -89,8 +89,29 @@ func TestRenderTableIncludesIPv6ForwardAccountingCounters(t *testing.T) { got := RenderTable(plan) wantLines := []string{ `tcp dport 12346 counter dnat ip6 to [2001:db8::20]:8443 comment "flvx forward:43 dnat tcp"`, - `ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`, - `ip6 saddr 2001:db8::20 tcp sport 8443 counter comment "flvx forward:43 from-target tcp"`, + `ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`, + `ct original proto-dst 12346 ip6 saddr 2001:db8::20 tcp sport 8443 counter comment "flvx forward:43 from-target tcp"`, + } + for _, want := range wantLines { + if !strings.Contains(got, want) { + t.Fatalf("RenderTable() missing %q\n%s", want, got) + } + } +} + +func TestRenderTableAccountingCountersIncludeOriginalPort(t *testing.T) { + plan := NodePlan{ + NodeID: 7, + Rules: []Rule{ + {ForwardID: 42, InPort: 12345, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp"}}, + {ForwardID: 43, InPort: 12346, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp"}}, + }, + } + + got := RenderTable(plan) + wantLines := []string{ + `ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`, + `ct original proto-dst 12346 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:43 to-target tcp"`, } for _, want := range wantLines { if !strings.Contains(got, want) {