Skip to main content

pg_proto/
transport.rs

1//! Buffered, cancellation-safe transport with one inbound interception workflow.
2//!
3//! Wire-family-specific readers feed shared typed or runtime-checked validation,
4//! while receipt and session projection remain separate operations.
5
6use 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/// Transport wrapper which retains bytes until each write has completed.
91#[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    /// Wraps an upstream-facing transport which receives backend messages.
103    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    /// Creates a backend-facing transport with a bounded tagged-frame size.
115    ///
116    /// # Errors
117    ///
118    /// Returns an error when the limit is outside `PostgreSQL`'s frame range.
119    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    /// Wraps a client-facing transport which receives frontend messages.
133    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    /// Creates a frontend-facing transport with a bounded tagged-frame size.
145    ///
146    /// # Errors
147    ///
148    /// Returns an error when the limit is outside `PostgreSQL`'s frame range.
149    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    /// Creates a frontend-facing transport with bounded tagged and pre-startup packets.
154    ///
155    /// # Errors
156    ///
157    /// Returns an error when either limit is outside `PostgreSQL`'s framing range.
158    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    /// Encodes a frame synchronously into the outbound buffer.
182    ///
183    /// # Errors
184    ///
185    /// Returns an error when the frame is too large to encode.
186    pub fn push(&mut self, frame: Frame) -> io::Result<()> {
187        self.inbound_codec.encode(frame, &mut self.outbound)
188    }
189
190    #[must_use]
191    /// Returns encoded bytes which have not yet been fully written.
192    pub fn pending(&self) -> &[u8] {
193        &self.outbound
194    }
195
196    /// Removes buffering and returns the underlying I/O transport.
197    pub fn into_inner(self) -> S {
198        self.io
199    }
200
201    /// Borrows the underlying I/O transport without disturbing codec buffers.
202    pub const fn get_ref(&self) -> &S {
203        &self.io
204    }
205
206    /// Mutably borrows the underlying I/O transport without disturbing codec buffers.
207    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    /// Returns the backend asynchronous-message demultiplexer.
217    pub const fn demux(&self) -> &Demux {
218        &self.demux
219    }
220
221    /// Returns mutable access to the backend asynchronous-message demultiplexer.
222    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    /// Writes all buffered bytes without consuming the connection.
282    ///
283    /// Completed partial writes are removed immediately. If this future is
284    /// cancelled, the connection remains owned by the caller and all unwritten
285    /// bytes remain buffered for the next call.
286    ///
287    /// # Errors
288    ///
289    /// Returns the underlying transport's write error or `WriteZero`.
290    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    /// Receives one typed message in this transport's inbound direction.
307    ///
308    /// # Errors
309    ///
310    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
311    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    /// Receives one raw first packet before tagged frontend framing begins.
336    ///
337    /// # Errors
338    ///
339    /// Returns malformed pre-startup data and underlying transport read errors.
340    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    /// Receives one decoded backend message while retaining partial input.
359    ///
360    /// # Errors
361    ///
362    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
363    pub async fn receive_backend(&mut self) -> io::Result<BackendMessage> {
364        self.receive_wire().await
365    }
366
367    /// Receives the next protocol-advancing message through the async demux.
368    ///
369    /// # Errors
370    ///
371    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
372    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    /// Projects an inspected or modified backend message into the session stream.
382    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    /// Adds an already-typed message to this connection's outbound buffer.
389    ///
390    /// # Errors
391    ///
392    /// Returns an error when the frame is too large to encode.
393    pub fn push_frame(&mut self, frame: Frame) -> io::Result<()> {
394        self.transport_mut().push(frame)
395    }
396
397    #[must_use]
398    /// Returns encoded output which has not yet been flushed.
399    pub fn pending_output(&self) -> &[u8] {
400        self.transport().pending()
401    }
402}
403
404impl<S, Cleanliness> Conn<Buffered<S, Backend>, PreStartup, Cleanliness> {
405    /// Buffers an `SSLRequest` and enters the raw single-byte reply phase.
406    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    /// Buffers a `GSSENCRequest` and enters the raw single-byte reply phase.
412    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    /// Buffers the server's raw `S` response and enters the TLS handshake phase.
422    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    /// Buffers the server's raw `N` response and returns to pre-startup choice.
428    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    /// Buffers the historical raw `E` response and terminates negotiation.
434    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    /// Buffers the server's raw `S` response and enters the GSS handshake phase.
446    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    /// Buffers the server's raw `N` response and returns to pre-startup choice.
454    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    /// Buffers the historical raw `E` response and terminates negotiation.
460    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    /// Receives and projects the server's raw SSL decision byte.
470    ///
471    /// # Errors
472    ///
473    /// Returns an I/O error or rejects a byte other than `S`, `N`, or `E`.
474    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    /// Receives the server decision and enforces the selected plaintext fallback policy.
486    ///
487    /// # Errors
488    ///
489    /// Returns an I/O error or rejects a byte other than `S`, `N`, or `E`.
490    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    /// Receives and projects the server's raw GSSENC decision byte.
503    ///
504    /// # Errors
505    ///
506    /// Returns an I/O error or rejects a byte other than `S`, `N`, or `E`.
507    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    /// Completes a client-side TLS handshake and changes the transport type.
525    ///
526    /// # Errors
527    ///
528    /// Returns a TLS handshake, certificate, channel-binding, or buffer-state error.
529    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    /// Completes a server-side TLS handshake and changes the transport type.
545    ///
546    /// # Errors
547    ///
548    /// Returns a TLS handshake, certificate, channel-binding, or buffer-state error.
549    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    /// Buffers the raw, untagged startup packet before normal framing begins.
564    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    /// Flushes buffered output while retaining ownership of the typed connection.
571    ///
572    /// # Errors
573    ///
574    /// Returns an error from the underlying transport.
575    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    /// Receives one backend message before demultiplexing or state advancement.
582    /// This is the interception point for proxy policy and message rewriting.
583    ///
584    /// # Errors
585    ///
586    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
587    pub async fn receive_backend_wire(&mut self) -> io::Result<BackendMessage> {
588        self.transport_mut().receive_backend().await
589    }
590
591    /// Receives one backend message through middleware indexed by this connection phase.
592    ///
593    /// Unlike [`Self::receive_backend_wire_with_middleware`], callers do not pass
594    /// a runtime protocol state. `Phase` selects the generated legal message set
595    /// and the server sender role at compile time.
596    ///
597    /// # Errors
598    ///
599    /// Returns an I/O or decoding error, an illegal peer message, a middleware
600    /// policy error, or a phase-legal replacement with an invalid wire shape.
601    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    /// Receives typed backend traffic until one protocol-advancing item remains.
625    ///
626    /// Asynchronous messages pass through the same middleware, are recorded by
627    /// the demultiplexer in wire order, and leave `Phase` unchanged.
628    ///
629    /// # Errors
630    ///
631    /// Returns the same failures as [`Self::receive_backend_typed`].
632    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    /// Receives one SSL or GSSENC decision through phase-typed middleware.
654    ///
655    /// `Phase` must be an encryption-reply phase; callers then consume the
656    /// connection with its existing typed `receive_reply` projection.
657    ///
658    /// # Errors
659    ///
660    /// Returns an I/O error, illegal decision, middleware policy error, or an
661    /// invalid replacement wire shape.
662    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    /// Receives, intercepts, and validates one backend message before projection.
686    ///
687    /// # Errors
688    ///
689    /// Returns an I/O, decoding, middleware-policy, or state-validation error.
690    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    /// Projects an inspected or modified message into the filtered session stream.
703    pub fn project_backend(&mut self, message: BackendMessage) -> Option<SessionItem> {
704        self.transport_mut().project_backend(message)
705    }
706
707    /// Receives the next message in the filtered session projection.
708    ///
709    /// # Errors
710    ///
711    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
712    pub async fn receive(&mut self) -> io::Result<SessionItem> {
713        self.transport_mut().receive_session().await
714    }
715
716    /// Receives backend messages through middleware before demultiplexing.
717    ///
718    /// Asynchronous messages are intercepted and then recorded by the demux;
719    /// this method continues until a protocol-advancing item is available.
720    ///
721    /// # Errors
722    ///
723    /// Returns an I/O, decoding, middleware-policy, or state-validation error.
724    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    /// Returns the latest upstream cancellation key observed during startup.
745    pub fn cancel_key(&self) -> Option<&CancelKey> {
746        self.transport().demux().cancel_key()
747    }
748
749    /// Returns the latest backend parameter values observed by the demux.
750    #[must_use]
751    pub fn parameters(&self) -> &BTreeMap<Bytes, Bytes> {
752        self.transport().demux().parameters()
753    }
754
755    /// Returns whether current parameters differ from the startup baseline.
756    #[must_use]
757    pub fn parameters_changed(&self) -> bool {
758        self.transport().demux().parameters_changed()
759    }
760
761    /// Returns the latest transaction status observed in `ReadyForQuery`.
762    #[must_use]
763    pub fn transaction_status(&self) -> Option<crate::codec::TransactionStatus> {
764        self.transport().demux().transaction_status()
765    }
766
767    /// Removes the oldest queued asynchronous notification.
768    pub fn pop_notification(&mut self) -> Option<Notification> {
769        self.transport_mut().demux_mut().pop_notification()
770    }
771
772    /// Removes the next tagged notice for prompt forwarding to the client.
773    pub fn pop_notice(&mut self) -> Option<TaggedNotice> {
774        self.transport_mut().demux_mut().pop_notice()
775    }
776
777    /// Removes the next ordered parameter update for forwarding to the client.
778    pub fn pop_parameter_status(&mut self) -> Option<ParameterStatus> {
779        self.transport_mut().demux_mut().pop_parameter_status()
780    }
781
782    /// Removes the next independent backend event in original wire order.
783    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    /// Receives one frontend message before any server-role state advancement.
790    ///
791    /// # Errors
792    ///
793    /// Returns decoding and underlying transport read errors, or `UnexpectedEof`.
794    pub async fn receive_frontend_wire(&mut self) -> io::Result<FrontendMessage> {
795        self.transport_mut().receive_wire().await
796    }
797
798    /// Receives one frontend message through middleware indexed by this connection phase.
799    ///
800    /// `Phase` selects the generated legal message set and the client sender role
801    /// at compile time, so no runtime protocol-state argument is accepted.
802    ///
803    /// # Errors
804    ///
805    /// Returns an I/O or decoding error, an illegal peer message, a middleware
806    /// policy error, or a phase-legal replacement with an invalid wire shape.
807    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    /// Receives, intercepts, and validates one frontend message before projection.
831    ///
832    /// # Errors
833    ///
834    /// Returns an I/O, decoding, middleware-policy, or state-validation error.
835    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    /// Receives a raw pre-startup packet before server-role state projection.
850    ///
851    /// # Errors
852    ///
853    /// Returns malformed pre-startup data and underlying transport read errors.
854    pub async fn receive_pre_startup_wire(&mut self) -> io::Result<PreStartupMessage> {
855        self.transport_mut().receive_pre_startup().await
856    }
857
858    /// Receives a client pre-startup packet through phase-typed middleware.
859    ///
860    /// # Errors
861    ///
862    /// Returns an I/O or decoding error, an illegal pre-startup packet, a
863    /// middleware policy error, or an invalid replacement wire shape.
864    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    /// Receives, intercepts, and validates one untagged pre-startup message.
887    ///
888    /// # Errors
889    ///
890    /// Returns an I/O, decoding, middleware-policy, or state-validation error.
891    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}