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) + } +}