Skip to main content

mz_pgwire/
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::net::IpAddr;
18
19use async_trait::async_trait;
20use bytes::{Buf, BufMut, BytesMut};
21use futures::{SinkExt, TryStreamExt, sink};
22use itertools::Itertools;
23use mz_adapter_types::connection::ConnectionId;
24use mz_ore::future::OreSinkExt;
25use mz_ore::netio::AsyncReady;
26use mz_pgwire_common::{
27    ChannelBinding, Conn, Cursor, DecodeState, ErrorResponse, FrontendMessage, GS2Header,
28    MAX_PREAUTH_FRAME_SIZE, Pgbuf, SASLClientFinalResponse, SASLInitialResponse, input_err,
29    parse_frame_len,
30};
31use tokio::io::{self, AsyncRead, AsyncWrite, Interest, Ready};
32use tokio::time::{self, Duration};
33use tokio_util::codec::{Decoder, Encoder, Framed};
34use tracing::trace;
35
36use crate::message::{BackendMessage, BackendMessageKind, SASLServerFinalMessageKinds};
37
38/// A connection that manages the encoding and decoding of pgwire frames.
39pub struct FramedConn<A> {
40    conn_id: ConnectionId,
41    peer_addr: Option<IpAddr>,
42    inner: sink::Buffer<Framed<Conn<A>, Codec>, BackendMessage>,
43}
44
45impl<A> FramedConn<A>
46where
47    A: AsyncRead + AsyncWrite + Unpin,
48{
49    /// Constructs a new framed connection.
50    ///
51    /// The underlying connection, `inner`, is expected to be something like a
52    /// TCP stream. Anything that implements [`AsyncRead`] and [`AsyncWrite`]
53    /// will do.
54    ///
55    /// The supplied `conn_id` is used to identify the connection in logging
56    /// messages.
57    pub fn new(conn_id: ConnectionId, peer_addr: Option<IpAddr>, inner: Conn<A>) -> FramedConn<A> {
58        FramedConn {
59            conn_id,
60            peer_addr,
61            inner: Framed::new(inner, Codec::new()).buffer(32),
62        }
63    }
64
65    /// Reads and decodes one frontend message from the client.
66    ///
67    /// Blocks until the client sends a complete message. If the client
68    /// terminates the stream, returns `None`. Returns an error if the client
69    /// sends a malformed message or if the connection underlying is broken.
70    ///
71    /// # Cancel safety
72    ///
73    /// This method is cancel safe. The returned future only holds onto a
74    /// reference to thea underlying stream, so dropping it will never lose a
75    /// value.
76    ///
77    /// <https://docs.rs/tokio-stream/latest/tokio_stream/trait.StreamExt.html#cancel-safety-1>
78    pub async fn recv(&mut self) -> Result<Option<FrontendMessage>, io::Error> {
79        let message = self.inner.try_next().await?;
80        match &message {
81            Some(message) => trace!("cid={} recv_name={}", self.conn_id, message.name()),
82            None => trace!("cid={} recv=<eof>", self.conn_id),
83        }
84        Ok(message)
85    }
86
87    /// Encodes and sends one backend message to the client.
88    ///
89    /// Note that the connection is not flushed after calling this method. You
90    /// must call [`FramedConn::flush`] explicitly. Returns an error if the
91    /// underlying connection is broken.
92    ///
93    /// Please use `StateMachine::send` instead if calling from `StateMachine`,
94    /// as it applies session-based filters before calling this method.
95    pub async fn send<M>(&mut self, message: M) -> Result<(), io::Error>
96    where
97        M: Into<BackendMessage>,
98    {
99        let message = message.into();
100        trace!(
101            "cid={} send={:?}",
102            self.conn_id,
103            BackendMessageKind::from(&message)
104        );
105        self.inner.enqueue(message).await
106    }
107
108    /// Encodes and sends the backend messages in the `messages` iterator to the
109    /// client.
110    ///
111    /// As with [`FramedConn::send`], the connection is not flushed after
112    /// calling this method. You must call [`FramedConn::flush`] explicitly.
113    /// Returns an error if the underlying connection is broken.
114    pub async fn send_all(
115        &mut self,
116        messages: impl IntoIterator<Item = BackendMessage>,
117    ) -> Result<(), io::Error> {
118        // N.B. we intentionally don't use `self.conn.send_all` here to avoid
119        // flushing the sink unnecessarily.
120        for m in messages {
121            self.send(m).await?;
122        }
123        Ok(())
124    }
125
126    /// Flushes all outstanding messages.
127    pub async fn flush(&mut self) -> Result<(), io::Error> {
128        self.inner.flush().await
129    }
130
131    /// Injects state that affects how certain backend messages are encoded.
132    ///
133    /// Specifically, the encoding of `BackendMessage::DataRow` depends upon the
134    /// types of the datums in the row. To avoid including the same type
135    /// information in each message, we use this side channel to install the
136    /// type information in the codec before sending any data row messages. This
137    /// violates the abstraction boundary a bit but results in much better
138    /// performance.
139    pub fn set_encode_state(
140        &mut self,
141        encode_state: Vec<(mz_pgrepr::Type, mz_pgwire_common::Format)>,
142        text_settings: mz_pgrepr::TextEncodeSettings,
143    ) {
144        let codec = self.inner.get_mut().codec_mut();
145        codec.encode_state = encode_state;
146        codec.text_settings = text_settings;
147    }
148
149    /// Raises the frame ceiling now that the client has authenticated.
150    ///
151    /// Until this is called the connection is held to
152    /// [`MAX_PREAUTH_FRAME_SIZE`], which fits any credential but not a query.
153    /// Must not be called before authentication succeeds, since that is what
154    /// keeps the two phases on their own ceilings.
155    pub fn allow_post_auth_frames(&mut self) {
156        self.inner.get_mut().codec_mut().max_frame_len = mz_ore::netio::MAX_FRAME_SIZE;
157    }
158
159    /// Waits for the connection to be closed.
160    ///
161    /// Returns a "connection closed" error when the connection is closed. If
162    /// another error occurs before the connection is closed, that error is
163    /// returned instead.
164    ///
165    /// Use this method when you have an unbounded stream of data to forward to
166    /// the connection and the protocol does not require the client to
167    /// periodically acknowledge receipt. If you don't call this method to
168    /// periodically check if the connection has closed, you may not notice that
169    /// the client has gone away for an unboundedly long amount of time; usually
170    /// not until the stream of data produces its next message and you attempt
171    /// to write the data to the connection.
172    pub async fn wait_closed(&self) -> io::Error
173    where
174        A: AsyncReady + Send + Sync,
175    {
176        loop {
177            time::sleep(Duration::from_secs(1)).await;
178
179            match self.ready(Interest::READABLE | Interest::WRITABLE).await {
180                Ok(ready) if ready.is_read_closed() || ready.is_write_closed() => {
181                    return io::Error::new(io::ErrorKind::Other, "connection closed");
182                }
183                Ok(_) => (),
184                Err(err) => return err,
185            }
186        }
187    }
188
189    /// Returns the ID associated with this connection.
190    pub fn conn_id(&self) -> &ConnectionId {
191        &self.conn_id
192    }
193
194    /// Returns the peer address of the connection.
195    pub fn peer_addr(&self) -> &Option<IpAddr> {
196        &self.peer_addr
197    }
198}
199
200impl<A> FramedConn<A>
201where
202    A: AsyncRead + AsyncWrite + Unpin,
203{
204    pub fn inner(&self) -> &Conn<A> {
205        self.inner.get_ref().get_ref()
206    }
207}
208
209#[async_trait]
210impl<A> AsyncReady for FramedConn<A>
211where
212    A: AsyncRead + AsyncWrite + AsyncReady + Send + Sync + Unpin,
213{
214    async fn ready(&self, interest: Interest) -> io::Result<Ready> {
215        self.inner.get_ref().get_ref().ready(interest).await
216    }
217}
218
219pub struct Codec {
220    decode_state: DecodeState,
221    encode_state: Vec<(mz_pgrepr::Type, mz_pgwire_common::Format)>,
222    /// The session's text encoding settings when `encode_state` was installed.
223    text_settings: mz_pgrepr::TextEncodeSettings,
224    /// Largest frame the client may declare, raised once it authenticates.
225    ///
226    /// One `Codec` serves a connection for its whole life, but the frames it
227    /// should accept change partway through: before authentication only a
228    /// credential is legitimate, while afterwards a single `Bind` parameter can
229    /// carry bulk data. A ceiling wide enough for the second is far too wide for
230    /// the first, so the bound starts tight and
231    /// [`FramedConn::allow_post_auth_frames`] widens it.
232    max_frame_len: usize,
233}
234
235impl Codec {
236    /// Creates a new `Codec` for a client that has not yet authenticated.
237    pub fn new() -> Codec {
238        Codec {
239            decode_state: DecodeState::Head,
240            max_frame_len: MAX_PREAUTH_FRAME_SIZE,
241            encode_state: vec![],
242            text_settings: mz_pgrepr::TextEncodeSettings::STABLE,
243        }
244    }
245}
246
247impl Default for Codec {
248    fn default() -> Codec {
249        Codec::new()
250    }
251}
252
253impl Encoder<BackendMessage> for Codec {
254    type Error = io::Error;
255
256    /// Encode a backend message into `dst`.
257    /// If this function returns an error result, `dst` is left unmodified.
258    fn encode(&mut self, msg: BackendMessage, dst: &mut BytesMut) -> Result<(), io::Error> {
259        // Record the starting position so we can truncate on error.
260        // This prevents partial messages from being left in the buffer,
261        // which could be sent to the client and cause "lost synchronization" errors.
262        let start = dst.len();
263        match self.encode_inner(msg, dst) {
264            Ok(()) => Ok(()),
265            Err(e) => {
266                dst.truncate(start);
267                Err(e)
268            }
269        }
270    }
271}
272
273impl Codec {
274    /// This is the meat of the encoding logic. It's a separate function so that errors returned by
275    /// `?` can be handled in the outer `encode` function.
276    fn encode_inner(&self, msg: BackendMessage, dst: &mut BytesMut) -> Result<(), io::Error> {
277        // Write type byte.
278        let byte = match &msg {
279            BackendMessage::AuthenticationOk => b'R',
280            BackendMessage::AuthenticationCleartextPassword
281            | BackendMessage::AuthenticationSASL
282            | BackendMessage::AuthenticationSASLContinue(_)
283            | BackendMessage::AuthenticationSASLFinal(_) => b'R',
284            BackendMessage::RowDescription(_) => b'T',
285            BackendMessage::DataRow(_) => b'D',
286            BackendMessage::CommandComplete { .. } => b'C',
287            BackendMessage::EmptyQueryResponse => b'I',
288            BackendMessage::ReadyForQuery(_) => b'Z',
289            BackendMessage::NoData => b'n',
290            BackendMessage::ParameterStatus(_, _) => b'S',
291            BackendMessage::PortalSuspended => b's',
292            BackendMessage::BackendKeyData { .. } => b'K',
293            BackendMessage::ParameterDescription(_) => b't',
294            BackendMessage::ParseComplete => b'1',
295            BackendMessage::BindComplete => b'2',
296            BackendMessage::CloseComplete => b'3',
297            BackendMessage::ErrorResponse(r) => {
298                if r.severity.is_error() {
299                    b'E'
300                } else {
301                    b'N'
302                }
303            }
304            BackendMessage::CopyInResponse { .. } => b'G',
305            BackendMessage::CopyOutResponse { .. } => b'H',
306            BackendMessage::CopyData(_) => b'd',
307            BackendMessage::CopyDone => b'c',
308        };
309        dst.put_u8(byte);
310
311        // Write message length placeholder. The true length is filled in later.
312        let base = dst.len();
313        dst.put_u32(0);
314
315        // Write message contents.
316        match msg {
317            BackendMessage::CopyInResponse {
318                overall_format,
319                column_formats,
320            }
321            | BackendMessage::CopyOutResponse {
322                overall_format,
323                column_formats,
324            } => {
325                dst.put_format_i8(overall_format);
326                if column_formats.len() > usize::try_from(i16::MAX).expect("i16::MAX is positive") {
327                    return Err(io::Error::new(
328                        io::ErrorKind::InvalidData,
329                        format!(
330                            "{} columns in COPY response, which exceeds {}",
331                            column_formats.len(),
332                            i16::MAX
333                        ),
334                    ));
335                }
336                dst.put_length_i16(column_formats.len())?;
337                for format in column_formats {
338                    dst.put_format_i16(format);
339                }
340            }
341            BackendMessage::CopyData(data) => {
342                dst.put_slice(&data);
343            }
344            BackendMessage::CopyDone => (),
345            BackendMessage::AuthenticationOk => {
346                dst.put_u32(0);
347            }
348            BackendMessage::AuthenticationCleartextPassword => {
349                dst.put_u32(3);
350            }
351            BackendMessage::AuthenticationSASL => {
352                dst.put_u32(10);
353                dst.put_string("SCRAM-SHA-256");
354                dst.put_u8(b'\0');
355            }
356            BackendMessage::AuthenticationSASLContinue(data) => {
357                dst.put_u32(11);
358                let data = format!(
359                    "r={},s={},i={}",
360                    data.nonce, data.salt, data.iteration_count
361                );
362                dst.put_slice(data.as_bytes());
363            }
364            BackendMessage::AuthenticationSASLFinal(data) => {
365                dst.put_u32(12);
366                let res = match data.kind {
367                    SASLServerFinalMessageKinds::Verifier(verifier) => {
368                        format!("v={}", verifier)
369                    }
370                };
371                dst.put_slice(res.as_bytes());
372                if !data.extensions.is_empty() {
373                    dst.put_slice(b",");
374                    dst.put_slice(data.extensions.join(",").as_bytes());
375                }
376            }
377            BackendMessage::RowDescription(fields) => {
378                if fields.len() > usize::try_from(i16::MAX).expect("i16::MAX is positive") {
379                    return Err(io::Error::new(
380                        io::ErrorKind::InvalidData,
381                        format!(
382                            "{} fields in row description, which exceeds {}",
383                            fields.len(),
384                            i16::MAX
385                        ),
386                    ));
387                }
388                dst.put_length_i16(fields.len())?;
389                for f in &fields {
390                    dst.put_string(&f.name.to_string());
391                    dst.put_u32(f.table_id);
392                    dst.put_u16(f.column_id);
393                    dst.put_u32(f.type_oid);
394                    dst.put_i16(f.type_len);
395                    dst.put_i32(f.type_mod);
396                    // TODO: make the format correct
397                    dst.put_format_i16(f.format);
398                }
399            }
400            BackendMessage::DataRow(fields) => {
401                if fields.len() > usize::try_from(i16::MAX).expect("i16::MAX is positive") {
402                    return Err(io::Error::new(
403                        io::ErrorKind::InvalidData,
404                        format!(
405                            "{} fields in data row, which exceeds {}",
406                            fields.len(),
407                            i16::MAX
408                        ),
409                    ));
410                }
411                dst.put_length_i16(fields.len())?;
412                for (f, (ty, format)) in fields.iter().zip_eq(&self.encode_state) {
413                    if let Some(f) = f {
414                        let base = dst.len();
415                        dst.put_u32(0);
416                        f.encode(ty, *format, dst, self.text_settings)?;
417                        let len = dst.len() - base - 4;
418                        let len = i32::try_from(len).map_err(|_| {
419                            io::Error::new(
420                                io::ErrorKind::InvalidData,
421                                "length of encoded data row field does not fit into an i32",
422                            )
423                        })?;
424                        dst[base..base + 4].copy_from_slice(&len.to_be_bytes());
425                    } else {
426                        dst.put_i32(-1);
427                    }
428                }
429            }
430            BackendMessage::CommandComplete { tag } => {
431                dst.put_string(&tag);
432            }
433            BackendMessage::ParseComplete => (),
434            BackendMessage::BindComplete => (),
435            BackendMessage::CloseComplete => (),
436            BackendMessage::EmptyQueryResponse => (),
437            BackendMessage::ReadyForQuery(status) => {
438                dst.put_u8(status.into());
439            }
440            BackendMessage::ParameterStatus(name, value) => {
441                dst.put_string(name);
442                dst.put_string(&value);
443            }
444            BackendMessage::PortalSuspended => (),
445            BackendMessage::NoData => (),
446            BackendMessage::BackendKeyData {
447                conn_id,
448                secret_key,
449            } => {
450                dst.put_u32(conn_id);
451                dst.put_u32(secret_key);
452            }
453            BackendMessage::ParameterDescription(params) => {
454                if params.len() > usize::from(u16::MAX) {
455                    return Err(io::Error::new(
456                        io::ErrorKind::InvalidData,
457                        format!(
458                            "{} params in parameter description, which exceeds {}",
459                            params.len(),
460                            u16::MAX
461                        ),
462                    ));
463                }
464                dst.put_length_u16(params.len())?;
465                for param in params {
466                    dst.put_u32(param.oid());
467                }
468            }
469            BackendMessage::ErrorResponse(ErrorResponse {
470                severity,
471                code,
472                message,
473                detail,
474                hint,
475                position,
476            }) => {
477                dst.put_u8(b'S');
478                dst.put_string(severity.as_str());
479                dst.put_u8(b'C');
480                dst.put_string(code.code());
481                dst.put_u8(b'M');
482                dst.put_string(&message);
483                if let Some(detail) = &detail {
484                    dst.put_u8(b'D');
485                    dst.put_string(detail);
486                }
487                if let Some(hint) = &hint {
488                    dst.put_u8(b'H');
489                    dst.put_string(hint);
490                }
491                if let Some(position) = &position {
492                    dst.put_u8(b'P');
493                    dst.put_string(&position.to_string());
494                }
495                dst.put_u8(b'\0');
496            }
497        }
498
499        let len = dst.len() - base;
500
501        // Overwrite length placeholder with true length.
502        let len = i32::try_from(len).map_err(|_| {
503            io::Error::new(
504                io::ErrorKind::InvalidData,
505                "length of encoded message does not fit into an i32",
506            )
507        })?;
508        dst[base..base + 4].copy_from_slice(&len.to_be_bytes());
509
510        Ok(())
511    }
512}
513
514impl Decoder for Codec {
515    type Item = FrontendMessage;
516    type Error = io::Error;
517
518    fn decode(&mut self, src: &mut BytesMut) -> Result<Option<FrontendMessage>, io::Error> {
519        loop {
520            match self.decode_state {
521                DecodeState::Head => {
522                    if src.len() < 5 {
523                        return Ok(None);
524                    }
525                    let msg_type = src[0];
526                    let frame_len = parse_frame_len(&src[1..], self.max_frame_len)?;
527                    src.advance(5);
528                    src.reserve(frame_len);
529                    self.decode_state = DecodeState::Data(msg_type, frame_len);
530                }
531
532                DecodeState::Data(msg_type, frame_len) => {
533                    if src.len() < frame_len {
534                        return Ok(None);
535                    }
536                    let buf = src.split_to(frame_len).freeze();
537                    let buf = Cursor::new(&buf);
538                    let msg = match msg_type {
539                        // Simple query flow.
540                        b'Q' => decode_query(buf)?,
541
542                        // Extended query flow.
543                        b'P' => decode_parse(buf)?,
544                        b'D' => decode_describe(buf)?,
545                        b'B' => decode_bind(buf)?,
546                        b'E' => decode_execute(buf)?,
547                        b'H' => decode_flush(buf)?,
548                        b'S' => decode_sync(buf)?,
549                        b'C' => decode_close(buf)?,
550
551                        // Termination.
552                        b'X' => decode_terminate(buf)?,
553
554                        // Authentication.
555                        b'p' => decode_auth(buf)?,
556
557                        // Copy from flow.
558                        b'f' => decode_copy_fail(buf)?,
559                        b'd' => decode_copy_data(buf, frame_len)?,
560                        b'c' => decode_copy_done(buf)?,
561
562                        // Invalid.
563                        _ => {
564                            return Err(io::Error::new(
565                                io::ErrorKind::InvalidData,
566                                format!("unknown message type {}", msg_type),
567                            ));
568                        }
569                    };
570                    src.reserve(5);
571                    self.decode_state = DecodeState::Head;
572                    return Ok(Some(msg));
573                }
574            }
575        }
576    }
577}
578
579fn decode_terminate(mut _buf: Cursor) -> Result<FrontendMessage, io::Error> {
580    // Nothing more to decode.
581    Ok(FrontendMessage::Terminate)
582}
583
584fn decode_auth(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
585    let mut value = Vec::new();
586    while let Ok(b) = buf.read_byte() {
587        value.push(b);
588    }
589    Ok(FrontendMessage::RawAuthentication(value))
590}
591
592fn expect(buf: &mut Cursor, expected: &[u8]) -> Result<(), io::Error> {
593    for i in 0..expected.len() {
594        if buf.read_byte()? != expected[i] {
595            return Err(input_err(format!(
596                "Invalid SASL initial response: expected '{}'",
597                std::str::from_utf8(expected).unwrap_or("invalid UTF-8")
598            )));
599        }
600    }
601    Ok(())
602}
603
604fn read_until_comma(buf: &mut Cursor) -> Result<Vec<u8>, io::Error> {
605    let mut v = Vec::new();
606    while let Ok(b) = buf.peek_byte() {
607        if b == b',' {
608            break;
609        }
610        v.push(buf.read_byte()?);
611    }
612    Ok(v)
613}
614
615// All SASL parsing is based on RFC 5802, [section 7](https://datatracker.ietf.org/doc/html/rfc5802#section-7)
616
617//   extensions = attr-val *("," attr-val)
618//                     ;; All extensions are optional,
619//                     ;; i.e., unrecognized attributes
620//                     ;; not defined in this document
621//                     ;; MUST be ignored.
622//   reserved-mext  = "m=" 1*(value-char)
623//                     ;; Reserved for signaling mandatory extensions.
624//                     ;; The exact syntax will be defined in
625//                     ;; the future.
626//   gs2-cbind-flag  = ("p=" cb-name) / "n" / "y"
627//                     ;; "n" -> client doesn't support channel binding.
628//                     ;; "y" -> client does support channel binding
629//                     ;;        but thinks the server does not.
630//                     ;; "p" -> client requires channel binding.
631//                     ;; The selected channel binding follows "p=".
632//
633//   gs2-header      = gs2-cbind-flag "," [ authzid ] ","
634//                     ;; GS2 header for SCRAM
635//                     ;; (the actual GS2 header includes an optional
636//                     ;; flag to indicate that the GSS mechanism is not
637//                     ;; "standard", but since SCRAM is "standard", we
638//                     ;; don't include that flag).
639//   client-first-message-bare =
640//                     [reserved-mext ","]
641//                     username "," nonce ["," extensions]
642//
643//   client-first-message =
644//                     gs2-header client-first-message-bare
645pub fn decode_sasl_client_first_message(mut buf: Cursor) -> Result<SASLInitialResponse, io::Error> {
646    // 1) GS2 cbind flag
647    let cbind_flag = match buf.read_byte()? {
648        b'n' => ChannelBinding::None,
649        b'y' => ChannelBinding::ClientSupported,
650        b'p' => {
651            // must be "p=" then cbname up to next comma
652            expect(&mut buf, b"=")?;
653            let cbname = String::from_utf8(read_until_comma(&mut buf)?)
654                .map_err(|_| input_err("invalid cbname utf8"))?;
655            ChannelBinding::Required(cbname)
656        }
657        other => {
658            return Err(input_err(format!(
659                "Invalid channel binding flag: {}",
660                other
661            )));
662        }
663    };
664    expect(&mut buf, b",")?;
665
666    // 2) Optional authzid: either empty, or "a=" up to next comma
667    let mut authzid = None;
668    if buf.peek_byte()? == b'a' {
669        expect(&mut buf, b"a=")?;
670        let a = String::from_utf8(read_until_comma(&mut buf)?)
671            .map_err(|_| input_err("invalid authzid utf8"))?;
672        authzid = Some(a);
673    }
674    expect(&mut buf, b",")?;
675
676    let mut client_first_message_bare_raw = String::new();
677
678    // 3) Optional reserved "m=" extension before n=
679    let mut reserved_mext = None;
680    if buf.peek_byte()? == b'm' {
681        expect(&mut buf, b"m=")?;
682        let mext_val = String::from_utf8(read_until_comma(&mut buf)?)
683            .map_err(|_| input_err("invalid m ext utf8"))?;
684        client_first_message_bare_raw.push_str(&format!("m={},", mext_val));
685        reserved_mext = Some(mext_val);
686        expect(&mut buf, b",")?;
687    }
688
689    // 4) Username: must be "n=" then saslname
690    expect(&mut buf, b"n=")?;
691    // Postgres doesn't use the username here, so we just consume
692    let username = String::from_utf8(read_until_comma(&mut buf)?)
693        .map_err(|_| input_err("invalid username utf8"))?;
694    expect(&mut buf, b",")?;
695    client_first_message_bare_raw.push_str(&format!("n={},", username));
696
697    // 5) Nonce: must be "r=" then value up to next comma or end
698    expect(&mut buf, b"r=")?;
699    let nonce = String::from_utf8(read_until_comma(&mut buf)?)
700        .map_err(|_| input_err("invalid nonce utf8"))?;
701    client_first_message_bare_raw.push_str(&format!("r={}", nonce));
702
703    // 6) Optional extensions: "," key=value chunks
704    let mut extensions = Vec::new();
705    while let Ok(b',') = buf.peek_byte().map(|b| b) {
706        expect(&mut buf, b",")?;
707        let ext = String::from_utf8(read_until_comma(&mut buf)?)
708            .map_err(|_| input_err("invalid ext utf8"))?;
709        if !ext.is_empty() {
710            client_first_message_bare_raw.push_str(&format!(",{}", ext));
711            extensions.push(ext);
712        }
713    }
714
715    Ok(SASLInitialResponse {
716        gs2_header: GS2Header {
717            cbind_flag,
718            authzid,
719        },
720        nonce,
721        extensions,
722        reserved_mext,
723        client_first_message_bare_raw,
724    })
725}
726
727pub fn decode_sasl_initial_response(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
728    let mechanism = buf.read_cstr()?;
729    let initial_resp_len = buf.read_i32()?;
730    if initial_resp_len < 0 {
731        // -1 means no response? We bail here
732        return Err(input_err("No initial response"));
733    }
734
735    let initial_response = decode_sasl_client_first_message(buf)?;
736    Ok(FrontendMessage::SASLInitialResponse {
737        gs2_header: initial_response.gs2_header.clone(),
738        mechanism: mechanism.to_owned(),
739        initial_response,
740    })
741}
742
743//   proof           = "p=" base64
744//
745//   channel-binding = "c=" base64
746//                     ;; base64 encoding of cbind-input.
747//   client-final-message-without-proof =
748//                     channel-binding "," nonce [","
749//                     extensions]
750//
751//   client-final-message =
752//                     client-final-message-without-proof "," proof
753pub fn decode_sasl_response(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
754    // --- client-final-message-without-proof ---
755    let mut client_final_message_bare_raw = String::new();
756    // channel-binding: "c=" <base64>, up to the next comma
757    expect(&mut buf, b"c=")?;
758    let channel_binding = String::from_utf8(read_until_comma(&mut buf)?)
759        .map_err(|_| input_err("invalid channel-binding utf8"))?;
760    expect(&mut buf, b",")?;
761    client_final_message_bare_raw.push_str(&format!("c={},", channel_binding));
762
763    // nonce: "r=" <printable>, up to the next comma
764    expect(&mut buf, b"r=")?;
765    let nonce = String::from_utf8(read_until_comma(&mut buf)?)
766        .map_err(|_| input_err("invalid nonce utf8"))?;
767    client_final_message_bare_raw.push_str(&format!("r={}", nonce));
768
769    // after reading channel-binding and nonce
770    let mut extensions = Vec::new();
771
772    // Keep reading ",<token>" until we see ",p="
773    while buf.peek_byte()? == b',' {
774        expect(&mut buf, b",")?;
775        if buf.peek_byte()? == b'p' {
776            break;
777        }
778        let ext = String::from_utf8(read_until_comma(&mut buf)?)
779            .map_err(|_| input_err("invalid extension utf8"))?;
780        if !ext.is_empty() {
781            client_final_message_bare_raw.push_str(&format!(",{}", ext));
782            extensions.push(ext);
783        }
784    }
785
786    // Proof is mandatory and last
787    expect(&mut buf, b"p=")?;
788    let proof = String::from_utf8(read_until_comma(&mut buf)?)
789        .map_err(|_| input_err("invalid proof utf8"))?;
790
791    Ok(FrontendMessage::SASLResponse(SASLClientFinalResponse {
792        channel_binding,
793        nonce,
794        extensions,
795        proof,
796        client_final_message_bare_raw,
797    }))
798}
799
800pub fn decode_password(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
801    Ok(FrontendMessage::Password {
802        password: buf.read_cstr()?.to_owned(),
803    })
804}
805
806fn decode_query(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
807    Ok(FrontendMessage::Query {
808        sql: buf.read_cstr()?.to_string(),
809    })
810}
811
812fn decode_parse(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
813    let name = buf.read_cstr()?;
814    let sql = buf.read_cstr()?;
815
816    // NOTE: the protocol calls the counts that precede repeated groups `Int16`,
817    // but PostgreSQL decodes them as unsigned, so they reach 65535. Reading one
818    // as signed makes it negative, which drops the group and leaves the cursor
819    // misaligned for the rest of the message.
820    let mut param_types = vec![];
821    for _ in 0..buf.read_u16()? {
822        param_types.push(buf.read_u32()?);
823    }
824
825    Ok(FrontendMessage::Parse {
826        name: name.into(),
827        sql: sql.into(),
828        param_types,
829    })
830}
831
832fn decode_close(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
833    match buf.read_byte()? {
834        b'S' => Ok(FrontendMessage::CloseStatement {
835            name: buf.read_cstr()?.to_owned(),
836        }),
837        b'P' => Ok(FrontendMessage::ClosePortal {
838            name: buf.read_cstr()?.to_owned(),
839        }),
840        b => Err(input_err(format!(
841            "invalid type byte in close message: {}",
842            b
843        ))),
844    }
845}
846
847fn decode_describe(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
848    let first_char = buf.read_byte()?;
849    let name = buf.read_cstr()?.to_string();
850    match first_char {
851        b'S' => Ok(FrontendMessage::DescribeStatement { name }),
852        b'P' => Ok(FrontendMessage::DescribePortal { name }),
853        other => Err(input_err(format!("Invalid describe type: {:#x?}", other))),
854    }
855}
856
857fn decode_bind(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
858    let portal_name = buf.read_cstr()?.to_string();
859    let statement_name = buf.read_cstr()?.to_string();
860
861    // The three counts below are `Int16` in the protocol but decoded as
862    // unsigned, for the reason spelled out in `decode_parse`.
863    let mut param_formats = Vec::new();
864    for _ in 0..buf.read_u16()? {
865        param_formats.push(buf.read_format()?);
866    }
867
868    let mut raw_params = Vec::new();
869    for _ in 0..buf.read_u16()? {
870        let len = buf.read_i32()?;
871        if len == -1 {
872            raw_params.push(None); // NULL
873        } else {
874            // TODO(benesch): this should use bytes::Bytes to avoid the copy.
875            let mut value = Vec::new();
876            for _ in 0..len {
877                value.push(buf.read_byte()?);
878            }
879            raw_params.push(Some(value));
880        }
881    }
882
883    let mut result_formats = Vec::new();
884    for _ in 0..buf.read_u16()? {
885        result_formats.push(buf.read_format()?);
886    }
887
888    Ok(FrontendMessage::Bind {
889        portal_name,
890        statement_name,
891        param_formats,
892        raw_params,
893        result_formats,
894    })
895}
896
897fn decode_execute(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
898    let portal_name = buf.read_cstr()?.to_string();
899    let max_rows = buf.read_i32()?;
900    Ok(FrontendMessage::Execute {
901        portal_name,
902        max_rows,
903    })
904}
905
906fn decode_flush(mut _buf: Cursor) -> Result<FrontendMessage, io::Error> {
907    // Nothing more to decode.
908    Ok(FrontendMessage::Flush)
909}
910
911fn decode_sync(mut _buf: Cursor) -> Result<FrontendMessage, io::Error> {
912    // Nothing more to decode.
913    Ok(FrontendMessage::Sync)
914}
915
916fn decode_copy_data(mut buf: Cursor, frame_len: usize) -> Result<FrontendMessage, io::Error> {
917    let mut data = Vec::with_capacity(frame_len);
918    for _ in 0..frame_len {
919        data.push(buf.read_byte()?);
920    }
921    Ok(FrontendMessage::CopyData(data))
922}
923
924fn decode_copy_done(mut _buf: Cursor) -> Result<FrontendMessage, io::Error> {
925    // Nothing more to decode.
926    Ok(FrontendMessage::CopyDone)
927}
928
929fn decode_copy_fail(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
930    Ok(FrontendMessage::CopyFail(buf.read_cstr()?.to_string()))
931}