plugin/forward: support DNS-over-QUIC upstreams (#8474)

Signed-off-by: houyuwushang <liuluoqianqiu@outlook.com>
This commit is contained in:
houyuwushang
2026-09-08 12:51:18 +08:00
committed by GitHub
parent 71a60e140b
commit e1d3fe6bc6
13 changed files with 1679 additions and 40 deletions

View File

@@ -1,6 +1,6 @@
// Package proxy implements a forwarding proxy with connection caching.
// It manages a pool of upstream connections (UDP and TCP) to reuse them for subsequent requests,
// reducing latency and handshake overhead. It supports in-band health checking.
// It reuses upstream DNS, DoT, DoH, and DoQ connections to reduce latency and
// handshake overhead. It supports in-band health checking.
package proxy
import (
@@ -267,6 +267,17 @@ func (p *Proxy) lookupDoH(ctx context.Context, state request.Request, _ Options)
return ret, localAddr, proto, nil
}
func (p *Proxy) lookupDoQ(ctx context.Context, state request.Request, _ Options) (*dns.Msg, net.Addr, string, error) {
// QUIC runs over UDP. Reporting udp keeps dnstap query_address and
// response_address consistent with the actual upstream socket.
const proto = "udp"
if p.doq == nil {
return nil, nil, proto, errors.New("proxy: DoQ transport is not initialized")
}
ret, localAddr, err := p.doq.exchange(ctx, state.Req, p.readTimeout)
return ret, localAddr, proto, err
}
// Connect selects an upstream, sends the request and waits for a response. It
// also returns CoreDNS's own outbound address on the upstream socket
// (localAddr) and the transport proto ("udp" or "tcp") actually used to reach
@@ -284,6 +295,8 @@ func (p *Proxy) Connect(ctx context.Context, state request.Request, opts Options
switch p.protocol {
case transport.HTTPS:
ret, localAddr, proto, err = p.lookupDoH(ctx, state, opts)
case transport.QUIC:
ret, localAddr, proto, err = p.lookupDoQ(ctx, state, opts)
case transport.DNS, transport.TLS:
ret, localAddr, proto, err = p.lookupDNS(ctx, state, opts)
default:

548
plugin/pkg/proxy/doq.go Normal file
View File

@@ -0,0 +1,548 @@
package proxy
import (
"context"
"crypto/tls"
"encoding/binary"
"errors"
"fmt"
"io"
"net"
"sync"
"time"
"github.com/coredns/coredns/plugin/pkg/transport"
"github.com/miekg/dns"
"github.com/quic-go/quic-go"
)
const (
doqALPN = "doq"
doqDialTimeout = 5 * time.Second
doqDefaultIdleTimeout = 30 * time.Second
doqProtocolError = quic.ApplicationErrorCode(0x2)
doqRequestCancelled = quic.StreamErrorCode(0x3)
)
var errDoQProtocol = errors.New("DNS-over-QUIC protocol error")
type doqConn struct {
conn *quic.Conn
transport *quic.Transport
created time.Time
lastUsed time.Time
active int
draining bool
closed bool
}
// doqTransport owns one reusable QUIC connection. Queries share the
// connection, but each query uses its own bidirectional stream as required by
// RFC 9250.
type doqTransport struct {
proxyName string
addr string
mu sync.Mutex
tlsConfig *tls.Config
localAddress net.IP
expire time.Duration
maxAge time.Duration
readTimeout time.Duration
current *doqConn
connections map[*doqConn]struct{}
dialDone chan struct{}
started bool
stopped bool
stop chan struct{}
stopOnce sync.Once
lifecycleCtx context.Context
cancelLifecycle context.CancelFunc
}
func newDoQTransport(proxyName, addr string) *doqTransport {
lifecycleCtx, cancel := context.WithCancel(context.Background()) // #nosec G118 -- stopTransport calls the stored cancel function
return &doqTransport{
proxyName: proxyName,
addr: addr,
expire: defaultExpire,
readTimeout: maxTimeout,
connections: make(map[*doqConn]struct{}),
stop: make(chan struct{}),
lifecycleCtx: lifecycleCtx,
cancelLifecycle: cancel,
}
}
func (t *doqTransport) setTLSConfig(cfg *tls.Config) {
t.mu.Lock()
defer t.mu.Unlock()
if cfg == nil {
t.tlsConfig = nil
return
}
t.tlsConfig = cfg.Clone()
}
func (t *doqTransport) setLocalAddress(addr net.IP) {
t.mu.Lock()
defer t.mu.Unlock()
t.localAddress = append(net.IP(nil), addr...)
}
func (t *doqTransport) setExpire(expire time.Duration) {
t.mu.Lock()
t.expire = expire
t.mu.Unlock()
}
func (t *doqTransport) setMaxAge(maxAge time.Duration) {
t.mu.Lock()
t.maxAge = maxAge
t.mu.Unlock()
}
func (t *doqTransport) setReadTimeout(timeout time.Duration) {
t.mu.Lock()
t.readTimeout = timeout
t.mu.Unlock()
}
func (t *doqTransport) start() {
t.mu.Lock()
if t.started || t.stopped {
t.mu.Unlock()
return
}
t.started = true
t.mu.Unlock()
go t.connManager()
}
func (t *doqTransport) connManager() {
ticker := time.NewTicker(defaultExpire)
defer ticker.Stop()
for {
select {
case now := <-ticker.C:
t.cleanup(now)
case <-t.stop:
return
}
}
}
func (t *doqTransport) stopTransport() {
t.stopOnce.Do(func() {
t.cancelLifecycle()
t.mu.Lock()
t.stopped = true
close(t.stop)
connections := make([]*doqConn, 0, len(t.connections))
for c := range t.connections {
if c.closed {
continue
}
c.closed = true
connections = append(connections, c)
}
t.current = nil
clear(t.connections)
t.mu.Unlock()
for _, c := range connections {
closeDoQConn(c, 0, "")
}
})
}
func (t *doqTransport) cleanup(now time.Time) {
var toClose []*doqConn
t.mu.Lock()
if c := t.current; c != nil {
dead := c.conn.Context().Err() != nil
expired := c.active == 0 && (t.expire == 0 || now.Sub(c.lastUsed) >= t.expire)
tooOld := t.maxAge > 0 && now.Sub(c.created) >= t.maxAge
if dead || expired || tooOld {
t.current = nil
c.draining = true
}
}
for c := range t.connections {
if c.draining && c.active == 0 && !c.closed {
c.closed = true
delete(t.connections, c)
toClose = append(toClose, c)
}
}
t.mu.Unlock()
for _, c := range toClose {
closeDoQConn(c, 0, "")
}
}
func (t *doqTransport) acquire(ctx context.Context) (*doqConn, bool, error) {
for {
t.cleanup(time.Now())
t.mu.Lock()
if t.stopped {
t.mu.Unlock()
return nil, false, errors.New(ErrTransportStopped)
}
if c := t.current; c != nil {
c.active++
t.mu.Unlock()
connCacheHitsCount.WithLabelValues(t.proxyName, t.addr, transport.QUIC).Inc()
return c, true, nil
}
if done := t.dialDone; done != nil {
t.mu.Unlock()
select {
case <-done:
continue
case <-t.stop:
return nil, false, errors.New(ErrTransportStopped)
case <-ctx.Done():
return nil, false, ctx.Err()
}
}
done := make(chan struct{})
t.dialDone = done
tlsConfig := cloneDoQTLSConfig(t.tlsConfig)
localAddress := append(net.IP(nil), t.localAddress...)
idleTimeout := max(doqDefaultIdleTimeout, t.expire, t.readTimeout)
t.mu.Unlock()
connCacheMissesCount.WithLabelValues(t.proxyName, t.addr, transport.QUIC).Inc()
dialCtx, cancelDial := context.WithCancel(ctx)
stopDial := context.AfterFunc(t.lifecycleCtx, cancelDial)
c, err := dialDoQ(dialCtx, t.addr, localAddress, tlsConfig, idleTimeout)
stopDial()
cancelDial()
t.mu.Lock()
t.dialDone = nil
close(done)
if err == nil && !t.stopped {
c.active = 1
t.current = c
t.connections[c] = struct{}{}
t.mu.Unlock()
return c, false, nil
}
stopped := t.stopped
t.mu.Unlock()
if c != nil {
closeDoQConn(c, 0, "")
}
if stopped {
return nil, false, errors.New(ErrTransportStopped)
}
return nil, false, err
}
}
func cloneDoQTLSConfig(cfg *tls.Config) *tls.Config {
if cfg == nil {
cfg = new(tls.Config)
} else {
cfg = cfg.Clone()
}
// DoQ uses a dedicated ALPN. Do not offer another application protocol on
// this connection.
cfg.NextProtos = []string{doqALPN}
return cfg
}
func dialDoQ(ctx context.Context, addr string, localAddress net.IP, tlsConfig *tls.Config, idleTimeout time.Duration) (*doqConn, error) {
remote, err := net.ResolveUDPAddr("udp", addr)
if err != nil {
return nil, err
}
network := "udp6"
if remote.IP.To4() != nil {
network = "udp4"
}
local := &net.UDPAddr{IP: localAddress}
packetConn, err := net.ListenUDP(network, local)
if err != nil {
return nil, err
}
quicTransport := &quic.Transport{Conn: packetConn}
quicConfig := &quic.Config{
HandshakeIdleTimeout: doqDialTimeout,
MaxIncomingStreams: -1,
MaxIncomingUniStreams: -1,
}
quicConfig.MaxIdleTimeout = idleTimeout
dialCtx, cancel := context.WithTimeout(ctx, doqDialTimeout)
defer cancel()
conn, err := quicTransport.Dial(dialCtx, remote, tlsConfig, quicConfig)
if err != nil {
_ = quicTransport.Close()
return nil, err
}
now := time.Now()
return &doqConn{
conn: conn,
transport: quicTransport,
created: now,
lastUsed: now,
}, nil
}
func (t *doqTransport) release(c *doqConn) {
var closeConn bool
t.mu.Lock()
if c.active > 0 {
c.active--
}
c.lastUsed = time.Now()
if c.draining && c.active == 0 && !c.closed {
c.closed = true
delete(t.connections, c)
closeConn = true
}
t.mu.Unlock()
if closeConn {
closeDoQConn(c, 0, "")
}
}
func (t *doqTransport) retire(c *doqConn, code quic.ApplicationErrorCode, reason string, abort bool) {
var closeConn bool
t.mu.Lock()
if t.current == c {
t.current = nil
}
c.draining = true
if c.active == 0 && !c.closed {
c.closed = true
delete(t.connections, c)
closeConn = true
}
t.mu.Unlock()
if abort {
_ = c.conn.CloseWithError(code, reason)
}
if closeConn {
closeDoQConn(c, code, reason)
}
}
func closeDoQConn(c *doqConn, code quic.ApplicationErrorCode, reason string) {
if c == nil {
return
}
if c.conn != nil {
_ = c.conn.CloseWithError(code, reason)
}
if c.transport != nil {
_ = c.transport.Close()
}
}
func (t *doqTransport) exchange(ctx context.Context, msg *dns.Msg, timeout time.Duration) (*dns.Msg, net.Addr, error) {
if isDNSZoneTransfer(msg) {
return nil, nil, fmt.Errorf("%w: zone transfers over DoQ require multi-message response support", ErrUnsupportedRequest)
}
query := msg.Copy()
query.Id = 0
removeEDNSTCPKeepalive(query)
wire, err := query.Pack()
if err != nil {
return nil, nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
}
if len(wire) > int(^uint16(0)) {
return nil, nil, fmt.Errorf("%w: DNS message is too large for DoQ", ErrInvalidRequest)
}
c, cached, err := t.acquire(ctx)
if err != nil {
return nil, nil, err
}
localAddr := c.conn.LocalAddr()
if timeout <= 0 {
t.mu.Lock()
timeout = t.readTimeout
t.mu.Unlock()
}
queryCtx := ctx
cancel := func() {}
if timeout > 0 {
queryCtx, cancel = context.WithTimeout(ctx, timeout)
}
defer cancel()
stream, err := c.conn.OpenStreamSync(queryCtx)
if err != nil {
if queryCtx.Err() != nil {
t.release(c)
return nil, localAddr, queryCtx.Err()
}
t.retire(c, 0, "", true)
t.release(c)
if cached {
return nil, localAddr, ErrCachedClosed
}
return nil, localAddr, err
}
defer t.release(c)
stopCancellation := context.AfterFunc(queryCtx, func() {
stream.CancelRead(doqRequestCancelled)
stream.CancelWrite(doqRequestCancelled)
})
defer stopCancellation()
if err = writeDOQMessage(stream, wire); err != nil {
stream.CancelRead(doqRequestCancelled)
stream.CancelWrite(doqRequestCancelled)
if queryCtx.Err() != nil {
return nil, localAddr, queryCtx.Err()
}
if c.conn.Context().Err() != nil {
t.retire(c, 0, "", true)
}
return nil, localAddr, err
}
if err = stream.Close(); err != nil {
stream.CancelRead(doqRequestCancelled)
if queryCtx.Err() != nil {
return nil, localAddr, queryCtx.Err()
}
if c.conn.Context().Err() != nil {
t.retire(c, 0, "", true)
}
return nil, localAddr, err
}
responseWire, err := readDOQMessage(stream)
if err != nil {
if errors.Is(err, errDoQProtocol) {
t.retire(c, doqProtocolError, err.Error(), true)
} else {
stream.CancelRead(doqRequestCancelled)
if c.conn.Context().Err() != nil {
t.retire(c, 0, "", true)
}
}
if queryCtx.Err() != nil {
return nil, localAddr, queryCtx.Err()
}
return nil, localAddr, err
}
if err = expectDOQFIN(stream); err != nil {
if errors.Is(err, errDoQProtocol) {
t.retire(c, doqProtocolError, err.Error(), true)
} else {
stream.CancelRead(doqRequestCancelled)
}
if queryCtx.Err() != nil {
return nil, localAddr, queryCtx.Err()
}
return nil, localAddr, err
}
response := new(dns.Msg)
if err = response.Unpack(responseWire); err != nil {
err = fmt.Errorf("%w: invalid DNS response: %v", errDoQProtocol, err)
t.retire(c, doqProtocolError, err.Error(), true)
return nil, localAddr, err
}
if response.Id != 0 {
err = fmt.Errorf("%w: response message ID is %d, want 0", errDoQProtocol, response.Id)
t.retire(c, doqProtocolError, err.Error(), true)
return nil, localAddr, err
}
response.Id = msg.Id
return response, localAddr, nil
}
func isDNSZoneTransfer(msg *dns.Msg) bool {
return len(msg.Question) == 1 && (msg.Question[0].Qtype == dns.TypeAXFR || msg.Question[0].Qtype == dns.TypeIXFR)
}
func removeEDNSTCPKeepalive(msg *dns.Msg) {
opt := msg.IsEdns0()
if opt == nil {
return
}
options := opt.Option[:0]
for _, option := range opt.Option {
if option.Option() != dns.EDNS0TCPKEEPALIVE {
options = append(options, option)
}
}
opt.Option = options
}
func writeDOQMessage(w io.Writer, msg []byte) error {
frame := make([]byte, 2+len(msg))
binary.BigEndian.PutUint16(frame, uint16(len(msg))) // #nosec G115 -- checked by caller
copy(frame[2:], msg)
for len(frame) > 0 {
n, err := w.Write(frame)
if err != nil {
return err
}
if n == 0 {
return io.ErrShortWrite
}
frame = frame[n:]
}
return nil
}
func readDOQMessage(r io.Reader) ([]byte, error) {
var sizeBytes [2]byte
if _, err := io.ReadFull(r, sizeBytes[:]); err != nil {
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
return nil, fmt.Errorf("%w: incomplete message length", errDoQProtocol)
}
return nil, err
}
size := binary.BigEndian.Uint16(sizeBytes[:])
if size == 0 {
return nil, fmt.Errorf("%w: zero-length DNS message", errDoQProtocol)
}
msg := make([]byte, int(size))
if _, err := io.ReadFull(r, msg); err != nil {
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
return nil, fmt.Errorf("%w: message ended before %d bytes", errDoQProtocol, size)
}
return nil, err
}
return msg, nil
}
func expectDOQFIN(r io.Reader) error {
var extra [1]byte
for {
n, err := r.Read(extra[:])
if n != 0 {
return fmt.Errorf("%w: multiple responses on one query stream", errDoQProtocol)
}
if errors.Is(err, io.EOF) {
return nil
}
if err != nil {
return err
}
}
}

View File

@@ -0,0 +1,670 @@
package proxy
import (
"bytes"
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"errors"
"fmt"
"math/big"
"net"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/coredns/coredns/plugin/pkg/transport"
"github.com/coredns/coredns/request"
"github.com/miekg/dns"
"github.com/quic-go/quic-go"
)
type doqTestHandler func(int64, *quic.Conn, *quic.Stream, *dns.Msg) error
type doqTestServer struct {
listener *quic.Listener
handler doqTestHandler
ctx context.Context
cancel context.CancelFunc
acceptDone chan struct{}
errors chan error
closeOnce sync.Once
wg sync.WaitGroup
mu sync.Mutex
conns map[*quic.Conn]struct{}
accepted atomic.Int64
streams atomic.Int64
}
func newDoQTestServer(t *testing.T, handler doqTestHandler) (*doqTestServer, *tls.Config) {
t.Helper()
serverTLS, clientTLS := makeDoQTestTLSConfigs(t)
listener, err := quic.ListenAddr("127.0.0.1:0", serverTLS, &quic.Config{MaxIncomingStreams: 256})
if err != nil {
t.Fatalf("quic.ListenAddr() failed: %v", err)
}
ctx, cancel := context.WithCancel(context.Background())
s := &doqTestServer{
listener: listener,
handler: handler,
ctx: ctx,
cancel: cancel,
acceptDone: make(chan struct{}),
errors: make(chan error, 64),
conns: make(map[*quic.Conn]struct{}),
}
go s.serve()
t.Cleanup(s.close)
return s, clientTLS
}
func (s *doqTestServer) addr() string { return s.listener.Addr().String() }
func (s *doqTestServer) serve() {
defer close(s.acceptDone)
for {
conn, err := s.listener.Accept(s.ctx)
if err != nil {
return
}
connNumber := s.accepted.Add(1)
s.mu.Lock()
s.conns[conn] = struct{}{}
s.mu.Unlock()
s.wg.Go(func() { s.serveConn(connNumber, conn) })
}
}
func (s *doqTestServer) serveConn(connNumber int64, conn *quic.Conn) {
defer func() {
s.mu.Lock()
delete(s.conns, conn)
s.mu.Unlock()
}()
for {
stream, err := conn.AcceptStream(s.ctx)
if err != nil {
return
}
s.streams.Add(1)
s.wg.Go(func() {
if err := s.serveStream(connNumber, conn, stream); err != nil {
select {
case s.errors <- err:
default:
}
}
})
}
}
func (s *doqTestServer) serveStream(connNumber int64, conn *quic.Conn, stream *quic.Stream) error {
_ = stream.SetDeadline(time.Now().Add(5 * time.Second))
wire, err := readDOQMessage(stream)
if err != nil {
return fmt.Errorf("read query: %w", err)
}
if err := expectDOQFIN(stream); err != nil {
return fmt.Errorf("read query FIN: %w", err)
}
query := new(dns.Msg)
if err := query.Unpack(wire); err != nil {
return fmt.Errorf("unpack query: %w", err)
}
return s.handler(connNumber, conn, stream, query)
}
func (s *doqTestServer) close() {
s.closeOnce.Do(func() {
s.cancel()
_ = s.listener.Close()
<-s.acceptDone
s.mu.Lock()
connections := make([]*quic.Conn, 0, len(s.conns))
for conn := range s.conns {
connections = append(connections, conn)
}
s.mu.Unlock()
for _, conn := range connections {
_ = conn.CloseWithError(0, "test shutdown")
}
done := make(chan struct{})
go func() {
s.wg.Wait()
close(done)
}()
select {
case <-done:
case <-time.After(2 * time.Second):
}
})
}
func makeDoQTestTLSConfigs(t *testing.T) (*tls.Config, *tls.Config) {
t.Helper()
privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatalf("ecdsa.GenerateKey() failed: %v", err)
}
template := &x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "doq.test"},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
DNSNames: []string{"doq.test"},
}
der, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
if err != nil {
t.Fatalf("x509.CreateCertificate() failed: %v", err)
}
cert := tls.Certificate{Certificate: [][]byte{der}, PrivateKey: privateKey}
roots := x509.NewCertPool()
parsed, err := x509.ParseCertificate(der)
if err != nil {
t.Fatalf("x509.ParseCertificate() failed: %v", err)
}
roots.AddCert(parsed)
return &tls.Config{
Certificates: []tls.Certificate{cert},
NextProtos: []string{doqALPN},
}, &tls.Config{
RootCAs: roots,
ServerName: "doq.test",
}
}
func writeDoQTestResponse(stream *quic.Stream, response *dns.Msg) error {
wire, err := response.Pack()
if err != nil {
return err
}
if err := writeDOQMessage(stream, wire); err != nil {
return err
}
return stream.Close()
}
func replyToDoQTestQuery(stream *quic.Stream, query *dns.Msg) error {
response := new(dns.Msg)
response.SetReply(query)
return writeDoQTestResponse(stream, response)
}
func doqTestRequest(name string, id uint16) request.Request {
query := new(dns.Msg)
query.SetQuestion(name, dns.TypeA)
query.Id = id
return request.Request{Req: query}
}
func TestProxyDoQExchange(t *testing.T) {
type observation struct {
id uint16
alpn string
hasKeepalive bool
}
observed := make(chan observation, 1)
server, clientTLS := newDoQTestServer(t, func(_ int64, conn *quic.Conn, stream *quic.Stream, query *dns.Msg) error {
obs := observation{id: query.Id, alpn: conn.ConnectionState().TLS.NegotiatedProtocol}
if opt := query.IsEdns0(); opt != nil {
for _, option := range opt.Option {
obs.hasKeepalive = obs.hasKeepalive || option.Option() == dns.EDNS0TCPKEEPALIVE
}
}
observed <- obs
response := new(dns.Msg)
response.SetReply(query)
record, err := dns.NewRR("example.org. 60 IN A 192.0.2.1")
if err != nil {
return err
}
response.Answer = []dns.RR{record}
return writeDoQTestResponse(stream, response)
})
p := NewProxy("TestProxyDoQExchange", server.addr(), transport.QUIC)
p.SetTLSConfig(clientTLS)
defer p.Stop()
state := doqTestRequest("example.org.", 0x1234)
state.Req.SetEdns0(1232, false)
state.Req.IsEdns0().Option = append(state.Req.IsEdns0().Option, &dns.EDNS0_TCP_KEEPALIVE{Code: dns.EDNS0TCPKEEPALIVE, Timeout: 10})
response, localAddr, proto, err := p.Connect(context.Background(), state, Options{ForceTCP: true})
if err != nil {
t.Fatalf("Connect() failed: %v", err)
}
if response.Id != 0x1234 {
t.Fatalf("response ID = %d, want %d", response.Id, 0x1234)
}
if state.Req.Id != 0x1234 {
t.Fatalf("request ID was mutated: got %d", state.Req.Id)
}
if len(state.Req.IsEdns0().Option) != 1 {
t.Fatal("the downstream EDNS TCP keepalive option was mutated")
}
if proto != "udp" {
t.Fatalf("reported protocol = %q, want udp", proto)
}
if _, ok := localAddr.(*net.UDPAddr); !ok {
t.Fatalf("local address type = %T, want *net.UDPAddr", localAddr)
}
if len(response.Answer) != 1 || response.Answer[0].String() != "example.org.\t60\tIN\tA\t192.0.2.1" {
t.Fatalf("unexpected answer: %v", response.Answer)
}
obs := <-observed
if obs.id != 0 {
t.Errorf("upstream query ID = %d, want 0", obs.id)
}
if obs.alpn != doqALPN {
t.Errorf("negotiated ALPN = %q, want %q", obs.alpn, doqALPN)
}
if obs.hasKeepalive {
t.Error("upstream query retained the EDNS TCP keepalive option")
}
}
func TestProxyDoQVerifiesServerName(t *testing.T) {
server, clientTLS := newDoQTestServer(t, func(_ int64, _ *quic.Conn, stream *quic.Stream, query *dns.Msg) error {
return replyToDoQTestQuery(stream, query)
})
badTLS := clientTLS.Clone()
badTLS.ServerName = "wrong.test"
p := NewProxy("TestProxyDoQVerifiesServerName", server.addr(), transport.QUIC)
p.SetTLSConfig(badTLS)
p.SetReadTimeout(time.Second)
defer p.Stop()
_, _, _, err := p.Connect(context.Background(), doqTestRequest("example.org.", 1), Options{})
if err == nil {
t.Fatal("Connect() succeeded with the wrong TLS server name")
}
var hostnameError x509.HostnameError
if !errors.As(err, &hostnameError) {
t.Fatalf("Connect() error = %T %v, want x509.HostnameError", err, err)
}
}
func TestProxyDoQSourceAddress(t *testing.T) {
remoteAddress := make(chan net.Addr, 1)
server, clientTLS := newDoQTestServer(t, func(_ int64, conn *quic.Conn, stream *quic.Stream, query *dns.Msg) error {
remoteAddress <- conn.RemoteAddr()
return replyToDoQTestQuery(stream, query)
})
p := NewProxy("TestProxyDoQSourceAddress", server.addr(), transport.QUIC)
p.SetTLSConfig(clientTLS)
p.SetLocalAddress(net.ParseIP("127.0.0.2"))
defer p.Stop()
_, localAddress, _, err := p.Connect(context.Background(), doqTestRequest("example.org.", 1), Options{})
if err != nil {
t.Fatalf("Connect() failed: %v", err)
}
localUDP, ok := localAddress.(*net.UDPAddr)
if !ok {
t.Fatalf("local address type = %T, want *net.UDPAddr", localAddress)
}
if got := localUDP.IP.String(); got != "127.0.0.2" {
t.Fatalf("local source address = %s, want 127.0.0.2", got)
}
remote := <-remoteAddress
remoteUDP, ok := remote.(*net.UDPAddr)
if !ok {
t.Fatalf("remote address type = %T, want *net.UDPAddr", remote)
}
if got := remoteUDP.IP.String(); got != "127.0.0.2" {
t.Fatalf("server observed source address = %s, want 127.0.0.2", got)
}
}
func TestProxyDoQReusesOneConnectionForConcurrentQueries(t *testing.T) {
const queries = 16
var active atomic.Int64
var maxActive atomic.Int64
var arrived atomic.Int64
release := make(chan struct{})
server, clientTLS := newDoQTestServer(t, func(_ int64, _ *quic.Conn, stream *quic.Stream, query *dns.Msg) error {
current := active.Add(1)
defer active.Add(-1)
for {
previous := maxActive.Load()
if current <= previous || maxActive.CompareAndSwap(previous, current) {
break
}
}
if arrived.Add(1) == queries {
close(release)
}
select {
case <-release:
case <-time.After(3 * time.Second):
return errors.New("concurrent queries did not arrive on time")
}
return replyToDoQTestQuery(stream, query)
})
p := NewProxy("TestProxyDoQConcurrent", server.addr(), transport.QUIC)
p.SetTLSConfig(clientTLS)
p.SetReadTimeout(4 * time.Second)
defer p.Stop()
var wg sync.WaitGroup
errs := make(chan error, queries)
for i := range queries {
wg.Go(func() {
state := doqTestRequest(fmt.Sprintf("q%d.example.", i), uint16(i+1))
response, _, _, err := p.Connect(context.Background(), state, Options{})
if err == nil && response.Id != uint16(i+1) {
err = fmt.Errorf("response ID = %d, want %d", response.Id, i+1)
}
errs <- err
})
}
wg.Wait()
close(errs)
for err := range errs {
if err != nil {
t.Fatalf("concurrent Connect() failed: %v", err)
}
}
if got := server.accepted.Load(); got != 1 {
t.Errorf("accepted connections = %d, want 1", got)
}
if got := server.streams.Load(); got != queries {
t.Errorf("accepted streams = %d, want %d", got, queries)
}
if got := maxActive.Load(); got != queries {
t.Errorf("maximum concurrent streams = %d, want %d", got, queries)
}
}
func TestProxyDoQCancellationDoesNotCloseConnection(t *testing.T) {
cancelledWrite := make(chan error, 1)
server, clientTLS := newDoQTestServer(t, func(_ int64, _ *quic.Conn, stream *quic.Stream, query *dns.Msg) error {
if query.Question[0].Name == "slow.example." {
time.Sleep(250 * time.Millisecond)
response := new(dns.Msg)
response.SetReply(query)
err := writeDoQTestResponse(stream, response)
cancelledWrite <- err
return nil
}
return replyToDoQTestQuery(stream, query)
})
p := NewProxy("TestProxyDoQCancellation", server.addr(), transport.QUIC)
p.SetTLSConfig(clientTLS)
p.SetReadTimeout(time.Second)
defer p.Stop()
ctx, cancel := context.WithTimeout(context.Background(), 75*time.Millisecond)
defer cancel()
started := time.Now()
_, _, _, err := p.Connect(ctx, doqTestRequest("slow.example.", 1), Options{})
if !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("slow Connect() error = %v, want context deadline exceeded", err)
}
if elapsed := time.Since(started); elapsed > 500*time.Millisecond {
t.Fatalf("canceled Connect() returned after %s", elapsed)
}
response, _, _, err := p.Connect(context.Background(), doqTestRequest("fast.example.", 2), Options{})
if err != nil {
t.Fatalf("Connect() after cancellation failed: %v", err)
}
if response.Id != 2 {
t.Fatalf("response ID = %d, want 2", response.Id)
}
if got := server.accepted.Load(); got != 1 {
t.Fatalf("connections after stream cancellation = %d, want 1", got)
}
select {
case err := <-cancelledWrite:
if err == nil {
t.Error("server write on the canceled stream unexpectedly succeeded")
}
case <-time.After(time.Second):
t.Fatal("server did not observe the canceled stream")
}
}
func TestProxyDoQReplacesClosedConnection(t *testing.T) {
server, clientTLS := newDoQTestServer(t, func(connNumber int64, conn *quic.Conn, stream *quic.Stream, query *dns.Msg) error {
if err := replyToDoQTestQuery(stream, query); err != nil {
return err
}
if connNumber == 1 {
go func() {
time.Sleep(10 * time.Millisecond)
_ = conn.CloseWithError(0, "rotate test connection")
}()
}
return nil
})
p := NewProxy("TestProxyDoQReplacesClosed", server.addr(), transport.QUIC)
p.SetTLSConfig(clientTLS)
p.SetReadTimeout(time.Second)
defer p.Stop()
if _, _, _, err := p.Connect(context.Background(), doqTestRequest("first.example.", 1), Options{}); err != nil {
t.Fatalf("first Connect() failed: %v", err)
}
p.doq.mu.Lock()
first := p.doq.current
p.doq.mu.Unlock()
if first == nil {
t.Fatal("first QUIC connection was not cached")
}
select {
case <-first.conn.Context().Done():
case <-time.After(time.Second):
t.Fatal("server did not close the first QUIC connection")
}
if _, _, _, err := p.Connect(context.Background(), doqTestRequest("second.example.", 2), Options{}); err != nil {
t.Fatalf("second Connect() failed: %v", err)
}
if got := server.accepted.Load(); got != 2 {
t.Fatalf("accepted connections = %d, want 2", got)
}
}
func TestProxyDoQProtocolErrorRetiresConnection(t *testing.T) {
server, clientTLS := newDoQTestServer(t, func(connNumber int64, _ *quic.Conn, stream *quic.Stream, query *dns.Msg) error {
response := new(dns.Msg)
response.SetReply(query)
if connNumber == 1 {
response.Id = 1
}
return writeDoQTestResponse(stream, response)
})
p := NewProxy("TestProxyDoQProtocolError", server.addr(), transport.QUIC)
p.SetTLSConfig(clientTLS)
p.SetReadTimeout(time.Second)
defer p.Stop()
_, _, _, err := p.Connect(context.Background(), doqTestRequest("bad.example.", 10), Options{})
if !errors.Is(err, errDoQProtocol) {
t.Fatalf("first Connect() error = %v, want DoQ protocol error", err)
}
response, _, _, err := p.Connect(context.Background(), doqTestRequest("good.example.", 11), Options{})
if err != nil {
t.Fatalf("Connect() after protocol error failed: %v", err)
}
if response.Id != 11 {
t.Fatalf("response ID = %d, want 11", response.Id)
}
if got := server.accepted.Load(); got != 2 {
t.Fatalf("accepted connections = %d, want 2", got)
}
}
func TestDoQHealthCheck(t *testing.T) {
query := make(chan *dns.Msg, 1)
server, clientTLS := newDoQTestServer(t, func(_ int64, _ *quic.Conn, stream *quic.Stream, msg *dns.Msg) error {
query <- msg.Copy()
return replyToDoQTestQuery(stream, msg)
})
p := NewProxy("TestDoQHealth", server.addr(), transport.QUIC)
p.SetTLSConfig(clientTLS)
defer p.Stop()
hc := p.GetHealthchecker()
hc.SetDomain("health.example.")
hc.SetRecursionDesired(false)
if err := hc.Check(p); err != nil {
t.Fatalf("health check failed: %v", err)
}
msg := <-query
if len(msg.Question) != 1 || msg.Question[0].Name != "health.example." || msg.Question[0].Qtype != dns.TypeNS {
t.Fatalf("unexpected health query: %v", msg.Question)
}
if msg.RecursionDesired {
t.Error("health query unexpectedly requested recursion")
}
if msg.Id != 0 {
t.Errorf("health query ID = %d, want 0", msg.Id)
}
}
func TestProxyDoQRejectsZoneTransferBeforeDial(t *testing.T) {
p := NewProxy("TestProxyDoQRejectsZoneTransfer", "127.0.0.1:1", transport.QUIC)
defer p.Stop()
for _, qtype := range []uint16{dns.TypeAXFR, dns.TypeIXFR} {
query := new(dns.Msg)
query.SetQuestion("example.org.", qtype)
_, _, _, err := p.Connect(context.Background(), request.Request{Req: query}, Options{})
if !errors.Is(err, ErrUnsupportedRequest) {
t.Errorf("Connect(%s) error = %v, want ErrUnsupportedRequest", dns.TypeToString[qtype], err)
}
}
if got := len(p.doq.connections); got != 0 {
t.Fatalf("zone transfer opened %d DoQ connections, want 0", got)
}
}
func TestDoQConnectionExpiryAndMaxAge(t *testing.T) {
tests := []struct {
name string
configure func(*Proxy)
}{
{
name: "expire",
configure: func(p *Proxy) {
p.SetExpire(20 * time.Millisecond)
},
},
{
name: "max age",
configure: func(p *Proxy) {
p.SetExpire(time.Hour)
p.SetMaxAge(20 * time.Millisecond)
},
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
server, clientTLS := newDoQTestServer(t, func(_ int64, _ *quic.Conn, stream *quic.Stream, query *dns.Msg) error {
return replyToDoQTestQuery(stream, query)
})
p := NewProxy("TestDoQLifetime", server.addr(), transport.QUIC)
p.SetTLSConfig(clientTLS)
p.SetReadTimeout(time.Second)
tc.configure(p)
defer p.Stop()
if _, _, _, err := p.Connect(context.Background(), doqTestRequest("first.example.", 1), Options{}); err != nil {
t.Fatalf("first Connect() failed: %v", err)
}
time.Sleep(30 * time.Millisecond)
if _, _, _, err := p.Connect(context.Background(), doqTestRequest("second.example.", 2), Options{}); err != nil {
t.Fatalf("second Connect() failed: %v", err)
}
if got := server.accepted.Load(); got != 2 {
t.Fatalf("accepted connections = %d, want 2", got)
}
})
}
}
func TestDoQStopCancelsDial(t *testing.T) {
packetConn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatalf("net.ListenUDP() failed: %v", err)
}
defer packetConn.Close()
p := NewProxy("TestDoQStopCancelsDial", packetConn.LocalAddr().String(), transport.QUIC)
p.SetTLSConfig(&tls.Config{ServerName: "doq.test"})
result := make(chan error, 1)
go func() {
_, _, _, err := p.Connect(context.Background(), doqTestRequest("example.org.", 1), Options{})
result <- err
}()
time.Sleep(50 * time.Millisecond)
p.Stop()
select {
case err := <-result:
if err == nil {
t.Fatal("Connect() unexpectedly succeeded")
}
case <-time.After(time.Second):
t.Fatal("Stop() did not cancel the in-progress DoQ dial")
}
p.Stop()
_, _, _, err = p.Connect(context.Background(), doqTestRequest("example.org.", 2), Options{})
if err == nil || err.Error() != ErrTransportStopped {
t.Fatalf("Connect() after Stop() error = %v, want %q", err, ErrTransportStopped)
}
}
func TestDoQFraming(t *testing.T) {
var framed bytes.Buffer
if err := writeDOQMessage(&framed, []byte{1, 2, 3}); err != nil {
t.Fatalf("writeDOQMessage() failed: %v", err)
}
if want := []byte{0, 3, 1, 2, 3}; !bytes.Equal(framed.Bytes(), want) {
t.Fatalf("framed message = %v, want %v", framed.Bytes(), want)
}
message, err := readDOQMessage(&framed)
if err != nil {
t.Fatalf("readDOQMessage() failed: %v", err)
}
if !bytes.Equal(message, []byte{1, 2, 3}) {
t.Fatalf("message = %v, want [1 2 3]", message)
}
invalid := [][]byte{
{},
{0},
{0, 0},
{0, 2, 1},
}
for _, wire := range invalid {
if _, err := readDOQMessage(bytes.NewReader(wire)); !errors.Is(err, errDoQProtocol) {
t.Errorf("readDOQMessage(%v) error = %v, want protocol error", wire, err)
}
}
if err := expectDOQFIN(bytes.NewReader(nil)); err != nil {
t.Errorf("expectDOQFIN(empty) failed: %v", err)
}
if err := expectDOQFIN(bytes.NewReader([]byte{1})); !errors.Is(err, errDoQProtocol) {
t.Errorf("expectDOQFIN(extra byte) error = %v, want protocol error", err)
}
}

View File

@@ -11,6 +11,8 @@ var (
ErrNoForward = errors.New("no forwarder defined")
// ErrCachedClosed means cached connection was closed by peer.
ErrCachedClosed = errors.New("cached connection was closed by peer")
// ErrUnsupportedRequest means the proxy transport cannot represent the request.
ErrUnsupportedRequest = errors.New("proxy: unsupported request")
)
// Options holds various Options that can be set.

View File

@@ -3,6 +3,7 @@ package proxy
import (
"context"
"crypto/tls"
"errors"
"net"
"net/http"
"sync/atomic"
@@ -73,12 +74,90 @@ func NewHealthChecker(proxyName, protocol string, recursionDesired bool, domain
domain: domain,
proxyName: proxyName,
}
case transport.QUIC:
return &doqHc{
recursionDesired: recursionDesired,
domain: domain,
proxyName: proxyName,
readTimeout: defaultTimeout,
writeTimeout: defaultTimeout,
}
}
log.Warningf("No healthchecker for transport %q", protocol)
return nil
}
// doqHc is a health checker for a DNS-over-QUIC endpoint. It uses the same
// reusable QUIC connection as normal forwarded queries.
type doqHc struct {
tlsConfig *tls.Config
recursionDesired bool
domain string
proxyName string
localAddress net.IP
readTimeout time.Duration
writeTimeout time.Duration
}
func (h *doqHc) Check(p *Proxy) error {
if p.doq == nil {
return errors.New("proxy: DoQ transport is not initialized")
}
ping := new(dns.Msg)
ping.SetQuestion(h.domain, dns.TypeNS)
ping.RecursionDesired = h.recursionDesired
timeout := max(h.readTimeout, h.writeTimeout)
if timeout <= 0 {
timeout = defaultTimeout
}
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
var err error
for range 2 {
_, _, err = p.doq.exchange(ctx, ping, timeout)
if !errors.Is(err, ErrCachedClosed) {
break
}
}
if err != nil {
healthcheckFailureCount.WithLabelValues(p.proxyName, p.addr).Inc()
p.incrementFails()
return err
}
atomic.StoreUint32(&p.fails, 0)
return nil
}
func (h *doqHc) SetTLSConfig(cfg *tls.Config) { h.tlsConfig = cfg }
func (h *doqHc) GetTLSConfig() *tls.Config { return h.tlsConfig }
func (h *doqHc) SetRecursionDesired(v bool) { h.recursionDesired = v }
func (h *doqHc) GetRecursionDesired() bool { return h.recursionDesired }
func (h *doqHc) SetDomain(domain string) { h.domain = domain }
func (h *doqHc) GetDomain() string { return h.domain }
func (h *doqHc) SetTCPTransport() {}
func (h *doqHc) GetReadTimeout() time.Duration {
return h.readTimeout
}
func (h *doqHc) SetReadTimeout(timeout time.Duration) {
h.readTimeout = timeout
}
func (h *doqHc) GetWriteTimeout() time.Duration {
return h.writeTimeout
}
func (h *doqHc) SetWriteTimeout(timeout time.Duration) {
h.writeTimeout = timeout
}
func (h *doqHc) SetLocalAddress(addr net.IP) {
h.localAddress = append(net.IP(nil), addr...)
}
func (h *doqHc) GetLocalAddress() net.IP {
return append(net.IP(nil), h.localAddress...)
}
func (h *dnsHc) SetTLSConfig(cfg *tls.Config) {
h.c.Net = "tcp-tls"
h.c.TLSConfig = cfg

View File

@@ -9,6 +9,7 @@ import (
"time"
"github.com/coredns/coredns/plugin/pkg/log"
"github.com/coredns/coredns/plugin/pkg/transport"
"github.com/coredns/coredns/plugin/pkg/up"
)
@@ -19,6 +20,7 @@ type Proxy struct {
proxyName string
transport *Transport
doq *doqTransport
protocol string
dohMethod string
@@ -45,6 +47,9 @@ func NewProxy(proxyName, addr, protocol string) *Proxy {
health: NewHealthChecker(proxyName, protocol, true, "."),
proxyName: proxyName,
}
if protocol == transport.QUIC {
p.doq = newDoQTransport(proxyName, addr)
}
runtime.SetFinalizer(p, (*Proxy).finalizer)
return p
@@ -55,18 +60,33 @@ func (p *Proxy) Addr() string { return p.addr }
// SetTLSConfig sets the TLS config in the lower p.transport and in the healthchecking client.
func (p *Proxy) SetTLSConfig(cfg *tls.Config) {
p.transport.SetTLSConfig(cfg)
p.health.SetTLSConfig(cfg)
if p.doq != nil {
p.doq.setTLSConfig(cfg)
}
if p.health != nil {
p.health.SetTLSConfig(cfg)
}
if p.transport.httpClient != nil {
p.transport.httpClient.Transport.(*http.Transport).TLSClientConfig = cfg
}
}
// SetExpire sets the expire duration in the lower p.transport.
func (p *Proxy) SetExpire(expire time.Duration) { p.transport.SetExpire(expire) }
func (p *Proxy) SetExpire(expire time.Duration) {
p.transport.SetExpire(expire)
if p.doq != nil {
p.doq.setExpire(expire)
}
}
// SetMaxAge sets the maximum connection lifetime in the lower p.transport.
// A value of 0 (default) disables max-age.
func (p *Proxy) SetMaxAge(maxAge time.Duration) { p.transport.SetMaxAge(maxAge) }
func (p *Proxy) SetMaxAge(maxAge time.Duration) {
p.transport.SetMaxAge(maxAge)
if p.doq != nil {
p.doq.setMaxAge(maxAge)
}
}
// SetMaxIdleConns sets the maximum idle connections per transport type.
// A value of 0 means unlimited (default).
@@ -126,18 +146,37 @@ func (p *Proxy) Down(maxfails uint32) bool {
return fails > maxfails
}
// Stop close stops the health checking goroutine.
func (p *Proxy) Stop() { p.probe.Stop() }
func (p *Proxy) finalizer() { p.transport.Stop() }
// Stop stops health checking and closes the DoQ transport, when configured.
func (p *Proxy) Stop() {
p.probe.Stop()
if p.doq != nil {
p.doq.stopTransport()
}
}
func (p *Proxy) finalizer() {
if p.doq != nil {
p.doq.stopTransport()
return
}
p.transport.Stop()
}
// Start starts the proxy's healthchecking.
func (p *Proxy) Start(duration time.Duration) {
p.probe.Start(duration)
if p.doq != nil {
p.doq.start()
return
}
p.transport.Start()
}
func (p *Proxy) SetReadTimeout(duration time.Duration) {
p.readTimeout = duration
if p.doq != nil {
p.doq.setReadTimeout(duration)
}
}
// incrementFails increments the number of fails safely.
@@ -153,6 +192,9 @@ func (p *Proxy) incrementFails() {
// SetLocalAddress sets the local address for the proxy, used as the source address for outbound connections.
func (p *Proxy) SetLocalAddress(addr net.IP) {
p.transport.SetLocalAddress(addr)
if p.doq != nil {
p.doq.setLocalAddress(addr)
}
if p.transport.httpClient != nil {
httpTransport := p.transport.httpClient.Transport.(*http.Transport)
if addr == nil {