[client] Cache the WireGuard interface check shared by ICE agents (#8001)

* [client] Take a WireGuard detector through the interface filter

The interface filter answers whether an interface is a WireGuard device by opening
a wgctrl client and asking for it, and it does that for every interface it is given.
Nothing about that call is tied to the caller, so it can be answered by a shared
object instead of being repeated, but the filter has no way to receive one.

InterfaceFilter and the constructors that build one now take a detector, and the ICE
config carries it so that every agent can be handed the same one. Nobody supplies a
detector yet: a nil one probes on every call, which is what the filter did before, so
this changes no behaviour.

* [client] Share one WireGuard detector across every ICE agent

Creating an ICE agent builds two interface filters, one for the agent and one for
the transport net it sits on, and each is asked about every host interface. For an
interface the disallow list does not settle, answering means opening a wgctrl client,
which builds a kernel and a userspace client and resolves the netlink family, and
then a round trip that usually just reports the device does not exist. An agent is
created per peer connection attempt, so on a large network that runs constantly:
on a routing peer with ~16000 peers it measured 2.40s of a 66.59s CPU profile, 3.6%,
split evenly between opening the client and the round trip.

The engine now owns a detector and passes it to every agent through the ICE config,
so the answer for an interface is reused instead of being asked again for each agent.
It is kept for a second, short enough that a WireGuard interface appearing is picked
up before ICE settles on candidates over it.

The callers that build one filter and keep it, the relay and the UDP mux, keep
passing nil and so keep probing, which costs them nothing at their rate.

* [client] Recheck the WireGuard cache inside the singleflight group

A caller that saw an expired entry could enter the singleflight group
after another caller had already refreshed the entry and left it, and
probe the interface a second time. Read the cache again inside the group
before probing.

This also makes the concurrent probe test independent of scheduling: a
late caller finds the fresh entry instead of starting a new probe.

* [client] Drop expired WireGuard detector entries

The detector lives as long as the engine and kept an entry for every
interface name it was ever asked about. On hosts that churn interfaces,
such as container veths, the map only grew. Remove expired entries when
a new answer is stored; the map holds a few dozen names at most, so the
sweep is cheap and runs at most once per interface per TTL.

* [client] Skip the disallow-list filter test on iOS

InterfaceFilter does not apply the disallow list on iOS, so the subtest
reaches the probe there and its no-probe assertion cannot hold.
This commit is contained in:
Riccardo Manfrin
2026-10-09 11:03:34 +02:00
committed by GitHub
parent 01d828d732
commit 3e85e40be2
23 changed files with 323 additions and 52 deletions
+9 -9
View File
@@ -51,7 +51,7 @@ func TestWGIface_UpdateAddr(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
addr := "100.64.0.1/8" addr := "100.64.0.1/8"
wgPort := 33100 wgPort := 33100
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
@@ -131,7 +131,7 @@ func getIfaceAddrs(ifaceName string) ([]net.Addr, error) {
func Test_CreateInterface(t *testing.T) { func Test_CreateInterface(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1)
wgIP := "10.99.99.1/32" wgIP := "10.99.99.1/32"
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
Address: wgaddr.MustParseWGAddress(wgIP), Address: wgaddr.MustParseWGAddress(wgIP),
@@ -171,7 +171,7 @@ func Test_Close(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
wgIP := "10.99.99.2/32" wgIP := "10.99.99.2/32"
wgPort := 33100 wgPort := 33100
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
@@ -213,7 +213,7 @@ func TestRecreation(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
wgIP := "10.99.99.2/32" wgIP := "10.99.99.2/32"
wgPort := 33100 wgPort := 33100
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
@@ -283,7 +283,7 @@ func Test_ConfigureInterface(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3)
wgIP := "10.99.99.5/30" wgIP := "10.99.99.5/30"
wgPort := 33100 wgPort := 33100
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
Address: wgaddr.MustParseWGAddress(wgIP), Address: wgaddr.MustParseWGAddress(wgIP),
@@ -335,7 +335,7 @@ func Test_ConfigureInterface(t *testing.T) {
func Test_UpdatePeer(t *testing.T) { func Test_UpdatePeer(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
wgIP := "10.99.99.9/30" wgIP := "10.99.99.9/30"
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
@@ -402,7 +402,7 @@ func Test_UpdatePeer(t *testing.T) {
func Test_RemovePeer(t *testing.T) { func Test_RemovePeer(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
wgIP := "10.99.99.13/30" wgIP := "10.99.99.13/30"
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
@@ -463,7 +463,7 @@ func Test_ConnectPeers(t *testing.T) {
peer2wgPort := 33200 peer2wgPort := 33200
keepAlive := 1 * time.Second keepAlive := 1 * time.Second
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
guid := fmt.Sprintf("{%s}", uuid.New().String()) guid := fmt.Sprintf("{%s}", uuid.New().String())
device.CustomWindowsGUIDString = strings.ToLower(guid) device.CustomWindowsGUIDString = strings.ToLower(guid)
@@ -499,7 +499,7 @@ func Test_ConnectPeers(t *testing.T) {
guid = fmt.Sprintf("{%s}", uuid.New().String()) guid = fmt.Sprintf("{%s}", uuid.New().String())
device.CustomWindowsGUIDString = strings.ToLower(guid) device.CustomWindowsGUIDString = strings.ToLower(guid)
newNet = stdnet.NewNet(context.Background(), testIFaceBlackList) newNet = stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
optsPeer2 := WGIFaceOpts{ optsPeer2 := WGIFaceOpts{
IFaceName: peer2ifaceName, IFaceName: peer2ifaceName,
+1 -1
View File
@@ -200,7 +200,7 @@ func (m *SingleSocketUDPMux) updateLocalAddresses() {
} }
if len(networks) > 0 { if len(networks) > 0 {
if m.params.Net == nil { if m.params.Net == nil {
m.params.Net = stdnet.NewNet(context.Background(), nil) m.params.Net = stdnet.NewNet(context.Background(), nil, nil)
} }
ips, err := localInterfaces(m.params.Net, m.params.InterfaceFilter, nil, networks, true) ips, err := localInterfaces(m.params.Net, m.params.InterfaceFilter, nil, networks, true)
@@ -247,7 +247,7 @@ func TestUpdateDNSServer(t *testing.T) {
for n, testCase := range testCases { for n, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) { t.Run(testCase.name, func(t *testing.T) {
privKey, _ := wgtypes.GenerateKey() privKey, _ := wgtypes.GenerateKey()
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := iface.WGIFaceOpts{ opts := iface.WGIFaceOpts{
IFaceName: fmt.Sprintf("utun230%d", n), IFaceName: fmt.Sprintf("utun230%d", n),
@@ -349,7 +349,7 @@ func TestDNSFakeResolverHandleUpdates(t *testing.T) {
defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) defer t.Setenv("NB_WG_KERNEL_DISABLED", ov)
t.Setenv("NB_WG_KERNEL_DISABLED", "true") t.Setenv("NB_WG_KERNEL_DISABLED", "true")
newNet := stdnet.NewNet(context.Background(), []string{"utun2301"}) newNet := stdnet.NewNet(context.Background(), []string{"utun2301"}, nil)
privKey, _ := wgtypes.GeneratePrivateKey() privKey, _ := wgtypes.GeneratePrivateKey()
opts := iface.WGIFaceOpts{ opts := iface.WGIFaceOpts{
+1 -1
View File
@@ -394,7 +394,7 @@ func createWgInterfaceWithBind(t *testing.T) (*iface.WGIface, error) {
defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) defer t.Setenv("NB_WG_KERNEL_DISABLED", ov)
t.Setenv("NB_WG_KERNEL_DISABLED", "true") t.Setenv("NB_WG_KERNEL_DISABLED", "true")
newNet := stdnet.NewNet(context.Background(), []string{"utun2301"}) newNet := stdnet.NewNet(context.Background(), []string{"utun2301"}, nil)
privKey, _ := wgtypes.GeneratePrivateKey() privKey, _ := wgtypes.GeneratePrivateKey()
+5
View File
@@ -57,6 +57,7 @@ import (
"github.com/netbirdio/netbird/client/internal/rosenpass" "github.com/netbirdio/netbird/client/internal/rosenpass"
"github.com/netbirdio/netbird/client/internal/routemanager" "github.com/netbirdio/netbird/client/internal/routemanager"
"github.com/netbirdio/netbird/client/internal/statemanager" "github.com/netbirdio/netbird/client/internal/statemanager"
"github.com/netbirdio/netbird/client/internal/stdnet"
"github.com/netbirdio/netbird/client/internal/syncstore" "github.com/netbirdio/netbird/client/internal/syncstore"
"github.com/netbirdio/netbird/client/internal/updater" "github.com/netbirdio/netbird/client/internal/updater"
"github.com/netbirdio/netbird/client/jobexec" "github.com/netbirdio/netbird/client/jobexec"
@@ -245,6 +246,9 @@ type Engine struct {
udpMux *udpmux.UniversalUDPMuxDefault udpMux *udpmux.UniversalUDPMuxDefault
// wgDetector is shared by every ICE agent through the ICE config.
wgDetector *stdnet.WGDetector
// networkSerial is the latest CurrentSerial (state ID) of the network sent by the Management service // networkSerial is the latest CurrentSerial (state ID) of the network sent by the Management service
networkSerial uint64 networkSerial uint64
@@ -362,6 +366,7 @@ func NewEngine(
mgmClient: services.MgmClient, mgmClient: services.MgmClient,
relayManager: services.RelayManager, relayManager: services.RelayManager,
peerStore: peerstore.NewConnStore(), peerStore: peerstore.NewConnStore(),
wgDetector: stdnet.NewWGDetector(),
syncMsgMux: &sync.Mutex{}, syncMsgMux: &sync.Mutex{},
config: config, config: config,
mobileDep: mobileDep, mobileDep: mobileDep,
+1
View File
@@ -15,5 +15,6 @@ func (e *Engine) createICEConfig() icemaker.Config {
UDPMux: e.udpMux.SingleSocketUDPMux, UDPMux: e.udpMux.SingleSocketUDPMux,
UDPMuxSrflx: e.udpMux, UDPMuxSrflx: e.udpMux,
NATExternalIPs: e.parseNATExternalIPMappings(), NATExternalIPs: e.parseNATExternalIPMappings(),
WGDetector: e.wgDetector,
} }
} }
+1
View File
@@ -13,6 +13,7 @@ func (e *Engine) createICEConfig() icemaker.Config {
InterfaceBlackList: e.config.IFaceBlackList, InterfaceBlackList: e.config.IFaceBlackList,
DisableIPv6Discovery: e.config.DisableIPv6Discovery, DisableIPv6Discovery: e.config.DisableIPv6Discovery,
NATExternalIPs: e.parseNATExternalIPMappings(), NATExternalIPs: e.parseNATExternalIPMappings(),
WGDetector: e.wgDetector,
} }
return cfg return cfg
} }
+1 -1
View File
@@ -7,5 +7,5 @@ import (
) )
func (e *Engine) newStdNet() *stdnet.Net { func (e *Engine) newStdNet() *stdnet.Net {
return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList) return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList, e.wgDetector)
} }
+1 -1
View File
@@ -3,5 +3,5 @@ package internal
import "github.com/netbirdio/netbird/client/internal/stdnet" import "github.com/netbirdio/netbird/client/internal/stdnet"
func (e *Engine) newStdNet() *stdnet.Net { func (e *Engine) newStdNet() *stdnet.Net {
return stdnet.NewNetWithDiscover(e.clientCtx, e.mobileDep.IFaceDiscover, e.config.IFaceBlackList) return stdnet.NewNetWithDiscover(e.clientCtx, e.mobileDep.IFaceDiscover, e.config.IFaceBlackList, e.wgDetector)
} }
+2 -2
View File
@@ -688,7 +688,7 @@ func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) {
StatusRecorder: peer.NewRecorder("https://mgm"), StatusRecorder: peer.NewRecorder("https://mgm"),
}, MobileDependency{}) }, MobileDependency{})
engine.ctx = ctx engine.ctx = ctx
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil)
opts := iface.WGIFaceOpts{ opts := iface.WGIFaceOpts{
IFaceName: wgIfaceName, IFaceName: wgIfaceName,
@@ -893,7 +893,7 @@ func TestEngine_UpdateNetworkMapWithDNSUpdate(t *testing.T) {
}, MobileDependency{}) }, MobileDependency{})
engine.ctx = ctx engine.ctx = ctx
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil)
opts := iface.WGIFaceOpts{ opts := iface.WGIFaceOpts{
IFaceName: wgIfaceName, IFaceName: wgIfaceName,
Address: wgaddr.MustParseWGAddress(wgAddr), Address: wgaddr.MustParseWGAddress(wgAddr),
+1 -1
View File
@@ -42,7 +42,7 @@ func TestNewConn_interfaceFilter(t *testing.T) {
ignore := []string{iface.WgInterfaceDefault, "tun0", "zt", "ZeroTier", "utun", "wg", "ts", ignore := []string{iface.WgInterfaceDefault, "tun0", "zt", "ZeroTier", "utun", "wg", "ts",
"Tailscale", "tailscale"} "Tailscale", "tailscale"}
filter := stdnet.InterfaceFilter(ignore) filter := stdnet.InterfaceFilter(ignore, nil)
for _, s := range ignore { for _, s := range ignore {
assert.Equal(t, filter(s), false) assert.Equal(t, filter(s), false)
+2 -2
View File
@@ -39,7 +39,7 @@ func NewAgent(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, c
iceFailedTimeout := iceFailedTimeout() iceFailedTimeout := iceFailedTimeout()
iceRelayAcceptanceMinWait := iceRelayAcceptanceMinWait() iceRelayAcceptanceMinWait := iceRelayAcceptanceMinWait()
transportNet := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList) transportNet := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList, config.WGDetector)
fac := logging.NewDefaultLoggerFactory() fac := logging.NewDefaultLoggerFactory()
@@ -50,7 +50,7 @@ func NewAgent(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, c
NetworkTypes: []ice.NetworkType{ice.NetworkTypeUDP4, ice.NetworkTypeUDP6}, NetworkTypes: []ice.NetworkType{ice.NetworkTypeUDP4, ice.NetworkTypeUDP6},
Urls: config.StunTurn.Load(), Urls: config.StunTurn.Load(),
CandidateTypes: candidateTypes, CandidateTypes: candidateTypes,
InterfaceFilter: stdnet.InterfaceFilter(config.InterfaceBlackList), InterfaceFilter: stdnet.InterfaceFilter(config.InterfaceBlackList, config.WGDetector),
UDPMux: config.UDPMux, UDPMux: config.UDPMux,
UDPMuxSrflx: config.UDPMuxSrflx, UDPMuxSrflx: config.UDPMuxSrflx,
NAT1To1IPs: config.NATExternalIPs, NAT1To1IPs: config.NATExternalIPs,
+6
View File
@@ -2,6 +2,8 @@ package ice
import ( import (
"github.com/pion/ice/v4" "github.com/pion/ice/v4"
"github.com/netbirdio/netbird/client/internal/stdnet"
) )
type Config struct { type Config struct {
@@ -17,4 +19,8 @@ type Config struct {
UDPMuxSrflx ice.UniversalUDPMux UDPMuxSrflx ice.UniversalUDPMux
NATExternalIPs []string NATExternalIPs []string
// WGDetector is shared by every agent so that the WireGuard check the interface
// filter performs is not repeated for each of them.
WGDetector *stdnet.WGDetector
} }
+2 -2
View File
@@ -8,6 +8,6 @@ import (
"github.com/netbirdio/netbird/client/internal/stdnet" "github.com/netbirdio/netbird/client/internal/stdnet"
) )
func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net { func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string, detector *stdnet.WGDetector) *stdnet.Net {
return stdnet.NewNet(ctx, ifaceBlacklist) return stdnet.NewNet(ctx, ifaceBlacklist, detector)
} }
+2 -2
View File
@@ -6,6 +6,6 @@ import (
"github.com/netbirdio/netbird/client/internal/stdnet" "github.com/netbirdio/netbird/client/internal/stdnet"
) )
func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net { func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string, detector *stdnet.WGDetector) *stdnet.Net {
return stdnet.NewNetWithDiscover(ctx, iFaceDiscover, ifaceBlacklist) return stdnet.NewNetWithDiscover(ctx, iFaceDiscover, ifaceBlacklist, detector)
} }
+2 -2
View File
@@ -201,7 +201,7 @@ func (p *StunTurnProbe) probeSTUN(ctx context.Context, uri *stun.URI) (addr stri
} }
}() }()
net := stdnet.NewNet(ctx, nil) net := stdnet.NewNet(ctx, nil, nil)
client, err := stun.DialURI(uri, &stun.DialConfig{ client, err := stun.DialURI(uri, &stun.DialConfig{
Net: net, Net: net,
@@ -286,7 +286,7 @@ func (p *StunTurnProbe) probeTURN(ctx context.Context, uri *stun.URI) (addr stri
} }
}() }()
net := stdnet.NewNet(ctx, nil) net := stdnet.NewNet(ctx, nil, nil)
cfg := &turn.ClientConfig{ cfg := &turn.ClientConfig{
STUNServerAddr: turnServerAddr, STUNServerAddr: turnServerAddr,
TURNServerAddr: turnServerAddr, TURNServerAddr: turnServerAddr,
+1 -1
View File
@@ -407,7 +407,7 @@ func TestManagerUpdateRoutes(t *testing.T) {
for n, testCase := range testCases { for n, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) { t.Run(testCase.name, func(t *testing.T) {
peerPrivateKey, _ := wgtypes.GeneratePrivateKey() peerPrivateKey, _ := wgtypes.GeneratePrivateKey()
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil)
opts := iface.WGIFaceOpts{ opts := iface.WGIFaceOpts{
IFaceName: fmt.Sprintf("utun43%d", n), IFaceName: fmt.Sprintf("utun43%d", n),
Address: wgaddr.MustParseWGAddress("100.65.65.2/24"), Address: wgaddr.MustParseWGAddress("100.65.65.2/24"),
@@ -437,7 +437,7 @@ func createWGInterface(t *testing.T, interfaceName, ipAddressCIDR string, listen
peerPrivateKey, err := wgtypes.GeneratePrivateKey() peerPrivateKey, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err) require.NoError(t, err)
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil)
opts := iface.WGIFaceOpts{ opts := iface.WGIFaceOpts{
IFaceName: interfaceName, IFaceName: interfaceName,
+5 -18
View File
@@ -3,17 +3,13 @@ package stdnet
import ( import (
"runtime" "runtime"
"strings" "strings"
log "github.com/sirupsen/logrus"
"golang.zx2c4.com/wireguard/wgctrl"
) )
// InterfaceFilter is a function passed to ICE Agent to filter out not allowed interfaces // InterfaceFilter is a function passed to ICE Agent to filter out not allowed interfaces
// to avoid building tunnel over them. // to avoid building tunnel over them. A nil detector probes the interface on every call,
func InterfaceFilter(disallowList []string) func(string) bool { // which is what the callers that build one filter for their whole lifetime want.
func InterfaceFilter(disallowList []string, detector *WGDetector) func(string) bool {
return func(iFace string) bool { return func(iFace string) bool {
if strings.HasPrefix(iFace, "lo") { if strings.HasPrefix(iFace, "lo") {
// hardcoded loopback check to support already installed agents // hardcoded loopback check to support already installed agents
return false return false
@@ -24,17 +20,8 @@ func InterfaceFilter(disallowList []string) func(string) bool {
return false return false
} }
} }
// look for unlisted WireGuard interfaces
wg, err := wgctrl.New()
if err != nil {
log.Debugf("trying to create a wgctrl client failed with: %v", err)
return true
}
defer func() {
_ = wg.Close()
}()
_, err = wg.Device(iFace) // look for unlisted WireGuard interfaces
return err != nil return !detector.IsWireGuard(iFace)
} }
} }
+4 -4
View File
@@ -45,12 +45,12 @@ type Net struct {
} }
// NewNetWithDiscover creates a new StdNet instance. // NewNetWithDiscover creates a new StdNet instance.
func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) *Net { func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string, detector *WGDetector) *Net {
if ctx == nil { if ctx == nil {
ctx = context.Background() ctx = context.Background()
} }
n := &Net{ n := &Net{
interfaceFilter: InterfaceFilter(disallowList), interfaceFilter: InterfaceFilter(disallowList, detector),
ctx: ctx, ctx: ctx,
} }
// current ExternalIFaceDiscover implement in android-client https://github.dev/netbirdio/android-client // current ExternalIFaceDiscover implement in android-client https://github.dev/netbirdio/android-client
@@ -64,13 +64,13 @@ func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover
} }
// NewNet creates a new StdNet instance. // NewNet creates a new StdNet instance.
func NewNet(ctx context.Context, disallowList []string) *Net { func NewNet(ctx context.Context, disallowList []string, detector *WGDetector) *Net {
if ctx == nil { if ctx == nil {
ctx = context.Background() ctx = context.Background()
} }
return &Net{ return &Net{
iFaceDiscover: pionDiscover{}, iFaceDiscover: pionDiscover{},
interfaceFilter: InterfaceFilter(disallowList), interfaceFilter: InterfaceFilter(disallowList, detector),
ctx: ctx, ctx: ctx,
} }
} }
+2 -2
View File
@@ -54,13 +54,13 @@ func TestNet_InterfacesDiscoversLazilyAndCaches(t *testing.T) {
} }
func TestNewNet_DoesNotDiscoverAtConstruction(t *testing.T) { func TestNewNet_DoesNotDiscoverAtConstruction(t *testing.T) {
n := NewNet(context.Background(), nil) n := NewNet(context.Background(), nil, nil)
require.NotNil(t, n) require.NotNil(t, n)
assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold") assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold")
} }
func TestNewNetWithDiscover_DoesNotDiscoverAtConstruction(t *testing.T) { func TestNewNetWithDiscover_DoesNotDiscoverAtConstruction(t *testing.T) {
n := NewNetWithDiscover(context.Background(), nil, nil) n := NewNetWithDiscover(context.Background(), nil, nil, nil)
require.NotNil(t, n) require.NotNil(t, n)
assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold") assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold")
} }
+119
View File
@@ -0,0 +1,119 @@
package stdnet
import (
"sync"
"time"
log "github.com/sirupsen/logrus"
"golang.org/x/sync/singleflight"
"golang.zx2c4.com/wireguard/wgctrl"
)
// wgDetectorTTL bounds how long a cached answer is trusted. An interface rarely
// becomes, or stops being, a WireGuard device, and the window only has to be short
// enough that ICE does not keep gathering candidates on one that just appeared.
const wgDetectorTTL = 1 * time.Second
type wgDetectorEntry struct {
isWireGuard bool
expireAt time.Time
}
// WGDetector answers whether an interface is a WireGuard device, remembering the
// answer for a short while.
//
// The question is asked once per interface for every ICE agent, and an agent is
// created per peer connection attempt, so on a large network the uncached form runs
// constantly. Answering it means opening a wgctrl client, which builds both a kernel
// and a userspace client and resolves the netlink family, and then a round trip that
// usually just reports the device does not exist.
//
// A detector is safe for concurrent use and is meant to be shared by every agent.
type WGDetector struct {
ttl time.Duration
// probe is replaced in tests; it is the call this type exists to avoid repeating.
probe func(string) bool
mu sync.RWMutex
cache map[string]wgDetectorEntry
sf singleflight.Group
}
// NewWGDetector returns a detector with the default time to live.
func NewWGDetector() *WGDetector {
return &WGDetector{
ttl: wgDetectorTTL,
probe: probeWireGuard,
cache: make(map[string]wgDetectorEntry),
}
}
// IsWireGuard reports whether the named interface is a WireGuard device. An interface
// it cannot ask about is reported as not being one, which leaves it available to ICE
// exactly as an uncached probe would.
func (d *WGDetector) IsWireGuard(iFace string) bool {
if d == nil {
return probeWireGuard(iFace)
}
if isWireGuard, ok := d.cached(iFace); ok {
return isWireGuard
}
result, _, _ := d.sf.Do(iFace, func() (interface{}, error) {
// A caller that saw the entry expire may get here after another caller already
// refreshed it and left the singleflight group.
if isWireGuard, ok := d.cached(iFace); ok {
return isWireGuard, nil
}
isWireGuard := d.probe(iFace)
d.store(iFace, isWireGuard)
return isWireGuard, nil
})
return result.(bool)
}
func (d *WGDetector) cached(iFace string) (isWireGuard, ok bool) {
d.mu.RLock()
defer d.mu.RUnlock()
entry, found := d.cache[iFace]
if !found || !time.Now().Before(entry.expireAt) {
return false, false
}
return entry.isWireGuard, true
}
// store records an answer and drops the expired ones, so names of interfaces that came
// and went, such as container veths, do not pile up for the life of the engine.
func (d *WGDetector) store(iFace string, isWireGuard bool) {
now := time.Now()
d.mu.Lock()
defer d.mu.Unlock()
for name, entry := range d.cache {
if !now.Before(entry.expireAt) {
delete(d.cache, name)
}
}
d.cache[iFace] = wgDetectorEntry{isWireGuard: isWireGuard, expireAt: now.Add(d.ttl)}
}
func probeWireGuard(iFace string) bool {
wg, err := wgctrl.New()
if err != nil {
log.Debugf("trying to create a wgctrl client failed with: %v", err)
return false
}
defer func() {
_ = wg.Close()
}()
_, err = wg.Device(iFace)
return err == nil
}
+152
View File
@@ -0,0 +1,152 @@
package stdnet
import (
"runtime"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// newCountingDetector returns a detector whose probe records how often it ran, so a test
// can assert on the thing this type exists for rather than on its return value alone.
func newCountingDetector(t *testing.T, ttl time.Duration, answer bool) (*WGDetector, *atomic.Int64) {
t.Helper()
var calls atomic.Int64
d := &WGDetector{
ttl: ttl,
cache: make(map[string]wgDetectorEntry),
probe: func(string) bool {
calls.Add(1)
return answer
},
}
return d, &calls
}
func TestWGDetectorAsksOncePerInterfaceWithinTheTTL(t *testing.T) {
d, calls := newCountingDetector(t, time.Minute, true)
for i := 0; i < 20; i++ {
assert.True(t, d.IsWireGuard("wt0"), "cached answer must not change")
}
assert.Equal(t, int64(1), calls.Load(), "the interface must be probed once within the TTL")
d.IsWireGuard("eth0")
assert.Equal(t, int64(2), calls.Load(), "a different interface is a different question and is probed on its own")
}
func TestWGDetectorReprobesAfterTheTTL(t *testing.T) {
d, calls := newCountingDetector(t, time.Millisecond, true)
require.True(t, d.IsWireGuard("wt0"), "first answer")
require.Equal(t, int64(1), calls.Load(), "first call probes")
time.Sleep(5 * time.Millisecond)
require.True(t, d.IsWireGuard("wt0"), "answer after expiry")
assert.Equal(t, int64(2), calls.Load(), "an expired entry must be probed again")
}
func TestWGDetectorCollapsesConcurrentProbes(t *testing.T) {
var calls atomic.Int64
release := make(chan struct{})
d := &WGDetector{
ttl: time.Minute,
cache: make(map[string]wgDetectorEntry),
probe: func(string) bool {
calls.Add(1)
<-release
return true
},
}
var wg sync.WaitGroup
for i := 0; i < 50; i++ {
wg.Add(1)
go func() {
defer wg.Done()
d.IsWireGuard("wt0")
}()
}
// The sleep only lets the callers pile up on the blocked probe so the collapse is
// exercised. The count does not depend on it: a caller that arrives after the probe
// finished finds the fresh entry, either before or inside the singleflight group.
time.Sleep(20 * time.Millisecond)
close(release)
wg.Wait()
assert.Equal(t, int64(1), calls.Load(), "concurrent callers must share one probe")
}
func TestWGDetectorNilProbesEveryTime(t *testing.T) {
var d *WGDetector
// A nil detector keeps the uncached behaviour, which is what the callers that build one
// filter for their whole lifetime rely on. It must not panic.
assert.NotPanics(t, func() { d.IsWireGuard("definitely-not-an-interface-0") },
"a nil detector must fall back to probing")
}
func TestInterfaceFilter(t *testing.T) {
wgDetector, calls := newCountingDetector(t, time.Minute, true)
plainDetector, _ := newCountingDetector(t, time.Minute, false)
t.Run("loopback is rejected without probing", func(t *testing.T) {
filter := InterfaceFilter(nil, wgDetector)
assert.False(t, filter("lo"), "loopback must never be offered to ICE")
assert.Equal(t, int64(0), calls.Load(), "a name settled by prefix must not reach the probe")
})
t.Run("a disallowed interface is rejected without probing", func(t *testing.T) {
if runtime.GOOS == "ios" {
t.Skip("the disallow list is not applied on iOS")
}
filter := InterfaceFilter([]string{"wt"}, wgDetector)
assert.False(t, filter("wt0"), "a blacklisted interface must be rejected")
assert.Equal(t, int64(0), calls.Load(), "a name settled by the disallow list must not reach the probe")
})
t.Run("an unlisted WireGuard interface is rejected", func(t *testing.T) {
filter := InterfaceFilter(nil, wgDetector)
assert.False(t, filter("somewg0"), "a WireGuard interface must not be used to build a tunnel")
})
t.Run("an ordinary interface is allowed", func(t *testing.T) {
filter := InterfaceFilter(nil, plainDetector)
assert.True(t, filter("eth0"), "a plain interface must remain available to ICE")
})
}
func TestInterfaceFilterSharesOneProbeAcrossFilters(t *testing.T) {
d, calls := newCountingDetector(t, time.Minute, false)
// Every ICE agent builds its own filter, twice, and each one is asked about every
// interface. Sharing the detector is what keeps that from repeating the probe.
for i := 0; i < 10; i++ {
filter := InterfaceFilter(nil, d)
require.True(t, filter("eth0"), "a plain interface stays allowed")
require.True(t, filter("eth1"), "a plain interface stays allowed")
}
assert.Equal(t, int64(2), calls.Load(), "one probe per interface, not per filter")
}
func TestWGDetectorDropsExpiredEntries(t *testing.T) {
d, _ := newCountingDetector(t, time.Millisecond, false)
for _, name := range []string{"veth1", "veth2", "veth3"} {
d.IsWireGuard(name)
}
time.Sleep(5 * time.Millisecond)
d.IsWireGuard("eth0")
d.mu.RLock()
defer d.mu.RUnlock()
assert.Len(t, d.cache, 1, "expired entries of vanished interfaces must be dropped")
assert.Contains(t, d.cache, "eth0", "the fresh answer must be kept")
}