From cb664d1a4bff89c27c900efecd6de07d1ced7afd Mon Sep 17 00:00:00 2001 From: Shtorm <108103062+shtorm-7@users.noreply.github.com> Date: Sun, 9 Aug 2026 15:01:36 +0300 Subject: [PATCH] Fix DNS Fallback --- dns/transport/fallback/fallback.go | 2 +- dns/transport/fallback/strategy.go | 48 +++++++++++++++++++++++------- examples/dns_fallback/client.json | 3 +- option/dns.go | 5 ++-- 4 files changed, 44 insertions(+), 14 deletions(-) diff --git a/dns/transport/fallback/fallback.go b/dns/transport/fallback/fallback.go index 47d8b0d1..8b71d8b2 100644 --- a/dns/transport/fallback/fallback.go +++ b/dns/transport/fallback/fallback.go @@ -43,7 +43,7 @@ func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, opt } servers[i] = server } - strategy, err := CreateStrategy(options.Strategy, servers, logger) + strategy, err := CreateStrategy(options.Strategy, servers, logger, options.Timeout.Build()) if err != nil { return nil, err } diff --git a/dns/transport/fallback/strategy.go b/dns/transport/fallback/strategy.go index 34580956..591e0f1c 100644 --- a/dns/transport/fallback/strategy.go +++ b/dns/transport/fallback/strategy.go @@ -2,18 +2,20 @@ package fallback import ( "context" + "time" mDNS "github.com/miekg/dns" "github.com/sagernet/sing-box/adapter" + C "github.com/sagernet/sing-box/constant" E "github.com/sagernet/sing/common/exceptions" "github.com/sagernet/sing/common/logger" ) 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) { - queryCtx, cancel := context.WithCancel(ctx) + ctx, cancel := context.WithTimeout(ctx, timeout) defer cancel() type result struct { response *mDNS.Msg @@ -22,10 +24,13 @@ func parallelStrategy(servers []adapter.DNSTransport, logger logger.ContextLogge results := make(chan result) for _, server := range servers { 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 { 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) { + ctx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() var lastErr error - for _, server := range servers { - response, err := server.Exchange(ctx, message) + for index, server := range servers { + exchangeCtx, exchangeCancel := context.WithTimeout(ctx, perAttemptTimeout(ctx, len(servers)-index)) + response, err := checkExchangeResponse(server.Exchange(exchangeCtx, message)) + exchangeCancel() if err != nil { + logger.InfoContext(ctx, E.Cause(err, "resolve failed for server ", server.Tag())) lastErr = err 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 { case "parallel": - return parallelStrategy(servers, logger), nil + return parallelStrategy(servers, logger, timeout), nil case "", "sequential": - return sequentialStrategy(servers, logger), nil + return sequentialStrategy(servers, logger, timeout), nil default: return nil, E.New("strategy not found: ", strategy) } diff --git a/examples/dns_fallback/client.json b/examples/dns_fallback/client.json index 4d2aaf8a..a2dd9de1 100644 --- a/examples/dns_fallback/client.json +++ b/examples/dns_fallback/client.json @@ -34,7 +34,8 @@ // - "parallel": query all servers concurrently. Returns // the first successful response (cancelling the rest), or the last // error if all servers failed. - "strategy": "sequential" + "strategy": "sequential", + "timeout": "10s" // overall budget for the whole fallback exchange } ], "disable_cache": true, diff --git a/option/dns.go b/option/dns.go index e2e97106..88853e69 100644 --- a/option/dns.go +++ b/option/dns.go @@ -427,6 +427,7 @@ type SDNSDNSServerOptions struct { } type FallbackDNSServerOptions struct { - Servers []string `json:"servers"` - Strategy string `json:"strategy,omitempty"` + Servers []string `json:"servers"` + Strategy string `json:"strategy,omitempty"` + Timeout badoption.Duration `json:"timeout,omitempty"` }