diff --git a/resolver.go b/resolver.go index a73bfa4..d84acf9 100644 --- a/resolver.go +++ b/resolver.go @@ -21,6 +21,7 @@ import ( "tailscale.com/net/tsaddr" "github.com/Control-D-Inc/ctrld/internal/dnscache" + ctrldnet "github.com/Control-D-Inc/ctrld/internal/net" ) const ( @@ -227,6 +228,10 @@ type osResolver struct { publicServers atomic.Pointer[[]string] group *singleflight.Group cache *sync.Map + // Per-resolver seams let tests exercise the production Resolve path without + // mutating process-wide resolver state. + exchangeDNS dnsExchangeFunc + localIP func(string) net.IP } type osResolverResult struct { @@ -342,21 +347,65 @@ func GetDefaultLocalIPv6() net.IP { return nil } -// customDNSExchange wraps the DNS exchange to use our debug dialer. -// It uses dns.ExchangeWithConn so that our custom dialer is used directly. -func customDNSExchange(ctx context.Context, msg *dns.Msg, server string, desiredLocalIP net.IP) (*dns.Msg, time.Duration, error) { +type dnsExchangeFunc func(context.Context, *dns.Msg, string, net.IP) (*dns.Msg, time.Duration, error) + +func exchangeDNS(ctx context.Context, msg *dns.Msg, server string, localIP net.IP) (*dns.Msg, time.Duration, error) { baseDialer := &net.Dialer{ Timeout: 3 * time.Second, Resolver: &net.Resolver{PreferGo: true}, } - if desiredLocalIP != nil { - baseDialer.LocalAddr = &net.UDPAddr{IP: desiredLocalIP, Port: 0} + if localIP != nil { + baseDialer.LocalAddr = &net.UDPAddr{IP: localIP, Port: 0} } dnsClient := &dns.Client{Net: "udp"} dnsClient.Dialer = baseDialer return dnsClient.ExchangeContext(ctx, msg, server) } +func defaultLocalIPForServer(server string) net.IP { + if runtime.GOOS != "darwin" { + return nil + } + host, _, err := net.SplitHostPort(server) + if err != nil { + return nil + } + ip := net.ParseIP(host) + if ip != nil && ip.To4() == nil { + return GetDefaultLocalIPv6() + } + return GetDefaultLocalIPv4() +} + +func preSendUnreachable(err error) bool { + var opErr *net.OpError + if !errors.As(err, &opErr) || (opErr.Op != "dial" && opErr.Op != "write") { + return false + } + return ctrldnet.IsUnreachable(err) +} + +// customDNSExchangeWith preserves the preferred source first. A route-selected +// retry is allowed only when the caller knows the server is an OS-selected resolver, +// not ctrld's synthetic public fallback. This includes public DNS pushed by a VPN: +// unbinding changes the source route, not the recipient. +func customDNSExchangeWith(ctx context.Context, msg *dns.Msg, server string, desiredLocalIP net.IP, allowRouteSelectedRetry bool, exchange dnsExchangeFunc) (*dns.Msg, time.Duration, error) { + answer, rtt, err := exchange(ctx, msg, server, desiredLocalIP) + if answer != nil || err == nil || ctx.Err() != nil || desiredLocalIP == nil || !allowRouteSelectedRetry || !preSendUnreachable(err) { + return answer, rtt, err + } + + Log(ctx, ProxyLogger.Load().Debug(), "OS resolver source binding is unreachable; retrying with route-selected source") + return exchange(ctx, msg.Copy(), server, nil) +} + +// allowRouteSelectedRetryForOSServer excludes only ctrld's synthetic public +// fallback. System-provided resolvers remain eligible even when their addresses +// are public, as with VPNs that push public DNS servers. +func allowRouteSelectedRetryForOSServer(server string) bool { + return server != controldPublicDnsWithPort +} + const hotCacheTTL = time.Second // Resolve resolves DNS queries using pre-configured nameservers. @@ -465,6 +514,14 @@ func (o *osResolver) resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error ch := make(chan *osResolverResult, numServers) wg := &sync.WaitGroup{} + exchange := o.exchangeDNS + if exchange == nil { + exchange = exchangeDNS + } + localIPForServer := o.localIP + if localIPForServer == nil { + localIPForServer = defaultLocalIPForServer + } wg.Add(numServers) go func() { wg.Wait() @@ -484,20 +541,8 @@ func (o *osResolver) resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error var answer *dns.Msg var err error - var localOSResolverIP net.IP - if runtime.GOOS == "darwin" { - host, _, err := net.SplitHostPort(server) - if err == nil { - ip := net.ParseIP(host) - if ip != nil && ip.To4() == nil { - // IPv6 nameserver; use default IPv6 address (if set) - localOSResolverIP = GetDefaultLocalIPv6() - } else { - localOSResolverIP = GetDefaultLocalIPv4() - } - } - } - answer, _, err = customDNSExchange(ctx, msg.Copy(), server, localOSResolverIP) + localOSResolverIP := localIPForServer(server) + answer, _, err = customDNSExchangeWith(ctx, msg.Copy(), server, localOSResolverIP, allowRouteSelectedRetryForOSServer(server), exchange) ch <- &osResolverResult{answer: answer, err: err, server: server, lan: isLan} }(server) } diff --git a/resolver_test.go b/resolver_test.go index 30a3af7..930871a 100644 --- a/resolver_test.go +++ b/resolver_test.go @@ -4,9 +4,12 @@ import ( "context" "crypto/rand" "encoding/hex" + "errors" "net" + "os" "sync" "sync/atomic" + "syscall" "testing" "time" @@ -70,6 +73,270 @@ func Test_osResolver_ResolveLanHostname(t *testing.T) { } } +func Test_customDNSExchangeWith_RetriesUnboundOnUnreachableSource(t *testing.T) { + tests := []struct { + name string + boundIP net.IP + server string + errno syscall.Errno + }{ + {"ipv4 network unreachable", net.ParseIP("192.0.2.10"), "192.0.2.53:53", syscall.ENETUNREACH}, + {"ipv4 host unreachable", net.ParseIP("192.0.2.10"), "192.0.2.53:53", syscall.EHOSTUNREACH}, + {"ipv6 network unreachable", net.ParseIP("2001:db8::10"), "[2001:db8::53]:53", syscall.ENETUNREACH}, + {"ipv6 host unreachable", net.ParseIP("2001:db8::10"), "[2001:db8::53]:53", syscall.EHOSTUNREACH}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + msg := new(dns.Msg) + msg.SetQuestion("internal.example.", dns.TypeA) + var localIPs []net.IP + var servers []string + exchange := func(_ context.Context, msg *dns.Msg, server string, localIP net.IP) (*dns.Msg, time.Duration, error) { + localIPs = append(localIPs, append(net.IP(nil), localIP...)) + servers = append(servers, server) + if localIP != nil { + return nil, 0, &net.OpError{Op: "write", Net: "udp", Err: &os.SyscallError{Syscall: "write", Err: tt.errno}} + } + answer := new(dns.Msg) + answer.SetReply(msg) + return answer, time.Millisecond, nil + } + + answer, _, err := customDNSExchangeWith(context.Background(), msg, tt.server, tt.boundIP, true, exchange) + if err != nil { + t.Fatal(err) + } + if answer == nil { + t.Fatal("expected answer from route-selected retry") + } + if len(localIPs) != 2 { + t.Fatalf("exchange calls: got %d, want 2", len(localIPs)) + } + if !localIPs[0].Equal(tt.boundIP) { + t.Fatalf("first source: got %v, want %v", localIPs[0], tt.boundIP) + } + if localIPs[1] != nil { + t.Fatalf("retry source: got %v, want route-selected nil", localIPs[1]) + } + if len(servers) != 2 || servers[0] != tt.server || servers[1] != tt.server { + t.Fatalf("exchange servers: got %v, want two attempts to %s", servers, tt.server) + } + }) + } +} + +func Test_customDNSExchangeWith_PreservesReachableBoundSource(t *testing.T) { + msg := new(dns.Msg) + msg.SetQuestion("internal.example.", dns.TypeA) + boundIP := net.ParseIP("192.0.2.10") + calls := 0 + exchange := func(_ context.Context, msg *dns.Msg, _ string, localIP net.IP) (*dns.Msg, time.Duration, error) { + calls++ + if !localIP.Equal(boundIP) { + t.Fatalf("source: got %v, want %v", localIP, boundIP) + } + answer := new(dns.Msg) + answer.SetReply(msg) + return answer, time.Millisecond, nil + } + + answer, _, err := customDNSExchangeWith(context.Background(), msg, "192.0.2.53:53", boundIP, true, exchange) + if err != nil { + t.Fatal(err) + } + if answer == nil { + t.Fatal("expected answer from bound exchange") + } + if calls != 1 { + t.Fatalf("exchange calls: got %d, want 1", calls) + } +} + +func Test_customDNSExchangeWith_DoesNotRetryOtherFailures(t *testing.T) { + msg := new(dns.Msg) + msg.SetQuestion("internal.example.", dns.TypeA) + calls := 0 + exchange := func(_ context.Context, _ *dns.Msg, _ string, _ net.IP) (*dns.Msg, time.Duration, error) { + calls++ + return nil, 0, context.DeadlineExceeded + } + + _, _, err := customDNSExchangeWith(context.Background(), msg, "192.0.2.53:53", net.ParseIP("192.0.2.10"), true, exchange) + if err == nil { + t.Fatal("expected exchange failure") + } + if calls != 1 { + t.Fatalf("exchange calls: got %d, want 1", calls) + } +} + +func Test_customDNSExchangeWith_DoesNotRetryWithoutBoundSource(t *testing.T) { + msg := new(dns.Msg) + msg.SetQuestion("internal.example.", dns.TypeA) + calls := 0 + exchange := func(_ context.Context, _ *dns.Msg, _ string, _ net.IP) (*dns.Msg, time.Duration, error) { + calls++ + return nil, 0, &net.OpError{Op: "write", Net: "udp", Err: syscall.EHOSTUNREACH} + } + + _, _, err := customDNSExchangeWith(context.Background(), msg, "192.0.2.53:53", nil, true, exchange) + if err == nil { + t.Fatal("expected exchange failure") + } + if calls != 1 { + t.Fatalf("exchange calls: got %d, want 1", calls) + } +} + +func Test_customDNSExchangeWith_DoesNotRetryCanceledContext(t *testing.T) { + msg := new(dns.Msg) + msg.SetQuestion("internal.example.", dns.TypeA) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + calls := 0 + exchange := func(_ context.Context, _ *dns.Msg, _ string, _ net.IP) (*dns.Msg, time.Duration, error) { + calls++ + return nil, 0, &net.OpError{Op: "write", Net: "udp", Err: syscall.EHOSTUNREACH} + } + + _, _, err := customDNSExchangeWith(ctx, msg, "192.0.2.53:53", net.ParseIP("192.0.2.10"), true, exchange) + if err == nil { + t.Fatal("expected exchange failure") + } + if calls != 1 { + t.Fatalf("exchange calls: got %d, want 1", calls) + } +} + +func Test_customDNSExchangeWith_ReturnsUnboundRetryFailure(t *testing.T) { + msg := new(dns.Msg) + msg.SetQuestion("internal.example.", dns.TypeA) + retryErr := errors.New("route-selected exchange failed") + calls := 0 + exchange := func(_ context.Context, _ *dns.Msg, _ string, localIP net.IP) (*dns.Msg, time.Duration, error) { + calls++ + if localIP != nil { + return nil, 0, &net.OpError{Op: "write", Net: "udp", Err: syscall.EHOSTUNREACH} + } + return nil, 0, retryErr + } + + _, _, err := customDNSExchangeWith(context.Background(), msg, "192.0.2.53:53", net.ParseIP("192.0.2.10"), true, exchange) + if !errors.Is(err, retryErr) { + t.Fatalf("exchange error: got %v, want retry error %v", err, retryErr) + } + if calls != 2 { + t.Fatalf("exchange calls: got %d, want 2", calls) + } +} + +func Test_customDNSExchangeWith_DoesNotRetryReadSideUnreachable(t *testing.T) { + msg := new(dns.Msg) + msg.SetQuestion("internal.example.", dns.TypeA) + calls := 0 + exchange := func(_ context.Context, _ *dns.Msg, _ string, _ net.IP) (*dns.Msg, time.Duration, error) { + calls++ + return nil, 0, &net.OpError{Op: "read", Net: "udp", Err: syscall.EHOSTUNREACH} + } + + _, _, err := customDNSExchangeWith(context.Background(), msg, "192.0.2.53:53", net.ParseIP("192.0.2.10"), true, exchange) + if err == nil { + t.Fatal("expected exchange failure") + } + if calls != 1 { + t.Fatalf("exchange calls: got %d, want 1", calls) + } +} + +func Test_osResolver_ResolveUsesRouteSelectedFallbackForLANServer(t *testing.T) { + const server = "10.0.0.53:53" + boundIP := net.ParseIP("192.0.2.10") + resolver := newResolverWithNameserver([]string{server}) + resolver.localIP = func(string) net.IP { return boundIP } + + var localIPs []net.IP + var servers []string + resolver.exchangeDNS = func(_ context.Context, msg *dns.Msg, gotServer string, localIP net.IP) (*dns.Msg, time.Duration, error) { + servers = append(servers, gotServer) + localIPs = append(localIPs, append(net.IP(nil), localIP...)) + if localIP != nil { + return nil, 0, &net.OpError{Op: "write", Net: "udp", Err: syscall.EHOSTUNREACH} + } + answer := new(dns.Msg) + answer.SetReply(msg) + return answer, time.Millisecond, nil + } + + msg := new(dns.Msg) + msg.SetQuestion("internal.example.", dns.TypeA) + answer, err := resolver.Resolve(context.Background(), msg) + if err != nil { + t.Fatal(err) + } + if answer == nil { + t.Fatal("expected answer from route-selected retry") + } + if len(localIPs) != 2 || !localIPs[0].Equal(boundIP) || localIPs[1] != nil { + t.Fatalf("exchange sources: got %v, want [%v ]", localIPs, boundIP) + } + if len(servers) != 2 || servers[0] != server || servers[1] != server { + t.Fatalf("exchange servers: got %v, want two attempts to %s", servers, server) + } +} + +// A VPN-pushed public DNS address is categorized as public by IP, but it is +// still a system-selected resolver and must get the same route-compatible retry. +func Test_osResolver_ResolveUsesRouteSelectedFallbackForPublicVPNServer(t *testing.T) { + const server = "192.0.2.53:53" + boundIP := net.ParseIP("198.51.100.10") + resolver := newResolverWithNameserver([]string{server}) + resolver.localIP = func(string) net.IP { return boundIP } + + var localIPs []net.IP + resolver.exchangeDNS = func(_ context.Context, msg *dns.Msg, _ string, localIP net.IP) (*dns.Msg, time.Duration, error) { + localIPs = append(localIPs, append(net.IP(nil), localIP...)) + if localIP != nil { + return nil, 0, &net.OpError{Op: "write", Net: "udp", Err: syscall.EHOSTUNREACH} + } + answer := new(dns.Msg) + answer.SetReply(msg) + return answer, time.Millisecond, nil + } + + msg := new(dns.Msg) + msg.SetQuestion("internal.example.", dns.TypeA) + answer, err := resolver.Resolve(context.Background(), msg) + if err != nil { + t.Fatal(err) + } + if answer == nil { + t.Fatal("expected answer from route-selected retry") + } + if len(localIPs) != 2 || !localIPs[0].Equal(boundIP) || localIPs[1] != nil { + t.Fatalf("exchange sources: got %v, want [%v ]", localIPs, boundIP) + } +} + +func Test_osResolver_ResolveDoesNotRetrySyntheticControlDFallbackUnbound(t *testing.T) { + resolver := newResolverWithNameserver([]string{controldPublicDnsWithPort}) + resolver.localIP = func(string) net.IP { return net.ParseIP("198.51.100.10") } + calls := 0 + resolver.exchangeDNS = func(_ context.Context, _ *dns.Msg, _ string, _ net.IP) (*dns.Msg, time.Duration, error) { + calls++ + return nil, 0, &net.OpError{Op: "write", Net: "udp", Err: syscall.EHOSTUNREACH} + } + + msg := new(dns.Msg) + msg.SetQuestion("internal.example.", dns.TypeA) + _, err := resolver.Resolve(context.Background(), msg) + if err == nil { + t.Fatal("expected exchange failure") + } + if calls != 1 { + t.Fatalf("exchange calls: got %d, want one bound synthetic fallback attempt", calls) + } +} + func Test_osResolver_ResolveWithNonSuccessAnswer(t *testing.T) { // Set up a LAN nameserver that returns a success response. lanPC, err := net.ListenPacket("udp", "127.0.0.1:0") // 127.0.0.1 is considered LAN (loopback)