From 3b8a5c15cced0b6fbdeabf3e4c76373a2390b91b Mon Sep 17 00:00:00 2001 From: Viktor Liu Date: Mon, 25 May 2026 12:14:27 +0200 Subject: [PATCH] Address CodeRabbit review on DNAT refcount --- .../iptables/dnat_refcount_linux_test.go | 12 ++++++--- .../nftables/dnat_refcount_linux_test.go | 6 +++-- client/firewall/nftables/router_linux.go | 25 +++++++++++-------- 3 files changed, 27 insertions(+), 16 deletions(-) diff --git a/client/firewall/iptables/dnat_refcount_linux_test.go b/client/firewall/iptables/dnat_refcount_linux_test.go index 26df653ae..6d1cebeb2 100644 --- a/client/firewall/iptables/dnat_refcount_linux_test.go +++ b/client/firewall/iptables/dnat_refcount_linux_test.go @@ -85,12 +85,14 @@ func TestIptablesDNAT_RefcountBalancedV4(t *testing.T) { r2, err := m.AddDNATRule(iptDnatV4(7082)) require.NoError(t, err, "add v4 dnat 2") - v4, _ = state.Counts() + v4, v6 = state.Counts() require.Equal(t, 2, v4, "v4 refcount after second add") + require.Equal(t, 0, v6, "v6 refcount unchanged") require.NoError(t, m.DeleteDNATRule(r1)) - v4, _ = state.Counts() + v4, v6 = state.Counts() require.Equal(t, 1, v4, "v4 refcount after first delete") + require.Equal(t, 0, v6, "v6 refcount unchanged") require.NoError(t, m.DeleteDNATRule(r2)) v4, v6 = state.Counts() @@ -114,11 +116,13 @@ func TestIptablesDNAT_RefcountBalancedV6(t *testing.T) { r2, err := m.AddDNATRule(iptDnatV6(9082)) require.NoError(t, err, "add v6 dnat 2") - _, v6 = state.Counts() + v4, v6 = state.Counts() + require.Equal(t, 0, v4, "v4 refcount unchanged") require.Equal(t, 2, v6, "v6 refcount after second add") require.NoError(t, m.DeleteDNATRule(r1)) - _, v6 = state.Counts() + v4, v6 = state.Counts() + require.Equal(t, 0, v4, "v4 refcount unchanged") require.Equal(t, 1, v6, "v6 refcount after first delete") require.NoError(t, m.DeleteDNATRule(r2)) diff --git a/client/firewall/nftables/dnat_refcount_linux_test.go b/client/firewall/nftables/dnat_refcount_linux_test.go index e4d1449d3..8df535976 100644 --- a/client/firewall/nftables/dnat_refcount_linux_test.go +++ b/client/firewall/nftables/dnat_refcount_linux_test.go @@ -94,8 +94,9 @@ func TestNftablesDNAT_RefcountBalancedV4(t *testing.T) { require.Equal(t, 0, v6, "v6 refcount unchanged") require.NoError(t, m.DeleteDNATRule(r1), "delete v4 dnat 1") - v4, _ = state.Counts() + v4, v6 = state.Counts() require.Equal(t, 1, v4, "v4 refcount after first delete") + require.Equal(t, 0, v6, "v6 refcount unchanged") require.NoError(t, m.DeleteDNATRule(r2), "delete v4 dnat 2") v4, v6 = state.Counts() @@ -124,7 +125,8 @@ func TestNftablesDNAT_RefcountBalancedV6(t *testing.T) { require.Equal(t, 2, v6, "v6 refcount after second add") require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat 1") - _, v6 = state.Counts() + v4, v6 = state.Counts() + require.Equal(t, 0, v4, "v4 refcount unchanged") require.Equal(t, 1, v6, "v6 refcount after first delete") require.NoError(t, m.DeleteDNATRule(r2), "delete v6 dnat 2") diff --git a/client/firewall/nftables/router_linux.go b/client/firewall/nftables/router_linux.go index 9e5fb431d..27f6e0a68 100644 --- a/client/firewall/nftables/router_linux.go +++ b/client/firewall/nftables/router_linux.go @@ -1560,7 +1560,18 @@ func (r *router) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { return nil, fmt.Errorf("convert protocol to number: %w", err) } + // Request forwarding before queueing rules: addDnatRedirect/addDnatMasq + // buffer netlink messages on r.conn that the next caller's Flush would + // commit if we returned without flushing them ourselves. + v6 := r.af.tableFamily == nftables.TableFamilyIPv6 + if err := r.ipFwdState.RequestForwarding(v6); err != nil { + return nil, fmt.Errorf("enable forwarding: %w", err) + } + if err := r.addDnatRedirect(rule, protoNum, ruleKey); err != nil { + if rerr := r.ipFwdState.ReleaseForwarding(v6); rerr != nil { + log.Warnf("rollback forwarding refcount: %v", rerr) + } return nil, err } @@ -1571,13 +1582,6 @@ func (r *router) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { // We also cannot just add "oif accept" there and filter in our own table as we don't know what is supposed to be allowed. // TODO: find chains with drop policies and add rules there - v6 := r.af.tableFamily == nftables.TableFamilyIPv6 - if err := r.ipFwdState.RequestForwarding(v6); err != nil { - delete(r.rules, ruleKey+dnatSuffix) - delete(r.rules, ruleKey+snatSuffix) - return nil, fmt.Errorf("enable forwarding: %w", err) - } - if err := r.conn.Flush(); err != nil { if rerr := r.ipFwdState.ReleaseForwarding(v6); rerr != nil { log.Warnf("rollback forwarding refcount: %v", rerr) @@ -1832,9 +1836,10 @@ func (r *router) DeleteDNATRule(rule firewall.Rule) error { if merr == nil { delete(r.rules, ruleKey+dnatSuffix) delete(r.rules, ruleKey+snatSuffix) - if err := r.ipFwdState.ReleaseForwarding(r.af.tableFamily == nftables.TableFamilyIPv6); err != nil { - log.Errorf("%v", err) - } + } + + if err := r.ipFwdState.ReleaseForwarding(r.af.tableFamily == nftables.TableFamilyIPv6); err != nil { + log.Errorf("%v", err) } return nberrors.FormatErrorOrNil(merr)