From 9c5a1546c971af89166dc445042ebac99b4a2f0c Mon Sep 17 00:00:00 2001 From: Viktor Liu Date: Mon, 10 Aug 2026 18:01:40 +0200 Subject: [PATCH] Address review comments --- .../iptables/dnat_refcount_linux_test.go | 67 ++++++++++--------- client/firewall/iptables/router_linux.go | 18 +++-- .../nftables/dnat_refcount_linux_test.go | 67 ++++++++++--------- .../routemanager/ipfwdstate/ipfwdstate.go | 20 ++++-- .../ipfwdstate_privileged_linux_test.go | 39 +++++++++++ 5 files changed, 135 insertions(+), 76 deletions(-) create mode 100644 client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go diff --git a/client/firewall/iptables/dnat_refcount_linux_test.go b/client/firewall/iptables/dnat_refcount_linux_test.go index 549539475..681bc0b99 100644 --- a/client/firewall/iptables/dnat_refcount_linux_test.go +++ b/client/firewall/iptables/dnat_refcount_linux_test.go @@ -6,6 +6,7 @@ import ( "net/netip" "testing" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" fw "github.com/netbirdio/netbird/client/firewall/manager" @@ -85,13 +86,13 @@ func TestIptablesRouting_RepeatedEnableSingleReference(t *testing.T) { require.NoError(t, m.EnableRouting(), "second enable") require.NoError(t, m.EnableRouting(), "third enable") v4, v6 := state.Counts() - require.Equal(t, 1, v4, "repeated enable holds a single v4 reference") - require.Equal(t, 1, v6, "repeated enable holds a single v6 reference") + assert.Equal(t, 1, v4, "repeated enable holds a single v4 reference") + assert.Equal(t, 1, v6, "repeated enable holds a single v6 reference") require.NoError(t, m.DisableRouting(), "disable") v4, v6 = state.Counts() - require.Equal(t, 0, v4, "single disable releases the v4 reference") - require.Equal(t, 0, v6, "single disable releases the v6 reference") + assert.Equal(t, 0, v4, "single disable releases the v4 reference") + assert.Equal(t, 0, v6, "single disable releases the v6 reference") } // TestIptablesRouting_DisableKeepsDNATReference verifies that an unpaired @@ -105,11 +106,11 @@ func TestIptablesRouting_DisableKeepsDNATReference(t *testing.T) { require.NoError(t, m.DisableRouting(), "unpaired disable") _, v6 := state.Counts() - require.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting") + assert.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting") require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat") _, v6 = state.Counts() - require.Equal(t, 0, v6, "delete releases the DNAT reference") + assert.Equal(t, 0, v6, "delete releases the DNAT reference") } // TestIptablesDNAT_RefcountBalancedV4 covers a Balanced Add/Delete pair on v4. @@ -120,24 +121,24 @@ func TestIptablesDNAT_RefcountBalancedV4(t *testing.T) { r1, err := m.AddDNATRule(iptDnatV4(7081)) require.NoError(t, err, "add v4 dnat 1") v4, v6 := state.Counts() - require.Equal(t, 1, v4, "v4 refcount after first add") - require.Equal(t, 0, v6, "v6 refcount unchanged") + assert.Equal(t, 1, v4, "v4 refcount after first add") + assert.Equal(t, 0, v6, "v6 refcount unchanged") r2, err := m.AddDNATRule(iptDnatV4(7082)) require.NoError(t, err, "add v4 dnat 2") v4, v6 = state.Counts() - require.Equal(t, 2, v4, "v4 refcount after second add") - require.Equal(t, 0, v6, "v6 refcount unchanged") + assert.Equal(t, 2, v4, "v4 refcount after second add") + assert.Equal(t, 0, v6, "v6 refcount unchanged") require.NoError(t, m.DeleteDNATRule(r1)) v4, v6 = state.Counts() - require.Equal(t, 1, v4, "v4 refcount after first delete") - require.Equal(t, 0, v6, "v6 refcount unchanged") + assert.Equal(t, 1, v4, "v4 refcount after first delete") + assert.Equal(t, 0, v6, "v6 refcount unchanged") require.NoError(t, m.DeleteDNATRule(r2)) v4, v6 = state.Counts() - require.Equal(t, 0, v4, "v4 refcount after second delete") - require.Equal(t, 0, v6, "v6 refcount unchanged") + assert.Equal(t, 0, v4, "v4 refcount after second delete") + assert.Equal(t, 0, v6, "v6 refcount unchanged") } // TestIptablesDNAT_RefcountBalancedV6 checks the v6 path increments v6 only and @@ -151,24 +152,24 @@ func TestIptablesDNAT_RefcountBalancedV6(t *testing.T) { r1, err := m.AddDNATRule(iptDnatV6(9081)) require.NoError(t, err, "add v6 dnat 1") v4, v6 := state.Counts() - require.Equal(t, 0, v4) - require.Equal(t, 1, v6, "v6 refcount after first add") + assert.Equal(t, 0, v4) + assert.Equal(t, 1, v6, "v6 refcount after first add") r2, err := m.AddDNATRule(iptDnatV6(9082)) require.NoError(t, err, "add v6 dnat 2") v4, v6 = state.Counts() - require.Equal(t, 0, v4, "v4 refcount unchanged") - require.Equal(t, 2, v6, "v6 refcount after second add") + assert.Equal(t, 0, v4, "v4 refcount unchanged") + assert.Equal(t, 2, v6, "v6 refcount after second add") require.NoError(t, m.DeleteDNATRule(r1)) v4, v6 = state.Counts() - require.Equal(t, 0, v4, "v4 refcount unchanged") - require.Equal(t, 1, v6, "v6 refcount after first delete") + assert.Equal(t, 0, v4, "v4 refcount unchanged") + assert.Equal(t, 1, v6, "v6 refcount after first delete") require.NoError(t, m.DeleteDNATRule(r2)) v4, v6 = state.Counts() - require.Equal(t, 0, v4) - require.Equal(t, 0, v6, "v6 refcount after second delete") + assert.Equal(t, 0, v4) + assert.Equal(t, 0, v6, "v6 refcount after second delete") } // TestIptablesDNAT_DuplicateAddNoLeak verifies the duplicate-rule path returns @@ -181,16 +182,16 @@ func TestIptablesDNAT_DuplicateAddNoLeak(t *testing.T) { r1, err := m.AddDNATRule(rule) require.NoError(t, err) v4, _ := state.Counts() - require.Equal(t, 1, v4) + assert.Equal(t, 1, v4) _, err = m.AddDNATRule(rule) require.NoError(t, err, "duplicate add") v4, _ = state.Counts() - require.Equal(t, 1, v4, "duplicate add must not increment") + assert.Equal(t, 1, v4, "duplicate add must not increment") require.NoError(t, m.DeleteDNATRule(r1)) v4, _ = state.Counts() - require.Equal(t, 0, v4, "single delete must drop to zero") + assert.Equal(t, 0, v4, "single delete must drop to zero") } // TestIptablesDNAT_DeleteMissingNoUnderflow verifies Delete on an unknown rule @@ -202,19 +203,19 @@ func TestIptablesDNAT_DeleteMissingNoUnderflow(t *testing.T) { phantom := iptDnatV4(7099) require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4") v4, v6 := state.Counts() - require.Equal(t, 0, v4) - require.Equal(t, 0, v6) + assert.Equal(t, 0, v4) + assert.Equal(t, 0, v6) phantom6 := iptDnatV6(9099) require.NoError(t, m.DeleteDNATRule(&phantom6), "delete missing v6") v4, v6 = state.Counts() - require.Equal(t, 0, v4) - require.Equal(t, 0, v6) + assert.Equal(t, 0, v4) + assert.Equal(t, 0, v6) r1, err := m.AddDNATRule(iptDnatV4(7100)) require.NoError(t, err) v4, _ = state.Counts() - require.Equal(t, 1, v4, "real add still increments after phantom delete") + assert.Equal(t, 1, v4, "real add still increments after phantom delete") require.NoError(t, m.DeleteDNATRule(r1)) } @@ -227,13 +228,13 @@ func TestIptablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) { r1, err := m.AddDNATRule(iptDnatV6(9083)) require.NoError(t, err) _, v6 := state.Counts() - require.Equal(t, 1, v6) + assert.Equal(t, 1, v6) require.NoError(t, m.DeleteDNATRule(r1), "first delete") _, v6 = state.Counts() - require.Equal(t, 0, v6) + assert.Equal(t, 0, v6) require.NoError(t, m.DeleteDNATRule(r1), "second delete must be no-op") _, v6 = state.Counts() - require.Equal(t, 0, v6, "double delete must not underflow") + assert.Equal(t, 0, v6, "double delete must not underflow") } diff --git a/client/firewall/iptables/router_linux.go b/client/firewall/iptables/router_linux.go index 8912242f4..01b18570c 100644 --- a/client/firewall/iptables/router_linux.go +++ b/client/firewall/iptables/router_linux.go @@ -893,26 +893,34 @@ func (r *router) DeleteDNATRule(rule firewall.Rule) error { if dnatRule, exists := r.rules[ruleKey+dnatSuffix]; exists { if err := r.iptablesClient.Delete(tableNat, chainRTRDR, dnatRule...); err != nil { merr = multierror.Append(merr, fmt.Errorf("delete DNAT rule: %w", err)) + } else { + delete(r.rules, ruleKey+dnatSuffix) } - delete(r.rules, ruleKey+dnatSuffix) } if snatRule, exists := r.rules[ruleKey+snatSuffix]; exists { if err := r.iptablesClient.Delete(tableNat, chainRTNAT, snatRule...); err != nil { merr = multierror.Append(merr, fmt.Errorf("delete SNAT rule: %w", err)) + } else { + delete(r.rules, ruleKey+snatSuffix) } - delete(r.rules, ruleKey+snatSuffix) } if fwdRule, exists := r.rules[ruleKey+fwdSuffix]; exists { if err := r.iptablesClient.Delete(tableFilter, chainRTFWDOUT, fwdRule...); err != nil { merr = multierror.Append(merr, fmt.Errorf("delete forward rule: %w", err)) + } else { + delete(r.rules, ruleKey+fwdSuffix) } - delete(r.rules, ruleKey+fwdSuffix) } - if err := r.ipFwdState.ReleaseForwarding(r.v6); err != nil { - log.Errorf("%v", err) + // Release the refcount only once all rules are gone from the kernel. On + // partial failure the failed entries stay in r.rules so a retry can remove + // them and release then. + if merr == nil { + if err := r.ipFwdState.ReleaseForwarding(r.v6); err != nil { + log.Errorf("%v", err) + } } r.updateState() diff --git a/client/firewall/nftables/dnat_refcount_linux_test.go b/client/firewall/nftables/dnat_refcount_linux_test.go index d2c08e70a..86079676f 100644 --- a/client/firewall/nftables/dnat_refcount_linux_test.go +++ b/client/firewall/nftables/dnat_refcount_linux_test.go @@ -6,6 +6,7 @@ import ( "net/netip" "testing" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" fw "github.com/netbirdio/netbird/client/firewall/manager" @@ -86,24 +87,24 @@ func TestNftablesDNAT_RefcountBalancedV4(t *testing.T) { r1, err := m.AddDNATRule(dnatV4(8081)) require.NoError(t, err, "add v4 dnat 1") v4, v6 := state.Counts() - require.Equal(t, 1, v4, "v4 refcount after first add") - require.Equal(t, 0, v6, "v6 refcount unchanged") + assert.Equal(t, 1, v4, "v4 refcount after first add") + assert.Equal(t, 0, v6, "v6 refcount unchanged") r2, err := m.AddDNATRule(dnatV4(8082)) require.NoError(t, err, "add v4 dnat 2") v4, v6 = state.Counts() - require.Equal(t, 2, v4, "v4 refcount after second add") - require.Equal(t, 0, v6, "v6 refcount unchanged") + assert.Equal(t, 2, v4, "v4 refcount after second add") + assert.Equal(t, 0, v6, "v6 refcount unchanged") require.NoError(t, m.DeleteDNATRule(r1), "delete v4 dnat 1") v4, v6 = state.Counts() - require.Equal(t, 1, v4, "v4 refcount after first delete") - require.Equal(t, 0, v6, "v6 refcount unchanged") + assert.Equal(t, 1, v4, "v4 refcount after first delete") + assert.Equal(t, 0, v6, "v6 refcount unchanged") require.NoError(t, m.DeleteDNATRule(r2), "delete v4 dnat 2") v4, v6 = state.Counts() - require.Equal(t, 0, v4, "v4 refcount after second delete") - require.Equal(t, 0, v6, "v6 refcount unchanged") + assert.Equal(t, 0, v4, "v4 refcount after second delete") + assert.Equal(t, 0, v6, "v6 refcount unchanged") } // TestNftablesDNAT_RefcountBalancedV6 verifies the v6 path increments v6 only @@ -117,24 +118,24 @@ func TestNftablesDNAT_RefcountBalancedV6(t *testing.T) { r1, err := m.AddDNATRule(dnatV6(9091)) require.NoError(t, err, "add v6 dnat 1") v4, v6 := state.Counts() - require.Equal(t, 0, v4, "v4 refcount unchanged") - require.Equal(t, 1, v6, "v6 refcount after first add") + assert.Equal(t, 0, v4, "v4 refcount unchanged") + assert.Equal(t, 1, v6, "v6 refcount after first add") r2, err := m.AddDNATRule(dnatV6(9092)) require.NoError(t, err, "add v6 dnat 2") v4, v6 = state.Counts() - require.Equal(t, 0, v4) - require.Equal(t, 2, v6, "v6 refcount after second add") + assert.Equal(t, 0, v4) + assert.Equal(t, 2, v6, "v6 refcount after second add") require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat 1") v4, v6 = state.Counts() - require.Equal(t, 0, v4, "v4 refcount unchanged") - require.Equal(t, 1, v6, "v6 refcount after first delete") + assert.Equal(t, 0, v4, "v4 refcount unchanged") + assert.Equal(t, 1, v6, "v6 refcount after first delete") require.NoError(t, m.DeleteDNATRule(r2), "delete v6 dnat 2") v4, v6 = state.Counts() - require.Equal(t, 0, v4) - require.Equal(t, 0, v6, "v6 refcount after second delete") + assert.Equal(t, 0, v4) + assert.Equal(t, 0, v6, "v6 refcount after second delete") } // TestNftablesDNAT_DuplicateAddNoLeak verifies that a duplicate Add (same @@ -147,17 +148,17 @@ func TestNftablesDNAT_DuplicateAddNoLeak(t *testing.T) { r1, err := m.AddDNATRule(rule) require.NoError(t, err, "add v4 dnat") v4, _ := state.Counts() - require.Equal(t, 1, v4) + assert.Equal(t, 1, v4) // duplicate add: same rule ID, must be a no-op for the refcount. _, err = m.AddDNATRule(rule) require.NoError(t, err, "duplicate add") v4, _ = state.Counts() - require.Equal(t, 1, v4, "duplicate add must not increment") + assert.Equal(t, 1, v4, "duplicate add must not increment") require.NoError(t, m.DeleteDNATRule(r1), "delete v4 dnat") v4, _ = state.Counts() - require.Equal(t, 0, v4, "single delete must drop to zero") + assert.Equal(t, 0, v4, "single delete must drop to zero") } // TestNftablesDNAT_DeleteMissingNoUnderflow verifies deleting a rule that was @@ -172,20 +173,20 @@ func TestNftablesDNAT_DeleteMissingNoUnderflow(t *testing.T) { phantom := dnatV4(8099) require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4 dnat") v4, v6 := state.Counts() - require.Equal(t, 0, v4, "v4 refcount unaffected by missing delete") - require.Equal(t, 0, v6, "v6 refcount unaffected") + assert.Equal(t, 0, v4, "v4 refcount unaffected by missing delete") + assert.Equal(t, 0, v6, "v6 refcount unaffected") phantom6 := dnatV6(9099) require.NoError(t, m.DeleteDNATRule(&phantom6), "delete missing v6 dnat") v4, v6 = state.Counts() - require.Equal(t, 0, v4) - require.Equal(t, 0, v6, "v6 refcount unaffected by missing delete") + assert.Equal(t, 0, v4) + assert.Equal(t, 0, v6, "v6 refcount unaffected by missing delete") // And after a phantom delete, a real add still results in count=1. r1, err := m.AddDNATRule(dnatV4(8100)) require.NoError(t, err, "add v4 dnat after phantom delete") v4, _ = state.Counts() - require.Equal(t, 1, v4, "real add still increments after phantom delete") + assert.Equal(t, 1, v4, "real add still increments after phantom delete") require.NoError(t, m.DeleteDNATRule(r1)) } @@ -200,13 +201,13 @@ func TestNftablesRouting_RepeatedEnableSingleReference(t *testing.T) { require.NoError(t, m.EnableRouting(), "second enable") require.NoError(t, m.EnableRouting(), "third enable") v4, v6 := state.Counts() - require.Equal(t, 1, v4, "repeated enable holds a single v4 reference") - require.Equal(t, 1, v6, "repeated enable holds a single v6 reference") + assert.Equal(t, 1, v4, "repeated enable holds a single v4 reference") + assert.Equal(t, 1, v6, "repeated enable holds a single v6 reference") require.NoError(t, m.DisableRouting(), "disable") v4, v6 = state.Counts() - require.Equal(t, 0, v4, "single disable releases the v4 reference") - require.Equal(t, 0, v6, "single disable releases the v6 reference") + assert.Equal(t, 0, v4, "single disable releases the v4 reference") + assert.Equal(t, 0, v6, "single disable releases the v6 reference") } // TestNftablesRouting_DisableKeepsDNATReference verifies that an unpaired @@ -220,11 +221,11 @@ func TestNftablesRouting_DisableKeepsDNATReference(t *testing.T) { require.NoError(t, m.DisableRouting(), "unpaired disable") _, v6 := state.Counts() - require.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting") + assert.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting") require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat") _, v6 = state.Counts() - require.Equal(t, 0, v6, "delete releases the DNAT reference") + assert.Equal(t, 0, v6, "delete releases the DNAT reference") } // TestNftablesDNAT_DoubleDeleteNoUnderflow verifies that deleting the same rule @@ -236,13 +237,13 @@ func TestNftablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) { r1, err := m.AddDNATRule(dnatV6(9093)) require.NoError(t, err) _, v6 := state.Counts() - require.Equal(t, 1, v6) + assert.Equal(t, 1, v6) require.NoError(t, m.DeleteDNATRule(r1), "first delete") _, v6 = state.Counts() - require.Equal(t, 0, v6) + assert.Equal(t, 0, v6) require.NoError(t, m.DeleteDNATRule(r1), "second delete must be no-op") _, v6 = state.Counts() - require.Equal(t, 0, v6, "double delete must not underflow") + assert.Equal(t, 0, v6, "double delete must not underflow") } diff --git a/client/internal/routemanager/ipfwdstate/ipfwdstate.go b/client/internal/routemanager/ipfwdstate/ipfwdstate.go index 3fdbc90d0..3d571e16b 100644 --- a/client/internal/routemanager/ipfwdstate/ipfwdstate.go +++ b/client/internal/routemanager/ipfwdstate/ipfwdstate.go @@ -28,6 +28,8 @@ type IPForwardingState struct { v6Saved map[string]int } +// NewIPForwardingState returns a state tracker for the IP-forwarding sysctls. +// wgIfaceName is excluded from the per-interface accept_ra handling. func NewIPForwardingState(wgIfaceName string) *IPForwardingState { return &IPForwardingState{wgIfaceName: wgIfaceName} } @@ -42,10 +44,10 @@ func (f *IPForwardingState) Counts() (v4, v6 int) { // RequestRouting takes the forwarding references for the routing path. It is // idempotent: while routing already holds a reference, further calls don't -// increment the refcounts. A v6 sysctl failure is logged and not returned so -// it can't take down v4 routing (the sysctl may be unwritable, e.g. read-only -// /proc/sys or IPv6 disabled on the kernel command line); v6 is retried on the -// next call. +// increment the refcounts, and a v4-only request releases a previously held v6 +// reference. A v6 sysctl failure is logged and not returned so it can't take +// down v4 routing (the sysctl may be unwritable, e.g. read-only /proc/sys or +// IPv6 disabled on the kernel command line); v6 is retried on the next call. func (f *IPForwardingState) RequestRouting(v6 bool) error { f.mu.Lock() defer f.mu.Unlock() @@ -57,7 +59,15 @@ func (f *IPForwardingState) RequestRouting(v6 bool) error { f.routingV4 = true } - if !v6 || f.routingV6 { + if !v6 { + if !f.routingV6 { + return nil + } + f.routingV6 = false + return f.releaseV6() + } + + if f.routingV6 { return nil } if err := f.requestV6(); err != nil { diff --git a/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go b/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go new file mode 100644 index 000000000..b4615ff02 --- /dev/null +++ b/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go @@ -0,0 +1,39 @@ +//go:build privileged + +package ipfwdstate + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// TestRequestRoutingV6ToV4Transition verifies that a v4-only routing request +// releases a previously held routing-owned v6 reference without touching +// references held by DNAT rules. +func TestRequestRoutingV6ToV4Transition(t *testing.T) { + f := NewIPForwardingState("wt-fwd-test") + + require.NoError(t, f.RequestRouting(true), "request routing with v6") + v4, v6 := f.Counts() + assert.Equal(t, 1, v4, "v4 reference held") + assert.Equal(t, 1, v6, "v6 reference held") + + require.NoError(t, f.RequestRouting(false), "request routing v4-only") + v4, v6 = f.Counts() + assert.Equal(t, 1, v4, "v4 reference kept") + assert.Equal(t, 0, v6, "routing-owned v6 reference released") + + // A DNAT-held reference survives a v4-only routing request. + require.NoError(t, f.RequestForwarding(true), "dnat v6 reference") + require.NoError(t, f.RequestRouting(false), "repeat v4-only request") + _, v6 = f.Counts() + assert.Equal(t, 1, v6, "dnat-held v6 reference survives") + require.NoError(t, f.ReleaseForwarding(true), "release dnat v6 reference") + + require.NoError(t, f.ReleaseRouting(), "release routing") + v4, v6 = f.Counts() + assert.Equal(t, 0, v4, "all v4 references released") + assert.Equal(t, 0, v6, "all v6 references released") +}