Merge commit from fork

Ensure DoH, DoH3, DoQ, and DNS-over-gRPC continue to reject RFC 2136
UPDATE requests before dispatch. This guards the common request
acceptance policy across every affected transport.

Signed-off-by: Ville Vesilehto <ville@vesilehto.fi>
This commit is contained in:
Ville Vesilehto
2026-07-19 06:36:05 +03:00
committed by GitHub
parent c1fe47bc3a
commit 57f73b4a43
5 changed files with 189 additions and 0 deletions

View File

@@ -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 {

View File

@@ -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{

View File

@@ -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{

View File

@@ -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.

View File

@@ -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