mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-07 10:16:38 +08:00
fix(nftables): harden traffic accounting edges
This commit is contained in:
@@ -2412,23 +2412,30 @@ func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) {
|
|||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
var nftNodeID int64
|
||||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
||||||
response.WriteJSON(w, response.Err(-2, modeErr.Error()))
|
response.WriteJSON(w, response.Err(-2, modeErr.Error()))
|
||||||
return
|
return
|
||||||
} else if nftMode && len(entryNodeIDs) > 0 {
|
} else if nftMode {
|
||||||
if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil {
|
if len(entryNodeIDs) == 0 {
|
||||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
response.WriteJSON(w, response.ErrDefault("nftables 转发缺少入口节点"))
|
||||||
return
|
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 {
|
if err := h.deleteForwardByID(id); err != nil {
|
||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if nftNodeID > 0 {
|
||||||
|
if err := h.reconcileNftablesNodeByRequest(nftNodeID); err != nil {
|
||||||
|
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
response.WriteJSON(w, response.OKEmpty())
|
response.WriteJSON(w, response.OKEmpty())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -2577,7 +2584,19 @@ func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) {
|
|||||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||||
continue
|
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++
|
f++
|
||||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||||
continue
|
continue
|
||||||
@@ -2585,9 +2604,16 @@ func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) {
|
|||||||
if err := h.deleteForwardByID(id); err != nil {
|
if err := h.deleteForwardByID(id); err != nil {
|
||||||
f++
|
f++
|
||||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||||
} else {
|
continue
|
||||||
s++
|
|
||||||
}
|
}
|
||||||
|
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}))
|
response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: s, FailCount: f, Failures: failures}))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import (
|
|||||||
"database/sql"
|
"database/sql"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -78,9 +79,13 @@ func (h *Handler) validateNftablesForwardRequest(tunnel *tunnelRecord, remoteAdd
|
|||||||
if len(entryNodeIDs) != 1 {
|
if len(entryNodeIDs) != 1 {
|
||||||
return errors.New("nftables 节点仅支持单入口隧道")
|
return errors.New("nftables 节点仅支持单入口隧道")
|
||||||
}
|
}
|
||||||
if _, err := runtimenft.ParseSingleTarget(remoteAddr); err != nil {
|
target, err := runtimenft.ParseSingleTarget(remoteAddr)
|
||||||
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
if net.ParseIP(strings.Trim(strings.TrimSpace(target.Host), "[]")) == nil {
|
||||||
|
return errors.New("nftables 节点仅支持 IP 目标地址")
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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) {
|
func TestForwardForceDeleteRemovesNftablesBindingAndReconciles(t *testing.T) {
|
||||||
fixture := setupNftablesHandler(t)
|
fixture := setupNftablesHandler(t)
|
||||||
h := fixture.handler
|
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) {
|
func TestForwardBatchRedeployUsesNftablesReconcile(t *testing.T) {
|
||||||
fixture := setupNftablesHandler(t)
|
fixture := setupNftablesHandler(t)
|
||||||
h := fixture.handler
|
h := fixture.handler
|
||||||
|
|||||||
@@ -94,6 +94,7 @@ func parseCounterRule(rule nftCounterRule) (CounterSample, bool, error) {
|
|||||||
var (
|
var (
|
||||||
counter nftCounter
|
counter nftCounter
|
||||||
hasCounter bool
|
hasCounter bool
|
||||||
|
comment = rule.Comment
|
||||||
)
|
)
|
||||||
|
|
||||||
for _, expr := range rule.Expr {
|
for _, expr := range rule.Expr {
|
||||||
@@ -104,12 +105,17 @@ func parseCounterRule(rule nftCounterRule) (CounterSample, bool, error) {
|
|||||||
hasCounter = true
|
hasCounter = true
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
if rawComment, ok := expr["comment"]; ok && strings.TrimSpace(comment) == "" {
|
||||||
|
if err := json.Unmarshal(rawComment, &comment); err != nil {
|
||||||
|
return CounterSample{}, false, err
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if !hasCounter {
|
if !hasCounter {
|
||||||
return CounterSample{}, false, nil
|
return CounterSample{}, false, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
sample, ok := ParseCounterComment(rule.Comment)
|
sample, ok := ParseCounterComment(comment)
|
||||||
if !ok {
|
if !ok {
|
||||||
return CounterSample{}, false, nil
|
return CounterSample{}, false, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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) {
|
func TestParseCounterSamplesMalformedJSONReturnsError(t *testing.T) {
|
||||||
if _, err := ParseCounterSamples([]byte(`{"nftables": [`)); err == nil {
|
if _, err := ParseCounterSamples([]byte(`{"nftables": [`)); err == nil {
|
||||||
t.Fatal("expected malformed JSON error")
|
t.Fatal("expected malformed JSON error")
|
||||||
|
|||||||
@@ -46,14 +46,16 @@ func RenderTable(plan NodePlan) string {
|
|||||||
}
|
}
|
||||||
targetHost := strings.Trim(strings.TrimSpace(rule.TargetHost), "[]")
|
targetHost := strings.Trim(strings.TrimSpace(rule.TargetHost), "[]")
|
||||||
for _, protocol := range normalizedProtocols(rule.Protocols) {
|
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,
|
family,
|
||||||
targetHost,
|
targetHost,
|
||||||
protocol,
|
protocol,
|
||||||
rule.TargetPort,
|
rule.TargetPort,
|
||||||
counterComment(rule.ForwardID, CounterDirectionToTarget, protocol),
|
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,
|
family,
|
||||||
targetHost,
|
targetHost,
|
||||||
protocol,
|
protocol,
|
||||||
|
|||||||
@@ -62,10 +62,10 @@ func TestRenderTableIncludesForwardAccountingCounters(t *testing.T) {
|
|||||||
wantLines := []string{
|
wantLines := []string{
|
||||||
`tcp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat tcp"`,
|
`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"`,
|
`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"`,
|
`ct original proto-dst 12345 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"`,
|
`ct original proto-dst 12345 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"`,
|
`ct original proto-dst 12345 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 saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`,
|
||||||
}
|
}
|
||||||
for _, want := range wantLines {
|
for _, want := range wantLines {
|
||||||
if !strings.Contains(got, want) {
|
if !strings.Contains(got, want) {
|
||||||
@@ -89,8 +89,29 @@ func TestRenderTableIncludesIPv6ForwardAccountingCounters(t *testing.T) {
|
|||||||
got := RenderTable(plan)
|
got := RenderTable(plan)
|
||||||
wantLines := []string{
|
wantLines := []string{
|
||||||
`tcp dport 12346 counter dnat ip6 to [2001:db8::20]:8443 comment "flvx forward:43 dnat tcp"`,
|
`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"`,
|
`ct original proto-dst 12346 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 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 {
|
for _, want := range wantLines {
|
||||||
if !strings.Contains(got, want) {
|
if !strings.Contains(got, want) {
|
||||||
|
|||||||
Reference in New Issue
Block a user