mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-01 19:19:07 +02:00
Merge branch 'main' into loopback-wg-proxy
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user