From 3a6852cbc2c365c7bfc9d0b5f74190b8384d0b3f Mon Sep 17 00:00:00 2001 From: Viktor Liu Date: Wed, 17 Jun 2026 15:11:19 +0200 Subject: [PATCH] Answer NODATA for missing non-address records and cover record lookups with tests --- client/internal/dns/resutil/resolve.go | 54 ++------- client/internal/dns/resutil/resolve_test.go | 127 ++++++++++++++++++++ client/internal/dnsfwd/forwarder_test.go | 87 +++++++++++--- 3 files changed, 208 insertions(+), 60 deletions(-) diff --git a/client/internal/dns/resutil/resolve.go b/client/internal/dns/resutil/resolve.go index e80e175ce..ab2721d7a 100644 --- a/client/internal/dns/resutil/resolve.go +++ b/client/internal/dns/resutil/resolve.go @@ -189,8 +189,7 @@ func getRcodeForNotFound(ctx context.Context, r resolver, domain string, origina } // RecordResolver is the host resolver surface used to forward non-address -// record queries. net.DefaultResolver satisfies it. The LookupNetIP method is -// used only to probe name existence when distinguishing NXDOMAIN from NODATA. +// record queries. net.DefaultResolver satisfies it. type RecordResolver interface { LookupMX(ctx context.Context, name string) ([]*net.MX, error) LookupTXT(ctx context.Context, name string) ([]string, error) @@ -198,7 +197,6 @@ type RecordResolver interface { LookupSRV(ctx context.Context, service, proto, name string) (string, []*net.SRV, error) LookupCNAME(ctx context.Context, host string) (string, error) LookupAddr(ctx context.Context, addr string) ([]string, error) - LookupNetIP(ctx context.Context, network, host string) ([]netip.Addr, error) } // LookupRecords resolves a non-address DNS record type through the host @@ -234,7 +232,7 @@ func recordHeader(fqdn string, rrtype uint16, ttl uint32) dns.RR_Header { func lookupMX(ctx context.Context, r RecordResolver, name, fqdn string, ttl uint32) ([]dns.RR, int) { recs, err := r.LookupMX(ctx, name) if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) + return nil, rcodeForRecordError(err) } rrs := make([]dns.RR, 0, len(recs)) for _, mx := range recs { @@ -250,7 +248,7 @@ func lookupMX(ctx context.Context, r RecordResolver, name, fqdn string, ttl uint func lookupTXT(ctx context.Context, r RecordResolver, name, fqdn string, ttl uint32) ([]dns.RR, int) { recs, err := r.LookupTXT(ctx, name) if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) + return nil, rcodeForRecordError(err) } rrs := make([]dns.RR, 0, len(recs)) for _, txt := range recs { @@ -265,7 +263,7 @@ func lookupTXT(ctx context.Context, r RecordResolver, name, fqdn string, ttl uin func lookupNS(ctx context.Context, r RecordResolver, name, fqdn string, ttl uint32) ([]dns.RR, int) { recs, err := r.LookupNS(ctx, name) if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) + return nil, rcodeForRecordError(err) } rrs := make([]dns.RR, 0, len(recs)) for _, ns := range recs { @@ -280,7 +278,7 @@ func lookupNS(ctx context.Context, r RecordResolver, name, fqdn string, ttl uint func lookupSRV(ctx context.Context, r RecordResolver, name, fqdn string, ttl uint32) ([]dns.RR, int) { _, recs, err := r.LookupSRV(ctx, "", "", name) if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) + return nil, rcodeForRecordError(err) } rrs := make([]dns.RR, 0, len(recs)) for _, srv := range recs { @@ -298,7 +296,7 @@ func lookupSRV(ctx context.Context, r RecordResolver, name, fqdn string, ttl uin func lookupCNAME(ctx context.Context, r RecordResolver, name, fqdn string, ttl uint32) ([]dns.RR, int) { cname, err := r.LookupCNAME(ctx, name) if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) + return nil, rcodeForRecordError(err) } // LookupCNAME returns the queried name itself when the name resolves but // has no CNAME record; that is a NODATA result, not a CNAME. @@ -318,7 +316,7 @@ func lookupPTR(ctx context.Context, r RecordResolver, name, fqdn string, ttl uin } names, err := r.LookupAddr(ctx, addr) if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) + return nil, rcodeForRecordError(err) } rrs := make([]dns.RR, 0, len(names)) for _, n := range names { @@ -374,41 +372,15 @@ func ptrQueryAddr(qname string) (string, bool) { } // rcodeForRecordError maps a non-address lookup error to a DNS rcode. A -// not-found result is disambiguated into NXDOMAIN or NODATA by probing whether -// the name has any address record. An inconclusive probe (timeout, server -// failure) is treated as existing so a transient failure cannot poison the name -// with NXDOMAIN. -func rcodeForRecordError(ctx context.Context, r RecordResolver, name string, err error) int { - var dnsErr *net.DNSError - if !errors.As(err, &dnsErr) || !dnsErr.IsNotFound { - return dns.RcodeServerFailure - } - - if nameHasAddress(ctx, r, name) { - return dns.RcodeSuccess - } - return dns.RcodeNameError -} - -// nameHasAddress reports whether the name resolves to any address record. It is -// used to distinguish NXDOMAIN from NODATA; inconclusive lookups report true to -// avoid emitting a poisoning NXDOMAIN. -func nameHasAddress(ctx context.Context, r RecordResolver, name string) bool { - _, err := r.LookupNetIP(ctx, "ip", name) - if err == nil { - return true - } - - var addrErr *net.AddrError - if errors.As(err, &addrErr) && addrErr.Err == errNoSuitableAddress { - return true - } - +// not-found result becomes NODATA rather than NXDOMAIN: net.DNSError.IsNotFound +// does not distinguish a missing name from a name that exists only with records +// of other types, so the name cannot be proven absent and must not be poisoned. +func rcodeForRecordError(err error) int { var dnsErr *net.DNSError if errors.As(err, &dnsErr) && dnsErr.IsNotFound { - return false + return dns.RcodeSuccess } - return true + return dns.RcodeServerFailure } // chunkTXT splits a TXT string into character-strings no longer than 255 bytes diff --git a/client/internal/dns/resutil/resolve_test.go b/client/internal/dns/resutil/resolve_test.go index 8ffca2dff..3384182d6 100644 --- a/client/internal/dns/resutil/resolve_test.go +++ b/client/internal/dns/resutil/resolve_test.go @@ -5,6 +5,7 @@ import ( "errors" "net" "net/netip" + "strings" "testing" "github.com/miekg/dns" @@ -152,3 +153,129 @@ func TestPtrQueryAddr(t *testing.T) { }) } } + +type mockRecordResolver struct { + mx []*net.MX + txt []string + ns []*net.NS + srv []*net.SRV + cname string + ptr []string + err error +} + +func (m *mockRecordResolver) LookupMX(context.Context, string) ([]*net.MX, error) { + return m.mx, m.err +} +func (m *mockRecordResolver) LookupTXT(context.Context, string) ([]string, error) { + return m.txt, m.err +} +func (m *mockRecordResolver) LookupNS(context.Context, string) ([]*net.NS, error) { + return m.ns, m.err +} +func (m *mockRecordResolver) LookupSRV(context.Context, string, string, string) (string, []*net.SRV, error) { + return "", m.srv, m.err +} +func (m *mockRecordResolver) LookupCNAME(context.Context, string) (string, error) { + return m.cname, m.err +} +func (m *mockRecordResolver) LookupAddr(context.Context, string) ([]string, error) { + return m.ptr, m.err +} + +func TestLookupRecords(t *testing.T) { + notFound := &net.DNSError{IsNotFound: true, Name: "example.com."} + + t.Run("MX success", func(t *testing.T) { + r := &mockRecordResolver{mx: []*net.MX{{Host: "mail.example.com.", Pref: 10}}} + rrs, rcode := LookupRecords(context.Background(), r, "example.com.", dns.TypeMX, 300) + assert.Equal(t, dns.RcodeSuccess, rcode) + require.Len(t, rrs, 1) + assert.Equal(t, "mail.example.com.", rrs[0].(*dns.MX).Mx) + }) + + t.Run("TXT short string is one character-string", func(t *testing.T) { + r := &mockRecordResolver{txt: []string{"v=spf1 -all"}} + rrs, rcode := LookupRecords(context.Background(), r, "example.com.", dns.TypeTXT, 300) + assert.Equal(t, dns.RcodeSuccess, rcode) + require.Len(t, rrs, 1) + assert.Equal(t, []string{"v=spf1 -all"}, rrs[0].(*dns.TXT).Txt) + }) + + t.Run("TXT chunks long strings", func(t *testing.T) { + long := strings.Repeat("a", 300) + r := &mockRecordResolver{txt: []string{long}} + rrs, rcode := LookupRecords(context.Background(), r, "example.com.", dns.TypeTXT, 300) + assert.Equal(t, dns.RcodeSuccess, rcode) + require.Len(t, rrs, 1) + txt := rrs[0].(*dns.TXT).Txt + require.Len(t, txt, 2, "300-byte string should split into two character-strings") + assert.Equal(t, 255, len(txt[0])) + assert.Equal(t, 45, len(txt[1])) + }) + + t.Run("NS success", func(t *testing.T) { + r := &mockRecordResolver{ns: []*net.NS{{Host: "ns1.example.com."}}} + rrs, rcode := LookupRecords(context.Background(), r, "example.com.", dns.TypeNS, 300) + assert.Equal(t, dns.RcodeSuccess, rcode) + require.Len(t, rrs, 1) + assert.Equal(t, "ns1.example.com.", rrs[0].(*dns.NS).Ns) + }) + + t.Run("SRV success", func(t *testing.T) { + r := &mockRecordResolver{srv: []*net.SRV{{Target: "sip.example.com.", Port: 5060}}} + rrs, rcode := LookupRecords(context.Background(), r, "_sip._tcp.example.com.", dns.TypeSRV, 300) + assert.Equal(t, dns.RcodeSuccess, rcode) + require.Len(t, rrs, 1) + assert.Equal(t, uint16(5060), rrs[0].(*dns.SRV).Port) + }) + + t.Run("CNAME success", func(t *testing.T) { + r := &mockRecordResolver{cname: "target.example.com."} + rrs, rcode := LookupRecords(context.Background(), r, "www.example.com.", dns.TypeCNAME, 300) + assert.Equal(t, dns.RcodeSuccess, rcode) + require.Len(t, rrs, 1) + assert.Equal(t, "target.example.com.", rrs[0].(*dns.CNAME).Target) + }) + + t.Run("CNAME equal to name is NODATA", func(t *testing.T) { + r := &mockRecordResolver{cname: "example.com."} + rrs, rcode := LookupRecords(context.Background(), r, "example.com.", dns.TypeCNAME, 300) + assert.Equal(t, dns.RcodeSuccess, rcode) + assert.Empty(t, rrs, "self-referential CNAME is NODATA") + }) + + t.Run("PTR success", func(t *testing.T) { + r := &mockRecordResolver{ptr: []string{"host.example.com."}} + rrs, rcode := LookupRecords(context.Background(), r, "4.3.2.1.in-addr.arpa.", dns.TypePTR, 300) + assert.Equal(t, dns.RcodeSuccess, rcode) + require.Len(t, rrs, 1) + assert.Equal(t, "host.example.com.", rrs[0].(*dns.PTR).Ptr) + }) + + t.Run("PTR malformed name is NODATA", func(t *testing.T) { + r := &mockRecordResolver{} + rrs, rcode := LookupRecords(context.Background(), r, "example.com.", dns.TypePTR, 300) + assert.Equal(t, dns.RcodeSuccess, rcode) + assert.Empty(t, rrs) + }) + + t.Run("not found is NODATA never NXDOMAIN", func(t *testing.T) { + r := &mockRecordResolver{err: notFound} + _, rcode := LookupRecords(context.Background(), r, "example.com.", dns.TypeMX, 300) + assert.Equal(t, dns.RcodeSuccess, rcode, "missing record must not poison the name") + }) + + t.Run("server failure maps to SERVFAIL", func(t *testing.T) { + r := &mockRecordResolver{err: &net.DNSError{Err: "server misbehaving", IsTemporary: true}} + _, rcode := LookupRecords(context.Background(), r, "example.com.", dns.TypeMX, 300) + assert.Equal(t, dns.RcodeServerFailure, rcode) + }) + + t.Run("unsupported type is NODATA", func(t *testing.T) { + r := &mockRecordResolver{} + rrs, rcode := LookupRecords(context.Background(), r, "example.com.", dns.TypeCAA, 300) + assert.Equal(t, dns.RcodeSuccess, rcode) + assert.Empty(t, rrs) + }) +} diff --git a/client/internal/dnsfwd/forwarder_test.go b/client/internal/dnsfwd/forwarder_test.go index 62e3bde2a..32b6c770a 100644 --- a/client/internal/dnsfwd/forwarder_test.go +++ b/client/internal/dnsfwd/forwarder_test.go @@ -134,22 +134,26 @@ func (m *MockResolver) LookupNetIP(ctx context.Context, network, host string) ([ func (m *MockResolver) LookupMX(ctx context.Context, name string) ([]*net.MX, error) { args := m.Called(ctx, name) - return args.Get(0).([]*net.MX), args.Error(1) + recs, _ := args.Get(0).([]*net.MX) + return recs, args.Error(1) } func (m *MockResolver) LookupTXT(ctx context.Context, name string) ([]string, error) { args := m.Called(ctx, name) - return args.Get(0).([]string), args.Error(1) + recs, _ := args.Get(0).([]string) + return recs, args.Error(1) } func (m *MockResolver) LookupNS(ctx context.Context, name string) ([]*net.NS, error) { args := m.Called(ctx, name) - return args.Get(0).([]*net.NS), args.Error(1) + recs, _ := args.Get(0).([]*net.NS) + return recs, args.Error(1) } func (m *MockResolver) LookupSRV(ctx context.Context, service, proto, name string) (string, []*net.SRV, error) { args := m.Called(ctx, service, proto, name) - return args.String(0), args.Get(1).([]*net.SRV), args.Error(2) + recs, _ := args.Get(1).([]*net.SRV) + return args.String(0), recs, args.Error(2) } func (m *MockResolver) LookupCNAME(ctx context.Context, host string) (string, error) { @@ -159,7 +163,8 @@ func (m *MockResolver) LookupCNAME(ctx context.Context, host string) (string, er func (m *MockResolver) LookupAddr(ctx context.Context, addr string) ([]string, error) { args := m.Called(ctx, addr) - return args.Get(0).([]string), args.Error(1) + recs, _ := args.Get(0).([]string) + return recs, args.Error(1) } func TestDNSForwarder_SubdomainAccessLogic(t *testing.T) { @@ -676,34 +681,78 @@ func TestDNSForwarder_RecordQueries(t *testing.T) { mockResolver.AssertExpectations(t) }) - t.Run("missing MX on existing name is NODATA", func(t *testing.T) { + t.Run("missing MX is NODATA not NXDOMAIN", func(t *testing.T) { mockResolver := &MockResolver{} forwarder := newRecordTestForwarder(t, mockResolver, "example.com") + // A not-found cannot prove the name is absent (it may exist with only + // other record types), so it must answer NODATA, never NXDOMAIN. mockResolver.On("LookupMX", mock.Anything, "example.com."). - Return([]*net.MX{}, notFound).Once() - // Existence probe: the name resolves to an address, so the name exists. - mockResolver.On("LookupNetIP", mock.Anything, "ip", "example.com."). - Return([]netip.Addr{netip.MustParseAddr("1.2.3.4")}, nil).Once() + Return(nil, notFound).Once() resp := runRecordQuery(t, forwarder, "example.com", dns.TypeMX) - assert.Equal(t, dns.RcodeSuccess, resp.Rcode, "existing name with no MX must be NODATA") + assert.Equal(t, dns.RcodeSuccess, resp.Rcode, "missing record must be NODATA") assert.Empty(t, resp.Answer) mockResolver.AssertExpectations(t) }) - t.Run("missing MX on nonexistent name is NXDOMAIN", func(t *testing.T) { + t.Run("NS records are forwarded", func(t *testing.T) { mockResolver := &MockResolver{} forwarder := newRecordTestForwarder(t, mockResolver, "example.com") - mockResolver.On("LookupMX", mock.Anything, "example.com."). - Return([]*net.MX{}, notFound).Once() - // Existence probe: the name has no address either, so it truly doesn't exist. - mockResolver.On("LookupNetIP", mock.Anything, "ip", "example.com."). - Return([]netip.Addr{}, notFound).Once() + mockResolver.On("LookupNS", mock.Anything, "example.com."). + Return([]*net.NS{{Host: "ns1.example.com."}}, nil).Once() - resp := runRecordQuery(t, forwarder, "example.com", dns.TypeMX) - assert.Equal(t, dns.RcodeNameError, resp.Rcode, "nonexistent name may be NXDOMAIN") + resp := runRecordQuery(t, forwarder, "example.com", dns.TypeNS) + require.Equal(t, dns.RcodeSuccess, resp.Rcode) + require.Len(t, resp.Answer, 1) + ns, ok := resp.Answer[0].(*dns.NS) + require.True(t, ok, "answer should be an NS record") + assert.Equal(t, "ns1.example.com.", ns.Ns) + mockResolver.AssertExpectations(t) + }) + + t.Run("missing NS is NODATA", func(t *testing.T) { + mockResolver := &MockResolver{} + forwarder := newRecordTestForwarder(t, mockResolver, "example.com") + + mockResolver.On("LookupNS", mock.Anything, "example.com."). + Return(nil, notFound).Once() + + resp := runRecordQuery(t, forwarder, "example.com", dns.TypeNS) + assert.Equal(t, dns.RcodeSuccess, resp.Rcode) + assert.Empty(t, resp.Answer) + mockResolver.AssertExpectations(t) + }) + + t.Run("SRV records are forwarded", func(t *testing.T) { + mockResolver := &MockResolver{} + forwarder := newRecordTestForwarder(t, mockResolver, "_sip._tcp.example.com") + + mockResolver.On("LookupSRV", mock.Anything, "", "", "_sip._tcp.example.com."). + Return("", []*net.SRV{{Target: "sip.example.com.", Port: 5060, Priority: 10, Weight: 5}}, nil).Once() + + resp := runRecordQuery(t, forwarder, "_sip._tcp.example.com", dns.TypeSRV) + require.Equal(t, dns.RcodeSuccess, resp.Rcode) + require.Len(t, resp.Answer, 1) + srv, ok := resp.Answer[0].(*dns.SRV) + require.True(t, ok, "answer should be an SRV record") + assert.Equal(t, "sip.example.com.", srv.Target) + assert.Equal(t, uint16(5060), srv.Port) + assert.Equal(t, uint16(10), srv.Priority) + mockResolver.AssertExpectations(t) + }) + + t.Run("missing SRV is NODATA", func(t *testing.T) { + mockResolver := &MockResolver{} + forwarder := newRecordTestForwarder(t, mockResolver, "_sip._tcp.example.com") + + mockResolver.On("LookupSRV", mock.Anything, "", "", "_sip._tcp.example.com."). + Return("", nil, notFound).Once() + + resp := runRecordQuery(t, forwarder, "_sip._tcp.example.com", dns.TypeSRV) + assert.Equal(t, dns.RcodeSuccess, resp.Rcode) + assert.Empty(t, resp.Answer) mockResolver.AssertExpectations(t) })