diff --git a/client/firewall/iptables/dnat_refcount_linux_test.go b/client/firewall/iptables/dnat_refcount_linux_test.go index 6d1cebeb2..549539475 100644 --- a/client/firewall/iptables/dnat_refcount_linux_test.go +++ b/client/firewall/iptables/dnat_refcount_linux_test.go @@ -1,3 +1,5 @@ +//go:build privileged + package iptables import ( @@ -72,6 +74,44 @@ func iptDnatV6(port uint16) fw.ForwardRule { } } +// TestIptablesRouting_RepeatedEnableSingleReference verifies that EnableRouting +// (called on every network-map update) holds at most one reference per family +// and a single DisableRouting drops both back to zero. +func TestIptablesRouting_RepeatedEnableSingleReference(t *testing.T) { + m := newIptRefcountManager(t, true) + state := m.router.ipFwdState + + require.NoError(t, m.EnableRouting(), "first enable") + 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") + + 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") +} + +// TestIptablesRouting_DisableKeepsDNATReference verifies that an unpaired +// DisableRouting does not release references held by active DNAT rules. +func TestIptablesRouting_DisableKeepsDNATReference(t *testing.T) { + m := newIptRefcountManager(t, true) + state := m.router.ipFwdState + + r1, err := m.AddDNATRule(iptDnatV6(9095)) + require.NoError(t, err, "add v6 dnat") + + require.NoError(t, m.DisableRouting(), "unpaired disable") + _, v6 := state.Counts() + require.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") +} + // TestIptablesDNAT_RefcountBalancedV4 covers a Balanced Add/Delete pair on v4. func TestIptablesDNAT_RefcountBalancedV4(t *testing.T) { m := newIptRefcountManager(t, false) diff --git a/client/firewall/iptables/manager_linux.go b/client/firewall/iptables/manager_linux.go index 90ba540f4..aa052d933 100644 --- a/client/firewall/iptables/manager_linux.go +++ b/client/firewall/iptables/manager_linux.go @@ -402,33 +402,12 @@ func (m *Manager) SetLogLevel(log.Level) { } func (m *Manager) EnableRouting() error { - if err := m.router.ipFwdState.RequestForwarding(false); err != nil { - return fmt.Errorf("enable IPv4 forwarding: %w", err) - } // v6 only when the overlay actually has v6. - if m.router6 == nil { - return nil - } - if err := m.router.ipFwdState.RequestForwarding(true); err != nil { - if rerr := m.router.ipFwdState.ReleaseForwarding(false); rerr != nil { - log.Warnf("rollback v4 forwarding: %v", rerr) - } - return fmt.Errorf("enable IPv6 forwarding: %w", err) - } - return nil + return m.router.ipFwdState.RequestRouting(m.router6 != nil) } func (m *Manager) DisableRouting() error { - var merr *multierror.Error - if err := m.router.ipFwdState.ReleaseForwarding(false); err != nil { - merr = multierror.Append(merr, fmt.Errorf("disable IPv4 forwarding: %w", err)) - } - if m.router6 != nil { - if err := m.router.ipFwdState.ReleaseForwarding(true); err != nil { - merr = multierror.Append(merr, fmt.Errorf("disable IPv6 forwarding: %w", err)) - } - } - return nberrors.FormatErrorOrNil(merr) + return m.router.ipFwdState.ReleaseRouting() } // AddDNATRule adds a DNAT rule diff --git a/client/firewall/iptables/router_linux.go b/client/firewall/iptables/router_linux.go index 4c1b4973a..8912242f4 100644 --- a/client/firewall/iptables/router_linux.go +++ b/client/firewall/iptables/router_linux.go @@ -836,21 +836,14 @@ func (r *router) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { for key, ruleInfo := range rules { if err := r.iptablesClient.Append(ruleInfo.table, ruleInfo.chain, ruleInfo.rule...); err != nil { - if rollbackErr := r.rollbackRules(rules); rollbackErr != nil { - log.Errorf("rollback failed: %v", rollbackErr) - } + r.cleanupFailedDNATAdd(rules) return nil, fmt.Errorf("add rule %s: %w", key, err) } r.rules[key] = ruleInfo.rule } if err := r.ipFwdState.RequestForwarding(r.v6); err != nil { - if rollbackErr := r.rollbackRules(rules); rollbackErr != nil { - log.Errorf("rollback failed: %v", rollbackErr) - } - for key := range rules { - delete(r.rules, key) - } + r.cleanupFailedDNATAdd(rules) return nil, fmt.Errorf("enable forwarding: %w", err) } @@ -858,6 +851,19 @@ func (r *router) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { return rule, nil } +// cleanupFailedDNATAdd removes the bookkeeping written by a partially applied +// AddDNATRule before rolling back the kernel rules, so no entries remain that +// never got a forwarding refcount. rollbackRules re-adds entries it failed to +// remove from the kernel. +func (r *router) cleanupFailedDNATAdd(rules map[string]ruleInfo) { + for key := range rules { + delete(r.rules, key) + } + if err := r.rollbackRules(rules); err != nil { + log.Errorf("rollback failed: %v", err) + } +} + func (r *router) rollbackRules(rules map[string]ruleInfo) error { var merr *multierror.Error for key, ruleInfo := range rules { diff --git a/client/firewall/nftables/dnat_refcount_linux_test.go b/client/firewall/nftables/dnat_refcount_linux_test.go index 8df535976..d2c08e70a 100644 --- a/client/firewall/nftables/dnat_refcount_linux_test.go +++ b/client/firewall/nftables/dnat_refcount_linux_test.go @@ -1,3 +1,5 @@ +//go:build privileged + package nftables import ( @@ -187,6 +189,44 @@ func TestNftablesDNAT_DeleteMissingNoUnderflow(t *testing.T) { require.NoError(t, m.DeleteDNATRule(r1)) } +// TestNftablesRouting_RepeatedEnableSingleReference verifies that EnableRouting +// (called on every network-map update) holds at most one reference per family +// and a single DisableRouting drops both back to zero. +func TestNftablesRouting_RepeatedEnableSingleReference(t *testing.T) { + m := newNftRefcountManager(t, true) + state := m.router.ipFwdState + + require.NoError(t, m.EnableRouting(), "first enable") + 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") + + 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") +} + +// TestNftablesRouting_DisableKeepsDNATReference verifies that an unpaired +// DisableRouting does not release references held by active DNAT rules. +func TestNftablesRouting_DisableKeepsDNATReference(t *testing.T) { + m := newNftRefcountManager(t, true) + state := m.router.ipFwdState + + r1, err := m.AddDNATRule(dnatV6(9095)) + require.NoError(t, err, "add v6 dnat") + + require.NoError(t, m.DisableRouting(), "unpaired disable") + _, v6 := state.Counts() + require.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") +} + // TestNftablesDNAT_DoubleDeleteNoUnderflow verifies that deleting the same rule // twice does not underflow the refcount (the second delete is a no-op). func TestNftablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) { diff --git a/client/firewall/nftables/manager_linux.go b/client/firewall/nftables/manager_linux.go index d25b60585..984b1c3ba 100644 --- a/client/firewall/nftables/manager_linux.go +++ b/client/firewall/nftables/manager_linux.go @@ -530,33 +530,12 @@ func (m *Manager) SetLogLevel(log.Level) { } func (m *Manager) EnableRouting() error { - if err := m.router.ipFwdState.RequestForwarding(false); err != nil { - return fmt.Errorf("enable IPv4 forwarding: %w", err) - } // v6 only when the overlay actually has v6. - if m.router6 == nil { - return nil - } - if err := m.router.ipFwdState.RequestForwarding(true); err != nil { - if rerr := m.router.ipFwdState.ReleaseForwarding(false); rerr != nil { - log.Warnf("rollback v4 forwarding: %v", rerr) - } - return fmt.Errorf("enable IPv6 forwarding: %w", err) - } - return nil + return m.router.ipFwdState.RequestRouting(m.router6 != nil) } func (m *Manager) DisableRouting() error { - var merr *multierror.Error - if err := m.router.ipFwdState.ReleaseForwarding(false); err != nil { - merr = multierror.Append(merr, fmt.Errorf("disable IPv4 forwarding: %w", err)) - } - if m.router6 != nil { - if err := m.router.ipFwdState.ReleaseForwarding(true); err != nil { - merr = multierror.Append(merr, fmt.Errorf("disable IPv6 forwarding: %w", err)) - } - } - return nberrors.FormatErrorOrNil(merr) + return m.router.ipFwdState.ReleaseRouting() } // Flush rule/chain/set operations from the buffer diff --git a/client/firewall/nftables/router_linux.go b/client/firewall/nftables/router_linux.go index cd409559e..d3e031c5f 100644 --- a/client/firewall/nftables/router_linux.go +++ b/client/firewall/nftables/router_linux.go @@ -1836,13 +1836,16 @@ func (r *router) DeleteDNATRule(rule firewall.Rule) error { } } + // Release the refcount only once the rules are gone from the kernel. On + // failure (including the refreshRulesMap error above) the rules and their + // map entries remain, keeping forwarding on until a retry removes them. 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) diff --git a/client/internal/routemanager/ipfwdstate/ipfwdstate.go b/client/internal/routemanager/ipfwdstate/ipfwdstate.go index 5a1bef031..3fdbc90d0 100644 --- a/client/internal/routemanager/ipfwdstate/ipfwdstate.go +++ b/client/internal/routemanager/ipfwdstate/ipfwdstate.go @@ -17,6 +17,13 @@ type IPForwardingState struct { v4Count int v6Count int + // routingV4/routingV6 track whether the routing path currently holds a + // reference, so repeated EnableRouting calls (one per network-map update) + // hold at most one reference per family and an unpaired DisableRouting + // can't release references held by DNAT rules. + routingV4 bool + routingV6 bool + wgIfaceName string v6Saved map[string]int } @@ -33,6 +40,51 @@ func (f *IPForwardingState) Counts() (v4, v6 int) { return f.v4Count, f.v6Count } +// 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. +func (f *IPForwardingState) RequestRouting(v6 bool) error { + f.mu.Lock() + defer f.mu.Unlock() + + if !f.routingV4 { + if err := f.requestV4(); err != nil { + return err + } + f.routingV4 = true + } + + if !v6 || f.routingV6 { + return nil + } + if err := f.requestV6(); err != nil { + log.Warnf("enable IPv6 forwarding for routing: %v", err) + return nil + } + f.routingV6 = true + return nil +} + +// ReleaseRouting releases the references RequestRouting holds. Calls without a +// held reference are no-ops. +func (f *IPForwardingState) ReleaseRouting() error { + f.mu.Lock() + defer f.mu.Unlock() + + if f.routingV4 { + f.routingV4 = false + f.releaseV4() + } + if f.routingV6 { + f.routingV6 = false + return f.releaseV6() + } + return nil +} + // RequestForwarding enables the family's forwarding sysctl on first request. func (f *IPForwardingState) RequestForwarding(v6 bool) error { f.mu.Lock() @@ -84,7 +136,17 @@ func (f *IPForwardingState) requestV6() error { } return fmt.Errorf("enable IPv6 forwarding: %w", err) } - f.v6Saved = saved + // A failed restore on a previous release keeps its saved values; those + // are the true originals, so keep them over what this enable captured. + if f.v6Saved == nil { + f.v6Saved = saved + } else { + for k, v := range saved { + if _, ok := f.v6Saved[k]; !ok { + f.v6Saved[k] = v + } + } + } log.Info("IPv6 forwarding enabled") } f.v6Count++ @@ -100,11 +162,13 @@ func (f *IPForwardingState) releaseV6() error { return nil } - saved := f.v6Saved - f.v6Saved = nil - if err := systemops.DisableV6IPForwarding(saved); err != nil { + // Keep the saved values on failure so a later release or enable/release + // cycle can still restore them; re-restoring an already-restored key is a + // no-op since the sysctl already holds the desired value. + if err := systemops.DisableV6IPForwarding(f.v6Saved); err != nil { return fmt.Errorf("disable IPv6 forwarding: %w", err) } + f.v6Saved = nil log.Info("IPv6 forwarding disabled") return nil } diff --git a/client/internal/routemanager/sysctl/sysctl_linux.go b/client/internal/routemanager/sysctl/sysctl_linux.go index 46b7c9fb7..bb131c691 100644 --- a/client/internal/routemanager/sysctl/sysctl_linux.go +++ b/client/internal/routemanager/sysctl/sysctl_linux.go @@ -58,11 +58,7 @@ func Setup(wgIface iface) (map[string]int, error) { continue } - // Escape '%' and '.' so they survive the dot-to-slash conversion in Set() - safeName := strings.ReplaceAll(intf.Name, "%", percentEscape) - safeName = strings.ReplaceAll(safeName, ".", dotEscape) - - i := fmt.Sprintf(rpFilterInterfacePath, safeName) + i := fmt.Sprintf(rpFilterInterfacePath, EscapeInterfaceName(intf.Name)) oldVal, err := Set(i, 2, true) if err != nil { result = multierror.Append(result, err) @@ -74,6 +70,13 @@ func Setup(wgIface iface) (map[string]int, error) { return keys, nberrors.FormatErrorOrNil(result) } +// EscapeInterfaceName escapes '%' and '.' in an interface name (e.g. VLANs +// like eth0.100) so the name survives the dot-to-slash conversion in Set. +func EscapeInterfaceName(name string) string { + safe := strings.ReplaceAll(name, "%", percentEscape) + return strings.ReplaceAll(safe, ".", dotEscape) +} + // Set sets a sysctl configuration, if onlyIfOne is true it will only set the new value if it's set to 1 func Set(key string, desiredValue int, onlyIfOne bool) (int, error) { path := strings.ReplaceAll(key, ".", "/") diff --git a/client/internal/routemanager/systemops/v6forwarding_linux.go b/client/internal/routemanager/systemops/v6forwarding_linux.go index 14e6b1324..c1e0d4588 100644 --- a/client/internal/routemanager/systemops/v6forwarding_linux.go +++ b/client/internal/routemanager/systemops/v6forwarding_linux.go @@ -19,6 +19,7 @@ const ( // acceptance on regardless, so RA-installed host defaults survive our // v6 forwarding flip. acceptRAInterfacePath = "net.ipv6.conf.%s.accept_ra" + acceptRADefaultPath = "net.ipv6.conf.default.accept_ra" acceptRAProcPathFormat = "/proc/sys/net/ipv6/conf/%s/accept_ra" ) @@ -51,6 +52,10 @@ func DisableV6IPForwarding(saved map[string]int) error { } func bumpAcceptRA(saved map[string]int, wgIfaceName string) { + // Also bump conf.default so interfaces created while forwarding is on + // (hotplug, new Wi-Fi/dock) inherit accept_ra=2 and keep accepting RAs. + bumpAcceptRAKey(saved, acceptRADefaultPath) + interfaces, err := net.Interfaces() if err != nil { log.Warnf("list interfaces for accept_ra: %v", err) @@ -65,18 +70,23 @@ func bumpAcceptRA(saved map[string]int, wgIfaceName string) { } func bumpAcceptRAForInterface(saved map[string]int, name string) { - key := fmt.Sprintf(acceptRAInterfacePath, name) // Build procfs path from name, not the dotted key: VLAN names like eth0.100. if _, err := os.Stat(fmt.Sprintf(acceptRAProcPathFormat, name)); err != nil { return } + bumpAcceptRAKey(saved, fmt.Sprintf(acceptRAInterfacePath, sysctl.EscapeInterfaceName(name))) +} + +func bumpAcceptRAKey(saved map[string]int, key string) { // onlyIfOne=true: leave admin overrides (0, 2) alone. oldVal, err := sysctl.Set(key, 2, true) if err != nil { log.Warnf("bump %s: %v", key, err) return } - if oldVal != 2 { + // With onlyIfOne, a write only happened when the old value was 1; values + // left untouched (0, 2) must not be recorded for restore. + if oldVal == 1 { saved[key] = oldVal } }