From 4abdb6560a24ea8eb338a3da893a3de252582368 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Wed, 22 Jul 2026 15:30:33 +0800 Subject: [PATCH] dns: Cache responses with client subnet --- dns/client.go | 68 ++++++++++++++++++++++------------- dns/extension_edns0_subnet.go | 21 +++++++++++ 2 files changed, 64 insertions(+), 25 deletions(-) diff --git a/dns/client.go b/dns/client.go index 9b314bbcd..904c99e5c 100644 --- a/dns/client.go +++ b/dns/client.go @@ -86,6 +86,24 @@ func NewClient(options ClientOptions) *Client { type dnsCacheKey struct { dns.Question transportTag string + clientSubnet netip.Prefix +} + +func (k dnsCacheKey) persistentName() string { + if !k.clientSubnet.IsValid() { + return k.transportTag + } + return k.transportTag + "\x00" + k.clientSubnet.String() +} + +func (c *Client) effectiveClientSubnet(message *dns.Msg, options adapter.DNSQueryOptions) netip.Prefix { + if options.ClientSubnet.IsValid() { + return options.ClientSubnet + } + if c.clientSubnet.IsValid() { + return c.clientSubnet + } + return clientSubnetFromMessage(message) } func (c *Client) Start() { @@ -168,6 +186,7 @@ type exchangeOperation struct { options adapter.DNSQueryOptions responseChecker func(response *dns.Msg) bool disableCache bool + cacheKey dnsCacheKey releaseCond func() } @@ -192,16 +211,16 @@ func (c *Client) beginExchange(ctx context.Context, transport adapter.DNSTranspo } return nil, FixedResponseStatus(message, dns.RcodeSuccess), exchangeDone, nil } - message = c.prepareExchangeMessage(message, options) - isSimpleRequest := len(message.Question) == 1 && len(message.Ns) == 0 && (len(message.Extra) == 0 || len(message.Extra) == 1 && message.Extra[0].Header().Rrtype == dns.TypeOPT && message.Extra[0].Header().Class > 0 && message.Extra[0].Header().Ttl == 0 && - len(message.Extra[0].(*dns.OPT).Option) == 0) && - !options.ClientSubnet.IsValid() + common.All(message.Extra[0].(*dns.OPT).Option, func(it dns.EDNS0) bool { + return it.Option() == dns.EDNS0SUBNET + })) + message = c.prepareExchangeMessage(message, options) disableCache := !isSimpleRequest || c.disableCache || options.DisableCache operation := &exchangeOperation{ message: message, @@ -212,7 +231,8 @@ func (c *Client) beginExchange(ctx context.Context, transport adapter.DNSTranspo disableCache: disableCache, } if !disableCache { - cacheKey := dnsCacheKey{Question: question, transportTag: transport.Tag()} + cacheKey := dnsCacheKey{Question: question, transportTag: transport.Tag(), clientSubnet: c.effectiveClientSubnet(message, options)} + operation.cacheKey = cacheKey cond, loaded := c.cacheLock.LoadOrStore(cacheKey, make(chan struct{})) if loaded { if !allowWait { @@ -229,10 +249,10 @@ func (c *Client) beginExchange(ctx context.Context, transport adapter.DNSTranspo close(cond) } } - response, ttl, isStale := c.loadResponse(question, transport) + response, ttl, isStale := c.loadResponse(cacheKey) if response != nil { if isStale && !options.DisableOptimisticCache { - c.backgroundRefreshDNS(transport, question, message.Copy(), options, responseChecker) + c.backgroundRefreshDNS(transport, cacheKey, message.Copy(), options, responseChecker) logOptimisticResponse(c.logger, ctx, response) response.Id = message.Id operation.release() @@ -283,7 +303,7 @@ func (c *Client) finishExchange(transport adapter.DNSTransport, operation *excha } timeToLive := applyResponseOptions(question, response, operation.options) if !disableCache { - c.storeCache(transport, question, response, timeToLive) + c.storeCache(operation.cacheKey, response, timeToLive) } response.Id = operation.messageId requestEDNSOpt := operation.message.IsEdns0() @@ -403,7 +423,7 @@ func sortAddresses(response4 []netip.Addr, response6 []netip.Addr, strategy C.Do } } -func (c *Client) storeCache(transport adapter.DNSTransport, question dns.Question, message *dns.Msg, timeToLive uint32) { +func (c *Client) storeCache(key dnsCacheKey, message *dns.Msg, timeToLive uint32) { if timeToLive == 0 { return } @@ -411,14 +431,13 @@ func (c *Client) storeCache(transport adapter.DNSTransport, question dns.Questio packed, err := message.Pack() if err == nil { expireAt := time.Now().Add(time.Second * time.Duration(timeToLive)) - c.dnsCache.SaveDNSCacheAsync(transport.Tag(), question.Name, question.Qtype, packed, expireAt, c.logger) + c.dnsCache.SaveDNSCacheAsync(key.persistentName(), key.Name, key.Qtype, packed, expireAt, c.logger) } return } if c.cache == nil { return } - key := dnsCacheKey{Question: question, transportTag: transport.Tag()} if c.disableExpire { c.cache.Add(key, message.Copy()) } else { @@ -457,7 +476,8 @@ func (c *Client) lookupToExchange(ctx context.Context, transport adapter.DNSTran func (c *Client) questionCache(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool) ([]netip.Addr, error) { question := message.Question[0] - response, _, isStale := c.loadResponse(question, transport) + cacheKey := dnsCacheKey{Question: question, transportTag: transport.Tag(), clientSubnet: c.effectiveClientSubnet(message, options)} + response, _, isStale := c.loadResponse(cacheKey) if response == nil { return nil, ErrNotCached } @@ -465,7 +485,7 @@ func (c *Client) questionCache(ctx context.Context, transport adapter.DNSTranspo if options.DisableOptimisticCache { return nil, ErrNotCached } - c.backgroundRefreshDNS(transport, question, c.prepareExchangeMessage(message.Copy(), options), options, responseChecker) + c.backgroundRefreshDNS(transport, cacheKey, c.prepareExchangeMessage(message.Copy(), options), options, responseChecker) logOptimisticResponse(c.logger, ctx, response) } if response.Rcode != dns.RcodeSuccess { @@ -474,14 +494,13 @@ func (c *Client) questionCache(ctx context.Context, transport adapter.DNSTranspo return MessageToAddresses(response), nil } -func (c *Client) loadResponse(question dns.Question, transport adapter.DNSTransport) (*dns.Msg, int, bool) { +func (c *Client) loadResponse(key dnsCacheKey) (*dns.Msg, int, bool) { if c.dnsCache != nil { - return c.loadPersistentResponse(question, transport) + return c.loadPersistentResponse(key) } if c.cache == nil { return nil, 0, false } - key := dnsCacheKey{Question: question, transportTag: transport.Tag()} if c.disableExpire { response, loaded := c.cache.Get(key) if !loaded { @@ -509,8 +528,8 @@ func (c *Client) loadResponse(question dns.Question, transport adapter.DNSTransp return response, nowTTL, false } -func (c *Client) loadPersistentResponse(question dns.Question, transport adapter.DNSTransport) (*dns.Msg, int, bool) { - rawMessage, expireAt, loaded := c.dnsCache.LoadDNSCache(transport.Tag(), question.Name, question.Qtype) +func (c *Client) loadPersistentResponse(key dnsCacheKey) (*dns.Msg, int, bool) { + rawMessage, expireAt, loaded := c.dnsCache.LoadDNSCache(key.persistentName(), key.Name, key.Qtype) if !loaded { return nil, 0, false } @@ -560,8 +579,7 @@ func applyResponseOptions(question dns.Question, response *dns.Msg, options adap return timeToLive } -func (c *Client) backgroundRefreshDNS(transport adapter.DNSTransport, question dns.Question, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool) { - key := dnsCacheKey{Question: question, transportTag: transport.Tag()} +func (c *Client) backgroundRefreshDNS(transport adapter.DNSTransport, key dnsCacheKey, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool) { _, loaded := c.backgroundRefresh.LoadOrStore(key, struct{}{}) if loaded { return @@ -572,7 +590,7 @@ func (c *Client) backgroundRefreshDNS(transport adapter.DNSTransport, question d response, err := c.exchangeToTransport(ctx, transport, message, options.Timeout) if err != nil { if c.logger != nil { - c.logger.DebugContext(ctx, "optimistic refresh failed for ", FqdnToDomain(question.Name), ": ", err) + c.logger.DebugContext(ctx, "optimistic refresh failed for ", FqdnToDomain(key.Name), ": ", err) } return } @@ -585,18 +603,18 @@ func (c *Client) backgroundRefreshDNS(transport adapter.DNSTransport, question d } if rejected { if c.logger != nil { - c.logger.DebugContext(ctx, "optimistic refresh rejected for ", FqdnToDomain(question.Name)) + c.logger.DebugContext(ctx, "optimistic refresh rejected for ", FqdnToDomain(key.Name)) } if c.rdrc != nil { - c.rdrc.SaveRDRCAsync(transport.Tag(), question.Name, question.Qtype, c.logger) + c.rdrc.SaveRDRCAsync(transport.Tag(), key.Name, key.Qtype, c.logger) } return } } else if response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError { return } - timeToLive := applyResponseOptions(question, response, options) - c.storeCache(transport, question, response, timeToLive) + timeToLive := applyResponseOptions(key.Question, response, options) + c.storeCache(key, response, timeToLive) logRefreshedResponse(c.logger, ctx, response, timeToLive) }() } diff --git a/dns/extension_edns0_subnet.go b/dns/extension_edns0_subnet.go index e804fb6cd..772580975 100644 --- a/dns/extension_edns0_subnet.go +++ b/dns/extension_edns0_subnet.go @@ -10,6 +10,27 @@ func SetClientSubnet(message *dns.Msg, clientSubnet netip.Prefix) *dns.Msg { return setClientSubnet(message, clientSubnet, true) } +func clientSubnetFromMessage(message *dns.Msg) netip.Prefix { + for _, record := range message.Extra { + optRecord, isOPTRecord := record.(*dns.OPT) + if !isOPTRecord { + continue + } + for _, option := range optRecord.Option { + subnetOption, isEDNS0Subnet := option.(*dns.EDNS0_SUBNET) + if !isEDNS0Subnet { + continue + } + address, addressLoaded := netip.AddrFromSlice(subnetOption.Address) + if !addressLoaded { + return netip.Prefix{} + } + return netip.PrefixFrom(address, int(subnetOption.SourceNetmask)) + } + } + return netip.Prefix{} +} + func setClientSubnet(message *dns.Msg, clientSubnet netip.Prefix, clone bool) *dns.Msg { var ( optRecord *dns.OPT