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:
@@ -18,6 +18,7 @@ import (
|
|||||||
type DNSRouter interface {
|
type DNSRouter interface {
|
||||||
Lifecycle
|
Lifecycle
|
||||||
Exchange(ctx context.Context, message *dns.Msg, options DNSQueryOptions) (*dns.Msg, error)
|
Exchange(ctx context.Context, message *dns.Msg, options DNSQueryOptions) (*dns.Msg, error)
|
||||||
|
ExchangeAsync(ctx context.Context, message *dns.Msg, options DNSQueryOptions, callback func(response *dns.Msg, err error))
|
||||||
Lookup(ctx context.Context, domain string, options DNSQueryOptions) ([]netip.Addr, error)
|
Lookup(ctx context.Context, domain string, options DNSQueryOptions) ([]netip.Addr, error)
|
||||||
ClearCache()
|
ClearCache()
|
||||||
LookupReverseMapping(ip netip.Addr) (string, bool)
|
LookupReverseMapping(ip netip.Addr) (string, bool)
|
||||||
@@ -27,6 +28,7 @@ type DNSRouter interface {
|
|||||||
type DNSClient interface {
|
type DNSClient interface {
|
||||||
Start()
|
Start()
|
||||||
Exchange(ctx context.Context, transport DNSTransport, message *dns.Msg, options DNSQueryOptions, responseChecker func(response *dns.Msg) bool) (*dns.Msg, error)
|
Exchange(ctx context.Context, transport DNSTransport, message *dns.Msg, options DNSQueryOptions, responseChecker func(response *dns.Msg) bool) (*dns.Msg, error)
|
||||||
|
ExchangeAsync(ctx context.Context, transport DNSTransport, message *dns.Msg, options DNSQueryOptions, responseChecker func(response *dns.Msg) bool, callback func(response *dns.Msg, err error))
|
||||||
Lookup(ctx context.Context, transport DNSTransport, domain string, options DNSQueryOptions, responseChecker func(response *dns.Msg) bool) ([]netip.Addr, error)
|
Lookup(ctx context.Context, transport DNSTransport, domain string, options DNSQueryOptions, responseChecker func(response *dns.Msg) bool) ([]netip.Addr, error)
|
||||||
ClearCache()
|
ClearCache()
|
||||||
}
|
}
|
||||||
@@ -84,6 +86,7 @@ type DNSTransport interface {
|
|||||||
// Exchanges that are currently using those connections may fail.
|
// Exchanges that are currently using those connections may fail.
|
||||||
Reset()
|
Reset()
|
||||||
Exchange(ctx context.Context, message *dns.Msg) (*dns.Msg, error)
|
Exchange(ctx context.Context, message *dns.Msg) (*dns.Msg, error)
|
||||||
|
ExchangeAsync(ctx context.Context, message *dns.Msg, callback func(response *dns.Msg, err error))
|
||||||
}
|
}
|
||||||
|
|
||||||
type DNSTransportWithPreferredDomain interface {
|
type DNSTransportWithPreferredDomain interface {
|
||||||
|
|||||||
+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 len(message.Question) == 0 {
|
||||||
if c.logger != nil {
|
if c.logger != nil {
|
||||||
c.logger.WarnContext(ctx, "bad question size: ", len(message.Question))
|
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]
|
question := message.Question[0]
|
||||||
if question.Qtype == dns.TypeA && options.Strategy == C.DomainStrategyIPv6Only || question.Qtype == dns.TypeAAAA && options.Strategy == C.DomainStrategyIPv4Only {
|
if question.Qtype == dns.TypeA && options.Strategy == C.DomainStrategyIPv6Only || question.Qtype == dns.TypeAAAA && options.Strategy == C.DomainStrategyIPv4Only {
|
||||||
if c.logger != nil {
|
if c.logger != nil {
|
||||||
c.logger.DebugContext(ctx, "strategy rejected")
|
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)
|
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) &&
|
len(message.Extra[0].(*dns.OPT).Option) == 0) &&
|
||||||
!options.ClientSubnet.IsValid()
|
!options.ClientSubnet.IsValid()
|
||||||
disableCache := !isSimpleRequest || c.disableCache || options.DisableCache
|
disableCache := !isSimpleRequest || c.disableCache || options.DisableCache
|
||||||
|
operation := &exchangeOperation{
|
||||||
|
message: message,
|
||||||
|
question: question,
|
||||||
|
messageId: message.Id,
|
||||||
|
options: options,
|
||||||
|
responseChecker: responseChecker,
|
||||||
|
disableCache: disableCache,
|
||||||
|
}
|
||||||
if !disableCache {
|
if !disableCache {
|
||||||
cacheKey := dnsCacheKey{Question: question, transportTag: transport.Tag()}
|
cacheKey := dnsCacheKey{Question: question, transportTag: transport.Tag()}
|
||||||
cond, loaded := c.cacheLock.LoadOrStore(cacheKey, make(chan struct{}))
|
cond, loaded := c.cacheLock.LoadOrStore(cacheKey, make(chan struct{}))
|
||||||
if loaded {
|
if loaded {
|
||||||
|
if !allowWait {
|
||||||
|
return nil, nil, exchangeWait, nil
|
||||||
|
}
|
||||||
select {
|
select {
|
||||||
case <-cond:
|
case <-cond:
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return nil, ctx.Err()
|
return nil, nil, exchangeDone, ctx.Err()
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
defer func() {
|
operation.releaseCond = func() {
|
||||||
c.cacheLock.Delete(cacheKey)
|
c.cacheLock.Delete(cacheKey)
|
||||||
close(cond)
|
close(cond)
|
||||||
}()
|
}
|
||||||
}
|
}
|
||||||
response, ttl, isStale := c.loadResponse(question, transport)
|
response, ttl, isStale := c.loadResponse(question, transport)
|
||||||
if response != nil {
|
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)
|
c.backgroundRefreshDNS(transport, question, message.Copy(), options, responseChecker)
|
||||||
logOptimisticResponse(c.logger, ctx, response)
|
logOptimisticResponse(c.logger, ctx, response)
|
||||||
response.Id = message.Id
|
response.Id = message.Id
|
||||||
return response, nil
|
operation.release()
|
||||||
|
return nil, response, exchangeDone, nil
|
||||||
} else if !isStale {
|
} else if !isStale {
|
||||||
logCachedResponse(c.logger, ctx, response, ttl)
|
logCachedResponse(c.logger, ctx, response, ttl)
|
||||||
response.Id = message.Id
|
response.Id = message.Id
|
||||||
return response, nil
|
operation.release()
|
||||||
|
return nil, response, exchangeDone, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
messageId := message.Id
|
contextTransport, transportTagLoaded := transportTagFromContext(ctx)
|
||||||
contextTransport, clientSubnetLoaded := transportTagFromContext(ctx)
|
if transportTagLoaded && transport.Tag() == contextTransport {
|
||||||
if clientSubnetLoaded && transport.Tag() == contextTransport {
|
operation.release()
|
||||||
return nil, E.New("DNS query loopback in transport[", contextTransport, "]")
|
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 {
|
if !disableCache && responseChecker != nil && c.rdrc != nil {
|
||||||
rejected := c.rdrc.LoadRDRC(transport.Tag(), question.Name, question.Qtype)
|
rejected := c.rdrc.LoadRDRC(transport.Tag(), question.Name, question.Qtype)
|
||||||
if rejected {
|
if rejected {
|
||||||
return nil, ErrResponseRejectedCached
|
operation.release()
|
||||||
|
return nil, nil, exchangeDone, ErrResponseRejectedCached
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
response, err := c.exchangeToTransport(ctx, transport, message, options.Timeout)
|
return operation, nil, exchangeReady, nil
|
||||||
if err != nil {
|
}
|
||||||
return nil, err
|
|
||||||
}
|
func (c *Client) finishExchange(transport adapter.DNSTransport, operation *exchangeOperation, response *dns.Msg) (*dns.Msg, error) {
|
||||||
disableCache = disableCache || (response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError)
|
ctx := operation.ctx
|
||||||
if responseChecker != nil {
|
question := operation.question
|
||||||
|
disableCache := operation.disableCache || (response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError)
|
||||||
|
if operation.responseChecker != nil {
|
||||||
var rejected bool
|
var rejected bool
|
||||||
if response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError {
|
if response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError {
|
||||||
rejected = true
|
rejected = true
|
||||||
} else {
|
} else {
|
||||||
rejected = !responseChecker(response)
|
rejected = !operation.responseChecker(response)
|
||||||
}
|
}
|
||||||
if rejected {
|
if rejected {
|
||||||
if !disableCache && c.rdrc != nil {
|
if !disableCache && c.rdrc != nil {
|
||||||
@@ -239,12 +281,12 @@ func (c *Client) Exchange(ctx context.Context, transport adapter.DNSTransport, m
|
|||||||
return response, ErrResponseRejected
|
return response, ErrResponseRejected
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
timeToLive := applyResponseOptions(question, response, options)
|
timeToLive := applyResponseOptions(question, response, operation.options)
|
||||||
if !disableCache {
|
if !disableCache {
|
||||||
c.storeCache(transport, question, response, timeToLive)
|
c.storeCache(transport, question, response, timeToLive)
|
||||||
}
|
}
|
||||||
response.Id = messageId
|
response.Id = operation.messageId
|
||||||
requestEDNSOpt := message.IsEdns0()
|
requestEDNSOpt := operation.message.IsEdns0()
|
||||||
responseEDNSOpt := response.IsEdns0()
|
responseEDNSOpt := response.IsEdns0()
|
||||||
if responseEDNSOpt != nil && (requestEDNSOpt == nil || requestEDNSOpt.Version() < responseEDNSOpt.Version()) {
|
if responseEDNSOpt != nil && (requestEDNSOpt == nil || requestEDNSOpt.Version() < responseEDNSOpt.Version()) {
|
||||||
response.Extra = common.Filter(response.Extra, func(it dns.RR) bool {
|
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
|
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) {
|
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)
|
domain = FqdnToDomain(domain)
|
||||||
dnsName := dns.Fqdn(domain)
|
dnsName := dns.Fqdn(domain)
|
||||||
@@ -562,6 +642,27 @@ func (c *Client) exchangeToTransport(ctx context.Context, transport adapter.DNST
|
|||||||
return nil, err
|
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 {
|
func MessageToAddresses(response *dns.Msg) []netip.Addr {
|
||||||
return adapter.DNSResponseAddresses(response)
|
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"
|
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)
|
metadata := adapter.ContextFrom(ctx)
|
||||||
if metadata == nil {
|
if metadata == nil {
|
||||||
panic("no context")
|
panic("no context")
|
||||||
}
|
}
|
||||||
effectiveOptions := options
|
for ; state.ruleIndex < len(rules); state.ruleIndex++ {
|
||||||
var evaluatedResponse *mDNS.Msg
|
currentRule := rules[state.ruleIndex]
|
||||||
var evaluatedTransport adapter.DNSTransport
|
|
||||||
for currentRuleIndex, currentRule := range rules {
|
|
||||||
metadata.ResetRuleCache()
|
metadata.ResetRuleCache()
|
||||||
metadata.DNSResponse = evaluatedResponse
|
metadata.DNSResponse = state.evaluatedResponse
|
||||||
metadata.DestinationAddressMatchFromResponse = false
|
metadata.DestinationAddressMatchFromResponse = false
|
||||||
if !currentRule.Match(metadata) {
|
if !currentRule.Match(metadata) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
r.logRuleMatch(ctx, currentRuleIndex, currentRule)
|
r.logRuleMatch(ctx, state.ruleIndex, currentRule)
|
||||||
switch action := currentRule.Action().(type) {
|
switch action := currentRule.Action().(type) {
|
||||||
case *R.RuleActionDNSRouteOptions:
|
case *R.RuleActionDNSRouteOptions:
|
||||||
r.applyDNSRouteOptions(&effectiveOptions, *action)
|
r.applyDNSRouteOptions(&state.effectiveOptions, *action)
|
||||||
case *R.RuleActionEvaluate:
|
case *R.RuleActionEvaluate:
|
||||||
queryOptions := effectiveOptions
|
queryOptions := state.effectiveOptions
|
||||||
transport, loaded := r.transport.Transport(action.Server)
|
transport, loaded := r.transport.Transport(action.Server)
|
||||||
if !loaded {
|
if !loaded {
|
||||||
r.logger.ErrorContext(ctx, "transport not found: ", action.Server)
|
r.logger.ErrorContext(ctx, "transport not found: ", action.Server)
|
||||||
evaluatedResponse = nil
|
state.evaluatedResponse = nil
|
||||||
evaluatedTransport = nil
|
state.evaluatedTransport = nil
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
r.applyDNSRouteOptions(&queryOptions, action.RuleActionDNSRouteOptions)
|
r.applyDNSRouteOptions(&queryOptions, action.RuleActionDNSRouteOptions)
|
||||||
exchangeOptions := queryOptions
|
return exchangeWithRulesResult{}, &dnsPendingExchange{transport: transport, options: queryOptions, evaluate: true}
|
||||||
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
|
|
||||||
case *R.RuleActionRespond:
|
case *R.RuleActionRespond:
|
||||||
if evaluatedResponse == nil {
|
if state.evaluatedResponse == nil {
|
||||||
return exchangeWithRulesResult{
|
return exchangeWithRulesResult{
|
||||||
err: E.New(dnsRespondMissingResponseMessage),
|
err: E.New(dnsRespondMissingResponseMessage),
|
||||||
}
|
}, nil
|
||||||
}
|
}
|
||||||
return exchangeWithRulesResult{
|
return exchangeWithRulesResult{
|
||||||
response: evaluatedResponse,
|
response: state.evaluatedResponse,
|
||||||
transport: evaluatedTransport,
|
transport: state.evaluatedTransport,
|
||||||
}
|
}, nil
|
||||||
case *R.RuleActionDNSRoute:
|
case *R.RuleActionDNSRoute:
|
||||||
queryOptions := effectiveOptions
|
queryOptions := state.effectiveOptions
|
||||||
transport, status := r.resolveDNSRoute(action.Server, action.RuleActionDNSRouteOptions, allowFakeIP, &queryOptions)
|
transport, status := r.resolveDNSRoute(action.Server, action.RuleActionDNSRouteOptions, allowFakeIP, &queryOptions)
|
||||||
switch status {
|
switch status {
|
||||||
case dnsRouteStatusMissing:
|
case dnsRouteStatusMissing:
|
||||||
@@ -472,16 +478,7 @@ func (r *Router) exchangeWithRules(ctx context.Context, rules []adapter.DNSRule,
|
|||||||
case dnsRouteStatusSkipped:
|
case dnsRouteStatusSkipped:
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
exchangeOptions := queryOptions
|
return exchangeWithRulesResult{}, &dnsPendingExchange{transport: transport, options: 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,
|
|
||||||
}
|
|
||||||
case *R.RuleActionReject:
|
case *R.RuleActionReject:
|
||||||
switch action.Method {
|
switch action.Method {
|
||||||
case C.RuleActionRejectMethodDefault:
|
case C.RuleActionRejectMethodDefault:
|
||||||
@@ -495,32 +492,80 @@ func (r *Router) exchangeWithRules(ctx context.Context, rules []adapter.DNSRule,
|
|||||||
Question: []mDNS.Question{message.Question[0]},
|
Question: []mDNS.Question{message.Question[0]},
|
||||||
},
|
},
|
||||||
rejectAction: action,
|
rejectAction: action,
|
||||||
}
|
}, nil
|
||||||
case C.RuleActionRejectMethodDrop:
|
case C.RuleActionRejectMethodDrop:
|
||||||
return exchangeWithRulesResult{
|
return exchangeWithRulesResult{
|
||||||
rejectAction: action,
|
rejectAction: action,
|
||||||
err: R.ErrDrop,
|
err: R.ErrDrop,
|
||||||
}
|
}, nil
|
||||||
}
|
}
|
||||||
case *R.RuleActionPredefined:
|
case *R.RuleActionPredefined:
|
||||||
return exchangeWithRulesResult{
|
return exchangeWithRulesResult{
|
||||||
response: action.Response(message),
|
response: action.Response(message),
|
||||||
}
|
}, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
transport := r.transport.Default()
|
return exchangeWithRulesResult{}, &dnsPendingExchange{transport: r.transport.Default(), options: state.effectiveOptions}
|
||||||
exchangeOptions := effectiveOptions
|
}
|
||||||
if exchangeOptions.Strategy == C.DomainStrategyAsIS {
|
|
||||||
exchangeOptions.Strategy = r.defaultDomainStrategy
|
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 r.resumeExchangeWithRules(ctx, rules, message, &state, allowFakeIP, pending)
|
||||||
return exchangeWithRulesResult{
|
}
|
||||||
response: response,
|
|
||||||
transport: transport,
|
func (r *Router) resumeExchangeWithRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, state *dnsRuleWalkState, allowFakeIP bool, pending *dnsPendingExchange) exchangeWithRulesResult {
|
||||||
err: err,
|
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 {
|
func (r *Router) resolveLookupStrategy(options adapter.DNSQueryOptions) C.DomainStrategy {
|
||||||
if options.LookupStrategy != C.DomainStrategyAsIS {
|
if options.LookupStrategy != C.DomainStrategyAsIS {
|
||||||
return options.LookupStrategy
|
return options.LookupStrategy
|
||||||
@@ -617,35 +662,35 @@ func (r *Router) lookupWithRulesType(ctx context.Context, rules []adapter.DNSRul
|
|||||||
return filterAddressesByQueryType(MessageToAddresses(exchangeResult.response), qType), nil
|
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 {
|
if len(message.Question) != 1 {
|
||||||
r.logger.WarnContext(ctx, "bad question size: ", len(message.Question))
|
r.logger.WarnContext(ctx, "bad question size: ", len(message.Question))
|
||||||
responseMessage := mDNS.Msg{
|
return nil, &mDNS.Msg{
|
||||||
MsgHdr: mDNS.MsgHdr{
|
MsgHdr: mDNS.MsgHdr{
|
||||||
Id: message.Id,
|
Id: message.Id,
|
||||||
Response: true,
|
Response: true,
|
||||||
Rcode: mDNS.RcodeFormatError,
|
Rcode: mDNS.RcodeFormatError,
|
||||||
},
|
},
|
||||||
Question: message.Question,
|
Question: message.Question,
|
||||||
}
|
}, nil
|
||||||
return &responseMessage, nil
|
|
||||||
}
|
}
|
||||||
r.rulesAccess.RLock()
|
r.rulesAccess.RLock()
|
||||||
if r.closing {
|
if r.closing {
|
||||||
r.rulesAccess.RUnlock()
|
r.rulesAccess.RUnlock()
|
||||||
return nil, E.New("dns router closed")
|
return nil, nil, E.New("dns router closed")
|
||||||
}
|
}
|
||||||
rules := r.rules
|
rules := r.rules
|
||||||
legacyDNSMode := r.legacyDNSMode
|
legacyDNSMode := r.legacyDNSMode
|
||||||
r.rulesAccess.RUnlock()
|
r.rulesAccess.RUnlock()
|
||||||
r.logger.DebugContext(ctx, "exchange ", FormatQuestion(message.Question[0].String()))
|
r.logger.DebugContext(ctx, "exchange ", FormatQuestion(message.Question[0].String()))
|
||||||
var (
|
ctx, metadata := adapter.ExtendContext(ctx)
|
||||||
response *mDNS.Msg
|
|
||||||
transport adapter.DNSTransport
|
|
||||||
err error
|
|
||||||
)
|
|
||||||
var metadata *adapter.InboundContext
|
|
||||||
ctx, metadata = adapter.ExtendContext(ctx)
|
|
||||||
metadata.Destination = M.Socksaddr{}
|
metadata.Destination = M.Socksaddr{}
|
||||||
metadata.QueryType = message.Question[0].Qtype
|
metadata.QueryType = message.Question[0].Qtype
|
||||||
metadata.DNSResponse = nil
|
metadata.DNSResponse = nil
|
||||||
@@ -657,76 +702,15 @@ func (r *Router) Exchange(ctx context.Context, message *mDNS.Msg, options adapte
|
|||||||
metadata.IPVersion = 6
|
metadata.IPVersion = 6
|
||||||
}
|
}
|
||||||
metadata.Domain = FqdnToDomain(message.Question[0].Name)
|
metadata.Domain = FqdnToDomain(message.Question[0].Name)
|
||||||
if options.Transport != nil {
|
return &dnsExchangeContext{
|
||||||
transport = options.Transport
|
ctx: ctx,
|
||||||
if options.Strategy == C.DomainStrategyAsIS {
|
rules: rules,
|
||||||
options.Strategy = r.defaultDomainStrategy
|
legacyDNSMode: legacyDNSMode,
|
||||||
}
|
metadata: metadata,
|
||||||
response, err = r.client.Exchange(ctx, transport, message, options, nil)
|
}, nil, nil
|
||||||
} else if !legacyDNSMode {
|
}
|
||||||
exchangeResult := r.exchangeWithRules(ctx, rules, message, options, true)
|
|
||||||
response, transport, err = exchangeResult.response, exchangeResult.transport, exchangeResult.err
|
func (r *Router) recordReverseMapping(message *mDNS.Msg, response *mDNS.Msg, transport adapter.DNSTransport) {
|
||||||
} 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
|
|
||||||
}
|
|
||||||
if r.dnsReverseMapping != nil && len(message.Question) > 0 && response != nil && len(response.Answer) > 0 {
|
if r.dnsReverseMapping != nil && len(message.Question) > 0 && response != nil && len(response.Answer) > 0 {
|
||||||
if transport == nil || transport.Type() != C.DNSTypeFakeIP {
|
if transport == nil || transport.Type() != C.DNSTypeFakeIP {
|
||||||
for _, answer := range response.Answer {
|
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
|
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) {
|
func (r *Router) Lookup(ctx context.Context, domain string, options adapter.DNSQueryOptions) ([]netip.Addr, error) {
|
||||||
r.rulesAccess.RLock()
|
r.rulesAccess.RLock()
|
||||||
if r.closing {
|
if r.closing {
|
||||||
|
|||||||
+119
-23
@@ -8,12 +8,14 @@ import (
|
|||||||
"runtime"
|
"runtime"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"syscall"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sagernet/sing-box/adapter"
|
"github.com/sagernet/sing-box/adapter"
|
||||||
C "github.com/sagernet/sing-box/constant"
|
C "github.com/sagernet/sing-box/constant"
|
||||||
"github.com/sagernet/sing-box/dns"
|
"github.com/sagernet/sing-box/dns"
|
||||||
|
"github.com/sagernet/sing-box/dns/transport"
|
||||||
"github.com/sagernet/sing-box/log"
|
"github.com/sagernet/sing-box/log"
|
||||||
"github.com/sagernet/sing-box/option"
|
"github.com/sagernet/sing-box/option"
|
||||||
"github.com/sagernet/sing-tun"
|
"github.com/sagernet/sing-tun"
|
||||||
@@ -54,6 +56,8 @@ type Transport struct {
|
|||||||
updatedAt time.Time
|
updatedAt time.Time
|
||||||
lastError error
|
lastError error
|
||||||
servers []M.Socksaddr
|
servers []M.Socksaddr
|
||||||
|
serverTransports []adapter.DNSTransport
|
||||||
|
refreshing atomic.Bool
|
||||||
search []string
|
search []string
|
||||||
ndots int
|
ndots int
|
||||||
attempts int
|
attempts int
|
||||||
@@ -100,7 +104,7 @@ func (t *Transport) Start(stage adapter.StartStage) error {
|
|||||||
t.interfaceCallback = t.networkManager.InterfaceMonitor().RegisterCallback(t.interfaceUpdated)
|
t.interfaceCallback = t.networkManager.InterfaceMonitor().RegisterCallback(t.interfaceUpdated)
|
||||||
}
|
}
|
||||||
go func() {
|
go func() {
|
||||||
_, err := t.fetch()
|
err := t.fetch()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, errInterfaceIsCellular) && t.optional {
|
if errors.Is(err, errInterfaceIsCellular) && t.optional {
|
||||||
t.logger.Debug(E.Cause(errInterfaceIsCellular, "dhcp: fetch DNS servers"))
|
t.logger.Debug(E.Cause(errInterfaceIsCellular, "dhcp: fetch DNS servers"))
|
||||||
@@ -116,6 +120,9 @@ func (t *Transport) Close() error {
|
|||||||
if t.interfaceCallback != nil {
|
if t.interfaceCallback != nil {
|
||||||
t.networkManager.InterfaceMonitor().UnregisterCallback(t.interfaceCallback)
|
t.networkManager.InterfaceMonitor().UnregisterCallback(t.interfaceCallback)
|
||||||
}
|
}
|
||||||
|
t.transportLock.Lock()
|
||||||
|
defer t.transportLock.Unlock()
|
||||||
|
t.closeServerTransports()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -124,51 +131,122 @@ func (t *Transport) Reset() {
|
|||||||
t.updatedAt = time.Time{}
|
t.updatedAt = time.Time{}
|
||||||
t.lastError = nil
|
t.lastError = nil
|
||||||
t.servers = nil
|
t.servers = nil
|
||||||
|
t.closeServerTransports()
|
||||||
t.transportLock.Unlock()
|
t.transportLock.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
func (t *Transport) closeServerTransports() {
|
||||||
servers, err := t.fetch()
|
for _, serverTransport := range t.serverTransports {
|
||||||
if err != nil {
|
serverTransport.Close()
|
||||||
return nil, E.Cause(err, "dhcp: fetch DNS servers")
|
|
||||||
}
|
}
|
||||||
if len(servers) == 0 {
|
t.serverTransports = nil
|
||||||
return nil, E.New("dhcp: empty DNS servers from response")
|
|
||||||
}
|
|
||||||
return t.Exchange0(ctx, message, servers)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *Transport) Exchange0(ctx context.Context, message *mDNS.Msg, servers []M.Socksaddr) (*mDNS.Msg, error) {
|
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
||||||
return t.exchangeSearch(ctx, servers, message, dns.FqdnToDomain(message.Question[0].Name))
|
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 {
|
func (t *Transport) Fetch() []M.Socksaddr {
|
||||||
servers, _ := t.fetch()
|
|
||||||
return servers
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *Transport) fetch() ([]M.Socksaddr, error) {
|
|
||||||
t.transportLock.RLock()
|
t.transportLock.RLock()
|
||||||
updatedAt := t.updatedAt
|
updatedAt := t.updatedAt
|
||||||
lastError := t.lastError
|
lastError := t.lastError
|
||||||
servers := t.servers
|
servers := t.servers
|
||||||
t.transportLock.RUnlock()
|
t.transportLock.RUnlock()
|
||||||
if lastError != nil {
|
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 {
|
if time.Since(updatedAt) < C.DHCPTTL {
|
||||||
return servers, nil
|
return nil
|
||||||
}
|
}
|
||||||
t.transportLock.Lock()
|
t.transportLock.Lock()
|
||||||
defer t.transportLock.Unlock()
|
defer t.transportLock.Unlock()
|
||||||
if time.Since(t.updatedAt) < C.DHCPTTL {
|
if time.Since(t.updatedAt) < C.DHCPTTL {
|
||||||
return t.servers, nil
|
return nil
|
||||||
}
|
}
|
||||||
err := t.updateServers()
|
return t.updateServers()
|
||||||
if err != nil {
|
}
|
||||||
return servers, err
|
|
||||||
|
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) {
|
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) {
|
func (t *Transport) interfaceUpdated(defaultInterface *control.Interface, flags int) {
|
||||||
|
t.transportLock.Lock()
|
||||||
err := t.updateServers()
|
err := t.updateServers()
|
||||||
|
t.transportLock.Unlock()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, errInterfaceIsCellular) && t.optional {
|
if errors.Is(err, errInterfaceIsCellular) && t.optional {
|
||||||
t.logger.Debug(E.Cause(errInterfaceIsCellular, "dhcp: update DNS servers"))
|
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) {
|
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, ","), "]")
|
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
|
t.servers = serverAddrs
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,151 +2,52 @@ package dhcp
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
|
||||||
"math/rand"
|
|
||||||
"strings"
|
"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-box/dns/transport"
|
||||||
"github.com/sagernet/sing/common/buf"
|
|
||||||
E "github.com/sagernet/sing/common/exceptions"
|
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"
|
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)
|
names := t.nameList(domain)
|
||||||
if len(names) == 0 {
|
if len(names) == 0 {
|
||||||
return nil, E.New("dhcp: invalid domain: ", domain)
|
callback(nil, E.New("invalid domain: ", domain))
|
||||||
|
return
|
||||||
}
|
}
|
||||||
originalQuestion := message.Question[0]
|
nameExchangers := make([]transport.AsyncExchanger, 0, len(names))
|
||||||
var (
|
|
||||||
nameErrorResponse *mDNS.Msg
|
|
||||||
lastErr error
|
|
||||||
)
|
|
||||||
for _, fqdn := range names {
|
for _, fqdn := range names {
|
||||||
response, err := t.tryOneName(ctx, servers, fqdn, message)
|
nameExchangers = append(nameExchangers, t.newNameExchanger(message, fqdn, serverTransports))
|
||||||
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
|
|
||||||
}
|
}
|
||||||
if nameErrorResponse != nil {
|
if len(serverTransports) == 1 || !(question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) {
|
||||||
return nameErrorResponse, nil
|
transport.ExchangeSequential(ctx, nameExchangers, nil, callback)
|
||||||
}
|
} else {
|
||||||
return nil, lastErr
|
transport.ExchangeRace(ctx, nameExchangers, callback)
|
||||||
}
|
|
||||||
|
|
||||||
// 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
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *Transport) tryOneName(ctx context.Context, servers []M.Socksaddr, fqdn string, message *mDNS.Msg) (*mDNS.Msg, error) {
|
func (t *Transport) newNameExchanger(message *mDNS.Msg, fqdn string, serverTransports []adapter.DNSTransport) transport.AsyncExchanger {
|
||||||
sLen := len(servers)
|
attemptExchangers := make([]transport.AsyncExchanger, 0, t.attempts*len(serverTransports))
|
||||||
var lastErr error
|
for range t.attempts {
|
||||||
for i := 0; i < t.attempts; i++ {
|
for _, serverTransport := range serverTransports {
|
||||||
for j := range sLen {
|
attemptExchangers = append(attemptExchangers, func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
|
||||||
server := servers[j]
|
serverTransport.ExchangeAsync(ctx, transport.NewFanOutRequest(message, fqdn, true), callback)
|
||||||
question := message.Question[0]
|
})
|
||||||
question.Name = fqdn
|
}
|
||||||
response, err := t.exchangeOne(ctx, server, question)
|
}
|
||||||
|
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 {
|
if err != nil {
|
||||||
lastErr = err
|
err = E.Cause(err, fqdn)
|
||||||
continue
|
|
||||||
}
|
}
|
||||||
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 {
|
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
|
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 {
|
func (t *Transport) Store() adapter.FakeIPStore {
|
||||||
return t.store
|
return t.store
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -104,3 +104,7 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
|
|||||||
Question: []mDNS.Question{question},
|
Question: []mDNS.Question{question},
|
||||||
}, nil
|
}, 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
|
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) {
|
func (t *HTTPSTransport) exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
||||||
exMessage := *message
|
exMessage := *message
|
||||||
exMessage.Id = 0
|
exMessage.Id = 0
|
||||||
|
|||||||
@@ -2,6 +2,8 @@ package local
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/sagernet/sing-box/adapter"
|
"github.com/sagernet/sing-box/adapter"
|
||||||
C "github.com/sagernet/sing-box/constant"
|
C "github.com/sagernet/sing-box/constant"
|
||||||
@@ -31,15 +33,19 @@ var (
|
|||||||
|
|
||||||
type Transport struct {
|
type Transport struct {
|
||||||
dns.TransportAdapter
|
dns.TransportAdapter
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
logger logger.ContextLogger
|
logger logger.ContextLogger
|
||||||
hosts *hosts.File
|
hosts *hosts.File
|
||||||
dialer N.Dialer
|
dialer N.Dialer
|
||||||
preferGo bool
|
preferGo bool
|
||||||
fallback bool
|
fallback bool
|
||||||
resolved ResolvedResolver
|
resolved ResolvedResolver
|
||||||
mdnsTransport adapter.DNSTransport
|
mdnsTransport adapter.DNSTransport
|
||||||
dhcpTransport dhcpTransport
|
dhcpTransport dhcpTransport
|
||||||
|
system systemResolver
|
||||||
|
serverSet atomic.Pointer[localServerSet]
|
||||||
|
serverSetAccess sync.Mutex
|
||||||
|
|
||||||
neighborResolver adapter.NeighborResolver
|
neighborResolver adapter.NeighborResolver
|
||||||
neighborSuffixes []string
|
neighborSuffixes []string
|
||||||
}
|
}
|
||||||
@@ -47,7 +53,6 @@ type Transport struct {
|
|||||||
type dhcpTransport interface {
|
type dhcpTransport interface {
|
||||||
adapter.DNSTransport
|
adapter.DNSTransport
|
||||||
Fetch() []M.Socksaddr
|
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) {
|
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 {
|
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)
|
return common.Close(t.resolved, t.dhcpTransport, t.mdnsTransport)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *Transport) Reset() {
|
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 {
|
if t.dhcpTransport != nil {
|
||||||
t.dhcpTransport.Reset()
|
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) {
|
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]
|
question := message.Question[0]
|
||||||
if t.hosts != nil && (question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) {
|
if t.hosts != nil && (question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) {
|
||||||
addresses := t.hosts.Lookup(dns.FqdnToDomain(question.Name))
|
addresses := t.hosts.Lookup(dns.FqdnToDomain(question.Name))
|
||||||
if len(addresses) > 0 {
|
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)
|
response := t.lookupNeighbor(message)
|
||||||
if response != nil {
|
if response != nil {
|
||||||
return response, nil
|
callback(response, nil)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
if mdns.IsLocalDomain(question.Name) {
|
if mdns.IsLocalDomain(question.Name) {
|
||||||
if C.IsDarwin {
|
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 {
|
if t.resolved != nil {
|
||||||
return t.resolved.Exchange(ctx, message)
|
t.resolved.ExchangeAsync(ctx, message, callback)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
if t.dhcpTransport != nil {
|
if t.dhcpTransport != nil {
|
||||||
servers := t.dhcpTransport.Fetch()
|
servers := t.dhcpTransport.Fetch()
|
||||||
if len(servers) > 0 {
|
if len(servers) > 0 {
|
||||||
return t.dhcpTransport.Exchange0(ctx, message, servers)
|
t.dhcpTransport.ExchangeAsync(ctx, message, callback)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if t.fallback {
|
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"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
"os"
|
"os"
|
||||||
|
"sync"
|
||||||
|
|
||||||
"github.com/sagernet/sing-box/dns"
|
"github.com/sagernet/sing-box/dns"
|
||||||
|
dnsTransport "github.com/sagernet/sing-box/dns/transport"
|
||||||
E "github.com/sagernet/sing/common/exceptions"
|
E "github.com/sagernet/sing/common/exceptions"
|
||||||
|
|
||||||
mDNS "github.com/miekg/dns"
|
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]
|
question := message.Question[0]
|
||||||
response, err := darwinLookupSystemDNS(ctx, question.Name, question.Qtype, question.Qclass)
|
t.system.exchangeAsync(ctx, question.Name, question.Qtype, question.Qclass, func(response *mDNS.Msg, err error) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
var rcodeError dns.RcodeError
|
var rcodeError dns.RcodeError
|
||||||
if errors.As(err, &rcodeError) {
|
if errors.As(err, &rcodeError) {
|
||||||
return dns.FixedResponseStatus(message, int(rcodeError)), nil
|
callback(dns.FixedResponseStatus(message, int(rcodeError)), nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
callback(nil, err)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
return nil, err
|
response.Id = message.Id
|
||||||
}
|
response.Response = true
|
||||||
response.Id = message.Id
|
response.RecursionAvailable = true
|
||||||
response.Response = true
|
callback(response, nil)
|
||||||
response.RecursionAvailable = true
|
})
|
||||||
return response, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// The mDNSResponder daemon speaks an undocumented binary protocol over a
|
// The mDNSResponder daemon speaks an undocumented binary protocol over a
|
||||||
// AF_UNIX SOCK_STREAM socket. The framing below is taken from the client
|
// 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
|
// stub of Apple's open-source mDNSResponder (mDNSShared/dnssd_ipc.h,
|
||||||
// dnssd_clientstub.c). All multi-byte fields are big-endian; for a one-shot
|
// dnssd_clientstub.c and uds_daemon.c). All multi-byte fields are
|
||||||
// query on a fresh, non-shared connection the request and every reply travel
|
// big-endian. A connection opened with connection_request acts as a shared
|
||||||
// over the single connected stream (no SCM_RIGHTS, no return socket).
|
// 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 (
|
const (
|
||||||
mdnsResponderSocketPath = "/var/run/mDNSResponder"
|
mdnsResponderSocketPath = "/var/run/mDNSResponder"
|
||||||
mdnsResponderSocketEnv = "DNSSD_UDS_PATH"
|
mdnsResponderSocketEnv = "DNSSD_UDS_PATH"
|
||||||
mdnsResponderVersion = 1
|
mdnsResponderVersion = 1
|
||||||
mdnsResponderHeaderLength = 28
|
mdnsResponderHeaderLength = 28
|
||||||
mdnsResponderQueryRequest = 8 // query_request
|
mdnsResponderConnectionRequest = 1 // connection_request
|
||||||
mdnsResponderQueryReply = 68 // query_reply_op
|
mdnsResponderQueryRequest = 8 // query_request
|
||||||
|
mdnsResponderCancelRequest = 63 // cancel_request
|
||||||
|
mdnsResponderQueryReply = 68 // query_reply_op
|
||||||
|
mdnsResponderAsyncErrorReply = 73 // async_error_op
|
||||||
|
|
||||||
mdnsResponderFlagMoreComing = 0x1
|
mdnsResponderFlagMoreComing = 0x1
|
||||||
mdnsResponderFlagAdd = 0x2
|
mdnsResponderFlagAdd = 0x2
|
||||||
mdnsResponderFlagReturnIntermediates = 0x1000
|
mdnsResponderFlagReturnIntermediates = 0x1000
|
||||||
|
mdnsResponderFlagShareConnection = 0x4000
|
||||||
mdnsResponderFlagTimeout = 0x10000
|
mdnsResponderFlagTimeout = 0x10000
|
||||||
|
|
||||||
|
mdnsResponderIPCFlagNoErrorSocket = 0x4 // IPC_FLAGS_NOERRSD
|
||||||
|
|
||||||
mdnsResponderErrNoError = 0
|
mdnsResponderErrNoError = 0
|
||||||
mdnsResponderErrNoSuchName = -65538
|
mdnsResponderErrNoSuchName = -65538
|
||||||
mdnsResponderErrNoSuchRecord = -65554
|
mdnsResponderErrNoSuchRecord = -65554
|
||||||
mdnsResponderErrTimeout = -65568
|
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)
|
socketPath := cmp.Or(os.Getenv(mdnsResponderSocketEnv), mdnsResponderSocketPath)
|
||||||
var dialer net.Dialer
|
var dialer net.Dialer
|
||||||
conn, err := dialer.DialContext(ctx, "unix", socketPath)
|
conn, err := dialer.DialContext(ctx, "unix", socketPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, E.Cause(err, "connect mDNSResponder")
|
return nil, E.Cause(err, "connect mDNSResponder")
|
||||||
}
|
}
|
||||||
defer conn.Close()
|
|
||||||
stopCancel := context.AfterFunc(ctx, func() {
|
stopCancel := context.AfterFunc(ctx, func() {
|
||||||
conn.Close()
|
conn.Close()
|
||||||
})
|
})
|
||||||
defer stopCancel()
|
err = writeConnectionRequest(conn)
|
||||||
|
stopCancel()
|
||||||
_, err = conn.Write(buildQueryRequest(name, qtype, qclass))
|
|
||||||
if err != nil {
|
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
|
var status [4]byte
|
||||||
_, err = io.ReadFull(conn, status[:])
|
_, err = io.ReadFull(conn, status[:])
|
||||||
if err != nil {
|
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[:]))
|
statusCode := int32(binary.BigEndian.Uint32(status[:]))
|
||||||
if statusCode != mdnsResponderErrNoError {
|
if statusCode != mdnsResponderErrNoError {
|
||||||
return nil, darwinResolverError(name, statusCode)
|
return E.New("mDNSResponder connection request failed: error ", statusCode)
|
||||||
}
|
}
|
||||||
|
return nil
|
||||||
return readQueryResponse(ctx, conn, name, qtype, qclass)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func readQueryResponse(ctx context.Context, conn net.Conn, name string, qtype, qclass uint16) (*mDNS.Msg, error) {
|
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 {
|
||||||
var answers []mDNS.RR
|
r.queryAccess.Lock()
|
||||||
var hasFinalAnswer bool
|
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 {
|
for {
|
||||||
reply, replyErr := readReply(conn)
|
operation, clientContext, data, err := readResponderReply(conn)
|
||||||
if replyErr != nil {
|
if err != nil {
|
||||||
return nil, contextError(ctx, E.Cause(replyErr, "read mDNSResponder reply"))
|
r.connection.Invalidate(conn, err)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
if reply.errorCode != mdnsResponderErrNoError {
|
switch operation {
|
||||||
if len(answers) == 0 {
|
case mdnsResponderQueryReply:
|
||||||
return nil, darwinResolverError(name, reply.errorCode)
|
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)
|
// On a shared connection MoreComing applies collectively to all operations
|
||||||
if record.Header().Rrtype == qtype {
|
// (dns_sd.h "Collective kDNSServiceFlagsMoreComing flag"): the daemon sets it
|
||||||
hasFinalAnswer = true
|
// 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 pending.hasFinalAnswer && reply.rrtype == pending.qtype {
|
||||||
if reply.flags&mdnsResponderFlagMoreComing != 0 {
|
pending.ready = true
|
||||||
continue
|
}
|
||||||
}
|
|
||||||
if hasFinalAnswer && reply.rrtype == qtype {
|
|
||||||
break
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if reply.flags&mdnsResponderFlagMoreComing == 0 {
|
||||||
response := new(mDNS.Msg)
|
completions = r.collectReadyLocked(completions)
|
||||||
response.Question = []mDNS.Question{{Name: mDNS.Fqdn(name), Qtype: qtype, Qclass: qclass}}
|
}
|
||||||
response.Answer = answers
|
r.queryAccess.Unlock()
|
||||||
return response, nil
|
for _, completion := range completions {
|
||||||
|
r.finish(completion.pending, completion.err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func buildQueryRequest(name string, qtype, qclass uint16) []byte {
|
func (r *systemResolver) completeQueryError(queryId uint64, flags uint32, errorCode int32) {
|
||||||
payload := make([]byte, 0, 8+len(name)+1+4)
|
var completions []systemCompletion
|
||||||
payload = binary.BigEndian.AppendUint32(payload, mdnsResponderFlagReturnIntermediates|mdnsResponderFlagTimeout)
|
r.queryAccess.Lock()
|
||||||
payload = binary.BigEndian.AppendUint32(payload, 0) // interfaceIndex
|
pending, loaded := r.queries[queryId]
|
||||||
payload = append(payload, name...)
|
if loaded {
|
||||||
payload = append(payload, 0) // C string terminator
|
delete(r.queries, queryId)
|
||||||
payload = binary.BigEndian.AppendUint16(payload, qtype)
|
completions = append(completions, systemCompletion{pending: pending, err: darwinResolverError(pending.name, errorCode)})
|
||||||
payload = binary.BigEndian.AppendUint16(payload, qclass)
|
}
|
||||||
|
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))
|
func (r *systemResolver) collectReadyLocked(completions []systemCompletion) []systemCompletion {
|
||||||
binary.BigEndian.PutUint32(message[0:], mdnsResponderVersion)
|
for queryId, pending := range r.queries {
|
||||||
binary.BigEndian.PutUint32(message[4:], uint32(len(payload)))
|
if pending.ready {
|
||||||
binary.BigEndian.PutUint32(message[8:], 0) // ipc_flags
|
delete(r.queries, queryId)
|
||||||
binary.BigEndian.PutUint32(message[12:], mdnsResponderQueryRequest)
|
completions = append(completions, systemCompletion{pending: pending})
|
||||||
// message[16:24] client_context and message[24:28] reg_index stay zero.
|
}
|
||||||
return append(message, payload...)
|
}
|
||||||
|
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 {
|
type mdnsResponderReply struct {
|
||||||
@@ -154,24 +406,8 @@ type mdnsResponderReply struct {
|
|||||||
rdata []byte
|
rdata []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func readReply(conn net.Conn) (mdnsResponderReply, error) {
|
func parseResponderReply(data []byte) (mdnsResponderReply, error) {
|
||||||
var reply mdnsResponderReply
|
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}
|
reader := replyReader{data: data}
|
||||||
reply.flags = reader.uint32()
|
reply.flags = reader.uint32()
|
||||||
reader.uint32() // interfaceIndex
|
reader.uint32() // interfaceIndex
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"net"
|
"net"
|
||||||
"os"
|
"os"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -26,9 +27,25 @@ func requireMDNSResponder(t *testing.T) {
|
|||||||
conn.Close()
|
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) {
|
func TestSystemExchangeLoopback(t *testing.T) {
|
||||||
requireMDNSResponder(t)
|
requireMDNSResponder(t)
|
||||||
transport := &Transport{}
|
transport := &Transport{}
|
||||||
|
defer transport.system.close()
|
||||||
for _, testCase := range []struct {
|
for _, testCase := range []struct {
|
||||||
qtype uint16
|
qtype uint16
|
||||||
expected net.IP
|
expected net.IP
|
||||||
@@ -39,7 +56,7 @@ func TestSystemExchangeLoopback(t *testing.T) {
|
|||||||
message := new(mDNS.Msg)
|
message := new(mDNS.Msg)
|
||||||
message.SetQuestion("localhost.", testCase.qtype)
|
message.SetQuestion("localhost.", testCase.qtype)
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
response, err := transport.systemExchange(ctx, message)
|
response, err := systemExchangeForTest(ctx, transport, message)
|
||||||
cancel()
|
cancel()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("%s localhost: %v", mDNS.TypeToString[testCase.qtype], err)
|
t.Fatalf("%s localhost: %v", mDNS.TypeToString[testCase.qtype], err)
|
||||||
@@ -67,13 +84,15 @@ func TestSystemExchangeLoopback(t *testing.T) {
|
|||||||
|
|
||||||
func TestSystemExchangeNoData(t *testing.T) {
|
func TestSystemExchangeNoData(t *testing.T) {
|
||||||
requireMDNSResponder(t)
|
requireMDNSResponder(t)
|
||||||
|
transport := &Transport{}
|
||||||
|
defer transport.system.close()
|
||||||
message := new(mDNS.Msg)
|
message := new(mDNS.Msg)
|
||||||
// localhost has no MX record, so the daemon reports NoSuchRecord, which must
|
// localhost has no MX record, so the daemon reports NoSuchRecord, which must
|
||||||
// surface as an empty NOERROR response rather than an error.
|
// surface as an empty NOERROR response rather than an error.
|
||||||
message.SetQuestion("localhost.", mDNS.TypeMX)
|
message.SetQuestion("localhost.", mDNS.TypeMX)
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
response, err := (&Transport{}).systemExchange(ctx, message)
|
response, err := systemExchangeForTest(ctx, transport, message)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("MX localhost: %v", err)
|
t.Fatalf("MX localhost: %v", err)
|
||||||
}
|
}
|
||||||
@@ -87,12 +106,14 @@ func TestSystemExchangeNoData(t *testing.T) {
|
|||||||
|
|
||||||
func TestSystemExchangeCancel(t *testing.T) {
|
func TestSystemExchangeCancel(t *testing.T) {
|
||||||
requireMDNSResponder(t)
|
requireMDNSResponder(t)
|
||||||
|
transport := &Transport{}
|
||||||
|
defer transport.system.close()
|
||||||
message := new(mDNS.Msg)
|
message := new(mDNS.Msg)
|
||||||
message.SetQuestion("localhost.", mDNS.TypeA)
|
message.SetQuestion("localhost.", mDNS.TypeA)
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
cancel()
|
cancel()
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
_, err := (&Transport{}).systemExchange(ctx, message)
|
_, err := systemExchangeForTest(ctx, transport, message)
|
||||||
elapsed := time.Since(start)
|
elapsed := time.Since(start)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Fatal("expected error for cancelled context")
|
t.Fatal("expected error for cancelled context")
|
||||||
@@ -101,3 +122,44 @@ func TestSystemExchangeCancel(t *testing.T) {
|
|||||||
t.Fatalf("cancellation too slow: %s", elapsed)
|
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"
|
mDNS "github.com/miekg/dns"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (t *Transport) systemExchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
type systemResolver struct{}
|
||||||
return nil, os.ErrInvalid
|
|
||||||
|
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
|
Start() error
|
||||||
Close() error
|
Close() error
|
||||||
Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, 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)
|
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() {
|
func (t *DBusResolvedResolver) loopUpdateStatus() {
|
||||||
signalChan := make(chan *dbus.Signal, 1)
|
signalChan := make(chan *dbus.Signal, 1)
|
||||||
t.systemBus.Signal(signalChan)
|
t.systemBus.Signal(signalChan)
|
||||||
|
|||||||
@@ -2,182 +2,117 @@ package local
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"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-box/dns/transport"
|
||||||
"github.com/sagernet/sing/common/buf"
|
|
||||||
E "github.com/sagernet/sing/common/exceptions"
|
E "github.com/sagernet/sing/common/exceptions"
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
M "github.com/sagernet/sing/common/metadata"
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
|
|
||||||
mDNS "github.com/miekg/dns"
|
mDNS "github.com/miekg/dns"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (t *Transport) exchange(ctx context.Context, message *mDNS.Msg, domain string) (*mDNS.Msg, error) {
|
type localServerSet struct {
|
||||||
systemConfig := getSystemDNSConfig(t.ctx)
|
config *dnsConfig
|
||||||
if systemConfig.singleRequest || !(message.Question[0].Qtype == mDNS.TypeA || message.Question[0].Qtype == mDNS.TypeAAAA) {
|
transports []adapter.DNSTransport
|
||||||
return t.exchangeSingleRequest(ctx, systemConfig, message, domain)
|
}
|
||||||
} else {
|
|
||||||
return t.exchangeParallel(ctx, systemConfig, message, domain)
|
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) {
|
func (t *Transport) serverSetFor(systemConfig *dnsConfig) (*localServerSet, error) {
|
||||||
var lastErr error
|
serverSet := t.serverSet.Load()
|
||||||
for _, fqdn := range systemConfig.nameList(domain) {
|
if serverSet != nil && serverSet.config == systemConfig {
|
||||||
response, err := t.tryOneName(ctx, systemConfig, fqdn, message)
|
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 {
|
if err != nil {
|
||||||
lastErr = err
|
for _, startedTransport := range transports {
|
||||||
continue
|
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) {
|
func (t *Transport) exchangeAsync(ctx context.Context, message *mDNS.Msg, domain string, callback func(response *mDNS.Msg, err error)) {
|
||||||
returned := make(chan struct{})
|
systemConfig := getSystemDNSConfig(t.ctx)
|
||||||
defer close(returned)
|
serverSet, err := t.serverSetFor(systemConfig)
|
||||||
type queryResult struct {
|
if err != nil {
|
||||||
response *mDNS.Msg
|
callback(nil, err)
|
||||||
err error
|
return
|
||||||
}
|
}
|
||||||
results := make(chan queryResult)
|
names := systemConfig.nameList(domain)
|
||||||
startRacer := func(ctx context.Context, fqdn string) {
|
if len(names) == 0 {
|
||||||
response, err := t.tryOneName(ctx, systemConfig, fqdn, message)
|
callback(nil, E.New("invalid domain: ", domain))
|
||||||
select {
|
return
|
||||||
case results <- queryResult{response, err}:
|
|
||||||
case <-returned:
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
queryCtx, queryCancel := context.WithCancel(ctx)
|
nameExchangers := make([]transport.AsyncExchanger, 0, len(names))
|
||||||
defer queryCancel()
|
for _, fqdn := range names {
|
||||||
var nameCount int
|
nameExchangers = append(nameExchangers, newNameExchanger(systemConfig, serverSet, message, fqdn))
|
||||||
for _, fqdn := range systemConfig.nameList(domain) {
|
|
||||||
nameCount++
|
|
||||||
go startRacer(queryCtx, fqdn)
|
|
||||||
}
|
}
|
||||||
var errors []error
|
question := message.Question[0]
|
||||||
for {
|
if systemConfig.singleRequest || !(question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) {
|
||||||
select {
|
transport.ExchangeSequential(ctx, nameExchangers, nil, callback)
|
||||||
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)
|
|
||||||
} else {
|
} 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) {
|
func newNameExchanger(systemConfig *dnsConfig, serverSet *localServerSet, message *mDNS.Msg, fqdn string) transport.AsyncExchanger {
|
||||||
conn, err := t.dialer.DialContext(ctx, N.NetworkUDP, server)
|
serverOffset := systemConfig.serverOffset()
|
||||||
if err != nil {
|
serverCount := uint32(len(serverSet.transports))
|
||||||
return nil, err
|
attemptExchangers := make([]transport.AsyncExchanger, 0, systemConfig.attempts*int(serverCount))
|
||||||
}
|
for i := 0; i < systemConfig.attempts; i++ {
|
||||||
defer conn.Close()
|
for j := range serverCount {
|
||||||
if deadline, loaded := ctx.Deadline(); loaded && !deadline.IsZero() {
|
serverTransport := serverSet.transports[(serverOffset+j)%serverCount]
|
||||||
newDeadline := time.Now().Add(timeout)
|
attemptExchangers = append(attemptExchangers, func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
|
||||||
if deadline.After(newDeadline) {
|
attemptCtx, cancel := context.WithTimeout(ctx, systemConfig.timeout)
|
||||||
deadline = newDeadline
|
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)
|
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
|
||||||
defer buf.Put(buffer)
|
transport.ExchangeSequential(ctx, attemptExchangers, nil, func(response *mDNS.Msg, err error) {
|
||||||
rawMessage, err := request.PackBuffer(buffer)
|
if err != nil {
|
||||||
if err != nil {
|
err = E.Cause(err, fqdn)
|
||||||
return nil, E.Cause(err, "pack request")
|
}
|
||||||
|
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")
|
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 {
|
type exchangeResult struct {
|
||||||
response *mDNS.Msg
|
response *mDNS.Msg
|
||||||
err error
|
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
|
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
|
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) {
|
func (t *Transport) exchange(ctx context.Context, message *mDNS.Msg, conn *quic.Conn) (*mDNS.Msg, error) {
|
||||||
stream, err := conn.OpenStreamSync(ctx)
|
stream, err := conn.OpenStreamSync(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
+34
-23
@@ -15,7 +15,6 @@ import (
|
|||||||
"github.com/sagernet/sing-box/option"
|
"github.com/sagernet/sing-box/option"
|
||||||
"github.com/sagernet/sing/common"
|
"github.com/sagernet/sing/common"
|
||||||
"github.com/sagernet/sing/common/buf"
|
"github.com/sagernet/sing/common/buf"
|
||||||
"github.com/sagernet/sing/common/bufio/deadline"
|
|
||||||
E "github.com/sagernet/sing/common/exceptions"
|
E "github.com/sagernet/sing/common/exceptions"
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
M "github.com/sagernet/sing/common/metadata"
|
||||||
N "github.com/sagernet/sing/common/network"
|
N "github.com/sagernet/sing/common/network"
|
||||||
@@ -31,8 +30,9 @@ func RegisterTCP(registry *dns.TransportRegistry) {
|
|||||||
|
|
||||||
type TCPTransport struct {
|
type TCPTransport struct {
|
||||||
dns.TransportAdapter
|
dns.TransportAdapter
|
||||||
dialer N.Dialer
|
dialer N.Dialer
|
||||||
serverAddr M.Socksaddr
|
serverAddr M.Socksaddr
|
||||||
|
multiplexer *queryMultiplexer
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewTCP(ctx context.Context, logger log.ContextLogger, tag string, options option.RemoteDNSServerOptions) (adapter.DNSTransport, error) {
|
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() {
|
if !serverAddr.IsValid() {
|
||||||
return nil, E.New("invalid server address: ", serverAddr)
|
return nil, E.New("invalid server address: ", serverAddr)
|
||||||
}
|
}
|
||||||
return &TCPTransport{
|
return NewTCPRaw(dns.NewTransportAdapterWithRemoteOptions(C.DNSTypeTCP, tag, options), transportDialer, serverAddr), nil
|
||||||
TransportAdapter: dns.NewTransportAdapterWithRemoteOptions(C.DNSTypeTCP, tag, options),
|
}
|
||||||
dialer: transportDialer,
|
|
||||||
|
func NewTCPRaw(adapter dns.TransportAdapter, dialer N.Dialer, serverAddr M.Socksaddr) *TCPTransport {
|
||||||
|
t := &TCPTransport{
|
||||||
|
TransportAdapter: adapter,
|
||||||
|
dialer: dialer,
|
||||||
serverAddr: serverAddr,
|
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 {
|
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 {
|
func (t *TCPTransport) Close() error {
|
||||||
return nil
|
return t.multiplexer.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TCPTransport) Reset() {
|
func (t *TCPTransport) Reset() {
|
||||||
|
t.multiplexer.Reset()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TCPTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
func (t *TCPTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
||||||
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, t.serverAddr)
|
return t.multiplexer.Exchange(ctx, message)
|
||||||
if err != nil {
|
}
|
||||||
return nil, E.Cause(err, "dial TCP connection")
|
|
||||||
}
|
func (t *TCPTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
|
||||||
defer conn.Close()
|
t.multiplexer.ExchangeAsync(ctx, message, callback)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func setConnDeadline(ctx context.Context, conn net.Conn, needClose bool) func() {
|
func setConnDeadline(ctx context.Context, conn net.Conn, needClose bool) func() {
|
||||||
|
|||||||
+25
-67
@@ -2,6 +2,7 @@ package transport
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"net"
|
||||||
|
|
||||||
"github.com/sagernet/sing-box/adapter"
|
"github.com/sagernet/sing-box/adapter"
|
||||||
"github.com/sagernet/sing-box/common/dialer"
|
"github.com/sagernet/sing-box/common/dialer"
|
||||||
@@ -11,7 +12,6 @@ import (
|
|||||||
"github.com/sagernet/sing-box/log"
|
"github.com/sagernet/sing-box/log"
|
||||||
"github.com/sagernet/sing-box/option"
|
"github.com/sagernet/sing-box/option"
|
||||||
"github.com/sagernet/sing/common"
|
"github.com/sagernet/sing/common"
|
||||||
"github.com/sagernet/sing/common/bufio/deadline"
|
|
||||||
E "github.com/sagernet/sing/common/exceptions"
|
E "github.com/sagernet/sing/common/exceptions"
|
||||||
"github.com/sagernet/sing/common/logger"
|
"github.com/sagernet/sing/common/logger"
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
M "github.com/sagernet/sing/common/metadata"
|
||||||
@@ -22,26 +22,16 @@ import (
|
|||||||
|
|
||||||
var _ adapter.DNSTransport = (*TLSTransport)(nil)
|
var _ adapter.DNSTransport = (*TLSTransport)(nil)
|
||||||
|
|
||||||
const tlsDNSMaxInflight = 8
|
|
||||||
|
|
||||||
func RegisterTLS(registry *dns.TransportRegistry) {
|
func RegisterTLS(registry *dns.TransportRegistry) {
|
||||||
dns.RegisterTransport[option.RemoteTLSDNSServerOptions](registry, C.DNSTypeTLS, NewTLS)
|
dns.RegisterTransport[option.RemoteTLSDNSServerOptions](registry, C.DNSTypeTLS, NewTLS)
|
||||||
}
|
}
|
||||||
|
|
||||||
type TLSTransport struct {
|
type TLSTransport struct {
|
||||||
dns.TransportAdapter
|
dns.TransportAdapter
|
||||||
logger logger.ContextLogger
|
logger logger.ContextLogger
|
||||||
|
|
||||||
dialer tls.Dialer
|
dialer tls.Dialer
|
||||||
serverAddr M.Socksaddr
|
serverAddr M.Socksaddr
|
||||||
tlsConfig tls.Config
|
multiplexer *queryMultiplexer
|
||||||
connections *ConnPool[*tlsDNSConn]
|
|
||||||
}
|
|
||||||
|
|
||||||
type tlsDNSConn struct {
|
|
||||||
tls.Conn
|
|
||||||
queryId uint16
|
|
||||||
needDeadlineClose bool
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewTLS(ctx context.Context, logger log.ContextLogger, tag string, options option.RemoteTLSDNSServerOptions) (adapter.DNSTransport, error) {
|
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 {
|
func NewTLSRaw(logger logger.ContextLogger, adapter dns.TransportAdapter, dialer N.Dialer, serverAddr M.Socksaddr, tlsConfig tls.Config) *TLSTransport {
|
||||||
return &TLSTransport{
|
t := &TLSTransport{
|
||||||
TransportAdapter: adapter,
|
TransportAdapter: adapter,
|
||||||
logger: logger,
|
logger: logger,
|
||||||
dialer: tls.NewDialer(dialer, tlsConfig),
|
dialer: tls.NewDialer(dialer, tlsConfig),
|
||||||
serverAddr: serverAddr,
|
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 {
|
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 {
|
func (t *TLSTransport) Close() error {
|
||||||
return t.connections.Close()
|
return t.multiplexer.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TLSTransport) Reset() {
|
func (t *TLSTransport) Reset() {
|
||||||
t.connections.Reset()
|
t.multiplexer.Reset()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TLSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
func (t *TLSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
||||||
var lastErr error
|
return t.multiplexer.Exchange(ctx, message)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TLSTransport) exchange(ctx context.Context, message *mDNS.Msg, conn *tlsDNSConn) (*mDNS.Msg, error) {
|
func (t *TLSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
|
||||||
defer setConnDeadline(ctx, conn, conn.needDeadlineClose)()
|
t.multiplexer.ExchangeAsync(ctx, message, callback)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|||||||
+78
-156
@@ -3,7 +3,6 @@ package transport
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"net"
|
"net"
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/sagernet/sing-box/adapter"
|
"github.com/sagernet/sing-box/adapter"
|
||||||
@@ -12,6 +11,7 @@ import (
|
|||||||
"github.com/sagernet/sing-box/dns"
|
"github.com/sagernet/sing-box/dns"
|
||||||
"github.com/sagernet/sing-box/log"
|
"github.com/sagernet/sing-box/log"
|
||||||
"github.com/sagernet/sing-box/option"
|
"github.com/sagernet/sing-box/option"
|
||||||
|
"github.com/sagernet/sing/common"
|
||||||
"github.com/sagernet/sing/common/buf"
|
"github.com/sagernet/sing/common/buf"
|
||||||
"github.com/sagernet/sing/common/bufio/deadline"
|
"github.com/sagernet/sing/common/bufio/deadline"
|
||||||
E "github.com/sagernet/sing/common/exceptions"
|
E "github.com/sagernet/sing/common/exceptions"
|
||||||
@@ -36,17 +36,7 @@ type UDPTransport struct {
|
|||||||
serverAddr M.Socksaddr
|
serverAddr M.Socksaddr
|
||||||
udpSize atomic.Int32
|
udpSize atomic.Int32
|
||||||
|
|
||||||
connection *ConnPool[net.Conn]
|
multiplexer *queryMultiplexer
|
||||||
|
|
||||||
callbackAccess sync.RWMutex
|
|
||||||
queryId uint16
|
|
||||||
callbacks map[uint16]*udpCallback
|
|
||||||
}
|
|
||||||
|
|
||||||
type udpCallback struct {
|
|
||||||
access sync.Mutex
|
|
||||||
response *mDNS.Msg
|
|
||||||
done chan struct{}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewUDP(ctx context.Context, logger log.ContextLogger, tag string, options option.RemoteDNSServerOptions) (adapter.DNSTransport, error) {
|
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,
|
logger: logger,
|
||||||
dialer: dialerInstance,
|
dialer: dialerInstance,
|
||||||
serverAddr: serverAddr,
|
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.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
|
return t
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -93,28 +84,16 @@ func (t *UDPTransport) Start(stage adapter.StartStage) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *UDPTransport) Close() error {
|
func (t *UDPTransport) Close() error {
|
||||||
return t.connection.Close()
|
return t.multiplexer.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *UDPTransport) Reset() {
|
func (t *UDPTransport) Reset() {
|
||||||
t.connection.Reset()
|
t.multiplexer.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")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *UDPTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
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 {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -125,6 +104,67 @@ func (t *UDPTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.M
|
|||||||
return response, nil
|
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) {
|
func (t *UDPTransport) exchangeTCP(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
||||||
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, t.serverAddr)
|
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, t.serverAddr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -142,121 +182,3 @@ func (t *UDPTransport) exchangeTCP(ctx context.Context, message *mDNS.Msg) (*mDN
|
|||||||
}
|
}
|
||||||
return response, nil
|
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()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -105,6 +105,12 @@ func (p *platformTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*m
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (p *platformTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
|
||||||
|
go func() {
|
||||||
|
callback(p.Exchange(ctx, message))
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
type Func interface {
|
type Func interface {
|
||||||
Invoke() error
|
Invoke() error
|
||||||
}
|
}
|
||||||
|
|||||||
+38
-40
@@ -40,28 +40,30 @@ func HandleStreamDNSRequest(ctx context.Context, router adapter.DNSRouter, conn
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
metadataInQuery := metadata
|
metadataInQuery := metadata
|
||||||
go func() error {
|
router.ExchangeAsync(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{}, func(response *mDNS.Msg, err error) {
|
||||||
response, err := router.Exchange(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{})
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
conn.Close()
|
conn.Close()
|
||||||
return err
|
return
|
||||||
}
|
}
|
||||||
responseLength := response.Len()
|
go writeStreamResponse(conn, response)
|
||||||
responseBuffer := buf.NewSize(3 + responseLength)
|
})
|
||||||
defer responseBuffer.Release()
|
|
||||||
responseBuffer.Resize(2, 0)
|
|
||||||
n, err := response.PackBuffer(responseBuffer.FreeBytes())
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
responseBuffer.Truncate(len(n))
|
|
||||||
binary.BigEndian.PutUint16(responseBuffer.ExtendHeader(2), uint16(len(n)))
|
|
||||||
_, err = conn.Write(responseBuffer.Bytes())
|
|
||||||
return err
|
|
||||||
}()
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func writeStreamResponse(conn net.Conn, response *mDNS.Msg) {
|
||||||
|
responseLength := response.Len()
|
||||||
|
responseBuffer := buf.NewSize(3 + responseLength)
|
||||||
|
defer responseBuffer.Release()
|
||||||
|
responseBuffer.Resize(2, 0)
|
||||||
|
n, err := response.PackBuffer(responseBuffer.FreeBytes())
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
responseBuffer.Truncate(len(n))
|
||||||
|
binary.BigEndian.PutUint16(responseBuffer.ExtendHeader(2), uint16(len(n)))
|
||||||
|
conn.Write(responseBuffer.Bytes())
|
||||||
|
}
|
||||||
|
|
||||||
func NewDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn N.PacketConn, cachedPackets []*N.PacketBuffer, metadata adapter.InboundContext) error {
|
func NewDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn N.PacketConn, cachedPackets []*N.PacketBuffer, metadata adapter.InboundContext) error {
|
||||||
metadata.Destination = M.Socksaddr{}
|
metadata.Destination = M.Socksaddr{}
|
||||||
var reader N.PacketReader = conn
|
var reader N.PacketReader = conn
|
||||||
@@ -123,24 +125,22 @@ func NewDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn
|
|||||||
timeout.Update()
|
timeout.Update()
|
||||||
}
|
}
|
||||||
metadataInQuery := metadata
|
metadataInQuery := metadata
|
||||||
go func() error {
|
router.ExchangeAsync(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{}, func(response *mDNS.Msg, err error) {
|
||||||
response, err := router.Exchange(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{})
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
cancel(err)
|
cancel(err)
|
||||||
return err
|
return
|
||||||
}
|
}
|
||||||
timeout.Update()
|
timeout.Update()
|
||||||
responseBuffer, err := dns.TruncateDNSMessage(&message, response, 1024)
|
responseBuffer, truncateErr := dns.TruncateDNSMessage(&message, response, 1024)
|
||||||
if err != nil {
|
if truncateErr != nil {
|
||||||
cancel(err)
|
cancel(truncateErr)
|
||||||
return err
|
return
|
||||||
}
|
}
|
||||||
err = conn.WritePacket(responseBuffer, destination)
|
writeErr := conn.WritePacket(responseBuffer, destination)
|
||||||
if err != nil {
|
if writeErr != nil {
|
||||||
cancel(err)
|
cancel(writeErr)
|
||||||
}
|
}
|
||||||
return err
|
})
|
||||||
}()
|
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
group.Cleanup(func() {
|
group.Cleanup(func() {
|
||||||
@@ -193,24 +193,22 @@ func newDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn
|
|||||||
timeout.Update()
|
timeout.Update()
|
||||||
}
|
}
|
||||||
metadataInQuery := metadata
|
metadataInQuery := metadata
|
||||||
go func() error {
|
router.ExchangeAsync(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{}, func(response *mDNS.Msg, err error) {
|
||||||
response, err := router.Exchange(adapter.WithContext(ctx, &metadataInQuery), &message, adapter.DNSQueryOptions{})
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
cancel(err)
|
cancel(err)
|
||||||
return err
|
return
|
||||||
}
|
}
|
||||||
timeout.Update()
|
timeout.Update()
|
||||||
responseBuffer, err := dns.TruncateDNSMessage(&message, response, 1024)
|
responseBuffer, truncateErr := dns.TruncateDNSMessage(&message, response, 1024)
|
||||||
if err != nil {
|
if truncateErr != nil {
|
||||||
cancel(err)
|
cancel(truncateErr)
|
||||||
return err
|
return
|
||||||
}
|
}
|
||||||
err = conn.WritePacket(responseBuffer, destination)
|
writeErr := conn.WritePacket(responseBuffer, destination)
|
||||||
if err != nil {
|
if writeErr != nil {
|
||||||
cancel(err)
|
cancel(writeErr)
|
||||||
}
|
}
|
||||||
return err
|
})
|
||||||
}()
|
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
group.Cleanup(func() {
|
group.Cleanup(func() {
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ package tailscale
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
@@ -276,48 +275,64 @@ func (t *DNSTransport) PreferredDomain(domain string) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *DNSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
func (t *DNSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
||||||
|
done := make(chan struct{})
|
||||||
|
var (
|
||||||
|
response *mDNS.Msg
|
||||||
|
err error
|
||||||
|
)
|
||||||
|
t.ExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
|
||||||
|
response = callbackResponse
|
||||||
|
err = callbackErr
|
||||||
|
close(done)
|
||||||
|
})
|
||||||
|
<-done
|
||||||
|
return response, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *DNSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
|
||||||
if len(message.Question) != 1 {
|
if len(message.Question) != 1 {
|
||||||
return nil, os.ErrInvalid
|
callback(nil, os.ErrInvalid)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
if t.acceptSearchDomain && mDNS.CountLabel(message.Question[0].Name) == 1 {
|
if t.acceptSearchDomain && mDNS.CountLabel(message.Question[0].Name) == 1 {
|
||||||
return t.exchangeWithSearchDomains(ctx, message)
|
t.exchangeWithSearchDomains(ctx, message, callback)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
t.access.RLock()
|
t.access.RLock()
|
||||||
acceptDefaultResolvers := t.acceptDefaultResolvers
|
acceptDefaultResolvers := t.acceptDefaultResolvers
|
||||||
t.access.RUnlock()
|
t.access.RUnlock()
|
||||||
return t.exchangeOnce(ctx, message, acceptDefaultResolvers)
|
t.exchangeOnce(ctx, message, acceptDefaultResolvers, callback)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *DNSTransport) exchangeWithSearchDomains(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
func (t *DNSTransport) exchangeWithSearchDomains(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
|
||||||
t.access.RLock()
|
t.access.RLock()
|
||||||
searchDomains := t.searchDomains
|
searchDomains := t.searchDomains
|
||||||
t.access.RUnlock()
|
t.access.RUnlock()
|
||||||
|
if len(searchDomains) == 0 {
|
||||||
|
callback(nil, dns.RcodeNameError)
|
||||||
|
return
|
||||||
|
}
|
||||||
originalQuestion := message.Question[0]
|
originalQuestion := message.Question[0]
|
||||||
singleLabel := strings.TrimSuffix(originalQuestion.Name, ".")
|
singleLabel := strings.TrimSuffix(originalQuestion.Name, ".")
|
||||||
var lastErr error
|
domainExchangers := make([]transport.AsyncExchanger, 0, len(searchDomains))
|
||||||
for _, searchDomain := range searchDomains {
|
for _, searchDomain := range searchDomains {
|
||||||
expandedName := singleLabel + "." + searchDomain
|
expandedName := singleLabel + "." + searchDomain
|
||||||
question := originalQuestion
|
domainExchangers = append(domainExchangers, func(exchangeCtx context.Context, exchangeCallback func(response *mDNS.Msg, err error)) {
|
||||||
question.Name = expandedName
|
question := originalQuestion
|
||||||
rewritten := *message
|
question.Name = expandedName
|
||||||
rewritten.Question = []mDNS.Question{question}
|
rewritten := *message
|
||||||
response, err := t.exchangeOnce(ctx, &rewritten, false)
|
rewritten.Question = []mDNS.Question{question}
|
||||||
if err == nil {
|
t.exchangeOnce(exchangeCtx, &rewritten, false, func(response *mDNS.Msg, err error) {
|
||||||
if response.Rcode == mDNS.RcodeNameError {
|
if err == nil {
|
||||||
continue
|
restoreOriginalQuestion(response, expandedName, originalQuestion)
|
||||||
}
|
}
|
||||||
restoreOriginalQuestion(response, expandedName, originalQuestion)
|
exchangeCallback(response, err)
|
||||||
return response, nil
|
})
|
||||||
}
|
})
|
||||||
if errors.Is(err, dns.RcodeNameError) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
lastErr = err
|
|
||||||
}
|
}
|
||||||
if lastErr != nil {
|
transport.ExchangeSequential(ctx, domainExchangers, func(response *mDNS.Msg, err error) bool {
|
||||||
return nil, lastErr
|
return err == nil && response.Rcode != mDNS.RcodeNameError
|
||||||
}
|
}, callback)
|
||||||
return nil, dns.RcodeNameError
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// RFC 1035 §4.1.1 requires the response Question to match the request byte-for-byte,
|
// RFC 1035 §4.1.1 requires the response Question to match the request byte-for-byte,
|
||||||
@@ -331,7 +346,7 @@ func restoreOriginalQuestion(response *mDNS.Msg, expandedName string, originalQu
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *DNSTransport) exchangeOnce(ctx context.Context, message *mDNS.Msg, allowDefaultResolvers bool) (*mDNS.Msg, error) {
|
func (t *DNSTransport) exchangeOnce(ctx context.Context, message *mDNS.Msg, allowDefaultResolvers bool, callback func(response *mDNS.Msg, err error)) {
|
||||||
question := message.Question[0]
|
question := message.Question[0]
|
||||||
|
|
||||||
t.access.RLock()
|
t.access.RLock()
|
||||||
@@ -348,58 +363,53 @@ func (t *DNSTransport) exchangeOnce(ctx context.Context, message *mDNS.Msg, allo
|
|||||||
return addr.Is4()
|
return addr.Is4()
|
||||||
})
|
})
|
||||||
if len(addresses4) > 0 {
|
if len(addresses4) > 0 {
|
||||||
return dns.FixedResponse(message.Id, question, addresses4, C.DefaultDNSTTL), nil
|
callback(dns.FixedResponse(message.Id, question, addresses4, C.DefaultDNSTTL), nil)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
case mDNS.TypeAAAA:
|
case mDNS.TypeAAAA:
|
||||||
addresses6 := common.Filter(addresses, func(addr netip.Addr) bool {
|
addresses6 := common.Filter(addresses, func(addr netip.Addr) bool {
|
||||||
return addr.Is6()
|
return addr.Is6()
|
||||||
})
|
})
|
||||||
if len(addresses6) > 0 {
|
if len(addresses6) > 0 {
|
||||||
return dns.FixedResponse(message.Id, question, addresses6, C.DefaultDNSTTL), nil
|
callback(dns.FixedResponse(message.Id, question, addresses6, C.DefaultDNSTTL), nil)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for domainSuffix, transports := range routes {
|
for domainSuffix, transports := range routes {
|
||||||
if mDNS.IsSubDomain(domainSuffix, question.Name) {
|
if mDNS.IsSubDomain(domainSuffix, question.Name) {
|
||||||
if len(transports) == 0 {
|
if len(transports) == 0 {
|
||||||
return &mDNS.Msg{
|
callback(&mDNS.Msg{
|
||||||
MsgHdr: mDNS.MsgHdr{
|
MsgHdr: mDNS.MsgHdr{
|
||||||
Id: message.Id,
|
Id: message.Id,
|
||||||
Rcode: mDNS.RcodeNameError,
|
Rcode: mDNS.RcodeNameError,
|
||||||
Response: true,
|
Response: true,
|
||||||
},
|
},
|
||||||
Question: []mDNS.Question{question},
|
Question: []mDNS.Question{question},
|
||||||
}, nil
|
}, nil)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
var lastErr error
|
transport.ExchangeSequential(ctx, resolverExchangers(transports, message), nil, callback)
|
||||||
for _, dnsTransport := range transports {
|
return
|
||||||
response, err := dnsTransport.Exchange(ctx, message)
|
|
||||||
if err != nil {
|
|
||||||
lastErr = err
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
return response, nil
|
|
||||||
}
|
|
||||||
return nil, lastErr
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if allowDefaultResolvers {
|
if allowDefaultResolvers {
|
||||||
if len(defaultResolvers) > 0 {
|
if len(defaultResolvers) > 0 {
|
||||||
var lastErr error
|
transport.ExchangeSequential(ctx, resolverExchangers(defaultResolvers, message), nil, callback)
|
||||||
for _, resolver := range defaultResolvers {
|
|
||||||
response, err := resolver.Exchange(ctx, message)
|
|
||||||
if err != nil {
|
|
||||||
lastErr = err
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
return response, nil
|
|
||||||
}
|
|
||||||
return nil, lastErr
|
|
||||||
} else {
|
} else {
|
||||||
return nil, E.New("missing default resolvers")
|
callback(nil, E.New("missing default resolvers"))
|
||||||
}
|
}
|
||||||
|
return
|
||||||
}
|
}
|
||||||
return nil, dns.RcodeNameError
|
callback(nil, dns.RcodeNameError)
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolverExchangers(resolvers []adapter.DNSTransport, message *mDNS.Msg) []transport.AsyncExchanger {
|
||||||
|
return common.Map(resolvers, func(resolver adapter.DNSTransport) transport.AsyncExchanger {
|
||||||
|
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
|
||||||
|
resolver.ExchangeAsync(ctx, message, callback)
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *DNSTransport) collectResolversLocked() []adapter.DNSTransport {
|
func (t *DNSTransport) collectResolversLocked() []adapter.DNSTransport {
|
||||||
|
|||||||
+6
-8
@@ -50,19 +50,17 @@ func (r *Router) HijackDNSPacket(ctx context.Context, payload []byte, writer N.P
|
|||||||
}
|
}
|
||||||
destination := metadata.Destination
|
destination := metadata.Destination
|
||||||
metadata.Destination = M.Socksaddr{}
|
metadata.Destination = M.Socksaddr{}
|
||||||
go func() {
|
r.dns.ExchangeAsync(adapter.WithContext(ctx, &metadata), &message, adapter.DNSQueryOptions{}, func(response *mDNS.Msg, exchangeErr error) {
|
||||||
exchangeErr := r.exchangeDNSPacket(ctx, &message, writer, metadata, destination)
|
if exchangeErr == nil {
|
||||||
|
exchangeErr = r.writeDNSPacketResponse(&message, response, writer, destination)
|
||||||
|
}
|
||||||
if exchangeErr != nil && !R.IsRejected(exchangeErr) && !E.IsClosedOrCanceled(exchangeErr) {
|
if exchangeErr != nil && !R.IsRejected(exchangeErr) && !E.IsClosedOrCanceled(exchangeErr) {
|
||||||
r.logger.ErrorContext(ctx, E.Cause(exchangeErr, "process DNS packet"))
|
r.logger.ErrorContext(ctx, E.Cause(exchangeErr, "process DNS packet"))
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Router) exchangeDNSPacket(ctx context.Context, message *mDNS.Msg, writer N.PacketWriter, metadata adapter.InboundContext, destination M.Socksaddr) error {
|
func (r *Router) writeDNSPacketResponse(message *mDNS.Msg, response *mDNS.Msg, writer N.PacketWriter, destination M.Socksaddr) error {
|
||||||
response, err := r.dns.Exchange(adapter.WithContext(ctx, &metadata), message, adapter.DNSQueryOptions{})
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
responseBuffer, err := dns.TruncateDNSMessage(message, response, 1024)
|
responseBuffer, err := dns.TruncateDNSMessage(message, response, 1024)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|||||||
@@ -210,6 +210,21 @@ func (t *Transport) PreferredDomain(domain string) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
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]
|
question := message.Question[0]
|
||||||
var selectedLink *TransportLink
|
var selectedLink *TransportLink
|
||||||
t.service.linkAccess.RLock()
|
t.service.linkAccess.RLock()
|
||||||
@@ -233,93 +248,58 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
|
|||||||
}
|
}
|
||||||
t.service.linkAccess.RUnlock()
|
t.service.linkAccess.RUnlock()
|
||||||
if selectedLink == nil {
|
if selectedLink == nil {
|
||||||
return dns.FixedResponseStatus(message, mDNS.RcodeNameError), nil
|
callback(dns.FixedResponseStatus(message, mDNS.RcodeNameError), nil)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
t.linkAccess.RLock()
|
t.linkAccess.RLock()
|
||||||
servers := t.linkServers[selectedLink]
|
servers := t.linkServers[selectedLink]
|
||||||
t.linkAccess.RUnlock()
|
t.linkAccess.RUnlock()
|
||||||
if len(servers.Servers) == 0 {
|
if servers == nil || len(servers.Servers) == 0 {
|
||||||
return dns.FixedResponseStatus(message, mDNS.RcodeNameError), nil
|
callback(dns.FixedResponseStatus(message, mDNS.RcodeNameError), nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
names := servers.Link.nameList(t.ndots, question.Name)
|
||||||
|
if len(names) == 0 {
|
||||||
|
callback(nil, E.New("invalid domain: ", question.Name))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
nameExchangers := make([]transport.AsyncExchanger, 0, len(names))
|
||||||
|
for _, fqdn := range names {
|
||||||
|
nameExchangers = append(nameExchangers, t.newNameExchanger(servers, message, fqdn))
|
||||||
}
|
}
|
||||||
if question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA {
|
if question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA {
|
||||||
return t.exchangeParallel(ctx, servers, message)
|
transport.ExchangeRace(ctx, nameExchangers, callback)
|
||||||
} else {
|
} else {
|
||||||
return t.exchangeSingleRequest(ctx, servers, message)
|
transport.ExchangeSequential(ctx, nameExchangers, nil, callback)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *Transport) exchangeSingleRequest(ctx context.Context, servers *LinkServers, message *mDNS.Msg) (*mDNS.Msg, error) {
|
func (t *Transport) newNameExchanger(servers *LinkServers, message *mDNS.Msg, fqdn string) transport.AsyncExchanger {
|
||||||
var lastErr error
|
|
||||||
for _, fqdn := range servers.Link.nameList(t.ndots, message.Question[0].Name) {
|
|
||||||
response, err := t.tryOneName(ctx, servers, message, fqdn)
|
|
||||||
if err != nil {
|
|
||||||
lastErr = err
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
return response, nil
|
|
||||||
}
|
|
||||||
return nil, lastErr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *Transport) tryOneName(ctx context.Context, servers *LinkServers, message *mDNS.Msg, fqdn string) (*mDNS.Msg, error) {
|
|
||||||
serverOffset := servers.ServerOffset(t.rotate)
|
serverOffset := servers.ServerOffset(t.rotate)
|
||||||
sLen := uint32(len(servers.Servers))
|
serverCount := uint32(len(servers.Servers))
|
||||||
var lastErr error
|
attemptExchangers := make([]transport.AsyncExchanger, 0, t.attempts*int(serverCount))
|
||||||
for i := 0; i < t.attempts; i++ {
|
for i := 0; i < t.attempts; i++ {
|
||||||
for j := range sLen {
|
for j := range serverCount {
|
||||||
server := servers.Servers[(serverOffset+j)%sLen]
|
server := servers.Servers[(serverOffset+j)%serverCount]
|
||||||
question := message.Question[0]
|
attemptExchangers = append(attemptExchangers, func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
|
||||||
question.Name = fqdn
|
question := message.Question[0]
|
||||||
exchangeMessage := *message
|
question.Name = fqdn
|
||||||
exchangeMessage.Question = []mDNS.Question{question}
|
exchangeMessage := *message
|
||||||
exchangeCtx, cancel := context.WithTimeout(ctx, t.timeout)
|
exchangeMessage.Question = []mDNS.Question{question}
|
||||||
response, err := server.Exchange(exchangeCtx, &exchangeMessage)
|
exchangeCtx, cancel := context.WithTimeout(ctx, t.timeout)
|
||||||
cancel()
|
server.ExchangeAsync(exchangeCtx, &exchangeMessage, func(response *mDNS.Msg, err error) {
|
||||||
|
cancel()
|
||||||
|
callback(response, err)
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
|
||||||
|
transport.ExchangeSequential(ctx, attemptExchangers, nil, func(response *mDNS.Msg, err error) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
lastErr = err
|
err = E.Cause(err, fqdn)
|
||||||
continue
|
|
||||||
}
|
}
|
||||||
return response, nil
|
callback(response, err)
|
||||||
}
|
})
|
||||||
}
|
|
||||||
return nil, E.Cause(lastErr, fqdn)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *Transport) exchangeParallel(ctx context.Context, servers *LinkServers, message *mDNS.Msg) (*mDNS.Msg, error) {
|
|
||||||
returned := make(chan struct{})
|
|
||||||
defer close(returned)
|
|
||||||
type queryResult struct {
|
|
||||||
response *mDNS.Msg
|
|
||||||
err error
|
|
||||||
}
|
|
||||||
results := make(chan queryResult)
|
|
||||||
startRacer := func(ctx context.Context, fqdn string) {
|
|
||||||
response, err := t.tryOneName(ctx, servers, message, fqdn)
|
|
||||||
select {
|
|
||||||
case results <- queryResult{response, err}:
|
|
||||||
case <-returned:
|
|
||||||
}
|
|
||||||
}
|
|
||||||
queryCtx, queryCancel := context.WithCancel(ctx)
|
|
||||||
defer queryCancel()
|
|
||||||
var nameCount int
|
|
||||||
for _, fqdn := range servers.Link.nameList(t.ndots, message.Question[0].Name) {
|
|
||||||
nameCount++
|
|
||||||
go startRacer(queryCtx, fqdn)
|
|
||||||
}
|
|
||||||
var errors []error
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
return nil, ctx.Err()
|
|
||||||
case result := <-results:
|
|
||||||
if result.err == nil {
|
|
||||||
return result.response, nil
|
|
||||||
}
|
|
||||||
errors = append(errors, result.err)
|
|
||||||
if len(errors) == nameCount {
|
|
||||||
return nil, E.Errors(errors...)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user