Skip to main content

pg_proto/
transport.rs

1//! Buffered, cancellation-safe outbound transport.
2
3use std::{collections::BTreeMap, io, sync::Arc};
4
5use bytes::{Buf, Bytes, BytesMut};
6use rustls::{
7    ClientConfig, ServerConfig,
8    pki_types::{CertificateDer, ServerName},
9};
10use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
11use tokio_util::codec::{Decoder, Encoder};
12
13use crate::{
14    Conn,
15    auth::TlsServerEndPoint,
16    codec::{Backend, BackendMessage, Direction, Frame, Frontend, FrontendMessage, PgCodec},
17    demux::{
18        CancelKey, Demux, Notification, OrderedAsyncEvent, ParameterStatus, SessionItem,
19        TaggedNotice,
20    },
21    middleware::{
22        AcceptsMessage, ClientRole, MessageMiddleware, Middleware, ReceiveError,
23        ReconstructableMessage as _, ServerRole, TypedMiddleware, TypedPhase, TypedReceiveError,
24    },
25    pre_startup::{
26        AwaitingSslReply, DEFAULT_MAX_PRE_STARTUP_PACKET_LEN, EncryptionReply, Negotiation,
27        PreStartup, PreStartupMessage, ServerSslDecision, SslMode, SslModeNegotiation,
28        TlsHandshake, decode_pre_startup_with_limit, gssenc_request_packet, ssl_request_packet,
29    },
30    tls::{ClientTls, ServerTls},
31};
32
33/// Transport wrapper which retains bytes until each write has completed.
34#[derive(Debug)]
35pub struct Buffered<S, D = Backend> {
36    io: S,
37    outbound: BytesMut,
38    inbound: BytesMut,
39    inbound_codec: PgCodec<D>,
40    max_pre_startup_packet_len: usize,
41    demux: Demux,
42}
43
44impl<S> Buffered<S, Backend> {
45    /// Wraps an upstream-facing transport which receives backend messages.
46    pub fn new(io: S) -> Self {
47        Self {
48            io,
49            outbound: BytesMut::new(),
50            inbound: BytesMut::new(),
51            inbound_codec: PgCodec::default(),
52            max_pre_startup_packet_len: DEFAULT_MAX_PRE_STARTUP_PACKET_LEN,
53            demux: Demux::default(),
54        }
55    }
56
57    /// Creates a backend-facing transport with a bounded tagged-frame size.
58    ///
59    /// # Errors
60    ///
61    /// Returns an error when the limit is outside `PostgreSQL`'s frame range.
62    pub fn with_max_frame_len(io: S, max_frame_len: usize) -> io::Result<Self> {
63        Ok(Self {
64            io,
65            outbound: BytesMut::new(),
66            inbound: BytesMut::new(),
67            inbound_codec: PgCodec::with_max_frame_len(max_frame_len)?,
68            max_pre_startup_packet_len: DEFAULT_MAX_PRE_STARTUP_PACKET_LEN,
69            demux: Demux::default(),
70        })
71    }
72}
73
74impl<S> Buffered<S, Frontend> {
75    /// Wraps a client-facing transport which receives frontend messages.
76    pub fn new_frontend(io: S) -> Self {
77        Self {
78            io,
79            outbound: BytesMut::new(),
80            inbound: BytesMut::new(),
81            inbound_codec: PgCodec::default(),
82            max_pre_startup_packet_len: DEFAULT_MAX_PRE_STARTUP_PACKET_LEN,
83            demux: Demux::default(),
84        }
85    }
86
87    /// Creates a frontend-facing transport with a bounded tagged-frame size.
88    ///
89    /// # Errors
90    ///
91    /// Returns an error when the limit is outside `PostgreSQL`'s frame range.
92    pub fn with_max_frame_len_frontend(io: S, max_frame_len: usize) -> io::Result<Self> {
93        Self::with_limits_frontend(io, max_frame_len, DEFAULT_MAX_PRE_STARTUP_PACKET_LEN)
94    }
95
96    /// Creates a frontend-facing transport with bounded tagged and pre-startup packets.
97    ///
98    /// # Errors
99    ///
100    /// Returns an error when either limit is outside `PostgreSQL`'s framing range.
101    pub fn with_limits_frontend(
102        io: S,
103        max_frame_len: usize,
104        max_pre_startup_packet_len: usize,
105    ) -> io::Result<Self> {
106        if !(8..=i32::MAX as usize).contains(&max_pre_startup_packet_len) {
107            return Err(io::Error::new(
108                io::ErrorKind::InvalidInput,
109                "pre-startup packet limit must be between 8 and i32::MAX bytes",
110            ));
111        }
112        Ok(Self {
113            io,
114            outbound: BytesMut::new(),
115            inbound: BytesMut::new(),
116            inbound_codec: PgCodec::with_max_frame_len(max_frame_len)?,
117            max_pre_startup_packet_len,
118            demux: Demux::default(),
119        })
120    }
121}
122
123impl<S, D> Buffered<S, D> {
124    /// Encodes a frame synchronously into the outbound buffer.
125    ///
126    /// # Errors
127    ///
128    /// Returns an error when the frame is too large to encode.
129    pub fn push(&mut self, frame: Frame) -> io::Result<()> {
130        self.inbound_codec.encode(frame, &mut self.outbound)
131    }
132
133    #[must_use]
134    /// Returns encoded bytes which have not yet been fully written.
135    pub fn pending(&self) -> &[u8] {
136        &self.outbound
137    }
138
139    /// Removes buffering and returns the underlying I/O transport.
140    pub fn into_inner(self) -> S {
141        self.io
142    }
143
144    /// Borrows the underlying I/O transport without disturbing codec buffers.
145    pub const fn get_ref(&self) -> &S {
146        &self.io
147    }
148
149    /// Mutably borrows the underlying I/O transport without disturbing codec buffers.
150    pub const fn get_mut(&mut self) -> &mut S {
151        &mut self.io
152    }
153
154    fn push_raw(&mut self, bytes: &[u8]) {
155        self.outbound.extend_from_slice(bytes);
156    }
157
158    #[must_use]
159    /// Returns the backend asynchronous-message demultiplexer.
160    pub const fn demux(&self) -> &Demux {
161        &self.demux
162    }
163
164    /// Returns mutable access to the backend asynchronous-message demultiplexer.
165    pub const fn demux_mut(&mut self) -> &mut Demux {
166        &mut self.demux
167    }
168}
169
170impl<S, D> Buffered<S, D>
171where
172    S: AsyncRead + AsyncWrite + Unpin,
173{
174    async fn connect_tls(
175        self,
176        server_name: ServerName<'static>,
177        config: Arc<ClientConfig>,
178    ) -> io::Result<Buffered<ClientTls<S>, D>> {
179        if !self.outbound.is_empty() || !self.inbound.is_empty() {
180            return Err(io::Error::new(
181                io::ErrorKind::InvalidInput,
182                "TLS upgrade requires empty plaintext buffers",
183            ));
184        }
185        Ok(Buffered {
186            io: crate::tls::connect(self.io, server_name, config).await?,
187            outbound: self.outbound,
188            inbound: self.inbound,
189            inbound_codec: self.inbound_codec,
190            max_pre_startup_packet_len: self.max_pre_startup_packet_len,
191            demux: self.demux,
192        })
193    }
194
195    async fn accept_tls(
196        self,
197        config: Arc<ServerConfig>,
198        leaf_certificate: CertificateDer<'static>,
199    ) -> io::Result<Buffered<ServerTls<S>, D>> {
200        if !self.outbound.is_empty() || !self.inbound.is_empty() {
201            return Err(io::Error::new(
202                io::ErrorKind::InvalidInput,
203                "TLS upgrade requires empty plaintext buffers",
204            ));
205        }
206        Ok(Buffered {
207            io: crate::tls::accept(self.io, config, &leaf_certificate).await?,
208            outbound: self.outbound,
209            inbound: self.inbound,
210            inbound_codec: self.inbound_codec,
211            max_pre_startup_packet_len: self.max_pre_startup_packet_len,
212            demux: self.demux,
213        })
214    }
215}
216
217impl<S: TlsServerEndPoint, D> TlsServerEndPoint for Buffered<S, D> {
218    fn tls_server_end_point(&self) -> &[u8] {
219        self.io.tls_server_end_point()
220    }
221}
222
223impl<S: AsyncWrite + Unpin, D> Buffered<S, D> {
224    /// Writes all buffered bytes without consuming the connection.
225    ///
226    /// Completed partial writes are removed immediately. If this future is
227    /// cancelled, the connection remains owned by the caller and all unwritten
228    /// bytes remain buffered for the next call.
229    ///
230    /// # Errors
231    ///
232    /// Returns the underlying transport's write error or `WriteZero`.
233    pub async fn flush(&mut self) -> io::Result<()> {
234        while !self.outbound.is_empty() {
235            let written = self.io.write(&self.outbound).await?;
236            if written == 0 {
237                return Err(io::Error::new(
238                    io::ErrorKind::WriteZero,
239                    "transport wrote zero buffered bytes",
240                ));
241            }
242            self.outbound.advance(written);
243        }
244        self.io.flush().await
245    }
246}
247
248impl<S: AsyncRead + Unpin, D: Direction> Buffered<S, D> {
249    /// Receives one typed message in this transport's inbound direction.
250    ///
251    /// # Errors
252    ///
253    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
254    pub async fn receive_wire(&mut self) -> io::Result<D::Message> {
255        loop {
256            if let Some(message) = self.inbound_codec.decode(&mut self.inbound)? {
257                return Ok(message);
258            }
259            if self.io.read_buf(&mut self.inbound).await? == 0 {
260                return Err(io::Error::new(
261                    io::ErrorKind::UnexpectedEof,
262                    "peer closed with no complete message",
263                ));
264            }
265        }
266    }
267}
268
269impl<S: AsyncRead + Unpin> Buffered<S, Backend> {
270    async fn receive_encryption_reply(&mut self) -> io::Result<EncryptionReply> {
271        let byte = self.io.read_u8().await?;
272        EncryptionReply::try_from(byte)
273            .map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid encryption reply"))
274    }
275}
276
277impl<S: AsyncRead + Unpin> Buffered<S, Frontend> {
278    /// Receives one raw first packet before tagged frontend framing begins.
279    ///
280    /// # Errors
281    ///
282    /// Returns malformed pre-startup data and underlying transport read errors.
283    pub async fn receive_pre_startup(&mut self) -> io::Result<PreStartupMessage> {
284        loop {
285            if let Some(message) =
286                decode_pre_startup_with_limit(&mut self.inbound, self.max_pre_startup_packet_len)?
287            {
288                return Ok(message);
289            }
290            if self.io.read_buf(&mut self.inbound).await? == 0 {
291                return Err(io::Error::new(
292                    io::ErrorKind::UnexpectedEof,
293                    "client closed with no complete pre-startup packet",
294                ));
295            }
296        }
297    }
298}
299
300impl<S: AsyncRead + Unpin> Buffered<S, Backend> {
301    /// Receives one decoded backend message while retaining partial input.
302    ///
303    /// # Errors
304    ///
305    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
306    pub async fn receive_backend(&mut self) -> io::Result<BackendMessage> {
307        self.receive_wire().await
308    }
309
310    /// Receives the next protocol-advancing message through the async demux.
311    ///
312    /// # Errors
313    ///
314    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
315    pub async fn receive_session(&mut self) -> io::Result<SessionItem> {
316        loop {
317            let message = self.receive_backend().await?;
318            if let Some(item) = self.project_backend(message) {
319                return Ok(item);
320            }
321        }
322    }
323
324    /// Projects an inspected or modified backend message into the session stream.
325    pub fn project_backend(&mut self, message: BackendMessage) -> Option<SessionItem> {
326        self.demux.route(message)
327    }
328}
329
330impl<S, D, Phase, Cleanliness> Conn<Buffered<S, D>, Phase, Cleanliness> {
331    /// Adds an already-typed message to this connection's outbound buffer.
332    ///
333    /// # Errors
334    ///
335    /// Returns an error when the frame is too large to encode.
336    pub fn push_frame(&mut self, frame: Frame) -> io::Result<()> {
337        self.transport_mut().push(frame)
338    }
339
340    #[must_use]
341    /// Returns encoded output which has not yet been flushed.
342    pub fn pending_output(&self) -> &[u8] {
343        self.transport().pending()
344    }
345}
346
347impl<S, Cleanliness> Conn<Buffered<S, Backend>, PreStartup, Cleanliness> {
348    /// Buffers an `SSLRequest` and enters the raw single-byte reply phase.
349    pub fn request_ssl(mut self) -> Conn<Buffered<S, Backend>, AwaitingSslReply, Cleanliness> {
350        self.transport_mut().push_raw(&ssl_request_packet());
351        self.transition()
352    }
353
354    /// Buffers a `GSSENCRequest` and enters the raw single-byte reply phase.
355    pub fn request_gss(
356        mut self,
357    ) -> Conn<Buffered<S, Backend>, crate::pre_startup::AwaitingGssReply, Cleanliness> {
358        self.transport_mut().push_raw(&gssenc_request_packet());
359        self.transition()
360    }
361}
362
363impl<S, Cleanliness> Conn<Buffered<S, Frontend>, ServerSslDecision, Cleanliness> {
364    /// Buffers the server's raw `S` response and enters the TLS handshake phase.
365    pub fn approve_ssl(mut self) -> Conn<Buffered<S, Frontend>, TlsHandshake, Cleanliness> {
366        self.transport_mut().push_raw(b"S");
367        self.transition()
368    }
369
370    /// Buffers the server's raw `N` response and returns to pre-startup choice.
371    pub fn decline_ssl(mut self) -> Conn<Buffered<S, Frontend>, PreStartup, Cleanliness> {
372        self.transport_mut().push_raw(b"N");
373        self.transition()
374    }
375
376    /// Buffers the historical raw `E` response and terminates negotiation.
377    pub fn reject_ssl_with_legacy_error(
378        mut self,
379    ) -> Conn<Buffered<S, Frontend>, crate::pre_startup::Terminated, Cleanliness> {
380        self.transport_mut().push_raw(b"E");
381        self.transition()
382    }
383}
384
385impl<S, Cleanliness>
386    Conn<Buffered<S, Frontend>, crate::pre_startup::ServerGssDecision, Cleanliness>
387{
388    /// Buffers the server's raw `S` response and enters the GSS handshake phase.
389    pub fn approve_gss(
390        mut self,
391    ) -> Conn<Buffered<S, Frontend>, crate::pre_startup::GssHandshake, Cleanliness> {
392        self.transport_mut().push_raw(b"S");
393        self.transition()
394    }
395
396    /// Buffers the server's raw `N` response and returns to pre-startup choice.
397    pub fn decline_gss(mut self) -> Conn<Buffered<S, Frontend>, PreStartup, Cleanliness> {
398        self.transport_mut().push_raw(b"N");
399        self.transition()
400    }
401
402    /// Buffers the historical raw `E` response and terminates negotiation.
403    pub fn reject_gss_with_legacy_error(
404        mut self,
405    ) -> Conn<Buffered<S, Frontend>, crate::pre_startup::Terminated, Cleanliness> {
406        self.transport_mut().push_raw(b"E");
407        self.transition()
408    }
409}
410
411impl<S: AsyncRead + Unpin, Cleanliness> Conn<Buffered<S, Backend>, AwaitingSslReply, Cleanliness> {
412    /// Receives and projects the server's raw SSL decision byte.
413    ///
414    /// # Errors
415    ///
416    /// Returns an I/O error or rejects a byte other than `S`, `N`, or `E`.
417    pub async fn receive_ssl_reply(
418        mut self,
419    ) -> io::Result<Negotiation<Buffered<S, Backend>, TlsHandshake, Cleanliness>> {
420        let reply = self.transport_mut().receive_encryption_reply().await?;
421        Ok(match reply {
422            EncryptionReply::Accepted => Negotiation::Accepted(self.transition()),
423            EncryptionReply::Rejected => Negotiation::Rejected(self.transition()),
424            EncryptionReply::LegacyError => Negotiation::LegacyError(self.transition()),
425        })
426    }
427
428    /// Receives the server decision and enforces the selected plaintext fallback policy.
429    ///
430    /// # Errors
431    ///
432    /// Returns an I/O error or rejects a byte other than `S`, `N`, or `E`.
433    pub async fn receive_ssl_reply_for_mode(
434        mut self,
435        mode: SslMode,
436    ) -> io::Result<SslModeNegotiation<Buffered<S, Backend>, Cleanliness>> {
437        let reply = self.transport_mut().receive_encryption_reply().await?;
438        Ok(self.apply_ssl_reply(reply, mode))
439    }
440}
441
442impl<S: AsyncRead + Unpin, Cleanliness>
443    Conn<Buffered<S, Backend>, crate::pre_startup::AwaitingGssReply, Cleanliness>
444{
445    /// Receives and projects the server's raw GSSENC decision byte.
446    ///
447    /// # Errors
448    ///
449    /// Returns an I/O error or rejects a byte other than `S`, `N`, or `E`.
450    pub async fn receive_gss_reply(
451        mut self,
452    ) -> io::Result<Negotiation<Buffered<S, Backend>, crate::pre_startup::GssHandshake, Cleanliness>>
453    {
454        let reply = self.transport_mut().receive_encryption_reply().await?;
455        Ok(match reply {
456            EncryptionReply::Accepted => Negotiation::Accepted(self.transition()),
457            EncryptionReply::Rejected => Negotiation::Rejected(self.transition()),
458            EncryptionReply::LegacyError => Negotiation::LegacyError(self.transition()),
459        })
460    }
461}
462
463impl<S, Cleanliness> Conn<Buffered<S, Backend>, TlsHandshake, Cleanliness>
464where
465    S: AsyncRead + AsyncWrite + Unpin,
466{
467    /// Completes a client-side TLS handshake and changes the transport type.
468    ///
469    /// # Errors
470    ///
471    /// Returns a TLS handshake, certificate, channel-binding, or buffer-state error.
472    pub async fn connect_tls(
473        self,
474        server_name: ServerName<'static>,
475        config: Arc<ClientConfig>,
476    ) -> io::Result<Conn<Buffered<ClientTls<S>, Backend>, PreStartup, Cleanliness>> {
477        let transport = self.into_transport();
478        Ok(Conn::new(transport.connect_tls(server_name, config).await?)
479            .transition::<PreStartup, Cleanliness>())
480    }
481}
482
483impl<S, Cleanliness> Conn<Buffered<S, Frontend>, TlsHandshake, Cleanliness>
484where
485    S: AsyncRead + AsyncWrite + Unpin,
486{
487    /// Completes a server-side TLS handshake and changes the transport type.
488    ///
489    /// # Errors
490    ///
491    /// Returns a TLS handshake, certificate, channel-binding, or buffer-state error.
492    pub async fn accept_tls(
493        self,
494        config: Arc<ServerConfig>,
495        leaf_certificate: CertificateDer<'static>,
496    ) -> io::Result<Conn<Buffered<ServerTls<S>, Frontend>, PreStartup, Cleanliness>> {
497        let transport = self.into_transport();
498        Ok(
499            Conn::new(transport.accept_tls(config, leaf_certificate).await?)
500                .transition::<PreStartup, Cleanliness>(),
501        )
502    }
503}
504
505impl<S, D, Cleanliness> Conn<Buffered<S, D>, crate::pre_startup::Startup, Cleanliness> {
506    /// Buffers the raw, untagged startup packet before normal framing begins.
507    pub fn push_startup_packet(&mut self, packet: &[u8]) {
508        self.transport_mut().outbound.extend_from_slice(packet);
509    }
510}
511
512impl<S: AsyncWrite + Unpin, D, Phase, Cleanliness> Conn<Buffered<S, D>, Phase, Cleanliness> {
513    /// Flushes buffered output while retaining ownership of the typed connection.
514    ///
515    /// # Errors
516    ///
517    /// Returns an error from the underlying transport.
518    pub async fn flush(&mut self) -> io::Result<()> {
519        self.transport_mut().flush().await
520    }
521}
522
523impl<S: AsyncRead + Unpin, Phase, Cleanliness> Conn<Buffered<S, Backend>, Phase, Cleanliness> {
524    /// Receives one backend message before demultiplexing or state advancement.
525    /// This is the interception point for proxy policy and message rewriting.
526    ///
527    /// # Errors
528    ///
529    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
530    pub async fn receive_backend_wire(&mut self) -> io::Result<BackendMessage> {
531        self.transport_mut().receive_backend().await
532    }
533
534    /// Receives one backend message through middleware indexed by this connection phase.
535    ///
536    /// Unlike [`Self::receive_backend_wire_with_middleware`], callers do not pass
537    /// a runtime protocol state. `Phase` selects the generated legal message set
538    /// and the server sender role at compile time.
539    ///
540    /// # Errors
541    ///
542    /// Returns an I/O or decoding error, an illegal peer message, a middleware
543    /// policy error, or a phase-legal replacement with an invalid wire shape.
544    pub async fn receive_backend_typed<State, Handler>(
545        &mut self,
546        middleware: &mut Middleware<State, Handler>,
547    ) -> Result<
548        <Phase as TypedPhase<ServerRole, BackendMessage>>::Message,
549        TypedReceiveError<Handler::Error, BackendMessage>,
550    >
551    where
552        Phase: TypedPhase<ServerRole, BackendMessage>,
553        Handler: TypedMiddleware<
554                ServerRole,
555                <Phase as TypedPhase<ServerRole, BackendMessage>>::ProtocolPhase,
556                <Phase as TypedPhase<ServerRole, BackendMessage>>::Message,
557                State,
558            >,
559    {
560        let message = self
561            .receive_backend_wire()
562            .await
563            .map_err(TypedReceiveError::Io)?;
564        let message = <Phase as TypedPhase<ServerRole, BackendMessage>>::Message::try_from(message)
565            .map_err(TypedReceiveError::Illegal)?;
566        let message = middleware
567            .intercept_typed::<
568                ServerRole,
569                <Phase as TypedPhase<ServerRole, BackendMessage>>::ProtocolPhase,
570                _,
571            >(message)
572            .await
573            .map_err(TypedReceiveError::Middleware)?;
574        if message.as_ref().is_reconstructable() {
575            Ok(message)
576        } else {
577            Err(TypedReceiveError::InvalidWire(message.into()))
578        }
579    }
580
581    /// Receives typed backend traffic until one protocol-advancing item remains.
582    ///
583    /// Asynchronous messages pass through the same middleware, are recorded by
584    /// the demultiplexer in wire order, and leave `Phase` unchanged.
585    ///
586    /// # Errors
587    ///
588    /// Returns the same failures as [`Self::receive_backend_typed`].
589    pub async fn receive_typed<State, Handler>(
590        &mut self,
591        middleware: &mut Middleware<State, Handler>,
592    ) -> Result<SessionItem, TypedReceiveError<Handler::Error, BackendMessage>>
593    where
594        Phase: TypedPhase<ServerRole, BackendMessage>,
595        Handler: TypedMiddleware<
596                ServerRole,
597                <Phase as TypedPhase<ServerRole, BackendMessage>>::ProtocolPhase,
598                <Phase as TypedPhase<ServerRole, BackendMessage>>::Message,
599                State,
600            >,
601    {
602        loop {
603            let message = self.receive_backend_typed(middleware).await?;
604            if let Some(item) = self.project_backend(message.into()) {
605                return Ok(item);
606            }
607        }
608    }
609
610    /// Receives one SSL or GSSENC decision through phase-typed middleware.
611    ///
612    /// `Phase` must be an encryption-reply phase; callers then consume the
613    /// connection with its existing typed `receive_reply` projection.
614    ///
615    /// # Errors
616    ///
617    /// Returns an I/O error, illegal decision, middleware policy error, or an
618    /// invalid replacement wire shape.
619    pub async fn receive_encryption_reply_typed<State, Handler>(
620        &mut self,
621        middleware: &mut Middleware<State, Handler>,
622    ) -> Result<
623        <Phase as TypedPhase<ServerRole, EncryptionReply>>::Message,
624        TypedReceiveError<Handler::Error, EncryptionReply>,
625    >
626    where
627        Phase: TypedPhase<ServerRole, EncryptionReply>,
628        Handler: TypedMiddleware<
629                ServerRole,
630                <Phase as TypedPhase<ServerRole, EncryptionReply>>::ProtocolPhase,
631                <Phase as TypedPhase<ServerRole, EncryptionReply>>::Message,
632                State,
633            >,
634    {
635        let message = self
636            .transport_mut()
637            .receive_encryption_reply()
638            .await
639            .map_err(TypedReceiveError::Io)?;
640        let message =
641            <Phase as TypedPhase<ServerRole, EncryptionReply>>::Message::try_from(message)
642                .map_err(TypedReceiveError::Illegal)?;
643        let message = middleware
644            .intercept_typed::<
645                ServerRole,
646                <Phase as TypedPhase<ServerRole, EncryptionReply>>::ProtocolPhase,
647                _,
648            >(message)
649            .await
650            .map_err(TypedReceiveError::Middleware)?;
651        if message.as_ref().is_reconstructable() {
652            Ok(message)
653        } else {
654            Err(TypedReceiveError::InvalidWire(message.into()))
655        }
656    }
657
658    /// Receives, intercepts, and validates one backend message before projection.
659    ///
660    /// # Errors
661    ///
662    /// Returns an I/O, decoding, middleware-policy, or state-validation error.
663    pub async fn receive_backend_wire_with_middleware<State, Handler, ProtocolState>(
664        &mut self,
665        middleware: &mut Middleware<State, Handler>,
666        protocol_state: &ProtocolState,
667    ) -> Result<BackendMessage, ReceiveError<Handler::Error, BackendMessage>>
668    where
669        Handler: MessageMiddleware<BackendMessage, State>,
670        ProtocolState: AcceptsMessage<BackendMessage>,
671    {
672        let message = self
673            .receive_backend_wire()
674            .await
675            .map_err(ReceiveError::Io)?;
676        middleware
677            .intercept_checked(protocol_state, message)
678            .await
679            .map_err(ReceiveError::Intercept)
680    }
681
682    /// Projects an inspected or modified message into the filtered session stream.
683    pub fn project_backend(&mut self, message: BackendMessage) -> Option<SessionItem> {
684        self.transport_mut().project_backend(message)
685    }
686
687    /// Receives the next message in the filtered session projection.
688    ///
689    /// # Errors
690    ///
691    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
692    pub async fn receive(&mut self) -> io::Result<SessionItem> {
693        self.transport_mut().receive_session().await
694    }
695
696    /// Receives backend messages through middleware before demultiplexing.
697    ///
698    /// Asynchronous messages are intercepted and then recorded by the demux;
699    /// this method continues until a protocol-advancing item is available.
700    ///
701    /// # Errors
702    ///
703    /// Returns an I/O, decoding, middleware-policy, or state-validation error.
704    pub async fn receive_with_middleware<State, Handler, ProtocolState>(
705        &mut self,
706        middleware: &mut Middleware<State, Handler>,
707        protocol_state: &ProtocolState,
708    ) -> Result<SessionItem, ReceiveError<Handler::Error, BackendMessage>>
709    where
710        Handler: MessageMiddleware<BackendMessage, State>,
711        ProtocolState: AcceptsMessage<BackendMessage>,
712    {
713        loop {
714            let message = self
715                .receive_backend_wire_with_middleware(middleware, protocol_state)
716                .await?;
717            if let Some(item) = self.project_backend(message) {
718                return Ok(item);
719            }
720        }
721    }
722
723    #[must_use]
724    /// Returns the latest upstream cancellation key observed during startup.
725    pub fn cancel_key(&self) -> Option<&CancelKey> {
726        self.transport().demux().cancel_key()
727    }
728
729    /// Returns the latest backend parameter values observed by the demux.
730    #[must_use]
731    pub fn parameters(&self) -> &BTreeMap<Bytes, Bytes> {
732        self.transport().demux().parameters()
733    }
734
735    /// Returns whether current parameters differ from the startup baseline.
736    #[must_use]
737    pub fn parameters_changed(&self) -> bool {
738        self.transport().demux().parameters_changed()
739    }
740
741    /// Returns the latest transaction status observed in `ReadyForQuery`.
742    #[must_use]
743    pub fn transaction_status(&self) -> Option<crate::codec::TransactionStatus> {
744        self.transport().demux().transaction_status()
745    }
746
747    /// Removes the oldest queued asynchronous notification.
748    pub fn pop_notification(&mut self) -> Option<Notification> {
749        self.transport_mut().demux_mut().pop_notification()
750    }
751
752    /// Removes the next tagged notice for prompt forwarding to the client.
753    pub fn pop_notice(&mut self) -> Option<TaggedNotice> {
754        self.transport_mut().demux_mut().pop_notice()
755    }
756
757    /// Removes the next ordered parameter update for forwarding to the client.
758    pub fn pop_parameter_status(&mut self) -> Option<ParameterStatus> {
759        self.transport_mut().demux_mut().pop_parameter_status()
760    }
761
762    /// Removes the next independent backend event in original wire order.
763    pub fn pop_async_event(&mut self) -> Option<OrderedAsyncEvent> {
764        self.transport_mut().demux_mut().pop_async_event()
765    }
766}
767
768impl<S: AsyncRead + Unpin, Phase, Cleanliness> Conn<Buffered<S, Frontend>, Phase, Cleanliness> {
769    /// Receives one frontend message before any server-role state advancement.
770    ///
771    /// # Errors
772    ///
773    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
774    pub async fn receive_frontend_wire(&mut self) -> io::Result<FrontendMessage> {
775        self.transport_mut().receive_wire().await
776    }
777
778    /// Receives one frontend message through middleware indexed by this connection phase.
779    ///
780    /// `Phase` selects the generated legal message set and the client sender role
781    /// at compile time, so no runtime protocol-state argument is accepted.
782    ///
783    /// # Errors
784    ///
785    /// Returns an I/O or decoding error, an illegal peer message, a middleware
786    /// policy error, or a phase-legal replacement with an invalid wire shape.
787    pub async fn receive_frontend_typed<State, Handler>(
788        &mut self,
789        middleware: &mut Middleware<State, Handler>,
790    ) -> Result<
791        <Phase as TypedPhase<ClientRole, FrontendMessage>>::Message,
792        TypedReceiveError<Handler::Error, FrontendMessage>,
793    >
794    where
795        Phase: TypedPhase<ClientRole, FrontendMessage>,
796        Handler: TypedMiddleware<
797                ClientRole,
798                <Phase as TypedPhase<ClientRole, FrontendMessage>>::ProtocolPhase,
799                <Phase as TypedPhase<ClientRole, FrontendMessage>>::Message,
800                State,
801            >,
802    {
803        let message = self
804            .receive_frontend_wire()
805            .await
806            .map_err(TypedReceiveError::Io)?;
807        let message =
808            <Phase as TypedPhase<ClientRole, FrontendMessage>>::Message::try_from(message)
809                .map_err(TypedReceiveError::Illegal)?;
810        let message = middleware
811            .intercept_typed::<
812                ClientRole,
813                <Phase as TypedPhase<ClientRole, FrontendMessage>>::ProtocolPhase,
814                _,
815            >(message)
816            .await
817            .map_err(TypedReceiveError::Middleware)?;
818        if message.as_ref().is_reconstructable() {
819            Ok(message)
820        } else {
821            Err(TypedReceiveError::InvalidWire(message.into()))
822        }
823    }
824
825    /// Receives, intercepts, and validates one frontend message before projection.
826    ///
827    /// # Errors
828    ///
829    /// Returns an I/O, decoding, middleware-policy, or state-validation error.
830    pub async fn receive_frontend_wire_with_middleware<State, Handler, ProtocolState>(
831        &mut self,
832        middleware: &mut Middleware<State, Handler>,
833        protocol_state: &ProtocolState,
834    ) -> Result<FrontendMessage, ReceiveError<Handler::Error, FrontendMessage>>
835    where
836        Handler: MessageMiddleware<FrontendMessage, State>,
837        ProtocolState: AcceptsMessage<FrontendMessage>,
838    {
839        let message = self
840            .receive_frontend_wire()
841            .await
842            .map_err(ReceiveError::Io)?;
843        middleware
844            .intercept_checked(protocol_state, message)
845            .await
846            .map_err(ReceiveError::Intercept)
847    }
848}
849
850impl<S: AsyncRead + Unpin, Cleanliness> Conn<Buffered<S, Frontend>, PreStartup, Cleanliness> {
851    /// Receives a raw pre-startup packet before server-role state projection.
852    ///
853    /// # Errors
854    ///
855    /// Returns malformed pre-startup data and underlying transport read errors.
856    pub async fn receive_pre_startup_wire(&mut self) -> io::Result<PreStartupMessage> {
857        self.transport_mut().receive_pre_startup().await
858    }
859
860    /// Receives a client pre-startup packet through phase-typed middleware.
861    ///
862    /// # Errors
863    ///
864    /// Returns an I/O or decoding error, an illegal pre-startup packet, a
865    /// middleware policy error, or an invalid replacement wire shape.
866    pub async fn receive_pre_startup_typed<State, Handler>(
867        &mut self,
868        middleware: &mut Middleware<State, Handler>,
869    ) -> Result<
870        <PreStartup as TypedPhase<ClientRole, PreStartupMessage>>::Message,
871        TypedReceiveError<Handler::Error, PreStartupMessage>,
872    >
873    where
874        Handler: TypedMiddleware<
875                ClientRole,
876                <PreStartup as TypedPhase<ClientRole, PreStartupMessage>>::ProtocolPhase,
877                <PreStartup as TypedPhase<ClientRole, PreStartupMessage>>::Message,
878                State,
879            >,
880    {
881        let message = self
882            .receive_pre_startup_wire()
883            .await
884            .map_err(TypedReceiveError::Io)?;
885        let message =
886            <PreStartup as TypedPhase<ClientRole, PreStartupMessage>>::Message::try_from(message)
887                .map_err(TypedReceiveError::Illegal)?;
888        let message = middleware
889            .intercept_typed::<
890                ClientRole,
891                <PreStartup as TypedPhase<ClientRole, PreStartupMessage>>::ProtocolPhase,
892                _,
893            >(message)
894            .await
895            .map_err(TypedReceiveError::Middleware)?;
896        if message.as_ref().is_reconstructable() {
897            Ok(message)
898        } else {
899            Err(TypedReceiveError::InvalidWire(message.into()))
900        }
901    }
902
903    /// Receives, intercepts, and validates one untagged pre-startup message.
904    ///
905    /// # Errors
906    ///
907    /// Returns an I/O, decoding, middleware-policy, or state-validation error.
908    pub async fn receive_pre_startup_wire_with_middleware<State, Handler, ProtocolState>(
909        &mut self,
910        middleware: &mut Middleware<State, Handler>,
911        protocol_state: &ProtocolState,
912    ) -> Result<PreStartupMessage, ReceiveError<Handler::Error, PreStartupMessage>>
913    where
914        Handler: MessageMiddleware<PreStartupMessage, State>,
915        ProtocolState: AcceptsMessage<PreStartupMessage>,
916    {
917        let message = self
918            .receive_pre_startup_wire()
919            .await
920            .map_err(ReceiveError::Io)?;
921        middleware
922            .intercept_checked(protocol_state, message)
923            .await
924            .map_err(ReceiveError::Intercept)
925    }
926}
927
928#[cfg(test)]
929mod tests {
930    use std::{
931        convert::Infallible,
932        future::Future,
933        pin::Pin,
934        task::{Context, Poll},
935    };
936
937    use bytes::Bytes;
938    use tokio::io::AsyncWrite;
939
940    use super::*;
941    use crate::{
942        grammar::{backend, frontend, server_pre_startup},
943        middleware::{InterceptError, Middleware, ReceiveError},
944    };
945
946    #[derive(Debug, Default)]
947    struct ShortWriter {
948        output: Vec<u8>,
949    }
950
951    impl AsyncWrite for ShortWriter {
952        fn poll_write(
953            mut self: Pin<&mut Self>,
954            _cx: &mut Context<'_>,
955            buffer: &[u8],
956        ) -> Poll<io::Result<usize>> {
957            let written = buffer.len().min(2);
958            self.output.extend_from_slice(&buffer[..written]);
959            Poll::Ready(Ok(written))
960        }
961
962        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
963            Poll::Ready(Ok(()))
964        }
965
966        fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
967            Poll::Ready(Ok(()))
968        }
969    }
970
971    #[tokio::test]
972    async fn flush_handles_partial_writes_without_losing_bytes() {
973        let frame = Frame {
974            tag: b'S',
975            body: Bytes::new(),
976        };
977        let mut transport = Buffered::new(ShortWriter::default());
978        transport.push(frame).expect("encodable frame");
979        assert_eq!(transport.pending(), &[b'S', 0, 0, 0, 4]);
980        transport.flush().await.expect("writable transport");
981        assert!(transport.pending().is_empty());
982        assert_eq!(transport.into_inner().output, [b'S', 0, 0, 0, 4]);
983    }
984
985    #[test]
986    fn buffered_transport_enforces_its_frame_limit_on_output() {
987        let mut transport = Buffered::<_, Backend>::with_max_frame_len((), 9).unwrap();
988        let error = transport
989            .push(Frame {
990                tag: b'Q',
991                body: Bytes::from_static(b"12345"),
992            })
993            .unwrap_err();
994        assert_eq!(error.kind(), io::ErrorKind::InvalidInput);
995        assert!(transport.pending().is_empty());
996    }
997
998    #[test]
999    fn cancelling_flush_retains_unwritten_bytes() {
1000        #[derive(Debug, Default)]
1001        struct PausingWriter {
1002            output: Vec<u8>,
1003            blocked: bool,
1004        }
1005
1006        impl AsyncWrite for PausingWriter {
1007            fn poll_write(
1008                mut self: Pin<&mut Self>,
1009                _cx: &mut Context<'_>,
1010                buffer: &[u8],
1011            ) -> Poll<io::Result<usize>> {
1012                if self.blocked {
1013                    return Poll::Pending;
1014                }
1015                let written = buffer.len().min(2);
1016                self.output.extend_from_slice(&buffer[..written]);
1017                self.blocked = true;
1018                Poll::Ready(Ok(written))
1019            }
1020
1021            fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
1022                Poll::Ready(Ok(()))
1023            }
1024
1025            fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
1026                Poll::Ready(Ok(()))
1027            }
1028        }
1029
1030        let mut transport = Buffered::new(PausingWriter::default());
1031        transport
1032            .push(Frame {
1033                tag: b'S',
1034                body: Bytes::new(),
1035            })
1036            .expect("encodable frame");
1037
1038        let mut flush = Box::pin(transport.flush());
1039        let waker = std::task::Waker::noop();
1040        let mut context = Context::from_waker(waker);
1041        assert!(flush.as_mut().poll(&mut context).is_pending());
1042        drop(flush);
1043
1044        assert_eq!(transport.pending(), &[0, 0, 4]);
1045        assert_eq!(transport.io.output, [b'S', 0]);
1046    }
1047
1048    #[tokio::test]
1049    async fn receive_filters_parameter_status_before_session_message() {
1050        let (client, mut server) = tokio::io::duplex(256);
1051        let mut wire = BytesMut::new();
1052        let mut encoder = PgCodec::<Backend>::default();
1053        encoder
1054            .encode(
1055                Frame {
1056                    tag: b'S',
1057                    body: Bytes::from_static(b"client_encoding\0UTF8\0"),
1058                },
1059                &mut wire,
1060            )
1061            .expect("encodable ParameterStatus");
1062        encoder
1063            .encode(
1064                Frame {
1065                    tag: b'Z',
1066                    body: Bytes::from_static(b"I"),
1067                },
1068                &mut wire,
1069            )
1070            .expect("encodable ReadyForQuery");
1071        server.write_all(&wire).await.expect("writable test peer");
1072
1073        let mut transport = Buffered::new(client);
1074        assert_eq!(
1075            transport.receive_session().await.expect("valid messages"),
1076            SessionItem::ReadyForQuery {
1077                status: crate::codec::TransactionStatus::Idle,
1078                parameters_changed: false,
1079            }
1080        );
1081        assert_eq!(
1082            transport
1083                .demux()
1084                .parameters()
1085                .get(&Bytes::from_static(b"client_encoding")),
1086            Some(&Bytes::from_static(b"UTF8"))
1087        );
1088        let conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1089        assert_eq!(
1090            conn.parameters().get(b"client_encoding".as_slice()),
1091            Some(&Bytes::from_static(b"UTF8"))
1092        );
1093        assert!(!conn.parameters_changed());
1094        assert_eq!(
1095            conn.transaction_status(),
1096            Some(crate::codec::TransactionStatus::Idle)
1097        );
1098        conn.into_transport();
1099    }
1100
1101    #[tokio::test]
1102    async fn wire_message_can_be_modified_before_projection() {
1103        let (client, mut server) = tokio::io::duplex(128);
1104        let original = BackendMessage::ParameterStatus {
1105            name: Bytes::from_static(b"application_name"),
1106            value: Bytes::from_static(b"upstream"),
1107        };
1108        let mut bytes = BytesMut::new();
1109        PgCodec::<Backend>::default()
1110            .encode(
1111                original.to_frame().expect("reconstructable message"),
1112                &mut bytes,
1113            )
1114            .expect("encodable message");
1115        server.write_all(&bytes).await.expect("writable test peer");
1116
1117        let mut transport = Buffered::new(client);
1118        let mut message = transport
1119            .receive_backend()
1120            .await
1121            .expect("decodable message");
1122        let BackendMessage::ParameterStatus { value, .. } = &mut message else {
1123            panic!("unexpected message")
1124        };
1125        *value = Bytes::from_static(b"proxy");
1126        assert!(transport.project_backend(message).is_none());
1127        assert_eq!(
1128            transport
1129                .demux()
1130                .parameters()
1131                .get(&Bytes::from_static(b"application_name")),
1132            Some(&Bytes::from_static(b"proxy"))
1133        );
1134    }
1135
1136    #[tokio::test]
1137    async fn typed_middleware_accepts_async_traffic_without_advancing_ready() {
1138        let (client, mut server) = tokio::io::duplex(128);
1139        let original = BackendMessage::ParameterStatus {
1140            name: Bytes::from_static(b"application_name"),
1141            value: Bytes::from_static(b"upstream"),
1142        };
1143        let mut bytes = BytesMut::new();
1144        PgCodec::<Backend>::default()
1145            .encode(
1146                original.to_frame().expect("reconstructable message"),
1147                &mut bytes,
1148            )
1149            .expect("encodable message");
1150        server.write_all(&bytes).await.expect("writable test peer");
1151
1152        let transport = Buffered::new(client);
1153        let mut conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1154        let mut middleware = Middleware::new(0_usize, async |seen: &mut usize, _message| {
1155            *seen += 1;
1156            let replacement = BackendMessage::ParameterStatus {
1157                name: Bytes::from_static(b"application_name"),
1158                value: Bytes::from_static(b"proxy"),
1159            };
1160            match crate::middleware::TypedBackendMessage::try_from(replacement) {
1161                Ok(replacement) => Ok::<_, Infallible>(replacement),
1162                Err(message) => panic!("parameter status must be asynchronous: {message:?}"),
1163            }
1164        });
1165
1166        let message = conn
1167            .receive_backend_typed(&mut middleware)
1168            .await
1169            .expect("typed asynchronous message");
1170        assert!(conn.project_backend(message.into()).is_none());
1171        assert_eq!(*middleware.state(), 1);
1172        assert_eq!(
1173            conn.parameters().get(b"application_name".as_slice()),
1174            Some(&Bytes::from_static(b"proxy"))
1175        );
1176        conn.into_transport();
1177    }
1178
1179    #[tokio::test]
1180    async fn typed_receive_projects_into_the_existing_next_connection_enum() {
1181        let (client, mut server) = tokio::io::duplex(128);
1182        let ready = BackendMessage::ReadyForQuery(crate::codec::TransactionStatus::Idle);
1183        let mut bytes = BytesMut::new();
1184        PgCodec::<Backend>::default()
1185            .encode(
1186                ready.to_frame().expect("reconstructable message"),
1187                &mut bytes,
1188            )
1189            .expect("encodable message");
1190        server.write_all(&bytes).await.expect("writable test peer");
1191
1192        let transport = Buffered::new(client);
1193        let conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1194        let (mut query, _) = conn
1195            .push_stateless_query(b"select 1")
1196            .expect("encodable query");
1197        let mut middleware = Middleware::new((), crate::middleware::Identity);
1198        let item = query
1199            .receive_typed(&mut middleware)
1200            .await
1201            .expect("phase-legal ready message");
1202
1203        let transition = query.offer(item).expect("typed next-state projection");
1204        let crate::session::SimpleTransition::Ready(crate::session::ReadyState::Clean(ready)) =
1205            transition
1206        else {
1207            panic!("idle readiness must return a clean ready connection");
1208        };
1209        ready.into_transport();
1210    }
1211
1212    #[tokio::test]
1213    async fn typed_receive_keeps_wire_shape_validation_at_runtime() {
1214        let (client, mut peer) = tokio::io::duplex(256);
1215        let query = FrontendMessage::Query(Bytes::from_static(b"select 1"));
1216        let mut bytes = BytesMut::new();
1217        PgCodec::<Frontend>::default()
1218            .encode(query.to_frame().expect("reconstructable query"), &mut bytes)
1219            .expect("encodable query");
1220        peer.write_all(&bytes).await.expect("writable test peer");
1221
1222        let transport = Buffered::<_, Frontend>::new_frontend(client);
1223        let mut conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1224        let invalid = FrontendMessage::Parse(crate::codec::Parse {
1225            statement: Bytes::from_static(b"invalid\0statement"),
1226            query: Bytes::from_static(b"select 2"),
1227            parameter_types: Vec::new(),
1228        });
1229        let mut middleware = Middleware::new((), async move |_state: &mut (), _message| {
1230            let invalid = invalid.clone();
1231            match backend::ReadyExternalMessage::try_from(invalid) {
1232                Ok(invalid) => Ok::<_, Infallible>(invalid),
1233                Err(message) => panic!("parse must be protocol-legal while ready: {message:?}"),
1234            }
1235        });
1236
1237        let result = conn.receive_frontend_typed(&mut middleware).await;
1238        assert!(matches!(
1239            result,
1240            Err(TypedReceiveError::InvalidWire(FrontendMessage::Parse(_)))
1241        ));
1242        conn.into_transport();
1243    }
1244
1245    #[tokio::test]
1246    async fn middleware_rewrites_backend_before_demux_bookkeeping() {
1247        let (client, mut server) = tokio::io::duplex(256);
1248        let mut wire = BytesMut::new();
1249        let mut encoder = PgCodec::<Backend>::default();
1250        for message in [
1251            BackendMessage::BackendKeyData {
1252                process_id: 7,
1253                secret_key: Bytes::from_static(b"old!"),
1254            },
1255            BackendMessage::ParameterStatus {
1256                name: Bytes::from_static(b"application_name"),
1257                value: Bytes::from_static(b"upstream"),
1258            },
1259            BackendMessage::ReadyForQuery(crate::codec::TransactionStatus::Idle),
1260        ] {
1261            encoder
1262                .encode(
1263                    message.to_frame().expect("reconstructable message"),
1264                    &mut wire,
1265                )
1266                .expect("encodable message");
1267        }
1268        server.write_all(&wire).await.expect("writable test peer");
1269
1270        let transport = Buffered::new(client);
1271        let mut conn: Conn<_, crate::auth::Ready> = Conn::new(transport).transition();
1272        let mut middleware = Middleware::new(0_usize, async |seen: &mut usize, mut message| {
1273            *seen += 1;
1274            if let BackendMessage::ParameterStatus { value, .. } = &mut message {
1275                *value = Bytes::from_static(b"proxy");
1276            }
1277            if let BackendMessage::BackendKeyData {
1278                process_id,
1279                secret_key,
1280            } = &mut message
1281            {
1282                *process_id = 9;
1283                *secret_key = Bytes::from_static(b"new!");
1284            }
1285            if let BackendMessage::ReadyForQuery(status) = &mut message {
1286                *status = crate::codec::TransactionStatus::InTransaction;
1287            }
1288            Ok::<_, Infallible>(message)
1289        });
1290
1291        assert!(matches!(
1292            conn.receive_with_middleware(&mut middleware, &frontend::RuntimeState::Simple)
1293                .await,
1294            Ok(SessionItem::Message(BackendMessage::BackendKeyData { .. }))
1295        ));
1296        assert!(matches!(
1297            conn.receive_with_middleware(&mut middleware, &frontend::RuntimeState::Simple)
1298                .await,
1299            Ok(SessionItem::ReadyForQuery { .. })
1300        ));
1301        assert_eq!(*middleware.state(), 3);
1302        assert_eq!(
1303            conn.parameters().get(b"application_name".as_slice()),
1304            Some(&Bytes::from_static(b"proxy"))
1305        );
1306        assert_eq!(
1307            conn.cancel_key(),
1308            Some(&CancelKey {
1309                process_id: 9,
1310                secret_key: Bytes::from_static(b"new!"),
1311            })
1312        );
1313        assert_eq!(
1314            conn.transaction_status(),
1315            Some(crate::codec::TransactionStatus::InTransaction)
1316        );
1317        conn.into_transport();
1318    }
1319
1320    #[tokio::test]
1321    async fn frontend_middleware_returns_illegal_replacement_before_projection() {
1322        let (proxy, mut client) = tokio::io::duplex(128);
1323        let original = FrontendMessage::CopyData(Bytes::from_static(b"row"));
1324        let mut bytes = BytesMut::new();
1325        PgCodec::<Frontend>::default()
1326            .encode(
1327                original.to_frame().expect("reconstructable message"),
1328                &mut bytes,
1329            )
1330            .expect("encodable message");
1331        client.write_all(&bytes).await.expect("writable client");
1332
1333        let mut conn = Conn::new(Buffered::<_, Frontend>::new_frontend(proxy));
1334        let mut middleware = Middleware::new((), async |_state: &mut (), _message| {
1335            Ok::<_, Infallible>(FrontendMessage::Query(Bytes::from_static(b"select 1")))
1336        });
1337        let result = conn
1338            .receive_frontend_wire_with_middleware(
1339                &mut middleware,
1340                &backend::RuntimeState::SimpleCopyIn,
1341            )
1342            .await;
1343
1344        assert!(matches!(
1345            result,
1346            Err(ReceiveError::Intercept(InterceptError::Invalid(
1347                FrontendMessage::Query(_)
1348            )))
1349        ));
1350        conn.into_transport();
1351    }
1352
1353    #[tokio::test]
1354    async fn middleware_replacement_reaches_the_forwarded_peer() {
1355        let (client_side, mut client) = tokio::io::duplex(128);
1356        let original = FrontendMessage::Query(Bytes::from_static(b"select plaintext"));
1357        let mut bytes = BytesMut::new();
1358        PgCodec::<Frontend>::default()
1359            .encode(
1360                original.to_frame().expect("reconstructable message"),
1361                &mut bytes,
1362            )
1363            .expect("encodable message");
1364        client.write_all(&bytes).await.expect("writable client");
1365
1366        let mut downstream = Conn::new(Buffered::<_, Frontend>::new_frontend(client_side));
1367        let mut middleware = Middleware::new((), async |_state: &mut (), _message| {
1368            Ok::<_, Infallible>(FrontendMessage::Query(Bytes::from_static(
1369                b"select encrypted",
1370            )))
1371        });
1372        let rewritten = downstream
1373            .receive_frontend_wire_with_middleware(&mut middleware, &backend::RuntimeState::Ready)
1374            .await
1375            .expect("legal rewritten query");
1376
1377        let (upstream_side, mut server) = tokio::io::duplex(128);
1378        let mut upstream = Buffered::<_, Backend>::new(upstream_side);
1379        upstream
1380            .push(rewritten.to_frame().expect("reconstructable replacement"))
1381            .expect("encodable replacement");
1382        upstream.flush().await.expect("writable upstream");
1383
1384        let mut received = BytesMut::new();
1385        server
1386            .read_buf(&mut received)
1387            .await
1388            .expect("readable upstream peer");
1389        assert_eq!(
1390            PgCodec::<Frontend>::default()
1391                .decode(&mut received)
1392                .expect("decodable frame"),
1393            Some(FrontendMessage::Query(Bytes::from_static(
1394                b"select encrypted"
1395            )))
1396        );
1397        downstream.into_transport();
1398    }
1399
1400    #[tokio::test]
1401    async fn pre_startup_middleware_can_replace_with_another_legal_choice() {
1402        let (proxy, mut client) = tokio::io::duplex(128);
1403        client
1404            .write_all(
1405                &PreStartupMessage::SslRequest
1406                    .to_packet()
1407                    .expect("encodable SSLRequest"),
1408            )
1409            .await
1410            .expect("writable client");
1411
1412        let mut conn = Conn::new(Buffered::<_, Frontend>::new_frontend(proxy));
1413        let replacement = PreStartupMessage::CancelRequest {
1414            process_id: 42,
1415            secret_key: Bytes::from_static(b"key!"),
1416        };
1417        let expected = replacement.clone();
1418        let mut middleware = Middleware::new((), async move |_state: &mut (), _message| {
1419            Ok::<_, Infallible>(replacement.clone())
1420        });
1421
1422        assert_eq!(
1423            conn.receive_pre_startup_wire_with_middleware(
1424                &mut middleware,
1425                &server_pre_startup::RuntimeState::PreStartup,
1426            )
1427            .await
1428            .expect("legal replacement"),
1429            expected
1430        );
1431        conn.into_transport();
1432    }
1433
1434    #[tokio::test]
1435    async fn client_facing_transport_intercepts_typed_frontend_messages() {
1436        let (proxy, mut client) = tokio::io::duplex(128);
1437        let message = FrontendMessage::Query(Bytes::from_static(b"select plaintext"));
1438        let mut bytes = BytesMut::new();
1439        PgCodec::<Frontend>::default()
1440            .encode(
1441                message.to_frame().expect("reconstructable Query"),
1442                &mut bytes,
1443            )
1444            .expect("encodable Query");
1445        client.write_all(&bytes).await.expect("writable client");
1446
1447        let mut transport = Buffered::<_, Frontend>::new_frontend(proxy);
1448        let mut intercepted = transport.receive_wire().await.expect("decodable Query");
1449        let FrontendMessage::Query(query) = &mut intercepted else {
1450            panic!("unexpected frontend message")
1451        };
1452        *query = Bytes::from_static(b"select encrypted");
1453        assert_eq!(
1454            intercepted,
1455            FrontendMessage::Query(Bytes::from_static(b"select encrypted"))
1456        );
1457    }
1458
1459    #[tokio::test]
1460    async fn client_facing_transport_projects_repeated_pre_startup_choice() {
1461        let (proxy, mut client) = tokio::io::duplex(256);
1462        let ssl = PreStartupMessage::SslRequest
1463            .to_packet()
1464            .expect("encodable SSLRequest");
1465        let startup = PreStartupMessage::Startup(crate::startup::StartupMessage {
1466            version: crate::startup::ProtocolVersion::V3_2,
1467            parameters: std::collections::BTreeMap::from([(
1468                Bytes::from_static(b"user"),
1469                Bytes::from_static(b"postgres"),
1470            )]),
1471        });
1472        let startup_packet = startup.to_packet().expect("encodable StartupMessage");
1473        client.write_all(&ssl).await.expect("writable client");
1474        client
1475            .write_all(&startup_packet)
1476            .await
1477            .expect("writable client");
1478
1479        let mut conn = Conn::new(Buffered::<_, Frontend>::new_frontend(proxy));
1480        let ssl = conn
1481            .receive_pre_startup_wire()
1482            .await
1483            .expect("decodable SSLRequest");
1484        let crate::pre_startup::PreStartupOffer::Ssl(decision) = conn.offer_pre_startup(ssl) else {
1485            panic!("unexpected pre-startup branch")
1486        };
1487        let (mut conn, reply) = decision.reject_ssl();
1488        assert_eq!(reply, b'N');
1489        let message = conn
1490            .receive_pre_startup_wire()
1491            .await
1492            .expect("decodable StartupMessage");
1493        assert_eq!(message, startup);
1494        let crate::pre_startup::PreStartupOffer::Startup { conn, .. } =
1495            conn.offer_pre_startup(message)
1496        else {
1497            panic!("unexpected pre-startup branch")
1498        };
1499        let _transport = conn.into_transport();
1500    }
1501
1502    #[tokio::test]
1503    async fn client_facing_transport_applies_its_pre_startup_limit() {
1504        let (proxy, mut client) = tokio::io::duplex(32);
1505        client
1506            .write_all(&17_u32.to_be_bytes())
1507            .await
1508            .expect("writable client");
1509
1510        let mut transport =
1511            Buffered::<_, Frontend>::with_limits_frontend(proxy, 64, 16).expect("valid limits");
1512        let error = transport
1513            .receive_pre_startup()
1514            .await
1515            .expect_err("declared packet exceeds the configured limit");
1516
1517        assert_eq!(error.kind(), io::ErrorKind::InvalidData);
1518    }
1519
1520    #[tokio::test]
1521    async fn upstream_transport_negotiates_raw_gssenc_reply() {
1522        let (proxy, mut server) = tokio::io::duplex(32);
1523        let mut pending = Conn::new(Buffered::new(proxy)).request_gss();
1524        pending.flush().await.expect("GSSENCRequest is writable");
1525
1526        let mut request = [0_u8; 8];
1527        server
1528            .read_exact(&mut request)
1529            .await
1530            .expect("server receives GSSENCRequest");
1531        assert_eq!(request, gssenc_request_packet());
1532        server
1533            .write_all(b"N")
1534            .await
1535            .expect("server writes decision");
1536
1537        let Negotiation::Rejected(plaintext) = pending
1538            .receive_gss_reply()
1539            .await
1540            .expect("valid GSSENC decision")
1541        else {
1542            panic!("expected plaintext fallback")
1543        };
1544        plaintext.into_transport();
1545    }
1546
1547    #[test]
1548    fn client_facing_transport_buffers_raw_gssenc_decision() {
1549        let conn = Conn::new(Buffered::<_, Frontend>::new_frontend(()));
1550        let crate::pre_startup::PreStartupOffer::Gss(decision) =
1551            conn.offer_pre_startup(PreStartupMessage::GssEncRequest)
1552        else {
1553            panic!("expected GSSENC decision")
1554        };
1555
1556        let handshake = decision.approve_gss();
1557        assert_eq!(handshake.pending_output(), b"S");
1558        handshake.into_transport();
1559
1560        let conn = Conn::new(Buffered::<_, Frontend>::new_frontend(()));
1561        let crate::pre_startup::PreStartupOffer::Gss(decision) =
1562            conn.offer_pre_startup(PreStartupMessage::GssEncRequest)
1563        else {
1564            panic!("expected GSSENC decision")
1565        };
1566        let terminated = decision.reject_gss_with_legacy_error();
1567        assert_eq!(terminated.pending_output(), b"E");
1568        terminated.into_transport();
1569    }
1570
1571    #[test]
1572    fn client_facing_transport_buffers_legacy_ssl_error() {
1573        let conn = Conn::new(Buffered::<_, Frontend>::new_frontend(()));
1574        let crate::pre_startup::PreStartupOffer::Ssl(decision) =
1575            conn.offer_pre_startup(PreStartupMessage::SslRequest)
1576        else {
1577            panic!("expected SSL decision")
1578        };
1579
1580        let terminated = decision.reject_ssl_with_legacy_error();
1581        assert_eq!(terminated.pending_output(), b"E");
1582        terminated.into_transport();
1583    }
1584}