plugin/cache: preserve AD when storing cache entries (#8438)

This commit is contained in:
Sueun Cho
2026-08-15 12:28:20 +09:00
committed by GitHub
parent ab7a8880a4
commit 29ef323f82
3 changed files with 141 additions and 14 deletions

View File

@@ -3,9 +3,11 @@ package cache
import ( import (
"context" "context"
"testing" "testing"
"time"
"github.com/coredns/coredns/plugin/pkg/dnstest" "github.com/coredns/coredns/plugin/pkg/dnstest"
"github.com/coredns/coredns/plugin/test" "github.com/coredns/coredns/plugin/test"
"github.com/coredns/coredns/request"
"github.com/miekg/dns" "github.com/miekg/dns"
) )
@@ -28,7 +30,8 @@ func TestCacheADBitNotPartitioned(t *testing.T) {
// The upstream answer is authenticated, so the +ad query must receive AD=1. // The upstream answer is authenticated, so the +ad query must receive AD=1.
t.Run("noad_then_ad", func(t *testing.T) { t.Run("noad_then_ad", func(t *testing.T) {
c := New() c := New()
c.Next = dnssecHandler() // sets AuthenticatedData=true on the reply h := &adBitHandler{}
c.Next = h
// First query: AD not requested, DO not set -> miss, populates cache. // First query: AD not requested, DO not set -> miss, populates cache.
noad := new(dns.Msg) noad := new(dns.Msg)
@@ -38,6 +41,12 @@ func TestCacheADBitNotPartitioned(t *testing.T) {
if rec.Msg.AuthenticatedData { if rec.Msg.AuthenticatedData {
t.Errorf("first query did not request AD, expected AuthenticatedData=false, got true") t.Errorf("first query did not request AD, expected AuthenticatedData=false, got true")
} }
if !h.requestedAD[0] {
t.Errorf("cache refresh should ask upstream for AD when populating an AD-shared entry")
}
if c.pcache.Len() != 1 {
t.Fatalf("expected first query to populate one cache entry, got %d", c.pcache.Len())
}
// Second query: AD requested, DO not set -> hit on the same key. // Second query: AD requested, DO not set -> hit on the same key.
// Must reflect the authenticated answer with AD=1. // Must reflect the authenticated answer with AD=1.
@@ -46,6 +55,9 @@ func TestCacheADBitNotPartitioned(t *testing.T) {
ad.AuthenticatedData = true ad.AuthenticatedData = true
rec = dnstest.NewRecorder(&test.ResponseWriter{}) rec = dnstest.NewRecorder(&test.ResponseWriter{})
c.ServeDNS(context.TODO(), rec, ad) c.ServeDNS(context.TODO(), rec, ad)
if h.calls != 1 {
t.Fatalf("expected second query to be served from cache, backend calls=%d", h.calls)
}
if !rec.Msg.AuthenticatedData { if !rec.Msg.AuthenticatedData {
t.Errorf("second query requested AD on an authenticated cached answer, expected AuthenticatedData=true, got false") t.Errorf("second query requested AD on an authenticated cached answer, expected AuthenticatedData=true, got false")
} }
@@ -55,7 +67,8 @@ func TestCacheADBitNotPartitioned(t *testing.T) {
// pin it so a fix for the forward case never breaks it. // pin it so a fix for the forward case never breaks it.
t.Run("ad_then_noad", func(t *testing.T) { t.Run("ad_then_noad", func(t *testing.T) {
c := New() c := New()
c.Next = dnssecHandler() h := &adBitHandler{}
c.Next = h
// First query: AD requested -> AD=1 expected. // First query: AD requested -> AD=1 expected.
ad := new(dns.Msg) ad := new(dns.Msg)
@@ -66,14 +79,117 @@ func TestCacheADBitNotPartitioned(t *testing.T) {
if !rec.Msg.AuthenticatedData { if !rec.Msg.AuthenticatedData {
t.Errorf("first query requested AD on an authenticated answer, expected AuthenticatedData=true, got false") t.Errorf("first query requested AD on an authenticated answer, expected AuthenticatedData=true, got false")
} }
if c.pcache.Len() != 1 {
t.Fatalf("expected first query to populate one cache entry, got %d", c.pcache.Len())
}
// Second query: AD not requested -> AD must be cleared for this client. // Second query: AD not requested -> AD must be cleared for this client.
noad := new(dns.Msg) noad := new(dns.Msg)
noad.SetQuestion("invent.example.org.", dns.TypeA) noad.SetQuestion("invent.example.org.", dns.TypeA)
rec = dnstest.NewRecorder(&test.ResponseWriter{}) rec = dnstest.NewRecorder(&test.ResponseWriter{})
c.ServeDNS(context.TODO(), rec, noad) c.ServeDNS(context.TODO(), rec, noad)
if h.calls != 1 {
t.Fatalf("expected second query to be served from cache, backend calls=%d", h.calls)
}
if rec.Msg.AuthenticatedData { if rec.Msg.AuthenticatedData {
t.Errorf("second query did not request AD, expected AuthenticatedData=false, got true") t.Errorf("second query did not request AD, expected AuthenticatedData=false, got true")
} }
}) })
} }
func TestCacheADBitStaleVerify(t *testing.T) {
c := New()
h := &adBitHandler{}
c.Next = h
c.staleUpTo = time.Hour
c.verifyStale = true
now := time.Now()
c.now = func() time.Time { return now }
req := new(dns.Msg)
req.SetQuestion("invent.example.org.", dns.TypeA)
rec := dnstest.NewRecorder(&test.ResponseWriter{})
c.ServeDNS(context.TODO(), rec, req)
if c.pcache.Len() != 1 {
t.Fatalf("expected first query to populate one cache entry, got %d", c.pcache.Len())
}
if !h.requestedAD[0] {
t.Errorf("cache refresh should ask upstream for AD when populating an AD-shared entry")
}
now = now.Add(2 * time.Minute)
ad := req.Copy()
ad.AuthenticatedData = true
rec = dnstest.NewRecorder(&test.ResponseWriter{})
c.ServeDNS(context.TODO(), rec, ad)
if h.calls != 2 {
t.Fatalf("expected stale verify to call backend, backend calls=%d", h.calls)
}
if !h.requestedAD[1] {
t.Errorf("stale verify should ask upstream for AD when refreshing an AD-shared entry")
}
if !rec.Msg.AuthenticatedData {
t.Errorf("AD-requesting client should receive AD from a refreshed authenticated answer, got false")
}
}
func TestCacheADBitPrefetchRequestsAD(t *testing.T) {
requestedAD := make(chan bool, 2)
c := New()
h := &adBitHandler{requestedADCh: requestedAD}
c.Next = h
c.prefetch = 1
c.percentage = 100
now := time.Now()
c.now = func() time.Time { return now }
req := new(dns.Msg)
req.SetQuestion("invent.example.org.", dns.TypeA)
rec := dnstest.NewRecorder(&test.ResponseWriter{})
c.ServeDNS(context.TODO(), rec, req)
if got := <-requestedAD; !got {
t.Errorf("cache refresh should ask upstream for AD when populating an AD-shared entry")
}
now = now.Add(time.Second)
rec = dnstest.NewRecorder(&test.ResponseWriter{})
c.ServeDNS(context.TODO(), rec, req.Copy())
select {
case got := <-requestedAD:
if !got {
t.Errorf("prefetch should ask upstream for AD when refreshing an AD-shared entry")
}
case <-time.After(time.Second):
t.Fatal("prefetch did not call backend")
}
}
type adBitHandler struct {
calls int
requestedAD []bool
requestedADCh chan bool
}
func (h *adBitHandler) Name() string { return "adBitHandler" }
func (h *adBitHandler) ServeDNS(_ context.Context, w dns.ResponseWriter, r *dns.Msg) (int, error) {
h.calls++
h.requestedAD = append(h.requestedAD, r.AuthenticatedData)
if h.requestedADCh != nil {
h.requestedADCh <- r.AuthenticatedData
}
m := new(dns.Msg)
m.SetReply(r)
state := request.Request{W: w, Req: r}
m.AuthenticatedData = r.AuthenticatedData || state.Do()
m.Answer = []dns.RR{test.A("invent.example.org. 60 IN A 192.0.2.1")}
w.WriteMsg(m)
return dns.RcodeSuccess, nil
}

17
plugin/cache/cache.go vendored
View File

@@ -348,11 +348,14 @@ var prefetchAddr = &net.TCPAddr{}
// short-circuits after caching when w.prefetch is true, and the nil-safe // short-circuits after caching when w.prefetch is true, and the nil-safe
// overrides below make the remaining dns.ResponseWriter methods well-defined. // overrides below make the remaining dns.ResponseWriter methods well-defined.
func newPrefetchResponseWriter(server string, req *dns.Msg, do, cd bool, c *Cache) *ResponseWriter { func newPrefetchResponseWriter(server string, req *dns.Msg, do, cd bool, c *Cache) *ResponseWriter {
req = req.Copy()
req.AuthenticatedData = true
cw := &ResponseWriter{ cw := &ResponseWriter{
Cache: c, Cache: c,
server: server, server: server,
do: do, do: do,
cd: cd, cd: cd,
ad: true,
prefetch: true, prefetch: true,
remoteAddr: prefetchAddr, remoteAddr: prefetchAddr,
} }
@@ -435,13 +438,6 @@ func (w *ResponseWriter) WriteMsg(res *dns.Msg) error {
res.Ns = filterRRSlice(res.Ns, ttl, false) res.Ns = filterRRSlice(res.Ns, ttl, false)
res.Extra = filterRRSlice(res.Extra, ttl, false) res.Extra = filterRRSlice(res.Extra, ttl, false)
if !w.do && !w.ad {
// unset AD bit if requester is not OK with DNSSEC
// But retain AD bit if requester set the AD bit in the request, per RFC6840 5.7-5.8
res.AuthenticatedData = false
}
w.lastResponse = res.Copy()
if hasKey && duration > 0 { if hasKey && duration > 0 {
if w.state.Match(res) { if w.state.Match(res) {
w.set(res, key, mt, duration) w.set(res, key, mt, duration)
@@ -453,6 +449,13 @@ func (w *ResponseWriter) WriteMsg(res *dns.Msg) error {
} }
} }
if !w.do && !w.ad {
// unset AD bit if requester is not OK with DNSSEC
// But retain AD bit if requester set the AD bit in the request, per RFC6840 5.7-5.8
res.AuthenticatedData = false
}
w.lastResponse = res.Copy()
if w.prefetch { if w.prefetch {
return nil return nil
} }

View File

@@ -36,9 +36,10 @@ func (c *Cache) ServeDNS(ctx context.Context, w dns.ResponseWriter, r *dns.Msg)
i := c.getIfNotStale(now, state, server) i := c.getIfNotStale(now, state, server)
if i == nil { if i == nil {
crr := &ResponseWriter{ResponseWriter: w, Cache: c, state: state, server: server, do: do, ad: ad, cd: cd, refreshState := authenticatedRefreshState(state)
crr := &ResponseWriter{ResponseWriter: w, Cache: c, state: refreshState, server: server, do: do, ad: ad, cd: cd,
nexcept: c.nexcept, pexcept: c.pexcept, wildcardFunc: wildcardFunc(ctx)} nexcept: c.nexcept, pexcept: c.pexcept, wildcardFunc: wildcardFunc(ctx)}
return c.doRefresh(ctx, state, crr) return c.doRefresh(ctx, refreshState, crr)
} }
ttl := i.ttl(now) ttl := i.ttl(now)
stale := ttl <= 0 stale := ttl <= 0
@@ -49,7 +50,8 @@ func (c *Cache) ServeDNS(ctx context.Context, w dns.ResponseWriter, r *dns.Msg)
nowFunc := c.now nowFunc := c.now
trackRefresh := failureRecheck > 0 trackRefresh := failureRecheck > 0
if !trackRefresh || i.beginRefresh(now, failureRecheck) { if !trackRefresh || i.beginRefresh(now, failureRecheck) {
crr := &ResponseWriter{ResponseWriter: w, Cache: c, state: state, server: server, do: do, cd: cd} refreshState := authenticatedRefreshState(state)
crr := &ResponseWriter{ResponseWriter: w, Cache: c, state: refreshState, server: server, do: do, ad: ad, cd: cd}
if c.verifyStaleTimeout > 0 { if c.verifyStaleTimeout > 0 {
// Background verify: cache the response but do not write to the wire. // Background verify: cache the response but do not write to the wire.
// On timeout, we serve the stale entry below and let the goroutine continue. // On timeout, we serve the stale entry below and let the goroutine continue.
@@ -57,14 +59,14 @@ func (c *Cache) ServeDNS(ctx context.Context, w dns.ResponseWriter, r *dns.Msg)
} }
cw := newVerifyStaleResponseWriter(crr) cw := newVerifyStaleResponseWriter(crr)
if c.verifyStaleTimeout == 0 { if c.verifyStaleTimeout == 0 {
ret, err := c.doRefresh(ctx, state, cw) ret, err := c.doRefresh(ctx, refreshState, cw)
if trackRefresh { if trackRefresh {
i.endRefresh(nowFunc(), failureRecheck, cw.refreshed) i.endRefresh(nowFunc(), failureRecheck, cw.refreshed)
} }
if cw.refreshed { if cw.refreshed {
return ret, err return ret, err
} }
} else if served, ret, err := c.verifyWithTimeout(ctx, state, w, cw, r, do, ad, i, failureRecheck, nowFunc); served { } else if served, ret, err := c.verifyWithTimeout(ctx, refreshState, w, cw, r, do, ad, i, failureRecheck, nowFunc); served {
return ret, err return ret, err
} }
} }
@@ -220,6 +222,12 @@ func (c *Cache) shouldPrefetch(i *item, now time.Time) bool {
return i.Hits() >= c.prefetch && i.ttl(now) <= threshold return i.Hits() >= c.prefetch && i.ttl(now) <= threshold
} }
func authenticatedRefreshState(state request.Request) request.Request {
state.Req = state.Req.Copy()
state.Req.AuthenticatedData = true
return state
}
// Name implements the Handler interface. // Name implements the Handler interface.
func (c *Cache) Name() string { return "cache" } func (c *Cache) Name() string { return "cache" }