1use std::{collections::BTreeMap, future::Future, io, sync::Arc};
7
8use bytes::{Buf, Bytes, BytesMut};
9use rustls::{
10 ClientConfig, ServerConfig,
11 pki_types::{CertificateDer, ServerName},
12};
13use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
14use tokio_util::codec::{Decoder, Encoder};
15
16use crate::{
17 Conn,
18 auth::TlsServerEndPoint,
19 codec::{Backend, BackendMessage, Direction, Frame, Frontend, FrontendMessage, PgCodec},
20 demux::{
21 CancelKey, Demux, Notification, OrderedAsyncEvent, ParameterStatus, SessionItem,
22 TaggedNotice,
23 },
24 middleware::{
25 AcceptsMessage, ClientRole, Inbound, MessageMiddleware, Middleware, PhaseAssociation,
26 ReceiveError, ReconstructableMessage, ServerRole, TypedMiddleware, TypedReceiveError,
27 },
28 pre_startup::{
29 AwaitingSslReply, DEFAULT_MAX_PRE_STARTUP_PACKET_LEN, EncryptionReply, Negotiation,
30 PreStartup, PreStartupMessage, ServerSslDecision, SslMode, SslModeNegotiation,
31 TlsHandshake, decode_pre_startup_with_limit, gssenc_request_packet, ssl_request_packet,
32 },
33 tls::{ClientTls, ServerTls},
34};
35
36async fn receive_typed<Role, Phase, Wire, State, Handler, Read>(
37 read: Read,
38 middleware: &mut Middleware<State, Handler>,
39) -> Result<
40 <Phase as PhaseAssociation<Inbound, Role, Wire>>::Message,
41 TypedReceiveError<Handler::Error, Wire>,
42>
43where
44 Read: Future<Output = io::Result<Wire>>,
45 Phase: PhaseAssociation<Inbound, Role, Wire>,
46 Wire: ReconstructableMessage,
47 Handler: TypedMiddleware<
48 Role,
49 <Phase as PhaseAssociation<Inbound, Role, Wire>>::ProtocolPhase,
50 <Phase as PhaseAssociation<Inbound, Role, Wire>>::Message,
51 State,
52 >,
53{
54 let message = read.await.map_err(TypedReceiveError::Io)?;
55 let message = <Phase as PhaseAssociation<Inbound, Role, Wire>>::Message::try_from(message)
56 .map_err(TypedReceiveError::Illegal)?;
57 let message = middleware
58 .intercept_typed::<
59 Role,
60 <Phase as PhaseAssociation<Inbound, Role, Wire>>::ProtocolPhase,
61 _,
62 >(message)
63 .await
64 .map_err(TypedReceiveError::Middleware)?;
65 if message.as_ref().is_reconstructable() {
66 Ok(message)
67 } else {
68 Err(TypedReceiveError::InvalidWire(message.into()))
69 }
70}
71
72async fn receive_checked<Message, State, Handler, ProtocolState, Read>(
73 read: Read,
74 middleware: &mut Middleware<State, Handler>,
75 protocol_state: &ProtocolState,
76) -> Result<Message, ReceiveError<Handler::Error, Message>>
77where
78 Read: Future<Output = io::Result<Message>>,
79 Message: ReconstructableMessage,
80 Handler: MessageMiddleware<Message, State>,
81 ProtocolState: AcceptsMessage<Message>,
82{
83 let message = read.await.map_err(ReceiveError::Io)?;
84 middleware
85 .intercept_checked(protocol_state, message)
86 .await
87 .map_err(ReceiveError::Intercept)
88}
89
90#[derive(Debug)]
92pub struct Buffered<S, D = Backend> {
93 io: S,
94 outbound: BytesMut,
95 inbound: BytesMut,
96 inbound_codec: PgCodec<D>,
97 max_pre_startup_packet_len: usize,
98 demux: Demux,
99}
100
101impl<S> Buffered<S, Backend> {
102 pub fn new(io: S) -> Self {
104 Self {
105 io,
106 outbound: BytesMut::new(),
107 inbound: BytesMut::new(),
108 inbound_codec: PgCodec::default(),
109 max_pre_startup_packet_len: DEFAULT_MAX_PRE_STARTUP_PACKET_LEN,
110 demux: Demux::default(),
111 }
112 }
113
114 pub fn with_max_frame_len(io: S, max_frame_len: usize) -> io::Result<Self> {
120 Ok(Self {
121 io,
122 outbound: BytesMut::new(),
123 inbound: BytesMut::new(),
124 inbound_codec: PgCodec::with_max_frame_len(max_frame_len)?,
125 max_pre_startup_packet_len: DEFAULT_MAX_PRE_STARTUP_PACKET_LEN,
126 demux: Demux::default(),
127 })
128 }
129}
130
131impl<S> Buffered<S, Frontend> {
132 pub fn new_frontend(io: S) -> Self {
134 Self {
135 io,
136 outbound: BytesMut::new(),
137 inbound: BytesMut::new(),
138 inbound_codec: PgCodec::default(),
139 max_pre_startup_packet_len: DEFAULT_MAX_PRE_STARTUP_PACKET_LEN,
140 demux: Demux::default(),
141 }
142 }
143
144 pub fn with_max_frame_len_frontend(io: S, max_frame_len: usize) -> io::Result<Self> {
150 Self::with_limits_frontend(io, max_frame_len, DEFAULT_MAX_PRE_STARTUP_PACKET_LEN)
151 }
152
153 pub fn with_limits_frontend(
159 io: S,
160 max_frame_len: usize,
161 max_pre_startup_packet_len: usize,
162 ) -> io::Result<Self> {
163 if !(8..=i32::MAX as usize).contains(&max_pre_startup_packet_len) {
164 return Err(io::Error::new(
165 io::ErrorKind::InvalidInput,
166 "pre-startup packet limit must be between 8 and i32::MAX bytes",
167 ));
168 }
169 Ok(Self {
170 io,
171 outbound: BytesMut::new(),
172 inbound: BytesMut::new(),
173 inbound_codec: PgCodec::with_max_frame_len(max_frame_len)?,
174 max_pre_startup_packet_len,
175 demux: Demux::default(),
176 })
177 }
178}
179
180impl<S, D> Buffered<S, D> {
181 pub fn push(&mut self, frame: Frame) -> io::Result<()> {
187 self.inbound_codec.encode(frame, &mut self.outbound)
188 }
189
190 #[must_use]
191 pub fn pending(&self) -> &[u8] {
193 &self.outbound
194 }
195
196 pub fn into_inner(self) -> S {
198 self.io
199 }
200
201 pub const fn get_ref(&self) -> &S {
203 &self.io
204 }
205
206 pub const fn get_mut(&mut self) -> &mut S {
208 &mut self.io
209 }
210
211 fn push_raw(&mut self, bytes: &[u8]) {
212 self.outbound.extend_from_slice(bytes);
213 }
214
215 #[must_use]
216 pub const fn demux(&self) -> &Demux {
218 &self.demux
219 }
220
221 pub const fn demux_mut(&mut self) -> &mut Demux {
223 &mut self.demux
224 }
225}
226
227impl<S, D> Buffered<S, D>
228where
229 S: AsyncRead + AsyncWrite + Unpin,
230{
231 async fn connect_tls(
232 self,
233 server_name: ServerName<'static>,
234 config: Arc<ClientConfig>,
235 ) -> io::Result<Buffered<ClientTls<S>, D>> {
236 if !self.outbound.is_empty() || !self.inbound.is_empty() {
237 return Err(io::Error::new(
238 io::ErrorKind::InvalidInput,
239 "TLS upgrade requires empty plaintext buffers",
240 ));
241 }
242 Ok(Buffered {
243 io: crate::tls::connect(self.io, server_name, config).await?,
244 outbound: self.outbound,
245 inbound: self.inbound,
246 inbound_codec: self.inbound_codec,
247 max_pre_startup_packet_len: self.max_pre_startup_packet_len,
248 demux: self.demux,
249 })
250 }
251
252 async fn accept_tls(
253 self,
254 config: Arc<ServerConfig>,
255 leaf_certificate: CertificateDer<'static>,
256 ) -> io::Result<Buffered<ServerTls<S>, D>> {
257 if !self.outbound.is_empty() || !self.inbound.is_empty() {
258 return Err(io::Error::new(
259 io::ErrorKind::InvalidInput,
260 "TLS upgrade requires empty plaintext buffers",
261 ));
262 }
263 Ok(Buffered {
264 io: crate::tls::accept(self.io, config, &leaf_certificate).await?,
265 outbound: self.outbound,
266 inbound: self.inbound,
267 inbound_codec: self.inbound_codec,
268 max_pre_startup_packet_len: self.max_pre_startup_packet_len,
269 demux: self.demux,
270 })
271 }
272}
273
274impl<S: TlsServerEndPoint, D> TlsServerEndPoint for Buffered<S, D> {
275 fn tls_server_end_point(&self) -> &[u8] {
276 self.io.tls_server_end_point()
277 }
278}
279
280impl<S: AsyncWrite + Unpin, D> Buffered<S, D> {
281 pub async fn flush(&mut self) -> io::Result<()> {
291 while !self.outbound.is_empty() {
292 let written = self.io.write(&self.outbound).await?;
293 if written == 0 {
294 return Err(io::Error::new(
295 io::ErrorKind::WriteZero,
296 "transport wrote zero buffered bytes",
297 ));
298 }
299 self.outbound.advance(written);
300 }
301 self.io.flush().await
302 }
303}
304
305impl<S: AsyncRead + Unpin, D: Direction> Buffered<S, D> {
306 pub async fn receive_wire(&mut self) -> io::Result<D::Message> {
312 loop {
313 if let Some(message) = self.inbound_codec.decode(&mut self.inbound)? {
314 return Ok(message);
315 }
316 if self.io.read_buf(&mut self.inbound).await? == 0 {
317 return Err(io::Error::new(
318 io::ErrorKind::UnexpectedEof,
319 "peer closed with no complete message",
320 ));
321 }
322 }
323 }
324}
325
326impl<S: AsyncRead + Unpin> Buffered<S, Backend> {
327 async fn receive_encryption_reply(&mut self) -> io::Result<EncryptionReply> {
328 let byte = self.io.read_u8().await?;
329 EncryptionReply::try_from(byte)
330 .map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid encryption reply"))
331 }
332}
333
334impl<S: AsyncRead + Unpin> Buffered<S, Frontend> {
335 pub async fn receive_pre_startup(&mut self) -> io::Result<PreStartupMessage> {
341 loop {
342 if let Some(message) =
343 decode_pre_startup_with_limit(&mut self.inbound, self.max_pre_startup_packet_len)?
344 {
345 return Ok(message);
346 }
347 if self.io.read_buf(&mut self.inbound).await? == 0 {
348 return Err(io::Error::new(
349 io::ErrorKind::UnexpectedEof,
350 "client closed with no complete pre-startup packet",
351 ));
352 }
353 }
354 }
355}
356
357impl<S: AsyncRead + Unpin> Buffered<S, Backend> {
358 pub async fn receive_backend(&mut self) -> io::Result<BackendMessage> {
364 self.receive_wire().await
365 }
366
367 pub async fn receive_session(&mut self) -> io::Result<SessionItem> {
373 loop {
374 let message = self.receive_backend().await?;
375 if let Some(item) = self.project_backend(message) {
376 return Ok(item);
377 }
378 }
379 }
380
381 pub fn project_backend(&mut self, message: BackendMessage) -> Option<SessionItem> {
383 self.demux.route(message)
384 }
385}
386
387impl<S, D, Phase, Cleanliness> Conn<Buffered<S, D>, Phase, Cleanliness> {
388 pub fn push_frame(&mut self, frame: Frame) -> io::Result<()> {
394 self.transport_mut().push(frame)
395 }
396
397 #[must_use]
398 pub fn pending_output(&self) -> &[u8] {
400 self.transport().pending()
401 }
402}
403
404impl<S, Cleanliness> Conn<Buffered<S, Backend>, PreStartup, Cleanliness> {
405 pub fn request_ssl(mut self) -> Conn<Buffered<S, Backend>, AwaitingSslReply, Cleanliness> {
407 self.transport_mut().push_raw(&ssl_request_packet());
408 self.transition()
409 }
410
411 pub fn request_gss(
413 mut self,
414 ) -> Conn<Buffered<S, Backend>, crate::pre_startup::AwaitingGssReply, Cleanliness> {
415 self.transport_mut().push_raw(&gssenc_request_packet());
416 self.transition()
417 }
418}
419
420impl<S, Cleanliness> Conn<Buffered<S, Frontend>, ServerSslDecision, Cleanliness> {
421 pub fn approve_ssl(mut self) -> Conn<Buffered<S, Frontend>, TlsHandshake, Cleanliness> {
423 self.transport_mut().push_raw(b"S");
424 self.transition()
425 }
426
427 pub fn decline_ssl(mut self) -> Conn<Buffered<S, Frontend>, PreStartup, Cleanliness> {
429 self.transport_mut().push_raw(b"N");
430 self.transition()
431 }
432
433 pub fn reject_ssl_with_legacy_error(
435 mut self,
436 ) -> Conn<Buffered<S, Frontend>, crate::pre_startup::Terminated, Cleanliness> {
437 self.transport_mut().push_raw(b"E");
438 self.transition()
439 }
440}
441
442impl<S, Cleanliness>
443 Conn<Buffered<S, Frontend>, crate::pre_startup::ServerGssDecision, Cleanliness>
444{
445 pub fn approve_gss(
447 mut self,
448 ) -> Conn<Buffered<S, Frontend>, crate::pre_startup::GssHandshake, Cleanliness> {
449 self.transport_mut().push_raw(b"S");
450 self.transition()
451 }
452
453 pub fn decline_gss(mut self) -> Conn<Buffered<S, Frontend>, PreStartup, Cleanliness> {
455 self.transport_mut().push_raw(b"N");
456 self.transition()
457 }
458
459 pub fn reject_gss_with_legacy_error(
461 mut self,
462 ) -> Conn<Buffered<S, Frontend>, crate::pre_startup::Terminated, Cleanliness> {
463 self.transport_mut().push_raw(b"E");
464 self.transition()
465 }
466}
467
468impl<S: AsyncRead + Unpin, Cleanliness> Conn<Buffered<S, Backend>, AwaitingSslReply, Cleanliness> {
469 pub async fn receive_ssl_reply(
475 mut self,
476 ) -> io::Result<Negotiation<Buffered<S, Backend>, TlsHandshake, Cleanliness>> {
477 let reply = self.transport_mut().receive_encryption_reply().await?;
478 Ok(match reply {
479 EncryptionReply::Accepted => Negotiation::Accepted(self.transition()),
480 EncryptionReply::Rejected => Negotiation::Rejected(self.transition()),
481 EncryptionReply::LegacyError => Negotiation::LegacyError(self.transition()),
482 })
483 }
484
485 pub async fn receive_ssl_reply_for_mode(
491 mut self,
492 mode: SslMode,
493 ) -> io::Result<SslModeNegotiation<Buffered<S, Backend>, Cleanliness>> {
494 let reply = self.transport_mut().receive_encryption_reply().await?;
495 Ok(self.apply_ssl_reply(reply, mode))
496 }
497}
498
499impl<S: AsyncRead + Unpin, Cleanliness>
500 Conn<Buffered<S, Backend>, crate::pre_startup::AwaitingGssReply, Cleanliness>
501{
502 pub async fn receive_gss_reply(
508 mut self,
509 ) -> io::Result<Negotiation<Buffered<S, Backend>, crate::pre_startup::GssHandshake, Cleanliness>>
510 {
511 let reply = self.transport_mut().receive_encryption_reply().await?;
512 Ok(match reply {
513 EncryptionReply::Accepted => Negotiation::Accepted(self.transition()),
514 EncryptionReply::Rejected => Negotiation::Rejected(self.transition()),
515 EncryptionReply::LegacyError => Negotiation::LegacyError(self.transition()),
516 })
517 }
518}
519
520impl<S, Cleanliness> Conn<Buffered<S, Backend>, TlsHandshake, Cleanliness>
521where
522 S: AsyncRead + AsyncWrite + Unpin,
523{
524 pub async fn connect_tls(
530 self,
531 server_name: ServerName<'static>,
532 config: Arc<ClientConfig>,
533 ) -> io::Result<Conn<Buffered<ClientTls<S>, Backend>, PreStartup, Cleanliness>> {
534 let transport = self.into_transport();
535 Ok(Conn::new(transport.connect_tls(server_name, config).await?)
536 .transition::<PreStartup, Cleanliness>())
537 }
538}
539
540impl<S, Cleanliness> Conn<Buffered<S, Frontend>, TlsHandshake, Cleanliness>
541where
542 S: AsyncRead + AsyncWrite + Unpin,
543{
544 pub async fn accept_tls(
550 self,
551 config: Arc<ServerConfig>,
552 leaf_certificate: CertificateDer<'static>,
553 ) -> io::Result<Conn<Buffered<ServerTls<S>, Frontend>, PreStartup, Cleanliness>> {
554 let transport = self.into_transport();
555 Ok(
556 Conn::new(transport.accept_tls(config, leaf_certificate).await?)
557 .transition::<PreStartup, Cleanliness>(),
558 )
559 }
560}
561
562impl<S, D, Cleanliness> Conn<Buffered<S, D>, crate::pre_startup::Startup, Cleanliness> {
563 pub fn push_startup_packet(&mut self, packet: &[u8]) {
565 self.transport_mut().outbound.extend_from_slice(packet);
566 }
567}
568
569impl<S: AsyncWrite + Unpin, D, Phase, Cleanliness> Conn<Buffered<S, D>, Phase, Cleanliness> {
570 pub async fn flush(&mut self) -> io::Result<()> {
576 self.transport_mut().flush().await
577 }
578}
579
580impl<S: AsyncRead + Unpin, Phase, Cleanliness> Conn<Buffered<S, Backend>, Phase, Cleanliness> {
581 pub async fn receive_backend_wire(&mut self) -> io::Result<BackendMessage> {
588 self.transport_mut().receive_backend().await
589 }
590
591 pub async fn receive_backend_typed<State, Handler>(
602 &mut self,
603 middleware: &mut Middleware<State, Handler>,
604 ) -> Result<
605 <Phase as PhaseAssociation<Inbound, ServerRole, BackendMessage>>::Message,
606 TypedReceiveError<Handler::Error, BackendMessage>,
607 >
608 where
609 Phase: PhaseAssociation<Inbound, ServerRole, BackendMessage>,
610 Handler: TypedMiddleware<
611 ServerRole,
612 <Phase as PhaseAssociation<Inbound, ServerRole, BackendMessage>>::ProtocolPhase,
613 <Phase as PhaseAssociation<Inbound, ServerRole, BackendMessage>>::Message,
614 State,
615 >,
616 {
617 receive_typed::<ServerRole, Phase, BackendMessage, _, _, _>(
618 self.receive_backend_wire(),
619 middleware,
620 )
621 .await
622 }
623
624 pub async fn receive_typed<State, Handler>(
633 &mut self,
634 middleware: &mut Middleware<State, Handler>,
635 ) -> Result<SessionItem, TypedReceiveError<Handler::Error, BackendMessage>>
636 where
637 Phase: PhaseAssociation<Inbound, ServerRole, BackendMessage>,
638 Handler: TypedMiddleware<
639 ServerRole,
640 <Phase as PhaseAssociation<Inbound, ServerRole, BackendMessage>>::ProtocolPhase,
641 <Phase as PhaseAssociation<Inbound, ServerRole, BackendMessage>>::Message,
642 State,
643 >,
644 {
645 loop {
646 let message = self.receive_backend_typed(middleware).await?;
647 if let Some(item) = self.project_backend(message.into()) {
648 return Ok(item);
649 }
650 }
651 }
652
653 pub async fn receive_encryption_reply_typed<State, Handler>(
663 &mut self,
664 middleware: &mut Middleware<State, Handler>,
665 ) -> Result<
666 <Phase as PhaseAssociation<Inbound, ServerRole, EncryptionReply>>::Message,
667 TypedReceiveError<Handler::Error, EncryptionReply>,
668 >
669 where
670 Phase: PhaseAssociation<Inbound, ServerRole, EncryptionReply>,
671 Handler: TypedMiddleware<
672 ServerRole,
673 <Phase as PhaseAssociation<Inbound, ServerRole, EncryptionReply>>::ProtocolPhase,
674 <Phase as PhaseAssociation<Inbound, ServerRole, EncryptionReply>>::Message,
675 State,
676 >,
677 {
678 receive_typed::<ServerRole, Phase, EncryptionReply, _, _, _>(
679 self.transport_mut().receive_encryption_reply(),
680 middleware,
681 )
682 .await
683 }
684
685 pub async fn receive_backend_wire_with_middleware<State, Handler, ProtocolState>(
691 &mut self,
692 middleware: &mut Middleware<State, Handler>,
693 protocol_state: &ProtocolState,
694 ) -> Result<BackendMessage, ReceiveError<Handler::Error, BackendMessage>>
695 where
696 Handler: MessageMiddleware<BackendMessage, State>,
697 ProtocolState: AcceptsMessage<BackendMessage>,
698 {
699 receive_checked(self.receive_backend_wire(), middleware, protocol_state).await
700 }
701
702 pub fn project_backend(&mut self, message: BackendMessage) -> Option<SessionItem> {
704 self.transport_mut().project_backend(message)
705 }
706
707 pub async fn receive(&mut self) -> io::Result<SessionItem> {
713 self.transport_mut().receive_session().await
714 }
715
716 pub async fn receive_with_middleware<State, Handler, ProtocolState>(
725 &mut self,
726 middleware: &mut Middleware<State, Handler>,
727 protocol_state: &ProtocolState,
728 ) -> Result<SessionItem, ReceiveError<Handler::Error, BackendMessage>>
729 where
730 Handler: MessageMiddleware<BackendMessage, State>,
731 ProtocolState: AcceptsMessage<BackendMessage>,
732 {
733 loop {
734 let message = self
735 .receive_backend_wire_with_middleware(middleware, protocol_state)
736 .await?;
737 if let Some(item) = self.project_backend(message) {
738 return Ok(item);
739 }
740 }
741 }
742
743 #[must_use]
744 pub fn cancel_key(&self) -> Option<&CancelKey> {
746 self.transport().demux().cancel_key()
747 }
748
749 #[must_use]
751 pub fn parameters(&self) -> &BTreeMap<Bytes, Bytes> {
752 self.transport().demux().parameters()
753 }
754
755 #[must_use]
757 pub fn parameters_changed(&self) -> bool {
758 self.transport().demux().parameters_changed()
759 }
760
761 #[must_use]
763 pub fn transaction_status(&self) -> Option<crate::codec::TransactionStatus> {
764 self.transport().demux().transaction_status()
765 }
766
767 pub fn pop_notification(&mut self) -> Option<Notification> {
769 self.transport_mut().demux_mut().pop_notification()
770 }
771
772 pub fn pop_notice(&mut self) -> Option<TaggedNotice> {
774 self.transport_mut().demux_mut().pop_notice()
775 }
776
777 pub fn pop_parameter_status(&mut self) -> Option<ParameterStatus> {
779 self.transport_mut().demux_mut().pop_parameter_status()
780 }
781
782 pub fn pop_async_event(&mut self) -> Option<OrderedAsyncEvent> {
784 self.transport_mut().demux_mut().pop_async_event()
785 }
786}
787
788impl<S: AsyncRead + Unpin, Phase, Cleanliness> Conn<Buffered<S, Frontend>, Phase, Cleanliness> {
789 pub async fn receive_frontend_wire(&mut self) -> io::Result<FrontendMessage> {
795 self.transport_mut().receive_wire().await
796 }
797
798 pub async fn receive_frontend_typed<State, Handler>(
808 &mut self,
809 middleware: &mut Middleware<State, Handler>,
810 ) -> Result<
811 <Phase as PhaseAssociation<Inbound, ClientRole, FrontendMessage>>::Message,
812 TypedReceiveError<Handler::Error, FrontendMessage>,
813 >
814 where
815 Phase: PhaseAssociation<Inbound, ClientRole, FrontendMessage>,
816 Handler: TypedMiddleware<
817 ClientRole,
818 <Phase as PhaseAssociation<Inbound, ClientRole, FrontendMessage>>::ProtocolPhase,
819 <Phase as PhaseAssociation<Inbound, ClientRole, FrontendMessage>>::Message,
820 State,
821 >,
822 {
823 receive_typed::<ClientRole, Phase, FrontendMessage, _, _, _>(
824 self.receive_frontend_wire(),
825 middleware,
826 )
827 .await
828 }
829
830 pub async fn receive_frontend_wire_with_middleware<State, Handler, ProtocolState>(
836 &mut self,
837 middleware: &mut Middleware<State, Handler>,
838 protocol_state: &ProtocolState,
839 ) -> Result<FrontendMessage, ReceiveError<Handler::Error, FrontendMessage>>
840 where
841 Handler: MessageMiddleware<FrontendMessage, State>,
842 ProtocolState: AcceptsMessage<FrontendMessage>,
843 {
844 receive_checked(self.receive_frontend_wire(), middleware, protocol_state).await
845 }
846}
847
848impl<S: AsyncRead + Unpin, Cleanliness> Conn<Buffered<S, Frontend>, PreStartup, Cleanliness> {
849 pub async fn receive_pre_startup_wire(&mut self) -> io::Result<PreStartupMessage> {
855 self.transport_mut().receive_pre_startup().await
856 }
857
858 pub async fn receive_pre_startup_typed<State, Handler>(
865 &mut self,
866 middleware: &mut Middleware<State, Handler>,
867 ) -> Result<
868 <PreStartup as PhaseAssociation<Inbound, ClientRole, PreStartupMessage>>::Message,
869 TypedReceiveError<Handler::Error, PreStartupMessage>,
870 >
871 where
872 Handler: TypedMiddleware<
873 ClientRole,
874 <PreStartup as PhaseAssociation<Inbound, ClientRole, PreStartupMessage>>::ProtocolPhase,
875 <PreStartup as PhaseAssociation<Inbound, ClientRole, PreStartupMessage>>::Message,
876 State,
877 >,
878 {
879 receive_typed::<ClientRole, PreStartup, PreStartupMessage, _, _, _>(
880 self.receive_pre_startup_wire(),
881 middleware,
882 )
883 .await
884 }
885
886 pub async fn receive_pre_startup_wire_with_middleware<State, Handler, ProtocolState>(
892 &mut self,
893 middleware: &mut Middleware<State, Handler>,
894 protocol_state: &ProtocolState,
895 ) -> Result<PreStartupMessage, ReceiveError<Handler::Error, PreStartupMessage>>
896 where
897 Handler: MessageMiddleware<PreStartupMessage, State>,
898 ProtocolState: AcceptsMessage<PreStartupMessage>,
899 {
900 receive_checked(self.receive_pre_startup_wire(), middleware, protocol_state).await
901 }
902}
903
904#[cfg(test)]
905mod tests {
906 use std::{
907 convert::Infallible,
908 future::Future,
909 pin::Pin,
910 task::{Context, Poll},
911 };
912
913 use bytes::Bytes;
914 use tokio::io::AsyncWrite;
915
916 use super::*;
917 use crate::{
918 grammar::{backend, frontend, pre_startup as pre_startup_grammar, server_pre_startup},
919 middleware::{InterceptError, Middleware, ReceiveError},
920 };
921
922 #[derive(Debug, Default)]
923 struct ShortWriter {
924 output: Vec<u8>,
925 }
926
927 impl AsyncWrite for ShortWriter {
928 fn poll_write(
929 mut self: Pin<&mut Self>,
930 _cx: &mut Context<'_>,
931 buffer: &[u8],
932 ) -> Poll<io::Result<usize>> {
933 let written = buffer.len().min(2);
934 self.output.extend_from_slice(&buffer[..written]);
935 Poll::Ready(Ok(written))
936 }
937
938 fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
939 Poll::Ready(Ok(()))
940 }
941
942 fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
943 Poll::Ready(Ok(()))
944 }
945 }
946
947 #[tokio::test]
948 async fn flush_handles_partial_writes_without_losing_bytes() {
949 let frame = Frame {
950 tag: b'S',
951 body: Bytes::new(),
952 };
953 let mut transport = Buffered::new(ShortWriter::default());
954 transport.push(frame).expect("encodable frame");
955 assert_eq!(transport.pending(), &[b'S', 0, 0, 0, 4]);
956 transport.flush().await.expect("writable transport");
957 assert!(transport.pending().is_empty());
958 assert_eq!(transport.into_inner().output, [b'S', 0, 0, 0, 4]);
959 }
960
961 #[test]
962 fn buffered_transport_enforces_its_frame_limit_on_output() {
963 let mut transport = Buffered::<_, Backend>::with_max_frame_len((), 9).unwrap();
964 let error = transport
965 .push(Frame {
966 tag: b'Q',
967 body: Bytes::from_static(b"12345"),
968 })
969 .unwrap_err();
970 assert_eq!(error.kind(), io::ErrorKind::InvalidInput);
971 assert!(transport.pending().is_empty());
972 }
973
974 #[test]
975 fn cancelling_flush_retains_unwritten_bytes() {
976 #[derive(Debug, Default)]
977 struct PausingWriter {
978 output: Vec<u8>,
979 blocked: bool,
980 }
981
982 impl AsyncWrite for PausingWriter {
983 fn poll_write(
984 mut self: Pin<&mut Self>,
985 _cx: &mut Context<'_>,
986 buffer: &[u8],
987 ) -> Poll<io::Result<usize>> {
988 if self.blocked {
989 return Poll::Pending;
990 }
991 let written = buffer.len().min(2);
992 self.output.extend_from_slice(&buffer[..written]);
993 self.blocked = true;
994 Poll::Ready(Ok(written))
995 }
996
997 fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
998 Poll::Ready(Ok(()))
999 }
1000
1001 fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
1002 Poll::Ready(Ok(()))
1003 }
1004 }
1005
1006 let mut transport = Buffered::new(PausingWriter::default());
1007 transport
1008 .push(Frame {
1009 tag: b'S',
1010 body: Bytes::new(),
1011 })
1012 .expect("encodable frame");
1013
1014 let mut flush = Box::pin(transport.flush());
1015 let waker = std::task::Waker::noop();
1016 let mut context = Context::from_waker(waker);
1017 assert!(flush.as_mut().poll(&mut context).is_pending());
1018 drop(flush);
1019
1020 assert_eq!(transport.pending(), &[0, 0, 4]);
1021 assert_eq!(transport.io.output, [b'S', 0]);
1022 }
1023
1024 #[tokio::test]
1025 async fn receive_filters_parameter_status_before_session_message() {
1026 let (client, mut server) = tokio::io::duplex(256);
1027 let mut wire = BytesMut::new();
1028 let mut encoder = PgCodec::<Backend>::default();
1029 encoder
1030 .encode(
1031 Frame {
1032 tag: b'S',
1033 body: Bytes::from_static(b"client_encoding\0UTF8\0"),
1034 },
1035 &mut wire,
1036 )
1037 .expect("encodable ParameterStatus");
1038 encoder
1039 .encode(
1040 Frame {
1041 tag: b'Z',
1042 body: Bytes::from_static(b"I"),
1043 },
1044 &mut wire,
1045 )
1046 .expect("encodable ReadyForQuery");
1047 server.write_all(&wire).await.expect("writable test peer");
1048
1049 let mut transport = Buffered::new(client);
1050 assert_eq!(
1051 transport.receive_session().await.expect("valid messages"),
1052 SessionItem::ReadyForQuery {
1053 status: crate::codec::TransactionStatus::Idle,
1054 parameters_changed: false,
1055 }
1056 );
1057 assert_eq!(
1058 transport
1059 .demux()
1060 .parameters()
1061 .get(&Bytes::from_static(b"client_encoding")),
1062 Some(&Bytes::from_static(b"UTF8"))
1063 );
1064 let conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1065 assert_eq!(
1066 conn.parameters().get(b"client_encoding".as_slice()),
1067 Some(&Bytes::from_static(b"UTF8"))
1068 );
1069 assert!(!conn.parameters_changed());
1070 assert_eq!(
1071 conn.transaction_status(),
1072 Some(crate::codec::TransactionStatus::Idle)
1073 );
1074 conn.into_transport();
1075 }
1076
1077 #[tokio::test]
1078 async fn wire_message_can_be_modified_before_projection() {
1079 let (client, mut server) = tokio::io::duplex(128);
1080 let original = BackendMessage::ParameterStatus {
1081 name: Bytes::from_static(b"application_name"),
1082 value: Bytes::from_static(b"upstream"),
1083 };
1084 let mut bytes = BytesMut::new();
1085 PgCodec::<Backend>::default()
1086 .encode(
1087 original.to_frame().expect("reconstructable message"),
1088 &mut bytes,
1089 )
1090 .expect("encodable message");
1091 server.write_all(&bytes).await.expect("writable test peer");
1092
1093 let mut transport = Buffered::new(client);
1094 let mut message = transport
1095 .receive_backend()
1096 .await
1097 .expect("decodable message");
1098 let BackendMessage::ParameterStatus { value, .. } = &mut message else {
1099 panic!("unexpected message")
1100 };
1101 *value = Bytes::from_static(b"proxy");
1102 assert!(transport.project_backend(message).is_none());
1103 assert_eq!(
1104 transport
1105 .demux()
1106 .parameters()
1107 .get(&Bytes::from_static(b"application_name")),
1108 Some(&Bytes::from_static(b"proxy"))
1109 );
1110 }
1111
1112 #[tokio::test]
1113 async fn typed_middleware_accepts_async_traffic_without_advancing_ready() {
1114 let (client, mut server) = tokio::io::duplex(128);
1115 let original = BackendMessage::ParameterStatus {
1116 name: Bytes::from_static(b"application_name"),
1117 value: Bytes::from_static(b"upstream"),
1118 };
1119 let mut bytes = BytesMut::new();
1120 PgCodec::<Backend>::default()
1121 .encode(
1122 original.to_frame().expect("reconstructable message"),
1123 &mut bytes,
1124 )
1125 .expect("encodable message");
1126 server.write_all(&bytes).await.expect("writable test peer");
1127
1128 let transport = Buffered::new(client);
1129 let mut conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1130 let mut middleware = Middleware::new(0_usize, async |seen: &mut usize, _message| {
1131 *seen += 1;
1132 let replacement = BackendMessage::ParameterStatus {
1133 name: Bytes::from_static(b"application_name"),
1134 value: Bytes::from_static(b"proxy"),
1135 };
1136 match crate::middleware::TypedBackendMessage::try_from(replacement) {
1137 Ok(replacement) => Ok::<_, Infallible>(replacement),
1138 Err(message) => panic!("parameter status must be asynchronous: {message:?}"),
1139 }
1140 });
1141
1142 let message = conn
1143 .receive_backend_typed(&mut middleware)
1144 .await
1145 .expect("typed asynchronous message");
1146 assert!(conn.project_backend(message.into()).is_none());
1147 assert_eq!(*middleware.state(), 1);
1148 assert_eq!(
1149 conn.parameters().get(b"application_name".as_slice()),
1150 Some(&Bytes::from_static(b"proxy"))
1151 );
1152 conn.into_transport();
1153 }
1154
1155 #[tokio::test]
1156 async fn typed_receive_projects_into_the_existing_next_connection_enum() {
1157 let (client, mut server) = tokio::io::duplex(128);
1158 let ready = BackendMessage::ReadyForQuery(crate::codec::TransactionStatus::Idle);
1159 let mut bytes = BytesMut::new();
1160 PgCodec::<Backend>::default()
1161 .encode(
1162 ready.to_frame().expect("reconstructable message"),
1163 &mut bytes,
1164 )
1165 .expect("encodable message");
1166 server.write_all(&bytes).await.expect("writable test peer");
1167
1168 let transport = Buffered::new(client);
1169 let conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1170 let (mut query, _) = conn
1171 .push_stateless_query(b"select 1")
1172 .expect("encodable query");
1173 let mut middleware = Middleware::new((), crate::middleware::Identity);
1174 let item = query
1175 .receive_typed(&mut middleware)
1176 .await
1177 .expect("phase-legal ready message");
1178
1179 let transition = query.offer(item).expect("typed next-state projection");
1180 let crate::session::SimpleTransition::Ready(crate::session::ReadyState::Clean(ready)) =
1181 transition
1182 else {
1183 panic!("idle readiness must return a clean ready connection");
1184 };
1185 ready.into_transport();
1186 }
1187
1188 #[tokio::test]
1189 async fn typed_pre_startup_receive_intercepts_without_projecting() {
1190 let (proxy, mut client) = tokio::io::duplex(128);
1191 client
1192 .write_all(
1193 &PreStartupMessage::SslRequest
1194 .to_packet()
1195 .expect("encodable SSLRequest"),
1196 )
1197 .await
1198 .expect("writable client");
1199
1200 let transport = Buffered::<_, Frontend>::new_frontend(proxy);
1201 let mut conn = Conn::new(transport);
1202 let mut middleware = Middleware::new(0_usize, async |seen: &mut usize, _message| {
1203 *seen += 1;
1204 match server_pre_startup::PreStartupExternalMessage::try_from(
1205 PreStartupMessage::CancelRequest {
1206 process_id: 42,
1207 secret_key: Bytes::from_static(b"key!"),
1208 },
1209 ) {
1210 Ok(replacement) => Ok::<_, Infallible>(replacement),
1211 Err(message) => panic!("cancel request must be legal before startup: {message:?}"),
1212 }
1213 });
1214
1215 let message = conn
1216 .receive_pre_startup_typed(&mut middleware)
1217 .await
1218 .expect("phase-legal pre-startup packet");
1219
1220 assert_eq!(
1221 PreStartupMessage::from(message),
1222 PreStartupMessage::CancelRequest {
1223 process_id: 42,
1224 secret_key: Bytes::from_static(b"key!"),
1225 }
1226 );
1227 assert_eq!(*middleware.state(), 1);
1228 let conn: Conn<_, PreStartup> = conn;
1229 conn.into_transport();
1230 }
1231
1232 #[tokio::test]
1233 async fn typed_encryption_reply_receive_intercepts_without_projecting() {
1234 let (client, mut server) = tokio::io::duplex(16);
1235 server.write_all(b"S").await.expect("writable test peer");
1236
1237 let transport = Buffered::new(client);
1238 let mut conn: Conn<_, AwaitingSslReply> = Conn::new(transport).transition();
1239 let mut middleware = Middleware::new(0_usize, async |seen: &mut usize, _message| {
1240 *seen += 1;
1241 match pre_startup_grammar::AwaitingSslReplyExternalMessage::try_from(
1242 EncryptionReply::Rejected,
1243 ) {
1244 Ok(replacement) => Ok::<_, Infallible>(replacement),
1245 Err(message) => panic!("SSL rejection must be phase legal: {message:?}"),
1246 }
1247 });
1248
1249 let message = conn
1250 .receive_encryption_reply_typed(&mut middleware)
1251 .await
1252 .expect("phase-legal SSL reply");
1253
1254 assert_eq!(EncryptionReply::from(message), EncryptionReply::Rejected);
1255 assert_eq!(*middleware.state(), 1);
1256 let conn: Conn<_, AwaitingSslReply> = conn;
1257 conn.into_transport();
1258 }
1259
1260 #[tokio::test]
1261 async fn typed_receive_preserves_large_replacement_allocation() {
1262 let (proxy, mut client) = tokio::io::duplex(128);
1263 let query = FrontendMessage::Query(Bytes::from_static(b"select 1"));
1264 let mut bytes = BytesMut::new();
1265 PgCodec::<Frontend>::default()
1266 .encode(query.to_frame().expect("reconstructable Query"), &mut bytes)
1267 .expect("encodable Query");
1268 client.write_all(&bytes).await.expect("writable client");
1269
1270 let transport = Buffered::<_, Frontend>::new_frontend(proxy);
1271 let mut conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1272 let replacement_payload = Bytes::from(vec![b'x'; 2 * 1024 * 1024]);
1273 let replacement_pointer = replacement_payload.as_ptr();
1274 let mut middleware = Middleware::new((), async move |_state: &mut (), _message| {
1275 match backend::ReadyExternalMessage::try_from(FrontendMessage::Query(
1276 replacement_payload.clone(),
1277 )) {
1278 Ok(replacement) => Ok::<_, Infallible>(replacement),
1279 Err(message) => panic!("query must be legal while ready: {message:?}"),
1280 }
1281 });
1282
1283 let message = conn
1284 .receive_frontend_typed(&mut middleware)
1285 .await
1286 .expect("phase-legal replacement");
1287 let FrontendMessage::Query(replacement) = FrontendMessage::from(message) else {
1288 panic!("expected replacement Query")
1289 };
1290 assert_eq!(replacement.as_ptr(), replacement_pointer);
1291 conn.into_transport();
1292 }
1293
1294 #[tokio::test]
1295 async fn typed_receive_keeps_wire_shape_validation_at_runtime() {
1296 let (client, mut peer) = tokio::io::duplex(256);
1297 let query = FrontendMessage::Query(Bytes::from_static(b"select 1"));
1298 let mut bytes = BytesMut::new();
1299 PgCodec::<Frontend>::default()
1300 .encode(query.to_frame().expect("reconstructable query"), &mut bytes)
1301 .expect("encodable query");
1302 peer.write_all(&bytes).await.expect("writable test peer");
1303
1304 let transport = Buffered::<_, Frontend>::new_frontend(client);
1305 let mut conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1306 let replacement_payload = Bytes::from(vec![b'x'; 2 * 1024 * 1024]);
1307 let replacement_pointer = replacement_payload.as_ptr();
1308 let invalid = FrontendMessage::Parse(crate::codec::Parse {
1309 statement: Bytes::from_static(b"invalid\0statement"),
1310 query: replacement_payload,
1311 parameter_types: Vec::new(),
1312 });
1313 let mut middleware = Middleware::new((), async move |_state: &mut (), _message| {
1314 let invalid = invalid.clone();
1315 match backend::ReadyExternalMessage::try_from(invalid) {
1316 Ok(invalid) => Ok::<_, Infallible>(invalid),
1317 Err(message) => panic!("parse must be protocol-legal while ready: {message:?}"),
1318 }
1319 });
1320
1321 let result = conn.receive_frontend_typed(&mut middleware).await;
1322 let Err(TypedReceiveError::InvalidWire(FrontendMessage::Parse(replacement))) = result
1323 else {
1324 panic!("expected the invalid replacement wire value")
1325 };
1326 assert_eq!(replacement.query.as_ptr(), replacement_pointer);
1327 conn.into_transport();
1328 }
1329
1330 #[tokio::test]
1331 async fn typed_receive_rejects_illegal_peer_input_before_middleware() {
1332 let (proxy, mut client) = tokio::io::duplex(2 * 1024 * 1024 + 64);
1333 let illegal_payload = Bytes::from(vec![b'x'; 2 * 1024 * 1024]);
1334 let illegal = FrontendMessage::CopyData(illegal_payload);
1335 let mut bytes = BytesMut::new();
1336 PgCodec::<Frontend>::default()
1337 .encode(
1338 illegal.to_frame().expect("reconstructable CopyData"),
1339 &mut bytes,
1340 )
1341 .expect("encodable CopyData");
1342 client.write_all(&bytes).await.expect("writable client");
1343
1344 let transport = Buffered::<_, Frontend>::new_frontend(proxy);
1345 let mut conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1346 let mut middleware = Middleware::new(0_usize, async |seen: &mut usize, message| {
1347 *seen += 1;
1348 Ok::<_, Infallible>(message)
1349 });
1350
1351 let result = conn.receive_frontend_typed(&mut middleware).await;
1352
1353 let Err(TypedReceiveError::Illegal(FrontendMessage::CopyData(illegal))) = result else {
1354 panic!("expected the phase-illegal CopyData")
1355 };
1356 assert_eq!(illegal.len(), 2 * 1024 * 1024);
1357 assert!(illegal.iter().all(|byte| *byte == b'x'));
1358 assert_eq!(*middleware.state(), 0);
1359 conn.into_transport();
1360 }
1361
1362 #[tokio::test]
1363 async fn typed_backend_receive_rejects_illegal_peer_input_before_middleware() {
1364 let (client, mut server) = tokio::io::duplex(64);
1365 let mut bytes = BytesMut::new();
1366 PgCodec::<Backend>::default()
1367 .encode(
1368 BackendMessage::ParseComplete
1369 .to_frame()
1370 .expect("reconstructable ParseComplete"),
1371 &mut bytes,
1372 )
1373 .expect("encodable ParseComplete");
1374 server.write_all(&bytes).await.expect("writable server");
1375
1376 let mut conn: Conn<_, crate::auth::Ready> = Conn::new(Buffered::new(client)).transition();
1377 let mut middleware = Middleware::new(0_usize, async |seen: &mut usize, message| {
1378 *seen += 1;
1379 Ok::<_, Infallible>(message)
1380 });
1381
1382 assert!(matches!(
1383 conn.receive_backend_typed(&mut middleware).await,
1384 Err(TypedReceiveError::Illegal(BackendMessage::ParseComplete))
1385 ));
1386 assert_eq!(*middleware.state(), 0);
1387 conn.into_transport();
1388 }
1389
1390 #[tokio::test]
1391 async fn typed_backend_and_pre_startup_receive_return_invalid_replacements() {
1392 let (backend_side, mut server) = tokio::io::duplex(128);
1393 let mut backend_bytes = BytesMut::new();
1394 PgCodec::<Backend>::default()
1395 .encode(
1396 BackendMessage::ParameterStatus {
1397 name: Bytes::from_static(b"application_name"),
1398 value: Bytes::from_static(b"upstream"),
1399 }
1400 .to_frame()
1401 .expect("reconstructable ParameterStatus"),
1402 &mut backend_bytes,
1403 )
1404 .expect("encodable ParameterStatus");
1405 server
1406 .write_all(&backend_bytes)
1407 .await
1408 .expect("writable server");
1409 let mut backend: Conn<_, crate::auth::Ready> =
1410 Conn::new(Buffered::new(backend_side)).transition();
1411 let mut backend_middleware = Middleware::new((), async |_state: &mut (), _message| {
1412 let invalid = BackendMessage::ParameterStatus {
1413 name: Bytes::from_static(b"invalid\0name"),
1414 value: Bytes::from_static(b"value"),
1415 };
1416 match crate::middleware::TypedBackendMessage::try_from(invalid) {
1417 Ok(replacement) => Ok::<_, Infallible>(replacement),
1418 Err(message) => panic!("ParameterStatus must be phase legal: {message:?}"),
1419 }
1420 });
1421 assert!(matches!(
1422 backend.receive_backend_typed(&mut backend_middleware).await,
1423 Err(TypedReceiveError::InvalidWire(
1424 BackendMessage::ParameterStatus { .. }
1425 ))
1426 ));
1427 backend.into_transport();
1428
1429 let (pre_startup_side, mut client) = tokio::io::duplex(128);
1430 client
1431 .write_all(
1432 &PreStartupMessage::SslRequest
1433 .to_packet()
1434 .expect("encodable SSLRequest"),
1435 )
1436 .await
1437 .expect("writable client");
1438 let mut pre_startup = Conn::new(Buffered::<_, Frontend>::new_frontend(pre_startup_side));
1439 let invalid_key = Bytes::from_static(b"bad");
1440 let mut pre_startup_middleware =
1441 Middleware::new((), async move |_state: &mut (), _message| {
1442 match server_pre_startup::PreStartupExternalMessage::try_from(
1443 PreStartupMessage::CancelRequest {
1444 process_id: 42,
1445 secret_key: invalid_key.clone(),
1446 },
1447 ) {
1448 Ok(replacement) => Ok::<_, Infallible>(replacement),
1449 Err(message) => {
1450 panic!("CancelRequest must be phase legal before startup: {message:?}")
1451 }
1452 }
1453 });
1454 assert!(matches!(
1455 pre_startup
1456 .receive_pre_startup_typed(&mut pre_startup_middleware)
1457 .await,
1458 Err(TypedReceiveError::InvalidWire(
1459 PreStartupMessage::CancelRequest { .. }
1460 ))
1461 ));
1462 pre_startup.into_transport();
1463 }
1464
1465 #[tokio::test]
1466 async fn typed_receive_preserves_middleware_errors_for_every_wire_family() {
1467 let (backend_side, mut server) = tokio::io::duplex(128);
1468 let mut backend_bytes = BytesMut::new();
1469 PgCodec::<Backend>::default()
1470 .encode(
1471 BackendMessage::ParameterStatus {
1472 name: Bytes::from_static(b"application_name"),
1473 value: Bytes::from_static(b"upstream"),
1474 }
1475 .to_frame()
1476 .expect("reconstructable ParameterStatus"),
1477 &mut backend_bytes,
1478 )
1479 .expect("encodable ParameterStatus");
1480 server
1481 .write_all(&backend_bytes)
1482 .await
1483 .expect("writable server");
1484 let mut backend: Conn<_, crate::auth::Ready> =
1485 Conn::new(Buffered::new(backend_side)).transition();
1486 let mut backend_middleware = Middleware::new((), async |_state: &mut (), _message| {
1487 Err::<crate::middleware::TypedBackendMessage<frontend::ReadyExternalMessage>, _>(
1488 "blocked",
1489 )
1490 });
1491 assert!(matches!(
1492 backend.receive_backend_typed(&mut backend_middleware).await,
1493 Err(TypedReceiveError::Middleware("blocked"))
1494 ));
1495 backend.into_transport();
1496
1497 let (frontend_side, mut client) = tokio::io::duplex(128);
1498 let mut frontend_bytes = BytesMut::new();
1499 PgCodec::<Frontend>::default()
1500 .encode(
1501 FrontendMessage::Query(Bytes::from_static(b"select 1"))
1502 .to_frame()
1503 .expect("reconstructable Query"),
1504 &mut frontend_bytes,
1505 )
1506 .expect("encodable Query");
1507 client
1508 .write_all(&frontend_bytes)
1509 .await
1510 .expect("writable client");
1511 let mut frontend: Conn<_, crate::auth::Ready> =
1512 Conn::new(Buffered::<_, Frontend>::new_frontend(frontend_side)).transition();
1513 let mut frontend_middleware = Middleware::new((), async |_state: &mut (), _message| {
1514 Err::<backend::ReadyExternalMessage, _>("blocked")
1515 });
1516 assert!(matches!(
1517 frontend
1518 .receive_frontend_typed(&mut frontend_middleware)
1519 .await,
1520 Err(TypedReceiveError::Middleware("blocked"))
1521 ));
1522 frontend.into_transport();
1523
1524 let (pre_startup_side, mut client) = tokio::io::duplex(128);
1525 client
1526 .write_all(
1527 &PreStartupMessage::SslRequest
1528 .to_packet()
1529 .expect("encodable SSLRequest"),
1530 )
1531 .await
1532 .expect("writable client");
1533 let mut pre_startup = Conn::new(Buffered::<_, Frontend>::new_frontend(pre_startup_side));
1534 let mut pre_startup_middleware = Middleware::new((), async |_state: &mut (), _message| {
1535 Err::<server_pre_startup::PreStartupExternalMessage, _>("blocked")
1536 });
1537 assert!(matches!(
1538 pre_startup
1539 .receive_pre_startup_typed(&mut pre_startup_middleware)
1540 .await,
1541 Err(TypedReceiveError::Middleware("blocked"))
1542 ));
1543 pre_startup.into_transport();
1544
1545 let (encryption_side, mut server) = tokio::io::duplex(16);
1546 server.write_all(b"S").await.expect("writable server");
1547 let mut encryption: Conn<_, AwaitingSslReply> =
1548 Conn::new(Buffered::new(encryption_side)).transition();
1549 let mut encryption_middleware = Middleware::new((), async |_state: &mut (), _message| {
1550 Err::<pre_startup_grammar::AwaitingSslReplyExternalMessage, _>("blocked")
1551 });
1552 assert!(matches!(
1553 encryption
1554 .receive_encryption_reply_typed(&mut encryption_middleware)
1555 .await,
1556 Err(TypedReceiveError::Middleware("blocked"))
1557 ));
1558 encryption.into_transport();
1559 }
1560
1561 #[tokio::test]
1562 async fn middleware_rewrites_backend_before_demux_bookkeeping() {
1563 let (client, mut server) = tokio::io::duplex(256);
1564 let mut wire = BytesMut::new();
1565 let mut encoder = PgCodec::<Backend>::default();
1566 for message in [
1567 BackendMessage::BackendKeyData {
1568 process_id: 7,
1569 secret_key: Bytes::from_static(b"old!"),
1570 },
1571 BackendMessage::ParameterStatus {
1572 name: Bytes::from_static(b"application_name"),
1573 value: Bytes::from_static(b"upstream"),
1574 },
1575 BackendMessage::ReadyForQuery(crate::codec::TransactionStatus::Idle),
1576 ] {
1577 encoder
1578 .encode(
1579 message.to_frame().expect("reconstructable message"),
1580 &mut wire,
1581 )
1582 .expect("encodable message");
1583 }
1584 server.write_all(&wire).await.expect("writable test peer");
1585
1586 let transport = Buffered::new(client);
1587 let mut conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1588 let mut middleware = Middleware::new(0_usize, async |seen: &mut usize, mut message| {
1589 *seen += 1;
1590 if let BackendMessage::ParameterStatus { value, .. } = &mut message {
1591 *value = Bytes::from_static(b"proxy");
1592 }
1593 if let BackendMessage::BackendKeyData {
1594 process_id,
1595 secret_key,
1596 } = &mut message
1597 {
1598 *process_id = 9;
1599 *secret_key = Bytes::from_static(b"new!");
1600 }
1601 if let BackendMessage::ReadyForQuery(status) = &mut message {
1602 *status = crate::codec::TransactionStatus::InTransaction;
1603 }
1604 Ok::<_, Infallible>(message)
1605 });
1606
1607 assert!(matches!(
1608 conn.receive_with_middleware(&mut middleware, &frontend::RuntimeState::Simple)
1609 .await,
1610 Ok(SessionItem::Message(BackendMessage::BackendKeyData { .. }))
1611 ));
1612 assert!(matches!(
1613 conn.receive_with_middleware(&mut middleware, &frontend::RuntimeState::Simple)
1614 .await,
1615 Ok(SessionItem::ReadyForQuery { .. })
1616 ));
1617 assert_eq!(*middleware.state(), 3);
1618 assert_eq!(
1619 conn.parameters().get(b"application_name".as_slice()),
1620 Some(&Bytes::from_static(b"proxy"))
1621 );
1622 assert_eq!(
1623 conn.cancel_key(),
1624 Some(&CancelKey {
1625 process_id: 9,
1626 secret_key: Bytes::from_static(b"new!"),
1627 })
1628 );
1629 assert_eq!(
1630 conn.transaction_status(),
1631 Some(crate::codec::TransactionStatus::InTransaction)
1632 );
1633 conn.into_transport();
1634 }
1635
1636 #[tokio::test]
1637 async fn frontend_middleware_returns_illegal_replacement_before_projection() {
1638 let (proxy, mut client) = tokio::io::duplex(128);
1639 let original = FrontendMessage::CopyData(Bytes::from_static(b"row"));
1640 let mut bytes = BytesMut::new();
1641 PgCodec::<Frontend>::default()
1642 .encode(
1643 original.to_frame().expect("reconstructable message"),
1644 &mut bytes,
1645 )
1646 .expect("encodable message");
1647 client.write_all(&bytes).await.expect("writable client");
1648
1649 let mut conn = Conn::new(Buffered::<_, Frontend>::new_frontend(proxy));
1650 let mut middleware = Middleware::new((), async |_state: &mut (), _message| {
1651 Ok::<_, Infallible>(FrontendMessage::Query(Bytes::from_static(b"select 1")))
1652 });
1653 let result = conn
1654 .receive_frontend_wire_with_middleware(
1655 &mut middleware,
1656 &backend::RuntimeState::SimpleCopyIn,
1657 )
1658 .await;
1659
1660 assert!(matches!(
1661 result,
1662 Err(ReceiveError::Intercept(InterceptError::Invalid(
1663 FrontendMessage::Query(_)
1664 )))
1665 ));
1666 conn.into_transport();
1667 }
1668
1669 #[tokio::test]
1670 async fn middleware_replacement_reaches_the_forwarded_peer() {
1671 let (client_side, mut client) = tokio::io::duplex(128);
1672 let original = FrontendMessage::Query(Bytes::from_static(b"select plaintext"));
1673 let mut bytes = BytesMut::new();
1674 PgCodec::<Frontend>::default()
1675 .encode(
1676 original.to_frame().expect("reconstructable message"),
1677 &mut bytes,
1678 )
1679 .expect("encodable message");
1680 client.write_all(&bytes).await.expect("writable client");
1681
1682 let mut downstream = Conn::new(Buffered::<_, Frontend>::new_frontend(client_side));
1683 let mut middleware = Middleware::new((), async |_state: &mut (), _message| {
1684 Ok::<_, Infallible>(FrontendMessage::Query(Bytes::from_static(
1685 b"select encrypted",
1686 )))
1687 });
1688 let rewritten = downstream
1689 .receive_frontend_wire_with_middleware(&mut middleware, &backend::RuntimeState::Ready)
1690 .await
1691 .expect("legal rewritten query");
1692
1693 let (upstream_side, mut server) = tokio::io::duplex(128);
1694 let mut upstream = Buffered::<_, Backend>::new(upstream_side);
1695 upstream
1696 .push(rewritten.to_frame().expect("reconstructable replacement"))
1697 .expect("encodable replacement");
1698 upstream.flush().await.expect("writable upstream");
1699
1700 let mut received = BytesMut::new();
1701 server
1702 .read_buf(&mut received)
1703 .await
1704 .expect("readable upstream peer");
1705 assert_eq!(
1706 PgCodec::<Frontend>::default()
1707 .decode(&mut received)
1708 .expect("decodable frame"),
1709 Some(FrontendMessage::Query(Bytes::from_static(
1710 b"select encrypted"
1711 )))
1712 );
1713 downstream.into_transport();
1714 }
1715
1716 #[tokio::test]
1717 async fn pre_startup_middleware_can_replace_with_another_legal_choice() {
1718 let (proxy, mut client) = tokio::io::duplex(128);
1719 client
1720 .write_all(
1721 &PreStartupMessage::SslRequest
1722 .to_packet()
1723 .expect("encodable SSLRequest"),
1724 )
1725 .await
1726 .expect("writable client");
1727
1728 let mut conn = Conn::new(Buffered::<_, Frontend>::new_frontend(proxy));
1729 let replacement = PreStartupMessage::CancelRequest {
1730 process_id: 42,
1731 secret_key: Bytes::from_static(b"key!"),
1732 };
1733 let expected = replacement.clone();
1734 let mut middleware = Middleware::new((), async move |_state: &mut (), _message| {
1735 Ok::<_, Infallible>(replacement.clone())
1736 });
1737
1738 assert_eq!(
1739 conn.receive_pre_startup_wire_with_middleware(
1740 &mut middleware,
1741 &server_pre_startup::RuntimeState::PreStartup,
1742 )
1743 .await
1744 .expect("legal replacement"),
1745 expected
1746 );
1747 conn.into_transport();
1748 }
1749
1750 #[tokio::test]
1751 async fn runtime_checked_receive_rejects_illegal_backend_and_pre_startup_messages() {
1752 let (backend_side, mut server) = tokio::io::duplex(64);
1753 let mut backend_bytes = BytesMut::new();
1754 PgCodec::<Backend>::default()
1755 .encode(
1756 BackendMessage::ParseComplete
1757 .to_frame()
1758 .expect("reconstructable ParseComplete"),
1759 &mut backend_bytes,
1760 )
1761 .expect("encodable ParseComplete");
1762 server
1763 .write_all(&backend_bytes)
1764 .await
1765 .expect("writable server");
1766 let mut backend = Conn::new(Buffered::new(backend_side));
1767 let mut backend_middleware = Middleware::new((), crate::middleware::Identity);
1768
1769 assert!(matches!(
1770 backend
1771 .receive_backend_wire_with_middleware(
1772 &mut backend_middleware,
1773 &frontend::RuntimeState::Ready,
1774 )
1775 .await,
1776 Err(ReceiveError::Intercept(InterceptError::Invalid(
1777 BackendMessage::ParseComplete
1778 )))
1779 ));
1780 backend.into_transport();
1781
1782 let (pre_startup_side, mut client) = tokio::io::duplex(64);
1783 client
1784 .write_all(
1785 &PreStartupMessage::SslRequest
1786 .to_packet()
1787 .expect("encodable SSLRequest"),
1788 )
1789 .await
1790 .expect("writable client");
1791 let mut pre_startup = Conn::new(Buffered::<_, Frontend>::new_frontend(pre_startup_side));
1792 let mut pre_startup_middleware = Middleware::new((), crate::middleware::Identity);
1793
1794 assert!(matches!(
1795 pre_startup
1796 .receive_pre_startup_wire_with_middleware(
1797 &mut pre_startup_middleware,
1798 &server_pre_startup::RuntimeState::SslDecision,
1799 )
1800 .await,
1801 Err(ReceiveError::Intercept(InterceptError::Invalid(
1802 PreStartupMessage::SslRequest
1803 )))
1804 ));
1805 pre_startup.into_transport();
1806 }
1807
1808 #[tokio::test]
1809 async fn client_facing_transport_intercepts_typed_frontend_messages() {
1810 let (proxy, mut client) = tokio::io::duplex(128);
1811 let message = FrontendMessage::Query(Bytes::from_static(b"select plaintext"));
1812 let mut bytes = BytesMut::new();
1813 PgCodec::<Frontend>::default()
1814 .encode(
1815 message.to_frame().expect("reconstructable Query"),
1816 &mut bytes,
1817 )
1818 .expect("encodable Query");
1819 client.write_all(&bytes).await.expect("writable client");
1820
1821 let mut transport = Buffered::<_, Frontend>::new_frontend(proxy);
1822 let mut intercepted = transport.receive_wire().await.expect("decodable Query");
1823 let FrontendMessage::Query(query) = &mut intercepted else {
1824 panic!("unexpected frontend message")
1825 };
1826 *query = Bytes::from_static(b"select encrypted");
1827 assert_eq!(
1828 intercepted,
1829 FrontendMessage::Query(Bytes::from_static(b"select encrypted"))
1830 );
1831 }
1832
1833 #[tokio::test]
1834 async fn client_facing_transport_projects_repeated_pre_startup_choice() {
1835 let (proxy, mut client) = tokio::io::duplex(256);
1836 let ssl = PreStartupMessage::SslRequest
1837 .to_packet()
1838 .expect("encodable SSLRequest");
1839 let startup = PreStartupMessage::Startup(crate::startup::StartupMessage {
1840 version: crate::startup::ProtocolVersion::V3_2,
1841 parameters: std::collections::BTreeMap::from([(
1842 Bytes::from_static(b"user"),
1843 Bytes::from_static(b"postgres"),
1844 )]),
1845 });
1846 let startup_packet = startup.to_packet().expect("encodable StartupMessage");
1847 client.write_all(&ssl).await.expect("writable client");
1848 client
1849 .write_all(&startup_packet)
1850 .await
1851 .expect("writable client");
1852
1853 let mut conn = Conn::new(Buffered::<_, Frontend>::new_frontend(proxy));
1854 let ssl = conn
1855 .receive_pre_startup_wire()
1856 .await
1857 .expect("decodable SSLRequest");
1858 let crate::pre_startup::PreStartupOffer::Ssl(decision) = conn.offer_pre_startup(ssl) else {
1859 panic!("unexpected pre-startup branch")
1860 };
1861 let (mut conn, reply) = decision.reject_ssl();
1862 assert_eq!(reply, b'N');
1863 let message = conn
1864 .receive_pre_startup_wire()
1865 .await
1866 .expect("decodable StartupMessage");
1867 assert_eq!(message, startup);
1868 let crate::pre_startup::PreStartupOffer::Startup { conn, .. } =
1869 conn.offer_pre_startup(message)
1870 else {
1871 panic!("unexpected pre-startup branch")
1872 };
1873 let _transport = conn.into_transport();
1874 }
1875
1876 #[tokio::test]
1877 async fn client_facing_transport_applies_its_pre_startup_limit() {
1878 let (proxy, mut client) = tokio::io::duplex(32);
1879 client
1880 .write_all(&17_u32.to_be_bytes())
1881 .await
1882 .expect("writable client");
1883
1884 let mut transport =
1885 Buffered::<_, Frontend>::with_limits_frontend(proxy, 64, 16).expect("valid limits");
1886 let error = transport
1887 .receive_pre_startup()
1888 .await
1889 .expect_err("declared packet exceeds the configured limit");
1890
1891 assert_eq!(error.kind(), io::ErrorKind::InvalidData);
1892 }
1893
1894 #[tokio::test]
1895 async fn upstream_transport_negotiates_raw_gssenc_reply() {
1896 let (proxy, mut server) = tokio::io::duplex(32);
1897 let mut pending = Conn::new(Buffered::new(proxy)).request_gss();
1898 pending.flush().await.expect("GSSENCRequest is writable");
1899
1900 let mut request = [0_u8; 8];
1901 server
1902 .read_exact(&mut request)
1903 .await
1904 .expect("server receives GSSENCRequest");
1905 assert_eq!(request, gssenc_request_packet());
1906 server
1907 .write_all(b"N")
1908 .await
1909 .expect("server writes decision");
1910
1911 let Negotiation::Rejected(plaintext) = pending
1912 .receive_gss_reply()
1913 .await
1914 .expect("valid GSSENC decision")
1915 else {
1916 panic!("expected plaintext fallback")
1917 };
1918 plaintext.into_transport();
1919 }
1920
1921 #[test]
1922 fn client_facing_transport_buffers_raw_gssenc_decision() {
1923 let conn = Conn::new(Buffered::<_, Frontend>::new_frontend(()));
1924 let crate::pre_startup::PreStartupOffer::Gss(decision) =
1925 conn.offer_pre_startup(PreStartupMessage::GssEncRequest)
1926 else {
1927 panic!("expected GSSENC decision")
1928 };
1929
1930 let handshake = decision.approve_gss();
1931 assert_eq!(handshake.pending_output(), b"S");
1932 handshake.into_transport();
1933
1934 let conn = Conn::new(Buffered::<_, Frontend>::new_frontend(()));
1935 let crate::pre_startup::PreStartupOffer::Gss(decision) =
1936 conn.offer_pre_startup(PreStartupMessage::GssEncRequest)
1937 else {
1938 panic!("expected GSSENC decision")
1939 };
1940 let terminated = decision.reject_gss_with_legacy_error();
1941 assert_eq!(terminated.pending_output(), b"E");
1942 terminated.into_transport();
1943 }
1944
1945 #[test]
1946 fn client_facing_transport_buffers_legacy_ssl_error() {
1947 let conn = Conn::new(Buffered::<_, Frontend>::new_frontend(()));
1948 let crate::pre_startup::PreStartupOffer::Ssl(decision) =
1949 conn.offer_pre_startup(PreStartupMessage::SslRequest)
1950 else {
1951 panic!("expected SSL decision")
1952 };
1953
1954 let terminated = decision.reject_ssl_with_legacy_error();
1955 assert_eq!(terminated.pending_output(), b"E");
1956 terminated.into_transport();
1957 }
1958}