mirror of
https://github.com/coredns/coredns.git
synced 2026-10-09 03:55:21 -04:00
plugin/forward: support DNS-over-QUIC upstreams (#8474)
Signed-off-by: houyuwushang <liuluoqianqiu@outlook.com>
This commit is contained in:
@@ -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
548
plugin/pkg/proxy/doq.go
Normal 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
|
||||
}
|
||||
}
|
||||
}
|
||||
670
plugin/pkg/proxy/doq_test.go
Normal file
670
plugin/pkg/proxy/doq_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user