Merge branch 'main' into loopback-wg-proxy

This commit is contained in:
Viktor Liu
2026-09-30 05:50:22 +09:00
committed by GitHub
223 changed files with 20038 additions and 12222 deletions
+226
View File
@@ -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
}
+263
View File
@@ -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)
}
}
+8 -2
View File
@@ -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
+74 -34
View File
@@ -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
}
}
+120 -92
View File
@@ -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) {
@@ -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
}
+2 -15
View File
@@ -6,27 +6,14 @@ import (
"fmt"
"os/exec"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/wincmd"
)
func (w *WGIface) Destroy() error {
netshCmd := GetSystem32Command("netsh")
netshCmd := wincmd.System32("netsh")
out, err := exec.Command(netshCmd, "interface", "set", "interface", w.Name(), "admin=disable").CombinedOutput()
if err != nil {
return fmt.Errorf("failed to remove interface %s: %w - %s", w.Name(), err, out)
}
return nil
}
// GetSystem32Command checks if a command can be found in the system path and returns it. In case it can't find it
// in the path it will return the full path of a command assuming C:\windows\system32 as the base path.
func GetSystem32Command(command string) string {
_, err := exec.LookPath(command)
if err == nil {
return command
}
log.Tracef("Command %s not found in PATH, using C:\\windows\\system32\\%s.exe path", command, command)
return "C:\\windows\\system32\\" + command + ".exe"
}