1use 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#[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 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 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 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 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 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 pub fn push(&mut self, frame: Frame) -> io::Result<()> {
130 self.inbound_codec.encode(frame, &mut self.outbound)
131 }
132
133 #[must_use]
134 pub fn pending(&self) -> &[u8] {
136 &self.outbound
137 }
138
139 pub fn into_inner(self) -> S {
141 self.io
142 }
143
144 pub const fn get_ref(&self) -> &S {
146 &self.io
147 }
148
149 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 pub const fn demux(&self) -> &Demux {
161 &self.demux
162 }
163
164 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 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 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 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 pub async fn receive_backend(&mut self) -> io::Result<BackendMessage> {
307 self.receive_wire().await
308 }
309
310 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 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 pub fn push_frame(&mut self, frame: Frame) -> io::Result<()> {
337 self.transport_mut().push(frame)
338 }
339
340 #[must_use]
341 pub fn pending_output(&self) -> &[u8] {
343 self.transport().pending()
344 }
345}
346
347impl<S, Cleanliness> Conn<Buffered<S, Backend>, PreStartup, Cleanliness> {
348 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 pub async fn receive_backend_wire(&mut self) -> io::Result<BackendMessage> {
531 self.transport_mut().receive_backend().await
532 }
533
534 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 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 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 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 pub fn project_backend(&mut self, message: BackendMessage) -> Option<SessionItem> {
684 self.transport_mut().project_backend(message)
685 }
686
687 pub async fn receive(&mut self) -> io::Result<SessionItem> {
693 self.transport_mut().receive_session().await
694 }
695
696 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 pub fn cancel_key(&self) -> Option<&CancelKey> {
726 self.transport().demux().cancel_key()
727 }
728
729 #[must_use]
731 pub fn parameters(&self) -> &BTreeMap<Bytes, Bytes> {
732 self.transport().demux().parameters()
733 }
734
735 #[must_use]
737 pub fn parameters_changed(&self) -> bool {
738 self.transport().demux().parameters_changed()
739 }
740
741 #[must_use]
743 pub fn transaction_status(&self) -> Option<crate::codec::TransactionStatus> {
744 self.transport().demux().transaction_status()
745 }
746
747 pub fn pop_notification(&mut self) -> Option<Notification> {
749 self.transport_mut().demux_mut().pop_notification()
750 }
751
752 pub fn pop_notice(&mut self) -> Option<TaggedNotice> {
754 self.transport_mut().demux_mut().pop_notice()
755 }
756
757 pub fn pop_parameter_status(&mut self) -> Option<ParameterStatus> {
759 self.transport_mut().demux_mut().pop_parameter_status()
760 }
761
762 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 pub async fn receive_frontend_wire(&mut self) -> io::Result<FrontendMessage> {
775 self.transport_mut().receive_wire().await
776 }
777
778 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 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 pub async fn receive_pre_startup_wire(&mut self) -> io::Result<PreStartupMessage> {
857 self.transport_mut().receive_pre_startup().await
858 }
859
860 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 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}