mirror of
https://github.com/shtorm-7/sing-box-extended.git
synced 2026-09-15 21:00:27 +00:00
refactor: Async DNS
This commit is contained in:
+125
-24
@@ -152,19 +152,45 @@ func normalizeTTL(response *dns.Msg, timeToLive uint32) {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool) (*dns.Msg, error) {
|
||||
type exchangeStatus int
|
||||
|
||||
const (
|
||||
exchangeReady exchangeStatus = iota
|
||||
exchangeDone
|
||||
exchangeWait
|
||||
)
|
||||
|
||||
type exchangeOperation struct {
|
||||
ctx context.Context
|
||||
message *dns.Msg
|
||||
question dns.Question
|
||||
messageId uint16
|
||||
options adapter.DNSQueryOptions
|
||||
responseChecker func(response *dns.Msg) bool
|
||||
disableCache bool
|
||||
releaseCond func()
|
||||
}
|
||||
|
||||
func (o *exchangeOperation) release() {
|
||||
if o.releaseCond != nil {
|
||||
o.releaseCond()
|
||||
o.releaseCond = nil
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) beginExchange(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool, allowWait bool) (*exchangeOperation, *dns.Msg, exchangeStatus, error) {
|
||||
if len(message.Question) == 0 {
|
||||
if c.logger != nil {
|
||||
c.logger.WarnContext(ctx, "bad question size: ", len(message.Question))
|
||||
}
|
||||
return FixedResponseStatus(message, dns.RcodeFormatError), nil
|
||||
return nil, FixedResponseStatus(message, dns.RcodeFormatError), exchangeDone, nil
|
||||
}
|
||||
question := message.Question[0]
|
||||
if question.Qtype == dns.TypeA && options.Strategy == C.DomainStrategyIPv6Only || question.Qtype == dns.TypeAAAA && options.Strategy == C.DomainStrategyIPv4Only {
|
||||
if c.logger != nil {
|
||||
c.logger.DebugContext(ctx, "strategy rejected")
|
||||
}
|
||||
return FixedResponseStatus(message, dns.RcodeSuccess), nil
|
||||
return nil, FixedResponseStatus(message, dns.RcodeSuccess), exchangeDone, nil
|
||||
}
|
||||
message = c.prepareExchangeMessage(message, options)
|
||||
|
||||
@@ -177,20 +203,31 @@ func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, m
|
||||
len(message.Extra[0].(*dns.OPT).Option) == 0) &&
|
||||
!options.ClientSubnet.IsValid()
|
||||
disableCache := !isSimpleRequest || c.disableCache || options.DisableCache
|
||||
operation := &exchangeOperation{
|
||||
message: message,
|
||||
question: question,
|
||||
messageId: message.Id,
|
||||
options: options,
|
||||
responseChecker: responseChecker,
|
||||
disableCache: disableCache,
|
||||
}
|
||||
if !disableCache {
|
||||
cacheKey := dnsCacheKey{Question: question, transportTag: transport.Tag()}
|
||||
cond, loaded := c.cacheLock.LoadOrStore(cacheKey, make(chan struct{}))
|
||||
if loaded {
|
||||
if !allowWait {
|
||||
return nil, nil, exchangeWait, nil
|
||||
}
|
||||
select {
|
||||
case <-cond:
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
return nil, nil, exchangeDone, ctx.Err()
|
||||
}
|
||||
} else {
|
||||
defer func() {
|
||||
operation.releaseCond = func() {
|
||||
c.cacheLock.Delete(cacheKey)
|
||||
close(cond)
|
||||
}()
|
||||
}
|
||||
}
|
||||
response, ttl, isStale := c.loadResponse(question, transport)
|
||||
if response != nil {
|
||||
@@ -198,38 +235,43 @@ func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, m
|
||||
c.backgroundRefreshDNS(transport, question, message.Copy(), options, responseChecker)
|
||||
logOptimisticResponse(c.logger, ctx, response)
|
||||
response.Id = message.Id
|
||||
return response, nil
|
||||
operation.release()
|
||||
return nil, response, exchangeDone, nil
|
||||
} else if !isStale {
|
||||
logCachedResponse(c.logger, ctx, response, ttl)
|
||||
response.Id = message.Id
|
||||
return response, nil
|
||||
operation.release()
|
||||
return nil, response, exchangeDone, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
messageId := message.Id
|
||||
contextTransport, clientSubnetLoaded := transportTagFromContext(ctx)
|
||||
if clientSubnetLoaded && transport.Tag() == contextTransport {
|
||||
return nil, E.New("DNS query loopback in transport[", contextTransport, "]")
|
||||
contextTransport, transportTagLoaded := transportTagFromContext(ctx)
|
||||
if transportTagLoaded && transport.Tag() == contextTransport {
|
||||
operation.release()
|
||||
return nil, nil, exchangeDone, E.New("DNS query loopback in transport[", contextTransport, "]")
|
||||
}
|
||||
ctx = contextWithTransportTag(ctx, transport.Tag())
|
||||
operation.ctx = contextWithTransportTag(ctx, transport.Tag())
|
||||
if !disableCache && responseChecker != nil && c.rdrc != nil {
|
||||
rejected := c.rdrc.LoadRDRC(transport.Tag(), question.Name, question.Qtype)
|
||||
if rejected {
|
||||
return nil, ErrResponseRejectedCached
|
||||
operation.release()
|
||||
return nil, nil, exchangeDone, ErrResponseRejectedCached
|
||||
}
|
||||
}
|
||||
response, err := c.exchangeToTransport(ctx, transport, message, options.Timeout)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
disableCache = disableCache || (response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError)
|
||||
if responseChecker != nil {
|
||||
return operation, nil, exchangeReady, nil
|
||||
}
|
||||
|
||||
func (c *Client) finishExchange(transport adapter.DNSTransport, operation *exchangeOperation, response *dns.Msg) (*dns.Msg, error) {
|
||||
ctx := operation.ctx
|
||||
question := operation.question
|
||||
disableCache := operation.disableCache || (response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError)
|
||||
if operation.responseChecker != nil {
|
||||
var rejected bool
|
||||
if response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError {
|
||||
rejected = true
|
||||
} else {
|
||||
rejected = !responseChecker(response)
|
||||
rejected = !operation.responseChecker(response)
|
||||
}
|
||||
if rejected {
|
||||
if !disableCache && c.rdrc != nil {
|
||||
@@ -239,12 +281,12 @@ func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, m
|
||||
return response, ErrResponseRejected
|
||||
}
|
||||
}
|
||||
timeToLive := applyResponseOptions(question, response, options)
|
||||
timeToLive := applyResponseOptions(question, response, operation.options)
|
||||
if !disableCache {
|
||||
c.storeCache(transport, question, response, timeToLive)
|
||||
}
|
||||
response.Id = messageId
|
||||
requestEDNSOpt := message.IsEdns0()
|
||||
response.Id = operation.messageId
|
||||
requestEDNSOpt := operation.message.IsEdns0()
|
||||
responseEDNSOpt := response.IsEdns0()
|
||||
if responseEDNSOpt != nil && (requestEDNSOpt == nil || requestEDNSOpt.Version() < responseEDNSOpt.Version()) {
|
||||
response.Extra = common.Filter(response.Extra, func(it dns.RR) bool {
|
||||
@@ -258,6 +300,44 @@ func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, m
|
||||
return response, nil
|
||||
}
|
||||
|
||||
func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool) (*dns.Msg, error) {
|
||||
operation, earlyResponse, status, err := c.beginExchange(ctx, transport, message, options, responseChecker, true)
|
||||
if status != exchangeReady {
|
||||
return earlyResponse, err
|
||||
}
|
||||
defer operation.release()
|
||||
response, err := c.exchangeToTransport(operation.ctx, transport, operation.message, options.Timeout)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return c.finishExchange(transport, operation, response)
|
||||
}
|
||||
|
||||
func (c *Client) ExchangeAsync(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool, callback func(response *dns.Msg, err error)) {
|
||||
operation, earlyResponse, status, err := c.beginExchange(ctx, transport, message, options, responseChecker, false)
|
||||
switch status {
|
||||
case exchangeDone:
|
||||
callback(earlyResponse, err)
|
||||
return
|
||||
case exchangeWait:
|
||||
go func() {
|
||||
callback(c.Exchange(ctx, transport, message, options, responseChecker))
|
||||
}()
|
||||
return
|
||||
}
|
||||
finish := func(response *dns.Msg, exchangeErr error) {
|
||||
if exchangeErr != nil {
|
||||
operation.release()
|
||||
callback(nil, exchangeErr)
|
||||
return
|
||||
}
|
||||
finishedResponse, finishErr := c.finishExchange(transport, operation, response)
|
||||
operation.release()
|
||||
callback(finishedResponse, finishErr)
|
||||
}
|
||||
c.exchangeToTransportAsync(operation.ctx, transport, operation.message, options.Timeout, finish)
|
||||
}
|
||||
|
||||
func (c *Client) Lookup(ctx context.Context, transport adapter.DNSTransport, domain string, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool) ([]netip.Addr, error) {
|
||||
domain = FqdnToDomain(domain)
|
||||
dnsName := dns.Fqdn(domain)
|
||||
@@ -562,6 +642,27 @@ func (c *Client) exchangeToTransport(ctx context.Context, transport adapter.DNST
|
||||
return nil, err
|
||||
}
|
||||
|
||||
func (c *Client) exchangeToTransportAsync(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, timeout time.Duration, callback func(response *dns.Msg, err error)) {
|
||||
if timeout == 0 {
|
||||
timeout = c.timeout
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(ctx, timeout)
|
||||
transport.ExchangeAsync(ctx, message, func(response *dns.Msg, err error) {
|
||||
cancel()
|
||||
if err == nil {
|
||||
stripDNSPadding(response)
|
||||
callback(response, nil)
|
||||
return
|
||||
}
|
||||
var rcodeError RcodeError
|
||||
if errors.As(err, &rcodeError) {
|
||||
callback(FixedResponseStatus(message, int(rcodeError)), nil)
|
||||
return
|
||||
}
|
||||
callback(nil, err)
|
||||
})
|
||||
}
|
||||
|
||||
func MessageToAddresses(response *dns.Msg) []netip.Addr {
|
||||
return adapter.DNSResponseAddresses(response)
|
||||
}
|
||||
|
||||
+230
-134
@@ -410,60 +410,66 @@ type exchangeWithRulesResult struct {
|
||||
|
||||
const dnsRespondMissingResponseMessage = "respond action requires an evaluated response from a preceding evaluate action"
|
||||
|
||||
func (r *Router) exchangeWithRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, options adapter.DNSQueryOptions, allowFakeIP bool) exchangeWithRulesResult {
|
||||
type dnsRuleWalkState struct {
|
||||
ruleIndex int
|
||||
effectiveOptions adapter.DNSQueryOptions
|
||||
evaluatedResponse *mDNS.Msg
|
||||
evaluatedTransport adapter.DNSTransport
|
||||
}
|
||||
|
||||
type dnsPendingExchange struct {
|
||||
transport adapter.DNSTransport
|
||||
options adapter.DNSQueryOptions
|
||||
evaluate bool
|
||||
}
|
||||
|
||||
func (r *Router) finalizeExchangeOptions(options adapter.DNSQueryOptions) adapter.DNSQueryOptions {
|
||||
if options.Strategy == C.DomainStrategyAsIS {
|
||||
options.Strategy = r.defaultDomainStrategy
|
||||
}
|
||||
return options
|
||||
}
|
||||
|
||||
func (r *Router) walkDNSRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, state *dnsRuleWalkState, allowFakeIP bool) (exchangeWithRulesResult, *dnsPendingExchange) {
|
||||
metadata := adapter.ContextFrom(ctx)
|
||||
if metadata == nil {
|
||||
panic("no context")
|
||||
}
|
||||
effectiveOptions := options
|
||||
var evaluatedResponse *mDNS.Msg
|
||||
var evaluatedTransport adapter.DNSTransport
|
||||
for currentRuleIndex, currentRule := range rules {
|
||||
for ; state.ruleIndex < len(rules); state.ruleIndex++ {
|
||||
currentRule := rules[state.ruleIndex]
|
||||
metadata.ResetRuleCache()
|
||||
metadata.DNSResponse = evaluatedResponse
|
||||
metadata.DNSResponse = state.evaluatedResponse
|
||||
metadata.DestinationAddressMatchFromResponse = false
|
||||
if !currentRule.Match(metadata) {
|
||||
continue
|
||||
}
|
||||
r.logRuleMatch(ctx, currentRuleIndex, currentRule)
|
||||
r.logRuleMatch(ctx, state.ruleIndex, currentRule)
|
||||
switch action := currentRule.Action().(type) {
|
||||
case *R.RuleActionDNSRouteOptions:
|
||||
r.applyDNSRouteOptions(&effectiveOptions, *action)
|
||||
r.applyDNSRouteOptions(&state.effectiveOptions, *action)
|
||||
case *R.RuleActionEvaluate:
|
||||
queryOptions := effectiveOptions
|
||||
queryOptions := state.effectiveOptions
|
||||
transport, loaded := r.transport.Transport(action.Server)
|
||||
if !loaded {
|
||||
r.logger.ErrorContext(ctx, "transport not found: ", action.Server)
|
||||
evaluatedResponse = nil
|
||||
evaluatedTransport = nil
|
||||
state.evaluatedResponse = nil
|
||||
state.evaluatedTransport = nil
|
||||
continue
|
||||
}
|
||||
r.applyDNSRouteOptions(&queryOptions, action.RuleActionDNSRouteOptions)
|
||||
exchangeOptions := queryOptions
|
||||
if exchangeOptions.Strategy == C.DomainStrategyAsIS {
|
||||
exchangeOptions.Strategy = r.defaultDomainStrategy
|
||||
}
|
||||
response, err := r.client.Exchange(adapter.OverrideContext(ctx), transport, message, exchangeOptions, nil)
|
||||
if err != nil {
|
||||
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for ", FormatQuestion(message.Question[0].String())))
|
||||
evaluatedResponse = nil
|
||||
evaluatedTransport = nil
|
||||
continue
|
||||
}
|
||||
evaluatedResponse = response
|
||||
evaluatedTransport = transport
|
||||
return exchangeWithRulesResult{}, &dnsPendingExchange{transport: transport, options: queryOptions, evaluate: true}
|
||||
case *R.RuleActionRespond:
|
||||
if evaluatedResponse == nil {
|
||||
if state.evaluatedResponse == nil {
|
||||
return exchangeWithRulesResult{
|
||||
err: E.New(dnsRespondMissingResponseMessage),
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
return exchangeWithRulesResult{
|
||||
response: evaluatedResponse,
|
||||
transport: evaluatedTransport,
|
||||
}
|
||||
response: state.evaluatedResponse,
|
||||
transport: state.evaluatedTransport,
|
||||
}, nil
|
||||
case *R.RuleActionDNSRoute:
|
||||
queryOptions := effectiveOptions
|
||||
queryOptions := state.effectiveOptions
|
||||
transport, status := r.resolveDNSRoute(action.Server, action.RuleActionDNSRouteOptions, allowFakeIP, &queryOptions)
|
||||
switch status {
|
||||
case dnsRouteStatusMissing:
|
||||
@@ -472,16 +478,7 @@ func (r *Router) exchangeWithRules(ctx context.Context, rules []adapter.DNSRule,
|
||||
case dnsRouteStatusSkipped:
|
||||
continue
|
||||
}
|
||||
exchangeOptions := queryOptions
|
||||
if exchangeOptions.Strategy == C.DomainStrategyAsIS {
|
||||
exchangeOptions.Strategy = r.defaultDomainStrategy
|
||||
}
|
||||
response, err := r.client.Exchange(adapter.OverrideContext(ctx), transport, message, exchangeOptions, nil)
|
||||
return exchangeWithRulesResult{
|
||||
response: response,
|
||||
transport: transport,
|
||||
err: err,
|
||||
}
|
||||
return exchangeWithRulesResult{}, &dnsPendingExchange{transport: transport, options: queryOptions}
|
||||
case *R.RuleActionReject:
|
||||
switch action.Method {
|
||||
case C.RuleActionRejectMethodDefault:
|
||||
@@ -495,32 +492,80 @@ func (r *Router) exchangeWithRules(ctx context.Context, rules []adapter.DNSRule,
|
||||
Question: []mDNS.Question{message.Question[0]},
|
||||
},
|
||||
rejectAction: action,
|
||||
}
|
||||
}, nil
|
||||
case C.RuleActionRejectMethodDrop:
|
||||
return exchangeWithRulesResult{
|
||||
rejectAction: action,
|
||||
err: R.ErrDrop,
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
case *R.RuleActionPredefined:
|
||||
return exchangeWithRulesResult{
|
||||
response: action.Response(message),
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
transport := r.transport.Default()
|
||||
exchangeOptions := effectiveOptions
|
||||
if exchangeOptions.Strategy == C.DomainStrategyAsIS {
|
||||
exchangeOptions.Strategy = r.defaultDomainStrategy
|
||||
return exchangeWithRulesResult{}, &dnsPendingExchange{transport: r.transport.Default(), options: state.effectiveOptions}
|
||||
}
|
||||
|
||||
func (r *Router) exchangeWithRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, options adapter.DNSQueryOptions, allowFakeIP bool) exchangeWithRulesResult {
|
||||
state := dnsRuleWalkState{effectiveOptions: options}
|
||||
result, pending := r.walkDNSRules(ctx, rules, message, &state, allowFakeIP)
|
||||
if pending == nil {
|
||||
return result
|
||||
}
|
||||
response, err := r.client.Exchange(adapter.OverrideContext(ctx), transport, message, exchangeOptions, nil)
|
||||
return exchangeWithRulesResult{
|
||||
response: response,
|
||||
transport: transport,
|
||||
err: err,
|
||||
return r.resumeExchangeWithRules(ctx, rules, message, &state, allowFakeIP, pending)
|
||||
}
|
||||
|
||||
func (r *Router) resumeExchangeWithRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, state *dnsRuleWalkState, allowFakeIP bool, pending *dnsPendingExchange) exchangeWithRulesResult {
|
||||
for {
|
||||
response, err := r.client.Exchange(adapter.OverrideContext(ctx), pending.transport, message, r.finalizeExchangeOptions(pending.options), nil)
|
||||
if !pending.evaluate {
|
||||
return exchangeWithRulesResult{
|
||||
response: response,
|
||||
transport: pending.transport,
|
||||
err: err,
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for ", FormatQuestion(message.Question[0].String())))
|
||||
state.evaluatedResponse = nil
|
||||
state.evaluatedTransport = nil
|
||||
} else {
|
||||
state.evaluatedResponse = response
|
||||
state.evaluatedTransport = pending.transport
|
||||
}
|
||||
state.ruleIndex++
|
||||
var result exchangeWithRulesResult
|
||||
result, pending = r.walkDNSRules(ctx, rules, message, state, allowFakeIP)
|
||||
if pending == nil {
|
||||
return result
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Router) exchangeWithRulesAsync(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, options adapter.DNSQueryOptions, allowFakeIP bool, callback func(result exchangeWithRulesResult)) {
|
||||
state := dnsRuleWalkState{effectiveOptions: options}
|
||||
result, pending := r.walkDNSRules(ctx, rules, message, &state, allowFakeIP)
|
||||
if pending == nil {
|
||||
callback(result)
|
||||
return
|
||||
}
|
||||
if pending.evaluate {
|
||||
go func() {
|
||||
callback(r.resumeExchangeWithRules(ctx, rules, message, &state, allowFakeIP, pending))
|
||||
}()
|
||||
return
|
||||
}
|
||||
r.client.ExchangeAsync(adapter.OverrideContext(ctx), pending.transport, message, r.finalizeExchangeOptions(pending.options), nil, func(response *mDNS.Msg, err error) {
|
||||
callback(exchangeWithRulesResult{
|
||||
response: response,
|
||||
transport: pending.transport,
|
||||
err: err,
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
func (r *Router) resolveLookupStrategy(options adapter.DNSQueryOptions) C.DomainStrategy {
|
||||
if options.LookupStrategy != C.DomainStrategyAsIS {
|
||||
return options.LookupStrategy
|
||||
@@ -617,35 +662,35 @@ func (r *Router) lookupWithRulesType(ctx context.Context, rules []adapter.DNSRul
|
||||
return filterAddressesByQueryType(MessageToAddresses(exchangeResult.response), qType), nil
|
||||
}
|
||||
|
||||
func (r *Router) Exchange(ctx context.Context, message *mDNS.Msg, options adapter.DNSQueryOptions) (*mDNS.Msg, error) {
|
||||
type dnsExchangeContext struct {
|
||||
ctx context.Context
|
||||
rules []adapter.DNSRule
|
||||
legacyDNSMode bool
|
||||
metadata *adapter.InboundContext
|
||||
}
|
||||
|
||||
func (r *Router) prepareExchange(ctx context.Context, message *mDNS.Msg) (*dnsExchangeContext, *mDNS.Msg, error) {
|
||||
if len(message.Question) != 1 {
|
||||
r.logger.WarnContext(ctx, "bad question size: ", len(message.Question))
|
||||
responseMessage := mDNS.Msg{
|
||||
return nil, &mDNS.Msg{
|
||||
MsgHdr: mDNS.MsgHdr{
|
||||
Id: message.Id,
|
||||
Response: true,
|
||||
Rcode: mDNS.RcodeFormatError,
|
||||
},
|
||||
Question: message.Question,
|
||||
}
|
||||
return &responseMessage, nil
|
||||
}, nil
|
||||
}
|
||||
r.rulesAccess.RLock()
|
||||
if r.closing {
|
||||
r.rulesAccess.RUnlock()
|
||||
return nil, E.New("dns router closed")
|
||||
return nil, nil, E.New("dns router closed")
|
||||
}
|
||||
rules := r.rules
|
||||
legacyDNSMode := r.legacyDNSMode
|
||||
r.rulesAccess.RUnlock()
|
||||
r.logger.DebugContext(ctx, "exchange ", FormatQuestion(message.Question[0].String()))
|
||||
var (
|
||||
response *mDNS.Msg
|
||||
transport adapter.DNSTransport
|
||||
err error
|
||||
)
|
||||
var metadata *adapter.InboundContext
|
||||
ctx, metadata = adapter.ExtendContext(ctx)
|
||||
ctx, metadata := adapter.ExtendContext(ctx)
|
||||
metadata.Destination = M.Socksaddr{}
|
||||
metadata.QueryType = message.Question[0].Qtype
|
||||
metadata.DNSResponse = nil
|
||||
@@ -657,76 +702,15 @@ func (r *Router) Exchange(ctx context.Context, message *mDNS.Msg, options adapte
|
||||
metadata.IPVersion = 6
|
||||
}
|
||||
metadata.Domain = FqdnToDomain(message.Question[0].Name)
|
||||
if options.Transport != nil {
|
||||
transport = options.Transport
|
||||
if options.Strategy == C.DomainStrategyAsIS {
|
||||
options.Strategy = r.defaultDomainStrategy
|
||||
}
|
||||
response, err = r.client.Exchange(ctx, transport, message, options, nil)
|
||||
} else if !legacyDNSMode {
|
||||
exchangeResult := r.exchangeWithRules(ctx, rules, message, options, true)
|
||||
response, transport, err = exchangeResult.response, exchangeResult.transport, exchangeResult.err
|
||||
} else {
|
||||
var (
|
||||
rule adapter.DNSRule
|
||||
ruleIndex int
|
||||
)
|
||||
ruleIndex = -1
|
||||
for {
|
||||
dnsCtx := adapter.OverrideContext(ctx)
|
||||
dnsOptions := options
|
||||
transport, rule, ruleIndex = r.matchDNS(ctx, rules, true, ruleIndex, isAddressQuery(message), &dnsOptions)
|
||||
if rule != nil {
|
||||
switch action := rule.Action().(type) {
|
||||
case *R.RuleActionReject:
|
||||
switch action.Method {
|
||||
case C.RuleActionRejectMethodDefault:
|
||||
return &mDNS.Msg{
|
||||
MsgHdr: mDNS.MsgHdr{
|
||||
Id: message.Id,
|
||||
Rcode: mDNS.RcodeRefused,
|
||||
Response: true,
|
||||
},
|
||||
Question: []mDNS.Question{message.Question[0]},
|
||||
}, nil
|
||||
case C.RuleActionRejectMethodDrop:
|
||||
return nil, R.ErrDrop
|
||||
}
|
||||
case *R.RuleActionPredefined:
|
||||
err = nil
|
||||
response = action.Response(message)
|
||||
goto done
|
||||
}
|
||||
}
|
||||
responseCheck := addressLimitResponseCheck(rule, metadata)
|
||||
if dnsOptions.Strategy == C.DomainStrategyAsIS {
|
||||
dnsOptions.Strategy = r.defaultDomainStrategy
|
||||
}
|
||||
response, err = r.client.Exchange(dnsCtx, transport, message, dnsOptions, responseCheck)
|
||||
var rejected bool
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrResponseRejectedCached) {
|
||||
rejected = true
|
||||
r.logger.DebugContext(ctx, E.Cause(err, "response rejected for ", FormatQuestion(message.Question[0].String())), " (cached)")
|
||||
} else if errors.Is(err, ErrResponseRejected) {
|
||||
rejected = true
|
||||
r.logger.DebugContext(ctx, E.Cause(err, "response rejected for ", FormatQuestion(message.Question[0].String())))
|
||||
} else if len(message.Question) > 0 {
|
||||
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for ", FormatQuestion(message.Question[0].String())))
|
||||
} else {
|
||||
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for <empty query>"))
|
||||
}
|
||||
}
|
||||
if responseCheck != nil && rejected {
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
done:
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dnsExchangeContext{
|
||||
ctx: ctx,
|
||||
rules: rules,
|
||||
legacyDNSMode: legacyDNSMode,
|
||||
metadata: metadata,
|
||||
}, nil, nil
|
||||
}
|
||||
|
||||
func (r *Router) recordReverseMapping(message *mDNS.Msg, response *mDNS.Msg, transport adapter.DNSTransport) {
|
||||
if r.dnsReverseMapping != nil && len(message.Question) > 0 && response != nil && len(response.Answer) > 0 {
|
||||
if transport == nil || transport.Type() != C.DNSTypeFakeIP {
|
||||
for _, answer := range response.Answer {
|
||||
@@ -739,9 +723,121 @@ done:
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Router) exchangeLegacy(ctx context.Context, exchangeCtx *dnsExchangeContext, message *mDNS.Msg, options adapter.DNSQueryOptions) (*mDNS.Msg, adapter.DNSTransport, error) {
|
||||
var (
|
||||
transport adapter.DNSTransport
|
||||
rule adapter.DNSRule
|
||||
ruleIndex int
|
||||
)
|
||||
ruleIndex = -1
|
||||
for {
|
||||
dnsCtx := adapter.OverrideContext(ctx)
|
||||
dnsOptions := options
|
||||
transport, rule, ruleIndex = r.matchDNS(ctx, exchangeCtx.rules, true, ruleIndex, isAddressQuery(message), &dnsOptions)
|
||||
if rule != nil {
|
||||
switch action := rule.Action().(type) {
|
||||
case *R.RuleActionReject:
|
||||
switch action.Method {
|
||||
case C.RuleActionRejectMethodDefault:
|
||||
return &mDNS.Msg{
|
||||
MsgHdr: mDNS.MsgHdr{
|
||||
Id: message.Id,
|
||||
Rcode: mDNS.RcodeRefused,
|
||||
Response: true,
|
||||
},
|
||||
Question: []mDNS.Question{message.Question[0]},
|
||||
}, nil, nil
|
||||
case C.RuleActionRejectMethodDrop:
|
||||
return nil, nil, R.ErrDrop
|
||||
}
|
||||
case *R.RuleActionPredefined:
|
||||
return action.Response(message), nil, nil
|
||||
}
|
||||
}
|
||||
responseCheck := addressLimitResponseCheck(rule, exchangeCtx.metadata)
|
||||
response, err := r.client.Exchange(dnsCtx, transport, message, r.finalizeExchangeOptions(dnsOptions), responseCheck)
|
||||
var rejected bool
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrResponseRejectedCached) {
|
||||
rejected = true
|
||||
r.logger.DebugContext(ctx, E.Cause(err, "response rejected for ", FormatQuestion(message.Question[0].String())), " (cached)")
|
||||
} else if errors.Is(err, ErrResponseRejected) {
|
||||
rejected = true
|
||||
r.logger.DebugContext(ctx, E.Cause(err, "response rejected for ", FormatQuestion(message.Question[0].String())))
|
||||
} else if len(message.Question) > 0 {
|
||||
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for ", FormatQuestion(message.Question[0].String())))
|
||||
} else {
|
||||
r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for <empty query>"))
|
||||
}
|
||||
}
|
||||
if responseCheck != nil && rejected {
|
||||
continue
|
||||
}
|
||||
return response, transport, err
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Router) Exchange(ctx context.Context, message *mDNS.Msg, options adapter.DNSQueryOptions) (*mDNS.Msg, error) {
|
||||
exchangeCtx, earlyResponse, err := r.prepareExchange(ctx, message)
|
||||
if exchangeCtx == nil {
|
||||
return earlyResponse, err
|
||||
}
|
||||
ctx = exchangeCtx.ctx
|
||||
var (
|
||||
response *mDNS.Msg
|
||||
transport adapter.DNSTransport
|
||||
)
|
||||
if options.Transport != nil {
|
||||
transport = options.Transport
|
||||
response, err = r.client.Exchange(ctx, transport, message, r.finalizeExchangeOptions(options), nil)
|
||||
} else if !exchangeCtx.legacyDNSMode {
|
||||
exchangeResult := r.exchangeWithRules(ctx, exchangeCtx.rules, message, options, true)
|
||||
response, transport, err = exchangeResult.response, exchangeResult.transport, exchangeResult.err
|
||||
} else {
|
||||
response, transport, err = r.exchangeLegacy(ctx, exchangeCtx, message, options)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r.recordReverseMapping(message, response, transport)
|
||||
return response, nil
|
||||
}
|
||||
|
||||
func (r *Router) ExchangeAsync(ctx context.Context, message *mDNS.Msg, options adapter.DNSQueryOptions, callback func(response *mDNS.Msg, err error)) {
|
||||
exchangeCtx, earlyResponse, err := r.prepareExchange(ctx, message)
|
||||
if exchangeCtx == nil {
|
||||
callback(earlyResponse, err)
|
||||
return
|
||||
}
|
||||
ctx = exchangeCtx.ctx
|
||||
if options.Transport != nil {
|
||||
transport := options.Transport
|
||||
r.client.ExchangeAsync(ctx, transport, message, r.finalizeExchangeOptions(options), nil, func(response *mDNS.Msg, exchangeErr error) {
|
||||
r.finishExchangeAsync(message, transport, response, exchangeErr, callback)
|
||||
})
|
||||
} else if !exchangeCtx.legacyDNSMode {
|
||||
r.exchangeWithRulesAsync(ctx, exchangeCtx.rules, message, options, true, func(result exchangeWithRulesResult) {
|
||||
r.finishExchangeAsync(message, result.transport, result.response, result.err, callback)
|
||||
})
|
||||
} else {
|
||||
go func() {
|
||||
response, transport, exchangeErr := r.exchangeLegacy(ctx, exchangeCtx, message, options)
|
||||
r.finishExchangeAsync(message, transport, response, exchangeErr, callback)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Router) finishExchangeAsync(message *mDNS.Msg, transport adapter.DNSTransport, response *mDNS.Msg, err error, callback func(response *mDNS.Msg, err error)) {
|
||||
if err != nil {
|
||||
callback(nil, err)
|
||||
return
|
||||
}
|
||||
r.recordReverseMapping(message, response, transport)
|
||||
callback(response, nil)
|
||||
}
|
||||
|
||||
func (r *Router) Lookup(ctx context.Context, domain string, options adapter.DNSQueryOptions) ([]netip.Addr, error) {
|
||||
r.rulesAccess.RLock()
|
||||
if r.closing {
|
||||
|
||||
+119
-23
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -10,4 +10,5 @@ type ResolvedResolver interface {
|
||||
Start() error
|
||||
Close() error
|
||||
Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error)
|
||||
ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error))
|
||||
}
|
||||
|
||||
@@ -159,6 +159,51 @@ func (t *DBusResolvedResolver) Exchange(ctx context.Context, message *mDNS.Msg)
|
||||
return t.exchangeServerSet(ctx, message, refreshedServerSet)
|
||||
}
|
||||
|
||||
func (t *DBusResolvedResolver) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
|
||||
serverSet := t.savedServerSet.Load()
|
||||
if serverSet == nil {
|
||||
go func() {
|
||||
callback(t.Exchange(ctx, message))
|
||||
}()
|
||||
return
|
||||
}
|
||||
t.exchangeServerSetAsync(ctx, message, serverSet, func(response *mDNS.Msg, err error) {
|
||||
if err == nil {
|
||||
callback(response, nil)
|
||||
return
|
||||
}
|
||||
go func() {
|
||||
t.updateStatus()
|
||||
refreshedServerSet := t.savedServerSet.Load()
|
||||
if refreshedServerSet == nil || refreshedServerSet == serverSet {
|
||||
callback(nil, err)
|
||||
return
|
||||
}
|
||||
t.exchangeServerSetAsync(ctx, message, refreshedServerSet, callback)
|
||||
}()
|
||||
})
|
||||
}
|
||||
|
||||
func (t *DBusResolvedResolver) exchangeServerSetAsync(ctx context.Context, message *mDNS.Msg, serverSet *resolvedServerSet, callback func(response *mDNS.Msg, err error)) {
|
||||
if len(serverSet.servers) == 0 {
|
||||
callback(nil, E.New("link has no DNS servers configured"))
|
||||
return
|
||||
}
|
||||
serverExchangers := make([]dnsTransport.AsyncExchanger, 0, len(serverSet.servers))
|
||||
for _, server := range serverSet.servers {
|
||||
serverExchangers = append(serverExchangers, func(exchangeCtx context.Context, exchangeCallback func(response *mDNS.Msg, err error)) {
|
||||
server.primaryTransport.ExchangeAsync(exchangeCtx, message, func(response *mDNS.Msg, exchangeErr error) {
|
||||
if exchangeErr != nil && server.fallbackTransport != nil {
|
||||
server.fallbackTransport.ExchangeAsync(exchangeCtx, message, exchangeCallback)
|
||||
return
|
||||
}
|
||||
exchangeCallback(response, exchangeErr)
|
||||
})
|
||||
})
|
||||
}
|
||||
dnsTransport.ExchangeSequential(ctx, serverExchangers, nil, callback)
|
||||
}
|
||||
|
||||
func (t *DBusResolvedResolver) loopUpdateStatus() {
|
||||
signalChan := make(chan *dbus.Signal, 1)
|
||||
t.systemBus.Signal(signalChan)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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):
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}()
|
||||
}
|
||||
|
||||
@@ -145,6 +145,12 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
|
||||
return nil, err
|
||||
}
|
||||
|
||||
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
|
||||
go func() {
|
||||
callback(t.Exchange(ctx, message))
|
||||
}()
|
||||
}
|
||||
|
||||
func (t *Transport) exchange(ctx context.Context, message *mDNS.Msg, conn *quic.Conn) (*mDNS.Msg, error) {
|
||||
stream, err := conn.OpenStreamSync(ctx)
|
||||
if err != nil {
|
||||
|
||||
+34
-23
@@ -15,7 +15,6 @@ import (
|
||||
"github.com/sagernet/sing-box/option"
|
||||
"github.com/sagernet/sing/common"
|
||||
"github.com/sagernet/sing/common/buf"
|
||||
"github.com/sagernet/sing/common/bufio/deadline"
|
||||
E "github.com/sagernet/sing/common/exceptions"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
@@ -31,8 +30,9 @@ func RegisterTCP(registry *dns.TransportRegistry) {
|
||||
|
||||
type TCPTransport struct {
|
||||
dns.TransportAdapter
|
||||
dialer N.Dialer
|
||||
serverAddr M.Socksaddr
|
||||
dialer N.Dialer
|
||||
serverAddr M.Socksaddr
|
||||
multiplexer *queryMultiplexer
|
||||
}
|
||||
|
||||
func NewTCP(ctx context.Context, logger log.ContextLogger, tag string, options option.RemoteDNSServerOptions) (adapter.DNSTransport, error) {
|
||||
@@ -47,11 +47,31 @@ func NewTCP(ctx context.Context, logger log.ContextLogger, tag string, options o
|
||||
if !serverAddr.IsValid() {
|
||||
return nil, E.New("invalid server address: ", serverAddr)
|
||||
}
|
||||
return &TCPTransport{
|
||||
TransportAdapter: dns.NewTransportAdapterWithRemoteOptions(C.DNSTypeTCP, tag, options),
|
||||
dialer: transportDialer,
|
||||
return NewTCPRaw(dns.NewTransportAdapterWithRemoteOptions(C.DNSTypeTCP, tag, options), transportDialer, serverAddr), nil
|
||||
}
|
||||
|
||||
func NewTCPRaw(adapter dns.TransportAdapter, dialer N.Dialer, serverAddr M.Socksaddr) *TCPTransport {
|
||||
t := &TCPTransport{
|
||||
TransportAdapter: adapter,
|
||||
dialer: dialer,
|
||||
serverAddr: serverAddr,
|
||||
}, nil
|
||||
}
|
||||
t.multiplexer = newQueryMultiplexer(queryMultiplexerOptions{
|
||||
dial: func(ctx context.Context) (net.Conn, error) {
|
||||
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, t.serverAddr)
|
||||
if err != nil {
|
||||
return nil, E.Cause(err, "dial TCP connection")
|
||||
}
|
||||
return conn, nil
|
||||
},
|
||||
write: func(conn net.Conn, message *mDNS.Msg, queryId uint16) error {
|
||||
return WriteMessage(conn, queryId, message)
|
||||
},
|
||||
readNext: func(conn net.Conn) (*mDNS.Msg, error) {
|
||||
return ReadMessage(conn)
|
||||
},
|
||||
})
|
||||
return t
|
||||
}
|
||||
|
||||
func (t *TCPTransport) Start(stage adapter.StartStage) error {
|
||||
@@ -62,28 +82,19 @@ func (t *TCPTransport) Start(stage adapter.StartStage) error {
|
||||
}
|
||||
|
||||
func (t *TCPTransport) Close() error {
|
||||
return nil
|
||||
return t.multiplexer.Close()
|
||||
}
|
||||
|
||||
func (t *TCPTransport) Reset() {
|
||||
t.multiplexer.Reset()
|
||||
}
|
||||
|
||||
func (t *TCPTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
||||
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, t.serverAddr)
|
||||
if err != nil {
|
||||
return nil, E.Cause(err, "dial TCP connection")
|
||||
}
|
||||
defer conn.Close()
|
||||
defer setConnDeadline(ctx, conn, deadline.NeedAdditionalReadDeadline(conn))()
|
||||
err = WriteMessage(conn, 0, message)
|
||||
if err != nil {
|
||||
return nil, E.Cause(err, "write request")
|
||||
}
|
||||
response, err := ReadMessage(conn)
|
||||
if err != nil {
|
||||
return nil, E.Cause(err, "read response")
|
||||
}
|
||||
return response, nil
|
||||
return t.multiplexer.Exchange(ctx, message)
|
||||
}
|
||||
|
||||
func (t *TCPTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
|
||||
t.multiplexer.ExchangeAsync(ctx, message, callback)
|
||||
}
|
||||
|
||||
func setConnDeadline(ctx context.Context, conn net.Conn, needClose bool) func() {
|
||||
|
||||
+25
-67
@@ -2,6 +2,7 @@ package transport
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
|
||||
"github.com/sagernet/sing-box/adapter"
|
||||
"github.com/sagernet/sing-box/common/dialer"
|
||||
@@ -11,7 +12,6 @@ import (
|
||||
"github.com/sagernet/sing-box/log"
|
||||
"github.com/sagernet/sing-box/option"
|
||||
"github.com/sagernet/sing/common"
|
||||
"github.com/sagernet/sing/common/bufio/deadline"
|
||||
E "github.com/sagernet/sing/common/exceptions"
|
||||
"github.com/sagernet/sing/common/logger"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
@@ -22,26 +22,16 @@ import (
|
||||
|
||||
var _ adapter.DNSTransport = (*TLSTransport)(nil)
|
||||
|
||||
const tlsDNSMaxInflight = 8
|
||||
|
||||
func RegisterTLS(registry *dns.TransportRegistry) {
|
||||
dns.RegisterTransport[option.RemoteTLSDNSServerOptions](registry, C.DNSTypeTLS, NewTLS)
|
||||
}
|
||||
|
||||
type TLSTransport struct {
|
||||
dns.TransportAdapter
|
||||
logger logger.ContextLogger
|
||||
|
||||
logger logger.ContextLogger
|
||||
dialer tls.Dialer
|
||||
serverAddr M.Socksaddr
|
||||
tlsConfig tls.Config
|
||||
connections *ConnPool[*tlsDNSConn]
|
||||
}
|
||||
|
||||
type tlsDNSConn struct {
|
||||
tls.Conn
|
||||
queryId uint16
|
||||
needDeadlineClose bool
|
||||
multiplexer *queryMultiplexer
|
||||
}
|
||||
|
||||
func NewTLS(ctx context.Context, logger log.ContextLogger, tag string, options option.RemoteTLSDNSServerOptions) (adapter.DNSTransport, error) {
|
||||
@@ -66,23 +56,28 @@ func NewTLS(ctx context.Context, logger log.ContextLogger, tag string, options o
|
||||
}
|
||||
|
||||
func NewTLSRaw(logger logger.ContextLogger, adapter dns.TransportAdapter, dialer N.Dialer, serverAddr M.Socksaddr, tlsConfig tls.Config) *TLSTransport {
|
||||
return &TLSTransport{
|
||||
t := &TLSTransport{
|
||||
TransportAdapter: adapter,
|
||||
logger: logger,
|
||||
dialer: tls.NewDialer(dialer, tlsConfig),
|
||||
serverAddr: serverAddr,
|
||||
tlsConfig: tlsConfig,
|
||||
connections: NewConnPool(ConnPoolOptions[*tlsDNSConn]{
|
||||
Mode: ConnPoolOrdered,
|
||||
MaxInflight: tlsDNSMaxInflight,
|
||||
IsAlive: func(conn *tlsDNSConn) bool {
|
||||
return conn != nil
|
||||
},
|
||||
Close: func(conn *tlsDNSConn, _ error) {
|
||||
conn.Close()
|
||||
},
|
||||
}),
|
||||
}
|
||||
t.multiplexer = newQueryMultiplexer(queryMultiplexerOptions{
|
||||
dial: func(ctx context.Context) (net.Conn, error) {
|
||||
conn, err := t.dialer.DialTLSContext(ctx, t.serverAddr)
|
||||
if err != nil {
|
||||
return nil, E.Cause(err, "dial TLS connection")
|
||||
}
|
||||
return conn, nil
|
||||
},
|
||||
write: func(conn net.Conn, message *mDNS.Msg, queryId uint16) error {
|
||||
return WriteMessage(conn, queryId, message)
|
||||
},
|
||||
readNext: func(conn net.Conn) (*mDNS.Msg, error) {
|
||||
return ReadMessage(conn)
|
||||
},
|
||||
})
|
||||
return t
|
||||
}
|
||||
|
||||
func (t *TLSTransport) Start(stage adapter.StartStage) error {
|
||||
@@ -93,54 +88,17 @@ func (t *TLSTransport) Start(stage adapter.StartStage) error {
|
||||
}
|
||||
|
||||
func (t *TLSTransport) Close() error {
|
||||
return t.connections.Close()
|
||||
return t.multiplexer.Close()
|
||||
}
|
||||
|
||||
func (t *TLSTransport) Reset() {
|
||||
t.connections.Reset()
|
||||
t.multiplexer.Reset()
|
||||
}
|
||||
|
||||
func (t *TLSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
||||
var lastErr error
|
||||
for range 2 {
|
||||
conn, created, err := t.connections.Acquire(ctx, func(ctx context.Context) (*tlsDNSConn, error) {
|
||||
tlsConn, err := t.dialer.DialTLSContext(ctx, t.serverAddr)
|
||||
if err != nil {
|
||||
return nil, E.Cause(err, "dial TLS connection")
|
||||
}
|
||||
return &tlsDNSConn{
|
||||
Conn: tlsConn,
|
||||
needDeadlineClose: deadline.NeedAdditionalReadDeadline(tlsConn.NetConn()),
|
||||
}, nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
response, err := t.exchange(ctx, message, conn)
|
||||
if err == nil {
|
||||
t.connections.Release(conn, true)
|
||||
return response, nil
|
||||
}
|
||||
lastErr = err
|
||||
t.logger.DebugContext(ctx, "discarded pooled connection: ", err)
|
||||
t.connections.Release(conn, false)
|
||||
if created {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return nil, lastErr
|
||||
return t.multiplexer.Exchange(ctx, message)
|
||||
}
|
||||
|
||||
func (t *TLSTransport) exchange(ctx context.Context, message *mDNS.Msg, conn *tlsDNSConn) (*mDNS.Msg, error) {
|
||||
defer setConnDeadline(ctx, conn, conn.needDeadlineClose)()
|
||||
conn.queryId++
|
||||
err := WriteMessage(conn, conn.queryId, message)
|
||||
if err != nil {
|
||||
return nil, E.Cause(err, "write request")
|
||||
}
|
||||
response, err := ReadMessage(conn)
|
||||
if err != nil {
|
||||
return nil, E.Cause(err, "read response")
|
||||
}
|
||||
return response, nil
|
||||
func (t *TLSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
|
||||
t.multiplexer.ExchangeAsync(ctx, message, callback)
|
||||
}
|
||||
|
||||
+78
-156
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user