diff --git a/plugin/tsig/tsig.go b/plugin/tsig/tsig.go index e240f856a..9ab67287b 100644 --- a/plugin/tsig/tsig.go +++ b/plugin/tsig/tsig.go @@ -32,6 +32,14 @@ func (t TSIGServer) Name() string { return pluginName } // ServeDNS implements plugin.Handler func (t *TSIGServer) ServeDNS(ctx context.Context, w dns.ResponseWriter, r *dns.Msg) (int, error) { + for i, rr := range r.Extra { + if rr.Header().Rrtype == dns.TypeTSIG && i != len(r.Extra)-1 { + resp := new(dns.Msg).SetRcode(r, dns.RcodeFormatError) + w.WriteMsg(resp) + return dns.RcodeSuccess, nil + } + } + var ( state = request.Request{Req: r, W: w} tsigRR = r.IsTsig() diff --git a/plugin/tsig/tsig_test.go b/plugin/tsig/tsig_test.go index 44a109994..29957499f 100644 --- a/plugin/tsig/tsig_test.go +++ b/plugin/tsig/tsig_test.go @@ -216,6 +216,37 @@ func TestServeDNS(t *testing.T) { } } +func TestServeDNSRejectsNonFinalTSIG(t *testing.T) { + nextCalled := false + tsig := TSIGServer{ + Zones: []string{"."}, + Next: test.HandlerFunc(func(ctx context.Context, w dns.ResponseWriter, r *dns.Msg) (int, error) { + nextCalled = true + return testHandler()(ctx, w, r) + }), + } + + r := new(dns.Msg) + r.SetQuestion("test.example.", dns.TypeA) + r.SetTsig("test.key.", dns.HmacSHA256, 300, time.Now().Unix()) + r.Extra = append(r.Extra, test.OPT(42, true)) + + w := dnstest.NewRecorder(&test.ResponseWriter{}) + if _, err := tsig.ServeDNS(context.Background(), w, r); err != nil { + t.Fatal(err) + } + + if nextCalled { + t.Error("expected non-final TSIG to be rejected before calling the next plugin") + } + if w.Msg == nil { + t.Fatal("expected a FORMERR response") + } + if w.Msg.Rcode != dns.RcodeFormatError { + t.Errorf("expected FORMERR, got %s", dns.RcodeToString[w.Msg.Rcode]) + } +} + func TestServeDNSTsigErrors(t *testing.T) { clientNow := time.Now().Unix()