dns: Cache responses with client subnet

This commit is contained in:
世界
2026-07-22 15:30:33 +08:00
parent 908e12c55b
commit 4abdb6560a
2 changed files with 64 additions and 25 deletions
+43 -25
View File
@@ -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)
}()
}
+21
View File
@@ -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