Fix search domain handling in DNS transports

This commit is contained in:
世界
2026-08-30 17:41:45 +08:00
parent 0e75f5f45c
commit 87c4f84a89
4 changed files with 63 additions and 78 deletions
+3 -9
View File
@@ -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 {
+54 -50
View File
@@ -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 {
+3 -10
View File
@@ -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 {
+3 -9
View File
@@ -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 {