Skip to main content

mz_pgwire_common/
codec.rs

1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10//! Encoding/decoding of messages in pgwire. See "[Frontend/Backend Protocol:
11//! Message Formats][1]" in the PostgreSQL reference for the specification.
12//!
13//! See the [crate docs](crate) for higher level concerns.
14//!
15//! [1]: https://www.postgresql.org/docs/11/protocol-message-formats.html
16
17use std::collections::BTreeMap;
18use std::error::Error;
19use std::{fmt, str};
20
21use byteorder::{ByteOrder, NetworkEndian};
22use bytes::{BufMut, BytesMut};
23use mz_ore::cast::{CastFrom, u64_to_usize};
24use mz_ore::netio::{self};
25use tokio::io::{self, AsyncRead, AsyncReadExt};
26
27use crate::FrontendMessage;
28use crate::format::Format;
29use crate::message::{FrontendStartupMessage, VERSION_CANCEL, VERSION_GSSENC, VERSION_SSL};
30
31pub const REJECT_ENCRYPTION: u8 = b'N';
32pub const ACCEPT_SSL_ENCRYPTION: u8 = b'S';
33
34/// Maximum allowed size for a request.
35pub const MAX_REQUEST_SIZE: usize = u64_to_usize(2 * bytesize::MB);
36
37/// Maximum size of a startup frame accepted directly from a client.
38///
39/// Matches PostgreSQL's `MAX_STARTUP_PACKET_LENGTH`, so any client that can
40/// complete a PostgreSQL handshake can complete this one. A startup frame
41/// carries only the protocol version and the connection parameters, of which
42/// `options` is the only one a client can make large.
43pub const MAX_STARTUP_FRAME_SIZE: usize = 10_000;
44
45/// Startup budget allowed on top of [`MAX_STARTUP_FRAME_SIZE`] for parameters a
46/// balancer appends while forwarding.
47///
48/// A balancer adds [`CONN_UUID_KEY`] and [`MZ_FORWARDED_FOR_KEY`] to the
49/// parameters before forwarding startup, so a frame that just fits the client
50/// budget arrives downstream larger than the client sent it. Without this
51/// allowance such a connection is accepted by the balancer and then rejected
52/// behind it, which surfaces as a proxy error with no obvious cause.
53///
54/// The two parameters cost at most 119 bytes: an 18-byte key with a 36-byte
55/// UUID, a 16-byte key with an address of up to 45 bytes, and a NUL after each
56/// of the four strings. The rest is headroom, so that adding a third forwarded
57/// parameter does not immediately require re-deriving this number. Going over it
58/// is caught by `test_forwarded_startup_frame_fits_downstream_budget`, which
59/// measures what a balancer actually puts on the wire rather than trusting this
60/// comment.
61///
62/// [`CONN_UUID_KEY`]: crate::CONN_UUID_KEY
63/// [`MZ_FORWARDED_FOR_KEY`]: crate::MZ_FORWARDED_FOR_KEY
64pub const FORWARDED_STARTUP_PARAM_ALLOWANCE: usize = 512;
65
66/// Maximum size of a startup frame accepted from a client that may be behind a
67/// balancer.
68pub const MAX_FORWARDED_STARTUP_FRAME_SIZE: usize =
69    MAX_STARTUP_FRAME_SIZE + FORWARDED_STARTUP_PARAM_ALLOWANCE;
70
71/// Maximum frame size accepted from a client that has not yet authenticated.
72///
73/// The only frames a client legitimately sends before authenticating are
74/// credentials: a password, or one leg of a SASL exchange. This sits far above
75/// any of those while leaving room for a bearer token, which can run to several
76/// kilobytes.
77pub const MAX_PREAUTH_FRAME_SIZE: usize = u64_to_usize(16 * bytesize::KIB);
78
79#[derive(Debug)]
80pub enum CodecError {
81    StringNoTerminator,
82}
83
84impl Error for CodecError {}
85
86impl fmt::Display for CodecError {
87    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
88        f.write_str(match self {
89            CodecError::StringNoTerminator => "The string does not have a terminator",
90        })
91    }
92}
93
94pub trait Pgbuf: BufMut {
95    fn put_string(&mut self, s: &str);
96    fn put_length_i16(&mut self, len: usize) -> Result<(), io::Error>;
97    fn put_length_u16(&mut self, len: usize) -> Result<(), io::Error>;
98    fn put_format_i8(&mut self, format: Format);
99    fn put_format_i16(&mut self, format: Format);
100}
101
102impl<B: BufMut> Pgbuf for B {
103    fn put_string(&mut self, s: &str) {
104        self.put(s.as_bytes());
105        self.put_u8(b'\0');
106    }
107
108    fn put_length_i16(&mut self, len: usize) -> Result<(), io::Error> {
109        let len = i16::try_from(len).map_err(|_| {
110            io::Error::new(io::ErrorKind::InvalidData, "length does not fit in an i16")
111        })?;
112        self.put_i16(len);
113        Ok(())
114    }
115
116    /// Writes a count field as unsigned, so it may exceed 32767. The protocol
117    /// calls these fields `Int16`, but PostgreSQL clients decode them as
118    /// unsigned.
119    fn put_length_u16(&mut self, len: usize) -> Result<(), io::Error> {
120        let len = u16::try_from(len).map_err(|_| {
121            io::Error::new(io::ErrorKind::InvalidData, "length does not fit in a u16")
122        })?;
123        self.put_u16(len);
124        Ok(())
125    }
126
127    fn put_format_i8(&mut self, format: Format) {
128        self.put_i8(format.into())
129    }
130
131    fn put_format_i16(&mut self, format: Format) {
132        self.put_i8(0);
133        self.put_format_i8(format);
134    }
135}
136
137/// Reads and decodes one startup message from the client.
138///
139/// `max_frame_len` bounds the frame the client may declare, including its own
140/// four-byte length field. It is checked before the body buffer is sized, so the
141/// buffer follows the caller's ceiling rather than the protocol maximum. Pass
142/// [`MAX_STARTUP_FRAME_SIZE`], or [`MAX_FORWARDED_STARTUP_FRAME_SIZE`] if a
143/// balancer may have appended parameters in transit.
144pub async fn decode_startup<A>(
145    mut conn: A,
146    max_frame_len: usize,
147) -> Result<Option<FrontendStartupMessage>, io::Error>
148where
149    A: AsyncRead + Unpin,
150{
151    let mut frame_len = [0; 4];
152    let nread = netio::read_exact_or_eof(&mut conn, &mut frame_len).await?;
153    match nread {
154        // Complete frame length. Continue.
155        4 => (),
156        // Connection closed cleanly. Indicate that the startup sequence has
157        // been terminated by the client.
158        0 => return Ok(None),
159        // Partial frame length. Likely a client bug or network glitch, so
160        // surface the unexpected EOF.
161        _ => return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "early eof")),
162    };
163    let frame_len = parse_frame_len(&frame_len, max_frame_len)?;
164
165    let mut buf = BytesMut::new();
166    buf.resize(frame_len, b'0');
167    conn.read_exact(&mut buf).await?;
168
169    let mut buf = Cursor::new(&buf);
170    let version = buf.read_i32()?;
171    let message = match version {
172        VERSION_CANCEL => FrontendStartupMessage::CancelRequest {
173            conn_id: buf.read_u32()?,
174            secret_key: buf.read_u32()?,
175        },
176        VERSION_SSL => FrontendStartupMessage::SslRequest,
177        VERSION_GSSENC => FrontendStartupMessage::GssEncRequest,
178        _ => {
179            let mut params = BTreeMap::new();
180            while buf.peek_byte()? != 0 {
181                let name = buf.read_cstr()?.to_owned();
182                let value = buf.read_cstr()?.to_owned();
183                params.insert(name, value);
184            }
185            FrontendStartupMessage::Startup { version, params }
186        }
187    };
188    Ok(Some(message))
189}
190
191impl FrontendStartupMessage {
192    /// Encodes self into dst.
193    pub fn encode(&self, dst: &mut BytesMut) -> Result<(), io::Error> {
194        // Write message length placeholder. The true length is filled in later.
195        let base = dst.len();
196        dst.put_u32(0);
197
198        // Write message contents.
199        match self {
200            FrontendStartupMessage::Startup { version, params } => {
201                dst.put_i32(*version);
202                for (k, v) in params {
203                    dst.put_string(k);
204                    dst.put_string(v);
205                }
206                dst.put_i8(0);
207            }
208            FrontendStartupMessage::CancelRequest {
209                conn_id,
210                secret_key,
211            } => {
212                dst.put_i32(VERSION_CANCEL);
213                dst.put_u32(*conn_id);
214                dst.put_u32(*secret_key);
215            }
216            FrontendStartupMessage::SslRequest {} => dst.put_i32(VERSION_SSL),
217            FrontendStartupMessage::GssEncRequest => panic!("unsupported"),
218        }
219
220        let len = dst.len() - base;
221
222        // Overwrite length placeholder with true length.
223        let len = i32::try_from(len).map_err(|_| {
224            io::Error::new(
225                io::ErrorKind::InvalidData,
226                "length of encoded message does not fit into an i32",
227            )
228        })?;
229        dst[base..base + 4].copy_from_slice(&len.to_be_bytes());
230
231        Ok(())
232    }
233}
234
235impl FrontendMessage {
236    /// Encodes self into dst.
237    pub fn encode(&self, dst: &mut BytesMut) -> Result<(), io::Error> {
238        // Write type byte.
239        let byte = match self {
240            FrontendMessage::Password { .. } => b'p',
241            _ => panic!("unsupported"),
242        };
243        dst.put_u8(byte);
244
245        // Write message length placeholder. The true length is filled in later.
246        let base = dst.len();
247        dst.put_u32(0);
248
249        // Write message contents.
250        match self {
251            FrontendMessage::Password { password } => {
252                dst.put_string(password);
253            }
254            _ => panic!("unsupported"),
255        }
256
257        let len = dst.len() - base;
258
259        // Overwrite length placeholder with true length.
260        let len = i32::try_from(len).map_err(|_| {
261            io::Error::new(
262                io::ErrorKind::InvalidData,
263                "length of encoded message does not fit into an i32",
264            )
265        })?;
266        dst[base..base + 4].copy_from_slice(&len.to_be_bytes());
267
268        Ok(())
269    }
270}
271
272#[derive(Debug)]
273pub enum DecodeState {
274    Head,
275    Data(u8, usize),
276}
277
278/// Parses a frame length header, rejecting frames larger than `max_frame_len`.
279///
280/// The ceiling is the caller's rather than a single global one, because the
281/// paths that parse frame lengths have requirements that differ by orders of
282/// magnitude: a startup frame is well under 10 KB, a pre-authentication
283/// credential exchange needs a few kilobytes, and only post-authentication
284/// query traffic needs room for bulk data. Stating it per call site keeps a
285/// buffer from ever being sized against a bound that belongs to a different
286/// path.
287///
288/// `max_frame_len` counts the frame including its own four-byte length field,
289/// matching the number the client declares. [`netio::MAX_FRAME_SIZE`] is the
290/// protocol ceiling; callers pass that or less.
291pub fn parse_frame_len(src: &[u8], max_frame_len: usize) -> Result<usize, io::Error> {
292    let n = usize::cast_from(NetworkEndian::read_u32(src));
293    if n > max_frame_len {
294        // The limit differs per call site, so naming both numbers is what makes a
295        // rejection diagnosable: "too big" alone no longer says which bound was hit.
296        return Err(io::Error::new(
297            io::ErrorKind::InvalidData,
298            format!("frame of {n} bytes exceeds the {max_frame_len} byte limit"),
299        ));
300    } else if n < 4 {
301        return Err(io::Error::new(
302            io::ErrorKind::InvalidInput,
303            "invalid frame length",
304        ));
305    }
306    Ok(n - 4)
307}
308
309/// Decodes data within pgwire messages.
310///
311/// The API provided is very similar to [`bytes::Buf`], but operations return
312/// errors rather than panicking. This is important for safety, as we don't want
313/// to crash if the user sends us malformed pgwire messages.
314///
315/// There are also some special-purpose methods, like [`Cursor::read_cstr`],
316/// that are specific to pgwire messages.
317#[derive(Debug)]
318pub struct Cursor<'a> {
319    buf: &'a [u8],
320}
321
322impl<'a> Cursor<'a> {
323    /// Constructs a new `Cursor` from a byte slice. The cursor will begin
324    /// decoding from the beginning of the slice.
325    pub fn new(buf: &'a [u8]) -> Cursor<'a> {
326        Cursor { buf }
327    }
328
329    /// Returns the next byte without advancing the cursor.
330    pub fn peek_byte(&self) -> Result<u8, io::Error> {
331        self.buf
332            .get(0)
333            .copied()
334            .ok_or_else(|| input_err("No byte to read"))
335    }
336
337    /// Returns the next byte, advancing the cursor by one byte.
338    pub fn read_byte(&mut self) -> Result<u8, io::Error> {
339        let byte = self.peek_byte()?;
340        self.advance(1);
341        Ok(byte)
342    }
343
344    /// Returns the next null-terminated string. The null character is not
345    /// included the returned string. The cursor is advanced past the null-
346    /// terminated string.
347    ///
348    /// If there is no null byte remaining in the string, returns
349    /// `CodecError::StringNoTerminator`. If the string is not valid UTF-8,
350    /// returns an `io::Error` with an error kind of
351    /// `io::ErrorKind::InvalidInput`.
352    ///
353    /// NOTE(benesch): it is possible that returning a string here is wrong, and
354    /// we should be returning bytes, so that we can support messages that are
355    /// not UTF-8 encoded. At the moment, we've not discovered a need for this,
356    /// though, and using proper strings is convenient.
357    pub fn read_cstr(&mut self) -> Result<&'a str, io::Error> {
358        if let Some(pos) = self.buf.iter().position(|b| *b == 0) {
359            let val = std::str::from_utf8(&self.buf[..pos]).map_err(input_err)?;
360            self.advance(pos + 1);
361            Ok(val)
362        } else {
363            Err(input_err(CodecError::StringNoTerminator))
364        }
365    }
366
367    /// Reads the next 32-bit signed integer, advancing the cursor by four
368    /// bytes.
369    pub fn read_i32(&mut self) -> Result<i32, io::Error> {
370        if self.buf.len() < 4 {
371            return Err(input_err("not enough buffer for an Int32"));
372        }
373        let val = NetworkEndian::read_i32(self.buf);
374        self.advance(4);
375        Ok(val)
376    }
377
378    /// Reads the next 16-bit unsigned integer, advancing the cursor by two
379    /// bytes.
380    pub fn read_u16(&mut self) -> Result<u16, io::Error> {
381        if self.buf.len() < 2 {
382            return Err(input_err("not enough buffer for an Int16"));
383        }
384        let val = NetworkEndian::read_u16(self.buf);
385        self.advance(2);
386        Ok(val)
387    }
388
389    /// Reads the next 32-bit unsigned integer, advancing the cursor by four
390    /// bytes.
391    pub fn read_u32(&mut self) -> Result<u32, io::Error> {
392        if self.buf.len() < 4 {
393            return Err(input_err("not enough buffer for an Int32"));
394        }
395        let val = NetworkEndian::read_u32(self.buf);
396        self.advance(4);
397        Ok(val)
398    }
399
400    /// Reads the next 16-bit format code, advancing the cursor by two bytes.
401    pub fn read_format(&mut self) -> Result<Format, io::Error> {
402        Format::try_from(self.read_u16()?)
403    }
404
405    /// Advances the cursor by `n` bytes.
406    pub fn advance(&mut self, n: usize) {
407        self.buf = &self.buf[n..]
408    }
409}
410
411/// Constructs an error indicating that the client has violated the pgwire
412/// protocol.
413pub fn input_err(source: impl Into<Box<dyn std::error::Error + Send + Sync>>) -> io::Error {
414    io::Error::new(io::ErrorKind::InvalidInput, source.into())
415}
416
417#[cfg(test)]
418mod tests {
419    use std::pin::Pin;
420    use std::task::{Context, Poll};
421    use std::time::Duration;
422
423    use crate::conn::{CONN_UUID_KEY, MZ_FORWARDED_FOR_KEY};
424    use crate::message::VERSION_3;
425    use tokio::io::ReadBuf;
426
427    use super::*;
428
429    /// The buffer the server offered the client to fill.
430    ///
431    /// The server can only hand out a slice it has already sized, so an offer is
432    /// evidence that the declared length was trusted before the body arrived,
433    /// and its absence is evidence that the frame was rejected first.
434    #[derive(Debug)]
435    struct Offer {
436        len: usize,
437    }
438
439    /// A client that writes a fixed prefix and then goes silent forever.
440    ///
441    /// Going silent means returning `Poll::Pending` without registering a waker,
442    /// so nothing can ever resume the read. Any stall is therefore a property of
443    /// the protocol handling, not a scheduling artifact.
444    struct SilentClient {
445        prefix: Vec<u8>,
446        sent: usize,
447        offer: Option<Offer>,
448    }
449
450    impl AsyncRead for SilentClient {
451        fn poll_read(
452            mut self: Pin<&mut Self>,
453            _cx: &mut Context<'_>,
454            buf: &mut ReadBuf<'_>,
455        ) -> Poll<io::Result<()>> {
456            let unsent = self.prefix.len() - self.sent;
457            if unsent > 0 {
458                let n = std::cmp::min(unsent, buf.remaining());
459                let start = self.sent;
460                buf.put_slice(&self.prefix[start..start + n]);
461                self.sent += n;
462                return Poll::Ready(Ok(()));
463            }
464            self.offer = Some(Offer {
465                len: buf.remaining(),
466            });
467            Poll::Pending
468        }
469    }
470
471    struct Attempt {
472        offer: Option<Offer>,
473        /// `None` if `decode_startup` was still waiting an hour on.
474        result: Option<Result<Option<FrontendStartupMessage>, io::Error>>,
475    }
476
477    /// Feeds `frame` to `decode_startup` under `max_frame_len`, then goes silent
478    /// and lets an hour of (virtual) time pass.
479    async fn attempt(frame: &[u8], max_frame_len: usize) -> Attempt {
480        let mut client = SilentClient {
481            prefix: frame.to_vec(),
482            sent: 0,
483            offer: None,
484        };
485        // The runtime has nothing to run once the client goes silent, so the
486        // paused clock jumps straight to the deadline and the hour is free.
487        let result = tokio::time::timeout(
488            Duration::from_secs(60 * 60),
489            decode_startup(&mut client, max_frame_len),
490        )
491        .await;
492        Attempt {
493            offer: client.offer,
494            result: result.ok(),
495        }
496    }
497
498    /// Startup parameters whose encoded frame is exactly `frame_len` bytes.
499    fn params_sized_to(frame_len: usize) -> BTreeMap<String, String> {
500        // Length, version, the key and its NUL, the value's NUL, terminator.
501        let overhead = 4 + 4 + "options".len() + 1 + 1 + 1;
502        BTreeMap::from([("options".to_string(), "x".repeat(frame_len - overhead))])
503    }
504
505    #[mz_ore::test(tokio::test(start_paused = true))]
506    async fn test_startup_frame_over_budget_is_rejected_before_allocating() {
507        for declared in [MAX_STARTUP_FRAME_SIZE + 1, 1 << 20, netio::MAX_FRAME_SIZE] {
508            let header = u32::try_from(declared)
509                .expect("fits in a frame-length field")
510                .to_be_bytes();
511            let attempt = attempt(&header, MAX_STARTUP_FRAME_SIZE).await;
512
513            let err = attempt
514                .result
515                .expect("decode_startup stalled instead of rejecting the frame")
516                .expect_err("oversized startup frame was accepted");
517            assert_eq!(
518                err.kind(),
519                io::ErrorKind::InvalidData,
520                "declared {declared}"
521            );
522            assert_eq!(
523                attempt.offer.as_ref().map(|offer| offer.len),
524                None,
525                "declared {declared}: a buffer was sized before the frame was rejected",
526            );
527        }
528    }
529
530    #[mz_ore::test(tokio::test(start_paused = true))]
531    async fn test_startup_frame_within_budget_is_accepted() {
532        let params = params_sized_to(MAX_STARTUP_FRAME_SIZE);
533        let mut frame = BytesMut::new();
534        FrontendStartupMessage::Startup {
535            version: VERSION_3,
536            params: params.clone(),
537        }
538        .encode(&mut frame)
539        .expect("encodes");
540        assert_eq!(frame.len(), MAX_STARTUP_FRAME_SIZE);
541
542        let message = attempt(&frame, MAX_STARTUP_FRAME_SIZE)
543            .await
544            .result
545            .expect("decode_startup stalled on a complete frame")
546            .expect("a startup frame at the budget was rejected");
547        match message {
548            Some(FrontendStartupMessage::Startup {
549                version,
550                params: decoded,
551            }) => {
552                assert_eq!(version, VERSION_3);
553                assert_eq!(decoded, params);
554            }
555            other => panic!("expected a startup message, got {other:?}"),
556        }
557    }
558
559    /// A balancer appends two parameters while forwarding, so a frame that just
560    /// fits the client budget grows in transit. The downstream bound has to
561    /// cover the difference, or the connection is accepted by the balancer and
562    /// rejected behind it.
563    #[mz_ore::test(tokio::test(start_paused = true))]
564    async fn test_forwarded_startup_params_fit_the_allowance() {
565        // Widest values the two parameters can carry: a hyphenated UUID, and an
566        // IPv4-mapped IPv6 address in its longest textual form.
567        const WIDEST_UUID: &str = "00000000-0000-0000-0000-000000000000";
568        const WIDEST_ADDR: &str = "ffff:ffff:ffff:ffff:ffff:ffff:255.255.255.255";
569
570        let mut params = params_sized_to(MAX_STARTUP_FRAME_SIZE);
571        params.insert(CONN_UUID_KEY.to_string(), WIDEST_UUID.to_string());
572        params.insert(MZ_FORWARDED_FOR_KEY.to_string(), WIDEST_ADDR.to_string());
573
574        let mut forwarded = BytesMut::new();
575        FrontendStartupMessage::Startup {
576            version: VERSION_3,
577            params,
578        }
579        .encode(&mut forwarded)
580        .expect("encodes");
581
582        assert!(
583            forwarded.len() <= MAX_FORWARDED_STARTUP_FRAME_SIZE,
584            "FORWARDED_STARTUP_PARAM_ALLOWANCE is too small: {} bytes forwarded \
585             against a {MAX_FORWARDED_STARTUP_FRAME_SIZE} byte bound",
586            forwarded.len(),
587        );
588        assert!(
589            attempt(&forwarded, MAX_FORWARDED_STARTUP_FRAME_SIZE)
590                .await
591                .result
592                .expect("decode_startup stalled on a complete frame")
593                .is_ok(),
594        );
595    }
596
597    /// A client that declares a frame within budget and then stops sending
598    /// leaves the server holding a buffer sized for that declaration, so the
599    /// budget is what bounds it.
600    #[mz_ore::test(tokio::test(start_paused = true))]
601    async fn test_startup_buffer_is_bounded_by_the_budget() {
602        let header = u32::try_from(MAX_STARTUP_FRAME_SIZE)
603            .expect("fits in a frame-length field")
604            .to_be_bytes();
605        let attempt = attempt(&header, MAX_STARTUP_FRAME_SIZE).await;
606
607        assert_eq!(
608            attempt.offer.as_ref().map(|offer| offer.len),
609            Some(MAX_STARTUP_FRAME_SIZE - 4),
610        );
611    }
612}