diff --git a/client/firewall/iptables/filter_linux.go b/client/firewall/iptables/filter_linux.go index 8955d9541..dc606da2d 100644 --- a/client/firewall/iptables/filter_linux.go +++ b/client/firewall/iptables/filter_linux.go @@ -358,6 +358,14 @@ func (r *family) applyNetwork(flag string, network firewall.Network, prefixes [] } if network.IsSet() { + // A destination set is populated later from DNS results, so unlike a + // source set it cannot be expanded into per-prefix rules. Without + // ipset such a rule is not expressible; report it instead of + // installing something broader than the policy allows. + if flag == "-d" && !r.ipsetSupported { + return nil, fmt.Errorf("destination set %s requires ipset (ip_set_hash_net and xt_set)", network.Set.HashedName()) + } + name := r.ipsetName(network.Set.HashedName()) if _, err := r.ipsetCounter.Increment(name, prefixes); err != nil { return nil, fmt.Errorf("create or get ipset: %w", err) diff --git a/client/firewall/iptables/manager_linux_test.go b/client/firewall/iptables/manager_linux_test.go index 98983aeb6..3f45bf398 100644 --- a/client/firewall/iptables/manager_linux_test.go +++ b/client/firewall/iptables/manager_linux_test.go @@ -5,16 +5,19 @@ package iptables import ( "fmt" "net/netip" + "slices" "strings" "testing" "time" "github.com/coreos/go-iptables/iptables" + "github.com/lrh3321/ipset-go" "github.com/stretchr/testify/require" fw "github.com/netbirdio/netbird/client/firewall/manager" "github.com/netbirdio/netbird/client/iface" "github.com/netbirdio/netbird/client/iface/wgaddr" + "github.com/netbirdio/netbird/shared/management/domain" ) var ifaceMock = &iFaceMock{ @@ -97,10 +100,7 @@ func TestIptablesManager(t *testing.T) { ok, err := ipv4Client.ChainExists("filter", chainACLInput) require.NoError(t, err, "failed check chain exists") - - if ok { - require.NoErrorf(t, err, "chain '%v' still exists after Close", chainACLInput) - } + require.Falsef(t, ok, "chain %q still exists after Close", chainACLInput) }) } @@ -285,14 +285,80 @@ func TestIptablesFilterIPSetFallback(t *testing.T) { // The rule must actually be present in the ACL chain (not silently dropped). checkRuleSpecs(t, ipv4Client, rr.chain, true, fs.specs...) + + // Every expanded peer rule keeps its own redirect-mark pairing. + require.NotNil(t, fs.mangleSpecs, "peer rule must carry a mangle pairing") + checkTableRuleSpecs(t, ipv4Client, tableMangle, chainRTPre, true, fs.mangleSpecs...) } require.NoError(t, manager.DeleteFilterRule(rule), "failed to delete fallback rule") for _, fs := range all { checkRuleSpecs(t, ipv4Client, rr.chain, false, fs.specs...) + checkTableRuleSpecs(t, ipv4Client, tableMangle, chainRTPre, false, fs.mangleSpecs...) } } +// TestIptablesFilterDestinationSetRequiresIPSet documents that a dynamic +// (domain) destination cannot be expressed without ipset: its prefixes are only +// known after DNS resolution, so there is nothing to expand into per-prefix +// rules. The call must report that rather than install a broader rule than the +// policy allows. +func TestIptablesFilterDestinationSetRequiresIPSet(t *testing.T) { + manager, err := Create(ifaceMock, iface.DefaultMTU) + require.NoError(t, err) + require.NoError(t, manager.Init(nil)) + + defer func() { + require.NoError(t, manager.Close(nil)) + }() + + manager.family4.ipsetSupported = false + + destination := fw.Network{Set: fw.NewDomainSet(domain.List{"example.com"})} + + _, err = manager.AddFilterRule(nil, []netip.Prefix{netip.MustParsePrefix("172.16.0.0/16")}, + destination, fw.ProtocolALL, nil, nil, fw.ActionAccept) + require.Error(t, err, "a domain destination is not expressible without ipset") + require.ErrorContains(t, err, "requires ipset") +} + +// TestIptablesNatRuleReAddKeepsSetReferences re-adds the same NAT rule the way +// a repeated network-map update does. The marking rule's set references must not +// grow, or RemoveNatRule can never drop the count to zero and the set stays in +// the kernel for the rest of the process lifetime. +func TestIptablesNatRuleReAddKeepsSetReferences(t *testing.T) { + manager, err := Create(ifaceMock, iface.DefaultMTU) + require.NoError(t, err) + require.NoError(t, manager.Init(nil)) + + defer func() { + require.NoError(t, manager.Close(nil)) + }() + + set := fw.NewDomainSet(domain.List{"example.com"}) + pair := fw.RouterPair{ + ID: "nat-reference-test", + Source: fw.Network{Prefix: netip.MustParsePrefix("100.0.0.0/16")}, + Destination: fw.Network{Set: set}, + Masquerade: true, + Dynamic: true, + } + + require.NoError(t, manager.AddNatRule(pair), "add nat rule") + name := manager.family4.ipsetName(set.HashedName()) + first, ok := manager.family4.ipsetCounter.Get(name) + require.True(t, ok, "the marking rule must hold a reference to its set") + + require.NoError(t, manager.AddNatRule(pair), "re-add nat rule") + second, ok := manager.family4.ipsetCounter.Get(name) + require.True(t, ok, "the set must still be referenced") + require.Equal(t, first.Count, second.Count, "re-adding the same rule must not add references") + + require.NoError(t, manager.RemoveNatRule(pair), "remove nat rule") + _, ok = manager.family4.ipsetCounter.Get(name) + require.False(t, ok, "removing the rule must drop the last reference") +} + // TestIptablesRouteFilterIPSetFallback covers the route ACL side of the // fallback: with a destination set, the expanded per-source rules land // in the route forward chain and are all removed on delete. @@ -340,9 +406,115 @@ func TestIptablesRouteFilterIPSetFallback(t *testing.T) { } } +// TestIptablesCloseRemovesAllState exercises a spread of rule kinds and then +// asserts Close puts every table it touches back exactly as it found it. A +// leaked chain, jump, or ipset survives the daemon and nothing can remove it +// afterwards, since the tracking that knew about it is gone. +func TestIptablesCloseRemovesAllState(t *testing.T) { + ipv4Client, err := iptables.NewWithProtocol(iptables.ProtocolIPv4) + require.NoError(t, err) + + before := snapshotIptables(t, ipv4Client) + + manager, err := Create(ifaceMock, iface.DefaultMTU) + require.NoError(t, err) + require.NoError(t, manager.Init(nil)) + + sources := []netip.Prefix{ + netip.MustParsePrefix("10.20.0.42/32"), + netip.MustParsePrefix("10.20.0.43/32"), + } + + // A multi-source peer rule: shared ipset plus the mangle redirect pairing. + _, err = manager.AddFilterRule(nil, sources, fw.Network{}, "tcp", + nil, &fw.Port{Values: []uint16{22}}, fw.ActionAccept) + require.NoError(t, err, "add peer rule") + + // A route rule with a dynamic destination: a second set, in the forward chain. + _, err = manager.AddFilterRule(nil, sources, + fw.Network{Set: fw.NewDomainSet(domain.List{"example.com"})}, + fw.ProtocolALL, nil, nil, fw.ActionDrop) + require.NoError(t, err, "add route rule") + + // NAT marking for a routed destination, both directions. + pair := fw.RouterPair{ + ID: "cleanup-test", + Source: fw.Network{Prefix: netip.MustParsePrefix("100.0.0.0/16")}, + Destination: fw.Network{Prefix: netip.MustParsePrefix("192.168.55.0/24")}, + Masquerade: true, + } + require.NoError(t, manager.AddNatRule(pair), "add nat rule") + require.NoError(t, manager.EnableRouting(), "enable routing") + + // A DNAT redirect, which also holds a forwarding reference. + dnat := fw.ForwardRule{ + Protocol: fw.ProtocolTCP, + DestinationPort: fw.Port{Values: []uint16{8080}}, + TranslatedAddress: netip.MustParseAddr("10.20.0.44"), + TranslatedPort: fw.Port{Values: []uint16{80}}, + } + dnatRule, err := manager.AddDNATRule(dnat) + require.NoError(t, err, "add dnat rule") + + require.NotEqual(t, before, snapshotIptables(t, ipv4Client), "the manager must have installed state") + + require.NoError(t, manager.DeleteDNATRule(dnatRule), "delete dnat rule") + require.NoError(t, manager.DisableRouting(), "disable routing") + require.NoError(t, manager.Close(nil), "close") + + after := snapshotIptables(t, ipv4Client) + require.Equal(t, before.chains, after.chains, "Close must remove every chain it created") + require.Equal(t, before.rules, after.rules, "Close must remove every rule it created") + require.Equal(t, before.sets, after.sets, "Close must destroy every ipset it created") +} + +// iptablesState is a snapshot of the tables the manager writes to, used to +// compare the kernel before and after a manager lifetime. +type iptablesState struct { + chains map[string][]string + rules map[string][]string + sets []string +} + +func snapshotIptables(t *testing.T, client *iptables.IPTables) iptablesState { + t.Helper() + + state := iptablesState{ + chains: map[string][]string{}, + rules: map[string][]string{}, + } + + for _, table := range []string{tableFilter, tableNat, tableMangle, tableRaw} { + chains, err := client.ListChains(table) + require.NoErrorf(t, err, "list chains in %s", table) + slices.Sort(chains) + state.chains[table] = chains + + for _, chain := range chains { + rules, err := client.List(table, chain) + require.NoErrorf(t, err, "list rules in %s/%s", table, chain) + state.rules[table+"/"+chain] = rules + } + } + + sets, err := ipset.ListAll() + require.NoError(t, err, "list ipsets") + for _, set := range sets { + state.sets = append(state.sets, set.SetName) + } + slices.Sort(state.sets) + + return state +} + func checkRuleSpecs(t *testing.T, ipv4Client *iptables.IPTables, chainName string, mustExists bool, rulespec ...string) { t.Helper() - exists, err := ipv4Client.Exists("filter", chainName, rulespec...) + checkTableRuleSpecs(t, ipv4Client, tableFilter, chainName, mustExists, rulespec...) +} + +func checkTableRuleSpecs(t *testing.T, ipv4Client *iptables.IPTables, table, chainName string, mustExists bool, rulespec ...string) { + t.Helper() + exists, err := ipv4Client.Exists(table, chainName, rulespec...) require.NoError(t, err, "failed to check rule") require.Falsef(t, !exists && mustExists, "rule '%v' does not exist", rulespec) require.Falsef(t, exists && !mustExists, "rule '%v' exist", rulespec) diff --git a/client/firewall/iptables/routing_linux.go b/client/firewall/iptables/routing_linux.go index 1d6fdb14a..80403a9ca 100644 --- a/client/firewall/iptables/routing_linux.go +++ b/client/firewall/iptables/routing_linux.go @@ -192,14 +192,23 @@ func (r *family) insertEstablishedRule(chain string) error { return nil } -func (r *family) addNatRule(pair firewall.RouterPair) error { +func (r *family) addNatRule(pair firewall.RouterPair) (err error) { ruleID := pair.GenKey(firewall.NatFormat) if rule, exists := r.rules[ruleID]; exists { - if err := r.iptablesClient.DeleteIfExists(tableMangle, chainRTPre, rule...); err != nil { - return fmt.Errorf("remove existing marking rule for %s: %w", pair.Destination, err) + if derr := r.iptablesClient.DeleteIfExists(tableMangle, chainRTPre, rule...); derr != nil { + return fmt.Errorf("remove existing marking rule for %s: %w", pair.Destination, derr) } delete(r.rules, ruleID) + + // Drop the replaced spec's set references only once the new spec has + // taken its own, so a set both specs share is not destroyed and + // recreated, which would lose the prefixes UpdateSet put in it. + defer func() { + if derr := r.decrementSetCounter(rule); derr != nil && err == nil { + err = fmt.Errorf("decrement ipset counter: %w", derr) + } + }() } markValue := nbnet.PreroutingFwmarkMasquerade diff --git a/client/firewall/nftables/filter_linux.go b/client/firewall/nftables/filter_linux.go index fa1ef1484..ebd238063 100644 --- a/client/firewall/nftables/filter_linux.go +++ b/client/firewall/nftables/filter_linux.go @@ -236,6 +236,12 @@ func (r *family) DeleteFilterRule(rule firewall.Rule) error { if pr.nftRule.Handle == 0 { log.Warnf("filter rule %s has no handle, removing stale entry", ruleID) + // The paired mangle rule can still be in the kernel with a live + // handle. Dropping the tracking entry without removing it would + // leave a prerouting rule that nothing can find again. + if err := r.deleteMangleRule(pr, ruleID); err != nil { + return err + } r.dropNetworkMatch(pr.nftRule.Exprs) delete(r.filters, ruleID) return nil @@ -244,11 +250,7 @@ func (r *family) DeleteFilterRule(rule firewall.Rule) error { if err := r.conn.DelRule(pr.nftRule); err != nil { log.Errorf("queue rule delete: %v", err) } - if pr.mangleRule != nil { - if err := r.conn.DelRule(pr.mangleRule); err != nil { - log.Errorf("queue mangle rule delete: %v", err) - } - } + r.queueMangleDelete(pr) if err := r.conn.Flush(); err != nil { return fmt.Errorf("flush delete %s: %w", ruleID, err) } @@ -258,6 +260,32 @@ func (r *family) DeleteFilterRule(rule firewall.Rule) error { return nil } +// deleteMangleRule removes the prerouting rule paired with a filter rule on +// its own, for the paths that drop the filter rule's tracking without queueing +// a delete for it. +func (r *family) deleteMangleRule(pr *Rule, ruleID firewall.RuleID) error { + if pr.mangleRule == nil || pr.mangleRule.Handle == 0 { + return nil + } + + r.queueMangleDelete(pr) + if err := r.conn.Flush(); err != nil { + return fmt.Errorf("flush mangle delete %s: %w", ruleID, err) + } + return nil +} + +// queueMangleDelete queues the delete of the rule's prerouting counterpart, if +// it has one. The caller commits it. +func (r *family) queueMangleDelete(pr *Rule) { + if pr.mangleRule == nil { + return + } + if err := r.conn.DelRule(pr.mangleRule); err != nil { + log.Errorf("queue mangle rule delete: %v", err) + } +} + func (r *family) decrementSetCounter(rule *nftables.Rule) error { if r.ipsetCounter == nil { return nil diff --git a/client/firewall/nftables/router_linux_test.go b/client/firewall/nftables/router_linux_test.go index e38ec846d..49dbfc8f2 100644 --- a/client/firewall/nftables/router_linux_test.go +++ b/client/firewall/nftables/router_linux_test.go @@ -540,6 +540,23 @@ func TestNftablesUpdateSetMergesOverlapping(t *testing.T) { netip.MustParsePrefix("192.168.1.1/32"), } require.NoError(t, r.UpdateSet(set, overlapping), "UpdateSet must merge overlapping prefixes") + + fetchedSet, err := r.conn.GetSetByName(r.workTable, set.HashedName()) + require.NoError(t, err, "fetch updated set") + elements, err := r.conn.GetSetElements(fetchedSet) + require.NoError(t, err, "get set elements") + + starts := make(map[string]bool) + for _, elem := range elements { + if elem.IntervalEnd { + continue + } + starts[netip.AddrFrom4(*(*[4]byte)(elem.Key)).String()] = true + } + // The /32s are covered by the /24, so the update adds one interval and + // leaves the one created earlier in place. + assert.Equal(t, map[string]bool{"10.0.0.0": true, "192.168.1.0": true}, starts, + "merged set must hold the original and the merged interval") } func TestNftablesCreateIpSet_IPv6(t *testing.T) { diff --git a/client/firewall/nftables/routing_linux.go b/client/firewall/nftables/routing_linux.go index 00d5f86f0..fadb4edea 100644 --- a/client/firewall/nftables/routing_linux.go +++ b/client/firewall/nftables/routing_linux.go @@ -23,26 +23,47 @@ func (r *family) AddNatRule(pair firewall.RouterPair) error { return fmt.Errorf(refreshRulesMapError, err) } + // Resolve every rule's match expressions before queueing any of them: a + // message buffered on the shared connection cannot be un-queued, so + // returning an error after queueing would leave the next caller's Flush + // to commit a rule nothing tracks. + var legacyExprs []expr.Any if r.legacyManagement { log.Warnf("This peer is connected to a NetBird Management service with an older version. Allowing all traffic for %s", pair.Destination) - if err := r.addLegacyRouteRule(pair); err != nil { - r.rollbackRules(pair) - return fmt.Errorf("add legacy routing rule: %w", err) + + var err error + legacyExprs, err = r.legacyRouteRuleExprs(pair) + if err != nil { + return fmt.Errorf("build legacy routing rule: %w", err) } } + inverse := firewall.GetInversePair(pair) + var natExprs, inverseExprs []expr.Any if pair.Masquerade { - if err := r.addNatRule(pair); err != nil { - r.rollbackRules(pair) - return fmt.Errorf("add nat rule: %w", err) + var err error + natExprs, err = r.natRuleExprs(pair) + if err != nil { + r.dropNetworkMatch(legacyExprs) + return fmt.Errorf("build nat rule: %w", err) } - if err := r.addNatRule(firewall.GetInversePair(pair)); err != nil { - r.rollbackRules(pair) - return fmt.Errorf("add inverse nat rule: %w", err) + inverseExprs, err = r.natRuleExprs(inverse) + if err != nil { + r.dropNetworkMatch(legacyExprs) + r.dropNetworkMatch(natExprs) + return fmt.Errorf("build inverse nat rule: %w", err) } } + if legacyExprs != nil { + r.queueLegacyRouteRule(pair, legacyExprs) + } + if pair.Masquerade { + r.queueNatRule(pair, natExprs) + r.queueNatRule(inverse, inverseExprs) + } + if err := r.conn.Flush(); err != nil { r.rollbackRules(pair) return fmt.Errorf("insert rules for %s: %w", pair.Destination, err) @@ -70,17 +91,19 @@ func (r *family) rollbackRules(pair firewall.RouterPair) { } } -// addNatRule inserts a nftables rule to the conn client flush queue -func (r *family) addNatRule(pair firewall.RouterPair) error { +// natRuleExprs resolves the match expressions of the pair's prerouting +// marking rule. It reserves the ipset references the matches need but queues +// nothing on the connection, so its error paths leave the connection clean. +func (r *family) natRuleExprs(pair firewall.RouterPair) ([]expr.Any, error) { sourceExp, err := r.applyNetwork(pair.Source, nil, true) if err != nil { - return fmt.Errorf("apply source: %w", err) + return nil, fmt.Errorf("apply source: %w", err) } destExp, err := r.applyNetwork(pair.Destination, nil, false) if err != nil { r.dropNetworkMatch(sourceExp) - return fmt.Errorf("apply destination: %w", err) + return nil, fmt.Errorf("apply destination: %w", err) } op := expr.CmpOpEq @@ -123,13 +146,19 @@ func (r *family) addNatRule(pair firewall.RouterPair) error { }, ) + return exprs, nil +} + +// queueNatRule replaces any tracked rule for the pair and queues the new +// prerouting marking rule on the connection. Failures are logged rather than +// returned: the caller has already queued messages that only a Flush can +// commit, so it must not return early. +func (r *family) queueNatRule(pair firewall.RouterPair, exprs []expr.Any) { ruleID := pair.GenKey(firewall.PreroutingFormat) if _, exists := r.rules[ruleID]; exists { if err := r.removeNatRule(pair); err != nil { - r.dropNetworkMatch(sourceExp) - r.dropNetworkMatch(destExp) - return fmt.Errorf("remove prerouting rule: %w", err) + log.Errorf("replace prerouting rule %s: %v", ruleID, err) } } @@ -141,8 +170,6 @@ func (r *family) addNatRule(pair firewall.RouterPair) error { Exprs: exprs, UserData: []byte(ruleID), }) - - return nil } func (r *family) addPostroutingRules() { @@ -308,27 +335,32 @@ func buildLegacyRouteRuleExpressions(sourceExp, destExp []expr.Any) []expr.Any { return exprs } -func (r *family) addLegacyRouteRule(pair firewall.RouterPair) error { +// legacyRouteRuleExprs resolves the match expressions of the pair's legacy +// forwarding rule, queueing nothing on the connection. +func (r *family) legacyRouteRuleExprs(pair firewall.RouterPair) ([]expr.Any, error) { sourceExp, err := r.applyNetwork(pair.Source, nil, true) if err != nil { - return fmt.Errorf("apply source: %w", err) + return nil, fmt.Errorf("apply source: %w", err) } destExp, err := r.applyNetwork(pair.Destination, nil, false) if err != nil { r.dropNetworkMatch(sourceExp) - return fmt.Errorf("apply destination: %w", err) + return nil, fmt.Errorf("apply destination: %w", err) } - exprs := buildLegacyRouteRuleExpressions(sourceExp, destExp) + return buildLegacyRouteRuleExpressions(sourceExp, destExp), nil +} +// queueLegacyRouteRule replaces any tracked rule for the pair and queues the +// new legacy forwarding rule. Failures are logged for the same reason as in +// queueNatRule. +func (r *family) queueLegacyRouteRule(pair firewall.RouterPair, exprs []expr.Any) { ruleID := pair.GenKey(firewall.ForwardingFormat) if _, exists := r.rules[ruleID]; exists { if err := r.removeLegacyRouteRule(pair); err != nil { - r.dropNetworkMatch(sourceExp) - r.dropNetworkMatch(destExp) - return fmt.Errorf("remove legacy routing rule: %w", err) + log.Errorf("replace legacy forwarding rule %s: %v", ruleID, err) } } @@ -338,7 +370,6 @@ func (r *family) addLegacyRouteRule(pair firewall.RouterPair) error { Exprs: exprs, UserData: []byte(ruleID), }) - return nil } // removeLegacyRouteRule removes a legacy routing rule for mgmt servers pre route acls @@ -395,14 +426,27 @@ func (r *family) RemoveAllLegacyRouteRules() error { } var merr *multierror.Error + var found bool for k, rule := range r.rules { if !strings.HasPrefix(string(k), firewall.ForwardingFormatPrefix) { continue } + found = true if err := r.deleteLegacyRuleEntry(k, rule); err != nil { merr = multierror.Append(merr, err) } } + + // Commit the queued deletes here instead of leaving them for whichever + // caller flushes next: the tracking entries are already gone, so an + // uncommitted delete would leave a rule in the kernel that nothing can + // find again. + if found { + if err := r.conn.Flush(); err != nil { + merr = multierror.Append(merr, fmt.Errorf(flushError, err)) + } + } + return nberrors.FormatErrorOrNil(merr) } diff --git a/client/firewall/uspfilter/peer_family_scope_test.go b/client/firewall/uspfilter/peer_family_scope_test.go index 13c3a616e..1cf3498fb 100644 --- a/client/firewall/uspfilter/peer_family_scope_test.go +++ b/client/firewall/uspfilter/peer_family_scope_test.go @@ -101,6 +101,4 @@ func TestRouteACL_MixedFamilyZeroSourcesStayFamilySafe(t *testing.T) { assert.True(t, pass, "v4 source must match the v4 destination rule via 0.0.0.0/0") _, pass = m.routeACLsPass(v6Src, netip.MustParseAddr("fd00:1::5"), 255, 0, 0) assert.True(t, pass, "v6 source must match the v6 destination rule via ::/0") - _, pass = m.routeACLsPass(v6Src, netip.MustParseAddr("10.0.0.5"), 255, 0, 0) - assert.True(t, pass, "v6 source still passes the v4 destination rule via ::/0 in the same source list") }