mirror of
https://github.com/coredns/coredns.git
synced 2026-10-09 03:55:21 -04:00
549 lines
12 KiB
Go
549 lines
12 KiB
Go
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
|
|
}
|
|
}
|
|
}
|