diff --git a/client/iface/configurer/allowedips.go b/client/iface/configurer/allowedips.go new file mode 100644 index 000000000..193197d4a --- /dev/null +++ b/client/iface/configurer/allowedips.go @@ -0,0 +1,226 @@ +package configurer + +import ( + "net" + "net/netip" + "slices" + "sync" + + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) + +// allowedIPStore mirrors the allowed IPs configured on each peer of a device. +// +// A configurer is the only writer of its device's peer set, so the mirror is authoritative +// by construction. It spares the paths that have to rewrite one peer's allowed IPs a full +// device dump just to recover prefixes the process already configured itself. +// +// An allowed IP belongs to exactly one peer: configuring a prefix on a peer takes it away +// from whichever peer held it before, and the configurer leaves that handover to the device +// rather than removing the prefix from the previous holder itself. The store tracks the +// owner of each prefix and performs the same handover, so rewriting one peer's list never +// takes a prefix back from the peer that owns it now. +// +// Its own lock guards the map alone, not the device write it accompanies. Consistency +// between the two rests on the caller serializing every configurer call, which WGIface +// does with its mutex; two unserialized writers would interleave a device write with the +// record of a different one. +// +// An operator reconfiguring the device out of band, through `wg set` or the UAPI socket, +// is the one way the mirror can still go stale. A peer missing from it falls back to the +// device, which reseats that peer's prefixes and their ownership; a peer that is present +// does not, so one recorded from empty while the device already held prefixes keeps only +// what was recorded, and the next endpoint removal drops the rest. +type allowedIPStore struct { + mu sync.RWMutex + peers map[wgtypes.Key][]netip.Prefix + owners map[netip.Prefix]wgtypes.Key +} + +func newAllowedIPStore() *allowedIPStore { + return &allowedIPStore{ + peers: make(map[wgtypes.Key][]netip.Prefix), + owners: make(map[netip.Prefix]wgtypes.Key), + } +} + +// get returns the prefixes recorded for a peer, and whether the peer is known at all. +// The caller receives a copy and may retain or modify it freely. +func (s *allowedIPStore) get(key wgtypes.Key) ([]netip.Prefix, bool) { + s.mu.RLock() + defer s.mu.RUnlock() + + prefixes, ok := s.peers[key] + if !ok { + return nil, false + } + return slices.Clone(prefixes), true +} + +// set replaces the prefixes recorded for a peer. +func (s *allowedIPStore) set(key wgtypes.Key, prefixes []netip.Prefix) { + s.mu.Lock() + defer s.mu.Unlock() + + k := key + s.releaseLocked(k) + + normalized := normalizePrefixes(prefixes) + for _, prefix := range normalized { + s.claimLocked(k, prefix) + } + s.peers[k] = normalized +} + +// add records prefixes on a peer without dropping the ones already there, matching the +// union semantics of a peer update that does not replace its allowed IPs. It records the +// peer if it is not known yet, so it belongs to the operations that create a peer on the +// device rather than to the update-only ones. +func (s *allowedIPStore) add(key wgtypes.Key, prefixes []netip.Prefix) { + s.mu.Lock() + defer s.mu.Unlock() + + s.mergeLocked(key, prefixes) +} + +// addExisting is add for an update-only device operation. Such an operation is a silent +// no-op when the peer is absent, so recording a peer here would leave the store claiming +// prefixes the device never took, and the peer would then be recreated by the next endpoint +// removal, stealing those allowed IPs from the peer that legitimately holds them. +func (s *allowedIPStore) addExisting(key wgtypes.Key, prefixes []netip.Prefix) { + s.mu.Lock() + defer s.mu.Unlock() + + k := key + if _, ok := s.peers[k]; !ok { + return + } + s.mergeLocked(k, prefixes) +} + +// ensure records a peer with no prefixes unless it is already known. A device operation +// that is not update-only creates the peer when it is absent, so it has to be recorded even +// when it configures nothing else; otherwise the peer exists on the device while the store +// treats it as unknown, and a prefix later handed over to it is not accounted for. +func (s *allowedIPStore) ensure(key wgtypes.Key) { + s.mu.Lock() + defer s.mu.Unlock() + + k := key + if _, ok := s.peers[k]; !ok { + s.peers[k] = nil + } +} + +// forget drops every prefix recorded for a peer. +func (s *allowedIPStore) forget(key wgtypes.Key) { + s.mu.Lock() + defer s.mu.Unlock() + + k := key + s.releaseLocked(k) + delete(s.peers, k) +} + +// reset drops every peer, mirroring a device reconfiguration that replaces the peer set. +func (s *allowedIPStore) reset() { + s.mu.Lock() + defer s.mu.Unlock() + + s.peers = make(map[wgtypes.Key][]netip.Prefix) + s.owners = make(map[netip.Prefix]wgtypes.Key) +} + +// mergeLocked unions normalized prefixes into a peer and transfers their ownership. +// The caller must hold s.mu for writing. +func (s *allowedIPStore) mergeLocked(k wgtypes.Key, prefixes []netip.Prefix) { + merged := s.peers[k] + for _, prefix := range prefixes { + prefix = normalizePrefix(prefix) + s.claimLocked(k, prefix) + if !slices.Contains(merged, prefix) { + merged = append(merged, prefix) + } + } + s.peers[k] = merged +} + +// claimLocked hands a prefix over to a peer, taking it from its previous owner the way the +// device does when the same prefix is configured on a second peer. +func (s *allowedIPStore) claimLocked(k wgtypes.Key, prefix netip.Prefix) { + if owner, ok := s.owners[prefix]; ok && owner != k { + s.peers[owner] = slices.DeleteFunc(s.peers[owner], func(p netip.Prefix) bool { + return p == prefix + }) + } + s.owners[prefix] = k +} + +// releaseLocked drops a peer's claim on every prefix it currently holds. +func (s *allowedIPStore) releaseLocked(k wgtypes.Key) { + for _, prefix := range s.peers[k] { + if s.owners[prefix] == k { + delete(s.owners, prefix) + } + } +} + +// normalizePrefix puts a prefix into the form the store recognises it by. It clears the +// host bits, which a device does on its own, so a caller passing 10.20.0.1/16 still matches +// the 10.20.0.0/16 read back from the device; and it unmaps a v4-mapped prefix so that it +// compares equal to, and marshals like, the plain v4 prefix for the same network. +// +// Masking comes first because it also decides the address family: only a prefix at least 96 +// bits long keeps the mapped marker through the mask, so a shorter prefix inside the mapped +// range is a genuine v6 prefix and unmapping it would yield an invalid v4 prefix. +func normalizePrefix(prefix netip.Prefix) netip.Prefix { + masked := prefix.Masked() + + addr := masked.Addr() + if !addr.Is4In6() { + return masked + } + return netip.PrefixFrom(addr.Unmap(), masked.Bits()-96) +} + +// normalizePrefixes returns a normalized copy without changing the caller's slice. +func normalizePrefixes(prefixes []netip.Prefix) []netip.Prefix { + normalized := make([]netip.Prefix, len(prefixes)) + for i, prefix := range prefixes { + normalized[i] = normalizePrefix(prefix) + } + return normalized +} + +// ipNetsToPrefixes converts addresses read back from a device. Unmap keeps a v4-mapped v6 +// address comparable to the plain v4 prefix the configurer was given. +func ipNetsToPrefixes(ipNets []net.IPNet) []netip.Prefix { + prefixes := make([]netip.Prefix, 0, len(ipNets)) + for _, ipNet := range ipNets { + addr, ok := netip.AddrFromSlice(ipNet.IP) + if !ok { + continue + } + + ones, maskBits := ipNet.Mask.Size() + // A device may report a v4 prefix as a v4-mapped address. Align the address form with + // the mask rather than unmapping on sight: a 32 bit mask always describes v4, while a + // 128 bit mask describes v4 only when it covers the mapped prefix, so a genuine v6 + // prefix inside the mapped range stays v6 instead of being dropped as invalid. + if addr.Is4In6() { + switch { + case maskBits == 32: + addr = addr.Unmap() + case maskBits == 128 && ones >= 96: + addr, ones = addr.Unmap(), ones-96 + } + } + + prefix := netip.PrefixFrom(addr, ones) + if !prefix.IsValid() { + continue + } + prefixes = append(prefixes, prefix.Masked()) + } + return prefixes +} diff --git a/client/iface/configurer/allowedips_test.go b/client/iface/configurer/allowedips_test.go new file mode 100644 index 000000000..1d272d8c4 --- /dev/null +++ b/client/iface/configurer/allowedips_test.go @@ -0,0 +1,263 @@ +package configurer + +import ( + "net" + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) + +// The store keys on the parsed key, so the tests use two distinct ones rather than names. +var ( + testPeer = wgtypes.Key{1} + otherPeer = wgtypes.Key{2} +) + +func TestAllowedIPStoreUnknownPeer(t *testing.T) { + s := newAllowedIPStore() + + prefixes, ok := s.get(testPeer) + assert.False(t, ok, "an unconfigured peer must be reported as unknown, not as one without prefixes") + assert.Nil(t, prefixes, "an unknown peer has no prefixes") +} + +func TestAllowedIPStoreAddUnions(t *testing.T) { + s := newAllowedIPStore() + overlay := netip.MustParsePrefix("100.64.0.1/32") + routed := netip.MustParsePrefix("10.20.0.0/16") + + s.set(testPeer, []netip.Prefix{overlay}) + // A peer update does not replace allowed IPs, and a repeated prefix must not be doubled. + s.add(testPeer, []netip.Prefix{overlay, routed}) + + prefixes, ok := s.get(testPeer) + require.True(t, ok, "peer must be known after set") + assert.Equal(t, []netip.Prefix{overlay, routed}, prefixes, "add must union rather than replace") +} + +func TestAllowedIPStoreGetReturnsCopy(t *testing.T) { + s := newAllowedIPStore() + overlay := netip.MustParsePrefix("100.64.0.1/32") + s.set(testPeer, []netip.Prefix{overlay}) + + prefixes, ok := s.get(testPeer) + require.True(t, ok, "peer must be known after set") + prefixes[0] = netip.MustParsePrefix("0.0.0.0/0") + + stored, _ := s.get(testPeer) + assert.Equal(t, []netip.Prefix{overlay}, stored, "a caller mutating the returned slice must not corrupt the store") +} + +func TestAllowedIPStoreForgetAndReset(t *testing.T) { + s := newAllowedIPStore() + s.set(testPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32")}) + s.set(otherPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.2/32")}) + + s.forget(testPeer) + _, ok := s.get(testPeer) + assert.False(t, ok, "a forgotten peer must be unknown") + _, ok = s.get(otherPeer) + assert.True(t, ok, "forgetting one peer must not touch the others") + + s.reset() + _, ok = s.get(otherPeer) + assert.False(t, ok, "reset must drop every peer") +} + +func TestIPNetsToPrefixes(t *testing.T) { + tests := []struct { + name string + ipNet net.IPNet + want string + }{ + { + name: "v4", + ipNet: net.IPNet{IP: net.IP{10, 20, 0, 0}, Mask: net.CIDRMask(16, 32)}, + want: "10.20.0.0/16", + }, + { + name: "v4 mapped under a 128 bit mask", + ipNet: net.IPNet{IP: net.ParseIP("10.20.0.0"), Mask: net.CIDRMask(112, 128)}, + want: "10.20.0.0/16", + }, + { + name: "v6", + ipNet: net.IPNet{IP: net.ParseIP("fd00::"), Mask: net.CIDRMask(64, 128)}, + want: "fd00::/64", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := ipNetsToPrefixes([]net.IPNet{tc.ipNet}) + require.Len(t, got, 1, "the address must be converted, not dropped") + assert.Equal(t, tc.want, got[0].String(), "converted prefix") + }) + } +} + +func TestIPNetsToPrefixesRoundTrip(t *testing.T) { + prefixes := []netip.Prefix{ + netip.MustParsePrefix("100.64.0.1/32"), + netip.MustParsePrefix("10.20.0.0/16"), + netip.MustParsePrefix("fd00::/64"), + } + + assert.Equal(t, prefixes, ipNetsToPrefixes(prefixesToIPNets(prefixes)), + "prefixes handed to a device must come back unchanged") +} + +func TestAllowedIPStoreNormalizesMappedPrefixes(t *testing.T) { + s := newAllowedIPStore() + v4 := netip.MustParsePrefix("10.20.0.0/16") + mapped := netip.PrefixFrom(netip.AddrFrom16(v4.Addr().As16()), 112) + + s.set(testPeer, []netip.Prefix{mapped}) + // A v4 rule only matches a v4-mapped address once it has been unmapped, so the store must + // hold the plain form and recognise the two spellings as the same prefix. + s.add(testPeer, []netip.Prefix{v4}) + + prefixes, ok := s.get(testPeer) + require.True(t, ok, "peer must be known after set") + assert.Equal(t, []netip.Prefix{v4}, prefixes, "a mapped prefix must be stored unmapped and not duplicated") +} + +func TestNormalizePrefix(t *testing.T) { + v4 := netip.MustParsePrefix("10.20.0.0/16") + v6 := netip.MustParsePrefix("fd00::/64") + + assert.Equal(t, v4, normalizePrefix(v4), "a plain v4 prefix is unchanged") + assert.Equal(t, v6, normalizePrefix(v6), "a real v6 prefix is unchanged") + assert.Equal(t, v4, normalizePrefix(netip.PrefixFrom(netip.AddrFrom16(v4.Addr().As16()), 112)), + "a mapped prefix under a 128 bit mask becomes plain v4") + // A prefix shorter than /96 inside the mapped range is a genuine v6 prefix. Unmapping it + // would pair a v4 address with a v6 sized mask, which is invalid, and the store would then + // record a zero prefix that can never recreate the allowed IP. + for _, tc := range []string{"::ffff:0:0/64", "::ffff:1.2.3.4/80", "::ffff:1.2.3.4/95"} { + got := normalizePrefix(netip.MustParsePrefix(tc)) + assert.True(t, got.IsValid(), "%s must normalize to a valid prefix", tc) + assert.False(t, got.Addr().Is4(), "%s must stay v6", tc) + } +} + +func TestAllowedIPStoreAddExistingDoesNotCreate(t *testing.T) { + s := newAllowedIPStore() + routed := netip.MustParsePrefix("10.20.0.0/16") + + // An update-only device operation on an absent peer is a silent no-op, so nothing may be + // recorded for a peer the store does not already know. + s.addExisting(testPeer, []netip.Prefix{routed}) + _, ok := s.get(testPeer) + assert.False(t, ok, "addExisting must not record an unknown peer") + + overlay := netip.MustParsePrefix("100.64.0.1/32") + s.set(testPeer, []netip.Prefix{overlay}) + s.addExisting(testPeer, []netip.Prefix{routed}) + + prefixes, _ := s.get(testPeer) + assert.Equal(t, []netip.Prefix{overlay, routed}, prefixes, "addExisting must union onto a known peer") +} + +func TestAllowedIPStoreHandsPrefixOverToTheNewOwner(t *testing.T) { + s := newAllowedIPStore() + routed := netip.MustParsePrefix("10.20.0.0/16") + other := otherPeer + + s.set(testPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32"), routed}) + s.set(other, []netip.Prefix{netip.MustParsePrefix("100.64.0.2/32")}) + + // The device takes an allowed IP away from its previous holder when it is configured on + // another peer, so the store must do the same rather than list it under both. + s.addExisting(other, []netip.Prefix{routed}) + + previous, _ := s.get(testPeer) + assert.NotContains(t, previous, routed, "the previous owner must lose the prefix") + current, _ := s.get(other) + assert.Contains(t, current, routed, "the new owner must hold the prefix") +} + +func TestAllowedIPStoreForgetReleasesOwnership(t *testing.T) { + s := newAllowedIPStore() + routed := netip.MustParsePrefix("10.20.0.0/16") + + s.set(testPeer, []netip.Prefix{routed}) + s.forget(testPeer) + s.set(otherPeer, []netip.Prefix{routed}) + + // A forgotten peer must not be resurrected as a key in the peer map by a later claim. + _, ok := s.get(testPeer) + assert.False(t, ok, "the forgotten peer must stay unknown") + current, _ := s.get(otherPeer) + assert.Equal(t, []netip.Prefix{routed}, current, "the new owner must hold the prefix") +} + +func TestNormalizePrefixClearsHostBits(t *testing.T) { + // A device stores a prefix masked, so a caller passing host bits must still match what a + // device fallback seeded, otherwise that prefix could never be removed by value. + assert.Equal(t, netip.MustParsePrefix("10.20.0.0/16"), + normalizePrefix(netip.MustParsePrefix("10.20.0.1/16")), "host bits must be cleared") + assert.Equal(t, netip.MustParsePrefix("fd00::/64"), + normalizePrefix(netip.MustParsePrefix("fd00::1/64")), "host bits must be cleared for v6") +} + +func TestIPNetsToPrefixesKeepsV6InTheMappedRange(t *testing.T) { + // ::ffff:0:0/64 reads as v4-mapped but is a genuine v6 prefix: unmapping it would leave a + // v4 address under a 64 bit mask, which is invalid, and the allowed IP would be dropped. + got := ipNetsToPrefixes([]net.IPNet{{ + IP: net.ParseIP("::ffff:0:0"), + Mask: net.CIDRMask(64, 128), + }}) + + require.Len(t, got, 1, "the prefix must be converted, not dropped") + assert.False(t, got[0].Addr().Is4(), "a v6 prefix in the mapped range must not become v4") + assert.Equal(t, 64, got[0].Bits(), "the prefix length must survive the conversion") +} + +func TestPrefixesToIPNetsNormalizes(t *testing.T) { + // net.IPNet prints a v4-mapped address as v4 but takes the length from its 16 byte + // mask, so an unnormalized ::ffff:10.1.2.3/64 reaches a userspace device as 10.1.2.3/0, + // an allowed IP that matches every v4 address. + tests := []struct { + name string + given string + want string + }{ + {name: "mapped below /96", given: "::ffff:10.1.2.3/64", want: "::/64"}, + {name: "mapped at /112", given: "::ffff:10.1.2.3/112", want: "10.1.0.0/16"}, + {name: "host bits are cleared", given: "10.20.0.1/16", want: "10.20.0.0/16"}, + {name: "v6 is untouched", given: "fd00::1/64", want: "fd00::/64"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := prefixesToIPNets([]netip.Prefix{netip.MustParsePrefix(tc.given)}) + require.Len(t, got, 1, "the prefix must be converted, not dropped") + assert.Equal(t, tc.want, got[0].String(), "what the device is given") + assert.NotEqual(t, 0, mustOnes(t, got[0]), "a device must never be given a zero length allowed IP") + }) + } +} + +func mustOnes(t *testing.T, ipNet net.IPNet) int { + t.Helper() + + ones, _ := ipNet.Mask.Size() + return ones +} + +// TestPrefixesToIPNetsAgreesWithTheStore pins the property the store depends on: what a +// device is given and what is recorded for it are the same prefix. +func TestPrefixesToIPNetsAgreesWithTheStore(t *testing.T) { + for _, given := range []string{"::ffff:10.1.2.3/64", "::ffff:10.1.2.3/112", "10.20.0.1/16", "fd00::1/64"} { + prefix := netip.MustParsePrefix(given) + + toDevice := prefixesToIPNets([]netip.Prefix{prefix}) + recorded := normalizePrefix(prefix) + + assert.Equal(t, recorded.String(), toDevice[0].String(), + "%s must reach the device in the form the store records", given) + } +} diff --git a/client/iface/configurer/common.go b/client/iface/configurer/common.go index 10162d703..40f8209e9 100644 --- a/client/iface/configurer/common.go +++ b/client/iface/configurer/common.go @@ -19,12 +19,18 @@ func buildPresharedKeyConfig(peerKey wgtypes.Key, psk wgtypes.Key, updateOnly bo } } +// prefixesToIPNets converts prefixes on their way to a device. It is the only place that +// conversion happens, so it also normalizes: the device is then given the same form the +// store records, and a v4-mapped prefix cannot reach net.IPNet, which prints such an +// address as v4 while taking the length from its 16 byte mask and so turns +// ::ffff:10.1.2.3/64 into 10.1.2.3/0 — an allowed IP matching every v4 address. func prefixesToIPNets(prefixes []netip.Prefix) []net.IPNet { ipNets := make([]net.IPNet, len(prefixes)) for i, prefix := range prefixes { + normalized := normalizePrefix(prefix) ipNets[i] = net.IPNet{ - IP: prefix.Addr().AsSlice(), // Convert netip.Addr to net.IP - Mask: net.CIDRMask(prefix.Bits(), prefix.Addr().BitLen()), // Create subnet mask + IP: normalized.Addr().AsSlice(), + Mask: net.CIDRMask(normalized.Bits(), normalized.Addr().BitLen()), } } return ipNets diff --git a/client/iface/configurer/kernel_unix.go b/client/iface/configurer/kernel_unix.go index da69c2a35..3a95249c1 100644 --- a/client/iface/configurer/kernel_unix.go +++ b/client/iface/configurer/kernel_unix.go @@ -6,6 +6,7 @@ import ( "fmt" "net" "net/netip" + "slices" "time" log "github.com/sirupsen/logrus" @@ -18,16 +19,22 @@ import ( type KernelConfigurer struct { deviceName string statsCache *statsCache + allowedIPs *allowedIPStore } +// NewKernelConfigurer creates a configurer with an empty allowed IP mirror +// and a statistics cache for the named kernel device. func NewKernelConfigurer(deviceName string) *KernelConfigurer { c := &KernelConfigurer{ deviceName: deviceName, + allowedIPs: newAllowedIPStore(), } c.statsCache = newStatsCache(statsCacheTTL, c.fetchStats) return c } +// ConfigureInterface sets the device key, port and firewall mark, replacing all peers. +// The allowed IP mirror is reset only after the device accepts the configuration. func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error { log.Debugf("adding Wireguard private key") key, err := wgtypes.ParseKey(privateKey) @@ -46,6 +53,8 @@ func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error if err != nil { return fmt.Errorf(`received error "%w" while configuring interface %s with port %d`, err, c.deviceName, port) } + + c.allowedIPs.reset() return nil } @@ -58,9 +67,20 @@ func (c *KernelConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, upda } cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly) - return c.configure(cfg) + if err := c.configure(cfg); err != nil { + return err + } + + // Without updateOnly this creates the peer when it is absent, so the store has to + // know about it even though no allowed IP was configured. + if !updateOnly { + c.allowedIPs.ensure(parsedPeerKey) + } + return nil } +// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set. +// Prefixes assigned to this peer are transferred from their previous owners. func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { @@ -83,19 +103,23 @@ func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, if err != nil { return fmt.Errorf(`received error "%w" while updating peer on interface %s with settings: allowed ips %s, endpoint %s`, err, c.deviceName, allowedIps, endpoint.String()) } + + c.allowedIPs.add(peerKeyParsed, allowedIps) return nil } +// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured. +// Neither the netlink API nor the userspace one can clear an endpoint in place, so the peer +// is removed and re-added with the allowed IPs it already had. func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return err } - // Get the existing peer to preserve its allowed IPs - existingPeer, err := c.getPeer(c.deviceName, peerKey) + allowedIPs, err := c.peerAllowedIPs(peerKeyParsed) if err != nil { - return fmt.Errorf("get peer: %w", err) + return err } removePeerCfg := wgtypes.PeerConfig{ @@ -104,26 +128,27 @@ func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error { } if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{removePeerCfg}}); err != nil { - return fmt.Errorf(`error removing peer %s from interface %s: %w`, peerKey, c.deviceName, err) + return fmt.Errorf("remove peer %s from interface %s: %w", peerKey, c.deviceName, err) } - //Re-add the peer without the endpoint but same AllowedIPs reAddPeerCfg := wgtypes.PeerConfig{ PublicKey: peerKeyParsed, - AllowedIPs: existingPeer.AllowedIPs, + AllowedIPs: prefixesToIPNets(allowedIPs), ReplaceAllowedIPs: true, } if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{reAddPeerCfg}}); err != nil { + c.allowedIPs.forget(peerKeyParsed) return fmt.Errorf( - `error re-adding peer %s to interface %s with allowed IPs %v: %w`, - peerKey, c.deviceName, existingPeer.AllowedIPs, err, + "re-add peer %s to interface %s with allowed IPs %v: %w", + peerKey, c.deviceName, allowedIPs, err, ) } return nil } +// RemovePeer removes a peer and forgets its allowed IPs after a successful device write. func (c *KernelConfigurer) RemovePeer(peerKey string) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { @@ -142,15 +167,13 @@ func (c *KernelConfigurer) RemovePeer(peerKey string) error { if err != nil { return fmt.Errorf(`received error "%w" while removing peer %s from interface %s`, err, peerKey, c.deviceName) } + + c.allowedIPs.forget(peerKeyParsed) return nil } +// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op. func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error { - ipNet := net.IPNet{ - IP: allowedIP.Addr().AsSlice(), - Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()), - } - peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return err @@ -159,7 +182,7 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) PublicKey: peerKeyParsed, UpdateOnly: true, ReplaceAllowedIPs: false, - AllowedIPs: []net.IPNet{ipNet}, + AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}), } config := wgtypes.Config{ @@ -169,52 +192,69 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) if err != nil { return fmt.Errorf(`received error "%w" while adding allowed Ip to peer on interface %s with settings: allowed ips %s`, err, c.deviceName, allowedIP) } + + c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP}) return nil } +// RemoveAllowedIP removes a prefix while preserving the peer's other allowed IPs. +// A prefix not assigned to the peer is a no-op. func (c *KernelConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error { - ipNet := net.IPNet{ - IP: allowedIP.Addr().AsSlice(), - Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()), - } - peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return fmt.Errorf("parse peer key: %w", err) } - existingPeer, err := c.getPeer(c.deviceName, peerKey) + currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed) if err != nil { - return fmt.Errorf("get peer: %w", err) + return err } - newAllowedIPs := existingPeer.AllowedIPs - - for i, existingAllowedIP := range existingPeer.AllowedIPs { - if existingAllowedIP.String() == ipNet.String() { - newAllowedIPs = append(existingPeer.AllowedIPs[:i], existingPeer.AllowedIPs[i+1:]...) //nolint:gocritic - break - } + idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP)) + if idx < 0 { + return nil } + newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1) peer := wgtypes.PeerConfig{ PublicKey: peerKeyParsed, UpdateOnly: true, ReplaceAllowedIPs: true, - AllowedIPs: newAllowedIPs, + AllowedIPs: prefixesToIPNets(newAllowedIPs), } config := wgtypes.Config{ Peers: []wgtypes.PeerConfig{peer}, } - err = c.configure(config) - if err != nil { + if err := c.configure(config); err != nil { return fmt.Errorf("remove allowed IP %s on interface %s: %w", allowedIP, c.deviceName, err) } + + c.allowedIPs.set(peerKeyParsed, newAllowedIPs) return nil } -func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) { +// peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device +// only for a peer the store has not seen. Dumping the device costs a netlink round trip +// proportional to the whole network map, and this runs on every relay and ICE transition. +func (c *KernelConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) { + if prefixes, ok := c.allowedIPs.get(peerKey); ok { + return prefixes, nil + } + + existingPeer, err := c.getPeer(c.deviceName, peerKey) + if err != nil { + return nil, fmt.Errorf("get peer: %w", err) + } + + prefixes := ipNetsToPrefixes(existingPeer.AllowedIPs) + c.allowedIPs.set(peerKey, prefixes) + return prefixes, nil +} + +// getPeer scans the device for one peer. wgtypes.Key is an array, so the comparison is a +// plain equality: Key.String would base64 encode into a fresh allocation for every peer. +func (c *KernelConfigurer) getPeer(ifaceName string, peerPubKey wgtypes.Key) (wgtypes.Peer, error) { wg, err := wgctrl.New() if err != nil { return wgtypes.Peer{}, fmt.Errorf("wgctl: %w", err) @@ -231,7 +271,7 @@ func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, return wgtypes.Peer{}, fmt.Errorf("get device %s: %w", ifaceName, err) } for _, peer := range wgDevice.Peers { - if peer.PublicKey.String() == peerPubKey { + if peer.PublicKey == peerPubKey { return peer, nil } } diff --git a/client/iface/configurer/usp.go b/client/iface/configurer/usp.go index 2be1b861e..334d99369 100644 --- a/client/iface/configurer/usp.go +++ b/client/iface/configurer/usp.go @@ -8,6 +8,7 @@ import ( "net/netip" "os" "runtime" + "slices" "strconv" "strings" "time" @@ -41,31 +42,38 @@ type WGUSPConfigurer struct { deviceName string activityRecorder *bind.ActivityRecorder statsCache *statsCache + allowedIPs *allowedIPStore uapiListener net.Listener } +// NewUSPConfigurer creates a userspace configurer and starts its UAPI listener. func NewUSPConfigurer(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer { wgCfg := &WGUSPConfigurer{ device: device, deviceName: deviceName, activityRecorder: activityRecorder, + allowedIPs: newAllowedIPStore(), } wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats) wgCfg.startUAPI() return wgCfg } +// NewUSPConfigurerNoUAPI creates a userspace configurer without a UAPI listener. func NewUSPConfigurerNoUAPI(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer { wgCfg := &WGUSPConfigurer{ device: device, deviceName: deviceName, activityRecorder: activityRecorder, + allowedIPs: newAllowedIPStore(), } wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats) return wgCfg } +// ConfigureInterface sets the device key, port and firewall mark, replacing all peers. +// The allowed IP mirror is reset only after the device accepts the configuration. func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error { log.Debugf("adding Wireguard private key") key, err := wgtypes.ParseKey(privateKey) @@ -80,7 +88,12 @@ func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error ListenPort: &port, } - return c.device.IpcSet(toWgUserspaceString(config)) + if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil { + return err + } + + c.allowedIPs.reset() + return nil } // SetPresharedKey sets the preshared key for a peer. @@ -92,14 +105,38 @@ func (c *WGUSPConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, updat } cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly) - return c.device.IpcSet(toWgUserspaceString(cfg)) + if err := c.device.IpcSet(toWgUserspaceString(cfg)); err != nil { + return err + } + + // Without updateOnly this creates the peer when it is absent, so the store has to + // know about it even though no allowed IP was configured. + if !updateOnly { + c.allowedIPs.ensure(parsedPeerKey) + } + return nil } +// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set. +// It validates the endpoint before writing and records changes after a successful write. func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return err } + + // Everything that can fail is done before the device is touched, so a failure here + // cannot leave the device holding a peer that the activity recorder and the allowed + // IP store never learned about. + var addrPort netip.AddrPort + if endpoint != nil { + addr, err := netip.ParseAddr(endpoint.IP.String()) + if err != nil { + return fmt.Errorf("parse endpoint address: %w", err) + } + addrPort = netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port)) + } + peer := wgtypes.PeerConfig{ PublicKey: peerKeyParsed, ReplaceAllowedIPs: false, @@ -119,47 +156,27 @@ func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, } if endpoint != nil { - addr, err := netip.ParseAddr(endpoint.IP.String()) - if err != nil { - return fmt.Errorf("failed to parse endpoint address: %w", err) - } - addrPort := netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port)) c.activityRecorder.UpsertAddress(peerKey, addrPort) } + + c.allowedIPs.add(peerKeyParsed, allowedIps) return nil } +// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured. +// The UAPI cannot clear an endpoint in place, so the peer is removed and re-added with the +// allowed IPs it already had. func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return fmt.Errorf("parse peer key: %w", err) } - ipcStr, err := c.device.IpcGet() + allowedIPs, err := c.peerAllowedIPs(peerKeyParsed) if err != nil { - return fmt.Errorf("get IPC config: %w", err) + return err } - // Parse current status to get allowed IPs for the peer - stats, err := parseStatus(c.deviceName, ipcStr) - if err != nil { - return fmt.Errorf("parse IPC config: %w", err) - } - - var allowedIPs []net.IPNet - found := false - for _, peer := range stats.Peers { - if peer.PublicKey == peerKey { - allowedIPs = peer.AllowedIPs - found = true - break - } - } - if !found { - return fmt.Errorf("peer %s not found", peerKey) - } - - // remove the peer from the WireGuard configuration peer := wgtypes.PeerConfig{ PublicKey: peerKeyParsed, Remove: true, @@ -169,14 +186,13 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error { Peers: []wgtypes.PeerConfig{peer}, } if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil { - return fmt.Errorf("failed to remove peer: %s", ipcErr) + return fmt.Errorf("remove peer: %w", ipcErr) } - // Build the peer config peer = wgtypes.PeerConfig{ PublicKey: peerKeyParsed, ReplaceAllowedIPs: true, - AllowedIPs: allowedIPs, + AllowedIPs: prefixesToIPNets(allowedIPs), } config = wgtypes.Config{ @@ -184,12 +200,15 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error { } if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil { - return fmt.Errorf("remove endpoint address: %w", err) + c.allowedIPs.forget(peerKeyParsed) + return fmt.Errorf("re-add peer without endpoint: %w", err) } return nil } +// RemovePeer removes a peer, then clears its activity and allowed IP records. +// A failed device write leaves both records intact. func (c *WGUSPConfigurer) RemovePeer(peerKey string) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { @@ -204,18 +223,17 @@ func (c *WGUSPConfigurer) RemovePeer(peerKey string) error { config := wgtypes.Config{ Peers: []wgtypes.PeerConfig{peer}, } - ipcErr := c.device.IpcSet(toWgUserspaceString(config)) - - c.activityRecorder.Remove(peerKey) - return ipcErr -} - -func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error { - ipNet := net.IPNet{ - IP: allowedIP.Addr().AsSlice(), - Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()), + if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil { + return ipcErr } + c.activityRecorder.Remove(peerKey) + c.allowedIPs.forget(peerKeyParsed) + return nil +} + +// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op. +func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return err @@ -224,79 +242,89 @@ func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) e PublicKey: peerKeyParsed, UpdateOnly: true, ReplaceAllowedIPs: false, - AllowedIPs: []net.IPNet{ipNet}, + AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}), } config := wgtypes.Config{ Peers: []wgtypes.PeerConfig{peer}, } - return c.device.IpcSet(toWgUserspaceString(config)) + if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil { + return err + } + + c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP}) + return nil } +// RemoveAllowedIP removes a prefix while preserving the peer's other allowed IPs. +// It returns ErrAllowedIPNotFound if the prefix is not assigned to the peer. func (c *WGUSPConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error { - ipc, err := c.device.IpcGet() - if err != nil { - return err - } - peerKeyParsed, err := wgtypes.ParseKey(peerKey) + if err != nil { + return fmt.Errorf("parse peer key: %w", err) + } + + currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed) if err != nil { return err } - hexKey := hex.EncodeToString(peerKeyParsed[:]) - lines := strings.Split(ipc, "\n") + idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP)) + if idx < 0 { + return ErrAllowedIPNotFound + } + newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1) peer := wgtypes.PeerConfig{ PublicKey: peerKeyParsed, UpdateOnly: true, ReplaceAllowedIPs: true, - AllowedIPs: []net.IPNet{}, + AllowedIPs: prefixesToIPNets(newAllowedIPs), } - foundPeer := false - removedAllowedIP := false - ip := allowedIP.String() - - for _, line := range lines { - line = strings.TrimSpace(line) - - // If we're within the details of the found peer and encounter another public key, - // this means we're starting another peer's details. So, reset the flag. - if strings.HasPrefix(line, "public_key=") && foundPeer { - foundPeer = false - } - - // Identify the peer with the specific public key - if line == fmt.Sprintf("public_key=%s", hexKey) { - foundPeer = true - } - - // If we're within the details of the found peer and find the specific allowed IP, skip this line - if foundPeer && line == "allowed_ip="+ip { - removedAllowedIP = true - continue - } - - // Append the line to the output string - if foundPeer && strings.HasPrefix(line, "allowed_ip=") { - allowedIPStr := strings.TrimPrefix(line, "allowed_ip=") - _, ipNet, err := net.ParseCIDR(allowedIPStr) - if err != nil { - return err - } - peer.AllowedIPs = append(peer.AllowedIPs, *ipNet) - } - } - - if !removedAllowedIP { - return ErrAllowedIPNotFound - } config := wgtypes.Config{ Peers: []wgtypes.PeerConfig{peer}, } - return c.device.IpcSet(toWgUserspaceString(config)) + if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil { + return fmt.Errorf("remove allowed IP %s: %w", allowedIP, err) + } + + c.allowedIPs.set(peerKeyParsed, newAllowedIPs) + return nil +} + +// peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device +// only for a peer the store has not seen. Reading them back means dumping and parsing the +// whole device configuration, and this runs on every relay and ICE transition. +func (c *WGUSPConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) { + if prefixes, ok := c.allowedIPs.get(peerKey); ok { + return prefixes, nil + } + + ipcStr, err := c.device.IpcGet() + if err != nil { + return nil, fmt.Errorf("get IPC config: %w", err) + } + + stats, err := parseStatus(c.deviceName, ipcStr) + if err != nil { + return nil, fmt.Errorf("parse IPC config: %w", err) + } + + // parseStatus reports keys in their textual form, so the comparison needs it once. + wanted := peerKey.String() + for _, peer := range stats.Peers { + if peer.PublicKey != wanted { + continue + } + + prefixes := ipNetsToPrefixes(peer.AllowedIPs) + c.allowedIPs.set(peerKey, prefixes) + return prefixes, nil + } + + return nil, ErrPeerNotFound } func (c *WGUSPConfigurer) FullStats() (*Stats, error) { diff --git a/client/iface/configurer/usp_allowedips_test.go b/client/iface/configurer/usp_allowedips_test.go new file mode 100644 index 000000000..fba0ca546 --- /dev/null +++ b/client/iface/configurer/usp_allowedips_test.go @@ -0,0 +1,318 @@ +package configurer + +import ( + "net" + "net/netip" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + wgconn "golang.zx2c4.com/wireguard/conn" + wgdevice "golang.zx2c4.com/wireguard/device" + "golang.zx2c4.com/wireguard/tun/tuntest" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" + + "github.com/netbirdio/netbird/client/iface/bind" +) + +// newTestUSPConfigurer builds a configurer over a real wireguard-go device backed by an +// in-memory TUN. The device stays down, so no socket is opened and no privileges are needed. +func newTestUSPConfigurer(t *testing.T) *WGUSPConfigurer { + t.Helper() + + tun := tuntest.NewChannelTUN() + dev := wgdevice.NewDevice(tun.TUN(), wgconn.NewDefaultBind(), wgdevice.NewLogger(wgdevice.LogLevelSilent, "")) + t.Cleanup(dev.Close) + + c := NewUSPConfigurerNoUAPI(dev, "wgtest0", bind.NewActivityRecorder()) + + key, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate device private key") + require.NoError(t, c.ConfigureInterface(key.String(), 0), "configure test device") + + return c +} + +// seedPeers adds count peers, each with a /32 overlay address, and returns their public keys. +func seedPeers(t *testing.T, c *WGUSPConfigurer, count int) []string { + t.Helper() + + keys := make([]string, 0, count) + for i := 0; i < count; i++ { + priv, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate peer private key") + pub := priv.PublicKey().String() + + addr := netip.PrefixFrom(netip.AddrFrom4([4]byte{100, 64, byte(i >> 8), byte(i)}), 32) + require.NoError(t, c.UpdatePeer(pub, []netip.Prefix{addr}, 25*time.Second, nil, nil), "add peer") + keys = append(keys, pub) + } + return keys +} + +func peerAllowedIPs(t *testing.T, c *WGUSPConfigurer, peerKey string) []string { + t.Helper() + + stats, err := c.FullStats() + require.NoError(t, err, "read device stats") + + for _, p := range stats.Peers { + if p.PublicKey != peerKey { + continue + } + got := make([]string, 0, len(p.AllowedIPs)) + for _, ipNet := range p.AllowedIPs { + got = append(got, ipNet.String()) + } + return got + } + t.Fatalf("peer %s not found on device", peerKey) + return nil +} + +// TestRemoveEndpointAddressPreservesRoutedAllowedIPs covers the prefixes the route manager +// attaches to a routing peer through AddAllowedIP. Those are not known to the peer.Conn that +// triggers the endpoint removal, so dropping them here would silently blackhole every route +// behind that peer on each relay or ICE disconnect. +func TestRemoveEndpointAddressPreservesRoutedAllowedIPs(t *testing.T) { + c := newTestUSPConfigurer(t) + peerKey := seedPeers(t, c, 3)[1] + + routed := []netip.Prefix{ + netip.MustParsePrefix("10.20.0.0/16"), + netip.MustParsePrefix("192.168.7.0/24"), + } + for _, prefix := range routed { + require.NoError(t, c.AddAllowedIP(peerKey, prefix), "add routed prefix") + } + + before := peerAllowedIPs(t, c, peerKey) + require.Len(t, before, 3, "peer should hold its overlay address plus both routed prefixes") + + require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address") + + assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey), + "allowed IPs must survive the endpoint removal unchanged") +} + +// TestRemoveEndpointAddressDoesNotScaleWithPeerCount is the regression guard for the actual +// defect: clearing one peer's endpoint used to dump and parse the whole device, so its cost +// grew with the size of the network map. On a routing peer with thousands of peers that dump +// runs on every relay and ICE transition, under the interface lock. +func TestRemoveEndpointAddressDoesNotScaleWithPeerCount(t *testing.T) { + measure := func(peerCount int) float64 { + c := newTestUSPConfigurer(t) + peerKey := seedPeers(t, c, peerCount)[peerCount/2] + + return testing.AllocsPerRun(5, func() { + require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address") + }) + } + + small := measure(64) + large := measure(1024) + + assert.Less(t, large, small*2, + "clearing one endpoint allocated %.0f objects with 1024 peers against %.0f with 64: the cost still scales with the peer count", + large, small) +} + +// TestRemoveEndpointAddressFallsBackToDevice covers a peer the store never saw, which is what +// an out-of-band reconfiguration of the device leaves behind. The device stays the source of +// truth in that case, so the allowed IPs must still be preserved. +func TestRemoveEndpointAddressFallsBackToDevice(t *testing.T) { + c := newTestUSPConfigurer(t) + peerKey := seedPeers(t, c, 3)[1] + require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix") + + before := peerAllowedIPs(t, c, peerKey) + c.allowedIPs.reset() + + require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address") + + assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey), + "allowed IPs recovered from the device must be preserved") + + recovered, ok := c.allowedIPs.get(mustParseKey(t, peerKey)) + assert.True(t, ok, "the fallback must seed the store so the next call skips the device dump") + assert.Len(t, recovered, 2, "seeded prefixes") +} + +func TestRemoveAllowedIPKeepsTheOtherPrefixes(t *testing.T) { + c := newTestUSPConfigurer(t) + peerKey := seedPeers(t, c, 3)[0] + routed := netip.MustParsePrefix("10.20.0.0/16") + require.NoError(t, c.AddAllowedIP(peerKey, routed), "add routed prefix") + require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("192.168.7.0/24")), "add routed prefix") + + require.NoError(t, c.RemoveAllowedIP(peerKey, routed), "remove routed prefix") + + assert.ElementsMatch(t, []string{"100.64.0.0/32", "192.168.7.0/24"}, peerAllowedIPs(t, c, peerKey), + "only the removed prefix should be gone") + + assert.ErrorIs(t, c.RemoveAllowedIP(peerKey, routed), ErrAllowedIPNotFound, + "removing a prefix that is no longer configured must be reported") +} + +// TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt covers the lazy connection window documented +// in #6863: AddAllowedIP is update-only, a silent no-op when the peer is absent, so it must not +// leave the store claiming prefixes the device never took. RemoveEndpointAddress re-adds a peer +// without update-only, so a phantom entry would create a peer the device had dropped, and a +// created peer would steal those allowed IPs from whichever peer legitimately holds them. +func TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt(t *testing.T) { + c := newTestUSPConfigurer(t) + seedPeers(t, c, 2) + + priv, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate peer private key") + absent := priv.PublicKey().String() + + require.NoError(t, c.AddAllowedIP(absent, netip.MustParsePrefix("10.20.0.0/16")), + "update-only add on an absent peer is a silent no-op") + + stats, err := c.FullStats() + require.NoError(t, err, "read device stats") + require.Len(t, stats.Peers, 2, "the absent peer must not have been created by AddAllowedIP") + + assert.ErrorIs(t, c.RemoveEndpointAddress(absent), ErrPeerNotFound, + "clearing the endpoint of a peer the device does not have must fail") + + stats, err = c.FullStats() + require.NoError(t, err, "read device stats") + assert.Len(t, stats.Peers, 2, "no peer may be created while clearing an endpoint") +} + +// TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer covers WireGuard's rule that an +// allowed IP belongs to exactly one peer: configuring a prefix on a peer takes it away from +// whichever peer held it before. UpdatePeer relies on that rule rather than removing the prefix +// from the previous holder itself, so a prefix handed over between peers must not come back. +func TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer(t *testing.T) { + c := newTestUSPConfigurer(t) + keys := seedPeers(t, c, 2) + peerA, peerB := keys[0], keys[1] + routed := netip.MustParsePrefix("10.20.0.0/16") + + require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A") + require.Contains(t, peerAllowedIPs(t, c, peerA), routed.String(), "A must hold the prefix") + + // The route moves to B. The device takes it away from A on its own. + require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B") + require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix") + require.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(), "the device must have taken it from A") + + require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint") + + assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(), + "clearing A's endpoint must not take the prefix back from B") + assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), + "B must still hold the prefix") +} + +// TestPresharedKeyCreatedPeerTakesPartInPrefixHandover covers a peer created by a preshared +// key write rather than by a peer update. Rosenpass applies a peer's first key without +// updateOnly, which creates the peer on the device, so a store that ignored that operation +// would treat the peer as unknown and would not account for a prefix later handed over to it. +func TestPresharedKeyCreatedPeerTakesPartInPrefixHandover(t *testing.T) { + c := newTestUSPConfigurer(t) + peerA := seedPeers(t, c, 1)[0] + routed := netip.MustParsePrefix("10.20.0.0/16") + require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A") + + priv, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate peer private key") + peerB := priv.PublicKey().String() + + psk, err := wgtypes.GenerateKey() + require.NoError(t, err, "generate preshared key") + require.NoError(t, c.SetPresharedKey(peerB, psk, false), "a first key creates the peer") + + require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B") + require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix") + + require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint") + + assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(), + "clearing A's endpoint must not take the prefix back from B") + assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must still hold the prefix") +} + +// TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice is the end to end form of the +// conversion: a v4-mapped prefix must not reach the device as a zero length allowed IP, +// which would route every v4 address to that peer. +func TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice(t *testing.T) { + c := newTestUSPConfigurer(t) + + priv, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate peer private key") + peerKey := priv.PublicKey().String() + + mapped := netip.MustParsePrefix("::ffff:10.1.2.3/112") + require.NoError(t, c.UpdatePeer(peerKey, []netip.Prefix{mapped}, 25*time.Second, nil, nil), "add peer") + + onDevice := peerAllowedIPs(t, c, peerKey) + assert.NotContains(t, onDevice, "0.0.0.0/0", "the device must not be given a catch-all allowed IP") + assert.Equal(t, []string{"10.1.0.0/16"}, onDevice, "the device holds the normalized prefix") + + recorded, ok := c.allowedIPs.get(mustParseKey(t, peerKey)) + require.True(t, ok, "the peer must be recorded") + require.Len(t, recorded, 1, "one prefix recorded") + assert.Equal(t, onDevice[0], recorded[0].String(), "device and store must agree") +} + +// TestUpdatePeerWithAnUnusableEndpointTouchesNothing pins the ordering: the endpoint is +// parsed before the device is configured, so a failure cannot leave the device holding a +// peer that the store never learned about, with the prefix handover skipped along with it. +func TestUpdatePeerWithAnUnusableEndpointTouchesNothing(t *testing.T) { + c := newTestUSPConfigurer(t) + seedPeers(t, c, 2) + + priv, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate peer private key") + peerKey := priv.PublicKey().String() + + // A three byte address has no textual form netip can parse back. + endpoint := &net.UDPAddr{IP: net.IP{1, 2, 3}, Port: 51820} + require.Error(t, c.UpdatePeer(peerKey, []netip.Prefix{netip.MustParsePrefix("10.30.0.0/16")}, + 25*time.Second, endpoint, nil), "an unusable endpoint must fail the update") + + stats, err := c.FullStats() + require.NoError(t, err, "read device stats") + assert.Len(t, stats.Peers, 2, "the peer must not have reached the device") + + _, ok := c.allowedIPs.get(mustParseKey(t, peerKey)) + assert.False(t, ok, "the peer must not have been recorded either") +} + +// TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses covers a removal that never reached the +// device. A single peer removal is one write, so a failure leaves the peer on the device +// exactly as it was, and the record still describes it; dropping it would only force the +// next caller to read the whole device back for an answer it already had. +func TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses(t *testing.T) { + c := newTestUSPConfigurer(t) + peerKey := seedPeers(t, c, 1)[0] + require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix") + + before, ok := c.allowedIPs.get(mustParseKey(t, peerKey)) + require.True(t, ok, "the peer must be recorded before the removal") + require.Len(t, before, 2, "overlay address plus routed prefix") + + // A closed device refuses every write, which is the shape of any failed removal. + c.device.Close() + + require.Error(t, c.RemovePeer(peerKey), "the removal must report the failure") + + after, ok := c.allowedIPs.get(mustParseKey(t, peerKey)) + require.True(t, ok, "a peer still on the device must stay recorded") + assert.Equal(t, before, after, "the record must describe the peer the device kept") +} + +// mustParseKey turns the textual key the configurer API takes into the form the store +// keys on. +func mustParseKey(t *testing.T, key string) wgtypes.Key { + t.Helper() + + parsed, err := wgtypes.ParseKey(key) + require.NoError(t, err, "parse peer key") + return parsed +}