From 5d424ea2acd5f61acdcca588a1525fc49d1e36a5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sun, 19 Jul 2026 16:37:16 +0800 Subject: [PATCH] refactor: Async DNS --- adapter/dns.go | 3 + dns/client.go | 149 +++++-- dns/router.go | 364 ++++++++++------- dns/transport/dhcp/dhcp.go | 142 +++++-- dns/transport/dhcp/dhcp_shared.go | 153 ++----- dns/transport/exchange_strategy.go | 153 +++++++ dns/transport/fakeip/fakeip.go | 4 + dns/transport/hosts/hosts.go | 4 + dns/transport/https.go | 6 + dns/transport/local/local.go | 75 +++- dns/transport/local/local_darwin.go | 418 +++++++++++++++----- dns/transport/local/local_darwin_test.go | 68 +++- dns/transport/local/local_other.go | 10 +- dns/transport/local/local_resolved.go | 1 + dns/transport/local/local_resolved_linux.go | 45 +++ dns/transport/local/local_shared.go | 239 ++++------- dns/transport/mdns/mdns.go | 6 + dns/transport/multiplexer.go | 209 ++++++++++ dns/transport/multiplexer_test.go | 164 ++++++++ dns/transport/quic/http3.go | 6 + dns/transport/quic/quic.go | 6 + dns/transport/tcp.go | 57 +-- dns/transport/tls.go | 92 ++--- dns/transport/udp.go | 234 ++++------- experimental/libbox/dns.go | 6 + protocol/dns/handle.go | 78 ++-- protocol/tailscale/dns_transport.go | 116 +++--- route/dns.go | 14 +- service/resolved/transport.go | 128 +++--- 29 files changed, 1956 insertions(+), 994 deletions(-) create mode 100644 dns/transport/exchange_strategy.go create mode 100644 dns/transport/multiplexer.go create mode 100644 dns/transport/multiplexer_test.go diff --git a/adapter/dns.go b/adapter/dns.go index a613de0c..b399d6b4 100644 --- a/adapter/dns.go +++ b/adapter/dns.go @@ -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 { diff --git a/dns/client.go b/dns/client.go index 03c2a0e4..9b314bbc 100644 --- a/dns/client.go +++ b/dns/client.go @@ -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) } diff --git a/dns/router.go b/dns/router.go index b1217b64..13d4abef 100644 --- a/dns/router.go +++ b/dns/router.go @@ -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 ")) - } - } - 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 ")) + } + } + 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 { diff --git a/dns/transport/dhcp/dhcp.go b/dns/transport/dhcp/dhcp.go index 4b97c723..3abc7cf5 100644 --- a/dns/transport/dhcp/dhcp.go +++ b/dns/transport/dhcp/dhcp.go @@ -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 } diff --git a/dns/transport/dhcp/dhcp_shared.go b/dns/transport/dhcp/dhcp_shared.go index 32621ef7..3123d8f9 100644 --- a/dns/transport/dhcp/dhcp_shared.go +++ b/dns/transport/dhcp/dhcp_shared.go @@ -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 { diff --git a/dns/transport/exchange_strategy.go b/dns/transport/exchange_strategy.go new file mode 100644 index 00000000..92faff1c --- /dev/null +++ b/dns/transport/exchange_strategy.go @@ -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 +} diff --git a/dns/transport/fakeip/fakeip.go b/dns/transport/fakeip/fakeip.go index 9aa41e58..75db1cf8 100644 --- a/dns/transport/fakeip/fakeip.go +++ b/dns/transport/fakeip/fakeip.go @@ -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 } diff --git a/dns/transport/hosts/hosts.go b/dns/transport/hosts/hosts.go index 4db7988c..2ba47063 100644 --- a/dns/transport/hosts/hosts.go +++ b/dns/transport/hosts/hosts.go @@ -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)) +} diff --git a/dns/transport/https.go b/dns/transport/https.go index db89799c..5baa782f 100644 --- a/dns/transport/https.go +++ b/dns/transport/https.go @@ -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 diff --git a/dns/transport/local/local.go b/dns/transport/local/local.go index 34897e12..d0dca79b 100644 --- a/dns/transport/local/local.go +++ b/dns/transport/local/local.go @@ -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) } diff --git a/dns/transport/local/local_darwin.go b/dns/transport/local/local_darwin.go index 47aa51fc..e033ba12 100644 --- a/dns/transport/local/local_darwin.go +++ b/dns/transport/local/local_darwin.go @@ -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 diff --git a/dns/transport/local/local_darwin_test.go b/dns/transport/local/local_darwin_test.go index 43aa59e3..7db18c8d 100644 --- a/dns/transport/local/local_darwin_test.go +++ b/dns/transport/local/local_darwin_test.go @@ -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) + } +} diff --git a/dns/transport/local/local_other.go b/dns/transport/local/local_other.go index 9bb3d777..832f1ded 100644 --- a/dns/transport/local/local_other.go +++ b/dns/transport/local/local_other.go @@ -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) } diff --git a/dns/transport/local/local_resolved.go b/dns/transport/local/local_resolved.go index e0128d6d..451ee365 100644 --- a/dns/transport/local/local_resolved.go +++ b/dns/transport/local/local_resolved.go @@ -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)) } diff --git a/dns/transport/local/local_resolved_linux.go b/dns/transport/local/local_resolved_linux.go index fc3ca2b7..b55213b4 100644 --- a/dns/transport/local/local_resolved_linux.go +++ b/dns/transport/local/local_resolved_linux.go @@ -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) diff --git a/dns/transport/local/local_shared.go b/dns/transport/local/local_shared.go index 07040911..cd011d78 100644 --- a/dns/transport/local/local_shared.go +++ b/dns/transport/local/local_shared.go @@ -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) } diff --git a/dns/transport/mdns/mdns.go b/dns/transport/mdns/mdns.go index 2db3390d..76851ebc 100644 --- a/dns/transport/mdns/mdns.go +++ b/dns/transport/mdns/mdns.go @@ -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 diff --git a/dns/transport/multiplexer.go b/dns/transport/multiplexer.go new file mode 100644 index 00000000..51f3c325 --- /dev/null +++ b/dns/transport/multiplexer.go @@ -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) + } +} diff --git a/dns/transport/multiplexer_test.go b/dns/transport/multiplexer_test.go new file mode 100644 index 00000000..413a82e7 --- /dev/null +++ b/dns/transport/multiplexer_test.go @@ -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): + } + } +} diff --git a/dns/transport/quic/http3.go b/dns/transport/quic/http3.go index 0a93e515..3a6c3fc1 100644 --- a/dns/transport/quic/http3.go +++ b/dns/transport/quic/http3.go @@ -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)) + }() +} diff --git a/dns/transport/quic/quic.go b/dns/transport/quic/quic.go index 8d45bd82..dc2e22fc 100644 --- a/dns/transport/quic/quic.go +++ b/dns/transport/quic/quic.go @@ -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 { diff --git a/dns/transport/tcp.go b/dns/transport/tcp.go index f8249437..45f3cda7 100644 --- a/dns/transport/tcp.go +++ b/dns/transport/tcp.go @@ -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() { diff --git a/dns/transport/tls.go b/dns/transport/tls.go index fdb48563..9ec6bf40 100644 --- a/dns/transport/tls.go +++ b/dns/transport/tls.go @@ -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) } diff --git a/dns/transport/udp.go b/dns/transport/udp.go index 7203b5ad..4d5370c4 100644 --- a/dns/transport/udp.go +++ b/dns/transport/udp.go @@ -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() - } -} diff --git a/experimental/libbox/dns.go b/experimental/libbox/dns.go index 75472188..7a3b8a47 100644 --- a/experimental/libbox/dns.go +++ b/experimental/libbox/dns.go @@ -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 } diff --git a/protocol/dns/handle.go b/protocol/dns/handle.go index d7d89ca8..72197140 100644 --- a/protocol/dns/handle.go +++ b/protocol/dns/handle.go @@ -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() { diff --git a/protocol/tailscale/dns_transport.go b/protocol/tailscale/dns_transport.go index b3119b55..42452227 100644 --- a/protocol/tailscale/dns_transport.go +++ b/protocol/tailscale/dns_transport.go @@ -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 { diff --git a/route/dns.go b/route/dns.go index 94f73b51..58707152 100644 --- a/route/dns.go +++ b/route/dns.go @@ -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 diff --git a/service/resolved/transport.go b/service/resolved/transport.go index 95229763..b00f31a3 100644 --- a/service/resolved/transport.go +++ b/service/resolved/transport.go @@ -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) + }) } }