From d5c1188843553dd231f79a66a0bcb0ddc0168a7c Mon Sep 17 00:00:00 2001 From: Yong Tang Date: Tue, 22 Sep 2026 06:23:20 -0700 Subject: [PATCH] plugin/rewrite: Limit rewrite cname recursion (#8570) This PR limit rewrite cname recursion with the existing DNS server loop counter. The issu was that rewrite cname can recurse indefinitely through internal lookups, causing crash at the end Signed-off-by: Yong Tang --- plugin/rewrite/cname_target.go | 13 +++++- plugin/rewrite/cname_target_test.go | 70 +++++++++++++++++++++++++++++ 2 files changed, 82 insertions(+), 1 deletion(-) diff --git a/plugin/rewrite/cname_target.go b/plugin/rewrite/cname_target.go index 5623e3a2b..6b18193e9 100644 --- a/plugin/rewrite/cname_target.go +++ b/plugin/rewrite/cname_target.go @@ -7,6 +7,7 @@ import ( "strconv" "strings" + "github.com/coredns/coredns/core/dnsserver" "github.com/coredns/coredns/plugin" "github.com/coredns/coredns/plugin/pkg/log" "github.com/coredns/coredns/plugin/pkg/upstream" @@ -15,6 +16,8 @@ import ( "github.com/miekg/dns" ) +const maxCNAMERewriteDepth = 8 + // UpstreamInt wraps the Upstream API for dependency injection during testing type UpstreamInt interface { Lookup(ctx context.Context, state request.Request, name string, typ uint16) (*dns.Msg, error) @@ -82,12 +85,20 @@ func (r *cnameTargetRuleWithReqState) RewriteResponse(res *dns.Msg, rr dns.RR) { if cname.Target != fromTarget { return } + + // Limit internal lookups that re-enter the server. + loop, _ := r.ctx.Value(dnsserver.LoopKey{}).(int) + if loop > maxCNAMERewriteDepth { + return + } + ctx := context.WithValue(r.ctx, dnsserver.LoopKey{}, loop+1) + // create upstream request with the new target with the same qtype r.state.Req.Question[0].Name = toTarget // upRes can be nil if the internal query path didn't write a response // (e.g. a plugin returned a success rcode without writing, dropped the query, // or the context was canceled). Guard upRes before dereferencing. - upRes, err := r.rule.Upstream.Lookup(r.ctx, r.state, toTarget, r.state.Req.Question[0].Qtype) + upRes, err := r.rule.Upstream.Lookup(ctx, r.state, toTarget, r.state.Req.Question[0].Qtype) if err != nil { log.Errorf("upstream lookup failed: %v", err) return diff --git a/plugin/rewrite/cname_target_test.go b/plugin/rewrite/cname_target_test.go index 543fc40fb..de5101382 100644 --- a/plugin/rewrite/cname_target_test.go +++ b/plugin/rewrite/cname_target_test.go @@ -314,3 +314,73 @@ func TestNewCNAMERuleNormalization(t *testing.T) { t.Errorf("expected toTarget to be normalized to 'vpce-123.amazonaws.com.', got %q", cnameRule.paramToTarget) } } + +type cnameLoopBackend struct{} + +func (cnameLoopBackend) Name() string { return "cname-loop" } + +func (cnameLoopBackend) ServeDNS(_ context.Context, w dns.ResponseWriter, r *dns.Msg) (int, error) { + m := new(dns.Msg) + m.SetReply(r) + m.Authoritative = true + + switch r.Question[0].Name { + case "victim.poc.internal.", "target.poc.internal.": + m.Answer = []dns.RR{ + test.CNAME(r.Question[0].Name + " 500 IN CNAME loop.poc.internal."), + test.A("loop.poc.internal. 500 IN A 192.0.2.1"), + } + } + + w.WriteMsg(m) + return dns.RcodeSuccess, nil +} + +type reentrantCNAMEUpstream struct { + top plugin.Handler + calls int + maxCalls int + exceeded bool +} + +func (u *reentrantCNAMEUpstream) Lookup(ctx context.Context, state request.Request, name string, typ uint16) (*dns.Msg, error) { + u.calls++ + if u.calls > u.maxCalls { + u.exceeded = true + return nil, errors.New("test recursion limit exceeded") + } + + req := state.NewWithQuestion(name, typ) + rec := dnstest.NewRecorder(state.W) + _, err := u.top.ServeDNS(ctx, rec, req.Req) + return rec.Msg, err +} + +func TestCNAMETargetRewriteLoop(t *testing.T) { + rule, err := newCNAMERule(Stop, ExactMatch, "loop.poc.internal.", "target.poc.internal.") + if err != nil { + t.Fatalf("newCNAMERule failed: %v", err) + } + + rw := &Rewrite{ + Next: cnameLoopBackend{}, + Rules: []Rule{rule}, + } + const maxExpectedLookups = 9 + upstream := &reentrantCNAMEUpstream{top: rw, maxCalls: maxExpectedLookups + 1} + rule.(*cnameTargetRule).Upstream = upstream + + req := new(dns.Msg) + req.SetQuestion("victim.poc.internal.", dns.TypeA) + rec := dnstest.NewRecorder(&test.ResponseWriter{}) + if _, err := rw.ServeDNS(context.Background(), rec, req); err != nil { + t.Fatalf("ServeDNS returned error: %v", err) + } + + if upstream.exceeded { + t.Fatalf("rewrite cname recursion exceeded %d internal lookups", upstream.maxCalls) + } + if upstream.calls > maxExpectedLookups { + t.Fatalf("rewrite cname recursion was not bounded: got %d internal lookups", upstream.calls) + } +}