diff --git a/core/dnsserver/server_grpc_test.go b/core/dnsserver/server_grpc_test.go index 3a74b1677..65424e6c5 100644 --- a/core/dnsserver/server_grpc_test.go +++ b/core/dnsserver/server_grpc_test.go @@ -278,6 +278,31 @@ func TestServergRPC_Query(t *testing.T) { } } +func TestServergRPC_QueryRejectsUpdate(t *testing.T) { + handler := new(updateResponsePlugin) + server, err := NewServergRPC("127.0.0.1:0", []*Config{ + testConfig("grpc", handler), + }) + if err != nil { + t.Fatalf("NewServergRPC() failed: %v", err) + } + + tcpAddr, err := net.ResolveTCPAddr("tcp", "127.0.0.1:12345") + if err != nil { + t.Fatalf("net.ResolveTCPAddr() failed: %v", err) + } + server.listenAddr = tcpAddr + ctx := peer.NewContext(context.Background(), &peer.Peer{Addr: tcpAddr}) + + _, err = server.Query(ctx, &pb.DnsPacket{Msg: mustPackRFC2136Update(t)}) + if err == nil { + t.Fatal("Query() accepted an RFC 2136 UPDATE") + } + if handler.called.Load() { + t.Fatal("RFC 2136 UPDATE reached the plugin chain") + } +} + func TestServergRPC_Query_ErrorCases(t *testing.T) { server, err := NewServergRPC("127.0.0.1:0", []*Config{testConfig("grpc", testPlugin{})}) if err != nil { diff --git a/core/dnsserver/server_https3_test.go b/core/dnsserver/server_https3_test.go index 182e2d995..50608c1bc 100644 --- a/core/dnsserver/server_https3_test.go +++ b/core/dnsserver/server_https3_test.go @@ -73,6 +73,30 @@ func TestCustomHTTP3RequestValidator(t *testing.T) { } } +func TestServerHTTPS3RejectsUpdate(t *testing.T) { + handler := new(updateResponsePlugin) + config := testConfig("https3", handler) + config.TLSConfig = &tls.Config{} + + server, err := NewServerHTTPS3("127.0.0.1:443", []*Config{config}) + if err != nil { + t.Fatalf("NewServerHTTPS3() failed: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/dns-query", bytes.NewReader(mustPackRFC2136Update(t))) + req.RemoteAddr = "127.0.0.1:12345" + recorder := httptest.NewRecorder() + + server.ServeHTTP(recorder, req) + + if recorder.Code != http.StatusBadRequest { + t.Fatalf("ServeHTTP() status = %d, want %d", recorder.Code, http.StatusBadRequest) + } + if handler.called.Load() { + t.Fatal("RFC 2136 UPDATE reached the plugin chain") + } +} + func TestNewServerHTTPS3WithCustomLimits(t *testing.T) { maxStreams := 50 c := Config{ diff --git a/core/dnsserver/server_https_test.go b/core/dnsserver/server_https_test.go index a27e4642b..9c8b36bfc 100644 --- a/core/dnsserver/server_https_test.go +++ b/core/dnsserver/server_https_test.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "crypto/tls" + "encoding/base64" "errors" "io" "net" @@ -81,6 +82,42 @@ func TestCustomHTTPRequestValidator(t *testing.T) { } } +func TestServerHTTPSRejectsUpdate(t *testing.T) { + for _, method := range []string{http.MethodGet, http.MethodPost} { + t.Run(method, func(t *testing.T) { + handler := new(updateResponsePlugin) + config := testConfig("https", handler) + config.TLSConfig = &tls.Config{} + + server, err := NewServerHTTPS("127.0.0.1:443", []*Config{config}) + if err != nil { + t.Fatalf("NewServerHTTPS() failed: %v", err) + } + + wire := mustPackRFC2136Update(t) + target := "/dns-query" + var body io.Reader + if method == http.MethodGet { + target += "?dns=" + base64.RawURLEncoding.EncodeToString(wire) + } else { + body = bytes.NewReader(wire) + } + req := httptest.NewRequest(method, target, body) + req.RemoteAddr = "127.0.0.1:12345" + recorder := httptest.NewRecorder() + + server.ServeHTTP(recorder, req) + + if recorder.Code != http.StatusBadRequest { + t.Fatalf("ServeHTTP() status = %d, want %d", recorder.Code, http.StatusBadRequest) + } + if handler.called.Load() { + t.Fatal("RFC 2136 UPDATE reached the plugin chain") + } + }) + } +} + func TestNewServerHTTPSWithCustomLimits(t *testing.T) { maxConnections := 100 c := Config{ diff --git a/core/dnsserver/server_quic_test.go b/core/dnsserver/server_quic_test.go index 2c198dbf2..bcbb269ed 100644 --- a/core/dnsserver/server_quic_test.go +++ b/core/dnsserver/server_quic_test.go @@ -553,6 +553,73 @@ func TestServerQUIC_ServeQUIC_TSIGBadSigSetsTsigStatus(t *testing.T) { } } +func TestServerQUIC_ServeQUICRejectsUpdate(t *testing.T) { + handler := new(updateResponsePlugin) + config := testConfig("quic", handler) + config.TLSConfig = mustMakeQUICServerTLSConfig(t) + + server, err := NewServerQUIC(transport.QUIC+"://127.0.0.1:0", []*Config{config}) + if err != nil { + t.Fatalf("NewServerQUIC() failed: %v", err) + } + + pc, err := server.ListenPacket() + if err != nil { + t.Fatalf("ListenPacket() failed: %v", err) + } + defer pc.Close() + + serveErrCh := make(chan error, 1) + go func() { + serveErrCh <- server.ServeQUIC() + }() + defer func() { + _ = server.Stop() + select { + case <-serveErrCh: + case <-time.After(2 * time.Second): + } + }() + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + conn, err := quic.DialAddr(ctx, pc.LocalAddr().String(), mustMakeQUICClientTLSConfig(), &quic.Config{}) + if err != nil { + t.Fatalf("quic.DialAddr() failed: %v", err) + } + defer conn.CloseWithError(DoQCodeNoError, "") + + stream, err := conn.OpenStreamSync(ctx) + if err != nil { + t.Fatalf("OpenStreamSync() failed: %v", err) + } + if err := stream.SetReadDeadline(time.Now().Add(5 * time.Second)); err != nil { + t.Fatalf("SetReadDeadline() failed: %v", err) + } + if _, err := stream.Write(AddPrefix(mustPackRFC2136Update(t))); err != nil { + t.Fatalf("stream.Write() failed: %v", err) + } + if err := stream.Close(); err != nil { + t.Fatalf("stream.Close() failed: %v", err) + } + + _, err = readDOQMessage(stream) + if err == nil { + t.Fatal("DoQ server accepted an RFC 2136 UPDATE") + } + var applicationErr *quic.ApplicationError + if !errors.As(err, &applicationErr) { + t.Fatalf("readDOQMessage() error = %T %v, want QUIC application error", err, err) + } + if applicationErr.ErrorCode != DoQCodeProtocolError { + t.Fatalf("QUIC application error code = %d, want %d", applicationErr.ErrorCode, DoQCodeProtocolError) + } + if handler.called.Load() { + t.Fatal("RFC 2136 UPDATE reached the plugin chain") + } +} + // echoPlugin answers every query with a minimal reply. It is used as a // negative control to prove a normal DoQ query is still served after the // per-stream read deadline was introduced. diff --git a/core/dnsserver/server_test.go b/core/dnsserver/server_test.go index cfb50ed0c..86c381c28 100644 --- a/core/dnsserver/server_test.go +++ b/core/dnsserver/server_test.go @@ -24,6 +24,42 @@ func (tp testPlugin) ServeDNS(_ctx context.Context, _w dns.ResponseWriter, _r *d func (tp testPlugin) Name() string { return "local" } +type updateResponsePlugin struct { + called atomic.Bool +} + +func (p *updateResponsePlugin) Name() string { return "update-response" } + +func (p *updateResponsePlugin) ServeDNS(_ context.Context, w dns.ResponseWriter, r *dns.Msg) (int, error) { + p.called.Store(true) + + m := new(dns.Msg) + m.SetReply(r) + if err := w.WriteMsg(m); err != nil { + return dns.RcodeServerFailure, err + } + return dns.RcodeSuccess, nil +} + +func mustPackRFC2136Update(t *testing.T) []byte { + t.Helper() + + m := new(dns.Msg).SetUpdate("example.com.") + rr, err := dns.NewRR("foo.example.com. 300 IN A 192.0.2.123") + if err != nil { + t.Fatalf("dns.NewRR() failed: %v", err) + } + m.Insert([]dns.RR{rr}) + // DNS-over-QUIC requires the DNS message ID to be zero. + m.Id = 0 + + wire, err := m.Pack() + if err != nil { + t.Fatalf("dns.Msg.Pack() failed: %v", err) + } + return wire +} + // blockingPlugin uses sync.Mutex to simulate extended processing. type blockingPlugin struct { sync.Mutex