1use 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
34pub const MAX_REQUEST_SIZE: usize = u64_to_usize(2 * bytesize::MB);
36
37pub const MAX_STARTUP_FRAME_SIZE: usize = 10_000;
44
45pub const FORWARDED_STARTUP_PARAM_ALLOWANCE: usize = 512;
65
66pub const MAX_FORWARDED_STARTUP_FRAME_SIZE: usize =
69 MAX_STARTUP_FRAME_SIZE + FORWARDED_STARTUP_PARAM_ALLOWANCE;
70
71pub 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 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
137pub 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 4 => (),
156 0 => return Ok(None),
159 _ => 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 pub fn encode(&self, dst: &mut BytesMut) -> Result<(), io::Error> {
194 let base = dst.len();
196 dst.put_u32(0);
197
198 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 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 pub fn encode(&self, dst: &mut BytesMut) -> Result<(), io::Error> {
238 let byte = match self {
240 FrontendMessage::Password { .. } => b'p',
241 _ => panic!("unsupported"),
242 };
243 dst.put_u8(byte);
244
245 let base = dst.len();
247 dst.put_u32(0);
248
249 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 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
278pub 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 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#[derive(Debug)]
318pub struct Cursor<'a> {
319 buf: &'a [u8],
320}
321
322impl<'a> Cursor<'a> {
323 pub fn new(buf: &'a [u8]) -> Cursor<'a> {
326 Cursor { buf }
327 }
328
329 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 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 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 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 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 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 pub fn read_format(&mut self) -> Result<Format, io::Error> {
402 Format::try_from(self.read_u16()?)
403 }
404
405 pub fn advance(&mut self, n: usize) {
407 self.buf = &self.buf[n..]
408 }
409}
410
411pub 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 #[derive(Debug)]
435 struct Offer {
436 len: usize,
437 }
438
439 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 result: Option<Result<Option<FrontendStartupMessage>, io::Error>>,
475 }
476
477 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 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 fn params_sized_to(frame_len: usize) -> BTreeMap<String, String> {
500 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 #[mz_ore::test(tokio::test(start_paused = true))]
564 async fn test_forwarded_startup_params_fit_the_allowance() {
565 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 #[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}