refactor: Async DNS

This commit is contained in:
世界
2026-08-30 17:41:43 +08:00
parent 90d75b9673
commit 5d424ea2ac
29 changed files with 1956 additions and 994 deletions
+3
View File
@@ -18,6 +18,7 @@ import (
type DNSRouter interface {
Lifecycle
Exchange(ctx context.Context, message *dns.Msg, options DNSQueryOptions) (*dns.Msg, error)
ExchangeAsync(ctx context.Context, message *dns.Msg, options DNSQueryOptions, callback func(response *dns.Msg, err error))
Lookup(ctx context.Context, domain string, options DNSQueryOptions) ([]netip.Addr, error)
ClearCache()
LookupReverseMapping(ip netip.Addr) (string, bool)
@@ -27,6 +28,7 @@ type DNSRouter interface {
type DNSClient interface {
Start()
Exchange(ctx context.Context, transport DNSTransport, message *dns.Msg, options DNSQueryOptions, responseChecker func(response *dns.Msg) bool) (*dns.Msg, error)
ExchangeAsync(ctx context.Context, transport DNSTransport, message *dns.Msg, options DNSQueryOptions, responseChecker func(response *dns.Msg) bool, callback func(response *dns.Msg, err error))
Lookup(ctx context.Context, transport DNSTransport, domain string, options DNSQueryOptions, responseChecker func(response *dns.Msg) bool) ([]netip.Addr, error)
ClearCache()
}
@@ -84,6 +86,7 @@ type DNSTransport interface {
// Exchanges that are currently using those connections may fail.
Reset()
Exchange(ctx context.Context, message *dns.Msg) (*dns.Msg, error)
ExchangeAsync(ctx context.Context, message *dns.Msg, callback func(response *dns.Msg, err error))
}
type DNSTransportWithPreferredDomain interface {
+125 -24
View File
@@ -152,19 +152,45 @@ func normalizeTTL(response *dns.Msg, timeToLive uint32) {
}
}
func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool) (*dns.Msg, error) {
type exchangeStatus int
const (
exchangeReady exchangeStatus = iota
exchangeDone
exchangeWait
)
type exchangeOperation struct {
ctx context.Context
message *dns.Msg
question dns.Question
messageId uint16
options adapter.DNSQueryOptions
responseChecker func(response *dns.Msg) bool
disableCache bool
releaseCond func()
}
func (o *exchangeOperation) release() {
if o.releaseCond != nil {
o.releaseCond()
o.releaseCond = nil
}
}
func (c *Client) beginExchange(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool, allowWait bool) (*exchangeOperation, *dns.Msg, exchangeStatus, error) {
if len(message.Question) == 0 {
if c.logger != nil {
c.logger.WarnContext(ctx, "bad question size: ", len(message.Question))
}
return FixedResponseStatus(message, dns.RcodeFormatError), nil
return nil, FixedResponseStatus(message, dns.RcodeFormatError), exchangeDone, nil
}
question := message.Question[0]
if question.Qtype == dns.TypeA && options.Strategy == C.DomainStrategyIPv6Only || question.Qtype == dns.TypeAAAA && options.Strategy == C.DomainStrategyIPv4Only {
if c.logger != nil {
c.logger.DebugContext(ctx, "strategy rejected")
}
return FixedResponseStatus(message, dns.RcodeSuccess), nil
return nil, FixedResponseStatus(message, dns.RcodeSuccess), exchangeDone, nil
}
message = c.prepareExchangeMessage(message, options)
@@ -177,20 +203,31 @@ func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, m
len(message.Extra[0].(*dns.OPT).Option) == 0) &&
!options.ClientSubnet.IsValid()
disableCache := !isSimpleRequest || c.disableCache || options.DisableCache
operation := &exchangeOperation{
message: message,
question: question,
messageId: message.Id,
options: options,
responseChecker: responseChecker,
disableCache: disableCache,
}
if !disableCache {
cacheKey := dnsCacheKey{Question: question, transportTag: transport.Tag()}
cond, loaded := c.cacheLock.LoadOrStore(cacheKey, make(chan struct{}))
if loaded {
if !allowWait {
return nil, nil, exchangeWait, nil
}
select {
case <-cond:
case <-ctx.Done():
return nil, ctx.Err()
return nil, nil, exchangeDone, ctx.Err()
}
} else {
defer func() {
operation.releaseCond = func() {
c.cacheLock.Delete(cacheKey)
close(cond)
}()
}
}
response, ttl, isStale := c.loadResponse(question, transport)
if response != nil {
@@ -198,38 +235,43 @@ func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, m
c.backgroundRefreshDNS(transport, question, message.Copy(), options, responseChecker)
logOptimisticResponse(c.logger, ctx, response)
response.Id = message.Id
return response, nil
operation.release()
return nil, response, exchangeDone, nil
} else if !isStale {
logCachedResponse(c.logger, ctx, response, ttl)
response.Id = message.Id
return response, nil
operation.release()
return nil, response, exchangeDone, nil
}
}
}
messageId := message.Id
contextTransport, clientSubnetLoaded := transportTagFromContext(ctx)
if clientSubnetLoaded && transport.Tag() == contextTransport {
return nil, E.New("DNS query loopback in transport[", contextTransport, "]")
contextTransport, transportTagLoaded := transportTagFromContext(ctx)
if transportTagLoaded && transport.Tag() == contextTransport {
operation.release()
return nil, nil, exchangeDone, E.New("DNS query loopback in transport[", contextTransport, "]")
}
ctx = contextWithTransportTag(ctx, transport.Tag())
operation.ctx = contextWithTransportTag(ctx, transport.Tag())
if !disableCache && responseChecker != nil && c.rdrc != nil {
rejected := c.rdrc.LoadRDRC(transport.Tag(), question.Name, question.Qtype)
if rejected {
return nil, ErrResponseRejectedCached
operation.release()
return nil, nil, exchangeDone, ErrResponseRejectedCached
}
}
response, err := c.exchangeToTransport(ctx, transport, message, options.Timeout)
if err != nil {
return nil, err
}
disableCache = disableCache || (response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError)
if responseChecker != nil {
return operation, nil, exchangeReady, nil
}
func (c *Client) finishExchange(transport adapter.DNSTransport, operation *exchangeOperation, response *dns.Msg) (*dns.Msg, error) {
ctx := operation.ctx
question := operation.question
disableCache := operation.disableCache || (response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError)
if operation.responseChecker != nil {
var rejected bool
if response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError {
rejected = true
} else {
rejected = !responseChecker(response)
rejected = !operation.responseChecker(response)
}
if rejected {
if !disableCache && c.rdrc != nil {
@@ -239,12 +281,12 @@ func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, m
return response, ErrResponseRejected
}
}
timeToLive := applyResponseOptions(question, response, options)
timeToLive := applyResponseOptions(question, response, operation.options)
if !disableCache {
c.storeCache(transport, question, response, timeToLive)
}
response.Id = messageId
requestEDNSOpt := message.IsEdns0()
response.Id = operation.messageId
requestEDNSOpt := operation.message.IsEdns0()
responseEDNSOpt := response.IsEdns0()
if responseEDNSOpt != nil && (requestEDNSOpt == nil || requestEDNSOpt.Version() < responseEDNSOpt.Version()) {
response.Extra = common.Filter(response.Extra, func(it dns.RR) bool {
@@ -258,6 +300,44 @@ func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, m
return response, nil
}
func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool) (*dns.Msg, error) {
operation, earlyResponse, status, err := c.beginExchange(ctx, transport, message, options, responseChecker, true)
if status != exchangeReady {
return earlyResponse, err
}
defer operation.release()
response, err := c.exchangeToTransport(operation.ctx, transport, operation.message, options.Timeout)
if err != nil {
return nil, err
}
return c.finishExchange(transport, operation, response)
}
func (c *Client) ExchangeAsync(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool, callback func(response *dns.Msg, err error)) {
operation, earlyResponse, status, err := c.beginExchange(ctx, transport, message, options, responseChecker, false)
switch status {
case exchangeDone:
callback(earlyResponse, err)
return
case exchangeWait:
go func() {
callback(c.Exchange(ctx, transport, message, options, responseChecker))
}()
return
}
finish := func(response *dns.Msg, exchangeErr error) {
if exchangeErr != nil {
operation.release()
callback(nil, exchangeErr)
return
}
finishedResponse, finishErr := c.finishExchange(transport, operation, response)
operation.release()
callback(finishedResponse, finishErr)
}
c.exchangeToTransportAsync(operation.ctx, transport, operation.message, options.Timeout, finish)
}
func (c *Client) Lookup(ctx context.Context, transport adapter.DNSTransport, domain string, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool) ([]netip.Addr, error) {
domain = FqdnToDomain(domain)
dnsName := dns.Fqdn(domain)
@@ -562,6 +642,27 @@ func (c *Client) exchangeToTransport(ctx context.Context, transport adapter.DNST
return nil, err
}
func (c *Client) exchangeToTransportAsync(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, timeout time.Duration, callback func(response *dns.Msg, err error)) {
if timeout == 0 {
timeout = c.timeout
}
ctx, cancel := context.WithTimeout(ctx, timeout)
transport.ExchangeAsync(ctx, message, func(response *dns.Msg, err error) {
cancel()
if err == nil {
stripDNSPadding(response)
callback(response, nil)
return
}
var rcodeError RcodeError
if errors.As(err, &rcodeError) {
callback(FixedResponseStatus(message, int(rcodeError)), nil)
return
}
callback(nil, err)
})
}
func MessageToAddresses(response *dns.Msg) []netip.Addr {
return adapter.DNSResponseAddresses(response)
}
+230 -134
View File
@@ -410,60 +410,66 @@ type exchangeWithRulesResult struct {
const dnsRespondMissingResponseMessage = "respond action requires an evaluated response from a preceding evaluate action"
func (r *Router) exchangeWithRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, options adapter.DNSQueryOptions, allowFakeIP bool) exchangeWithRulesResult {
type dnsRuleWalkState struct {
ruleIndex int
effectiveOptions adapter.DNSQueryOptions
evaluatedResponse *mDNS.Msg
evaluatedTransport adapter.DNSTransport
}
type dnsPendingExchange struct {
transport adapter.DNSTransport
options adapter.DNSQueryOptions
evaluate bool
}
func (r *Router) finalizeExchangeOptions(options adapter.DNSQueryOptions) adapter.DNSQueryOptions {
if options.Strategy == C.DomainStrategyAsIS {
options.Strategy = r.defaultDomainStrategy
}
return options
}
func (r *Router) walkDNSRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, state *dnsRuleWalkState, allowFakeIP bool) (exchangeWithRulesResult, *dnsPendingExchange) {
metadata := adapter.ContextFrom(ctx)
if metadata == nil {
panic("no context")
}
effectiveOptions := options
var evaluatedResponse *mDNS.Msg
var evaluatedTransport adapter.DNSTransport
for currentRuleIndex, currentRule := range rules {
for ; state.ruleIndex < len(rules); state.ruleIndex++ {
currentRule := rules[state.ruleIndex]
metadata.ResetRuleCache()
metadata.DNSResponse = evaluatedResponse
metadata.DNSResponse = state.evaluatedResponse
metadata.DestinationAddressMatchFromResponse = false
if !currentRule.Match(metadata) {
continue
}
r.logRuleMatch(ctx, currentRuleIndex, currentRule)
r.logRuleMatch(ctx, state.ruleIndex, currentRule)
switch action := currentRule.Action().(type) {
case *R.RuleActionDNSRouteOptions:
r.applyDNSRouteOptions(&effectiveOptions, *action)
r.applyDNSRouteOptions(&state.effectiveOptions, *action)
case *R.RuleActionEvaluate:
queryOptions := effectiveOptions
queryOptions := state.effectiveOptions
transport, loaded := r.transport.Transport(action.Server)
if !loaded {
r.logger.ErrorContext(ctx, "transport not found: ", action.Server)
evaluatedResponse = nil
evaluatedTransport = nil
state.evaluatedResponse = nil
state.evaluatedTransport = nil
continue
}
r.applyDNSRouteOptions(&queryOptions, action.RuleActionDNSRouteOptions)
exchangeOptions := queryOptions
if exchangeOptions.Strategy == C.DomainStrategyAsIS {
exchangeOptions.Strategy = r.defaultDomainStrategy
}
response, err := r.client.Exchange(adapter.OverrideContext(ctx), transport, message, exchangeOptions, nil)
if err != nil {
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for ", FormatQuestion(message.Question[0].String())))
evaluatedResponse = nil
evaluatedTransport = nil
continue
}
evaluatedResponse = response
evaluatedTransport = transport
return exchangeWithRulesResult{}, &dnsPendingExchange{transport: transport, options: queryOptions, evaluate: true}
case *R.RuleActionRespond:
if evaluatedResponse == nil {
if state.evaluatedResponse == nil {
return exchangeWithRulesResult{
err: E.New(dnsRespondMissingResponseMessage),
}
}, nil
}
return exchangeWithRulesResult{
response: evaluatedResponse,
transport: evaluatedTransport,
}
response: state.evaluatedResponse,
transport: state.evaluatedTransport,
}, nil
case *R.RuleActionDNSRoute:
queryOptions := effectiveOptions
queryOptions := state.effectiveOptions
transport, status := r.resolveDNSRoute(action.Server, action.RuleActionDNSRouteOptions, allowFakeIP, &queryOptions)
switch status {
case dnsRouteStatusMissing:
@@ -472,16 +478,7 @@ func (r *Router) exchangeWithRules(ctx context.Context, rules []adapter.DNSRule,
case dnsRouteStatusSkipped:
continue
}
exchangeOptions := queryOptions
if exchangeOptions.Strategy == C.DomainStrategyAsIS {
exchangeOptions.Strategy = r.defaultDomainStrategy
}
response, err := r.client.Exchange(adapter.OverrideContext(ctx), transport, message, exchangeOptions, nil)
return exchangeWithRulesResult{
response: response,
transport: transport,
err: err,
}
return exchangeWithRulesResult{}, &dnsPendingExchange{transport: transport, options: queryOptions}
case *R.RuleActionReject:
switch action.Method {
case C.RuleActionRejectMethodDefault:
@@ -495,32 +492,80 @@ func (r *Router) exchangeWithRules(ctx context.Context, rules []adapter.DNSRule,
Question: []mDNS.Question{message.Question[0]},
},
rejectAction: action,
}
}, nil
case C.RuleActionRejectMethodDrop:
return exchangeWithRulesResult{
rejectAction: action,
err: R.ErrDrop,
}
}, nil
}
case *R.RuleActionPredefined:
return exchangeWithRulesResult{
response: action.Response(message),
}
}, nil
}
}
transport := r.transport.Default()
exchangeOptions := effectiveOptions
if exchangeOptions.Strategy == C.DomainStrategyAsIS {
exchangeOptions.Strategy = r.defaultDomainStrategy
return exchangeWithRulesResult{}, &dnsPendingExchange{transport: r.transport.Default(), options: state.effectiveOptions}
}
func (r *Router) exchangeWithRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, options adapter.DNSQueryOptions, allowFakeIP bool) exchangeWithRulesResult {
state := dnsRuleWalkState{effectiveOptions: options}
result, pending := r.walkDNSRules(ctx, rules, message, &state, allowFakeIP)
if pending == nil {
return result
}
response, err := r.client.Exchange(adapter.OverrideContext(ctx), transport, message, exchangeOptions, nil)
return exchangeWithRulesResult{
response: response,
transport: transport,
err: err,
return r.resumeExchangeWithRules(ctx, rules, message, &state, allowFakeIP, pending)
}
func (r *Router) resumeExchangeWithRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, state *dnsRuleWalkState, allowFakeIP bool, pending *dnsPendingExchange) exchangeWithRulesResult {
for {
response, err := r.client.Exchange(adapter.OverrideContext(ctx), pending.transport, message, r.finalizeExchangeOptions(pending.options), nil)
if !pending.evaluate {
return exchangeWithRulesResult{
response: response,
transport: pending.transport,
err: err,
}
}
if err != nil {
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for ", FormatQuestion(message.Question[0].String())))
state.evaluatedResponse = nil
state.evaluatedTransport = nil
} else {
state.evaluatedResponse = response
state.evaluatedTransport = pending.transport
}
state.ruleIndex++
var result exchangeWithRulesResult
result, pending = r.walkDNSRules(ctx, rules, message, state, allowFakeIP)
if pending == nil {
return result
}
}
}
func (r *Router) exchangeWithRulesAsync(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, options adapter.DNSQueryOptions, allowFakeIP bool, callback func(result exchangeWithRulesResult)) {
state := dnsRuleWalkState{effectiveOptions: options}
result, pending := r.walkDNSRules(ctx, rules, message, &state, allowFakeIP)
if pending == nil {
callback(result)
return
}
if pending.evaluate {
go func() {
callback(r.resumeExchangeWithRules(ctx, rules, message, &state, allowFakeIP, pending))
}()
return
}
r.client.ExchangeAsync(adapter.OverrideContext(ctx), pending.transport, message, r.finalizeExchangeOptions(pending.options), nil, func(response *mDNS.Msg, err error) {
callback(exchangeWithRulesResult{
response: response,
transport: pending.transport,
err: err,
})
})
}
func (r *Router) resolveLookupStrategy(options adapter.DNSQueryOptions) C.DomainStrategy {
if options.LookupStrategy != C.DomainStrategyAsIS {
return options.LookupStrategy
@@ -617,35 +662,35 @@ func (r *Router) lookupWithRulesType(ctx context.Context, rules []adapter.DNSRul
return filterAddressesByQueryType(MessageToAddresses(exchangeResult.response), qType), nil
}
func (r *Router) Exchange(ctx context.Context, message *mDNS.Msg, options adapter.DNSQueryOptions) (*mDNS.Msg, error) {
type dnsExchangeContext struct {
ctx context.Context
rules []adapter.DNSRule
legacyDNSMode bool
metadata *adapter.InboundContext
}
func (r *Router) prepareExchange(ctx context.Context, message *mDNS.Msg) (*dnsExchangeContext, *mDNS.Msg, error) {
if len(message.Question) != 1 {
r.logger.WarnContext(ctx, "bad question size: ", len(message.Question))
responseMessage := mDNS.Msg{
return nil, &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Id: message.Id,
Response: true,
Rcode: mDNS.RcodeFormatError,
},
Question: message.Question,
}
return &responseMessage, nil
}, nil
}
r.rulesAccess.RLock()
if r.closing {
r.rulesAccess.RUnlock()
return nil, E.New("dns router closed")
return nil, nil, E.New("dns router closed")
}
rules := r.rules
legacyDNSMode := r.legacyDNSMode
r.rulesAccess.RUnlock()
r.logger.DebugContext(ctx, "exchange ", FormatQuestion(message.Question[0].String()))
var (
response *mDNS.Msg
transport adapter.DNSTransport
err error
)
var metadata *adapter.InboundContext
ctx, metadata = adapter.ExtendContext(ctx)
ctx, metadata := adapter.ExtendContext(ctx)
metadata.Destination = M.Socksaddr{}
metadata.QueryType = message.Question[0].Qtype
metadata.DNSResponse = nil
@@ -657,76 +702,15 @@ func (r *Router) Exchange(ctx context.Context, message *mDNS.Msg, options adapte
metadata.IPVersion = 6
}
metadata.Domain = FqdnToDomain(message.Question[0].Name)
if options.Transport != nil {
transport = options.Transport
if options.Strategy == C.DomainStrategyAsIS {
options.Strategy = r.defaultDomainStrategy
}
response, err = r.client.Exchange(ctx, transport, message, options, nil)
} else if !legacyDNSMode {
exchangeResult := r.exchangeWithRules(ctx, rules, message, options, true)
response, transport, err = exchangeResult.response, exchangeResult.transport, exchangeResult.err
} else {
var (
rule adapter.DNSRule
ruleIndex int
)
ruleIndex = -1
for {
dnsCtx := adapter.OverrideContext(ctx)
dnsOptions := options
transport, rule, ruleIndex = r.matchDNS(ctx, rules, true, ruleIndex, isAddressQuery(message), &dnsOptions)
if rule != nil {
switch action := rule.Action().(type) {
case *R.RuleActionReject:
switch action.Method {
case C.RuleActionRejectMethodDefault:
return &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Id: message.Id,
Rcode: mDNS.RcodeRefused,
Response: true,
},
Question: []mDNS.Question{message.Question[0]},
}, nil
case C.RuleActionRejectMethodDrop:
return nil, R.ErrDrop
}
case *R.RuleActionPredefined:
err = nil
response = action.Response(message)
goto done
}
}
responseCheck := addressLimitResponseCheck(rule, metadata)
if dnsOptions.Strategy == C.DomainStrategyAsIS {
dnsOptions.Strategy = r.defaultDomainStrategy
}
response, err = r.client.Exchange(dnsCtx, transport, message, dnsOptions, responseCheck)
var rejected bool
if err != nil {
if errors.Is(err, ErrResponseRejectedCached) {
rejected = true
r.logger.DebugContext(ctx, E.Cause(err, "response rejected for ", FormatQuestion(message.Question[0].String())), " (cached)")
} else if errors.Is(err, ErrResponseRejected) {
rejected = true
r.logger.DebugContext(ctx, E.Cause(err, "response rejected for ", FormatQuestion(message.Question[0].String())))
} else if len(message.Question) > 0 {
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for ", FormatQuestion(message.Question[0].String())))
} else {
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for <empty query>"))
}
}
if responseCheck != nil && rejected {
continue
}
break
}
}
done:
if err != nil {
return nil, err
}
return &dnsExchangeContext{
ctx: ctx,
rules: rules,
legacyDNSMode: legacyDNSMode,
metadata: metadata,
}, nil, nil
}
func (r *Router) recordReverseMapping(message *mDNS.Msg, response *mDNS.Msg, transport adapter.DNSTransport) {
if r.dnsReverseMapping != nil && len(message.Question) > 0 && response != nil && len(response.Answer) > 0 {
if transport == nil || transport.Type() != C.DNSTypeFakeIP {
for _, answer := range response.Answer {
@@ -739,9 +723,121 @@ done:
}
}
}
}
func (r *Router) exchangeLegacy(ctx context.Context, exchangeCtx *dnsExchangeContext, message *mDNS.Msg, options adapter.DNSQueryOptions) (*mDNS.Msg, adapter.DNSTransport, error) {
var (
transport adapter.DNSTransport
rule adapter.DNSRule
ruleIndex int
)
ruleIndex = -1
for {
dnsCtx := adapter.OverrideContext(ctx)
dnsOptions := options
transport, rule, ruleIndex = r.matchDNS(ctx, exchangeCtx.rules, true, ruleIndex, isAddressQuery(message), &dnsOptions)
if rule != nil {
switch action := rule.Action().(type) {
case *R.RuleActionReject:
switch action.Method {
case C.RuleActionRejectMethodDefault:
return &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Id: message.Id,
Rcode: mDNS.RcodeRefused,
Response: true,
},
Question: []mDNS.Question{message.Question[0]},
}, nil, nil
case C.RuleActionRejectMethodDrop:
return nil, nil, R.ErrDrop
}
case *R.RuleActionPredefined:
return action.Response(message), nil, nil
}
}
responseCheck := addressLimitResponseCheck(rule, exchangeCtx.metadata)
response, err := r.client.Exchange(dnsCtx, transport, message, r.finalizeExchangeOptions(dnsOptions), responseCheck)
var rejected bool
if err != nil {
if errors.Is(err, ErrResponseRejectedCached) {
rejected = true
r.logger.DebugContext(ctx, E.Cause(err, "response rejected for ", FormatQuestion(message.Question[0].String())), " (cached)")
} else if errors.Is(err, ErrResponseRejected) {
rejected = true
r.logger.DebugContext(ctx, E.Cause(err, "response rejected for ", FormatQuestion(message.Question[0].String())))
} else if len(message.Question) > 0 {
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for ", FormatQuestion(message.Question[0].String())))
} else {
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for <empty query>"))
}
}
if responseCheck != nil && rejected {
continue
}
return response, transport, err
}
}
func (r *Router) Exchange(ctx context.Context, message *mDNS.Msg, options adapter.DNSQueryOptions) (*mDNS.Msg, error) {
exchangeCtx, earlyResponse, err := r.prepareExchange(ctx, message)
if exchangeCtx == nil {
return earlyResponse, err
}
ctx = exchangeCtx.ctx
var (
response *mDNS.Msg
transport adapter.DNSTransport
)
if options.Transport != nil {
transport = options.Transport
response, err = r.client.Exchange(ctx, transport, message, r.finalizeExchangeOptions(options), nil)
} else if !exchangeCtx.legacyDNSMode {
exchangeResult := r.exchangeWithRules(ctx, exchangeCtx.rules, message, options, true)
response, transport, err = exchangeResult.response, exchangeResult.transport, exchangeResult.err
} else {
response, transport, err = r.exchangeLegacy(ctx, exchangeCtx, message, options)
}
if err != nil {
return nil, err
}
r.recordReverseMapping(message, response, transport)
return response, nil
}
func (r *Router) ExchangeAsync(ctx context.Context, message *mDNS.Msg, options adapter.DNSQueryOptions, callback func(response *mDNS.Msg, err error)) {
exchangeCtx, earlyResponse, err := r.prepareExchange(ctx, message)
if exchangeCtx == nil {
callback(earlyResponse, err)
return
}
ctx = exchangeCtx.ctx
if options.Transport != nil {
transport := options.Transport
r.client.ExchangeAsync(ctx, transport, message, r.finalizeExchangeOptions(options), nil, func(response *mDNS.Msg, exchangeErr error) {
r.finishExchangeAsync(message, transport, response, exchangeErr, callback)
})
} else if !exchangeCtx.legacyDNSMode {
r.exchangeWithRulesAsync(ctx, exchangeCtx.rules, message, options, true, func(result exchangeWithRulesResult) {
r.finishExchangeAsync(message, result.transport, result.response, result.err, callback)
})
} else {
go func() {
response, transport, exchangeErr := r.exchangeLegacy(ctx, exchangeCtx, message, options)
r.finishExchangeAsync(message, transport, response, exchangeErr, callback)
}()
}
}
func (r *Router) finishExchangeAsync(message *mDNS.Msg, transport adapter.DNSTransport, response *mDNS.Msg, err error, callback func(response *mDNS.Msg, err error)) {
if err != nil {
callback(nil, err)
return
}
r.recordReverseMapping(message, response, transport)
callback(response, nil)
}
func (r *Router) Lookup(ctx context.Context, domain string, options adapter.DNSQueryOptions) ([]netip.Addr, error) {
r.rulesAccess.RLock()
if r.closing {
+119 -23
View File
@@ -8,12 +8,14 @@ import (
"runtime"
"strings"
"sync"
"sync/atomic"
"syscall"
"time"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/dns/transport"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing-tun"
@@ -54,6 +56,8 @@ type Transport struct {
updatedAt time.Time
lastError error
servers []M.Socksaddr
serverTransports []adapter.DNSTransport
refreshing atomic.Bool
search []string
ndots int
attempts int
@@ -100,7 +104,7 @@ func (t *Transport) Start(stage adapter.StartStage) error {
t.interfaceCallback = t.networkManager.InterfaceMonitor().RegisterCallback(t.interfaceUpdated)
}
go func() {
_, err := t.fetch()
err := t.fetch()
if err != nil {
if errors.Is(err, errInterfaceIsCellular) && t.optional {
t.logger.Debug(E.Cause(errInterfaceIsCellular, "dhcp: fetch DNS servers"))
@@ -116,6 +120,9 @@ func (t *Transport) Close() error {
if t.interfaceCallback != nil {
t.networkManager.InterfaceMonitor().UnregisterCallback(t.interfaceCallback)
}
t.transportLock.Lock()
defer t.transportLock.Unlock()
t.closeServerTransports()
return nil
}
@@ -124,51 +131,122 @@ func (t *Transport) Reset() {
t.updatedAt = time.Time{}
t.lastError = nil
t.servers = nil
t.closeServerTransports()
t.transportLock.Unlock()
}
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
servers, err := t.fetch()
if err != nil {
return nil, E.Cause(err, "dhcp: fetch DNS servers")
func (t *Transport) closeServerTransports() {
for _, serverTransport := range t.serverTransports {
serverTransport.Close()
}
if len(servers) == 0 {
return nil, E.New("dhcp: empty DNS servers from response")
}
return t.Exchange0(ctx, message, servers)
t.serverTransports = nil
}
func (t *Transport) Exchange0(ctx context.Context, message *mDNS.Msg, servers []M.Socksaddr) (*mDNS.Msg, error) {
return t.exchangeSearch(ctx, servers, message, dns.FqdnToDomain(message.Question[0].Name))
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
done := make(chan struct{})
var (
response *mDNS.Msg
err error
)
t.ExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
response = callbackResponse
err = callbackErr
close(done)
})
<-done
return response, err
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
t.transportLock.RLock()
updatedAt := t.updatedAt
lastError := t.lastError
serverTransports := t.serverTransports
t.transportLock.RUnlock()
if lastError != nil {
callback(nil, E.Cause(lastError, "dhcp: fetch DNS servers"))
return
}
if len(serverTransports) == 0 {
go t.exchangeCold(ctx, message, callback)
return
}
if time.Since(updatedAt) >= C.DHCPTTL {
t.startRefresh()
}
t.exchangeWithTransports(ctx, message, serverTransports, callback)
}
func (t *Transport) exchangeCold(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
err := t.fetch()
if err != nil {
callback(nil, E.Cause(err, "dhcp: fetch DNS servers"))
return
}
t.transportLock.RLock()
serverTransports := t.serverTransports
t.transportLock.RUnlock()
if len(serverTransports) == 0 {
callback(nil, E.New("dhcp: empty DNS servers from response"))
return
}
t.exchangeWithTransports(ctx, message, serverTransports, callback)
}
func (t *Transport) Fetch() []M.Socksaddr {
servers, _ := t.fetch()
return servers
}
func (t *Transport) fetch() ([]M.Socksaddr, error) {
t.transportLock.RLock()
updatedAt := t.updatedAt
lastError := t.lastError
servers := t.servers
t.transportLock.RUnlock()
if lastError != nil {
return nil, lastError
return nil
}
if len(servers) > 0 && time.Since(updatedAt) >= C.DHCPTTL {
t.startRefresh()
}
return servers
}
func (t *Transport) fetch() error {
t.transportLock.RLock()
updatedAt := t.updatedAt
lastError := t.lastError
t.transportLock.RUnlock()
if lastError != nil {
return lastError
}
if time.Since(updatedAt) < C.DHCPTTL {
return servers, nil
return nil
}
t.transportLock.Lock()
defer t.transportLock.Unlock()
if time.Since(t.updatedAt) < C.DHCPTTL {
return t.servers, nil
return nil
}
err := t.updateServers()
if err != nil {
return servers, err
return t.updateServers()
}
func (t *Transport) startRefresh() {
if !t.refreshing.CompareAndSwap(false, true) {
return
}
return t.servers, nil
go func() {
defer t.refreshing.Store(false)
t.transportLock.Lock()
defer t.transportLock.Unlock()
if time.Since(t.updatedAt) < C.DHCPTTL {
return
}
err := t.updateServers()
if err != nil {
if errors.Is(err, errInterfaceIsCellular) && t.optional {
t.logger.Debug(E.Cause(err, "dhcp: refresh DNS servers"))
} else {
t.logger.Error(E.Cause(err, "dhcp: refresh DNS servers"))
}
}
}()
}
func (t *Transport) fetchInterface() (*control.Interface, error) {
@@ -222,7 +300,9 @@ func (t *Transport) updateServers() error {
}
func (t *Transport) interfaceUpdated(defaultInterface *control.Interface, flags int) {
t.transportLock.Lock()
err := t.updateServers()
t.transportLock.Unlock()
if err != nil {
if errors.Is(err, errInterfaceIsCellular) && t.optional {
t.logger.Debug(E.Cause(errInterfaceIsCellular, "dhcp: update DNS servers"))
@@ -332,6 +412,22 @@ func (t *Transport) recreateServers(iface *control.Interface, dhcpPacket *dhcpv4
if len(serverAddrs) > 0 && !slices.Equal(t.servers, serverAddrs) {
t.logger.Info("dhcp: updated DNS servers from ", iface.Name, ": [", strings.Join(common.Map(serverAddrs, M.Socksaddr.String), ","), "], search: [", strings.Join(t.search, ","), "]")
}
if !slices.Equal(t.servers, serverAddrs) || t.serverTransports == nil {
t.closeServerTransports()
serverTransports := make([]adapter.DNSTransport, 0, len(serverAddrs))
for _, serverAddr := range serverAddrs {
serverTransport := transport.NewUDPRaw(t.logger, dns.NewTransportAdapter(C.DNSTypeUDP, "", nil), t.dialer, serverAddr)
err := serverTransport.Start(adapter.StartStateStart)
if err != nil {
for _, startedTransport := range serverTransports {
startedTransport.Close()
}
return E.Cause(err, "initialize transport for ", serverAddr)
}
serverTransports = append(serverTransports, serverTransport)
}
t.serverTransports = serverTransports
}
t.servers = serverAddrs
return nil
}
+27 -126
View File
@@ -2,151 +2,52 @@ package dhcp
import (
"context"
"errors"
"math/rand"
"strings"
"syscall"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/dns/transport"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
mDNS "github.com/miekg/dns"
)
func (t *Transport) exchangeSearch(ctx context.Context, servers []M.Socksaddr, message *mDNS.Msg, domain string) (*mDNS.Msg, error) {
func (t *Transport) exchangeWithTransports(ctx context.Context, message *mDNS.Msg, serverTransports []adapter.DNSTransport, callback func(response *mDNS.Msg, err error)) {
question := message.Question[0]
domain := dns.FqdnToDomain(question.Name)
names := t.nameList(domain)
if len(names) == 0 {
return nil, E.New("dhcp: invalid domain: ", domain)
callback(nil, E.New("invalid domain: ", domain))
return
}
originalQuestion := message.Question[0]
var (
nameErrorResponse *mDNS.Msg
lastErr error
)
nameExchangers := make([]transport.AsyncExchanger, 0, len(names))
for _, fqdn := range names {
response, err := t.tryOneName(ctx, servers, fqdn, message)
if err != nil {
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
nameExchangers = append(nameExchangers, t.newNameExchanger(message, fqdn, serverTransports))
}
if nameErrorResponse != nil {
return nameErrorResponse, nil
}
return nil, lastErr
}
// 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
}
if len(serverTransports) == 1 || !(question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) {
transport.ExchangeSequential(ctx, nameExchangers, nil, callback)
} else {
transport.ExchangeRace(ctx, nameExchangers, callback)
}
}
func (t *Transport) tryOneName(ctx context.Context, servers []M.Socksaddr, fqdn string, message *mDNS.Msg) (*mDNS.Msg, error) {
sLen := len(servers)
var lastErr error
for i := 0; i < t.attempts; i++ {
for j := range sLen {
server := servers[j]
question := message.Question[0]
question.Name = fqdn
response, err := t.exchangeOne(ctx, server, question)
func (t *Transport) newNameExchanger(message *mDNS.Msg, fqdn string, serverTransports []adapter.DNSTransport) transport.AsyncExchanger {
attemptExchangers := make([]transport.AsyncExchanger, 0, t.attempts*len(serverTransports))
for range t.attempts {
for _, serverTransport := range serverTransports {
attemptExchangers = append(attemptExchangers, func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
serverTransport.ExchangeAsync(ctx, transport.NewFanOutRequest(message, fqdn, true), callback)
})
}
}
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
transport.ExchangeSequential(ctx, attemptExchangers, nil, func(response *mDNS.Msg, err error) {
if err != nil {
lastErr = err
continue
err = E.Cause(err, fqdn)
}
return response, nil
}
callback(response, err)
})
}
return nil, E.Cause(lastErr, fqdn)
}
func (t *Transport) exchangeOne(ctx context.Context, server M.Socksaddr, question mDNS.Question) (*mDNS.Msg, error) {
if server.Port == 0 {
server.Port = 53
}
request := &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Id: uint16(rand.Uint32()),
RecursionDesired: true,
AuthenticatedData: true,
},
Question: []mDNS.Question{question},
Compress: true,
}
request.SetEdns0(buf.UDPBufferSize, false)
return t.exchangeUDP(ctx, server, request)
}
func (t *Transport) exchangeUDP(ctx context.Context, server M.Socksaddr, request *mDNS.Msg) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkUDP, server)
if err != nil {
return nil, err
}
defer conn.Close()
if deadline, loaded := ctx.Deadline(); loaded && !deadline.IsZero() {
conn.SetDeadline(deadline)
}
buffer := buf.Get(buf.UDPBufferSize)
defer buf.Put(buffer)
rawMessage, err := request.PackBuffer(buffer)
if err != nil {
return nil, E.Cause(err, "pack request")
}
_, err = conn.Write(rawMessage)
if err != nil {
if errors.Is(err, syscall.EMSGSIZE) {
return t.exchangeTCP(ctx, server, request)
}
return nil, E.Cause(err, "write request")
}
n, err := conn.Read(buffer)
if err != nil {
if errors.Is(err, syscall.EMSGSIZE) {
return t.exchangeTCP(ctx, server, request)
}
return nil, E.Cause(err, "read response")
}
var response mDNS.Msg
err = response.Unpack(buffer[:n])
if err != nil {
return nil, E.Cause(err, "unpack response")
}
if response.Truncated {
return t.exchangeTCP(ctx, server, request)
}
return &response, nil
}
func (t *Transport) exchangeTCP(ctx context.Context, server M.Socksaddr, request *mDNS.Msg) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, server)
if err != nil {
return nil, err
}
defer conn.Close()
if deadline, loaded := ctx.Deadline(); loaded && !deadline.IsZero() {
conn.SetDeadline(deadline)
}
err = transport.WriteMessage(conn, 0, request)
if err != nil {
return nil, err
}
return transport.ReadMessage(conn)
}
func (t *Transport) nameList(name string) []string {
+153
View File
@@ -0,0 +1,153 @@
package transport
import (
"context"
"sync"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
mDNS "github.com/miekg/dns"
)
type AsyncExchanger = func(ctx context.Context, callback func(response *mDNS.Msg, err error))
// ExchangeSequential tries exchangers in order until accept returns true
// (nil accept means err == nil); the last result is delivered as-is.
func ExchangeSequential(ctx context.Context, exchangers []AsyncExchanger, accept func(response *mDNS.Msg, err error) bool, callback func(response *mDNS.Msg, err error)) {
if len(exchangers) == 0 {
callback(nil, E.New("missing exchangers"))
return
}
if accept == nil {
accept = func(response *mDNS.Msg, err error) bool {
return err == nil
}
}
sequential := &sequentialExchange{
ctx: ctx,
exchangers: exchangers,
accept: accept,
callback: callback,
}
sequential.run(0)
}
type sequentialExchange struct {
ctx context.Context
exchangers []AsyncExchanger
accept func(response *mDNS.Msg, err error) bool
callback func(response *mDNS.Msg, err error)
}
func (s *sequentialExchange) run(index int) {
for index < len(s.exchangers) {
ctxErr := s.ctx.Err()
if ctxErr != nil {
s.callback(nil, ctxErr)
return
}
currentIndex := index
state := &sequentialCallState{}
s.exchangers[currentIndex](s.ctx, func(response *mDNS.Msg, err error) {
if currentIndex == len(s.exchangers)-1 || s.accept(response, err) {
s.callback(response, err)
return
}
state.access.Lock()
if state.returned {
state.access.Unlock()
s.run(currentIndex + 1)
return
}
state.continued = true
state.access.Unlock()
})
state.access.Lock()
state.returned = true
continued := state.continued
state.access.Unlock()
if !continued {
return
}
index = currentIndex + 1
}
}
type sequentialCallState struct {
access sync.Mutex
returned bool
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"))
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
}
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 {
question := message.Question[0]
question.Name = fqdn
request := &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Id: message.Id,
RecursionDesired: true,
AuthenticatedData: authenticatedData,
},
Question: []mDNS.Question{question},
Compress: true,
}
request.SetEdns0(buf.UDPBufferSize, false)
return request
}
+4
View File
@@ -74,6 +74,10 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
return dns.FixedResponse(message.Id, question, []netip.Addr{address}, C.DefaultDNSTTL), nil
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
callback(t.Exchange(ctx, message))
}
func (t *Transport) Store() adapter.FakeIPStore {
return t.store
}
+4
View File
@@ -104,3 +104,7 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
Question: []mDNS.Question{question},
}, nil
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
callback(t.Exchange(ctx, message))
}
+6
View File
@@ -171,6 +171,12 @@ func (t *HTTPSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS
return response, nil
}
func (t *HTTPSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
go func() {
callback(t.Exchange(ctx, message))
}()
}
func (t *HTTPSTransport) exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
exMessage := *message
exMessage.Id = 0
+57 -18
View File
@@ -2,6 +2,8 @@ package local
import (
"context"
"sync"
"sync/atomic"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
@@ -31,15 +33,19 @@ var (
type Transport struct {
dns.TransportAdapter
ctx context.Context
logger logger.ContextLogger
hosts *hosts.File
dialer N.Dialer
preferGo bool
fallback bool
resolved ResolvedResolver
mdnsTransport adapter.DNSTransport
dhcpTransport dhcpTransport
ctx context.Context
logger logger.ContextLogger
hosts *hosts.File
dialer N.Dialer
preferGo bool
fallback bool
resolved ResolvedResolver
mdnsTransport adapter.DNSTransport
dhcpTransport dhcpTransport
system systemResolver
serverSet atomic.Pointer[localServerSet]
serverSetAccess sync.Mutex
neighborResolver adapter.NeighborResolver
neighborSuffixes []string
}
@@ -47,7 +53,6 @@ type Transport struct {
type dhcpTransport interface {
adapter.DNSTransport
Fetch() []M.Socksaddr
Exchange0(ctx context.Context, message *mDNS.Msg, servers []M.Socksaddr) (*mDNS.Msg, error)
}
func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, options option.LocalDNSServerOptions) (adapter.DNSTransport, error) {
@@ -127,10 +132,22 @@ func (t *Transport) Start(stage adapter.StartStage) error {
}
func (t *Transport) Close() error {
serverSet := t.serverSet.Swap(nil)
if serverSet != nil {
serverSet.Close()
}
t.system.close()
return common.Close(t.resolved, t.dhcpTransport, t.mdnsTransport)
}
func (t *Transport) Reset() {
serverSet := t.serverSet.Load()
if serverSet != nil {
for _, serverTransport := range serverSet.transports {
serverTransport.Reset()
}
}
t.system.reset()
if t.dhcpTransport != nil {
t.dhcpTransport.Reset()
}
@@ -149,34 +166,56 @@ func (t *Transport) PreferredDomain(domain string) bool {
}
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
done := make(chan struct{})
var (
response *mDNS.Msg
err error
)
t.ExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
response = callbackResponse
err = callbackErr
close(done)
})
<-done
return response, err
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
question := message.Question[0]
if t.hosts != nil && (question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) {
addresses := t.hosts.Lookup(dns.FqdnToDomain(question.Name))
if len(addresses) > 0 {
return dns.FixedResponse(message.Id, question, addresses, C.DefaultDNSTTL), nil
callback(dns.FixedResponse(message.Id, question, addresses, C.DefaultDNSTTL), nil)
return
}
}
response := t.lookupNeighbor(message)
if response != nil {
return response, nil
callback(response, nil)
return
}
if mdns.IsLocalDomain(question.Name) {
if C.IsDarwin {
return t.systemExchange(ctx, message)
t.systemExchangeAsync(ctx, message, callback)
return
}
return t.mdnsTransport.Exchange(ctx, message)
t.mdnsTransport.ExchangeAsync(ctx, message, callback)
return
}
if t.resolved != nil {
return t.resolved.Exchange(ctx, message)
t.resolved.ExchangeAsync(ctx, message, callback)
return
}
if t.dhcpTransport != nil {
servers := t.dhcpTransport.Fetch()
if len(servers) > 0 {
return t.dhcpTransport.Exchange0(ctx, message, servers)
t.dhcpTransport.ExchangeAsync(ctx, message, callback)
return
}
}
if t.fallback {
return t.systemExchange(ctx, message)
t.systemExchangeAsync(ctx, message, callback)
return
}
return t.exchange(ctx, message, question.Name)
t.exchangeAsync(ctx, message, question.Name, callback)
}
+327 -91
View File
@@ -10,138 +10,390 @@ import (
"io"
"net"
"os"
"sync"
"github.com/sagernet/sing-box/dns"
dnsTransport "github.com/sagernet/sing-box/dns/transport"
E "github.com/sagernet/sing/common/exceptions"
mDNS "github.com/miekg/dns"
)
func (t *Transport) systemExchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
func (t *Transport) systemExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
question := message.Question[0]
response, err := darwinLookupSystemDNS(ctx, question.Name, question.Qtype, question.Qclass)
if err != nil {
var rcodeError dns.RcodeError
if errors.As(err, &rcodeError) {
return dns.FixedResponseStatus(message, int(rcodeError)), nil
t.system.exchangeAsync(ctx, question.Name, question.Qtype, question.Qclass, func(response *mDNS.Msg, err error) {
if err != nil {
var rcodeError dns.RcodeError
if errors.As(err, &rcodeError) {
callback(dns.FixedResponseStatus(message, int(rcodeError)), nil)
return
}
callback(nil, err)
return
}
return nil, err
}
response.Id = message.Id
response.Response = true
response.RecursionAvailable = true
return response, nil
response.Id = message.Id
response.Response = true
response.RecursionAvailable = true
callback(response, nil)
})
}
// The mDNSResponder daemon speaks an undocumented binary protocol over a
// AF_UNIX SOCK_STREAM socket. The framing below is taken from the client
// stub of Apple's open-source mDNSResponder (mDNSShared/dnssd_ipc.h and
// dnssd_clientstub.c). All multi-byte fields are big-endian; for a one-shot
// query on a fresh, non-shared connection the request and every reply travel
// over the single connected stream (no SCM_RIGHTS, no return socket).
// stub of Apple's open-source mDNSResponder (mDNSShared/dnssd_ipc.h,
// dnssd_clientstub.c and uds_daemon.c). All multi-byte fields are
// big-endian. A connection opened with connection_request acts as a shared
// connection (DNSServiceCreateConnection): subsequent requests on the same
// stream carry a unique client_context in header bytes 16-24, which the
// daemon echoes back in every reply, allowing concurrent queries to be
// demultiplexed. With IPC_FLAGS_NOERRSD set the daemon does not expect the
// SCM_RIGHTS error-return socket used by Apple's stub; request errors are
// instead delivered as async_error_op replies, and success produces no
// acknowledgment at all. A query is cancelled by sending cancel_request
// with the same client_context and no payload.
const (
mdnsResponderSocketPath = "/var/run/mDNSResponder"
mdnsResponderSocketEnv = "DNSSD_UDS_PATH"
mdnsResponderVersion = 1
mdnsResponderHeaderLength = 28
mdnsResponderQueryRequest = 8 // query_request
mdnsResponderQueryReply = 68 // query_reply_op
mdnsResponderSocketPath = "/var/run/mDNSResponder"
mdnsResponderSocketEnv = "DNSSD_UDS_PATH"
mdnsResponderVersion = 1
mdnsResponderHeaderLength = 28
mdnsResponderConnectionRequest = 1 // connection_request
mdnsResponderQueryRequest = 8 // query_request
mdnsResponderCancelRequest = 63 // cancel_request
mdnsResponderQueryReply = 68 // query_reply_op
mdnsResponderAsyncErrorReply = 73 // async_error_op
mdnsResponderFlagMoreComing = 0x1
mdnsResponderFlagAdd = 0x2
mdnsResponderFlagReturnIntermediates = 0x1000
mdnsResponderFlagShareConnection = 0x4000
mdnsResponderFlagTimeout = 0x10000
mdnsResponderIPCFlagNoErrorSocket = 0x4 // IPC_FLAGS_NOERRSD
mdnsResponderErrNoError = 0
mdnsResponderErrNoSuchName = -65538
mdnsResponderErrNoSuchRecord = -65554
mdnsResponderErrTimeout = -65568
mdnsResponderMaxReplyLength = 1 << 20
)
func darwinLookupSystemDNS(ctx context.Context, name string, qtype, qclass uint16) (*mDNS.Msg, error) {
type systemResolver struct {
initOnce sync.Once
connection *dnsTransport.ConnPool[net.Conn]
queryAccess sync.Mutex
queryId uint64
queries map[uint64]*systemPendingQuery
}
type systemPendingQuery struct {
conn net.Conn
name string
qtype uint16
qclass uint16
answers []mDNS.RR
hasFinalAnswer bool
ready bool
callback func(response *mDNS.Msg, err error)
stopContext func() bool
stopConn func() bool
}
type systemCompletion struct {
pending *systemPendingQuery
err error
}
func (r *systemResolver) init() {
r.queries = make(map[uint64]*systemPendingQuery)
r.connection = dnsTransport.NewConnPool(dnsTransport.ConnPoolOptions[net.Conn]{
Mode: dnsTransport.ConnPoolSingle,
IsAlive: func(conn net.Conn) bool {
return conn != nil
},
Close: func(conn net.Conn, cause error) {
conn.Close()
},
})
}
func (r *systemResolver) close() {
r.initOnce.Do(r.init)
_ = r.connection.Close()
}
func (r *systemResolver) reset() {
r.initOnce.Do(r.init)
r.connection.Reset()
}
func (r *systemResolver) exchangeAsync(ctx context.Context, name string, qtype uint16, qclass uint16, callback func(response *mDNS.Msg, err error)) {
r.initOnce.Do(r.init)
for firstAttempt := true; ; firstAttempt = false {
conn, connCtx, created, err := r.connection.AcquireShared(ctx, r.dial)
if err != nil {
callback(nil, err)
return
}
if created {
go r.recvLoop(conn)
}
queryId := r.register(ctx, connCtx, conn, name, qtype, qclass, callback)
_, writeErr := conn.Write(buildQueryRequest(queryId, name, qtype, qclass))
if writeErr == nil {
return
}
pending := r.take(queryId)
r.connection.Invalidate(conn, writeErr)
if pending == nil {
return
}
if !created && firstAttempt {
continue
}
callback(nil, E.Cause(writeErr, "write mDNSResponder query"))
return
}
}
func (r *systemResolver) dial(ctx context.Context) (net.Conn, error) {
socketPath := cmp.Or(os.Getenv(mdnsResponderSocketEnv), mdnsResponderSocketPath)
var dialer net.Dialer
conn, err := dialer.DialContext(ctx, "unix", socketPath)
if err != nil {
return nil, E.Cause(err, "connect mDNSResponder")
}
defer conn.Close()
stopCancel := context.AfterFunc(ctx, func() {
conn.Close()
})
defer stopCancel()
_, err = conn.Write(buildQueryRequest(name, qtype, qclass))
err = writeConnectionRequest(conn)
stopCancel()
if err != nil {
return nil, contextError(ctx, E.Cause(err, "write mDNSResponder query"))
conn.Close()
return nil, contextError(ctx, err)
}
return conn, nil
}
func writeConnectionRequest(conn net.Conn) error {
_, err := conn.Write(appendResponderHeader(make([]byte, 0, mdnsResponderHeaderLength), mdnsResponderConnectionRequest, 0, 0, 0))
if err != nil {
return E.Cause(err, "write mDNSResponder connection request")
}
var status [4]byte
_, err = io.ReadFull(conn, status[:])
if err != nil {
return nil, contextError(ctx, E.Cause(err, "read mDNSResponder status"))
return E.Cause(err, "read mDNSResponder connection status")
}
statusCode := int32(binary.BigEndian.Uint32(status[:]))
if statusCode != mdnsResponderErrNoError {
return nil, darwinResolverError(name, statusCode)
return E.New("mDNSResponder connection request failed: error ", statusCode)
}
return readQueryResponse(ctx, conn, name, qtype, qclass)
return nil
}
func readQueryResponse(ctx context.Context, conn net.Conn, name string, qtype, qclass uint16) (*mDNS.Msg, error) {
var answers []mDNS.RR
var hasFinalAnswer bool
func (r *systemResolver) register(ctx context.Context, connCtx context.Context, conn net.Conn, name string, qtype uint16, qclass uint16, callback func(response *mDNS.Msg, err error)) uint64 {
r.queryAccess.Lock()
defer r.queryAccess.Unlock()
r.queryId++
queryId := r.queryId
pending := &systemPendingQuery{
conn: conn,
name: name,
qtype: qtype,
qclass: qclass,
callback: callback,
}
r.queries[queryId] = pending
pending.stopContext = context.AfterFunc(ctx, func() {
r.cancelQuery(queryId, ctx)
})
pending.stopConn = context.AfterFunc(connCtx, func() {
r.completeConnClosed(queryId, connCtx)
})
return queryId
}
func (r *systemResolver) take(queryId uint64) *systemPendingQuery {
r.queryAccess.Lock()
pending, loaded := r.queries[queryId]
if !loaded {
r.queryAccess.Unlock()
return nil
}
delete(r.queries, queryId)
r.queryAccess.Unlock()
pending.stopContext()
pending.stopConn()
return pending
}
func (r *systemResolver) cancelQuery(queryId uint64, ctx context.Context) {
pending := r.take(queryId)
if pending == nil {
return
}
_, err := pending.conn.Write(appendResponderHeader(make([]byte, 0, mdnsResponderHeaderLength), mdnsResponderCancelRequest, 0, queryId, 0))
if err != nil {
r.connection.Invalidate(pending.conn, err)
} else {
r.connection.Release(pending.conn, true)
}
pending.callback(nil, ctx.Err())
}
func (r *systemResolver) completeConnClosed(queryId uint64, connCtx context.Context) {
pending := r.take(queryId)
if pending == nil {
return
}
pending.callback(nil, context.Cause(connCtx))
}
func (r *systemResolver) finish(pending *systemPendingQuery, err error) {
pending.stopContext()
pending.stopConn()
r.connection.Release(pending.conn, true)
if err != nil {
pending.callback(nil, err)
return
}
pending.callback(&mDNS.Msg{
Question: []mDNS.Question{{Name: mDNS.Fqdn(pending.name), Qtype: pending.qtype, Qclass: pending.qclass}},
Answer: pending.answers,
}, nil)
}
func (r *systemResolver) recvLoop(conn net.Conn) {
for {
reply, replyErr := readReply(conn)
if replyErr != nil {
return nil, contextError(ctx, E.Cause(replyErr, "read mDNSResponder reply"))
operation, clientContext, data, err := readResponderReply(conn)
if err != nil {
r.connection.Invalidate(conn, err)
return
}
if reply.errorCode != mdnsResponderErrNoError {
if len(answers) == 0 {
return nil, darwinResolverError(name, reply.errorCode)
switch operation {
case mdnsResponderQueryReply:
reply, parseErr := parseResponderReply(data)
if parseErr != nil {
r.connection.Invalidate(conn, parseErr)
return
}
r.handleQueryReply(clientContext, reply)
case mdnsResponderAsyncErrorReply:
if len(data) >= 12 {
r.completeQueryError(clientContext, binary.BigEndian.Uint32(data[0:4]), int32(binary.BigEndian.Uint32(data[8:12])))
}
break
}
if reply.flags&mdnsResponderFlagAdd != 0 && len(reply.rdata) > 0 {
record, buildErr := buildResourceRecord(reply)
if buildErr == nil {
answers = append(answers, record)
if record.Header().Rrtype == qtype {
hasFinalAnswer = true
}
}
// On a shared connection MoreComing applies collectively to all operations
// (dns_sd.h "Collective kDNSServiceFlagsMoreComing flag"): the daemon sets it
// whenever another reply, for any query, is queued behind this one. A reply
// without it is therefore a connection-wide flush point, at which every query
// that already collected its final answer is completed.
func (r *systemResolver) handleQueryReply(queryId uint64, reply mdnsResponderReply) {
var completions []systemCompletion
r.queryAccess.Lock()
pending, loaded := r.queries[queryId]
if loaded {
if reply.errorCode != mdnsResponderErrNoError {
delete(r.queries, queryId)
if len(pending.answers) > 0 {
completions = append(completions, systemCompletion{pending: pending})
} else {
completions = append(completions, systemCompletion{pending: pending, err: darwinResolverError(pending.name, reply.errorCode)})
}
} else {
if reply.flags&mdnsResponderFlagAdd != 0 && len(reply.rdata) > 0 {
record, buildErr := buildResourceRecord(reply)
if buildErr == nil {
pending.answers = append(pending.answers, record)
if record.Header().Rrtype == pending.qtype {
pending.hasFinalAnswer = true
}
}
}
}
if reply.flags&mdnsResponderFlagMoreComing != 0 {
continue
}
if hasFinalAnswer && reply.rrtype == qtype {
break
if pending.hasFinalAnswer && reply.rrtype == pending.qtype {
pending.ready = true
}
}
}
response := new(mDNS.Msg)
response.Question = []mDNS.Question{{Name: mDNS.Fqdn(name), Qtype: qtype, Qclass: qclass}}
response.Answer = answers
return response, nil
if reply.flags&mdnsResponderFlagMoreComing == 0 {
completions = r.collectReadyLocked(completions)
}
r.queryAccess.Unlock()
for _, completion := range completions {
r.finish(completion.pending, completion.err)
}
}
func buildQueryRequest(name string, qtype, qclass uint16) []byte {
payload := make([]byte, 0, 8+len(name)+1+4)
payload = binary.BigEndian.AppendUint32(payload, mdnsResponderFlagReturnIntermediates|mdnsResponderFlagTimeout)
payload = binary.BigEndian.AppendUint32(payload, 0) // interfaceIndex
payload = append(payload, name...)
payload = append(payload, 0) // C string terminator
payload = binary.BigEndian.AppendUint16(payload, qtype)
payload = binary.BigEndian.AppendUint16(payload, qclass)
func (r *systemResolver) completeQueryError(queryId uint64, flags uint32, errorCode int32) {
var completions []systemCompletion
r.queryAccess.Lock()
pending, loaded := r.queries[queryId]
if loaded {
delete(r.queries, queryId)
completions = append(completions, systemCompletion{pending: pending, err: darwinResolverError(pending.name, errorCode)})
}
if flags&mdnsResponderFlagMoreComing == 0 {
completions = r.collectReadyLocked(completions)
}
r.queryAccess.Unlock()
for _, completion := range completions {
r.finish(completion.pending, completion.err)
}
}
message := make([]byte, mdnsResponderHeaderLength, mdnsResponderHeaderLength+len(payload))
binary.BigEndian.PutUint32(message[0:], mdnsResponderVersion)
binary.BigEndian.PutUint32(message[4:], uint32(len(payload)))
binary.BigEndian.PutUint32(message[8:], 0) // ipc_flags
binary.BigEndian.PutUint32(message[12:], mdnsResponderQueryRequest)
// message[16:24] client_context and message[24:28] reg_index stay zero.
return append(message, payload...)
func (r *systemResolver) collectReadyLocked(completions []systemCompletion) []systemCompletion {
for queryId, pending := range r.queries {
if pending.ready {
delete(r.queries, queryId)
completions = append(completions, systemCompletion{pending: pending})
}
}
return completions
}
func appendResponderHeader(buffer []byte, operation uint32, dataLength int, clientContext uint64, ipcFlags uint32) []byte {
buffer = binary.BigEndian.AppendUint32(buffer, mdnsResponderVersion)
buffer = binary.BigEndian.AppendUint32(buffer, uint32(dataLength))
buffer = binary.BigEndian.AppendUint32(buffer, ipcFlags)
buffer = binary.BigEndian.AppendUint32(buffer, operation)
buffer = binary.BigEndian.AppendUint64(buffer, clientContext)
buffer = binary.BigEndian.AppendUint32(buffer, 0) // reg_index
return buffer
}
func buildQueryRequest(queryId uint64, name string, qtype uint16, qclass uint16) []byte {
payloadLength := 4 + 4 + len(name) + 1 + 2 + 2
message := make([]byte, 0, mdnsResponderHeaderLength+payloadLength)
message = appendResponderHeader(message, mdnsResponderQueryRequest, payloadLength, queryId, mdnsResponderIPCFlagNoErrorSocket)
message = binary.BigEndian.AppendUint32(message, mdnsResponderFlagShareConnection|mdnsResponderFlagReturnIntermediates|mdnsResponderFlagTimeout)
message = binary.BigEndian.AppendUint32(message, 0) // interfaceIndex
message = append(message, name...)
message = append(message, 0) // C string terminator
message = binary.BigEndian.AppendUint16(message, qtype)
message = binary.BigEndian.AppendUint16(message, qclass)
return message
}
func readResponderReply(conn net.Conn) (operation uint32, clientContext uint64, data []byte, err error) {
var header [mdnsResponderHeaderLength]byte
_, err = io.ReadFull(conn, header[:])
if err != nil {
return
}
dataLength := binary.BigEndian.Uint32(header[4:8])
if dataLength > mdnsResponderMaxReplyLength {
err = E.New("oversized mDNSResponder reply: ", dataLength)
return
}
operation = binary.BigEndian.Uint32(header[12:16])
clientContext = binary.BigEndian.Uint64(header[16:24])
data = make([]byte, dataLength)
_, err = io.ReadFull(conn, data)
return
}
type mdnsResponderReply struct {
@@ -154,24 +406,8 @@ type mdnsResponderReply struct {
rdata []byte
}
func readReply(conn net.Conn) (mdnsResponderReply, error) {
func parseResponderReply(data []byte) (mdnsResponderReply, error) {
var reply mdnsResponderReply
var header [mdnsResponderHeaderLength]byte
_, err := io.ReadFull(conn, header[:])
if err != nil {
return reply, err
}
dataLength := binary.BigEndian.Uint32(header[4:8])
operation := binary.BigEndian.Uint32(header[12:16])
if operation != mdnsResponderQueryReply {
return reply, E.New("unexpected mDNSResponder reply op ", operation)
}
data := make([]byte, dataLength)
_, err = io.ReadFull(conn, data)
if err != nil {
return reply, err
}
reader := replyReader{data: data}
reply.flags = reader.uint32()
reader.uint32() // interfaceIndex
+65 -3
View File
@@ -7,6 +7,7 @@ import (
"context"
"net"
"os"
"sync"
"testing"
"time"
@@ -26,9 +27,25 @@ func requireMDNSResponder(t *testing.T) {
conn.Close()
}
func systemExchangeForTest(ctx context.Context, transport *Transport, message *mDNS.Msg) (*mDNS.Msg, error) {
done := make(chan struct{})
var (
response *mDNS.Msg
err error
)
transport.systemExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
response = callbackResponse
err = callbackErr
close(done)
})
<-done
return response, err
}
func TestSystemExchangeLoopback(t *testing.T) {
requireMDNSResponder(t)
transport := &Transport{}
defer transport.system.close()
for _, testCase := range []struct {
qtype uint16
expected net.IP
@@ -39,7 +56,7 @@ func TestSystemExchangeLoopback(t *testing.T) {
message := new(mDNS.Msg)
message.SetQuestion("localhost.", testCase.qtype)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
response, err := transport.systemExchange(ctx, message)
response, err := systemExchangeForTest(ctx, transport, message)
cancel()
if err != nil {
t.Fatalf("%s localhost: %v", mDNS.TypeToString[testCase.qtype], err)
@@ -67,13 +84,15 @@ func TestSystemExchangeLoopback(t *testing.T) {
func TestSystemExchangeNoData(t *testing.T) {
requireMDNSResponder(t)
transport := &Transport{}
defer transport.system.close()
message := new(mDNS.Msg)
// localhost has no MX record, so the daemon reports NoSuchRecord, which must
// surface as an empty NOERROR response rather than an error.
message.SetQuestion("localhost.", mDNS.TypeMX)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
response, err := (&Transport{}).systemExchange(ctx, message)
response, err := systemExchangeForTest(ctx, transport, message)
if err != nil {
t.Fatalf("MX localhost: %v", err)
}
@@ -87,12 +106,14 @@ func TestSystemExchangeNoData(t *testing.T) {
func TestSystemExchangeCancel(t *testing.T) {
requireMDNSResponder(t)
transport := &Transport{}
defer transport.system.close()
message := new(mDNS.Msg)
message.SetQuestion("localhost.", mDNS.TypeA)
ctx, cancel := context.WithCancel(context.Background())
cancel()
start := time.Now()
_, err := (&Transport{}).systemExchange(ctx, message)
_, err := systemExchangeForTest(ctx, transport, message)
elapsed := time.Since(start)
if err == nil {
t.Fatal("expected error for cancelled context")
@@ -101,3 +122,44 @@ func TestSystemExchangeCancel(t *testing.T) {
t.Fatalf("cancellation too slow: %s", elapsed)
}
}
func TestSystemExchangeConcurrent(t *testing.T) {
requireMDNSResponder(t)
transport := &Transport{}
defer transport.system.close()
var waitGroup sync.WaitGroup
errors := make(chan error, 16)
for i := range 16 {
qtype := mDNS.TypeA
if i%2 == 1 {
qtype = mDNS.TypeAAAA
}
waitGroup.Add(1)
go func() {
defer waitGroup.Done()
message := new(mDNS.Msg)
message.SetQuestion("localhost.", qtype)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
response, exchangeErr := systemExchangeForTest(ctx, transport, message)
if exchangeErr != nil {
errors <- exchangeErr
return
}
if len(response.Answer) == 0 {
errors <- context.DeadlineExceeded
}
}()
}
waitGroup.Wait()
close(errors)
for exchangeErr := range errors {
t.Fatal("concurrent query failed: ", exchangeErr)
}
transport.system.queryAccess.Lock()
pendingCount := len(transport.system.queries)
transport.system.queryAccess.Unlock()
if pendingCount != 0 {
t.Fatalf("expected no pending queries after completion, got %d", pendingCount)
}
}
+8 -2
View File
@@ -9,6 +9,12 @@ import (
mDNS "github.com/miekg/dns"
)
func (t *Transport) systemExchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
return nil, os.ErrInvalid
type systemResolver struct{}
func (r *systemResolver) close() {}
func (r *systemResolver) reset() {}
func (t *Transport) systemExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
callback(nil, os.ErrInvalid)
}
+1
View File
@@ -10,4 +10,5 @@ type ResolvedResolver interface {
Start() error
Close() error
Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error)
ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error))
}
@@ -159,6 +159,51 @@ func (t *DBusResolvedResolver) Exchange(ctx context.Context, message *mDNS.Msg)
return t.exchangeServerSet(ctx, message, refreshedServerSet)
}
func (t *DBusResolvedResolver) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
serverSet := t.savedServerSet.Load()
if serverSet == nil {
go func() {
callback(t.Exchange(ctx, message))
}()
return
}
t.exchangeServerSetAsync(ctx, message, serverSet, func(response *mDNS.Msg, err error) {
if err == nil {
callback(response, nil)
return
}
go func() {
t.updateStatus()
refreshedServerSet := t.savedServerSet.Load()
if refreshedServerSet == nil || refreshedServerSet == serverSet {
callback(nil, err)
return
}
t.exchangeServerSetAsync(ctx, message, refreshedServerSet, callback)
}()
})
}
func (t *DBusResolvedResolver) exchangeServerSetAsync(ctx context.Context, message *mDNS.Msg, serverSet *resolvedServerSet, callback func(response *mDNS.Msg, err error)) {
if len(serverSet.servers) == 0 {
callback(nil, E.New("link has no DNS servers configured"))
return
}
serverExchangers := make([]dnsTransport.AsyncExchanger, 0, len(serverSet.servers))
for _, server := range serverSet.servers {
serverExchangers = append(serverExchangers, func(exchangeCtx context.Context, exchangeCallback func(response *mDNS.Msg, err error)) {
server.primaryTransport.ExchangeAsync(exchangeCtx, message, func(response *mDNS.Msg, exchangeErr error) {
if exchangeErr != nil && server.fallbackTransport != nil {
server.fallbackTransport.ExchangeAsync(exchangeCtx, message, exchangeCallback)
return
}
exchangeCallback(response, exchangeErr)
})
})
}
dnsTransport.ExchangeSequential(ctx, serverExchangers, nil, callback)
}
func (t *DBusResolvedResolver) loopUpdateStatus() {
signalChan := make(chan *dbus.Signal, 1)
t.systemBus.Signal(signalChan)
+87 -152
View File
@@ -2,182 +2,117 @@ package local
import (
"context"
"errors"
"math/rand"
"syscall"
"time"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/dns/transport"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
mDNS "github.com/miekg/dns"
)
func (t *Transport) exchange(ctx context.Context, message *mDNS.Msg, domain string) (*mDNS.Msg, error) {
systemConfig := getSystemDNSConfig(t.ctx)
if systemConfig.singleRequest || !(message.Question[0].Qtype == mDNS.TypeA || message.Question[0].Qtype == mDNS.TypeAAAA) {
return t.exchangeSingleRequest(ctx, systemConfig, message, domain)
} else {
return t.exchangeParallel(ctx, systemConfig, message, domain)
type localServerSet struct {
config *dnsConfig
transports []adapter.DNSTransport
}
func (s *localServerSet) Close() {
for _, serverTransport := range s.transports {
serverTransport.Close()
}
}
func (t *Transport) exchangeSingleRequest(ctx context.Context, systemConfig *dnsConfig, message *mDNS.Msg, domain string) (*mDNS.Msg, error) {
var lastErr error
for _, fqdn := range systemConfig.nameList(domain) {
response, err := t.tryOneName(ctx, systemConfig, fqdn, message)
func (t *Transport) serverSetFor(systemConfig *dnsConfig) (*localServerSet, error) {
serverSet := t.serverSet.Load()
if serverSet != nil && serverSet.config == systemConfig {
return serverSet, nil
}
t.serverSetAccess.Lock()
defer t.serverSetAccess.Unlock()
serverSet = t.serverSet.Load()
if serverSet != nil && serverSet.config == systemConfig {
return serverSet, nil
}
transports := make([]adapter.DNSTransport, 0, len(systemConfig.servers))
for _, server := range systemConfig.servers {
serverAddr := M.ParseSocksaddr(server)
if serverAddr.Port == 0 {
serverAddr.Port = 53
}
var serverTransport adapter.DNSTransport
if systemConfig.useTCP {
serverTransport = transport.NewTCPRaw(dns.NewTransportAdapter(C.DNSTypeTCP, "", nil), t.dialer, serverAddr)
} else {
serverTransport = transport.NewUDPRaw(t.logger, dns.NewTransportAdapter(C.DNSTypeUDP, "", nil), t.dialer, serverAddr)
}
err := serverTransport.Start(adapter.StartStateStart)
if err != nil {
lastErr = err
continue
for _, startedTransport := range transports {
startedTransport.Close()
}
return nil, E.Cause(err, "initialize transport for ", serverAddr)
}
return response, nil
transports = append(transports, serverTransport)
}
return nil, lastErr
newServerSet := &localServerSet{
config: systemConfig,
transports: transports,
}
oldServerSet := t.serverSet.Swap(newServerSet)
if oldServerSet != nil {
oldServerSet.Close()
}
return newServerSet, nil
}
func (t *Transport) exchangeParallel(ctx context.Context, systemConfig *dnsConfig, message *mDNS.Msg, domain string) (*mDNS.Msg, error) {
returned := make(chan struct{})
defer close(returned)
type queryResult struct {
response *mDNS.Msg
err error
func (t *Transport) exchangeAsync(ctx context.Context, message *mDNS.Msg, domain string, callback func(response *mDNS.Msg, err error)) {
systemConfig := getSystemDNSConfig(t.ctx)
serverSet, err := t.serverSetFor(systemConfig)
if err != nil {
callback(nil, err)
return
}
results := make(chan queryResult)
startRacer := func(ctx context.Context, fqdn string) {
response, err := t.tryOneName(ctx, systemConfig, fqdn, message)
select {
case results <- queryResult{response, err}:
case <-returned:
}
names := systemConfig.nameList(domain)
if len(names) == 0 {
callback(nil, E.New("invalid domain: ", domain))
return
}
queryCtx, queryCancel := context.WithCancel(ctx)
defer queryCancel()
var nameCount int
for _, fqdn := range systemConfig.nameList(domain) {
nameCount++
go startRacer(queryCtx, fqdn)
nameExchangers := make([]transport.AsyncExchanger, 0, len(names))
for _, fqdn := range names {
nameExchangers = append(nameExchangers, newNameExchanger(systemConfig, serverSet, message, 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...)
}
}
}
}
func (t *Transport) tryOneName(ctx context.Context, config *dnsConfig, fqdn string, message *mDNS.Msg) (*mDNS.Msg, error) {
serverOffset := config.serverOffset()
sLen := uint32(len(config.servers))
var lastErr error
for i := 0; i < config.attempts; i++ {
for j := range sLen {
server := config.servers[(serverOffset+j)%sLen]
question := message.Question[0]
question.Name = fqdn
response, err := t.exchangeOne(ctx, M.ParseSocksaddr(server), question, config.timeout, config.useTCP, config.trustAD)
if err != nil {
lastErr = err
continue
}
return response, nil
}
}
return nil, E.Cause(lastErr, fqdn)
}
func (t *Transport) exchangeOne(ctx context.Context, server M.Socksaddr, question mDNS.Question, timeout time.Duration, useTCP, ad bool) (*mDNS.Msg, error) {
if server.Port == 0 {
server.Port = 53
}
request := &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Id: uint16(rand.Uint32()),
RecursionDesired: true,
AuthenticatedData: ad,
},
Question: []mDNS.Question{question},
Compress: true,
}
request.SetEdns0(buf.UDPBufferSize, false)
if !useTCP {
return t.exchangeUDP(ctx, server, request, timeout)
question := message.Question[0]
if systemConfig.singleRequest || !(question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) {
transport.ExchangeSequential(ctx, nameExchangers, nil, callback)
} else {
return t.exchangeTCP(ctx, server, request, timeout)
transport.ExchangeRace(ctx, nameExchangers, callback)
}
}
func (t *Transport) exchangeUDP(ctx context.Context, server M.Socksaddr, request *mDNS.Msg, timeout time.Duration) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkUDP, server)
if err != nil {
return nil, err
}
defer conn.Close()
if deadline, loaded := ctx.Deadline(); loaded && !deadline.IsZero() {
newDeadline := time.Now().Add(timeout)
if deadline.After(newDeadline) {
deadline = newDeadline
func newNameExchanger(systemConfig *dnsConfig, serverSet *localServerSet, message *mDNS.Msg, fqdn string) transport.AsyncExchanger {
serverOffset := systemConfig.serverOffset()
serverCount := uint32(len(serverSet.transports))
attemptExchangers := make([]transport.AsyncExchanger, 0, systemConfig.attempts*int(serverCount))
for i := 0; i < systemConfig.attempts; i++ {
for j := range serverCount {
serverTransport := serverSet.transports[(serverOffset+j)%serverCount]
attemptExchangers = append(attemptExchangers, func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
attemptCtx, cancel := context.WithTimeout(ctx, systemConfig.timeout)
serverTransport.ExchangeAsync(attemptCtx, transport.NewFanOutRequest(message, fqdn, systemConfig.trustAD), func(response *mDNS.Msg, err error) {
cancel()
callback(response, err)
})
})
}
conn.SetDeadline(deadline)
}
buffer := buf.Get(buf.UDPBufferSize)
defer buf.Put(buffer)
rawMessage, err := request.PackBuffer(buffer)
if err != nil {
return nil, E.Cause(err, "pack request")
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
transport.ExchangeSequential(ctx, attemptExchangers, nil, func(response *mDNS.Msg, err error) {
if err != nil {
err = E.Cause(err, fqdn)
}
callback(response, err)
})
}
_, err = conn.Write(rawMessage)
if err != nil {
if errors.Is(err, syscall.EMSGSIZE) {
return t.exchangeTCP(ctx, server, request, timeout)
}
return nil, E.Cause(err, "write request")
}
n, err := conn.Read(buffer)
if err != nil {
if errors.Is(err, syscall.EMSGSIZE) {
return t.exchangeTCP(ctx, server, request, timeout)
}
return nil, E.Cause(err, "read response")
}
var response mDNS.Msg
err = response.Unpack(buffer[:n])
if err != nil {
return nil, E.Cause(err, "unpack response")
}
if response.Truncated {
return t.exchangeTCP(ctx, server, request, timeout)
}
return &response, nil
}
func (t *Transport) exchangeTCP(ctx context.Context, server M.Socksaddr, request *mDNS.Msg, timeout time.Duration) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, server)
if err != nil {
return nil, err
}
defer conn.Close()
if deadline, loaded := ctx.Deadline(); loaded && !deadline.IsZero() {
newDeadline := time.Now().Add(timeout)
if deadline.After(newDeadline) {
deadline = newDeadline
}
conn.SetDeadline(deadline)
}
err = transport.WriteMessage(conn, 0, request)
if err != nil {
return nil, err
}
return transport.ReadMessage(conn)
}
+6
View File
@@ -159,6 +159,12 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
return nil, E.New("mdns: query timeout")
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
go func() {
callback(t.Exchange(ctx, message))
}()
}
type exchangeResult struct {
response *mDNS.Msg
err error
+209
View File
@@ -0,0 +1,209 @@
package transport
import (
"context"
"errors"
"net"
"sync"
"sync/atomic"
E "github.com/sagernet/sing/common/exceptions"
mDNS "github.com/miekg/dns"
)
type queryMultiplexerOptions struct {
dial func(ctx context.Context) (net.Conn, error)
write func(conn net.Conn, message *mDNS.Msg, queryId uint16) error
readNext func(conn net.Conn) (*mDNS.Msg, error)
}
type queryMultiplexer struct {
options queryMultiplexerOptions
connection *ConnPool[*multiplexConn]
queryAccess sync.Mutex
queryId uint16
queries map[uint16]*pendingQuery
}
type multiplexConn struct {
net.Conn
readEpoch atomic.Uint64
}
type pendingQuery struct {
conn *multiplexConn
originalId uint16
readEpoch uint64
callback func(response *mDNS.Msg, err error)
stopContext func() bool
stopConn func() bool
}
func newQueryMultiplexer(options queryMultiplexerOptions) *queryMultiplexer {
return &queryMultiplexer{
options: options,
queries: make(map[uint16]*pendingQuery),
connection: NewConnPool(ConnPoolOptions[*multiplexConn]{
Mode: ConnPoolSingle,
IsAlive: func(conn *multiplexConn) bool {
return conn != nil
},
Close: func(conn *multiplexConn, cause error) {
conn.Close()
},
}),
}
}
func (m *queryMultiplexer) Close() error {
return m.connection.Close()
}
func (m *queryMultiplexer) Reset() {
m.connection.Reset()
}
func (m *queryMultiplexer) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
done := make(chan struct{})
var (
response *mDNS.Msg
err error
)
m.ExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
response = callbackResponse
err = callbackErr
close(done)
})
<-done
return response, err
}
func (m *queryMultiplexer) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
for firstAttempt := true; ; firstAttempt = false {
conn, connCtx, created, err := m.connection.AcquireShared(ctx, m.dialConn)
if err != nil {
callback(nil, err)
return
}
if created {
go m.recvLoop(conn)
}
queryId, err := m.register(ctx, connCtx, conn, message.Id, callback)
if err != nil {
m.connection.Release(conn, true)
callback(nil, err)
return
}
writeErr := m.options.write(conn, message, queryId)
if writeErr == nil {
return
}
pending := m.take(queryId)
m.connection.Invalidate(conn, writeErr)
if pending == nil {
return
}
if !created && firstAttempt {
continue
}
callback(nil, E.Cause(writeErr, "write request"))
return
}
}
func (m *queryMultiplexer) dialConn(ctx context.Context) (*multiplexConn, error) {
conn, err := m.options.dial(ctx)
if err != nil {
return nil, err
}
return &multiplexConn{Conn: conn}, nil
}
func (m *queryMultiplexer) register(ctx context.Context, connCtx context.Context, conn *multiplexConn, originalId uint16, callback func(response *mDNS.Msg, err error)) (uint16, error) {
m.queryAccess.Lock()
defer m.queryAccess.Unlock()
start := m.queryId
for {
m.queryId++
if _, exists := m.queries[m.queryId]; !exists {
break
}
if m.queryId == start {
return 0, E.New("no available query ID")
}
}
queryId := m.queryId
pending := &pendingQuery{
conn: conn,
originalId: originalId,
readEpoch: conn.readEpoch.Load(),
callback: callback,
}
m.queries[queryId] = pending
pending.stopContext = context.AfterFunc(ctx, func() {
m.completeContextDone(queryId, ctx)
})
pending.stopConn = context.AfterFunc(connCtx, func() {
m.complete(queryId, nil, context.Cause(connCtx), false)
})
return queryId, nil
}
func (m *queryMultiplexer) take(queryId uint16) *pendingQuery {
m.queryAccess.Lock()
pending, loaded := m.queries[queryId]
if !loaded {
m.queryAccess.Unlock()
return nil
}
delete(m.queries, queryId)
m.queryAccess.Unlock()
pending.stopContext()
pending.stopConn()
return pending
}
func (m *queryMultiplexer) complete(queryId uint16, response *mDNS.Msg, err error, releaseConn bool) {
pending := m.take(queryId)
if pending == nil {
return
}
if releaseConn {
m.connection.Release(pending.conn, true)
}
if response != nil {
response.Id = pending.originalId
}
pending.callback(response, err)
}
func (m *queryMultiplexer) completeContextDone(queryId uint16, ctx context.Context) {
pending := m.take(queryId)
if pending == nil {
return
}
err := ctx.Err()
if errors.Is(err, context.DeadlineExceeded) && pending.conn.readEpoch.Load() == pending.readEpoch {
m.connection.Invalidate(pending.conn, err)
} else {
m.connection.Release(pending.conn, true)
}
pending.callback(nil, err)
}
func (m *queryMultiplexer) recvLoop(conn *multiplexConn) {
for {
message, err := m.options.readNext(conn)
if err != nil {
m.connection.Invalidate(conn, err)
return
}
conn.readEpoch.Add(1)
if message == nil {
continue
}
m.complete(message.Id, message, nil, true)
}
}
+164
View File
@@ -0,0 +1,164 @@
package transport
import (
"context"
"errors"
"io"
"net"
"testing"
"time"
mDNS "github.com/miekg/dns"
)
func TestMultiplexerTimeoutInvalidatesConn(t *testing.T) {
t.Parallel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer listener.Close()
accepted := make(chan net.Conn, 16)
go func() {
for {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
return
}
accepted <- conn
}
}()
multiplexer := newQueryMultiplexer(queryMultiplexerOptions{
dial: func(ctx context.Context) (net.Conn, error) {
return net.Dial("tcp", listener.Addr().String())
},
write: func(conn net.Conn, message *mDNS.Msg, queryId uint16) error {
return WriteMessage(conn, queryId, message)
},
readNext: func(conn net.Conn) (*mDNS.Msg, error) {
return ReadMessage(conn)
},
})
defer multiplexer.Close()
message := new(mDNS.Msg)
message.SetQuestion("example.com.", mDNS.TypeA)
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
start := time.Now()
_, err = multiplexer.Exchange(ctx, message)
elapsed := time.Since(start)
if !errors.Is(err, context.DeadlineExceeded) {
t.Fatal("expected deadline exceeded, got ", err)
}
if elapsed > 2*time.Second {
t.Fatal("timeout not enforced, took ", elapsed)
}
firstConn := <-accepted
firstConn.SetReadDeadline(time.Now().Add(2 * time.Second))
_, err = io.Copy(io.Discard, firstConn)
if err != nil {
t.Fatal("expected the client side to close the connection, got ", err)
}
ctx2, cancel2 := context.WithTimeout(context.Background(), time.Second)
defer cancel2()
multiplexer.Exchange(ctx2, message)
select {
case <-accepted:
case <-time.After(time.Second):
t.Fatal("expected a fresh connection for the second query")
}
}
func TestMultiplexerSlowQueryKeepsActiveConn(t *testing.T) {
t.Parallel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer listener.Close()
accepted := make(chan net.Conn, 16)
go func() {
for {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
return
}
accepted <- conn
go func() {
for {
request, readErr := ReadMessage(conn)
if readErr != nil {
return
}
if request.Question[0].Name == "slow.example.com." {
continue
}
response := new(mDNS.Msg)
response.SetReply(request)
WriteMessage(conn, request.Id, response)
}
}()
}
}()
multiplexer := newQueryMultiplexer(queryMultiplexerOptions{
dial: func(ctx context.Context) (net.Conn, error) {
return net.Dial("tcp", listener.Addr().String())
},
write: func(conn net.Conn, message *mDNS.Msg, queryId uint16) error {
return WriteMessage(conn, queryId, message)
},
readNext: func(conn net.Conn) (*mDNS.Msg, error) {
return ReadMessage(conn)
},
})
defer multiplexer.Close()
slowMessage := new(mDNS.Msg)
slowMessage.SetQuestion("slow.example.com.", mDNS.TypeA)
slowCtx, slowCancel := context.WithTimeout(context.Background(), time.Second)
defer slowCancel()
slowDone := make(chan error, 1)
go func() {
_, slowErr := multiplexer.Exchange(slowCtx, slowMessage)
slowDone <- slowErr
}()
select {
case <-accepted:
case <-time.After(time.Second):
t.Fatal("expected a connection for the slow query")
}
fastMessage := new(mDNS.Msg)
fastMessage.SetQuestion("fast.example.com.", mDNS.TypeA)
exchangeFast := func() {
fastCtx, fastCancel := context.WithTimeout(context.Background(), time.Second)
defer fastCancel()
_, fastErr := multiplexer.Exchange(fastCtx, fastMessage)
if fastErr != nil {
t.Fatal("fast query failed: ", fastErr)
}
}
deadline := time.Now().Add(3 * time.Second)
for {
if !time.Now().Before(deadline) {
t.Fatal("slow query did not complete")
}
exchangeFast()
select {
case slowErr := <-slowDone:
if !errors.Is(slowErr, context.DeadlineExceeded) {
t.Fatal("expected deadline exceeded for slow query, got ", slowErr)
}
exchangeFast()
if len(accepted) > 0 {
t.Fatal("slow query timeout must not replace the active connection")
}
return
case <-time.After(50 * time.Millisecond):
}
}
}
+6
View File
@@ -209,3 +209,9 @@ func (t *HTTP3Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS
}
return &responseMessage, nil
}
func (t *HTTP3Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
go func() {
callback(t.Exchange(ctx, message))
}()
}
+6
View File
@@ -145,6 +145,12 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
return nil, err
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
go func() {
callback(t.Exchange(ctx, message))
}()
}
func (t *Transport) exchange(ctx context.Context, message *mDNS.Msg, conn *quic.Conn) (*mDNS.Msg, error) {
stream, err := conn.OpenStreamSync(ctx)
if err != nil {
+34 -23
View File
@@ -15,7 +15,6 @@ import (
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio/deadline"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
@@ -31,8 +30,9 @@ func RegisterTCP(registry *dns.TransportRegistry) {
type TCPTransport struct {
dns.TransportAdapter
dialer N.Dialer
serverAddr M.Socksaddr
dialer N.Dialer
serverAddr M.Socksaddr
multiplexer *queryMultiplexer
}
func NewTCP(ctx context.Context, logger log.ContextLogger, tag string, options option.RemoteDNSServerOptions) (adapter.DNSTransport, error) {
@@ -47,11 +47,31 @@ func NewTCP(ctx context.Context, logger log.ContextLogger, tag string, options o
if !serverAddr.IsValid() {
return nil, E.New("invalid server address: ", serverAddr)
}
return &TCPTransport{
TransportAdapter: dns.NewTransportAdapterWithRemoteOptions(C.DNSTypeTCP, tag, options),
dialer: transportDialer,
return NewTCPRaw(dns.NewTransportAdapterWithRemoteOptions(C.DNSTypeTCP, tag, options), transportDialer, serverAddr), nil
}
func NewTCPRaw(adapter dns.TransportAdapter, dialer N.Dialer, serverAddr M.Socksaddr) *TCPTransport {
t := &TCPTransport{
TransportAdapter: adapter,
dialer: dialer,
serverAddr: serverAddr,
}, nil
}
t.multiplexer = newQueryMultiplexer(queryMultiplexerOptions{
dial: func(ctx context.Context) (net.Conn, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial TCP connection")
}
return conn, nil
},
write: func(conn net.Conn, message *mDNS.Msg, queryId uint16) error {
return WriteMessage(conn, queryId, message)
},
readNext: func(conn net.Conn) (*mDNS.Msg, error) {
return ReadMessage(conn)
},
})
return t
}
func (t *TCPTransport) Start(stage adapter.StartStage) error {
@@ -62,28 +82,19 @@ func (t *TCPTransport) Start(stage adapter.StartStage) error {
}
func (t *TCPTransport) Close() error {
return nil
return t.multiplexer.Close()
}
func (t *TCPTransport) Reset() {
t.multiplexer.Reset()
}
func (t *TCPTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial TCP connection")
}
defer conn.Close()
defer setConnDeadline(ctx, conn, deadline.NeedAdditionalReadDeadline(conn))()
err = WriteMessage(conn, 0, message)
if err != nil {
return nil, E.Cause(err, "write request")
}
response, err := ReadMessage(conn)
if err != nil {
return nil, E.Cause(err, "read response")
}
return response, nil
return t.multiplexer.Exchange(ctx, message)
}
func (t *TCPTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
t.multiplexer.ExchangeAsync(ctx, message, callback)
}
func setConnDeadline(ctx context.Context, conn net.Conn, needClose bool) func() {
+25 -67
View File
@@ -2,6 +2,7 @@ package transport
import (
"context"
"net"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/dialer"
@@ -11,7 +12,6 @@ import (
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/bufio/deadline"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
@@ -22,26 +22,16 @@ import (
var _ adapter.DNSTransport = (*TLSTransport)(nil)
const tlsDNSMaxInflight = 8
func RegisterTLS(registry *dns.TransportRegistry) {
dns.RegisterTransport[option.RemoteTLSDNSServerOptions](registry, C.DNSTypeTLS, NewTLS)
}
type TLSTransport struct {
dns.TransportAdapter
logger logger.ContextLogger
logger logger.ContextLogger
dialer tls.Dialer
serverAddr M.Socksaddr
tlsConfig tls.Config
connections *ConnPool[*tlsDNSConn]
}
type tlsDNSConn struct {
tls.Conn
queryId uint16
needDeadlineClose bool
multiplexer *queryMultiplexer
}
func NewTLS(ctx context.Context, logger log.ContextLogger, tag string, options option.RemoteTLSDNSServerOptions) (adapter.DNSTransport, error) {
@@ -66,23 +56,28 @@ func NewTLS(ctx context.Context, logger log.ContextLogger, tag string, options o
}
func NewTLSRaw(logger logger.ContextLogger, adapter dns.TransportAdapter, dialer N.Dialer, serverAddr M.Socksaddr, tlsConfig tls.Config) *TLSTransport {
return &TLSTransport{
t := &TLSTransport{
TransportAdapter: adapter,
logger: logger,
dialer: tls.NewDialer(dialer, tlsConfig),
serverAddr: serverAddr,
tlsConfig: tlsConfig,
connections: NewConnPool(ConnPoolOptions[*tlsDNSConn]{
Mode: ConnPoolOrdered,
MaxInflight: tlsDNSMaxInflight,
IsAlive: func(conn *tlsDNSConn) bool {
return conn != nil
},
Close: func(conn *tlsDNSConn, _ error) {
conn.Close()
},
}),
}
t.multiplexer = newQueryMultiplexer(queryMultiplexerOptions{
dial: func(ctx context.Context) (net.Conn, error) {
conn, err := t.dialer.DialTLSContext(ctx, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial TLS connection")
}
return conn, nil
},
write: func(conn net.Conn, message *mDNS.Msg, queryId uint16) error {
return WriteMessage(conn, queryId, message)
},
readNext: func(conn net.Conn) (*mDNS.Msg, error) {
return ReadMessage(conn)
},
})
return t
}
func (t *TLSTransport) Start(stage adapter.StartStage) error {
@@ -93,54 +88,17 @@ func (t *TLSTransport) Start(stage adapter.StartStage) error {
}
func (t *TLSTransport) Close() error {
return t.connections.Close()
return t.multiplexer.Close()
}
func (t *TLSTransport) Reset() {
t.connections.Reset()
t.multiplexer.Reset()
}
func (t *TLSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
var lastErr error
for range 2 {
conn, created, err := t.connections.Acquire(ctx, func(ctx context.Context) (*tlsDNSConn, error) {
tlsConn, err := t.dialer.DialTLSContext(ctx, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial TLS connection")
}
return &tlsDNSConn{
Conn: tlsConn,
needDeadlineClose: deadline.NeedAdditionalReadDeadline(tlsConn.NetConn()),
}, nil
})
if err != nil {
return nil, err
}
response, err := t.exchange(ctx, message, conn)
if err == nil {
t.connections.Release(conn, true)
return response, nil
}
lastErr = err
t.logger.DebugContext(ctx, "discarded pooled connection: ", err)
t.connections.Release(conn, false)
if created {
return nil, err
}
}
return nil, lastErr
return t.multiplexer.Exchange(ctx, message)
}
func (t *TLSTransport) exchange(ctx context.Context, message *mDNS.Msg, conn *tlsDNSConn) (*mDNS.Msg, error) {
defer setConnDeadline(ctx, conn, conn.needDeadlineClose)()
conn.queryId++
err := WriteMessage(conn, conn.queryId, message)
if err != nil {
return nil, E.Cause(err, "write request")
}
response, err := ReadMessage(conn)
if err != nil {
return nil, E.Cause(err, "read response")
}
return response, nil
func (t *TLSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
t.multiplexer.ExchangeAsync(ctx, message, callback)
}
+78 -156
View File
@@ -3,7 +3,6 @@ package transport
import (
"context"
"net"
"sync"
"sync/atomic"
"github.com/sagernet/sing-box/adapter"
@@ -12,6 +11,7 @@ import (
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio/deadline"
E "github.com/sagernet/sing/common/exceptions"
@@ -36,17 +36,7 @@ type UDPTransport struct {
serverAddr M.Socksaddr
udpSize atomic.Int32
connection *ConnPool[net.Conn]
callbackAccess sync.RWMutex
queryId uint16
callbacks map[uint16]*udpCallback
}
type udpCallback struct {
access sync.Mutex
response *mDNS.Msg
done chan struct{}
multiplexer *queryMultiplexer
}
func NewUDP(ctx context.Context, logger log.ContextLogger, tag string, options option.RemoteDNSServerOptions) (adapter.DNSTransport, error) {
@@ -70,18 +60,19 @@ func NewUDPRaw(logger logger.ContextLogger, adapter dns.TransportAdapter, dialer
logger: logger,
dialer: dialerInstance,
serverAddr: serverAddr,
callbacks: make(map[uint16]*udpCallback),
connection: NewConnPool(ConnPoolOptions[net.Conn]{
Mode: ConnPoolSingle,
IsAlive: func(conn net.Conn) bool {
return conn != nil
},
Close: func(conn net.Conn, cause error) {
conn.Close()
},
}),
}
t.udpSize.Store(2048)
t.multiplexer = newQueryMultiplexer(queryMultiplexerOptions{
dial: func(ctx context.Context) (net.Conn, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkUDP, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial UDP connection")
}
return conn, nil
},
write: t.writeQuery,
readNext: t.readResponse,
})
return t
}
@@ -93,28 +84,16 @@ func (t *UDPTransport) Start(stage adapter.StartStage) error {
}
func (t *UDPTransport) Close() error {
return t.connection.Close()
return t.multiplexer.Close()
}
func (t *UDPTransport) Reset() {
t.connection.Reset()
}
func (t *UDPTransport) nextAvailableQueryId() (uint16, error) {
start := t.queryId
for {
t.queryId++
if _, exists := t.callbacks[t.queryId]; !exists {
return t.queryId, nil
}
if t.queryId == start {
return 0, E.New("no available query ID")
}
}
t.multiplexer.Reset()
}
func (t *UDPTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
response, err := t.exchange(ctx, message)
t.updateUDPSize(message)
response, err := t.multiplexer.Exchange(ctx, message)
if err != nil {
return nil, err
}
@@ -125,6 +104,67 @@ func (t *UDPTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.M
return response, nil
}
func (t *UDPTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
t.updateUDPSize(message)
t.multiplexer.ExchangeAsync(ctx, message, func(response *mDNS.Msg, err error) {
if err == nil && response.Truncated {
t.logger.InfoContext(ctx, "response truncated, retrying with TCP")
go func() {
callback(t.exchangeTCP(ctx, message))
}()
return
}
callback(response, err)
})
}
func (t *UDPTransport) updateUDPSize(message *mDNS.Msg) {
edns0Opt := message.IsEdns0()
if edns0Opt == nil {
return
}
udpSize := int32(edns0Opt.UDPSize())
for {
current := t.udpSize.Load()
if udpSize <= current {
return
}
if t.udpSize.CompareAndSwap(current, udpSize) {
t.Reset()
return
}
}
}
func (t *UDPTransport) writeQuery(conn net.Conn, message *mDNS.Msg, queryId uint16) error {
buffer := buf.NewSize(1 + message.Len())
defer buffer.Release()
exMessage := *message
exMessage.Compress = true
exMessage.Id = queryId
rawMessage, err := exMessage.PackBuffer(buffer.FreeBytes())
if err != nil {
return err
}
return common.Error(conn.Write(rawMessage))
}
func (t *UDPTransport) readResponse(conn net.Conn) (*mDNS.Msg, error) {
buffer := buf.NewSize(int(t.udpSize.Load()))
defer buffer.Release()
_, err := buffer.ReadOnceFrom(conn)
if err != nil {
return nil, err
}
var message mDNS.Msg
err = message.Unpack(buffer.Bytes())
if err != nil {
t.logger.Debug("discarded malformed UDP response: ", err)
return nil, nil
}
return &message, nil
}
func (t *UDPTransport) exchangeTCP(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, t.serverAddr)
if err != nil {
@@ -142,121 +182,3 @@ func (t *UDPTransport) exchangeTCP(ctx context.Context, message *mDNS.Msg) (*mDN
}
return response, nil
}
func (t *UDPTransport) exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
if edns0Opt := message.IsEdns0(); edns0Opt != nil {
udpSize := int32(edns0Opt.UDPSize())
for {
current := t.udpSize.Load()
if udpSize <= current {
break
}
if t.udpSize.CompareAndSwap(current, udpSize) {
t.Reset()
break
}
}
}
conn, connCtx, created, err := t.connection.AcquireShared(ctx, func(ctx context.Context) (net.Conn, error) {
rawConn, err := t.dialer.DialContext(ctx, N.NetworkUDP, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial UDP connection")
}
return rawConn, nil
})
if err != nil {
return nil, err
}
if created {
go t.recvLoop(conn)
}
callback := &udpCallback{
done: make(chan struct{}),
}
t.callbackAccess.Lock()
queryId, err := t.nextAvailableQueryId()
if err != nil {
t.callbackAccess.Unlock()
t.connection.Release(conn, true)
return nil, err
}
t.callbacks[queryId] = callback
t.callbackAccess.Unlock()
defer func() {
t.callbackAccess.Lock()
delete(t.callbacks, queryId)
t.callbackAccess.Unlock()
}()
buffer := buf.NewSize(1 + message.Len())
defer buffer.Release()
exMessage := *message
exMessage.Compress = true
originalId := message.Id
exMessage.Id = queryId
rawMessage, err := exMessage.PackBuffer(buffer.FreeBytes())
if err != nil {
return nil, err
}
_, err = conn.Write(rawMessage)
if err != nil {
t.connection.Invalidate(conn, err)
return nil, E.Cause(err, "write request")
}
select {
case <-callback.done:
t.connection.Release(conn, true)
callback.response.Id = originalId
return callback.response, nil
case <-connCtx.Done():
return nil, context.Cause(connCtx)
case <-ctx.Done():
t.connection.Release(conn, true)
return nil, ctx.Err()
}
}
func (t *UDPTransport) recvLoop(conn net.Conn) {
for {
buffer := buf.NewSize(int(t.udpSize.Load()))
_, err := buffer.ReadOnceFrom(conn)
if err != nil {
buffer.Release()
t.connection.Invalidate(conn, err)
return
}
var message mDNS.Msg
err = message.Unpack(buffer.Bytes())
buffer.Release()
if err != nil {
t.logger.Debug("discarded malformed UDP response: ", err)
continue
}
t.callbackAccess.RLock()
callback, loaded := t.callbacks[message.Id]
t.callbackAccess.RUnlock()
if !loaded {
continue
}
callback.access.Lock()
select {
case <-callback.done:
default:
callback.response = &message
close(callback.done)
}
callback.access.Unlock()
}
}
+6
View File
@@ -105,6 +105,12 @@ func (p *platformTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*m
}
}
func (p *platformTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
go func() {
callback(p.Exchange(ctx, message))
}()
}
type Func interface {
Invoke() error
}
+38 -40
View File
@@ -40,28 +40,30 @@ func HandleStreamDNSRequest(ctx context.Context, router adapter.DNSRouter, conn
return err
}
metadataInQuery := metadata
go func() error {
response, err := router.Exchange(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{})
router.ExchangeAsync(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{}, func(response *mDNS.Msg, err error) {
if err != nil {
conn.Close()
return err
return
}
responseLength := response.Len()
responseBuffer := buf.NewSize(3 + responseLength)
defer responseBuffer.Release()
responseBuffer.Resize(2, 0)
n, err := response.PackBuffer(responseBuffer.FreeBytes())
if err != nil {
return err
}
responseBuffer.Truncate(len(n))
binary.BigEndian.PutUint16(responseBuffer.ExtendHeader(2), uint16(len(n)))
_, err = conn.Write(responseBuffer.Bytes())
return err
}()
go writeStreamResponse(conn, response)
})
return nil
}
func writeStreamResponse(conn net.Conn, response *mDNS.Msg) {
responseLength := response.Len()
responseBuffer := buf.NewSize(3 + responseLength)
defer responseBuffer.Release()
responseBuffer.Resize(2, 0)
n, err := response.PackBuffer(responseBuffer.FreeBytes())
if err != nil {
return
}
responseBuffer.Truncate(len(n))
binary.BigEndian.PutUint16(responseBuffer.ExtendHeader(2), uint16(len(n)))
conn.Write(responseBuffer.Bytes())
}
func NewDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn N.PacketConn, cachedPackets []*N.PacketBuffer, metadata adapter.InboundContext) error {
metadata.Destination = M.Socksaddr{}
var reader N.PacketReader = conn
@@ -123,24 +125,22 @@ func NewDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn
timeout.Update()
}
metadataInQuery := metadata
go func() error {
response, err := router.Exchange(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{})
router.ExchangeAsync(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{}, func(response *mDNS.Msg, err error) {
if err != nil {
cancel(err)
return err
return
}
timeout.Update()
responseBuffer, err := dns.TruncateDNSMessage(&message, response, 1024)
if err != nil {
cancel(err)
return err
responseBuffer, truncateErr := dns.TruncateDNSMessage(&message, response, 1024)
if truncateErr != nil {
cancel(truncateErr)
return
}
err = conn.WritePacket(responseBuffer, destination)
if err != nil {
cancel(err)
writeErr := conn.WritePacket(responseBuffer, destination)
if writeErr != nil {
cancel(writeErr)
}
return err
}()
})
}
})
group.Cleanup(func() {
@@ -193,24 +193,22 @@ func newDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn
timeout.Update()
}
metadataInQuery := metadata
go func() error {
response, err := router.Exchange(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{})
router.ExchangeAsync(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{}, func(response *mDNS.Msg, err error) {
if err != nil {
cancel(err)
return err
return
}
timeout.Update()
responseBuffer, err := dns.TruncateDNSMessage(&message, response, 1024)
if err != nil {
cancel(err)
return err
responseBuffer, truncateErr := dns.TruncateDNSMessage(&message, response, 1024)
if truncateErr != nil {
cancel(truncateErr)
return
}
err = conn.WritePacket(responseBuffer, destination)
if err != nil {
cancel(err)
writeErr := conn.WritePacket(responseBuffer, destination)
if writeErr != nil {
cancel(writeErr)
}
return err
}()
})
}
})
group.Cleanup(func() {
+63 -53
View File
@@ -4,7 +4,6 @@ package tailscale
import (
"context"
"errors"
"net"
"net/http"
"net/netip"
@@ -276,48 +275,64 @@ func (t *DNSTransport) PreferredDomain(domain string) bool {
}
func (t *DNSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
done := make(chan struct{})
var (
response *mDNS.Msg
err error
)
t.ExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
response = callbackResponse
err = callbackErr
close(done)
})
<-done
return response, err
}
func (t *DNSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
if len(message.Question) != 1 {
return nil, os.ErrInvalid
callback(nil, os.ErrInvalid)
return
}
if t.acceptSearchDomain && mDNS.CountLabel(message.Question[0].Name) == 1 {
return t.exchangeWithSearchDomains(ctx, message)
t.exchangeWithSearchDomains(ctx, message, callback)
return
}
t.access.RLock()
acceptDefaultResolvers := t.acceptDefaultResolvers
t.access.RUnlock()
return t.exchangeOnce(ctx, message, acceptDefaultResolvers)
t.exchangeOnce(ctx, message, acceptDefaultResolvers, callback)
}
func (t *DNSTransport) exchangeWithSearchDomains(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
func (t *DNSTransport) exchangeWithSearchDomains(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
t.access.RLock()
searchDomains := t.searchDomains
t.access.RUnlock()
if len(searchDomains) == 0 {
callback(nil, dns.RcodeNameError)
return
}
originalQuestion := message.Question[0]
singleLabel := strings.TrimSuffix(originalQuestion.Name, ".")
var lastErr error
domainExchangers := make([]transport.AsyncExchanger, 0, len(searchDomains))
for _, searchDomain := range searchDomains {
expandedName := singleLabel + "." + searchDomain
question := originalQuestion
question.Name = expandedName
rewritten := *message
rewritten.Question = []mDNS.Question{question}
response, err := t.exchangeOnce(ctx, &rewritten, false)
if err == nil {
if response.Rcode == mDNS.RcodeNameError {
continue
}
restoreOriginalQuestion(response, expandedName, originalQuestion)
return response, nil
}
if errors.Is(err, dns.RcodeNameError) {
continue
}
lastErr = err
domainExchangers = append(domainExchangers, func(exchangeCtx context.Context, exchangeCallback func(response *mDNS.Msg, err error)) {
question := originalQuestion
question.Name = expandedName
rewritten := *message
rewritten.Question = []mDNS.Question{question}
t.exchangeOnce(exchangeCtx, &rewritten, false, func(response *mDNS.Msg, err error) {
if err == nil {
restoreOriginalQuestion(response, expandedName, originalQuestion)
}
exchangeCallback(response, err)
})
})
}
if lastErr != nil {
return nil, lastErr
}
return nil, dns.RcodeNameError
transport.ExchangeSequential(ctx, domainExchangers, func(response *mDNS.Msg, err error) bool {
return err == nil && response.Rcode != mDNS.RcodeNameError
}, callback)
}
// RFC 1035 §4.1.1 requires the response Question to match the request byte-for-byte,
@@ -331,7 +346,7 @@ func restoreOriginalQuestion(response *mDNS.Msg, expandedName string, originalQu
}
}
func (t *DNSTransport) exchangeOnce(ctx context.Context, message *mDNS.Msg, allowDefaultResolvers bool) (*mDNS.Msg, error) {
func (t *DNSTransport) exchangeOnce(ctx context.Context, message *mDNS.Msg, allowDefaultResolvers bool, callback func(response *mDNS.Msg, err error)) {
question := message.Question[0]
t.access.RLock()
@@ -348,58 +363,53 @@ func (t *DNSTransport) exchangeOnce(ctx context.Context, message *mDNS.Msg, allo
return addr.Is4()
})
if len(addresses4) > 0 {
return dns.FixedResponse(message.Id, question, addresses4, C.DefaultDNSTTL), nil
callback(dns.FixedResponse(message.Id, question, addresses4, C.DefaultDNSTTL), nil)
return
}
case mDNS.TypeAAAA:
addresses6 := common.Filter(addresses, func(addr netip.Addr) bool {
return addr.Is6()
})
if len(addresses6) > 0 {
return dns.FixedResponse(message.Id, question, addresses6, C.DefaultDNSTTL), nil
callback(dns.FixedResponse(message.Id, question, addresses6, C.DefaultDNSTTL), nil)
return
}
}
}
for domainSuffix, transports := range routes {
if mDNS.IsSubDomain(domainSuffix, question.Name) {
if len(transports) == 0 {
return &mDNS.Msg{
callback(&mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Id: message.Id,
Rcode: mDNS.RcodeNameError,
Response: true,
},
Question: []mDNS.Question{question},
}, nil
}, nil)
return
}
var lastErr error
for _, dnsTransport := range transports {
response, err := dnsTransport.Exchange(ctx, message)
if err != nil {
lastErr = err
continue
}
return response, nil
}
return nil, lastErr
transport.ExchangeSequential(ctx, resolverExchangers(transports, message), nil, callback)
return
}
}
if allowDefaultResolvers {
if len(defaultResolvers) > 0 {
var lastErr error
for _, resolver := range defaultResolvers {
response, err := resolver.Exchange(ctx, message)
if err != nil {
lastErr = err
continue
}
return response, nil
}
return nil, lastErr
transport.ExchangeSequential(ctx, resolverExchangers(defaultResolvers, message), nil, callback)
} else {
return nil, E.New("missing default resolvers")
callback(nil, E.New("missing default resolvers"))
}
return
}
return nil, dns.RcodeNameError
callback(nil, dns.RcodeNameError)
}
func resolverExchangers(resolvers []adapter.DNSTransport, message *mDNS.Msg) []transport.AsyncExchanger {
return common.Map(resolvers, func(resolver adapter.DNSTransport) transport.AsyncExchanger {
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
resolver.ExchangeAsync(ctx, message, callback)
}
})
}
func (t *DNSTransport) collectResolversLocked() []adapter.DNSTransport {
+6 -8
View File
@@ -50,19 +50,17 @@ func (r *Router) HijackDNSPacket(ctx context.Context, payload []byte, writer N.P
}
destination := metadata.Destination
metadata.Destination = M.Socksaddr{}
go func() {
exchangeErr := r.exchangeDNSPacket(ctx, &message, writer, metadata, destination)
r.dns.ExchangeAsync(adapter.WithContext(ctx, &metadata), &message, adapter.DNSQueryOptions{}, func(response *mDNS.Msg, exchangeErr error) {
if exchangeErr == nil {
exchangeErr = r.writeDNSPacketResponse(&message, response, writer, destination)
}
if exchangeErr != nil && !R.IsRejected(exchangeErr) && !E.IsClosedOrCanceled(exchangeErr) {
r.logger.ErrorContext(ctx, E.Cause(exchangeErr, "process DNS packet"))
}
}()
})
}
func (r *Router) exchangeDNSPacket(ctx context.Context, message *mDNS.Msg, writer N.PacketWriter, metadata adapter.InboundContext, destination M.Socksaddr) error {
response, err := r.dns.Exchange(adapter.WithContext(ctx, &metadata), message, adapter.DNSQueryOptions{})
if err != nil {
return err
}
func (r *Router) writeDNSPacketResponse(message *mDNS.Msg, response *mDNS.Msg, writer N.PacketWriter, destination M.Socksaddr) error {
responseBuffer, err := dns.TruncateDNSMessage(message, response, 1024)
if err != nil {
return err
+54 -74
View File
@@ -210,6 +210,21 @@ func (t *Transport) PreferredDomain(domain string) bool {
}
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
done := make(chan struct{})
var (
response *mDNS.Msg
err error
)
t.ExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
response = callbackResponse
err = callbackErr
close(done)
})
<-done
return response, err
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
question := message.Question[0]
var selectedLink *TransportLink
t.service.linkAccess.RLock()
@@ -233,93 +248,58 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
}
t.service.linkAccess.RUnlock()
if selectedLink == nil {
return dns.FixedResponseStatus(message, mDNS.RcodeNameError), nil
callback(dns.FixedResponseStatus(message, mDNS.RcodeNameError), nil)
return
}
t.linkAccess.RLock()
servers := t.linkServers[selectedLink]
t.linkAccess.RUnlock()
if len(servers.Servers) == 0 {
return dns.FixedResponseStatus(message, mDNS.RcodeNameError), nil
if servers == nil || len(servers.Servers) == 0 {
callback(dns.FixedResponseStatus(message, mDNS.RcodeNameError), nil)
return
}
names := servers.Link.nameList(t.ndots, question.Name)
if len(names) == 0 {
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 {
return t.exchangeParallel(ctx, servers, message)
transport.ExchangeRace(ctx, nameExchangers, callback)
} else {
return t.exchangeSingleRequest(ctx, servers, message)
transport.ExchangeSequential(ctx, nameExchangers, nil, callback)
}
}
func (t *Transport) exchangeSingleRequest(ctx context.Context, servers *LinkServers, message *mDNS.Msg) (*mDNS.Msg, error) {
var lastErr error
for _, fqdn := range servers.Link.nameList(t.ndots, message.Question[0].Name) {
response, err := t.tryOneName(ctx, servers, message, fqdn)
if err != nil {
lastErr = err
continue
}
return response, nil
}
return nil, lastErr
}
func (t *Transport) tryOneName(ctx context.Context, servers *LinkServers, message *mDNS.Msg, fqdn string) (*mDNS.Msg, error) {
func (t *Transport) newNameExchanger(servers *LinkServers, message *mDNS.Msg, fqdn string) transport.AsyncExchanger {
serverOffset := servers.ServerOffset(t.rotate)
sLen := uint32(len(servers.Servers))
var lastErr error
serverCount := uint32(len(servers.Servers))
attemptExchangers := make([]transport.AsyncExchanger, 0, t.attempts*int(serverCount))
for i := 0; i < t.attempts; i++ {
for j := range sLen {
server := servers.Servers[(serverOffset+j)%sLen]
question := message.Question[0]
question.Name = fqdn
exchangeMessage := *message
exchangeMessage.Question = []mDNS.Question{question}
exchangeCtx, cancel := context.WithTimeout(ctx, t.timeout)
response, err := server.Exchange(exchangeCtx, &exchangeMessage)
cancel()
for j := range serverCount {
server := servers.Servers[(serverOffset+j)%serverCount]
attemptExchangers = append(attemptExchangers, func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
question := message.Question[0]
question.Name = fqdn
exchangeMessage := *message
exchangeMessage.Question = []mDNS.Question{question}
exchangeCtx, cancel := context.WithTimeout(ctx, t.timeout)
server.ExchangeAsync(exchangeCtx, &exchangeMessage, func(response *mDNS.Msg, err error) {
cancel()
callback(response, err)
})
})
}
}
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
transport.ExchangeSequential(ctx, attemptExchangers, nil, func(response *mDNS.Msg, err error) {
if err != nil {
lastErr = err
continue
err = E.Cause(err, fqdn)
}
return response, nil
}
}
return nil, E.Cause(lastErr, fqdn)
}
func (t *Transport) exchangeParallel(ctx context.Context, servers *LinkServers, message *mDNS.Msg) (*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, message, fqdn)
select {
case results <- queryResult{response, err}:
case <-returned:
}
}
queryCtx, queryCancel := context.WithCancel(ctx)
defer queryCancel()
var nameCount int
for _, fqdn := range servers.Link.nameList(t.ndots, message.Question[0].Name) {
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...)
}
}
callback(response, err)
})
}
}