diff --git a/plugin/rewrite/reverter.go b/plugin/rewrite/reverter.go index a99a6b076..8350b32c4 100644 --- a/plugin/rewrite/reverter.go +++ b/plugin/rewrite/reverter.go @@ -1,6 +1,8 @@ package rewrite import ( + "fmt" + "github.com/miekg/dns" ) @@ -84,6 +86,9 @@ func NewResponseReverter(w dns.ResponseWriter, r *dns.Msg, policy RevertPolicy) // WriteMsg records the status code and calls the underlying ResponseWriter's WriteMsg method. func (r *ResponseReverter) WriteMsg(res1 *dns.Msg) error { + if res1 == nil { + return fmt.Errorf("rewrite: response message is nil") + } // Deep copy 'res' as to not (e.g). rewrite a message that's also stored in the cache. res := res1.Copy() diff --git a/plugin/rewrite/reverter_test.go b/plugin/rewrite/reverter_test.go index b235fed3c..8800c8549 100644 --- a/plugin/rewrite/reverter_test.go +++ b/plugin/rewrite/reverter_test.go @@ -370,3 +370,22 @@ func noQuestionMsgPrinter(_ context.Context, w dns.ResponseWriter, _ *dns.Msg) ( return dns.RcodeSuccess, nil } + +func TestResponseReverterWriteMsgNilResponse(t *testing.T) { + req := new(dns.Msg) + req.SetQuestion("service.example.org.", dns.TypeA) + + rec := dnstest.NewRecorder(&test.ResponseWriter{}) + rw := NewResponseReverter(rec, req, NewRevertPolicy(false, false)) + + defer func() { + if r := recover(); r != nil { + t.Fatalf("ResponseReverter.WriteMsg panicked on nil response: %v", r) + } + }() + + err := rw.WriteMsg(nil) + if err == nil { + t.Error("Expected error when passing nil response to WriteMsg, got nil") + } +}