Add search domain support for Tailscale DNS

This commit is contained in:
世界
2026-04-20 09:39:32 +08:00
parent dae5dcf632
commit 3329f5deed
4 changed files with 74 additions and 5 deletions
+45 -3
View File
@@ -4,6 +4,7 @@ package tailscale
import (
"context"
"errors"
"net"
"net/http"
"net/netip"
@@ -28,6 +29,7 @@ import (
"github.com/sagernet/sing/service"
nDNS "github.com/sagernet/tailscale/net/dns"
"github.com/sagernet/tailscale/types/dnstype"
"github.com/sagernet/tailscale/util/dnsname"
"github.com/sagernet/tailscale/wgengine/router"
"github.com/sagernet/tailscale/wgengine/wgcfg"
@@ -46,12 +48,14 @@ type DNSTransport struct {
logger logger.ContextLogger
endpointTag string
acceptDefaultResolvers bool
acceptSearchDomain bool
dnsRouter adapter.DNSRouter
endpointManager adapter.EndpointManager
endpoint *Endpoint
routePrefixes []netip.Prefix
routes map[string][]adapter.DNSTransport
hosts map[string][]netip.Addr
searchDomains []string
defaultResolvers []adapter.DNSTransport
}
@@ -65,6 +69,7 @@ func NewDNSTransport(ctx context.Context, logger log.ContextLogger, tag string,
logger: logger,
endpointTag: options.Endpoint,
acceptDefaultResolvers: options.AcceptDefaultResolvers,
acceptSearchDomain: options.AcceptSearchDomain,
dnsRouter: service.FromContext[adapter.DNSRouter](ctx),
endpointManager: service.FromContext[adapter.EndpointManager](ctx),
}, nil
@@ -122,6 +127,9 @@ func (t *DNSTransport) updateDNSServers(routeConfig *router.Config, dnsConfig *n
for domain, addresses := range dnsConfig.Hosts {
hosts[domain.WithTrailingDot()] = addresses
}
searchDomains := common.Map(dnsConfig.SearchDomains, func(it dnsname.FQDN) string {
return it.WithTrailingDot()
})
var defaultResolvers []adapter.DNSTransport
for _, resolver := range dnsConfig.DefaultResolvers {
myResolver, err := t.createResolver(directDialerOnce, resolver)
@@ -132,12 +140,13 @@ func (t *DNSTransport) updateDNSServers(routeConfig *router.Config, dnsConfig *n
}
t.routes = routes
t.hosts = hosts
t.searchDomains = searchDomains
t.defaultResolvers = defaultResolvers
if len(defaultResolvers) > 0 {
t.logger.Info("updated ", len(routes), " routes, ", len(hosts), " hosts, default resolvers: ",
t.logger.Info("updated ", len(routes), " routes, ", len(hosts), " hosts, ", len(searchDomains), " search domains, default resolvers: ",
strings.Join(common.Map(dnsConfig.DefaultResolvers, func(it *dnstype.Resolver) string { return it.Addr }), " "))
} else {
t.logger.Info("updated ", len(routes), " routes, ", len(hosts), " hosts")
t.logger.Info("updated ", len(routes), " routes, ", len(hosts), " hosts, ", len(searchDomains), " search domains")
}
return nil
}
@@ -218,6 +227,39 @@ func (t *DNSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.M
if len(message.Question) != 1 {
return nil, os.ErrInvalid
}
if t.acceptSearchDomain && mDNS.CountLabel(message.Question[0].Name) == 1 {
return t.exchangeWithSearchDomains(ctx, message)
}
return t.exchangeOnce(ctx, message, t.acceptDefaultResolvers)
}
func (t *DNSTransport) exchangeWithSearchDomains(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
singleLabel := strings.TrimSuffix(message.Question[0].Name, ".")
var lastErr error
for _, searchDomain := range t.searchDomains {
question := message.Question[0]
question.Name = singleLabel + "." + searchDomain
rewritten := *message
rewritten.Question = []mDNS.Question{question}
response, err := t.exchangeOnce(ctx, &rewritten, false)
if err == nil {
if response.Rcode == mDNS.RcodeNameError {
continue
}
return response, nil
}
if errors.Is(err, dns.RcodeNameError) {
continue
}
lastErr = err
}
if lastErr != nil {
return nil, lastErr
}
return nil, dns.RcodeNameError
}
func (t *DNSTransport) exchangeOnce(ctx context.Context, message *mDNS.Msg, allowDefaultResolvers bool) (*mDNS.Msg, error) {
question := message.Question[0]
addresses, hostsLoaded := t.hosts[question.Name]
if hostsLoaded {
@@ -262,7 +304,7 @@ func (t *DNSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.M
return nil, lastErr
}
}
if t.acceptDefaultResolvers {
if allowDefaultResolvers {
if len(t.defaultResolvers) > 0 {
var lastErr error
for _, resolver := range t.defaultResolvers {