mirror of
https://github.com/shtorm-7/sing-box-extended.git
synced 2026-09-15 21:00:27 +00:00
Fix search domain handling in DNS transports
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user