diff --git a/plugin/rewrite/edns0.go b/plugin/rewrite/edns0.go index ddaaf97e5..1c84563d2 100644 --- a/plugin/rewrite/edns0.go +++ b/plugin/rewrite/edns0.go @@ -44,6 +44,8 @@ type edns0SetResponseRule struct { code uint16 } +func (r *edns0SetResponseRule) revertRequestExtra() {} + func (r *edns0SetResponseRule) RewriteResponse(res *dns.Msg, _ dns.RR) { ednsOpt := res.IsEdns0() if ednsOpt == nil { @@ -62,6 +64,8 @@ type edns0ReplaceResponseRule[T dns.EDNS0] struct { source T } +func (r *edns0ReplaceResponseRule[T]) revertRequestExtra() {} + func (r *edns0ReplaceResponseRule[T]) RewriteResponse(res *dns.Msg, _ dns.RR) { ednsOpt := res.IsEdns0() if ednsOpt == nil { diff --git a/plugin/rewrite/reverter.go b/plugin/rewrite/reverter.go index 3204263ff..16811076c 100644 --- a/plugin/rewrite/reverter.go +++ b/plugin/rewrite/reverter.go @@ -44,6 +44,11 @@ type ResponseRule interface { RewriteResponse(res *dns.Msg, rr dns.RR) } +type requestExtraRevertRule interface { + ResponseRule + revertRequestExtra() +} + // ResponseRules describes an ordered list of response rules to apply // after a name rewrite type ResponseRules = []ResponseRule @@ -54,6 +59,7 @@ type ResponseRules = []ResponseRule type ResponseReverter struct { dns.ResponseWriter originalQuestion dns.Question + request *dns.Msg ResponseRules ResponseRules revertPolicy RevertPolicy } @@ -63,6 +69,7 @@ func NewResponseReverter(w dns.ResponseWriter, r *dns.Msg, policy RevertPolicy) return &ResponseReverter{ ResponseWriter: w, originalQuestion: r.Question[0], + request: r, revertPolicy: policy, } } @@ -90,9 +97,48 @@ func (r *ResponseReverter) WriteMsg(res1 *dns.Msg) error { r.rewriteResourceRecord(res, rr) } } + return r.writeWithRevertedRequestExtra(res) +} + +func (r *ResponseReverter) writeWithRevertedRequestExtra(res *dns.Msg) error { + if r.request == nil || !r.hasRequestExtraRevertRule() { + return r.ResponseWriter.WriteMsg(res) + } + + currentExtra := r.request.Extra + req := new(dns.Msg) + req.Extra = copyRRs(currentExtra) + for _, rr := range req.Extra { + r.rewriteRequestExtra(req, rr) + } + r.request.Extra = req.Extra + defer func() { + r.request.Extra = currentExtra + }() + return r.ResponseWriter.WriteMsg(res) } +func (r *ResponseReverter) hasRequestExtraRevertRule() bool { + for _, rule := range r.ResponseRules { + if _, ok := rule.(requestExtraRevertRule); ok { + return true + } + } + return false +} + +func copyRRs(rrs []dns.RR) []dns.RR { + if len(rrs) == 0 { + return nil + } + copied := make([]dns.RR, len(rrs)) + for i, rr := range rrs { + copied[i] = dns.Copy(rr) + } + return copied +} + func (r *ResponseReverter) rewriteResourceRecord(res *dns.Msg, rr dns.RR) { // The reverting rules need to be done in reversed order. for i := len(r.ResponseRules) - 1; i >= 0; i-- { @@ -100,6 +146,17 @@ func (r *ResponseReverter) rewriteResourceRecord(res *dns.Msg, rr dns.RR) { } } +func (r *ResponseReverter) rewriteRequestExtra(req *dns.Msg, rr dns.RR) { + // The reverting rules need to be done in reversed order. + for i := len(r.ResponseRules) - 1; i >= 0; i-- { + rule, ok := r.ResponseRules[i].(requestExtraRevertRule) + if !ok { + continue + } + rule.RewriteResponse(req, rr) + } +} + // Write is a wrapper that records the size of the message that gets written. func (r *ResponseReverter) Write(buf []byte) (int, error) { n, err := r.ResponseWriter.Write(buf) diff --git a/plugin/rewrite/rewrite_test.go b/plugin/rewrite/rewrite_test.go index cf51672ca..b24801132 100644 --- a/plugin/rewrite/rewrite_test.go +++ b/plugin/rewrite/rewrite_test.go @@ -1153,8 +1153,56 @@ func TestRewriteEDNS0Unset(t *testing.T) { rec := dnstest.NewRecorder(&test.ResponseWriter{}) rw.ServeDNS(ctx, rec, m) - if !optsEqual(o.Option, tc.toOpts) { - t.Errorf("Test %d: Expected %v but got %v", i, tc.toOpts, o) + respOpt := rec.Msg.IsEdns0() + if respOpt == nil { + t.Errorf("Test %d: EDNS0 options not set", i) + continue + } + if !optsEqual(respOpt.Option, tc.toOpts) { + t.Errorf("Test %d: Expected %v but got %v", i, tc.toOpts, respOpt) } } } + +func TestRewriteEDNS0RevertDoesNotLeakThroughScrubWriter(t *testing.T) { + rw := Rewrite{ + Next: plugin.HandlerFunc(func(_ctx context.Context, w dns.ResponseWriter, r *dns.Msg) (int, error) { + resp := new(dns.Msg) + resp.SetReply(r) + return 0, w.WriteMsg(resp) + }), + RevertPolicy: NewRevertPolicy(false, false), + } + + r, err := newEdns0Rule("stop", "local", "set", "0xffee", "0xabcdef", "revert") + if err != nil { + t.Fatalf("Error creating test rule: %s", err) + } + rw.Rules = []Rule{r} + + m := new(dns.Msg) + m.SetQuestion("example.com.", dns.TypeA) + m.SetEdns0(4096, false) + m.IsEdns0().Option = append(m.IsEdns0().Option, &dns.EDNS0_COOKIE{Code: dns.EDNS0COOKIE, Cookie: "abcdef0123456789"}) + + rec := dnstest.NewRecorder(&test.ResponseWriter{}) + scrub := request.NewScrubWriter(m, rec) + rw.ServeDNS(context.TODO(), scrub, m) + + o := rec.Msg.IsEdns0() + if o == nil { + t.Fatal("expected EDNS0 option record in response") + } + var foundCookie bool + for _, opt := range o.Option { + if opt.Option() == 0xffee { + t.Fatalf("expected rewritten EDNS0 option to be reverted, got %v", o.Option) + } + if opt.Option() == dns.EDNS0COOKIE { + foundCookie = true + } + } + if !foundCookie { + t.Fatalf("expected original EDNS0 cookie option to be preserved, got %v", o.Option) + } +}