diff --git a/dns/transport/dhcp/dhcp.go b/dns/transport/dhcp/dhcp.go index c06ad5b9..f9702da4 100644 --- a/dns/transport/dhcp/dhcp.go +++ b/dns/transport/dhcp/dhcp.go @@ -128,13 +128,7 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, } func (t *Transport) Exchange0(ctx context.Context, message *mDNS.Msg, servers []M.Socksaddr) (*mDNS.Msg, error) { - question := message.Question[0] - domain := dns.FqdnToDomain(question.Name) - if len(servers) == 1 || !(message.Question[0].Qtype == mDNS.TypeA || message.Question[0].Qtype == mDNS.TypeAAAA) { - return t.exchangeSingleRequest(ctx, servers, message, domain) - } else { - return t.exchangeParallel(ctx, servers, message, domain) - } + return t.exchangeSearch(ctx, servers, message, dns.FqdnToDomain(message.Question[0].Name)) } func (t *Transport) Fetch() []M.Socksaddr { diff --git a/dns/transport/dhcp/dhcp_shared.go b/dns/transport/dhcp/dhcp_shared.go index 16e319ba..32621ef7 100644 --- a/dns/transport/dhcp/dhcp_shared.go +++ b/dns/transport/dhcp/dhcp_shared.go @@ -16,54 +16,43 @@ import ( mDNS "github.com/miekg/dns" ) -func (t *Transport) exchangeSingleRequest(ctx context.Context, servers []M.Socksaddr, message *mDNS.Msg, domain string) (*mDNS.Msg, error) { - var lastErr error - for _, fqdn := range t.nameList(domain) { +func (t *Transport) exchangeSearch(ctx context.Context, servers []M.Socksaddr, message *mDNS.Msg, domain string) (*mDNS.Msg, error) { + names := t.nameList(domain) + if len(names) == 0 { + return nil, E.New("dhcp: invalid domain: ", domain) + } + originalQuestion := message.Question[0] + var ( + nameErrorResponse *mDNS.Msg + lastErr error + ) + for _, fqdn := range names { response, err := t.tryOneName(ctx, servers, fqdn, message) if err != nil { - lastErr = err + lastErr = E.Errors(lastErr, err) + continue + } + restoreOriginalQuestion(response, fqdn, originalQuestion) + if response.Rcode == mDNS.RcodeNameError { + if nameErrorResponse == nil || fqdn == originalQuestion.Name { + nameErrorResponse = response + } continue } return response, nil } + if nameErrorResponse != nil { + return nameErrorResponse, nil + } return nil, lastErr } -func (t *Transport) exchangeParallel(ctx context.Context, servers []M.Socksaddr, message *mDNS.Msg, domain string) (*mDNS.Msg, error) { - returned := make(chan struct{}) - defer close(returned) - type queryResult struct { - response *mDNS.Msg - err error - } - results := make(chan queryResult) - startRacer := func(ctx context.Context, fqdn string) { - response, err := t.tryOneName(ctx, servers, fqdn, message) - select { - case results <- queryResult{response, err}: - case <-returned: - } - } - queryCtx, queryCancel := context.WithCancel(ctx) - defer queryCancel() - var nameCount int - for _, fqdn := range t.nameList(domain) { - nameCount++ - go startRacer(queryCtx, fqdn) - } - var errors []error - for { - select { - case <-ctx.Done(): - return nil, ctx.Err() - case result := <-results: - if result.err == nil { - return result.response, nil - } - errors = append(errors, result.err) - if len(errors) == nameCount { - return nil, E.Errors(errors...) - } +// Stub resolvers discard Answer RRs whose owner name does not match the question. +func restoreOriginalQuestion(response *mDNS.Msg, fqdn string, question mDNS.Question) { + response.Question = []mDNS.Question{question} + for _, record := range response.Answer { + if strings.EqualFold(record.Header().Name, fqdn) { + record.Header().Name = question.Name } } }