diff --git a/dns/doq.go b/dns/doq.go index 401e1bd7..02558321 100644 --- a/dns/doq.go +++ b/dns/doq.go @@ -5,6 +5,7 @@ import ( "encoding/binary" "errors" "fmt" + "io" "net" "runtime" "strconv" @@ -153,18 +154,27 @@ func (doq *dnsOverQUIC) ResetConnection() { // exchangeQUIC attempts to open a QUIC connection, send the DNS message // through it and return the response it got from the server. func (doq *dnsOverQUIC) exchangeQUIC(ctx context.Context, msg *D.Msg) (resp *D.Msg, err error) { + // All DNS messages (queries and responses) sent over DoQ connections MUST + // be encoded as a 2-octet length field followed by the message content as + // specified in [RFC1035]. + // Note: we do not support receiving multiple messages over a single connection. + buf := pool.Get(2 + MaxMsgSize) + defer pool.Put(buf) + b, err := msg.PackBuffer(buf[2:]) + if err != nil { + return nil, fmt.Errorf("failed to pack DNS message for DoQ: %w", err) + } + if len(b) > MaxMsgSize { + return nil, fmt.Errorf("DNS message is too large: %d > %d", len(b), MaxMsgSize) + } + binary.BigEndian.PutUint16(buf, uint16(len(b))) + var conn *quic.Conn conn, err = doq.getConnection(ctx, true) if err != nil { return nil, err } - var buf []byte - buf, err = msg.Pack() - if err != nil { - return nil, fmt.Errorf("failed to pack DNS message for DoQ: %w", err) - } - var stream *quic.Stream stream, err = doq.openStream(ctx, conn) if err != nil { @@ -176,7 +186,7 @@ func (doq *dnsOverQUIC) exchangeQUIC(ctx context.Context, msg *D.Msg) (resp *D.M }) defer stop() - _, err = stream.Write(AddPrefix(buf)) + _, err = stream.Write(buf[:2+len(b)]) if err != nil { return nil, fmt.Errorf("failed to write to a QUIC stream: %w", err) } @@ -187,16 +197,31 @@ func (doq *dnsOverQUIC) exchangeQUIC(ctx context.Context, msg *D.Msg) (resp *D.M // write-direction of the stream, but does not prevent reading from it. _ = stream.Close() - return doq.readMsg(stream) -} + // -- reading the response --- + var respLen uint16 + err = binary.Read(stream, binary.BigEndian, &respLen) + if err != nil { + return nil, fmt.Errorf("reading response length from %s: %w", doq.Address(), err) + } + if respLen == 0 { + return nil, fmt.Errorf("received empty response from %s", doq.Address()) + } + if respLen > MaxMsgSize { + return nil, fmt.Errorf("received response that is too large: %d > %d", respLen, MaxMsgSize) + } -// AddPrefix adds a 2-byte prefix with the DNS message length. -func AddPrefix(b []byte) (m []byte) { - m = make([]byte, 2+len(b)) - binary.BigEndian.PutUint16(m, uint16(len(b))) - copy(m[2:], b) + _, err = io.ReadFull(stream, buf[:respLen]) + if err != nil { + return nil, fmt.Errorf("reading response from %s: %w", doq.Address(), err) + } - return m + resp = new(D.Msg) + err = resp.Unpack(buf[:respLen]) + if err != nil { + return nil, fmt.Errorf("unpacking response from %s: %w", doq.Address(), err) + } + + return resp, nil } // shouldRetry checks what error we received and decides whether it is required @@ -369,30 +394,6 @@ func (doq *dnsOverQUIC) closeConnWithError(err error) { doq.conn = nil } -// readMsg reads the incoming DNS message from the QUIC stream. -func (doq *dnsOverQUIC) readMsg(stream *quic.Stream) (m *D.Msg, err error) { - respBuf := pool.Get(MaxMsgSize) - defer pool.Put(respBuf) - - n, err := stream.Read(respBuf) - if err != nil && n == 0 { - return nil, fmt.Errorf("reading response from %s: %w", doq.Address(), err) - } - - // All DNS messages (queries and responses) sent over DoQ connections MUST - // be encoded as a 2-octet length field followed by the message content as - // specified in [RFC1035]. - // IMPORTANT: Note, that we ignore this prefix here as this implementation - // does not support receiving multiple messages over a single connection. - m = new(D.Msg) - err = m.Unpack(respBuf[2:]) - if err != nil { - return nil, fmt.Errorf("unpacking response from %s: %w", doq.Address(), err) - } - - return m, nil -} - // newQUICTokenStore creates a new quic.TokenStore that is necessary to have // in order to benefit from 0-RTT. func newQUICTokenStore() (s quic.TokenStore) {