Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
141 changes: 75 additions & 66 deletions src/zdns/lookup.go
Original file line number Diff line number Diff line change
Expand Up @@ -891,25 +891,15 @@ func (r *Resolver) cachedLookup(ctx context.Context, q Question, nameServer *Nam
var result *SingleQueryResult
var rawResp *dns.Msg
var status Status
if r.dnsOverHTTPSEnabled {
r.verboseLog(depth, "****WIRE LOOKUP*** ", DoHProtocol, " ", dns.TypeToString[q.Type], " ", q.Name, " ", nameServer)
result, rawResp, status, err = doDoHLookup(lookupCtx, connInfo.httpsClient, q, nameServer, requestIteration, r.ednsOptions, r.dnsSecEnabled, r.checkingDisabledBit, r.authenticatedDataBit)
} else if r.dnsOverTLSEnabled {
r.verboseLog(depth, "****WIRE LOOKUP*** ", DoTProtocol, " ", dns.TypeToString[q.Type], " ", q.Name, " ", nameServer)
result, rawResp, status, err = doDoTLookup(lookupCtx, connInfo, q, nameServer, r.rootCAs, r.verifyServerCert, requestIteration, r.ednsOptions, r.dnsSecEnabled, r.checkingDisabledBit, r.authenticatedDataBit)
} else if connInfo.udpClient != nil {
r.verboseLog(depth, "****WIRE LOOKUP*** ", UDPProtocol, " ", dns.TypeToString[q.Type], " ", q.Name, " ", nameServer)
result, rawResp, status, err = wireLookupUDP(lookupCtx, connInfo, q, nameServer, r.ednsOptions, requestIteration, r.dnsSecEnabled, r.checkingDisabledBit, r.authenticatedDataBit)
if status == StatusTruncated && connInfo.tcpClient != nil {
// result truncated, try again with TCP
r.verboseLog(depth, "****WIRE LOOKUP*** ", TCPProtocol, " ", dns.TypeToString[q.Type], " ", q.Name, " ", nameServer)
result, rawResp, status, err = wireLookupTCP(lookupCtx, connInfo, q, nameServer, r.ednsOptions, requestIteration, r.dnsSecEnabled, r.checkingDisabledBit, r.authenticatedDataBit)
result, rawResp, status, err = r.performWireLookup(lookupCtx, connInfo, q, nameServer, requestIteration, depth, true)
if err == nil && shouldRetryWithoutEDNS(rawResp, r.ednsOptions, r.dnsSecEnabled) {
fallbackCtx, fallbackCancel := context.WithTimeout(ctx, r.networkTimeout)
defer fallbackCancel()

if rateLimitErr := r.rateLimit.wait(fallbackCtx, *nameServer); rateLimitErr != nil {
return &SingleQueryResult{}, false, StatusError, trace, fmt.Errorf("rate limiter error for EDNS fallback against nameserver %s: %w", nameServer, rateLimitErr)
}
} else if connInfo.tcpClient != nil {
r.verboseLog(depth, "****WIRE LOOKUP*** ", TCPProtocol, " ", dns.TypeToString[q.Type], " ", q.Name, " ", nameServer)
result, rawResp, status, err = wireLookupTCP(lookupCtx, connInfo, q, nameServer, r.ednsOptions, requestIteration, r.dnsSecEnabled, r.checkingDisabledBit, r.authenticatedDataBit)
} else {
return &SingleQueryResult{}, false, StatusError, trace, errors.New("no connection info for nameserver")
result, rawResp, status, err = r.performWireLookup(fallbackCtx, connInfo, q, nameServer, requestIteration, depth, false)
}

if err != nil {
Expand Down Expand Up @@ -943,23 +933,44 @@ func (r *Resolver) cachedLookup(ctx context.Context, q Question, nameServer *Nam
return result, isCached, status, trace, err
}

func doDoTLookup(ctx context.Context, connInfo *ConnectionInfo, q Question, nameServer *NameServer, rootCAs *x509.CertPool, shouldVerifyServerCert, recursive bool, ednsOptions []dns.EDNS0, dnssec, checkingDisabled, authenticatedData bool) (*SingleQueryResult, *dns.Msg, Status, error) {
func (r *Resolver) performWireLookup(ctx context.Context, connInfo *ConnectionInfo, q Question, nameServer *NameServer, recursive bool, depth int, useEDNS bool) (*SingleQueryResult, *dns.Msg, Status, error) {
logPrefix := "****WIRE LOOKUP*** "
if !useEDNS {
logPrefix = "****WIRE LOOKUP WITHOUT EDNS*** "
}

if r.dnsOverHTTPSEnabled {
r.verboseLog(depth, logPrefix, DoHProtocol, " ", dns.TypeToString[q.Type], " ", q.Name, " ", nameServer)
return doDoHLookup(ctx, connInfo.httpsClient, q, nameServer, recursive, r.ednsOptions, r.dnsSecEnabled, r.checkingDisabledBit, r.authenticatedDataBit, useEDNS)
}
if r.dnsOverTLSEnabled {
r.verboseLog(depth, logPrefix, DoTProtocol, " ", dns.TypeToString[q.Type], " ", q.Name, " ", nameServer)
return doDoTLookup(ctx, connInfo, q, nameServer, r.rootCAs, r.verifyServerCert, recursive, r.ednsOptions, r.dnsSecEnabled, r.checkingDisabledBit, r.authenticatedDataBit, useEDNS)
}
if connInfo.udpClient != nil {
r.verboseLog(depth, logPrefix, UDPProtocol, " ", dns.TypeToString[q.Type], " ", q.Name, " ", nameServer)
result, rawResp, status, err := wireLookupUDP(ctx, connInfo, q, nameServer, r.ednsOptions, recursive, r.dnsSecEnabled, r.checkingDisabledBit, r.authenticatedDataBit, useEDNS)
if status == StatusTruncated && connInfo.tcpClient != nil {
// result truncated, try again with TCP
r.verboseLog(depth, logPrefix, TCPProtocol, " ", dns.TypeToString[q.Type], " ", q.Name, " ", nameServer)
return wireLookupTCP(ctx, connInfo, q, nameServer, r.ednsOptions, recursive, r.dnsSecEnabled, r.checkingDisabledBit, r.authenticatedDataBit, useEDNS)
}
return result, rawResp, status, err
}
if connInfo.tcpClient != nil {
r.verboseLog(depth, logPrefix, TCPProtocol, " ", dns.TypeToString[q.Type], " ", q.Name, " ", nameServer)
return wireLookupTCP(ctx, connInfo, q, nameServer, r.ednsOptions, recursive, r.dnsSecEnabled, r.checkingDisabledBit, r.authenticatedDataBit, useEDNS)
}
return &SingleQueryResult{}, nil, StatusError, errors.New("no connection info for nameserver")
}

func doDoTLookup(ctx context.Context, connInfo *ConnectionInfo, q Question, nameServer *NameServer, rootCAs *x509.CertPool, shouldVerifyServerCert, recursive bool, ednsOptions []dns.EDNS0, dnssec, checkingDisabled, authenticatedData, useEDNS bool) (*SingleQueryResult, *dns.Msg, Status, error) {
if util.HasCtxExpired(ctx) {
return nil, nil, StatusTimeout, errors.New("context expired")
}
m := new(dns.Msg)
m.SetQuestion(dotName(q.Name), q.Type)
m.Question[0].Qclass = q.Class
m.RecursionDesired = recursive
m.CheckingDisabled = checkingDisabled
m.AuthenticatedData = authenticatedData
m := makeWireQuery(q, recursive, ednsOptions, dnssec, checkingDisabled, authenticatedData, useEDNS)
m.Id = 12345

m.SetEdns0(1232, dnssec)
if ednsOpt := m.IsEdns0(); ednsOpt != nil {
ednsOpt.Option = append(ednsOpt.Option, ednsOptions...)
}

// if tlsConn is nil or if this is a new nameserver, create a new connection
var isConnNew bool
if connInfo.tlsConn != nil {
Expand Down Expand Up @@ -1038,18 +1049,8 @@ func doDoTLookup(ctx context.Context, connInfo *ConnectionInfo, q Question, name
return constructSingleQueryResultFromDNSMsg(&res, responseMsg)
}

func doDoHLookup(ctx context.Context, httpClient *http.Client, q Question, nameServer *NameServer, recursive bool, ednsOptions []dns.EDNS0, dnssec, checkingDisabled, authenticatedData bool) (*SingleQueryResult, *dns.Msg, Status, error) {
m := new(dns.Msg)
m.SetQuestion(dotName(q.Name), q.Type)
m.Question[0].Qclass = q.Class
m.RecursionDesired = recursive
m.CheckingDisabled = checkingDisabled
m.AuthenticatedData = authenticatedData

m.SetEdns0(1232, dnssec)
if ednsOpt := m.IsEdns0(); ednsOpt != nil {
ednsOpt.Option = append(ednsOpt.Option, ednsOptions...)
}
func doDoHLookup(ctx context.Context, httpClient *http.Client, q Question, nameServer *NameServer, recursive bool, ednsOptions []dns.EDNS0, dnssec, checkingDisabled, authenticatedData, useEDNS bool) (*SingleQueryResult, *dns.Msg, Status, error) {
m := makeWireQuery(q, recursive, ednsOptions, dnssec, checkingDisabled, authenticatedData, useEDNS)
bytes, err := m.Pack()
if err != nil {
return nil, nil, StatusError, errors.Wrap(err, "could not pack DNS message")
Expand Down Expand Up @@ -1111,21 +1112,11 @@ func doDoHLookup(ctx context.Context, httpClient *http.Client, q Question, nameS
}

// wireLookupTCP performs a DNS lookup on-the-wire over TCP with the given parameters
func wireLookupTCP(ctx context.Context, connInfo *ConnectionInfo, q Question, nameServer *NameServer, ednsOptions []dns.EDNS0, recursive, dnssec, checkingDisabled, authenticatedData bool) (*SingleQueryResult, *dns.Msg, Status, error) {
func wireLookupTCP(ctx context.Context, connInfo *ConnectionInfo, q Question, nameServer *NameServer, ednsOptions []dns.EDNS0, recursive, dnssec, checkingDisabled, authenticatedData, useEDNS bool) (*SingleQueryResult, *dns.Msg, Status, error) {
res := SingleQueryResult{Answers: []any{}, Authorities: []any{}, Additionals: []any{}}
res.Resolver = nameServer.String()

m := new(dns.Msg)
m.SetQuestion(dotName(q.Name), q.Type)
m.Question[0].Qclass = q.Class
m.RecursionDesired = recursive
m.CheckingDisabled = checkingDisabled
m.AuthenticatedData = authenticatedData

m.SetEdns0(1232, dnssec)
if ednsOpt := m.IsEdns0(); ednsOpt != nil {
ednsOpt.Option = append(ednsOpt.Option, ednsOptions...)
}
m := makeWireQuery(q, recursive, ednsOptions, dnssec, checkingDisabled, authenticatedData, useEDNS)

var r *dns.Msg
var err error
Expand Down Expand Up @@ -1172,22 +1163,12 @@ func wireLookupTCP(ctx context.Context, connInfo *ConnectionInfo, q Question, na
}

// wireLookupUDP performs a DNS lookup on-the-wire over UDP with the given parameters
func wireLookupUDP(ctx context.Context, connInfo *ConnectionInfo, q Question, nameServer *NameServer, ednsOptions []dns.EDNS0, recursive, dnssec, checkingDisabled, authenticatedData bool) (*SingleQueryResult, *dns.Msg, Status, error) {
func wireLookupUDP(ctx context.Context, connInfo *ConnectionInfo, q Question, nameServer *NameServer, ednsOptions []dns.EDNS0, recursive, dnssec, checkingDisabled, authenticatedData, useEDNS bool) (*SingleQueryResult, *dns.Msg, Status, error) {
res := SingleQueryResult{Answers: []any{}, Authorities: []any{}, Additionals: []any{}}
res.Resolver = nameServer.String()
res.Protocol = "udp"

m := new(dns.Msg)
m.SetQuestion(dotName(q.Name), q.Type)
m.Question[0].Qclass = q.Class
m.RecursionDesired = recursive
m.CheckingDisabled = checkingDisabled
m.AuthenticatedData = authenticatedData

m.SetEdns0(1232, dnssec)
if ednsOpt := m.IsEdns0(); ednsOpt != nil {
ednsOpt.Option = append(ednsOpt.Option, ednsOptions...)
}
m := makeWireQuery(q, recursive, ednsOptions, dnssec, checkingDisabled, authenticatedData, useEDNS)

var r *dns.Msg
var err error
Expand Down Expand Up @@ -1218,6 +1199,34 @@ func wireLookupUDP(ctx context.Context, connInfo *ConnectionInfo, q Question, na
return constructSingleQueryResultFromDNSMsg(&res, r)
}

func makeWireQuery(q Question, recursive bool, ednsOptions []dns.EDNS0, dnssec, checkingDisabled, authenticatedData, useEDNS bool) *dns.Msg {
m := new(dns.Msg)
m.SetQuestion(dotName(q.Name), q.Type)
m.Question[0].Qclass = q.Class
m.RecursionDesired = recursive
m.CheckingDisabled = checkingDisabled
m.AuthenticatedData = authenticatedData

if useEDNS {
m.SetEdns0(1232, dnssec)
if ednsOpt := m.IsEdns0(); ednsOpt != nil {
ednsOpt.Option = append(ednsOpt.Option, ednsOptions...)
}
}
return m
}

func shouldRetryWithoutEDNS(response *dns.Msg, ednsOptions []dns.EDNS0, dnssec bool) bool {
// RFC 6891 section 7 specifies FORMERR without an OPT record as the response
// from a server that does not implement EDNS. Queries that require an EDNS
// feature cannot be retried without changing their requested semantics.
return response != nil &&
response.Rcode == dns.RcodeFormatError &&
response.IsEdns0() == nil &&
!dnssec &&
len(ednsOptions) == 0
}

// fills out all the fields in a SingleQueryResult from a dns.Msg directly.
func constructSingleQueryResultFromDNSMsg(res *SingleQueryResult, r *dns.Msg) (*SingleQueryResult, *dns.Msg, Status, error) {
if r.Rcode != dns.RcodeSuccess {
Expand Down
Loading