Skip to main content

mz_balancerd/
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
10use async_trait::async_trait;
11use bytes::{Buf, BufMut, BytesMut};
12use bytesize::ByteSize;
13use futures::{SinkExt, TryStreamExt, sink};
14use mz_ore::cast::CastFrom;
15use mz_ore::future::OreSinkExt;
16use mz_ore::netio::AsyncReady;
17use mz_pgwire_common::{
18    Conn, Cursor, DecodeState, ErrorResponse, FrontendMessage, MAX_PREAUTH_FRAME_SIZE,
19    MAX_REQUEST_SIZE, Pgbuf, parse_frame_len,
20};
21use tokio::io::{self, AsyncRead, AsyncWrite, Interest, Ready};
22use tokio_util::codec::{Decoder, Encoder, Framed};
23
24/// Internal representation of a backend [message].
25///
26/// [message]: https://www.postgresql.org/docs/11/protocol-message-formats.html
27#[derive(Debug)]
28pub enum BackendMessage {
29    AuthenticationCleartextPassword,
30    ErrorResponse(ErrorResponse),
31}
32
33impl From<ErrorResponse> for BackendMessage {
34    fn from(err: ErrorResponse) -> BackendMessage {
35        BackendMessage::ErrorResponse(err)
36    }
37}
38
39/// A connection that manages the encoding and decoding of pgwire frames.
40///
41/// This decodes at most one frame per connection, the client's credential, and
42/// is bounded by [`MAX_PREAUTH_FRAME_SIZE`] throughout. Once the destination is
43/// resolved the connection is spliced and the remaining bytes are proxied
44/// without being framed.
45pub struct FramedConn<A> {
46    inner: sink::Buffer<Framed<Conn<A>, Codec>, BackendMessage>,
47}
48
49impl<A> FramedConn<A>
50where
51    A: AsyncRead + AsyncWrite + Unpin,
52{
53    /// Constructs a new framed connection.
54    ///
55    /// The underlying connection, `inner`, is expected to be something like a
56    /// TCP stream. Anything that implements [`AsyncRead`] and [`AsyncWrite`]
57    /// will do.
58    pub fn new(inner: Conn<A>) -> FramedConn<A> {
59        FramedConn {
60            inner: Framed::new(inner, Codec::new()).buffer(32),
61        }
62    }
63
64    /// Reads and decodes one frontend message from the client.
65    ///
66    /// Blocks until the client sends a complete message. If the client
67    /// terminates the stream, returns `None`. Returns an error if the client
68    /// sends a malformed message or if the connection underlying is broken.
69    ///
70    /// # Cancel safety
71    ///
72    /// This method is cancel safe. The returned future only holds onto a
73    /// reference to thea underlying stream, so dropping it will never lose a
74    /// value.
75    ///
76    /// <https://docs.rs/tokio-stream/latest/tokio_stream/trait.StreamExt.html#cancel-safety-1>
77    pub async fn recv(&mut self) -> Result<Option<FrontendMessage>, io::Error> {
78        let message = self.inner.try_next().await?;
79        Ok(message)
80    }
81
82    /// Encodes and sends one backend message to the client.
83    ///
84    /// Note that the connection is not flushed after calling this method. You
85    /// must call [`FramedConn::flush`] explicitly. Returns an error if the
86    /// underlying connection is broken.
87    ///
88    /// Please use `StateMachine::send` instead if calling from `StateMachine`,
89    /// as it applies session-based filters before calling this method.
90    pub async fn send<M>(&mut self, message: M) -> Result<(), io::Error>
91    where
92        M: Into<BackendMessage>,
93    {
94        let message = message.into();
95        self.inner.enqueue(message).await
96    }
97
98    /// Flushes all outstanding messages.
99    pub async fn flush(&mut self) -> Result<(), io::Error> {
100        self.inner.flush().await
101    }
102}
103
104impl<A> FramedConn<A>
105where
106    A: AsyncRead + AsyncWrite + Unpin,
107{
108    pub fn inner(&self) -> &Conn<A> {
109        self.inner.get_ref().get_ref()
110    }
111    pub fn inner_mut(&mut self) -> &mut Conn<A> {
112        self.inner.get_mut().get_mut()
113    }
114}
115
116#[async_trait]
117impl<A> AsyncReady for FramedConn<A>
118where
119    A: AsyncRead + AsyncWrite + AsyncReady + Send + Sync + Unpin,
120{
121    async fn ready(&self, interest: Interest) -> io::Result<Ready> {
122        self.inner.get_ref().get_ref().ready(interest).await
123    }
124}
125
126struct Codec {
127    decode_state: DecodeState,
128}
129
130impl Codec {
131    /// Creates a new `Codec`.
132    pub fn new() -> Codec {
133        Codec {
134            decode_state: DecodeState::Head,
135        }
136    }
137}
138
139impl Default for Codec {
140    fn default() -> Codec {
141        Codec::new()
142    }
143}
144
145impl Encoder<BackendMessage> for Codec {
146    type Error = io::Error;
147
148    /// Encode a backend message into `dst`.
149    /// If this function returns an error result, `dst` is left unmodified.
150    fn encode(&mut self, msg: BackendMessage, dst: &mut BytesMut) -> Result<(), io::Error> {
151        // Record the starting position so we can truncate on error.
152        // This prevents partial messages from being left in the buffer,
153        // which could be sent to the client and cause "lost synchronization" errors.
154        let start = dst.len();
155        match self.encode_inner(msg, dst) {
156            Ok(()) => Ok(()),
157            Err(e) => {
158                dst.truncate(start);
159                Err(e)
160            }
161        }
162    }
163}
164
165impl Codec {
166    /// This is the meat of the encoding logic. It's a separate function so that errors returned by
167    /// `?` can be handled in the outer `encode` function.
168    fn encode_inner(&self, msg: BackendMessage, dst: &mut BytesMut) -> Result<(), io::Error> {
169        // Write type byte.
170        let byte = match &msg {
171            BackendMessage::AuthenticationCleartextPassword => b'R',
172            BackendMessage::ErrorResponse(r) => {
173                if r.severity.is_error() {
174                    b'E'
175                } else {
176                    b'N'
177                }
178            }
179        };
180        dst.put_u8(byte);
181
182        // Write message length placeholder. The true length is filled in later.
183        let base = dst.len();
184        dst.put_u32(0);
185
186        // Write message contents.
187        match msg {
188            BackendMessage::AuthenticationCleartextPassword => {
189                dst.put_u32(3);
190            }
191            BackendMessage::ErrorResponse(ErrorResponse {
192                severity,
193                code,
194                message,
195                detail,
196                hint,
197                position,
198            }) => {
199                dst.put_u8(b'S');
200                dst.put_string(severity.as_str());
201                dst.put_u8(b'C');
202                dst.put_string(code.code());
203                dst.put_u8(b'M');
204                dst.put_string(&message);
205                if let Some(detail) = &detail {
206                    dst.put_u8(b'D');
207                    dst.put_string(detail);
208                }
209                if let Some(hint) = &hint {
210                    dst.put_u8(b'H');
211                    dst.put_string(hint);
212                }
213                if let Some(position) = &position {
214                    dst.put_u8(b'P');
215                    dst.put_string(&position.to_string());
216                }
217                dst.put_u8(b'\0');
218            }
219        }
220
221        let len = dst.len() - base;
222
223        // Overwrite length placeholder with true length.
224        let len = i32::try_from(len).map_err(|_| {
225            io::Error::new(
226                io::ErrorKind::InvalidData,
227                "length of encoded message does not fit into an i32",
228            )
229        })?;
230        dst[base..base + 4].copy_from_slice(&len.to_be_bytes());
231
232        Ok(())
233    }
234}
235
236impl Decoder for Codec {
237    type Item = FrontendMessage;
238    type Error = io::Error;
239
240    fn decode(&mut self, src: &mut BytesMut) -> Result<Option<FrontendMessage>, io::Error> {
241        if src.len() > MAX_REQUEST_SIZE {
242            return Err(io::Error::new(
243                io::ErrorKind::InvalidData,
244                format!(
245                    "request larger than {}",
246                    ByteSize::b(u64::cast_from(MAX_REQUEST_SIZE))
247                ),
248            ));
249        }
250        loop {
251            match self.decode_state {
252                DecodeState::Head => {
253                    if src.len() < 5 {
254                        return Ok(None);
255                    }
256                    let msg_type = src[0];
257                    let frame_len = parse_frame_len(&src[1..], MAX_PREAUTH_FRAME_SIZE)?;
258                    src.advance(5);
259                    src.reserve(frame_len);
260                    self.decode_state = DecodeState::Data(msg_type, frame_len);
261                }
262
263                DecodeState::Data(msg_type, frame_len) => {
264                    if src.len() < frame_len {
265                        return Ok(None);
266                    }
267                    let buf = src.split_to(frame_len).freeze();
268                    let buf = Cursor::new(&buf);
269                    let msg = match msg_type {
270                        // Termination.
271                        b'X' => decode_terminate(buf)?,
272
273                        // Authentication.
274                        b'p' => decode_password(buf)?,
275
276                        // Invalid.
277                        _ => {
278                            return Err(io::Error::new(
279                                io::ErrorKind::InvalidData,
280                                format!("unknown message type {}", msg_type),
281                            ));
282                        }
283                    };
284                    src.reserve(5);
285                    self.decode_state = DecodeState::Head;
286                    return Ok(Some(msg));
287                }
288            }
289        }
290    }
291}
292
293fn decode_terminate(mut _buf: Cursor) -> Result<FrontendMessage, io::Error> {
294    // Nothing more to decode.
295    Ok(FrontendMessage::Terminate)
296}
297
298fn decode_password(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
299    Ok(FrontendMessage::Password {
300        password: buf.read_cstr()?.to_owned(),
301    })
302}