From e659350eaf6e104e69982b72c1ac9174fe32442d Mon Sep 17 00:00:00 2001 From: Viktor Liu Date: Wed, 17 Jun 2026 14:30:23 +0200 Subject: [PATCH] Split record lookups into per-type helpers and reduce forwarder param count --- client/internal/dns/resutil/resolve.go | 196 ++++++++++++++----------- client/internal/dnsfwd/forwarder.go | 36 ++--- 2 files changed, 128 insertions(+), 104 deletions(-) diff --git a/client/internal/dns/resutil/resolve.go b/client/internal/dns/resutil/resolve.go index e2c321477..e80e175ce 100644 --- a/client/internal/dns/resutil/resolve.go +++ b/client/internal/dns/resutil/resolve.go @@ -211,103 +211,125 @@ func LookupRecords(ctx context.Context, r RecordResolver, name string, qtype uin switch qtype { case dns.TypeMX: - recs, err := r.LookupMX(ctx, name) - if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) - } - rrs := make([]dns.RR, 0, len(recs)) - for _, mx := range recs { - rrs = append(rrs, &dns.MX{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypeMX, Class: dns.ClassINET, Ttl: ttl}, - Preference: mx.Pref, - Mx: dns.Fqdn(mx.Host), - }) - } - return rrs, dns.RcodeSuccess - + return lookupMX(ctx, r, name, fqdn, ttl) case dns.TypeTXT: - recs, err := r.LookupTXT(ctx, name) - if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) - } - rrs := make([]dns.RR, 0, len(recs)) - for _, txt := range recs { - rrs = append(rrs, &dns.TXT{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypeTXT, Class: dns.ClassINET, Ttl: ttl}, - Txt: chunkTXT(txt), - }) - } - return rrs, dns.RcodeSuccess - + return lookupTXT(ctx, r, name, fqdn, ttl) case dns.TypeNS: - recs, err := r.LookupNS(ctx, name) - if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) - } - rrs := make([]dns.RR, 0, len(recs)) - for _, ns := range recs { - rrs = append(rrs, &dns.NS{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: ttl}, - Ns: dns.Fqdn(ns.Host), - }) - } - return rrs, dns.RcodeSuccess - + return lookupNS(ctx, r, name, fqdn, ttl) case dns.TypeSRV: - _, recs, err := r.LookupSRV(ctx, "", "", name) - if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) - } - rrs := make([]dns.RR, 0, len(recs)) - for _, srv := range recs { - rrs = append(rrs, &dns.SRV{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypeSRV, Class: dns.ClassINET, Ttl: ttl}, - Priority: srv.Priority, - Weight: srv.Weight, - Port: srv.Port, - Target: dns.Fqdn(srv.Target), - }) - } - return rrs, dns.RcodeSuccess - + return lookupSRV(ctx, r, name, fqdn, ttl) case dns.TypeCNAME: - cname, err := r.LookupCNAME(ctx, name) - if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) - } - // LookupCNAME returns the queried name itself when the name resolves - // but has no CNAME record; that is a NODATA result, not a CNAME. - if strings.EqualFold(dns.Fqdn(cname), fqdn) { - return nil, dns.RcodeSuccess - } - return []dns.RR{&dns.CNAME{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: ttl}, - Target: dns.Fqdn(cname), - }}, dns.RcodeSuccess - + return lookupCNAME(ctx, r, name, fqdn, ttl) case dns.TypePTR: - addr, ok := ptrQueryAddr(name) - if !ok { - return nil, dns.RcodeSuccess - } - names, err := r.LookupAddr(ctx, addr) - if err != nil { - return nil, rcodeForRecordError(ctx, r, name, err) - } - rrs := make([]dns.RR, 0, len(names)) - for _, n := range names { - rrs = append(rrs, &dns.PTR{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypePTR, Class: dns.ClassINET, Ttl: ttl}, - Ptr: dns.Fqdn(n), - }) - } - return rrs, dns.RcodeSuccess - + return lookupPTR(ctx, r, name, fqdn, ttl) default: return nil, dns.RcodeSuccess } } +func recordHeader(fqdn string, rrtype uint16, ttl uint32) dns.RR_Header { + return dns.RR_Header{Name: fqdn, Rrtype: rrtype, Class: dns.ClassINET, Ttl: ttl} +} + +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) + } + rrs := make([]dns.RR, 0, len(recs)) + for _, mx := range recs { + rrs = append(rrs, &dns.MX{ + Hdr: recordHeader(fqdn, dns.TypeMX, ttl), + Preference: mx.Pref, + Mx: dns.Fqdn(mx.Host), + }) + } + return rrs, dns.RcodeSuccess +} + +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) + } + rrs := make([]dns.RR, 0, len(recs)) + for _, txt := range recs { + rrs = append(rrs, &dns.TXT{ + Hdr: recordHeader(fqdn, dns.TypeTXT, ttl), + Txt: chunkTXT(txt), + }) + } + return rrs, dns.RcodeSuccess +} + +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) + } + rrs := make([]dns.RR, 0, len(recs)) + for _, ns := range recs { + rrs = append(rrs, &dns.NS{ + Hdr: recordHeader(fqdn, dns.TypeNS, ttl), + Ns: dns.Fqdn(ns.Host), + }) + } + return rrs, dns.RcodeSuccess +} + +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) + } + rrs := make([]dns.RR, 0, len(recs)) + for _, srv := range recs { + rrs = append(rrs, &dns.SRV{ + Hdr: recordHeader(fqdn, dns.TypeSRV, ttl), + Priority: srv.Priority, + Weight: srv.Weight, + Port: srv.Port, + Target: dns.Fqdn(srv.Target), + }) + } + return rrs, dns.RcodeSuccess +} + +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) + } + // LookupCNAME returns the queried name itself when the name resolves but + // has no CNAME record; that is a NODATA result, not a CNAME. + if strings.EqualFold(dns.Fqdn(cname), fqdn) { + return nil, dns.RcodeSuccess + } + return []dns.RR{&dns.CNAME{ + Hdr: recordHeader(fqdn, dns.TypeCNAME, ttl), + Target: dns.Fqdn(cname), + }}, dns.RcodeSuccess +} + +func lookupPTR(ctx context.Context, r RecordResolver, name, fqdn string, ttl uint32) ([]dns.RR, int) { + addr, ok := ptrQueryAddr(name) + if !ok { + return nil, dns.RcodeSuccess + } + names, err := r.LookupAddr(ctx, addr) + if err != nil { + return nil, rcodeForRecordError(ctx, r, name, err) + } + rrs := make([]dns.RR, 0, len(names)) + for _, n := range names { + rrs = append(rrs, &dns.PTR{ + Hdr: recordHeader(fqdn, dns.TypePTR, ttl), + Ptr: dns.Fqdn(n), + }) + } + return rrs, dns.RcodeSuccess +} + // ptrQueryAddr converts a reverse-DNS query name (in-addr.arpa or ip6.arpa) // into the address string expected by net.Resolver.LookupAddr. It reports false // when the name is not a well-formed reverse name. diff --git a/client/internal/dnsfwd/forwarder.go b/client/internal/dnsfwd/forwarder.go index 09a0129fa..905c53be5 100644 --- a/client/internal/dnsfwd/forwarder.go +++ b/client/internal/dnsfwd/forwarder.go @@ -220,9 +220,9 @@ func (f *DNSForwarder) handleDNSQuery(logger *log.Entry, w dns.ResponseWriter, q switch question.Qtype { case dns.TypeA, dns.TypeAAAA: - f.handleAddressQuery(ctx, logger, w, question, resp, qname, mostSpecificResId, matchingEntries, startTime) + f.handleAddressQuery(ctx, logger, w, resp, mostSpecificResId, matchingEntries, startTime) case dns.TypeMX, dns.TypeTXT, dns.TypeNS, dns.TypeSRV, dns.TypeCNAME, dns.TypePTR: - f.handleRecordQuery(ctx, logger, w, question, resp, qname, startTime) + f.handleRecordQuery(ctx, logger, w, resp, startTime) default: // The domain is routed here, so any other type is answered NODATA // (NOERROR, empty answer) rather than falling back to a resolver that @@ -237,30 +237,20 @@ func (f *DNSForwarder) handleDNSQuery(logger *log.Entry, w dns.ResponseWriter, q } } -// attachEDE adds an Extended DNS Error (RFC 8914) option to the response, -// creating the OPT pseudo-record if the response does not already carry one. -func attachEDE(resp *dns.Msg, code uint16, text string) { - opt := resp.IsEdns0() - if opt == nil { - resp.SetEdns0(dns.DefaultMsgSize, false) - opt = resp.IsEdns0() - } - opt.Option = append(opt.Option, &dns.EDNS0_EDE{InfoCode: code, ExtraText: text}) -} - // handleAddressQuery resolves A/AAAA queries, programs the firewall sets and // resolved-IP state, and caches the answer for resilience on upstream failure. func (f *DNSForwarder) handleAddressQuery( ctx context.Context, logger *log.Entry, w dns.ResponseWriter, - question dns.Question, resp *dns.Msg, - qname string, mostSpecificResId route.ResID, matchingEntries []*ForwarderEntry, startTime time.Time, ) { + question := resp.Question[0] + qname := strings.ToLower(question.Name) + network := resutil.NetworkForQtype(question.Qtype) result := resutil.LookupIP(ctx, f.resolver, network, qname, question.Qtype) if result.Err != nil { @@ -282,11 +272,12 @@ func (f *DNSForwarder) handleRecordQuery( ctx context.Context, logger *log.Entry, w dns.ResponseWriter, - question dns.Question, resp *dns.Msg, - qname string, startTime time.Time, ) { + question := resp.Question[0] + qname := strings.ToLower(question.Name) + records, rcode := resutil.LookupRecords(ctx, f.resolver, qname, question.Qtype, f.ttl) resp.Rcode = rcode resp.Answer = append(resp.Answer, records...) @@ -476,3 +467,14 @@ func (f *DNSForwarder) getMatchingEntries(domain string) (route.ResID, []*Forwar return selectedResId, matches } + +// attachEDE adds an Extended DNS Error (RFC 8914) option to the response, +// creating the OPT pseudo-record if the response does not already carry one. +func attachEDE(resp *dns.Msg, code uint16, text string) { + opt := resp.IsEdns0() + if opt == nil { + resp.SetEdns0(dns.DefaultMsgSize, false) + opt = resp.IsEdns0() + } + opt.Option = append(opt.Option, &dns.EDNS0_EDE{InfoCode: code, ExtraText: text}) +}