mirror of
https://github.com/Control-D-Inc/ctrld.git
synced 2026-07-29 01:18:48 +02:00
TestDoHResolve_{OversizedBody_Rejected,NonOKStatus_BoundedErrorBody,
OversizedBody_DoH3} asserted how many bytes the test server managed to
write before the client tore down the connection. That count reflects
kernel socket send buffers and HTTP/2 flow-control windows, which vary
by OS and load, so the server could buffer the whole body before
teardown and fail the assertion. It flaked on the Windows CI runner, but
reproduces on Linux too.
Replace the server-side byte counter with a deterministic synchronization
point. The handler writes exactly the read cap (dohMaxResponseSize+1 for
the body, dohMaxErrorBodySize for the error path), flushes, then blocks
without ever returning, so the response stream never gets an EOF. The
test then requires Resolve to return the size/status error before the
handler is released: ctrld's bounded read (io.LimitReader) returns after
the capped prefix, while a read to EOF would block on the withheld stream
and trip the deadline.
This removes the socket-buffer timing dependence and, unlike asserting on
the returned error alone, still fails if the caps are removed -- verified
by reverting both reads in doh.go to io.ReadAll(resp.Body), which makes
all three tests time out.
482 lines
14 KiB
Go
482 lines
14 KiB
Go
package ctrld
|
|
|
|
import (
|
|
"context"
|
|
"crypto/tls"
|
|
"crypto/x509"
|
|
"crypto/x509/pkix"
|
|
"errors"
|
|
"net"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"net/url"
|
|
"runtime"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/miekg/dns"
|
|
"github.com/quic-go/quic-go/http3"
|
|
)
|
|
|
|
func Test_dohOsHeaderValue(t *testing.T) {
|
|
val := dohOsHeaderValue
|
|
if val == "" {
|
|
t.Fatalf("empty %s", dohOsHeader)
|
|
}
|
|
t.Log(val)
|
|
|
|
encodedOs := EncodeOsNameMap[runtime.GOOS]
|
|
if encodedOs == "" {
|
|
t.Fatalf("missing encoding value for: %q", runtime.GOOS)
|
|
}
|
|
decodedOs := DecodeOsNameMap[encodedOs]
|
|
if decodedOs == "" {
|
|
t.Fatalf("missing decoding value for: %q", runtime.GOOS)
|
|
}
|
|
}
|
|
|
|
func Test_wrapUrlError(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
err error
|
|
wantErr string
|
|
}{
|
|
{
|
|
name: "No wrapping for non-URL errors",
|
|
err: errors.New("plain error"),
|
|
wantErr: "plain error",
|
|
},
|
|
{
|
|
name: "URL error without TLS error",
|
|
err: &url.Error{
|
|
Op: "Get",
|
|
URL: "https://example.com",
|
|
Err: errors.New("underlying error"),
|
|
},
|
|
wantErr: "Get \"https://example.com\": underlying error",
|
|
},
|
|
{
|
|
name: "TLS error with missing unverified certificate data",
|
|
err: &url.Error{
|
|
Op: "Get",
|
|
URL: "https://example.com",
|
|
Err: &tls.CertificateVerificationError{
|
|
UnverifiedCertificates: nil,
|
|
Err: &x509.UnknownAuthorityError{},
|
|
},
|
|
},
|
|
wantErr: `Get "https://example.com": tls: failed to verify certificate: x509: certificate signed by unknown authority`,
|
|
},
|
|
{
|
|
name: "TLS error with valid certificate data",
|
|
err: &url.Error{
|
|
Op: "Get",
|
|
URL: "https://example.com",
|
|
Err: &tls.CertificateVerificationError{
|
|
UnverifiedCertificates: []*x509.Certificate{
|
|
{
|
|
Subject: pkix.Name{
|
|
CommonName: "BadSubjectCN",
|
|
Organization: []string{"BadSubjectOrg"},
|
|
},
|
|
Issuer: pkix.Name{
|
|
CommonName: "BadIssuerCN",
|
|
Organization: []string{"BadIssuerOrg"},
|
|
},
|
|
},
|
|
},
|
|
Err: &x509.UnknownAuthorityError{},
|
|
},
|
|
},
|
|
wantErr: `Get "https://example.com": tls: failed to verify certificate: x509: certificate signed by unknown authority: BadSubjectCN, BadSubjectOrg, BadIssuerOrg`,
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
gotErr := wrapUrlError(tt.err)
|
|
if gotErr.Error() != tt.wantErr {
|
|
t.Errorf("wrapCertificateVerificationError() error = %v, want %v", gotErr, tt.wantErr)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func Test_ClientCertificateVerificationError(t *testing.T) {
|
|
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", "application/dns-message")
|
|
})
|
|
tlsServer, cert := testTLSServer(t, handler)
|
|
tlsServerUrl, err := url.Parse(tlsServer.URL)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
quicServer := newTestQUICServer(t)
|
|
http3Server := newTestHTTP3Server(t, handler)
|
|
|
|
tests := []struct {
|
|
name string
|
|
uc *UpstreamConfig
|
|
}{
|
|
{
|
|
"doh",
|
|
&UpstreamConfig{
|
|
Name: "doh",
|
|
Type: ResolverTypeDOH,
|
|
Endpoint: tlsServer.URL,
|
|
Timeout: 1000,
|
|
},
|
|
},
|
|
{
|
|
"doh3",
|
|
&UpstreamConfig{
|
|
Name: "doh3",
|
|
Type: ResolverTypeDOH3,
|
|
Endpoint: http3Server.addr,
|
|
Timeout: 5000,
|
|
},
|
|
},
|
|
{
|
|
"doq",
|
|
&UpstreamConfig{
|
|
Name: "doq",
|
|
Type: ResolverTypeDOQ,
|
|
Endpoint: quicServer.addr,
|
|
Timeout: 5000,
|
|
},
|
|
},
|
|
{
|
|
"dot",
|
|
&UpstreamConfig{
|
|
Name: "dot",
|
|
Type: ResolverTypeDOT,
|
|
Endpoint: net.JoinHostPort(tlsServerUrl.Hostname(), tlsServerUrl.Port()),
|
|
Timeout: 1000,
|
|
},
|
|
},
|
|
}
|
|
|
|
ctx := context.Background()
|
|
for _, tc := range tests {
|
|
tc := tc
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
t.Parallel()
|
|
tc.uc.Init(ctx)
|
|
tc.uc.SetupBootstrapIP(ctx)
|
|
r, err := NewResolver(ctx, tc.uc)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
msg := new(dns.Msg)
|
|
msg.SetQuestion("verify.controld.com.", dns.TypeA)
|
|
msg.RecursionDesired = true
|
|
_, err = r.Resolve(ctx, msg)
|
|
// Verify the error contains the expected certificate information
|
|
if err == nil {
|
|
t.Fatal("expected certificate verification error, got nil")
|
|
}
|
|
|
|
// You can check the error contains information about the test certificate
|
|
if !strings.Contains(err.Error(), cert.Issuer.CommonName) {
|
|
t.Fatalf("error should contain issuer information %q, got: %v", cert.Issuer.CommonName, err)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
// testTLSServer creates an HTTPS test server with a self-signed certificate
|
|
// returns the server and its certificate for verification testing
|
|
// testTLSServer creates an HTTPS test server with a self-signed certificate
|
|
func testTLSServer(t *testing.T, handler http.Handler) (*httptest.Server, *x509.Certificate) {
|
|
t.Helper()
|
|
|
|
testCert := generateTestCertificate(t)
|
|
|
|
// Create a test server
|
|
server := httptest.NewUnstartedServer(handler)
|
|
server.TLS = &tls.Config{
|
|
Certificates: []tls.Certificate{testCert.tlsCert},
|
|
MinVersion: tls.VersionTLS12,
|
|
}
|
|
server.StartTLS()
|
|
|
|
// Add cleanup
|
|
t.Cleanup(server.Close)
|
|
|
|
return server, testCert.cert
|
|
}
|
|
|
|
// testHTTP3Server represents a structure for an HTTP/3 test server with its server instance, TLS certificate, and address.
|
|
type testHTTP3Server struct {
|
|
server *http3.Server
|
|
cert *x509.Certificate
|
|
addr string
|
|
}
|
|
|
|
// newTestHTTP3Server creates and starts a test HTTP/3 server with a given handler and returns the server instance.
|
|
func newTestHTTP3Server(t *testing.T, handler http.Handler) *testHTTP3Server {
|
|
t.Helper()
|
|
|
|
testCert := generateTestCertificate(t)
|
|
|
|
// First create a listener to get the actual port
|
|
udpAddr := &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 0}
|
|
udpConn, err := net.ListenUDP("udp", udpAddr)
|
|
if err != nil {
|
|
t.Fatalf("failed to create UDP listener: %v", err)
|
|
}
|
|
|
|
// Get the actual address
|
|
actualAddr := udpConn.LocalAddr().String()
|
|
|
|
// Create TLS config
|
|
tlsConfig := &tls.Config{
|
|
Certificates: []tls.Certificate{testCert.tlsCert},
|
|
NextProtos: []string{"h3"}, // HTTP/3 protocol identifier
|
|
MinVersion: tls.VersionTLS12,
|
|
}
|
|
|
|
// Create HTTP/3 server
|
|
server := &http3.Server{
|
|
Handler: handler,
|
|
TLSConfig: tlsConfig,
|
|
}
|
|
|
|
// Start the server with the existing UDP connection
|
|
go func() {
|
|
if err := server.Serve(udpConn); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
|
t.Logf("HTTP/3 server error: %v", err)
|
|
}
|
|
}()
|
|
|
|
h3Server := &testHTTP3Server{
|
|
server: server,
|
|
cert: testCert.cert,
|
|
addr: actualAddr,
|
|
}
|
|
|
|
// Add cleanup
|
|
t.Cleanup(func() {
|
|
server.Close()
|
|
udpConn.Close()
|
|
})
|
|
|
|
// Wait a bit for the server to be ready
|
|
time.Sleep(100 * time.Millisecond)
|
|
|
|
return h3Server
|
|
}
|
|
|
|
// blockingBodyHandler writes exactly nbytes of body with the given status,
|
|
// flushes them, then blocks until release is closed WITHOUT ever returning.
|
|
// Because the handler does not return, the response stream is never terminated
|
|
// (no EOF/FIN). A client that stops after a bounded prefix therefore completes,
|
|
// while a client that reads to EOF blocks. Tests set nbytes to the exact read
|
|
// cap so the client consumes the whole written body (no half-written frame is
|
|
// left blocking on flow control) yet still never sees EOF.
|
|
func blockingBodyHandler(status, nbytes int, release <-chan struct{}) http.HandlerFunc {
|
|
return func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", headerApplicationDNS)
|
|
w.WriteHeader(status)
|
|
if _, err := w.Write(make([]byte, nbytes)); err != nil {
|
|
return
|
|
}
|
|
if f, ok := w.(http.Flusher); ok {
|
|
f.Flush()
|
|
}
|
|
<-release
|
|
}
|
|
}
|
|
|
|
// requireBoundedResolve asserts that r.Resolve returns the expected size/status
|
|
// error while the server is still withholding EOF (the handler is blocked in
|
|
// blockingBodyHandler). Returning under those conditions proves ctrld read only
|
|
// a bounded prefix of the body: a resolver that instead read to EOF would block
|
|
// on the withheld stream and trip the deadline. This is the deterministic
|
|
// regression guard for the issue-312 OOM protections, replacing the earlier
|
|
// flaky server-side byte counter (issue-561).
|
|
func requireBoundedResolve(t *testing.T, r Resolver, msg *dns.Msg, wantErrSubstr string) {
|
|
t.Helper()
|
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
defer cancel()
|
|
|
|
type result struct {
|
|
answer *dns.Msg
|
|
err error
|
|
}
|
|
done := make(chan result, 1)
|
|
go func() {
|
|
answer, err := r.Resolve(ctx, msg)
|
|
done <- result{answer, err}
|
|
}()
|
|
|
|
select {
|
|
case res := <-done:
|
|
if res.err == nil {
|
|
t.Fatalf("Resolve unexpectedly succeeded; answer=%v", res.answer)
|
|
}
|
|
if !strings.Contains(res.err.Error(), wantErrSubstr) {
|
|
t.Fatalf("error %q does not contain %q", res.err, wantErrSubstr)
|
|
}
|
|
if res.answer != nil {
|
|
t.Fatalf("Resolve returned non-nil answer alongside error: %v", res.answer)
|
|
}
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("Resolve did not return while the server withheld EOF: the body is being read to EOF instead of a bounded prefix (issue-312 OOM protection missing)")
|
|
}
|
|
}
|
|
|
|
// dohUpstreamForTLSServer wires an UpstreamConfig at a local httptest TLS
|
|
// server, trusting its self-signed certificate. BootstrapIP is set so no
|
|
// real DNS lookup runs.
|
|
func dohUpstreamForTLSServer(t *testing.T, srv *httptest.Server) *UpstreamConfig {
|
|
t.Helper()
|
|
pool := x509.NewCertPool()
|
|
pool.AddCert(srv.Certificate())
|
|
u, err := url.Parse(srv.URL)
|
|
if err != nil {
|
|
t.Fatalf("parse server URL: %v", err)
|
|
}
|
|
uc := &UpstreamConfig{
|
|
Name: "doh-oversize",
|
|
Type: ResolverTypeDOH,
|
|
Endpoint: srv.URL + "/dns-query",
|
|
BootstrapIP: u.Hostname(),
|
|
Timeout: 2000,
|
|
}
|
|
uc.SetCertPool(pool)
|
|
uc.Init(context.Background())
|
|
return uc
|
|
}
|
|
|
|
// doh3UpstreamForAddr wires an UpstreamConfig at a local HTTP/3 server,
|
|
// trusting its self-signed certificate.
|
|
func doh3UpstreamForAddr(t *testing.T, addr string, cert *x509.Certificate) *UpstreamConfig {
|
|
t.Helper()
|
|
pool := x509.NewCertPool()
|
|
pool.AddCert(cert)
|
|
host, _, err := net.SplitHostPort(addr)
|
|
if err != nil {
|
|
t.Fatalf("split host/port %q: %v", addr, err)
|
|
}
|
|
uc := &UpstreamConfig{
|
|
Name: "doh3-oversize",
|
|
Type: ResolverTypeDOH3,
|
|
Endpoint: "h3://" + addr + "/dns-query",
|
|
BootstrapIP: host,
|
|
Timeout: 5000,
|
|
}
|
|
uc.SetCertPool(pool)
|
|
uc.Init(context.Background())
|
|
return uc
|
|
}
|
|
|
|
// TestDoHResolve_OversizedBody_Rejected locks in the fix for
|
|
// github.com/Control-D-Inc/ctrld/issues/312: a malicious DoH upstream
|
|
// returning a body larger than the DNS protocol allows must be rejected
|
|
// with an explicit size error rather than buffered into ctrld memory.
|
|
func TestDoHResolve_OversizedBody_Rejected(t *testing.T) {
|
|
// Write exactly the LimitReader cap, then withhold EOF. ctrld's bounded
|
|
// read (io.LimitReader of dohMaxResponseSize+1) returns after this prefix;
|
|
// an unbounded read would block on the missing EOF and trip the deadline.
|
|
release := make(chan struct{})
|
|
defer close(release)
|
|
srv := httptest.NewUnstartedServer(blockingBodyHandler(http.StatusOK, dohMaxResponseSize+1, release))
|
|
testCert := generateTestCertificate(t)
|
|
srv.TLS = &tls.Config{
|
|
Certificates: []tls.Certificate{testCert.tlsCert},
|
|
NextProtos: []string{"h2", "http/1.1"},
|
|
MinVersion: tls.VersionTLS12,
|
|
}
|
|
srv.StartTLS()
|
|
t.Cleanup(srv.Close)
|
|
|
|
uc := dohUpstreamForTLSServer(t, srv)
|
|
r, err := NewResolver(context.Background(), uc)
|
|
if err != nil {
|
|
t.Fatalf("NewResolver: %v", err)
|
|
}
|
|
|
|
msg := new(dns.Msg)
|
|
msg.SetQuestion("example.com.", dns.TypeA)
|
|
msg.RecursionDesired = true
|
|
|
|
requireBoundedResolve(t, r, msg, "maximum DNS message size")
|
|
}
|
|
|
|
// TestDoHResolve_NonOKStatus_BoundedErrorBody locks in that a non-200
|
|
// response with a huge body does not pull the body fully into ctrld
|
|
// memory just to format an error string.
|
|
func TestDoHResolve_NonOKStatus_BoundedErrorBody(t *testing.T) {
|
|
// Same synchronization as the oversized-body test, but at the error-body
|
|
// cap: the non-200 path reads through an io.LimitReader of
|
|
// dohMaxErrorBodySize, so it must return after this prefix without EOF.
|
|
release := make(chan struct{})
|
|
defer close(release)
|
|
srv := httptest.NewUnstartedServer(blockingBodyHandler(http.StatusBadGateway, dohMaxErrorBodySize, release))
|
|
testCert := generateTestCertificate(t)
|
|
srv.TLS = &tls.Config{
|
|
Certificates: []tls.Certificate{testCert.tlsCert},
|
|
NextProtos: []string{"h2", "http/1.1"},
|
|
MinVersion: tls.VersionTLS12,
|
|
}
|
|
srv.StartTLS()
|
|
t.Cleanup(srv.Close)
|
|
|
|
uc := dohUpstreamForTLSServer(t, srv)
|
|
r, err := NewResolver(context.Background(), uc)
|
|
if err != nil {
|
|
t.Fatalf("NewResolver: %v", err)
|
|
}
|
|
|
|
msg := new(dns.Msg)
|
|
msg.SetQuestion("example.com.", dns.TypeA)
|
|
msg.RecursionDesired = true
|
|
|
|
requireBoundedResolve(t, r, msg, "status: 502")
|
|
}
|
|
|
|
// TestDoHResolve_OversizedBody_DoH3 mirrors the DoH oversized-body check
|
|
// on the HTTP/3 transport, since github-312 specifically reproduced the
|
|
// OOM via DoH3.
|
|
func TestDoHResolve_OversizedBody_DoH3(t *testing.T) {
|
|
release := make(chan struct{})
|
|
defer close(release)
|
|
testCert := generateTestCertificate(t)
|
|
udpConn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 0})
|
|
if err != nil {
|
|
t.Fatalf("udp listen: %v", err)
|
|
}
|
|
h3 := &http3.Server{
|
|
Handler: blockingBodyHandler(http.StatusOK, dohMaxResponseSize+1, release),
|
|
TLSConfig: &tls.Config{
|
|
Certificates: []tls.Certificate{testCert.tlsCert},
|
|
NextProtos: []string{"h3"},
|
|
MinVersion: tls.VersionTLS12,
|
|
},
|
|
}
|
|
go func() {
|
|
if err := h3.Serve(udpConn); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
|
t.Logf("h3 server: %v", err)
|
|
}
|
|
}()
|
|
t.Cleanup(func() {
|
|
_ = h3.Close()
|
|
_ = udpConn.Close()
|
|
})
|
|
time.Sleep(100 * time.Millisecond)
|
|
|
|
uc := doh3UpstreamForAddr(t, udpConn.LocalAddr().String(), testCert.cert)
|
|
r, err := NewResolver(context.Background(), uc)
|
|
if err != nil {
|
|
t.Fatalf("NewResolver: %v", err)
|
|
}
|
|
|
|
msg := new(dns.Msg)
|
|
msg.SetQuestion("example.com.", dns.TypeA)
|
|
msg.RecursionDesired = true
|
|
|
|
requireBoundedResolve(t, r, msg, "maximum DNS message size")
|
|
}
|