diff --git a/client/iface/iface_test.go b/client/iface/iface_test.go index cb50ca4a1..690a0f66f 100644 --- a/client/iface/iface_test.go +++ b/client/iface/iface_test.go @@ -51,7 +51,7 @@ func TestWGIface_UpdateAddr(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) addr := "100.64.0.1/8" wgPort := 33100 - newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -131,7 +131,7 @@ func getIfaceAddrs(ifaceName string) ([]net.Addr, error) { func Test_CreateInterface(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1) wgIP := "10.99.99.1/32" - newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil) opts := WGIFaceOpts{ IFaceName: ifaceName, Address: wgaddr.MustParseWGAddress(wgIP), @@ -171,7 +171,7 @@ func Test_Close(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2) wgIP := "10.99.99.2/32" wgPort := 33100 - newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -213,7 +213,7 @@ func TestRecreation(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2) wgIP := "10.99.99.2/32" wgPort := 33100 - newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -283,7 +283,7 @@ func Test_ConfigureInterface(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3) wgIP := "10.99.99.5/30" wgPort := 33100 - newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil) opts := WGIFaceOpts{ IFaceName: ifaceName, Address: wgaddr.MustParseWGAddress(wgIP), @@ -335,7 +335,7 @@ func Test_ConfigureInterface(t *testing.T) { func Test_UpdatePeer(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) wgIP := "10.99.99.9/30" - newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -402,7 +402,7 @@ func Test_UpdatePeer(t *testing.T) { func Test_RemovePeer(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) wgIP := "10.99.99.13/30" - newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -463,7 +463,7 @@ func Test_ConnectPeers(t *testing.T) { peer2wgPort := 33200 keepAlive := 1 * time.Second - newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil) guid := fmt.Sprintf("{%s}", uuid.New().String()) device.CustomWindowsGUIDString = strings.ToLower(guid) @@ -499,7 +499,7 @@ func Test_ConnectPeers(t *testing.T) { guid = fmt.Sprintf("{%s}", uuid.New().String()) device.CustomWindowsGUIDString = strings.ToLower(guid) - newNet = stdnet.NewNet(context.Background(), testIFaceBlackList) + newNet = stdnet.NewNet(context.Background(), testIFaceBlackList, nil) optsPeer2 := WGIFaceOpts{ IFaceName: peer2ifaceName, diff --git a/client/iface/udpmux/mux.go b/client/iface/udpmux/mux.go index 68cecc953..3aa0b4d88 100644 --- a/client/iface/udpmux/mux.go +++ b/client/iface/udpmux/mux.go @@ -200,7 +200,7 @@ func (m *SingleSocketUDPMux) updateLocalAddresses() { } if len(networks) > 0 { 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) diff --git a/client/internal/dns/server_privileged_test.go b/client/internal/dns/server_privileged_test.go index 270e3bf91..b135c3661 100644 --- a/client/internal/dns/server_privileged_test.go +++ b/client/internal/dns/server_privileged_test.go @@ -247,7 +247,7 @@ func TestUpdateDNSServer(t *testing.T) { for n, testCase := range testCases { t.Run(testCase.name, func(t *testing.T) { privKey, _ := wgtypes.GenerateKey() - newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil) opts := iface.WGIFaceOpts{ IFaceName: fmt.Sprintf("utun230%d", n), @@ -349,7 +349,7 @@ func TestDNSFakeResolverHandleUpdates(t *testing.T) { defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) 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() opts := iface.WGIFaceOpts{ diff --git a/client/internal/dns/server_test.go b/client/internal/dns/server_test.go index 414890158..777594272 100644 --- a/client/internal/dns/server_test.go +++ b/client/internal/dns/server_test.go @@ -394,7 +394,7 @@ func createWgInterfaceWithBind(t *testing.T) (*iface.WGIface, error) { defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) 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() diff --git a/client/internal/engine.go b/client/internal/engine.go index e167f7589..ccff974d8 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -57,6 +57,7 @@ import ( "github.com/netbirdio/netbird/client/internal/rosenpass" "github.com/netbirdio/netbird/client/internal/routemanager" "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/updater" "github.com/netbirdio/netbird/client/jobexec" @@ -245,6 +246,9 @@ type Engine struct { 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 uint64 @@ -362,6 +366,7 @@ func NewEngine( mgmClient: services.MgmClient, relayManager: services.RelayManager, peerStore: peerstore.NewConnStore(), + wgDetector: stdnet.NewWGDetector(), syncMsgMux: &sync.Mutex{}, config: config, mobileDep: mobileDep, diff --git a/client/internal/engine_generic.go b/client/internal/engine_generic.go index 34a75e45b..e293f1636 100644 --- a/client/internal/engine_generic.go +++ b/client/internal/engine_generic.go @@ -15,5 +15,6 @@ func (e *Engine) createICEConfig() icemaker.Config { UDPMux: e.udpMux.SingleSocketUDPMux, UDPMuxSrflx: e.udpMux, NATExternalIPs: e.parseNATExternalIPMappings(), + WGDetector: e.wgDetector, } } diff --git a/client/internal/engine_js.go b/client/internal/engine_js.go index dce3c57fb..0243b7b37 100644 --- a/client/internal/engine_js.go +++ b/client/internal/engine_js.go @@ -13,6 +13,7 @@ func (e *Engine) createICEConfig() icemaker.Config { InterfaceBlackList: e.config.IFaceBlackList, DisableIPv6Discovery: e.config.DisableIPv6Discovery, NATExternalIPs: e.parseNATExternalIPMappings(), + WGDetector: e.wgDetector, } return cfg } diff --git a/client/internal/engine_stdnet.go b/client/internal/engine_stdnet.go index 86f6d297a..d6e2b1aa4 100644 --- a/client/internal/engine_stdnet.go +++ b/client/internal/engine_stdnet.go @@ -7,5 +7,5 @@ import ( ) func (e *Engine) newStdNet() *stdnet.Net { - return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList) + return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList, e.wgDetector) } diff --git a/client/internal/engine_stdnet_android.go b/client/internal/engine_stdnet_android.go index b14deeadf..7ef996203 100644 --- a/client/internal/engine_stdnet_android.go +++ b/client/internal/engine_stdnet_android.go @@ -3,5 +3,5 @@ package internal import "github.com/netbirdio/netbird/client/internal/stdnet" 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) } diff --git a/client/internal/engine_test.go b/client/internal/engine_test.go index 14076c051..bfde05c04 100644 --- a/client/internal/engine_test.go +++ b/client/internal/engine_test.go @@ -688,7 +688,7 @@ func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) { StatusRecorder: peer.NewRecorder("https://mgm"), }, MobileDependency{}) engine.ctx = ctx - newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil) opts := iface.WGIFaceOpts{ IFaceName: wgIfaceName, @@ -893,7 +893,7 @@ func TestEngine_UpdateNetworkMapWithDNSUpdate(t *testing.T) { }, MobileDependency{}) engine.ctx = ctx - newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil) opts := iface.WGIFaceOpts{ IFaceName: wgIfaceName, Address: wgaddr.MustParseWGAddress(wgAddr), diff --git a/client/internal/peer/conn_test.go b/client/internal/peer/conn_test.go index b709d5e40..8f3d6216d 100644 --- a/client/internal/peer/conn_test.go +++ b/client/internal/peer/conn_test.go @@ -42,7 +42,7 @@ func TestNewConn_interfaceFilter(t *testing.T) { ignore := []string{iface.WgInterfaceDefault, "tun0", "zt", "ZeroTier", "utun", "wg", "ts", "Tailscale", "tailscale"} - filter := stdnet.InterfaceFilter(ignore) + filter := stdnet.InterfaceFilter(ignore, nil) for _, s := range ignore { assert.Equal(t, filter(s), false) diff --git a/client/internal/peer/ice/agent.go b/client/internal/peer/ice/agent.go index 6cd8c48de..a919cded5 100644 --- a/client/internal/peer/ice/agent.go +++ b/client/internal/peer/ice/agent.go @@ -39,7 +39,7 @@ func NewAgent(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, c iceFailedTimeout := iceFailedTimeout() iceRelayAcceptanceMinWait := iceRelayAcceptanceMinWait() - transportNet := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList) + transportNet := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList, config.WGDetector) fac := logging.NewDefaultLoggerFactory() @@ -50,7 +50,7 @@ func NewAgent(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, c NetworkTypes: []ice.NetworkType{ice.NetworkTypeUDP4, ice.NetworkTypeUDP6}, Urls: config.StunTurn.Load(), CandidateTypes: candidateTypes, - InterfaceFilter: stdnet.InterfaceFilter(config.InterfaceBlackList), + InterfaceFilter: stdnet.InterfaceFilter(config.InterfaceBlackList, config.WGDetector), UDPMux: config.UDPMux, UDPMuxSrflx: config.UDPMuxSrflx, NAT1To1IPs: config.NATExternalIPs, diff --git a/client/internal/peer/ice/config.go b/client/internal/peer/ice/config.go index dd5d67403..9d6b19925 100644 --- a/client/internal/peer/ice/config.go +++ b/client/internal/peer/ice/config.go @@ -2,6 +2,8 @@ package ice import ( "github.com/pion/ice/v4" + + "github.com/netbirdio/netbird/client/internal/stdnet" ) type Config struct { @@ -17,4 +19,8 @@ type Config struct { UDPMuxSrflx ice.UniversalUDPMux 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 } diff --git a/client/internal/peer/ice/stdnet.go b/client/internal/peer/ice/stdnet.go index 0c819ff66..fbbe4f388 100644 --- a/client/internal/peer/ice/stdnet.go +++ b/client/internal/peer/ice/stdnet.go @@ -8,6 +8,6 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" ) -func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net { - return stdnet.NewNet(ctx, ifaceBlacklist) +func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string, detector *stdnet.WGDetector) *stdnet.Net { + return stdnet.NewNet(ctx, ifaceBlacklist, detector) } diff --git a/client/internal/peer/ice/stdnet_android.go b/client/internal/peer/ice/stdnet_android.go index 2962ecf66..72f0c329e 100644 --- a/client/internal/peer/ice/stdnet_android.go +++ b/client/internal/peer/ice/stdnet_android.go @@ -6,6 +6,6 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" ) -func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net { - return stdnet.NewNetWithDiscover(ctx, iFaceDiscover, ifaceBlacklist) +func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string, detector *stdnet.WGDetector) *stdnet.Net { + return stdnet.NewNetWithDiscover(ctx, iFaceDiscover, ifaceBlacklist, detector) } diff --git a/client/internal/relay/relay.go b/client/internal/relay/relay.go index f0c65301e..e39eed39a 100644 --- a/client/internal/relay/relay.go +++ b/client/internal/relay/relay.go @@ -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{ 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{ STUNServerAddr: turnServerAddr, TURNServerAddr: turnServerAddr, diff --git a/client/internal/routemanager/manager_test.go b/client/internal/routemanager/manager_test.go index a1624cf46..766ea1c61 100644 --- a/client/internal/routemanager/manager_test.go +++ b/client/internal/routemanager/manager_test.go @@ -407,7 +407,7 @@ func TestManagerUpdateRoutes(t *testing.T) { for n, testCase := range testCases { t.Run(testCase.name, func(t *testing.T) { peerPrivateKey, _ := wgtypes.GeneratePrivateKey() - newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil) opts := iface.WGIFaceOpts{ IFaceName: fmt.Sprintf("utun43%d", n), Address: wgaddr.MustParseWGAddress("100.65.65.2/24"), diff --git a/client/internal/routemanager/systemops/systemops_generic_test.go b/client/internal/routemanager/systemops/systemops_generic_test.go index 5b569ebd6..f117ff751 100644 --- a/client/internal/routemanager/systemops/systemops_generic_test.go +++ b/client/internal/routemanager/systemops/systemops_generic_test.go @@ -437,7 +437,7 @@ func createWGInterface(t *testing.T, interfaceName, ipAddressCIDR string, listen peerPrivateKey, err := wgtypes.GeneratePrivateKey() require.NoError(t, err) - newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil) opts := iface.WGIFaceOpts{ IFaceName: interfaceName, diff --git a/client/internal/stdnet/filter.go b/client/internal/stdnet/filter.go index e45714001..07025ff1e 100644 --- a/client/internal/stdnet/filter.go +++ b/client/internal/stdnet/filter.go @@ -3,17 +3,13 @@ package stdnet import ( "runtime" "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 -// to avoid building tunnel over them. -func InterfaceFilter(disallowList []string) func(string) bool { - +// to avoid building tunnel over them. A nil detector probes the interface on every call, +// 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 { - if strings.HasPrefix(iFace, "lo") { // hardcoded loopback check to support already installed agents return false @@ -24,17 +20,8 @@ func InterfaceFilter(disallowList []string) func(string) bool { 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) - return err != nil + // look for unlisted WireGuard interfaces + return !detector.IsWireGuard(iFace) } } diff --git a/client/internal/stdnet/stdnet.go b/client/internal/stdnet/stdnet.go index c3a9d3d97..32f030612 100644 --- a/client/internal/stdnet/stdnet.go +++ b/client/internal/stdnet/stdnet.go @@ -45,12 +45,12 @@ type Net struct { } // 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 { ctx = context.Background() } n := &Net{ - interfaceFilter: InterfaceFilter(disallowList), + interfaceFilter: InterfaceFilter(disallowList, detector), ctx: ctx, } // 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. -func NewNet(ctx context.Context, disallowList []string) *Net { +func NewNet(ctx context.Context, disallowList []string, detector *WGDetector) *Net { if ctx == nil { ctx = context.Background() } return &Net{ iFaceDiscover: pionDiscover{}, - interfaceFilter: InterfaceFilter(disallowList), + interfaceFilter: InterfaceFilter(disallowList, detector), ctx: ctx, } } diff --git a/client/internal/stdnet/stdnet_test.go b/client/internal/stdnet/stdnet_test.go index 822972f39..2b16a50c0 100644 --- a/client/internal/stdnet/stdnet_test.go +++ b/client/internal/stdnet/stdnet_test.go @@ -54,13 +54,13 @@ func TestNet_InterfacesDiscoversLazilyAndCaches(t *testing.T) { } func TestNewNet_DoesNotDiscoverAtConstruction(t *testing.T) { - n := NewNet(context.Background(), nil) + n := NewNet(context.Background(), nil, nil) require.NotNil(t, n) assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold") } func TestNewNetWithDiscover_DoesNotDiscoverAtConstruction(t *testing.T) { - n := NewNetWithDiscover(context.Background(), nil, nil) + n := NewNetWithDiscover(context.Background(), nil, nil, nil) require.NotNil(t, n) assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold") } diff --git a/client/internal/stdnet/wgdetector.go b/client/internal/stdnet/wgdetector.go new file mode 100644 index 000000000..f34e469e1 --- /dev/null +++ b/client/internal/stdnet/wgdetector.go @@ -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 +} diff --git a/client/internal/stdnet/wgdetector_test.go b/client/internal/stdnet/wgdetector_test.go new file mode 100644 index 000000000..68260ae59 --- /dev/null +++ b/client/internal/stdnet/wgdetector_test.go @@ -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") +}