diff --git a/README.md b/README.md index ac4a8d4..635b357 100644 --- a/README.md +++ b/README.md @@ -19,12 +19,14 @@ make test # start the test postgres (localhost:5443) and run the suite ``` To try DNS locally, pass the DNS variables to `make run` (the dev Compose file -has no Tidepool service). Use a free port; 5353 is mDNS and usually taken. The -dev zone root is `localhost`: +has no Tidepool service). DNS listens on UDP and TCP at the same address and +answers SOA and NS at each label apex. Use a free port; 5353 is mDNS and +usually taken. The dev zone root is `localhost`: ```sh DNS_LISTEN=127.0.0.1:5300 DNS_PUBLIC_IPV4=127.0.0.1 make run dig @127.0.0.1 -p 5300 TXT _atproto...localhost +dig +tcp @127.0.0.1 -p 5300 NS .localhost ``` Production does not migrate on server startup. If operating without a registry, @@ -314,10 +316,10 @@ Two classes, and the difference matters at boot: | `DATABASE_URL` | local dev postgres | bridge state | | `LISTEN_ADDR` | `:8091` | HTTP bind address | | `BRIDGE_HOSTNAME` | `localhost` | public domain of the bridge; anchors handles and the PDS endpoint in minted DID docs | -| `DNS_LISTEN` | *(empty; disabled)* | UDP listen address (e.g. `:53`); empty disables the DNS server | +| `DNS_LISTEN` | *(empty; disabled)* | UDP and TCP listen address (e.g. `:53`); empty disables the DNS server; answers SOA and NS at each label apex | | `DNS_PUBLIC_IPV4` | *(unset)* | public IPv4 address; required when `DNS_LISTEN` is set | | `DNS_PUBLIC_IPV6` | *(unset)* | optional public IPv6 address when DNS is enabled | -| `DNS_NAMESERVERS` | `ns1`/`ns2` under `BRIDGE_HOSTNAME` when DNS is enabled | comma-separated nameserver hostnames (first is the SOA primary) | +| `DNS_NAMESERVERS` | `ns1`/`ns2` under `BRIDGE_HOSTNAME` when DNS is enabled | comma-separated nameserver hostnames returned in apex NS answers (first is the SOA primary) | | `BRIDGE_SCHEME` | `https` | scheme of the bridge's own AP URLs (actor id, inbox, activity ids). `http` is dev-only — the e2e harness federates with a debug-mode Lemmy over plain HTTP | | `PLC_DIRECTORY_URL` | `http://localhost:3002` (local, `make plc-up`) | did:plc directory; production uses `https://plc.directory` | | `BRIDGE_KEK` | fixed public dev key | 32-byte key-encryption key (64 hex chars or base64) sealing per-actor signing keys and the escrow rotation key at rest (AES-256-GCM) | diff --git a/cmd/tidepool/dns_test.go b/cmd/tidepool/dns_test.go index 753da9f..2536a3c 100644 --- a/cmd/tidepool/dns_test.go +++ b/cmd/tidepool/dns_test.go @@ -2,8 +2,10 @@ package main import ( "context" + "errors" "net" "net/netip" + "syscall" "testing" "time" @@ -21,11 +23,29 @@ func (r *dnsStartupResolver) ResolveHandle(_ context.Context, handle string) (st return r.handles[handle], nil } +func freeStartupDNSAddress(t *testing.T) string { + t.Helper() + for attempt := 0; attempt < 10; attempt++ { + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + address := listener.Addr().String() + packet, err := net.ListenPacket("udp", address) + if err == nil { + require.NoError(t, packet.Close()) + require.NoError(t, listener.Close()) + return address + } + require.NoError(t, listener.Close()) + if !errors.Is(err, syscall.EADDRINUSE) { + require.NoError(t, err) + } + } + t.Fatal("could not find a loopback port free for both UDP and TCP") + return "" +} + func TestStartDNSServerUsesConfiguredZoneAndResolver(t *testing.T) { - probe, err := net.ListenPacket("udp", "127.0.0.1:0") - require.NoError(t, err) - address := probe.LocalAddr().String() - require.NoError(t, probe.Close()) + address := freeStartupDNSAddress(t) ctx, cancel := context.WithCancel(t.Context()) defer cancel() @@ -85,6 +105,51 @@ func TestStartDNSServerUsesConfiguredZoneAndResolver(t *testing.T) { } } +func TestStartDNSServerReleasesUDPWhenTCPBindFails(t *testing.T) { + var listener net.Listener + for attempt := 0; attempt < 10; attempt++ { + candidate, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + packet, err := net.ListenPacket("udp", candidate.Addr().String()) + if err == nil { + require.NoError(t, packet.Close()) + listener = candidate + break + } + require.NoError(t, candidate.Close()) + if !errors.Is(err, syscall.EADDRINUSE) { + require.NoError(t, err) + } + } + require.NotNil(t, listener) + defer listener.Close() + address := listener.Addr().String() + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + cfg := &config.Config{ + DNSListen: address, + BridgeHostname: "tdpl.example", + DNSPublicIPv4: netip.MustParseAddr("127.0.0.1"), + DNSNameservers: []string{"ns1.tdpl.example", "ns2.tdpl.example"}, + } + done, err := startDNSServer(ctx, cfg, &dnsStartupResolver{}, testLogger()) + defer func() { + cancel() + if done != nil { + select { + case <-done: + case <-time.After(2 * time.Second): + t.Error("DNS server did not stop after cancellation") + } + } + }() + require.Error(t, err, "occupied TCP port must fail startup synchronously") + require.Nil(t, done) + packet, err := net.ListenPacket("udp", address) + require.NoError(t, err, "TCP bind failure must release UDP before context cancellation") + require.NoError(t, packet.Close()) +} + func TestStartDNSServerDisabledReturnsNilChannel(t *testing.T) { dnsErrors, err := startDNSServer(t.Context(), &config.Config{}, &dnsStartupResolver{}, testLogger()) require.NoError(t, err) diff --git a/cmd/tidepool/main.go b/cmd/tidepool/main.go index 21ab73d..c01b219 100644 --- a/cmd/tidepool/main.go +++ b/cmd/tidepool/main.go @@ -1017,8 +1017,8 @@ func startConsumer( return done, engine, nil } -// startDNSServer serves handle TXT lookups when DNS_LISTEN is set. The -// returned channel carries ServeUDP's unexpected-stop error; it is nil when +// startDNSServer serves handle DNS over UDP and TCP when DNS_LISTEN is set. The +// returned channel carries Serve's unexpected-stop error; it is nil when // DNS is disabled, so receiving from it in run's select blocks forever. func startDNSServer(ctx context.Context, cfg *config.Config, resolver identity.Resolver, logger *slog.Logger) (<-chan error, error) { if cfg.DNSListen == "" { @@ -1036,7 +1036,7 @@ func startDNSServer(ctx context.Context, cfg *config.Config, resolver identity.R if err != nil { return nil, fmt.Errorf("dns server %s: %w", cfg.DNSListen, err) } - serveErrors, err := dns.ServeUDP(ctx, cfg.DNSListen, handler) + serveErrors, err := dns.Serve(ctx, cfg.DNSListen, handler, logger) if err != nil { return nil, fmt.Errorf("dns server %s: %w", cfg.DNSListen, err) } diff --git a/internal/dns/handler.go b/internal/dns/handler.go index c1f44c5..17ff232 100644 --- a/internal/dns/handler.go +++ b/internal/dns/handler.go @@ -8,9 +8,11 @@ import ( "net" "net/netip" "strings" + "sync" "time" miekgdns "github.com/miekg/dns" + "golang.org/x/net/netutil" internalerrors "tidepool/internal/errors" "tidepool/internal/identity" @@ -18,11 +20,35 @@ import ( const handleLookupTimeout = 2 * time.Second +const ( + zoneRecordTTL = 3600 + soaMinimumTTL = 300 + soaRefresh = 3600 + soaRetry = 600 + soaExpire = 1209600 + handleTXTTTL = 300 +) + // maxConcurrentHandleLookups bounds how many DNS queries may run a handle // lookup at once, so spoofed UDP floods cannot exhaust the shared database // connection pool. Queries over the limit are answered SERVFAIL. const maxConcurrentHandleLookups = 8 +// advertisedUDPPayloadSize is the UDP payload size the server reads and +// advertises in its OPT record (the DNS Flag Day 2020 recommendation). +const advertisedUDPPayloadSize = 1232 + +// TCP limits keep a connection flood from exhausting the file descriptors, +// memory and CPU that the DNS server shares with the rest of the process. +const ( + maxTCPConnections = 256 + tcpReadTimeout = 2 * time.Second + tcpIdleTimeout = 3 * time.Second + maxTCPQueriesPerConnection = 16 + acceptRetryInitialDelay = 5 * time.Millisecond + acceptRetryMaximumDelay = time.Second +) + // Options configures an authoritative DNS handler for bridged handles. type Options struct { ZoneRoot string @@ -36,11 +62,11 @@ type Options struct { // Handler answers DNS queries in the configured handle subzones. type Handler struct { - zoneRoot string - nameserver string - serial uint32 - resolver identity.Resolver - logger *slog.Logger + zoneRoot string + nameservers []string + serial uint32 + resolver identity.Resolver + logger *slog.Logger // lookupSlots holds one token per in-flight handle lookup. lookupSlots chan struct{} } @@ -80,7 +106,7 @@ func NewHandler(options Options) (*Handler, error) { } return &Handler{ zoneRoot: zoneRoot, - nameserver: nameservers[0], + nameservers: nameservers, serial: options.Serial, resolver: options.Resolver, logger: logger, @@ -105,11 +131,16 @@ func normalizeDomainName(name string) (string, error) { return normalized, nil } -// ServeDNS answers handle TXT queries and returns NODATA for other in-zone queries. +// ServeDNS answers handle TXT and label-apex SOA/NS queries, returning NODATA +// for other in-zone questions. func (h *Handler) ServeDNS(writer miekgdns.ResponseWriter, request *miekgdns.Msg) { response := new(miekgdns.Msg) response.SetReply(request) response.RecursionAvailable = false + // SetReply does not copy the request's OPT record (RFC 6891 section 7). + if request.IsEdns0() != nil { + response.SetEdns0(advertisedUDPPayloadSize, false) + } if len(request.Question) != 1 { response.Rcode = miekgdns.RcodeFormatError h.writeResponse(writer, response) @@ -148,7 +179,7 @@ func (h *Handler) ServeDNS(writer miekgdns.ResponseWriter, request *miekgdns.Msg switch { case err == nil: response.Answer = []miekgdns.RR{&miekgdns.TXT{ - Hdr: miekgdns.RR_Header{Name: question.Name, Rrtype: miekgdns.TypeTXT, Class: miekgdns.ClassINET, Ttl: 300}, + Hdr: miekgdns.RR_Header{Name: question.Name, Rrtype: miekgdns.TypeTXT, Class: miekgdns.ClassINET, Ttl: handleTXTTTL}, Txt: []string{"did=" + did}, }} case internalerrors.IsNotFound(err) || internalerrors.IsValidation(err): @@ -161,7 +192,23 @@ func (h *Handler) ServeDNS(writer miekgdns.ResponseWriter, request *miekgdns.Msg return } - // Task 02: answer SOA and NS queries at the label apex here. + if len(labels) == 1 { + switch question.Qtype { + case miekgdns.TypeSOA: + response.Answer = []miekgdns.RR{h.soa(apex)} + h.writeResponse(writer, response) + return + case miekgdns.TypeNS: + for _, nameserver := range h.nameservers { + response.Answer = append(response.Answer, &miekgdns.NS{ + Hdr: miekgdns.RR_Header{Name: apex, Rrtype: miekgdns.TypeNS, Class: miekgdns.ClassINET, Ttl: zoneRecordTTL}, + Ns: nameserver, + }) + } + h.writeResponse(writer, response) + return + } + } // Task 03: answer A and AAAA queries here. h.addSOA(response, apex) h.writeResponse(writer, response) @@ -184,55 +231,102 @@ func (h *Handler) writeResponse(writer miekgdns.ResponseWriter, response *miekgd failure.MsgHdr = response.MsgHdr failure.Question = response.Question failure.Rcode = miekgdns.RcodeServerFailure + if opt := response.IsEdns0(); opt != nil { + failure.Extra = []miekgdns.RR{opt} + } if err := writer.WriteMsg(failure); err != nil { h.logger.Warn("DNS SERVFAIL write failed", "error", err) } } func (h *Handler) addSOA(response *miekgdns.Msg, apex string) { - response.Ns = []miekgdns.RR{&miekgdns.SOA{ - Hdr: miekgdns.RR_Header{Name: apex, Rrtype: miekgdns.TypeSOA, Class: miekgdns.ClassINET, Ttl: 3600}, - Ns: h.nameserver, + response.Ns = []miekgdns.RR{h.soa(apex)} +} + +func (h *Handler) soa(apex string) *miekgdns.SOA { + return &miekgdns.SOA{ + Hdr: miekgdns.RR_Header{Name: apex, Rrtype: miekgdns.TypeSOA, Class: miekgdns.ClassINET, Ttl: zoneRecordTTL}, + Ns: h.nameservers[0], Mbox: "hostmaster." + h.zoneRoot, Serial: h.serial, - Refresh: 3600, - Retry: 600, - Expire: 1209600, - Minttl: 300, - }} + Refresh: soaRefresh, + Retry: soaRetry, + Expire: soaExpire, + Minttl: soaMinimumTTL, + } } -// ServeUDP binds address synchronously and returns the bind or startup error. -// Once started it serves in the background until ctx is cancelled. The returned -// channel receives one non-nil error if serving stops for any reason other than -// ctx cancellation, and is closed once the server has stopped. -func ServeUDP(ctx context.Context, address string, handler miekgdns.Handler) (<-chan error, error) { - packet, err := net.ListenPacket("udp", address) - if err != nil { - return nil, fmt.Errorf("bind DNS UDP %s: %w", address, err) +// servePacketConn serves DNS on packet and takes ownership of it. +func servePacketConn(ctx context.Context, packet net.PacketConn, handler miekgdns.Handler) (<-chan error, error) { + server := &miekgdns.Server{PacketConn: packet, Handler: handler, UDPSize: advertisedUDPPayloadSize} + return serveServer(ctx, server, packet.Close, packet.LocalAddr()) +} + +// serveListener serves DNS on listener with bounded connections, timeouts and +// queries per connection, and takes ownership of it. +func serveListener(ctx context.Context, listener net.Listener, handler miekgdns.Handler, logger *slog.Logger) (<-chan error, error) { + // The backoff wraps the limit so a retrying Accept does not hold a slot. + limited := newAcceptBackoffListener(netutil.LimitListener(listener, maxTCPConnections), logger) + server := &miekgdns.Server{ + Listener: limited, + Handler: handler, + ReadTimeout: tcpReadTimeout, + IdleTimeout: func() time.Duration { return tcpIdleTimeout }, + MaxTCPQueries: maxTCPQueriesPerConnection, } - done, err := servePacketConn(ctx, packet, handler) - if err != nil { - return nil, fmt.Errorf("start DNS UDP %s: %w", address, err) + return serveServer(ctx, server, limited.Close, listener.Addr()) +} + +// acceptBackoffListener retries temporary Accept errors such as EMFILE with +// exponential backoff and a warning, where miekg would retry at once and +// silently. Other errors, including the one Accept returns after Close, pass +// through. +type acceptBackoffListener struct { + net.Listener + logger *slog.Logger + closeOnce sync.Once + closed chan struct{} +} + +func newAcceptBackoffListener(listener net.Listener, logger *slog.Logger) *acceptBackoffListener { + return &acceptBackoffListener{Listener: listener, logger: logger, closed: make(chan struct{})} +} + +func (l *acceptBackoffListener) Accept() (net.Conn, error) { + var delay time.Duration + for { + conn, err := l.Listener.Accept() + var netError net.Error + if err == nil || !errors.As(err, &netError) || !netError.Temporary() { //nolint:staticcheck // Temporary is the test miekg applies to Accept errors. + return conn, err + } + delay = min(max(delay*2, acceptRetryInitialDelay), acceptRetryMaximumDelay) + l.logger.Warn("DNS TCP accept failed; retrying", "error", err, "retry_in", delay) + timer := time.NewTimer(delay) + select { + case <-timer.C: + case <-l.closed: + timer.Stop() + return nil, net.ErrClosed + } } - return done, nil } -// servePacketConn serves DNS on packet with the contract of ServeUDP. It takes -// ownership of packet and closes it when serving stops. When a serve failure -// races with ctx cancellation, the failure is reported only if ctx.Err() is -// still nil at that point: a stop during shutdown is treated as the shutdown. -func servePacketConn(ctx context.Context, packet net.PacketConn, handler miekgdns.Handler) (<-chan error, error) { +// Close interrupts a backoff in progress and closes the wrapped listener. +func (l *acceptBackoffListener) Close() error { + l.closeOnce.Do(func() { close(l.closed) }) + return l.Listener.Close() +} + +// serveServer waits for startup and reports an unexpected stop unless the +// context is cancelled. It closes the socket once serving stops. +func serveServer(ctx context.Context, server *miekgdns.Server, closeSocket func() error, address net.Addr) (<-chan error, error) { started := make(chan struct{}) stopped := make(chan error, 1) - server := &miekgdns.Server{ - PacketConn: packet, - Handler: handler, - NotifyStartedFunc: func() { close(started) }, - } + server.NotifyStartedFunc = func() { close(started) } go func() { err := server.ActivateAndServe() - _ = packet.Close() + _ = closeSocket() stopped <- err }() select { @@ -259,7 +353,104 @@ func servePacketConn(ctx context.Context, packet net.PacketConn, handler miekgdn if err == nil { err = errors.New("DNS server stopped unexpectedly") } - done <- fmt.Errorf("serve DNS on %s: %w", packet.LocalAddr(), err) + done <- fmt.Errorf("serve DNS on %s: %w", address, err) + } + }() + return done, nil +} + +// Serve binds UDP and TCP on the same address before starting either server. +// Its channel reports one unexpected transport failure and closes after both +// transports stop; cancellation shuts both down without reporting an error. +// TCP accept failures that are retried are logged to logger. +func Serve(ctx context.Context, address string, handler miekgdns.Handler, logger *slog.Logger) (<-chan error, error) { + if logger == nil { + logger = slog.Default() + } + packet, err := net.ListenPacket("udp", address) + if err != nil { + return nil, fmt.Errorf("bind DNS UDP %s: %w", address, err) + } + tcpAddress := address + if _, port, splitErr := net.SplitHostPort(address); splitErr == nil && port == "0" { + tcpAddress = packet.LocalAddr().String() + } + listener, err := net.Listen("tcp", tcpAddress) + if err != nil { + _ = packet.Close() + return nil, fmt.Errorf("bind DNS TCP %s: %w", tcpAddress, err) + } + done, err := serveConns(ctx, packet, listener, handler, logger) + if err != nil { + return nil, fmt.Errorf("start DNS on %s: %w", address, err) + } + return done, nil +} + +// udpResponseWriter applies the request's UDP payload limit before sending. +type udpResponseWriter struct { + miekgdns.ResponseWriter + size int +} + +func (writer udpResponseWriter) WriteMsg(response *miekgdns.Msg) error { + response.Truncate(writer.size) + return writer.ResponseWriter.WriteMsg(response) +} + +type udpHandler struct{ miekgdns.Handler } + +func (handler udpHandler) ServeDNS(writer miekgdns.ResponseWriter, request *miekgdns.Msg) { + size := miekgdns.MinMsgSize + requestOPT := request.IsEdns0() + if requestOPT != nil { + size = int(requestOPT.UDPSize()) + if size < miekgdns.MinMsgSize { + size = miekgdns.MinMsgSize + } + } + handler.Handler.ServeDNS(udpResponseWriter{ResponseWriter: writer, size: size}, request) +} + +// serveConns takes ownership of both sockets and reports the first unexpected +// stop, closing its channel after both transports have stopped. +func serveConns(ctx context.Context, packet net.PacketConn, listener net.Listener, handler miekgdns.Handler, logger *slog.Logger) (<-chan error, error) { + serveContext, cancel := context.WithCancel(ctx) + udpDone, err := servePacketConn(serveContext, packet, udpHandler{handler}) + if err != nil { + cancel() + _ = listener.Close() + return nil, fmt.Errorf("start DNS UDP: %w", err) + } + tcpDone, err := serveListener(serveContext, listener, handler, logger) + if err != nil { + cancel() + for range udpDone { + } + return nil, fmt.Errorf("start DNS TCP: %w", err) + } + + done := make(chan error, 1) + go func() { + defer close(done) + defer cancel() + for udpDone != nil || tcpDone != nil { + select { + case err, open := <-udpDone: + if !open { + udpDone = nil + } else if serveContext.Err() == nil { + done <- err + cancel() + } + case err, open := <-tcpDone: + if !open { + tcpDone = nil + } else if serveContext.Err() == nil { + done <- err + cancel() + } + } } }() return done, nil diff --git a/internal/dns/handler_test.go b/internal/dns/handler_test.go index bd3d241..45c321b 100644 --- a/internal/dns/handler_test.go +++ b/internal/dns/handler_test.go @@ -101,7 +101,7 @@ func TestNewHandlerOptions(t *testing.T) { handler, err = NewHandler(options) require.NoError(t, err) require.Equal(t, "tdpl.example.", handler.zoneRoot, "zone root must be stored trimmed, lowercased and fully qualified") - require.Equal(t, "ns1.tdpl.example.", handler.nameserver, "nameserver must be stored trimmed, lowercased and fully qualified") + require.Equal(t, []string{"ns1.tdpl.example.", "ns2.tdpl.example."}, handler.nameservers, "nameservers must be stored trimmed, lowercased and fully qualified") } func TestHandlerNameClassification(t *testing.T) { @@ -299,6 +299,67 @@ func TestHandlerNameClassification(t *testing.T) { } } +func TestHandlerApexSOAAndNS(t *testing.T) { + for _, tc := range []struct { + name string + question string + questionType uint16 + wantAnswer int + wantAuthority int + }{ + {"B1 apex SOA", "lemmy-world.tdpl.example.", miekgdns.TypeSOA, 1, 0}, + {"B1 nested SOA NODATA", "alice.lemmy-world.tdpl.example.", miekgdns.TypeSOA, 0, 1}, + {"B2 apex NS", "lemmy-world.tdpl.example.", miekgdns.TypeNS, 2, 0}, + {"B2 nested NS NODATA", "alice.lemmy-world.tdpl.example.", miekgdns.TypeNS, 0, 1}, + } { + t.Run(tc.name, func(t *testing.T) { + resolver := &recordingResolver{} + handler, err := NewHandler(handlerOptions(resolver)) + require.NoError(t, err) + request := new(miekgdns.Msg) + request.SetQuestion(tc.question, tc.questionType) + writer := &recordingDNSWriter{} + handler.ServeDNS(writer, request) + require.NotNil(t, writer.response) + response := writer.response + require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) + require.True(t, response.Authoritative) + require.Len(t, response.Answer, tc.wantAnswer) + require.Len(t, response.Ns, tc.wantAuthority) + require.Empty(t, response.Extra) + require.Empty(t, resolver.asked) + + if tc.questionType == miekgdns.TypeNS && tc.wantAnswer != 0 { + for index, wantTarget := range []string{"ns1.tdpl.example.", "ns2.tdpl.example."} { + ns, ok := response.Answer[index].(*miekgdns.NS) + require.True(t, ok, "answer %d must be NS, got %T", index, response.Answer[index]) + require.Equal(t, "lemmy-world.tdpl.example.", ns.Hdr.Name) + require.Equal(t, uint16(miekgdns.ClassINET), ns.Hdr.Class) + require.Equal(t, uint32(3600), ns.Hdr.Ttl) + require.Equal(t, wantTarget, ns.Ns) + } + return + } + + var record miekgdns.RR + if tc.wantAnswer != 0 { + record = response.Answer[0] + } else { + record = response.Ns[0] + } + soa, ok := record.(*miekgdns.SOA) + require.True(t, ok, "expected SOA, got %T", record) + require.Equal(t, "lemmy-world.tdpl.example.", soa.Hdr.Name) + require.Equal(t, uint16(miekgdns.ClassINET), soa.Hdr.Class) + require.Equal(t, uint32(3600), soa.Hdr.Ttl) + require.Equal(t, "ns1.tdpl.example.", soa.Ns) + require.Equal(t, "hostmaster.tdpl.example.", soa.Mbox) + require.Equal(t, uint32(2024100101), soa.Serial) + require.Equal(t, uint32(300), soa.Minttl) + }) + } +} + type failingDNSWriter struct { recordingDNSWriter writeErrors []error @@ -322,6 +383,7 @@ func TestHandlerLogsWriteFailures(t *testing.T) { wantLevel string wantMessages int wantRetry bool + requestEDNS bool }{ { name: "pack failure logs error and answers SERVFAIL", @@ -329,6 +391,13 @@ func TestHandlerLogsWriteFailures(t *testing.T) { wantLevel: "level=ERROR", wantRetry: true, }, + { + name: "pack failure of an EDNS query answers SERVFAIL with an OPT record", + writeErrors: []error{miekgdns.ErrRdata}, + wantLevel: "level=ERROR", + wantRetry: true, + requestEDNS: true, + }, { name: "socket failure logs warning", writeErrors: []error{&net.OpError{Op: "write", Net: "udp", Err: errors.New("network unreachable")}}, @@ -344,6 +413,9 @@ func TestHandlerLogsWriteFailures(t *testing.T) { request := new(miekgdns.Msg) request.SetQuestion("alice.lemmy-world.tdpl.example.", miekgdns.TypeTXT) + if tc.requestEDNS { + request.SetEdns0(4096, false) + } writer := &failingDNSWriter{writeErrors: tc.writeErrors} handler.ServeDNS(writer, request) @@ -358,7 +430,14 @@ func TestHandlerLogsWriteFailures(t *testing.T) { require.Equal(t, request.Question, failure.Question) require.Empty(t, failure.Answer) require.Empty(t, failure.Ns) - require.Empty(t, failure.Extra) + if tc.requestEDNS { + require.Len(t, failure.Extra, 1, "the SERVFAIL must carry only the OPT record") + opt := failure.IsEdns0() + require.NotNil(t, opt, "an EDNS query's SERVFAIL must carry an OPT record") + require.Equal(t, uint16(advertisedUDPPayloadSize), opt.UDPSize()) + } else { + require.Empty(t, failure.Extra) + } } else { require.Len(t, writer.written, 1, "a socket failure must not be retried") } diff --git a/internal/dns/server_test.go b/internal/dns/server_test.go index 4d28580..2da6eff 100644 --- a/internal/dns/server_test.go +++ b/internal/dns/server_test.go @@ -1,9 +1,14 @@ package dns import ( + "bytes" "context" "errors" + "fmt" + "log/slog" "net" + "strings" + "sync" "syscall" "testing" "time" @@ -12,45 +17,55 @@ import ( "github.com/stretchr/testify/require" ) -func freeUDPAddress(t *testing.T) string { +func freeDNSAddress(t *testing.T) string { t.Helper() - packet, err := net.ListenPacket("udp", "127.0.0.1:0") - require.NoError(t, err) - address := packet.LocalAddr().String() - require.NoError(t, packet.Close()) - return address + for attempt := 0; attempt < 10; attempt++ { + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + address := listener.Addr().String() + packet, err := net.ListenPacket("udp", address) + if err == nil { + require.NoError(t, packet.Close()) + require.NoError(t, listener.Close()) + return address + } + require.NoError(t, listener.Close()) + if !errors.Is(err, syscall.EADDRINUSE) { + require.NoError(t, err) + } + } + t.Fatal("could not find a loopback port free for both UDP and TCP") + return "" } -func TestServeUDP(t *testing.T) { - t.Run("answers handle TXT and releases socket on cancellation", func(t *testing.T) { +func TestServe(t *testing.T) { + t.Run("answers handle TXT over UDP and TCP and releases both sockets on cancellation", func(t *testing.T) { resolver := &recordingResolver{handles: map[string]string{"alice.lemmy-world.tdpl.example": "did:plc:abc"}} handler, err := NewHandler(handlerOptions(resolver)) require.NoError(t, err) ctx, cancel := context.WithCancel(t.Context()) defer cancel() - address, done := serveOnFreeUDPAddress(t, ctx, handler) + address, done := serveOnFreeDNSAddress(t, ctx, handler) request := new(miekgdns.Msg) request.SetQuestion("_atproto.alice.lemmy-world.tdpl.example.", miekgdns.TypeTXT) - response, _, err := (&miekgdns.Client{Net: "udp", Timeout: time.Second}).Exchange(request, address) - require.NoError(t, err, "ServeUDP must receive the query on its bound loopback socket") - require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) - require.True(t, response.Authoritative) - require.Len(t, response.Answer, 1) - txt, ok := response.Answer[0].(*miekgdns.TXT) - require.True(t, ok) - require.Equal(t, []string{"did=did:plc:abc"}, txt.Txt) + for _, network := range []string{"udp", "tcp"} { + t.Run(network, func(t *testing.T) { + response, _, err := (&miekgdns.Client{Net: network, Timeout: time.Second}).Exchange(request, address) + require.NoError(t, err, "Serve must receive the query over %s", network) + require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) + require.True(t, response.Authoritative) + require.Len(t, response.Answer, 1) + txt, ok := response.Answer[0].(*miekgdns.TXT) + require.True(t, ok) + require.Equal(t, []string{"did=did:plc:abc"}, txt.Txt) + }) + } cancel() requireClosedWithoutError(t, done) - require.Eventually(t, func() bool { - probe, err := net.ListenPacket("udp", address) - if err != nil { - return false - } - return probe.Close() == nil - }, 2*time.Second, 10*time.Millisecond, "cancellation must release the UDP socket") + requireBothSocketsAvailable(t, address) }) t.Run("occupied UDP address fails synchronously", func(t *testing.T) { @@ -61,7 +76,7 @@ func TestServeUDP(t *testing.T) { require.NoError(t, err) ctx, cancel := context.WithCancel(t.Context()) defer cancel() - done, err := ServeUDP(ctx, packet.LocalAddr().String(), handler) + done, err := Serve(ctx, packet.LocalAddr().String(), handler, slog.New(slog.DiscardHandler)) require.Error(t, err) require.Nil(t, done) }) @@ -71,12 +86,12 @@ func TestServeUDP(t *testing.T) { require.NoError(t, err) ctx, cancel := context.WithCancel(t.Context()) defer cancel() - address, done := serveOnFreeUDPAddress(t, ctx, handler) + address, done := serveOnFreeDNSAddress(t, ctx, handler) probe, err := net.ListenPacket("udp", address) if err == nil { _ = probe.Close() } - require.Error(t, err, "ServeUDP must bind before returning") + require.Error(t, err, "Serve must bind before returning") cancel() requireClosedWithoutError(t, done) require.Eventually(t, func() bool { @@ -89,7 +104,23 @@ func TestServeUDP(t *testing.T) { }) } -func serveOnFreeUDPAddress(t *testing.T, ctx context.Context, handler miekgdns.Handler) (string, <-chan error) { +func requireBothSocketsAvailable(t *testing.T, address string) { + t.Helper() + require.Eventually(t, func() bool { + packet, err := net.ListenPacket("udp", address) + if err != nil { + return false + } + defer packet.Close() + listener, err := net.Listen("tcp", address) + if err != nil { + return false + } + return listener.Close() == nil + }, 2*time.Second, 10*time.Millisecond, "cancellation must release both UDP and TCP sockets") +} + +func serveOnFreeDNSAddress(t *testing.T, ctx context.Context, handler miekgdns.Handler) (string, <-chan error) { t.Helper() var ( address string @@ -97,8 +128,8 @@ func serveOnFreeUDPAddress(t *testing.T, ctx context.Context, handler miekgdns.H err error ) for attempt := 0; attempt < 5; attempt++ { - address = freeUDPAddress(t) - done, err = ServeUDP(ctx, address, handler) + address = freeDNSAddress(t) + done, err = Serve(ctx, address, handler, slog.New(slog.DiscardHandler)) if !errors.Is(err, syscall.EADDRINUSE) { break } @@ -118,6 +149,167 @@ func requireClosedWithoutError(t *testing.T, done <-chan error) { } } +func largeNSFixture(t *testing.T) string { + t.Helper() + resolver := &recordingResolver{handles: map[string]string{"alice.lemmy-world.tdpl.example": "did:plc:abc"}} + options := handlerOptions(resolver) + options.Nameservers = make([]string, 40) + for index := range options.Nameservers { + options.Nameservers[index] = fmt.Sprintf("ns.distinct-nameserver-domain-label-%02d.example", index) + } + handler, err := NewHandler(options) + require.NoError(t, err) + ctx, cancel := context.WithCancel(t.Context()) + address, done := serveOnFreeDNSAddress(t, ctx, handler) + t.Cleanup(func() { + cancel() + requireClosedWithoutError(t, done) + }) + + request := new(miekgdns.Msg) + request.SetQuestion("lemmy-world.tdpl.example.", miekgdns.TypeNS) + response, _, err := (&miekgdns.Client{Net: "tcp", Timeout: time.Second}).Exchange(request, address) + require.NoError(t, err, "fixture precondition: full NS reply must be available over TCP") + require.Len(t, response.Answer, 40, "fixture precondition: full NS reply") + response.Compress = true + packed, err := response.Pack() + require.NoError(t, err) + require.Greater(t, len(packed), 1300, "fixture precondition: packed full reply must exceed 1300 bytes") + require.Less(t, len(packed), 4000, "fixture precondition: packed full reply must be under 4000 bytes") + return address +} + +func rawUDPQuery(t *testing.T, address string, request *miekgdns.Msg) (*miekgdns.Msg, int) { + t.Helper() + packed, err := request.Pack() + require.NoError(t, err) + conn, err := net.Dial("udp", address) + require.NoError(t, err) + defer conn.Close() + require.NoError(t, conn.SetDeadline(time.Now().Add(time.Second))) + _, err = conn.Write(packed) + require.NoError(t, err) + buffer := make([]byte, 65535) + count, err := conn.Read(buffer) + require.NoError(t, err) + response := new(miekgdns.Msg) + require.NoError(t, response.Unpack(buffer[:count])) + return response, count +} + +func TestServeTCPFullNSAnswers(t *testing.T) { + address := largeNSFixture(t) + for _, tc := range []struct { + name string + ednsSize uint16 + }{ + {"without OPT", 0}, + {"with OPT advertising 1232", 1232}, + } { + t.Run(tc.name, func(t *testing.T) { + request := new(miekgdns.Msg) + request.SetQuestion("lemmy-world.tdpl.example.", miekgdns.TypeNS) + if tc.ednsSize != 0 { + request.SetEdns0(tc.ednsSize, false) + } + response, _, err := (&miekgdns.Client{Net: "tcp", Timeout: time.Second}).Exchange(request, address) + require.NoError(t, err) + require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) + require.False(t, response.Truncated, "TCP must return the full reply") + require.Len(t, response.Answer, 40) + if tc.ednsSize != 0 { + opt := response.IsEdns0() + require.NotNil(t, opt, "TCP replies to EDNS queries must carry an OPT record") + require.Equal(t, uint16(advertisedUDPPayloadSize), opt.UDPSize()) + } else { + require.Nil(t, response.IsEdns0(), "replies to non-EDNS queries must not carry an OPT record") + } + for _, record := range response.Answer { + _, ok := record.(*miekgdns.NS) + require.True(t, ok, "expected NS, got %T", record) + } + }) + } +} + +func TestServeUDPWithoutEDNSTruncatesNS(t *testing.T) { + address := largeNSFixture(t) + request := new(miekgdns.Msg) + request.SetQuestion("lemmy-world.tdpl.example.", miekgdns.TypeNS) + response, count := rawUDPQuery(t, address, request) + require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) + require.True(t, response.Truncated) + require.LessOrEqual(t, count, 512, "received datagram must fit legacy UDP") + require.Nil(t, response.IsEdns0(), "replies to non-EDNS queries must not carry an OPT record") + + request.SetQuestion("_atproto.alice.lemmy-world.tdpl.example.", miekgdns.TypeTXT) + response, count = rawUDPQuery(t, address, request) + require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) + require.False(t, response.Truncated, "small UDP replies must not be truncated") + require.LessOrEqual(t, count, 512) + require.Nil(t, response.IsEdns0(), "replies to non-EDNS queries must not carry an OPT record") + require.Len(t, response.Answer, 1) + txt, ok := response.Answer[0].(*miekgdns.TXT) + require.True(t, ok) + require.Equal(t, []string{"did=did:plc:abc"}, txt.Txt) +} + +func TestServeUDPWithLargeEDNSReturnsFullNS(t *testing.T) { + address := largeNSFixture(t) + request := new(miekgdns.Msg) + request.SetQuestion("lemmy-world.tdpl.example.", miekgdns.TypeNS) + request.SetEdns0(4096, false) + response, count := rawUDPQuery(t, address, request) + require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) + require.False(t, response.Truncated) + require.Len(t, response.Answer, 40) + require.LessOrEqual(t, count, 4096) + opt := response.IsEdns0() + require.NotNil(t, opt, "EDNS0 replies must carry an OPT record") + require.Equal(t, uint16(advertisedUDPPayloadSize), opt.UDPSize(), "the OPT must advertise the server's own receive size") +} + +func TestServeUDPWithSmallEDNSTruncatesNS(t *testing.T) { + address := largeNSFixture(t) + request := new(miekgdns.Msg) + request.SetQuestion("lemmy-world.tdpl.example.", miekgdns.TypeNS) + request.SetEdns0(1232, false) + response, count := rawUDPQuery(t, address, request) + require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) + require.True(t, response.Truncated) + require.Greater(t, count, 512, "EDNS0 must allow more than legacy UDP") + require.LessOrEqual(t, count, 1232) + opt := response.IsEdns0() + require.NotNil(t, opt, "EDNS0 replies must carry an OPT record") + require.Equal(t, uint16(advertisedUDPPayloadSize), opt.UDPSize(), "the OPT must advertise the server's own receive size") +} + +func TestServeUDPAnswersQueriesLargerThanLegacySize(t *testing.T) { + resolver := &recordingResolver{handles: map[string]string{"alice.lemmy-world.tdpl.example": "did:plc:abc"}} + handler, err := NewHandler(handlerOptions(resolver)) + require.NoError(t, err) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + address, _ := serveOnFreeDNSAddress(t, ctx, handler) + + request := new(miekgdns.Msg) + request.SetQuestion("_atproto.alice.lemmy-world.tdpl.example.", miekgdns.TypeTXT) + request.SetEdns0(4096, false) + opt := request.IsEdns0() + opt.Option = append(opt.Option, &miekgdns.EDNS0_PADDING{Padding: make([]byte, 800)}) + packed, err := request.Pack() + require.NoError(t, err) + require.Greater(t, len(packed), 512, "fixture precondition: query must exceed legacy UDP size") + require.LessOrEqual(t, len(packed), advertisedUDPPayloadSize, "fixture precondition: query must fit the advertised size") + + response, _ := rawUDPQuery(t, address, request) + require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) + require.Len(t, response.Answer, 1) + txt, ok := response.Answer[0].(*miekgdns.TXT) + require.True(t, ok) + require.Equal(t, []string{"did=did:plc:abc"}, txt.Txt) +} + var errInjectedRead = errors.New("injected read failure") // failingPacketConn reads nothing until fail is closed, then returns a @@ -163,3 +355,223 @@ func TestServePacketConnReportsServeFailure(t *testing.T) { t.Fatal("server channel was not closed after the serve failure") } } + +var errInjectedAccept = errors.New("injected accept failure") + +type failingListener struct { + net.Listener + fail chan struct{} +} + +func (l *failingListener) Accept() (net.Conn, error) { + <-l.fail + return nil, errInjectedAccept +} + +func TestServeConnsReportsEitherTransportFailure(t *testing.T) { + for _, network := range []string{"udp", "tcp"} { + t.Run(network, func(t *testing.T) { + packet, err := net.ListenPacket("udp", "127.0.0.1:0") + require.NoError(t, err) + defer packet.Close() + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer listener.Close() + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + var fail chan struct{} + if network == "udp" { + failing := &failingPacketConn{PacketConn: packet, fail: make(chan struct{})} + fail = failing.fail + packet = failing + } else { + failing := &failingListener{Listener: listener, fail: make(chan struct{})} + fail = failing.fail + listener = failing + } + handler, err := NewHandler(handlerOptions(&recordingResolver{})) + require.NoError(t, err) + done, err := serveConns(ctx, packet, listener, handler, slog.New(slog.DiscardHandler)) + require.NoError(t, err) + select { + case reported, open := <-done: + t.Fatalf("server stopped before %s failed: error=%v open=%v", network, reported, open) + default: + } + close(fail) + select { + case reported, open := <-done: + require.True(t, open, "failure must be sent before the channel closes") + require.Error(t, reported) + case <-time.After(2 * time.Second): + t.Fatal("transport failure was not reported") + } + select { + case reported, open := <-done: + require.False(t, open, "the channel must close after one error, got %v", reported) + case <-time.After(2 * time.Second): + t.Fatal("server channel was not closed after the failure") + } + }) + } +} + +// dialTCPAndQuery opens a DNS TCP connection and checks it is served. +func dialTCPAndQuery(t *testing.T, address string, request *miekgdns.Msg) *miekgdns.Conn { + t.Helper() + conn, err := miekgdns.Dial("tcp", address) + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + require.NoError(t, conn.SetDeadline(time.Now().Add(time.Second))) + require.NoError(t, conn.WriteMsg(request)) + response, err := conn.ReadMsg() + require.NoError(t, err) + require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) + return conn +} + +func TestServeTCPLimitsConcurrentConnections(t *testing.T) { + handler, err := NewHandler(handlerOptions(&recordingResolver{})) + require.NoError(t, err) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + address, _ := serveOnFreeDNSAddress(t, ctx, handler) + request := new(miekgdns.Msg) + request.SetQuestion("lemmy-world.tdpl.example.", miekgdns.TypeSOA) + + held := make([]*miekgdns.Conn, 0, maxTCPConnections) + for range maxTCPConnections { + held = append(held, dialTCPAndQuery(t, address, request)) + } + + extra, err := miekgdns.Dial("tcp", address) + require.NoError(t, err) + defer func() { _ = extra.Close() }() + require.NoError(t, extra.SetDeadline(time.Now().Add(300*time.Millisecond))) + require.NoError(t, extra.WriteMsg(request)) + _, err = extra.ReadMsg() + var netError net.Error + require.ErrorAs(t, err, &netError, "a connection over the limit must not be served") + require.True(t, netError.Timeout(), "a connection over the limit must wait, got %v", err) + + require.NoError(t, held[0].Close()) + require.NoError(t, extra.SetDeadline(time.Now().Add(time.Second))) + response, err := extra.ReadMsg() + require.NoError(t, err, "a waiting connection must be served once a slot frees") + require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) +} + +func TestServeTCPClosesConnectionAfterQueryLimit(t *testing.T) { + handler, err := NewHandler(handlerOptions(&recordingResolver{})) + require.NoError(t, err) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + address, _ := serveOnFreeDNSAddress(t, ctx, handler) + request := new(miekgdns.Msg) + request.SetQuestion("lemmy-world.tdpl.example.", miekgdns.TypeSOA) + + conn, err := miekgdns.Dial("tcp", address) + require.NoError(t, err) + defer func() { _ = conn.Close() }() + require.NoError(t, conn.SetDeadline(time.Now().Add(2*time.Second))) + for query := range maxTCPQueriesPerConnection { + require.NoError(t, conn.WriteMsg(request), "query %d", query) + response, err := conn.ReadMsg() + require.NoError(t, err, "query %d must be answered", query) + require.Equal(t, miekgdns.RcodeSuccess, response.Rcode) + } + _ = conn.WriteMsg(request) + _, err = conn.ReadMsg() + require.Error(t, err, "the connection must be closed after the query limit") + var netError net.Error + if errors.As(err, &netError) { + require.False(t, netError.Timeout(), "the server must close the connection, not leave it idle") + } +} + +type temporaryAcceptError struct{} + +func (temporaryAcceptError) Error() string { return "too many open files" } +func (temporaryAcceptError) Timeout() bool { return false } +func (temporaryAcceptError) Temporary() bool { return true } + +// scriptedListener returns the scripted Accept errors, then a connection or, +// when the script is exhausted and repeat is set, the last error forever. +type scriptedListener struct { + mutex sync.Mutex + errors []error + repeat bool + accepted int + conn net.Conn +} + +func (l *scriptedListener) Accept() (net.Conn, error) { + l.mutex.Lock() + defer l.mutex.Unlock() + l.accepted++ + if len(l.errors) == 0 { + return l.conn, nil + } + err := l.errors[0] + if !l.repeat || len(l.errors) > 1 { + l.errors = l.errors[1:] + } + return nil, err +} + +func (l *scriptedListener) acceptCount() int { + l.mutex.Lock() + defer l.mutex.Unlock() + return l.accepted +} + +func (l *scriptedListener) Close() error { return nil } +func (l *scriptedListener) Addr() net.Addr { return &net.TCPAddr{} } + +func TestAcceptBackoffListener(t *testing.T) { + t.Run("retries temporary errors with a warning", func(t *testing.T) { + var logs bytes.Buffer + client, server := net.Pipe() + defer func() { _ = client.Close() }() + defer func() { _ = server.Close() }() + inner := &scriptedListener{errors: []error{temporaryAcceptError{}, temporaryAcceptError{}}, conn: server} + listener := newAcceptBackoffListener(inner, slog.New(slog.NewTextHandler(&logs, nil))) + + conn, err := listener.Accept() + require.NoError(t, err) + require.Same(t, server, conn) + require.Equal(t, 3, inner.acceptCount()) + require.Equal(t, 2, strings.Count(logs.String(), "level=WARN"), logs.String()) + require.Contains(t, logs.String(), "too many open files") + }) + + t.Run("passes other errors through without a warning", func(t *testing.T) { + var logs bytes.Buffer + inner := &scriptedListener{errors: []error{errInjectedAccept}} + listener := newAcceptBackoffListener(inner, slog.New(slog.NewTextHandler(&logs, nil))) + + _, err := listener.Accept() + require.ErrorIs(t, err, errInjectedAccept) + require.Equal(t, 1, inner.acceptCount()) + require.Empty(t, logs.String()) + }) + + t.Run("close interrupts the backoff", func(t *testing.T) { + inner := &scriptedListener{errors: []error{temporaryAcceptError{}}, repeat: true} + listener := newAcceptBackoffListener(inner, slog.New(slog.DiscardHandler)) + returned := make(chan error, 1) + go func() { + _, err := listener.Accept() + returned <- err + }() + // Eight failures put the next backoff at 640ms. + require.Eventually(t, func() bool { return inner.acceptCount() >= 8 }, 2*time.Second, time.Millisecond) + require.NoError(t, listener.Close()) + select { + case err := <-returned: + require.ErrorIs(t, err, net.ErrClosed) + case <-time.After(200 * time.Millisecond): + t.Fatal("Close must interrupt the accept backoff") + } + }) +}