diff --git a/dns/transport/dhcp/dhcp_shared.go b/dns/transport/dhcp/dhcp_shared.go index f840846b..a08c9bb5 100644 --- a/dns/transport/dhcp/dhcp_shared.go +++ b/dns/transport/dhcp/dhcp_shared.go @@ -20,15 +20,9 @@ func (t *Transport) exchangeWithTransports(ctx context.Context, message *mDNS.Ms callback(nil, E.New("invalid domain: ", domain)) return } - nameExchangers := make([]transport.AsyncExchanger, 0, len(names)) - for _, fqdn := range names { - nameExchangers = append(nameExchangers, t.newNameExchanger(message, fqdn, state.serverTransports)) - } - if len(state.serverTransports) == 1 || !(question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) { - transport.ExchangeSequential(ctx, nameExchangers, nil, callback) - } else { - transport.ExchangeRace(ctx, nameExchangers, callback) - } + transport.ExchangeNames(ctx, names, question, func(fqdn string) transport.AsyncExchanger { + return t.newNameExchanger(message, fqdn, state.serverTransports) + }, callback) } func (t *Transport) newNameExchanger(message *mDNS.Msg, fqdn string, serverTransports []adapter.DNSTransport) transport.AsyncExchanger { diff --git a/dns/transport/exchange_strategy.go b/dns/transport/exchange_strategy.go index 92faff1c..c85d3bb6 100644 --- a/dns/transport/exchange_strategy.go +++ b/dns/transport/exchange_strategy.go @@ -2,8 +2,10 @@ package transport import ( "context" + "strings" "sync" + "github.com/sagernet/sing/common" "github.com/sagernet/sing/common/buf" E "github.com/sagernet/sing/common/exceptions" @@ -80,60 +82,62 @@ type sequentialCallState struct { continued bool } -// ExchangeRace runs all exchangers concurrently; the first success wins and -// cancels the rest, and when all fail the errors are aggregated. -func ExchangeRace(ctx context.Context, exchangers []AsyncExchanger, callback func(response *mDNS.Msg, err error)) { - if len(exchangers) == 0 { - callback(nil, E.New("missing exchangers")) +func ExchangeNames(ctx context.Context, names []string, question mDNS.Question, exchangerFor func(fqdn string) AsyncExchanger, callback func(response *mDNS.Msg, err error)) { + if len(names) == 0 { + callback(nil, E.New("missing name candidates")) return } - if len(exchangers) == 1 { - exchangers[0](ctx, callback) - return - } - raceCtx, raceCancel := context.WithCancel(ctx) - state := &raceState{ - cancel: raceCancel, - remaining: len(exchangers), - callback: callback, - } - for _, exchanger := range exchangers { - exchanger(raceCtx, state.complete) - } -} - -type raceState struct { - access sync.Mutex - done bool - remaining int - errors []error - cancel context.CancelFunc - callback func(response *mDNS.Msg, err error) -} - -func (s *raceState) complete(response *mDNS.Msg, err error) { - s.access.Lock() - if s.done { - s.access.Unlock() - return - } - if err != nil { - s.errors = append(s.errors, err) - if len(s.errors) < s.remaining { - s.access.Unlock() - return + search := &nameSearchExchange{question: question} + nameExchangers := common.Map(names, func(fqdn string) AsyncExchanger { + return search.wrap(fqdn, exchangerFor(fqdn)) + }) + ExchangeSequential(ctx, nameExchangers, func(response *mDNS.Msg, err error) bool { + return err == nil && response.Rcode != mDNS.RcodeNameError + }, func(response *mDNS.Msg, err error) { + if err != nil || response.Rcode == mDNS.RcodeNameError { + search.access.Lock() + nameErrorResponse := search.nameErrorResponse + search.access.Unlock() + if nameErrorResponse != nil { + response, err = nameErrorResponse, nil + } + } + callback(response, err) + }) +} + +type nameSearchExchange struct { + question mDNS.Question + access sync.Mutex + nameErrorResponse *mDNS.Msg +} + +func (s *nameSearchExchange) wrap(fqdn string, exchanger AsyncExchanger) AsyncExchanger { + return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) { + exchanger(ctx, func(response *mDNS.Msg, err error) { + if err == nil { + restoreOriginalQuestion(response, fqdn, s.question) + if response.Rcode == mDNS.RcodeNameError { + s.access.Lock() + if s.nameErrorResponse == nil || fqdn == s.question.Name { + s.nameErrorResponse = response + } + s.access.Unlock() + } + } + callback(response, err) + }) + } +} + +// 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 } - raceErrors := s.errors - s.done = true - s.access.Unlock() - s.cancel() - s.callback(nil, E.Errors(raceErrors...)) - return } - s.done = true - s.access.Unlock() - s.cancel() - s.callback(response, nil) } func NewFanOutRequest(message *mDNS.Msg, fqdn string, authenticatedData bool) *mDNS.Msg { diff --git a/dns/transport/local/local_shared.go b/dns/transport/local/local_shared.go index 271776c3..b86f6d78 100644 --- a/dns/transport/local/local_shared.go +++ b/dns/transport/local/local_shared.go @@ -75,16 +75,9 @@ func (t *Transport) exchangeAsync(ctx context.Context, message *mDNS.Msg, domain callback(nil, E.New("invalid domain: ", domain)) return } - nameExchangers := make([]transport.AsyncExchanger, 0, len(names)) - for _, fqdn := range names { - nameExchangers = append(nameExchangers, newNameExchanger(systemConfig, serverSet, message, fqdn)) - } - question := message.Question[0] - if systemConfig.SingleRequest || !(question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) { - transport.ExchangeSequential(ctx, nameExchangers, nil, callback) - } else { - transport.ExchangeRace(ctx, nameExchangers, callback) - } + transport.ExchangeNames(ctx, names, message.Question[0], func(fqdn string) transport.AsyncExchanger { + return newNameExchanger(systemConfig, serverSet, message, fqdn) + }, callback) } func newNameExchanger(systemConfig *systemconfig.Config, serverSet *localServerSet, message *mDNS.Msg, fqdn string) transport.AsyncExchanger { diff --git a/service/resolved/transport.go b/service/resolved/transport.go index 4e53896f..858cb6b0 100644 --- a/service/resolved/transport.go +++ b/service/resolved/transport.go @@ -310,15 +310,9 @@ func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callba callback(nil, E.New("invalid domain: ", question.Name)) return } - nameExchangers := make([]transport.AsyncExchanger, 0, len(names)) - for _, fqdn := range names { - nameExchangers = append(nameExchangers, t.newNameExchanger(servers, message, fqdn)) - } - if question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA { - transport.ExchangeRace(ctx, nameExchangers, callback) - } else { - transport.ExchangeSequential(ctx, nameExchangers, nil, callback) - } + transport.ExchangeNames(ctx, names, question, func(fqdn string) transport.AsyncExchanger { + return t.newNameExchanger(servers, message, fqdn) + }, callback) } func (t *Transport) newNameExchanger(servers *LinkServers, message *mDNS.Msg, fqdn string) transport.AsyncExchanger {