diff --git a/cmd/cli/vpn_dns.go b/cmd/cli/vpn_dns.go index 707e2ab..50d81bd 100644 --- a/cmd/cli/vpn_dns.go +++ b/cmd/cli/vpn_dns.go @@ -52,6 +52,9 @@ type vpnDNSManager struct { // state for one guarded refresh cycle because Windows can briefly report an // intermediate empty adapter/DNS state after sleep/wake or reconnect. retainedAfterEmptyDiscovery bool + // discoverVPNDNS is injected for tests so Refresh does not depend on the + // runner host's real VPN/virtual adapter state. + discoverVPNDNS func(context.Context) []ctrld.VPNDNSConfig // Called when VPN DNS server list changes, to update intercept exemptions. onServersChanged vpnDNSExemptFunc } @@ -63,6 +66,7 @@ func newVPNDNSManager(logger *atomic.Pointer[ctrld.Logger], exemptFunc vpnDNSExe return &vpnDNSManager{ routes: make(map[string][]string), logger: logger, + discoverVPNDNS: ctrld.DiscoverVPNDNS, onServersChanged: exemptFunc, } } @@ -74,7 +78,11 @@ func (m *vpnDNSManager) Refresh(ctx context.Context, guardAgainstNoNameservers . guardedRefresh := len(guardAgainstNoNameservers) > 0 && guardAgainstNoNameservers[0] ctrld.Log(ctx, logger.Debug(), "Refreshing VPN DNS configurations") - configs := ctrld.DiscoverVPNDNS(ctx) + discoverVPNDNS := m.discoverVPNDNS + if discoverVPNDNS == nil { + discoverVPNDNS = ctrld.DiscoverVPNDNS + } + configs := discoverVPNDNS(ctx) // Detect exit mode: if the default route goes through a VPN DNS interface, // the VPN is routing ALL traffic (exit node / full tunnel). This is more diff --git a/cmd/cli/vpn_dns_test.go b/cmd/cli/vpn_dns_test.go index de893ec..32f6c24 100644 --- a/cmd/cli/vpn_dns_test.go +++ b/cmd/cli/vpn_dns_test.go @@ -21,6 +21,7 @@ func TestVPNDNSRefreshRetainsStateForOneGuardedEmptyDiscovery(t *testing.T) { gotExemptions = exemptions return nil }) + m.discoverVPNDNS = func(context.Context) []ctrld.VPNDNSConfig { return nil } m.configs = []ctrld.VPNDNSConfig{{ InterfaceName: "Ethernet 6", Servers: []string{"10.25.37.21", "10.25.37.22"}, @@ -47,6 +48,7 @@ func TestVPNDNSRefreshClearsOnSecondGuardedEmptyDiscovery(t *testing.T) { gotExemptions = exemptions return nil }) + m.discoverVPNDNS = func(context.Context) []ctrld.VPNDNSConfig { return nil } m.configs = []ctrld.VPNDNSConfig{{ InterfaceName: "Ethernet 6", Servers: []string{"10.25.37.21"},