Fix DNS Fallback

This commit is contained in:
Shtorm
2026-08-09 15:01:36 +03:00
parent f682ceb8e2
commit cb664d1a4b
4 changed files with 44 additions and 14 deletions

View File

@@ -43,7 +43,7 @@ func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, opt
} }
servers[i] = server servers[i] = server
} }
strategy, err := CreateStrategy(options.Strategy, servers, logger) strategy, err := CreateStrategy(options.Strategy, servers, logger, options.Timeout.Build())
if err != nil { if err != nil {
return nil, err return nil, err
} }

View File

@@ -2,18 +2,20 @@ package fallback
import ( import (
"context" "context"
"time"
mDNS "github.com/miekg/dns" mDNS "github.com/miekg/dns"
"github.com/sagernet/sing-box/adapter" "github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
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"
) )
type ExchangeStrategy = func(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) type ExchangeStrategy = func(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error)
func parallelStrategy(servers []adapter.DNSTransport, logger logger.ContextLogger) ExchangeStrategy { func parallelStrategy(servers []adapter.DNSTransport, logger logger.ContextLogger, timeout time.Duration) ExchangeStrategy {
return func(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) { return func(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
queryCtx, cancel := context.WithCancel(ctx) ctx, cancel := context.WithTimeout(ctx, timeout)
defer cancel() defer cancel()
type result struct { type result struct {
response *mDNS.Msg response *mDNS.Msg
@@ -22,10 +24,13 @@ func parallelStrategy(servers []adapter.DNSTransport, logger logger.ContextLogge
results := make(chan result) results := make(chan result)
for _, server := range servers { for _, server := range servers {
go func() { go func() {
response, err := server.Exchange(queryCtx, message) response, err := checkExchangeResponse(server.Exchange(ctx, message))
if err != nil {
logger.InfoContext(ctx, E.Cause(err, "resolve failed for server ", server.Tag()))
}
select { select {
case results <- result{response, err}: case results <- result{response, err}:
case <-queryCtx.Done(): case <-ctx.Done():
} }
}() }()
} }
@@ -46,12 +51,17 @@ func parallelStrategy(servers []adapter.DNSTransport, logger logger.ContextLogge
} }
} }
func sequentialStrategy(servers []adapter.DNSTransport, logger logger.ContextLogger) ExchangeStrategy { func sequentialStrategy(servers []adapter.DNSTransport, logger logger.ContextLogger, timeout time.Duration) ExchangeStrategy {
return func(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) { return func(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
ctx, cancel := context.WithTimeout(ctx, timeout)
defer cancel()
var lastErr error var lastErr error
for _, server := range servers { for index, server := range servers {
response, err := server.Exchange(ctx, message) exchangeCtx, exchangeCancel := context.WithTimeout(ctx, perAttemptTimeout(ctx, len(servers)-index))
response, err := checkExchangeResponse(server.Exchange(exchangeCtx, message))
exchangeCancel()
if err != nil { if err != nil {
logger.InfoContext(ctx, E.Cause(err, "resolve failed for server ", server.Tag()))
lastErr = err lastErr = err
continue continue
} }
@@ -61,12 +71,30 @@ func sequentialStrategy(servers []adapter.DNSTransport, logger logger.ContextLog
} }
} }
func CreateStrategy(strategy string, servers []adapter.DNSTransport, logger logger.ContextLogger) (ExchangeStrategy, error) { func checkExchangeResponse(response *mDNS.Msg, err error) (*mDNS.Msg, error) {
if err != nil {
return nil, err
}
if response.Rcode != mDNS.RcodeSuccess && response.Rcode != mDNS.RcodeNameError {
return nil, E.New("bad response rcode: ", mDNS.RcodeToString[response.Rcode])
}
return response, nil
}
func perAttemptTimeout(ctx context.Context, remaining int) time.Duration {
deadline, _ := ctx.Deadline()
return time.Until(deadline) / time.Duration(remaining)
}
func CreateStrategy(strategy string, servers []adapter.DNSTransport, logger logger.ContextLogger, timeout time.Duration) (ExchangeStrategy, error) {
if timeout <= 0 {
timeout = C.DNSTimeout
}
switch strategy { switch strategy {
case "parallel": case "parallel":
return parallelStrategy(servers, logger), nil return parallelStrategy(servers, logger, timeout), nil
case "", "sequential": case "", "sequential":
return sequentialStrategy(servers, logger), nil return sequentialStrategy(servers, logger, timeout), nil
default: default:
return nil, E.New("strategy not found: ", strategy) return nil, E.New("strategy not found: ", strategy)
} }

View File

@@ -34,7 +34,8 @@
// - "parallel": query all servers concurrently. Returns // - "parallel": query all servers concurrently. Returns
// the first successful response (cancelling the rest), or the last // the first successful response (cancelling the rest), or the last
// error if all servers failed. // error if all servers failed.
"strategy": "sequential" "strategy": "sequential",
"timeout": "10s" // overall budget for the whole fallback exchange
} }
], ],
"disable_cache": true, "disable_cache": true,

View File

@@ -429,4 +429,5 @@ type SDNSDNSServerOptions struct {
type FallbackDNSServerOptions struct { type FallbackDNSServerOptions struct {
Servers []string `json:"servers"` Servers []string `json:"servers"`
Strategy string `json:"strategy,omitempty"` Strategy string `json:"strategy,omitempty"`
Timeout badoption.Duration `json:"timeout,omitempty"`
} }