Skip to main content

qail_pg/driver/
io.rs

1//! Core I/O operations for PostgreSQL connection.
2//!
3//! This module provides low-level send/receive methods.
4
5use super::{PgBytesRow, PgConnection, PgError, PgResult, is_ignorable_session_message};
6use crate::protocol::{BackendMessage, FrontendMessage, PgEncoder};
7use bytes::{Bytes, BytesMut};
8use tokio::io::{AsyncReadExt, AsyncWriteExt};
9
10pub(crate) const MAX_MESSAGE_SIZE: usize = 64 * 1024 * 1024; // 64 MB — prevents OOM from malicious server messages
11
12/// Default read timeout for individual socket reads.
13/// Prevents Slowloris DoS where a server sends partial data then goes silent.
14const DEFAULT_READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
15/// Default write timeout for individual socket writes/flushes.
16/// Prevents indefinitely blocked writes from pinning pool slots.
17const DEFAULT_WRITE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
18const READ_SPARE_LOW_WATERMARK: usize = 64 * 1024;
19
20#[inline]
21fn reserve_read_spare_capacity(buffer: &mut BytesMut) {
22    let spare = buffer.capacity().saturating_sub(buffer.len());
23    if spare < READ_SPARE_LOW_WATERMARK {
24        let target_spare = READ_SPARE_LOW_WATERMARK.max(buffer.capacity());
25        buffer.reserve(target_spare.saturating_sub(spare));
26    }
27}
28
29#[inline]
30fn parse_data_row_payload_owned(payload: &[u8]) -> PgResult<Vec<Option<Vec<u8>>>> {
31    if payload.len() < 2 {
32        return Err(PgError::Protocol("DataRow payload too short".into()));
33    }
34
35    let raw_count = i16::from_be_bytes([payload[0], payload[1]]);
36    if raw_count < 0 {
37        return Err(PgError::Protocol(format!(
38            "DataRow invalid column count: {}",
39            raw_count
40        )));
41    }
42    let column_count = raw_count as usize;
43    if column_count > (payload.len() - 2) / 4 + 1 {
44        return Err(PgError::Protocol(format!(
45            "DataRow claims {} columns but payload is only {} bytes",
46            column_count,
47            payload.len()
48        )));
49    }
50
51    let mut columns = Vec::with_capacity(column_count);
52    let mut pos = 2;
53    for _ in 0..column_count {
54        if pos + 4 > payload.len() {
55            return Err(PgError::Protocol(
56                "DataRow truncated: missing column length".into(),
57            ));
58        }
59
60        let len = i32::from_be_bytes([
61            payload[pos],
62            payload[pos + 1],
63            payload[pos + 2],
64            payload[pos + 3],
65        ]);
66        pos += 4;
67
68        if len == -1 {
69            columns.push(None);
70            continue;
71        }
72        if len < -1 {
73            return Err(PgError::Protocol(format!(
74                "DataRow invalid column length: {}",
75                len
76            )));
77        }
78
79        let len = len as usize;
80        if len > payload.len().saturating_sub(pos) {
81            return Err(PgError::Protocol(
82                "DataRow truncated: column data exceeds payload".into(),
83            ));
84        }
85        columns.push(Some(payload[pos..pos + len].to_vec()));
86        pos += len;
87    }
88
89    if pos != payload.len() {
90        return Err(PgError::Protocol("DataRow has trailing bytes".into()));
91    }
92
93    Ok(columns)
94}
95
96#[inline]
97fn parse_data_row_payload_reuse(
98    payload: &[u8],
99    columns: &mut Vec<Option<Vec<u8>>>,
100) -> PgResult<()> {
101    if payload.len() < 2 {
102        return Err(PgError::Protocol("DataRow payload too short".into()));
103    }
104
105    let raw_count = i16::from_be_bytes([payload[0], payload[1]]);
106    if raw_count < 0 {
107        return Err(PgError::Protocol(format!(
108            "DataRow invalid column count: {}",
109            raw_count
110        )));
111    }
112    let column_count = raw_count as usize;
113    if column_count > (payload.len() - 2) / 4 + 1 {
114        return Err(PgError::Protocol(format!(
115            "DataRow claims {} columns but payload is only {} bytes",
116            column_count,
117            payload.len()
118        )));
119    }
120
121    let previous_len = columns.len();
122    if previous_len < column_count {
123        columns.reserve(column_count - previous_len);
124    }
125
126    let mut pos = 2usize;
127    for idx in 0..column_count {
128        if pos + 4 > payload.len() {
129            return Err(PgError::Protocol(
130                "DataRow truncated: missing column length".into(),
131            ));
132        }
133
134        let len = i32::from_be_bytes([
135            payload[pos],
136            payload[pos + 1],
137            payload[pos + 2],
138            payload[pos + 3],
139        ]);
140        pos += 4;
141
142        if len == -1 {
143            if idx < previous_len {
144                columns[idx] = None;
145            } else {
146                columns.push(None);
147            }
148            continue;
149        }
150        if len < -1 {
151            return Err(PgError::Protocol(format!(
152                "DataRow invalid column length: {}",
153                len
154            )));
155        }
156
157        let len = len as usize;
158        if len > payload.len().saturating_sub(pos) {
159            return Err(PgError::Protocol(
160                "DataRow truncated: column data exceeds payload".into(),
161            ));
162        }
163        let value = &payload[pos..pos + len];
164        pos += len;
165
166        if idx < previous_len {
167            match &mut columns[idx] {
168                Some(buf) => {
169                    buf.clear();
170                    buf.extend_from_slice(value);
171                }
172                None => columns[idx] = Some(value.to_vec()),
173            }
174        } else {
175            columns.push(Some(value.to_vec()));
176        }
177    }
178
179    if columns.len() > column_count {
180        columns.truncate(column_count);
181    }
182
183    if pos != payload.len() {
184        return Err(PgError::Protocol("DataRow has trailing bytes".into()));
185    }
186
187    Ok(())
188}
189
190#[inline]
191fn parse_data_row_payload_zerocopy(payload: Bytes, row: &mut PgBytesRow) -> PgResult<()> {
192    if payload.len() < 2 {
193        return Err(PgError::Protocol("DataRow payload too short".into()));
194    }
195
196    let raw_count = i16::from_be_bytes([payload[0], payload[1]]);
197    if raw_count < 0 {
198        return Err(PgError::Protocol(format!(
199            "DataRow invalid column count: {}",
200            raw_count
201        )));
202    }
203    let column_count = raw_count as usize;
204    if column_count > (payload.len() - 2) / 4 + 1 {
205        return Err(PgError::Protocol(format!(
206            "DataRow claims {} columns but payload is only {} bytes",
207            column_count,
208            payload.len()
209        )));
210    }
211
212    row.payload = payload;
213    row.spans.clear();
214    if row.spans.capacity() < column_count {
215        row.spans.reserve(column_count - row.spans.capacity());
216    }
217
218    let mut pos = 2usize;
219    for _ in 0..column_count {
220        if pos + 4 > row.payload.len() {
221            return Err(PgError::Protocol(
222                "DataRow truncated: missing column length".into(),
223            ));
224        }
225
226        let len = i32::from_be_bytes([
227            row.payload[pos],
228            row.payload[pos + 1],
229            row.payload[pos + 2],
230            row.payload[pos + 3],
231        ]);
232        pos += 4;
233
234        if len == -1 {
235            row.spans.push(None);
236            continue;
237        }
238        if len < -1 {
239            return Err(PgError::Protocol(format!(
240                "DataRow invalid column length: {}",
241                len
242            )));
243        }
244
245        let len = len as usize;
246        if len > row.payload.len().saturating_sub(pos) {
247            return Err(PgError::Protocol(
248                "DataRow truncated: column data exceeds payload".into(),
249            ));
250        }
251        row.spans.push(Some((pos, len)));
252        pos += len;
253    }
254
255    if pos != row.payload.len() {
256        return Err(PgError::Protocol("DataRow has trailing bytes".into()));
257    }
258
259    Ok(())
260}
261
262#[inline]
263fn parse_first_column_payload_zerocopy(payload: Bytes) -> PgResult<Option<Bytes>> {
264    if payload.len() < 2 {
265        return Err(PgError::Protocol("DataRow payload too short".into()));
266    }
267
268    let raw_count = i16::from_be_bytes([payload[0], payload[1]]);
269    if raw_count < 0 {
270        return Err(PgError::Protocol(format!(
271            "DataRow invalid column count: {}",
272            raw_count
273        )));
274    }
275    let column_count = raw_count as usize;
276    if column_count > (payload.len() - 2) / 4 + 1 {
277        return Err(PgError::Protocol(format!(
278            "DataRow claims {} columns but payload is only {} bytes",
279            column_count,
280            payload.len()
281        )));
282    }
283
284    let mut pos = 2usize;
285    let mut first_column = None;
286
287    for idx in 0..column_count {
288        if pos + 4 > payload.len() {
289            return Err(PgError::Protocol(
290                "DataRow truncated: missing column length".into(),
291            ));
292        }
293
294        let len = i32::from_be_bytes([
295            payload[pos],
296            payload[pos + 1],
297            payload[pos + 2],
298            payload[pos + 3],
299        ]);
300        pos += 4;
301
302        if len == -1 {
303            if idx == 0 {
304                first_column = None;
305            }
306            continue;
307        }
308        if len < -1 {
309            return Err(PgError::Protocol(format!(
310                "DataRow invalid column length: {}",
311                len
312            )));
313        }
314
315        let len = len as usize;
316        if len > payload.len().saturating_sub(pos) {
317            return Err(PgError::Protocol(
318                "DataRow truncated: column data exceeds payload".into(),
319            ));
320        }
321
322        if idx == 0 {
323            first_column = Some(payload.slice(pos..pos + len));
324        }
325        pos += len;
326    }
327
328    if pos != payload.len() {
329        return Err(PgError::Protocol("DataRow has trailing bytes".into()));
330    }
331
332    Ok(first_column)
333}
334
335#[inline]
336fn parse_first_four_columns_payload_zerocopy(
337    payload: Bytes,
338    columns: &mut [Option<Bytes>; 4],
339) -> PgResult<()> {
340    if payload.len() < 2 {
341        return Err(PgError::Protocol("DataRow payload too short".into()));
342    }
343
344    let raw_count = i16::from_be_bytes([payload[0], payload[1]]);
345    if raw_count < 0 {
346        return Err(PgError::Protocol(format!(
347            "DataRow invalid column count: {}",
348            raw_count
349        )));
350    }
351    let column_count = raw_count as usize;
352    if column_count > (payload.len() - 2) / 4 + 1 {
353        return Err(PgError::Protocol(format!(
354            "DataRow claims {} columns but payload is only {} bytes",
355            column_count,
356            payload.len()
357        )));
358    }
359    if column_count != 4 {
360        return Err(PgError::Protocol(format!(
361            "DataRow fast-path expects exactly 4 columns, got {}",
362            column_count
363        )));
364    }
365
366    let mut pos = 2usize;
367    for slot in columns.iter_mut() {
368        if pos + 4 > payload.len() {
369            return Err(PgError::Protocol(
370                "DataRow truncated: missing column length".into(),
371            ));
372        }
373
374        let len = i32::from_be_bytes([
375            payload[pos],
376            payload[pos + 1],
377            payload[pos + 2],
378            payload[pos + 3],
379        ]);
380        pos += 4;
381
382        if len == -1 {
383            *slot = None;
384            continue;
385        }
386        if len < -1 {
387            return Err(PgError::Protocol(format!(
388                "DataRow invalid column length: {}",
389                len
390            )));
391        }
392
393        let len = len as usize;
394        if len > payload.len().saturating_sub(pos) {
395            return Err(PgError::Protocol(
396                "DataRow truncated: column data exceeds payload".into(),
397            ));
398        }
399        *slot = Some(payload.slice(pos..pos + len));
400        pos += len;
401    }
402
403    if pos != payload.len() {
404        return Err(PgError::Protocol("DataRow has trailing bytes".into()));
405    }
406
407    Ok(())
408}
409
410impl PgConnection {
411    #[inline]
412    fn stream_requires_flush(&self) -> bool {
413        use super::stream::PgStream;
414
415        match &self.stream {
416            PgStream::Tcp(_) => false,
417            PgStream::Tls(_) => true,
418            #[cfg(all(target_os = "linux", feature = "io_uring"))]
419            PgStream::Uring(_) => false,
420            #[cfg(unix)]
421            PgStream::Unix(_) => false,
422            #[cfg(all(feature = "enterprise-gssapi", target_os = "linux"))]
423            PgStream::GssEnc(_) => true,
424        }
425    }
426
427    #[inline]
428    pub(crate) fn mark_io_desynced(&mut self) {
429        self.io_desynced = true;
430    }
431
432    #[inline]
433    pub(crate) fn is_io_desynced(&self) -> bool {
434        self.io_desynced
435    }
436
437    #[inline]
438    fn protocol_desync<T>(&mut self, msg: String) -> PgResult<T> {
439        self.mark_io_desynced();
440        Err(PgError::Protocol(msg))
441    }
442
443    #[inline]
444    fn protocol_desync_error<T>(&mut self, err: PgError) -> PgResult<T> {
445        match err {
446            PgError::Protocol(msg) => self.protocol_desync(msg),
447            err => {
448                self.mark_io_desynced();
449                Err(err)
450            }
451        }
452    }
453
454    #[inline]
455    fn connection_desync<T>(&mut self, msg: String) -> PgResult<T> {
456        self.mark_io_desynced();
457        Err(PgError::Connection(msg))
458    }
459
460    /// Send queued statement `Close` messages and drain until `ReadyForQuery`.
461    ///
462    /// We ignore `26000 prepared statement ... does not exist` because this
463    /// can happen after failover or server-side invalidation, and in that case
464    /// local state is already being reconciled by retry paths.
465    async fn flush_pending_statement_closes(&mut self) -> PgResult<()> {
466        if self.draining_statement_closes || self.pending_statement_closes.is_empty() {
467            return Ok(());
468        }
469
470        self.draining_statement_closes = true;
471        let close_names = std::mem::take(&mut self.pending_statement_closes);
472
473        let estimated_payload_len: usize = close_names
474            .iter()
475            .map(|name| 16usize.saturating_add(name.len()))
476            .sum();
477        let mut buf = BytesMut::with_capacity(estimated_payload_len.saturating_add(5));
478        for stmt_name in &close_names {
479            let close_msg = PgEncoder::try_encode_close(false, stmt_name)
480                .map_err(|e| PgError::Encode(e.to_string()))?;
481            buf.extend_from_slice(&close_msg);
482        }
483        PgEncoder::encode_sync_to(&mut buf);
484
485        if let Err(err) = self
486            .write_all_with_timeout_inner(&buf, "pending statement close write")
487            .await
488        {
489            self.draining_statement_closes = false;
490            return Err(err);
491        }
492        if let Err(err) = self
493            .flush_with_timeout("pending statement close flush")
494            .await
495        {
496            self.draining_statement_closes = false;
497            return Err(err);
498        }
499
500        let mut error: Option<PgError> = None;
501        loop {
502            let msg = match self.recv().await {
503                Ok(msg) => msg,
504                Err(err) => {
505                    self.draining_statement_closes = false;
506                    return Err(err);
507                }
508            };
509            match msg {
510                BackendMessage::CloseComplete => {}
511                BackendMessage::ReadyForQuery(_) => {
512                    self.draining_statement_closes = false;
513                    if let Some(err) = error {
514                        return Err(err);
515                    }
516                    return Ok(());
517                }
518                BackendMessage::ErrorResponse(err_fields) => {
519                    if error.is_none() {
520                        let code_26000 = err_fields.code.eq_ignore_ascii_case("26000");
521                        let msg_lower = err_fields.message.to_ascii_lowercase();
522                        let missing_prepared = msg_lower.contains("prepared statement")
523                            && msg_lower.contains("does not exist");
524                        if !(code_26000 && missing_prepared) {
525                            error = Some(PgError::QueryServer(err_fields.into()));
526                        }
527                    }
528                }
529                msg if is_ignorable_session_message(&msg) => {}
530                other => {
531                    self.draining_statement_closes = false;
532                    return self.protocol_desync(format!(
533                        "Unexpected backend message during pending statement close drain: {:?}",
534                        other
535                    ));
536                }
537            }
538        }
539    }
540
541    /// Write all bytes with a timeout guard.
542    ///
543    /// Prevents stuck kernel send buffers or dead sockets from hanging forever.
544    pub(crate) async fn write_all_with_timeout(
545        &mut self,
546        bytes: &[u8],
547        operation: &str,
548    ) -> PgResult<()> {
549        if !self.draining_statement_closes && !self.pending_statement_closes.is_empty() {
550            self.flush_pending_statement_closes().await?;
551        }
552        self.write_all_with_timeout_inner(bytes, operation).await
553    }
554
555    async fn write_all_with_timeout_inner(
556        &mut self,
557        bytes: &[u8],
558        operation: &str,
559    ) -> PgResult<()> {
560        if bytes.is_empty() {
561            return Err(PgError::Encode(
562                "refusing to send empty frontend payload".to_string(),
563            ));
564        }
565        use super::stream::PgStream;
566        let mut mark_desync = false;
567        let result = match &mut self.stream {
568            PgStream::Tcp(stream) => {
569                match tokio::time::timeout(DEFAULT_WRITE_TIMEOUT, stream.write_all(bytes)).await {
570                    Ok(Ok(())) => Ok(()),
571                    Ok(Err(e)) => {
572                        mark_desync = true;
573                        Err(PgError::Connection(format!("Write error: {}", e)))
574                    }
575                    Err(_) => {
576                        mark_desync = true;
577                        Err(PgError::Timeout(format!(
578                            "{} timeout after {:?}",
579                            operation, DEFAULT_WRITE_TIMEOUT
580                        )))
581                    }
582                }
583            }
584            PgStream::Tls(stream) => {
585                match tokio::time::timeout(DEFAULT_WRITE_TIMEOUT, stream.write_all(bytes)).await {
586                    Ok(Ok(())) => Ok(()),
587                    Ok(Err(e)) => {
588                        mark_desync = true;
589                        Err(PgError::Connection(format!("Write error: {}", e)))
590                    }
591                    Err(_) => {
592                        mark_desync = true;
593                        Err(PgError::Timeout(format!(
594                            "{} timeout after {:?}",
595                            operation, DEFAULT_WRITE_TIMEOUT
596                        )))
597                    }
598                }
599            }
600            #[cfg(all(target_os = "linux", feature = "io_uring"))]
601            PgStream::Uring(stream) => {
602                match tokio::time::timeout(DEFAULT_WRITE_TIMEOUT, stream.write_all(bytes)).await {
603                    Ok(Ok(())) => Ok(()),
604                    Ok(Err(e)) => {
605                        mark_desync = true;
606                        Err(PgError::Connection(format!("Write error: {}", e)))
607                    }
608                    Err(_) => {
609                        mark_desync = true;
610                        let _ = stream.abort_inflight();
611                        Err(PgError::Timeout(format!(
612                            "{} timeout after {:?}",
613                            operation, DEFAULT_WRITE_TIMEOUT
614                        )))
615                    }
616                }
617            }
618            #[cfg(unix)]
619            PgStream::Unix(stream) => {
620                match tokio::time::timeout(DEFAULT_WRITE_TIMEOUT, stream.write_all(bytes)).await {
621                    Ok(Ok(())) => Ok(()),
622                    Ok(Err(e)) => {
623                        mark_desync = true;
624                        Err(PgError::Connection(format!("Write error: {}", e)))
625                    }
626                    Err(_) => {
627                        mark_desync = true;
628                        Err(PgError::Timeout(format!(
629                            "{} timeout after {:?}",
630                            operation, DEFAULT_WRITE_TIMEOUT
631                        )))
632                    }
633                }
634            }
635            #[cfg(all(feature = "enterprise-gssapi", target_os = "linux"))]
636            PgStream::GssEnc(stream) => {
637                match tokio::time::timeout(DEFAULT_WRITE_TIMEOUT, stream.write_all(bytes)).await {
638                    Ok(Ok(())) => Ok(()),
639                    Ok(Err(e)) => {
640                        mark_desync = true;
641                        Err(PgError::Connection(format!("Write error: {}", e)))
642                    }
643                    Err(_) => {
644                        mark_desync = true;
645                        Err(PgError::Timeout(format!(
646                            "{} timeout after {:?}",
647                            operation, DEFAULT_WRITE_TIMEOUT
648                        )))
649                    }
650                }
651            }
652        };
653        if mark_desync {
654            self.mark_io_desynced();
655        }
656        result
657    }
658
659    /// Flush with a timeout guard.
660    pub(crate) async fn flush_with_timeout(&mut self, operation: &str) -> PgResult<()> {
661        if !self.stream_requires_flush() {
662            return Ok(());
663        }
664
665        use super::stream::PgStream;
666        let mut mark_desync = false;
667        let result = match &mut self.stream {
668            PgStream::Tcp(stream) => {
669                match tokio::time::timeout(DEFAULT_WRITE_TIMEOUT, stream.flush()).await {
670                    Ok(Ok(())) => Ok(()),
671                    Ok(Err(e)) => {
672                        mark_desync = true;
673                        Err(PgError::Connection(format!("Flush error: {}", e)))
674                    }
675                    Err(_) => {
676                        mark_desync = true;
677                        Err(PgError::Timeout(format!(
678                            "{} timeout after {:?}",
679                            operation, DEFAULT_WRITE_TIMEOUT
680                        )))
681                    }
682                }
683            }
684            PgStream::Tls(stream) => {
685                match tokio::time::timeout(DEFAULT_WRITE_TIMEOUT, stream.flush()).await {
686                    Ok(Ok(())) => Ok(()),
687                    Ok(Err(e)) => {
688                        mark_desync = true;
689                        Err(PgError::Connection(format!("Flush error: {}", e)))
690                    }
691                    Err(_) => {
692                        mark_desync = true;
693                        Err(PgError::Timeout(format!(
694                            "{} timeout after {:?}",
695                            operation, DEFAULT_WRITE_TIMEOUT
696                        )))
697                    }
698                }
699            }
700            #[cfg(all(target_os = "linux", feature = "io_uring"))]
701            PgStream::Uring(stream) => {
702                match tokio::time::timeout(DEFAULT_WRITE_TIMEOUT, stream.flush()).await {
703                    Ok(Ok(())) => Ok(()),
704                    Ok(Err(e)) => {
705                        mark_desync = true;
706                        Err(PgError::Connection(format!("Flush error: {}", e)))
707                    }
708                    Err(_) => {
709                        mark_desync = true;
710                        let _ = stream.abort_inflight();
711                        Err(PgError::Timeout(format!(
712                            "{} timeout after {:?}",
713                            operation, DEFAULT_WRITE_TIMEOUT
714                        )))
715                    }
716                }
717            }
718            #[cfg(unix)]
719            PgStream::Unix(stream) => {
720                match tokio::time::timeout(DEFAULT_WRITE_TIMEOUT, stream.flush()).await {
721                    Ok(Ok(())) => Ok(()),
722                    Ok(Err(e)) => {
723                        mark_desync = true;
724                        Err(PgError::Connection(format!("Flush error: {}", e)))
725                    }
726                    Err(_) => {
727                        mark_desync = true;
728                        Err(PgError::Timeout(format!(
729                            "{} timeout after {:?}",
730                            operation, DEFAULT_WRITE_TIMEOUT
731                        )))
732                    }
733                }
734            }
735            #[cfg(all(feature = "enterprise-gssapi", target_os = "linux"))]
736            PgStream::GssEnc(stream) => {
737                match tokio::time::timeout(DEFAULT_WRITE_TIMEOUT, stream.flush()).await {
738                    Ok(Ok(())) => Ok(()),
739                    Ok(Err(e)) => {
740                        mark_desync = true;
741                        Err(PgError::Connection(format!("Flush error: {}", e)))
742                    }
743                    Err(_) => {
744                        mark_desync = true;
745                        Err(PgError::Timeout(format!(
746                            "{} timeout after {:?}",
747                            operation, DEFAULT_WRITE_TIMEOUT
748                        )))
749                    }
750                }
751            }
752        };
753        if mark_desync {
754            self.mark_io_desynced();
755        }
756        result
757    }
758
759    /// Send a frontend message.
760    pub async fn send(&mut self, msg: FrontendMessage) -> PgResult<()> {
761        let bytes = msg
762            .encode_checked()
763            .map_err(|e| PgError::Encode(e.to_string()))?;
764        self.send_bytes(&bytes).await?;
765        Ok(())
766    }
767
768    /// Loops until a complete message is available.
769    /// Automatically buffers NotificationResponse messages for LISTEN/NOTIFY.
770    pub async fn recv(&mut self) -> PgResult<BackendMessage> {
771        loop {
772            // Try to decode from buffer first
773            if self.buffer.len() >= 5 {
774                let msg_len = u32::from_be_bytes([
775                    self.buffer[1],
776                    self.buffer[2],
777                    self.buffer[3],
778                    self.buffer[4],
779                ]) as usize;
780
781                if msg_len < 4 {
782                    return self.protocol_desync(format!(
783                        "Invalid message length: {} (minimum 4)",
784                        msg_len
785                    ));
786                }
787
788                if msg_len > MAX_MESSAGE_SIZE {
789                    return self.protocol_desync(format!(
790                        "Message too large: {} bytes (max {})",
791                        msg_len, MAX_MESSAGE_SIZE
792                    ));
793                }
794
795                if self.buffer.len() > msg_len {
796                    // We have a complete message - zero-copy split
797                    let msg_bytes = self.buffer.split_to(msg_len + 1);
798                    let (msg, _) = match BackendMessage::decode(&msg_bytes) {
799                        Ok(decoded) => decoded,
800                        Err(e) => return self.protocol_desync(e),
801                    };
802
803                    // Intercept async notifications — buffer them instead of returning
804                    if let BackendMessage::NotificationResponse {
805                        process_id,
806                        channel,
807                        payload,
808                    } = msg
809                    {
810                        self.notifications
811                            .push_back(super::notification::Notification {
812                                process_id,
813                                channel,
814                                payload,
815                            });
816                        continue; // Keep reading for the actual response
817                    }
818
819                    return Ok(msg);
820                }
821            }
822
823            let n = self.read_with_timeout().await?;
824            if n == 0 {
825                return self.connection_desync("Connection closed".to_string());
826            }
827        }
828    }
829
830    /// Receive a backend message with idle-friendly timeout behavior.
831    ///
832    /// For long-lived idle streams (e.g. logical replication), an empty
833    /// buffer uses no-timeout reads so inactivity does not fail the stream.
834    /// If a backend frame is already partially buffered, switch back to the
835    /// normal read timeout to fail-closed on partial-frame stalls.
836    pub(crate) async fn recv_without_timeout(&mut self) -> PgResult<BackendMessage> {
837        loop {
838            if self.buffer.len() >= 5 {
839                let msg_len = u32::from_be_bytes([
840                    self.buffer[1],
841                    self.buffer[2],
842                    self.buffer[3],
843                    self.buffer[4],
844                ]) as usize;
845
846                if msg_len < 4 {
847                    return self.protocol_desync(format!(
848                        "Invalid message length: {} (minimum 4)",
849                        msg_len
850                    ));
851                }
852
853                if msg_len > MAX_MESSAGE_SIZE {
854                    return self.protocol_desync(format!(
855                        "Message too large: {} bytes (max {})",
856                        msg_len, MAX_MESSAGE_SIZE
857                    ));
858                }
859
860                if self.buffer.len() > msg_len {
861                    let msg_bytes = self.buffer.split_to(msg_len + 1);
862                    let (msg, _) = match BackendMessage::decode(&msg_bytes) {
863                        Ok(decoded) => decoded,
864                        Err(e) => return self.protocol_desync(e),
865                    };
866
867                    if let BackendMessage::NotificationResponse {
868                        process_id,
869                        channel,
870                        payload,
871                    } = msg
872                    {
873                        self.notifications
874                            .push_back(super::notification::Notification {
875                                process_id,
876                                channel,
877                                payload,
878                            });
879                        continue;
880                    }
881
882                    return Ok(msg);
883                }
884            }
885
886            let n = if self.buffer.is_empty() {
887                self.read_without_timeout().await?
888            } else {
889                self.read_with_timeout().await?
890            };
891            if n == 0 {
892                return self.connection_desync("Connection closed".to_string());
893            }
894        }
895    }
896
897    /// Read from the socket with a timeout guard.
898    /// Returns the number of bytes read, or an error if the timeout fires.
899    /// This prevents Slowloris DoS attacks where a malicious server sends
900    /// partial data then goes silent, causing the driver to hang forever.
901    #[inline]
902    pub(crate) async fn read_with_timeout(&mut self) -> PgResult<usize> {
903        reserve_read_spare_capacity(&mut self.buffer);
904
905        use super::stream::PgStream;
906        let (stream, buffer) = (&mut self.stream, &mut self.buffer);
907        let mut mark_desync = false;
908        let result = match stream {
909            PgStream::Tcp(stream) => {
910                match tokio::time::timeout(DEFAULT_READ_TIMEOUT, stream.read_buf(buffer)).await {
911                    Ok(Ok(n)) => Ok(n),
912                    Ok(Err(e)) => {
913                        mark_desync = true;
914                        Err(PgError::Connection(format!("Read error: {}", e)))
915                    }
916                    Err(_) => {
917                        mark_desync = true;
918                        Err(PgError::Connection(format!(
919                            "Read timeout after {:?} — possible Slowloris attack or dead connection",
920                            DEFAULT_READ_TIMEOUT
921                        )))
922                    }
923                }
924            }
925            PgStream::Tls(stream) => {
926                match tokio::time::timeout(DEFAULT_READ_TIMEOUT, stream.read_buf(buffer)).await {
927                    Ok(Ok(n)) => Ok(n),
928                    Ok(Err(e)) => {
929                        mark_desync = true;
930                        Err(PgError::Connection(format!("Read error: {}", e)))
931                    }
932                    Err(_) => {
933                        mark_desync = true;
934                        Err(PgError::Connection(format!(
935                            "Read timeout after {:?} — possible Slowloris attack or dead connection",
936                            DEFAULT_READ_TIMEOUT
937                        )))
938                    }
939                }
940            }
941            #[cfg(all(target_os = "linux", feature = "io_uring"))]
942            PgStream::Uring(stream) => {
943                match tokio::time::timeout(DEFAULT_READ_TIMEOUT, stream.read_into(buffer, 131072))
944                    .await
945                {
946                    Ok(Ok(n)) => Ok(n),
947                    Ok(Err(e)) => {
948                        mark_desync = true;
949                        Err(PgError::Connection(format!("Read error: {}", e)))
950                    }
951                    Err(_) => {
952                        mark_desync = true;
953                        let _ = stream.abort_inflight();
954                        Err(PgError::Connection(format!(
955                            "Read timeout after {:?} — possible Slowloris attack or dead connection",
956                            DEFAULT_READ_TIMEOUT
957                        )))
958                    }
959                }
960            }
961            #[cfg(unix)]
962            PgStream::Unix(stream) => {
963                match tokio::time::timeout(DEFAULT_READ_TIMEOUT, stream.read_buf(buffer)).await {
964                    Ok(Ok(n)) => Ok(n),
965                    Ok(Err(e)) => {
966                        mark_desync = true;
967                        Err(PgError::Connection(format!("Read error: {}", e)))
968                    }
969                    Err(_) => {
970                        mark_desync = true;
971                        Err(PgError::Connection(format!(
972                            "Read timeout after {:?} — possible Slowloris attack or dead connection",
973                            DEFAULT_READ_TIMEOUT
974                        )))
975                    }
976                }
977            }
978            #[cfg(all(feature = "enterprise-gssapi", target_os = "linux"))]
979            PgStream::GssEnc(stream) => {
980                match tokio::time::timeout(DEFAULT_READ_TIMEOUT, stream.read_buf(buffer)).await {
981                    Ok(Ok(n)) => Ok(n),
982                    Ok(Err(e)) => {
983                        mark_desync = true;
984                        Err(PgError::Connection(format!("Read error: {}", e)))
985                    }
986                    Err(_) => {
987                        mark_desync = true;
988                        Err(PgError::Connection(format!(
989                            "Read timeout after {:?} — possible Slowloris attack or dead connection",
990                            DEFAULT_READ_TIMEOUT
991                        )))
992                    }
993                }
994            }
995        };
996        if mark_desync {
997            self.mark_io_desynced();
998        }
999        result
1000    }
1001
1002    /// Read from socket without timeout guard.
1003    ///
1004    /// Used for long-idle LISTEN/NOTIFY connections.
1005    pub(crate) async fn read_without_timeout(&mut self) -> PgResult<usize> {
1006        reserve_read_spare_capacity(&mut self.buffer);
1007
1008        use super::stream::PgStream;
1009        let (stream, buffer) = (&mut self.stream, &mut self.buffer);
1010        let read_result = match stream {
1011            PgStream::Tcp(stream) => stream.read_buf(buffer).await,
1012            PgStream::Tls(stream) => stream.read_buf(buffer).await,
1013            #[cfg(all(target_os = "linux", feature = "io_uring"))]
1014            PgStream::Uring(stream) => stream.read_into(buffer, 131072).await,
1015            #[cfg(unix)]
1016            PgStream::Unix(stream) => stream.read_buf(buffer).await,
1017            #[cfg(all(feature = "enterprise-gssapi", target_os = "linux"))]
1018            PgStream::GssEnc(stream) => stream.read_buf(buffer).await,
1019        };
1020
1021        match read_result {
1022            Ok(n) => Ok(n),
1023            Err(e) => {
1024                self.mark_io_desynced();
1025                Err(PgError::Connection(format!("Read error: {}", e)))
1026            }
1027        }
1028    }
1029
1030    /// Send raw bytes to the stream.
1031    /// Includes flush for TLS safety — TLS buffers internally and
1032    /// needs flush to push encrypted data to the underlying TCP socket.
1033    pub async fn send_bytes(&mut self, bytes: &[u8]) -> PgResult<()> {
1034        self.write_all_with_timeout(bytes, "send raw bytes").await?;
1035        self.flush_with_timeout("flush raw bytes").await?;
1036        Ok(())
1037    }
1038
1039    #[inline]
1040    fn decode_fast_message_type(&mut self, msg_bytes: BytesMut) -> PgResult<Option<u8>> {
1041        let msg_type = msg_bytes[0];
1042        let (msg, _) = match BackendMessage::decode(&msg_bytes) {
1043            Ok(decoded) => decoded,
1044            Err(e) => return self.protocol_desync(e),
1045        };
1046        match msg {
1047            BackendMessage::ErrorResponse(err) => Err(PgError::QueryServer(err.into())),
1048            BackendMessage::NotificationResponse {
1049                process_id,
1050                channel,
1051                payload,
1052            } => {
1053                self.notifications
1054                    .push_back(super::notification::Notification {
1055                        process_id,
1056                        channel,
1057                        payload,
1058                    });
1059                Ok(None)
1060            }
1061            _ => Ok(Some(msg_type)),
1062        }
1063    }
1064
1065    // ==================== BUFFERED WRITE API (High Performance) ====================
1066
1067    /// Buffer bytes for later flush (NO SYSCALL).
1068    /// Use flush_write_buf() to send all buffered data.
1069    #[inline]
1070    pub fn buffer_bytes(&mut self, bytes: &[u8]) {
1071        self.write_buf.extend_from_slice(bytes);
1072    }
1073
1074    /// Flush the write buffer to the stream (single write_all + flush).
1075    /// The flush is critical for TLS connections.
1076    pub async fn flush_write_buf(&mut self) -> PgResult<()> {
1077        if !self.write_buf.is_empty() {
1078            let payload = self.write_buf.split().freeze();
1079            self.write_all_with_timeout(&payload, "flush write buffer")
1080                .await?;
1081            self.flush_with_timeout("flush write buffer").await?;
1082        }
1083        Ok(())
1084    }
1085
1086    /// FAST receive - returns only message type byte, skips parsing.
1087    /// This is ~10x faster than recv() for pipelining benchmarks.
1088    /// Returns: message_type
1089    #[inline]
1090    pub(crate) async fn recv_msg_type_fast(&mut self) -> PgResult<u8> {
1091        loop {
1092            if self.buffer.len() >= 5 {
1093                let msg_len = u32::from_be_bytes([
1094                    self.buffer[1],
1095                    self.buffer[2],
1096                    self.buffer[3],
1097                    self.buffer[4],
1098                ]) as usize;
1099
1100                if msg_len < 4 {
1101                    return self.protocol_desync(format!(
1102                        "Invalid message length: {} (minimum 4)",
1103                        msg_len
1104                    ));
1105                }
1106
1107                if msg_len > MAX_MESSAGE_SIZE {
1108                    return self.protocol_desync(format!(
1109                        "Message too large: {} bytes (max {})",
1110                        msg_len, MAX_MESSAGE_SIZE
1111                    ));
1112                }
1113
1114                if self.buffer.len() > msg_len {
1115                    let msg_bytes = self.buffer.split_to(msg_len + 1);
1116                    if let Some(msg_type) = self.decode_fast_message_type(msg_bytes)? {
1117                        return Ok(msg_type);
1118                    }
1119                    continue;
1120                }
1121            }
1122
1123            let n = self.read_with_timeout().await?;
1124            if n == 0 {
1125                return self.connection_desync("Connection closed".to_string());
1126            }
1127        }
1128    }
1129
1130    /// FAST receive for result consumption - inline DataRow parsing.
1131    /// Returns: (msg_type, Option<row_data>)
1132    /// For 'D' (DataRow): returns parsed columns
1133    /// For other types: returns None
1134    /// This avoids BackendMessage enum allocation for non-DataRow messages.
1135    #[inline]
1136    pub(crate) async fn recv_with_data_fast(
1137        &mut self,
1138    ) -> PgResult<(u8, Option<Vec<Option<Vec<u8>>>>)> {
1139        loop {
1140            if self.buffer.len() >= 5 {
1141                let msg_len = u32::from_be_bytes([
1142                    self.buffer[1],
1143                    self.buffer[2],
1144                    self.buffer[3],
1145                    self.buffer[4],
1146                ]) as usize;
1147
1148                if msg_len < 4 {
1149                    return self.protocol_desync(format!(
1150                        "Invalid message length: {} (minimum 4)",
1151                        msg_len
1152                    ));
1153                }
1154
1155                if msg_len > MAX_MESSAGE_SIZE {
1156                    return self.protocol_desync(format!(
1157                        "Message too large: {} bytes (max {})",
1158                        msg_len, MAX_MESSAGE_SIZE
1159                    ));
1160                }
1161
1162                if self.buffer.len() > msg_len {
1163                    let msg_type = self.buffer[0];
1164
1165                    // Fast path: DataRow - parse inline
1166                    if msg_type == b'D' {
1167                        let parse_result = {
1168                            let payload = &self.buffer[5..msg_len + 1];
1169                            parse_data_row_payload_owned(payload)
1170                        };
1171
1172                        let _ = self.buffer.split_to(msg_len + 1);
1173                        match parse_result {
1174                            Ok(columns) => return Ok((msg_type, Some(columns))),
1175                            Err(err) => return self.protocol_desync_error(err),
1176                        }
1177                    }
1178
1179                    let msg_bytes = self.buffer.split_to(msg_len + 1);
1180                    if let Some(msg_type) = self.decode_fast_message_type(msg_bytes)? {
1181                        return Ok((msg_type, None));
1182                    }
1183                    continue;
1184                }
1185            }
1186
1187            let n = self.read_with_timeout().await?;
1188            if n == 0 {
1189                return self.connection_desync("Connection closed".to_string());
1190            }
1191        }
1192    }
1193
1194    /// FAST receive for result consumption into a reusable row buffer.
1195    ///
1196    /// This preserves owned row semantics while reusing allocations across
1197    /// `DataRow` messages.
1198    #[inline]
1199    pub(crate) async fn recv_fill_data_row_fast(
1200        &mut self,
1201        row_buf: &mut Vec<Option<Vec<u8>>>,
1202    ) -> PgResult<u8> {
1203        loop {
1204            if self.buffer.len() >= 5 {
1205                let msg_len = u32::from_be_bytes([
1206                    self.buffer[1],
1207                    self.buffer[2],
1208                    self.buffer[3],
1209                    self.buffer[4],
1210                ]) as usize;
1211
1212                if msg_len < 4 {
1213                    return self.protocol_desync(format!(
1214                        "Invalid message length: {} (minimum 4)",
1215                        msg_len
1216                    ));
1217                }
1218
1219                if msg_len > MAX_MESSAGE_SIZE {
1220                    return self.protocol_desync(format!(
1221                        "Message too large: {} bytes (max {})",
1222                        msg_len, MAX_MESSAGE_SIZE
1223                    ));
1224                }
1225
1226                if self.buffer.len() > msg_len {
1227                    let msg_type = self.buffer[0];
1228
1229                    if msg_type == b'D' {
1230                        let parse_result = {
1231                            let payload = &self.buffer[5..msg_len + 1];
1232                            parse_data_row_payload_reuse(payload, row_buf)
1233                        };
1234
1235                        let _ = self.buffer.split_to(msg_len + 1);
1236                        if let Err(err) = parse_result {
1237                            return self.protocol_desync_error(err);
1238                        }
1239                        return Ok(msg_type);
1240                    }
1241
1242                    let msg_bytes = self.buffer.split_to(msg_len + 1);
1243                    if let Some(msg_type) = self.decode_fast_message_type(msg_bytes)? {
1244                        return Ok(msg_type);
1245                    }
1246                    continue;
1247                }
1248            }
1249
1250            let n = self.read_with_timeout().await?;
1251            if n == 0 {
1252                return self.connection_desync("Connection closed".to_string());
1253            }
1254        }
1255    }
1256
1257    /// FAST receive for result consumption into a reusable zero-copy row.
1258    #[inline]
1259    pub(crate) async fn recv_fill_zerocopy_row_fast(
1260        &mut self,
1261        row: &mut PgBytesRow,
1262    ) -> PgResult<u8> {
1263        loop {
1264            if self.buffer.len() >= 5 {
1265                let msg_len = u32::from_be_bytes([
1266                    self.buffer[1],
1267                    self.buffer[2],
1268                    self.buffer[3],
1269                    self.buffer[4],
1270                ]) as usize;
1271
1272                if msg_len < 4 {
1273                    return self.protocol_desync(format!(
1274                        "Invalid message length: {} (minimum 4)",
1275                        msg_len
1276                    ));
1277                }
1278
1279                if msg_len > MAX_MESSAGE_SIZE {
1280                    return self.protocol_desync(format!(
1281                        "Message too large: {} bytes (max {})",
1282                        msg_len, MAX_MESSAGE_SIZE
1283                    ));
1284                }
1285
1286                if self.buffer.len() > msg_len {
1287                    let msg_type = self.buffer[0];
1288
1289                    if msg_type == b'D' {
1290                        let msg_bytes = self.buffer.split_to(msg_len + 1).freeze();
1291                        let payload = msg_bytes.slice(5..);
1292                        if let Err(err) = parse_data_row_payload_zerocopy(payload, row) {
1293                            return self.protocol_desync_error(err);
1294                        }
1295                        return Ok(msg_type);
1296                    }
1297
1298                    let msg_bytes = self.buffer.split_to(msg_len + 1);
1299                    if let Some(msg_type) = self.decode_fast_message_type(msg_bytes)? {
1300                        return Ok(msg_type);
1301                    }
1302                    continue;
1303                }
1304            }
1305
1306            let n = self.read_with_timeout().await?;
1307            if n == 0 {
1308                return self.connection_desync("Connection closed".to_string());
1309            }
1310        }
1311    }
1312
1313    /// FAST receive for scalar result consumption into a reusable first-column buffer.
1314    #[inline]
1315    pub(crate) async fn recv_fill_first_column_zerocopy_fast(
1316        &mut self,
1317        first_column: &mut Option<Bytes>,
1318    ) -> PgResult<u8> {
1319        loop {
1320            if self.buffer.len() >= 5 {
1321                let msg_len = u32::from_be_bytes([
1322                    self.buffer[1],
1323                    self.buffer[2],
1324                    self.buffer[3],
1325                    self.buffer[4],
1326                ]) as usize;
1327
1328                if msg_len < 4 {
1329                    return self.protocol_desync(format!(
1330                        "Invalid message length: {} (minimum 4)",
1331                        msg_len
1332                    ));
1333                }
1334
1335                if msg_len > MAX_MESSAGE_SIZE {
1336                    return self.protocol_desync(format!(
1337                        "Message too large: {} bytes (max {})",
1338                        msg_len, MAX_MESSAGE_SIZE
1339                    ));
1340                }
1341
1342                if self.buffer.len() > msg_len {
1343                    let msg_type = self.buffer[0];
1344
1345                    if msg_type == b'D' {
1346                        let msg_bytes = self.buffer.split_to(msg_len + 1).freeze();
1347                        let payload = msg_bytes.slice(5..);
1348                        match parse_first_column_payload_zerocopy(payload) {
1349                            Ok(column) => *first_column = column,
1350                            Err(err) => return self.protocol_desync_error(err),
1351                        }
1352                        return Ok(msg_type);
1353                    }
1354
1355                    let msg_bytes = self.buffer.split_to(msg_len + 1);
1356                    if let Some(msg_type) = self.decode_fast_message_type(msg_bytes)? {
1357                        return Ok(msg_type);
1358                    }
1359                    continue;
1360                }
1361            }
1362
1363            let n = self.read_with_timeout().await?;
1364            if n == 0 {
1365                return self.connection_desync("Connection closed".to_string());
1366            }
1367        }
1368    }
1369
1370    /// FAST receive for fixed 4-column scalar result sets.
1371    #[inline]
1372    pub(crate) async fn recv_fill_first_four_columns_zerocopy_fast(
1373        &mut self,
1374        columns: &mut [Option<Bytes>; 4],
1375    ) -> PgResult<u8> {
1376        loop {
1377            if self.buffer.len() >= 5 {
1378                let msg_len = u32::from_be_bytes([
1379                    self.buffer[1],
1380                    self.buffer[2],
1381                    self.buffer[3],
1382                    self.buffer[4],
1383                ]) as usize;
1384
1385                if msg_len < 4 {
1386                    return self.protocol_desync(format!(
1387                        "Invalid message length: {} (minimum 4)",
1388                        msg_len
1389                    ));
1390                }
1391
1392                if msg_len > MAX_MESSAGE_SIZE {
1393                    return self.protocol_desync(format!(
1394                        "Message too large: {} bytes (max {})",
1395                        msg_len, MAX_MESSAGE_SIZE
1396                    ));
1397                }
1398
1399                if self.buffer.len() > msg_len {
1400                    let msg_type = self.buffer[0];
1401
1402                    if msg_type == b'D' {
1403                        let msg_bytes = self.buffer.split_to(msg_len + 1).freeze();
1404                        let payload = msg_bytes.slice(5..);
1405                        if let Err(err) =
1406                            parse_first_four_columns_payload_zerocopy(payload, columns)
1407                        {
1408                            return self.protocol_desync_error(err);
1409                        }
1410                        return Ok(msg_type);
1411                    }
1412
1413                    let msg_bytes = self.buffer.split_to(msg_len + 1);
1414                    if let Some(msg_type) = self.decode_fast_message_type(msg_bytes)? {
1415                        return Ok(msg_type);
1416                    }
1417                    continue;
1418                }
1419            }
1420
1421            let n = self.read_with_timeout().await?;
1422            if n == 0 {
1423                return self.connection_desync("Connection closed".to_string());
1424            }
1425        }
1426    }
1427
1428    /// ZERO-COPY receive for DataRow.
1429    /// Uses bytes::Bytes for reference-counted slicing instead of Vec copy.
1430    /// Returns: (msg_type, Option<row_data>)
1431    /// For 'D' (DataRow): returns Bytes slices (no copy!)
1432    /// For other types: returns None
1433    #[inline]
1434    pub(crate) async fn recv_data_zerocopy(
1435        &mut self,
1436    ) -> PgResult<(u8, Option<Vec<Option<bytes::Bytes>>>)> {
1437        use bytes::Buf;
1438
1439        loop {
1440            if self.buffer.len() >= 5 {
1441                let msg_len = u32::from_be_bytes([
1442                    self.buffer[1],
1443                    self.buffer[2],
1444                    self.buffer[3],
1445                    self.buffer[4],
1446                ]) as usize;
1447
1448                if msg_len < 4 {
1449                    return self.protocol_desync(format!(
1450                        "Invalid message length: {} (minimum 4)",
1451                        msg_len
1452                    ));
1453                }
1454
1455                if msg_len > MAX_MESSAGE_SIZE {
1456                    return self.protocol_desync(format!(
1457                        "Message too large: {} bytes (max {})",
1458                        msg_len, MAX_MESSAGE_SIZE
1459                    ));
1460                }
1461
1462                if self.buffer.len() > msg_len {
1463                    let msg_type = self.buffer[0];
1464
1465                    // Fast path: DataRow - ZERO-COPY using Bytes
1466                    if msg_type == b'D' {
1467                        // Split off the entire message
1468                        let mut msg_bytes = self.buffer.split_to(msg_len + 1);
1469
1470                        // Skip type byte (1) + length (4) = 5 bytes
1471                        msg_bytes.advance(5);
1472
1473                        if msg_bytes.len() >= 2 {
1474                            let raw_count = msg_bytes.get_i16();
1475                            if raw_count < 0 {
1476                                return self.protocol_desync(format!(
1477                                    "DataRow invalid column count: {}",
1478                                    raw_count
1479                                ));
1480                            }
1481                            let column_count = raw_count as usize;
1482                            if column_count > msg_bytes.remaining() / 4 + 1 {
1483                                return self.protocol_desync(format!(
1484                                    "DataRow claims {} columns but payload is only {} bytes",
1485                                    column_count,
1486                                    msg_bytes.remaining() + 2
1487                                ));
1488                            }
1489                            let mut columns = Vec::with_capacity(column_count);
1490
1491                            for _ in 0..column_count {
1492                                if msg_bytes.remaining() < 4 {
1493                                    return self.protocol_desync(
1494                                        "DataRow truncated: missing column length".into(),
1495                                    );
1496                                }
1497
1498                                let len = msg_bytes.get_i32();
1499
1500                                if len == -1 {
1501                                    columns.push(None);
1502                                } else {
1503                                    if len < -1 {
1504                                        return self.protocol_desync(format!(
1505                                            "DataRow invalid column length: {}",
1506                                            len
1507                                        ));
1508                                    }
1509                                    let len = len as usize;
1510                                    if msg_bytes.remaining() < len {
1511                                        return self.protocol_desync(
1512                                            "DataRow truncated: column data exceeds payload".into(),
1513                                        );
1514                                    }
1515                                    let col_data = msg_bytes.split_to(len).freeze();
1516                                    columns.push(Some(col_data));
1517                                }
1518                            }
1519
1520                            if msg_bytes.remaining() != 0 {
1521                                return self.protocol_desync("DataRow has trailing bytes".into());
1522                            }
1523
1524                            return Ok((msg_type, Some(columns)));
1525                        }
1526                        return self.protocol_desync("DataRow payload too short".into());
1527                    }
1528
1529                    let msg_bytes = self.buffer.split_to(msg_len + 1);
1530                    if let Some(msg_type) = self.decode_fast_message_type(msg_bytes)? {
1531                        return Ok((msg_type, None));
1532                    }
1533                    continue;
1534                }
1535            }
1536
1537            let n = self.read_with_timeout().await?;
1538            if n == 0 {
1539                return self.connection_desync("Connection closed".to_string());
1540            }
1541        }
1542    }
1543
1544    /// ULTRA-FAST receive for 2-column DataRow (id, name pattern).
1545    /// Uses fixed-size array instead of Vec allocation.
1546    /// Returns: (msg_type, Option<(col0, col1)>)
1547    #[inline(always)]
1548    pub(crate) async fn recv_data_ultra(
1549        &mut self,
1550    ) -> PgResult<(u8, Option<(bytes::Bytes, bytes::Bytes)>)> {
1551        use bytes::Buf;
1552
1553        loop {
1554            if self.buffer.len() >= 5 {
1555                let msg_len = u32::from_be_bytes([
1556                    self.buffer[1],
1557                    self.buffer[2],
1558                    self.buffer[3],
1559                    self.buffer[4],
1560                ]) as usize;
1561
1562                if msg_len < 4 {
1563                    return self.protocol_desync(format!(
1564                        "Invalid message length: {} (minimum 4)",
1565                        msg_len
1566                    ));
1567                }
1568
1569                if msg_len > MAX_MESSAGE_SIZE {
1570                    return self.protocol_desync(format!(
1571                        "Message too large: {} bytes (max {})",
1572                        msg_len, MAX_MESSAGE_SIZE
1573                    ));
1574                }
1575
1576                if self.buffer.len() > msg_len {
1577                    let msg_type = self.buffer[0];
1578
1579                    if msg_type == b'D' {
1580                        let mut msg_bytes = self.buffer.split_to(msg_len + 1);
1581                        msg_bytes.advance(5); // Skip type + length
1582
1583                        // Bounds checks to prevent panic on truncated DataRow
1584                        if msg_bytes.remaining() < 2 {
1585                            return self.protocol_desync(
1586                                "DataRow ultra: too short for column count".into(),
1587                            );
1588                        }
1589
1590                        // Read column count (expect 2)
1591                        let col_count = msg_bytes.get_i16();
1592                        if col_count != 2 {
1593                            return self.protocol_desync(format!(
1594                                "DataRow ultra expects exactly 2 columns, got {}",
1595                                col_count
1596                            ));
1597                        }
1598
1599                        if msg_bytes.remaining() < 4 {
1600                            return self.protocol_desync(
1601                                "DataRow ultra: truncated before col0 length".into(),
1602                            );
1603                        }
1604                        let len0 = msg_bytes.get_i32();
1605                        let col0 = if len0 > 0 {
1606                            let len0 = len0 as usize;
1607                            if msg_bytes.remaining() < len0 {
1608                                return self.protocol_desync(
1609                                    "DataRow ultra: col0 data exceeds payload".into(),
1610                                );
1611                            }
1612                            msg_bytes.split_to(len0).freeze()
1613                        } else if len0 == 0 {
1614                            bytes::Bytes::new()
1615                        } else if len0 == -1 {
1616                            return self.protocol_desync(
1617                                "DataRow ultra does not support NULL columns".into(),
1618                            );
1619                        } else {
1620                            return self.protocol_desync(format!(
1621                                "DataRow ultra: invalid col0 length {}",
1622                                len0
1623                            ));
1624                        };
1625
1626                        if msg_bytes.remaining() < 4 {
1627                            return self.protocol_desync(
1628                                "DataRow ultra: truncated before col1 length".into(),
1629                            );
1630                        }
1631                        let len1 = msg_bytes.get_i32();
1632                        let col1 = if len1 > 0 {
1633                            let len1 = len1 as usize;
1634                            if msg_bytes.remaining() < len1 {
1635                                return self.protocol_desync(
1636                                    "DataRow ultra: col1 data exceeds payload".into(),
1637                                );
1638                            }
1639                            msg_bytes.split_to(len1).freeze()
1640                        } else if len1 == 0 {
1641                            bytes::Bytes::new()
1642                        } else if len1 == -1 {
1643                            return self.protocol_desync(
1644                                "DataRow ultra does not support NULL columns".into(),
1645                            );
1646                        } else {
1647                            return self.protocol_desync(format!(
1648                                "DataRow ultra: invalid col1 length {}",
1649                                len1
1650                            ));
1651                        };
1652
1653                        if msg_bytes.remaining() != 0 {
1654                            return self.protocol_desync(
1655                                "DataRow ultra: trailing bytes after expected columns".into(),
1656                            );
1657                        }
1658
1659                        return Ok((msg_type, Some((col0, col1))));
1660                    }
1661
1662                    let msg_bytes = self.buffer.split_to(msg_len + 1);
1663                    if let Some(msg_type) = self.decode_fast_message_type(msg_bytes)? {
1664                        return Ok((msg_type, None));
1665                    }
1666                    continue;
1667                }
1668            }
1669
1670            let n = self.read_with_timeout().await?;
1671            if n == 0 {
1672                return self.connection_desync("Connection closed".to_string());
1673            }
1674        }
1675    }
1676}
1677
1678#[cfg(test)]
1679mod tests {
1680    use super::*;
1681
1682    #[cfg(unix)]
1683    fn test_conn() -> PgConnection {
1684        use crate::driver::connection::StatementCache;
1685        use crate::driver::stream::PgStream;
1686        use std::collections::{HashMap, VecDeque};
1687        use std::num::NonZeroUsize;
1688        use tokio::net::UnixStream;
1689
1690        let (unix_stream, _peer) = UnixStream::pair().expect("unix stream pair");
1691        PgConnection {
1692            stream: PgStream::Unix(unix_stream),
1693            buffer: BytesMut::with_capacity(1024),
1694            write_buf: BytesMut::with_capacity(1024),
1695            sql_buf: BytesMut::with_capacity(256),
1696            params_buf: Vec::new(),
1697            prepared_statements: HashMap::new(),
1698            stmt_cache: StatementCache::new(NonZeroUsize::new(2).expect("non-zero")),
1699            column_info_cache: HashMap::new(),
1700            process_id: 0,
1701            cancel_key_bytes: Vec::new(),
1702            requested_protocol_minor: PgConnection::default_protocol_minor(),
1703            negotiated_protocol_minor: PgConnection::default_protocol_minor(),
1704            notifications: VecDeque::new(),
1705            replication_stream_active: false,
1706            replication_mode_enabled: false,
1707            last_replication_wal_end: None,
1708            io_desynced: false,
1709            pending_statement_closes: Vec::new(),
1710            draining_statement_closes: false,
1711        }
1712    }
1713
1714    fn build_data_row_payload(columns: &[Option<&[u8]>]) -> Bytes {
1715        let mut payload = Vec::new();
1716        payload.extend_from_slice(&(columns.len() as i16).to_be_bytes());
1717        for column in columns {
1718            match column {
1719                Some(bytes) => {
1720                    payload.extend_from_slice(&(bytes.len() as i32).to_be_bytes());
1721                    payload.extend_from_slice(bytes);
1722                }
1723                None => payload.extend_from_slice(&(-1i32).to_be_bytes()),
1724            }
1725        }
1726        Bytes::from(payload)
1727    }
1728
1729    fn push_data_row_frame(conn: &mut PgConnection, payload: &[u8]) {
1730        let msg_len = payload.len() + 4;
1731        conn.buffer.extend_from_slice(b"D");
1732        conn.buffer
1733            .extend_from_slice(&(msg_len as u32).to_be_bytes());
1734        conn.buffer.extend_from_slice(payload);
1735    }
1736
1737    fn push_one_column_datarow_without_column_length(conn: &mut PgConnection) {
1738        push_data_row_frame(conn, &[0, 1]);
1739    }
1740
1741    fn assert_protocol_error_contains(err: PgError, expected: &str) {
1742        match err {
1743            PgError::Protocol(msg) => assert!(
1744                msg.contains(expected),
1745                "expected protocol error containing {expected:?}, got {msg:?}"
1746            ),
1747            err => panic!("expected protocol error containing {expected:?}, got {err:?}"),
1748        }
1749    }
1750
1751    #[test]
1752    fn parse_first_four_columns_payload_zerocopy_reads_values() {
1753        let payload = build_data_row_payload(&[Some(b"10"), None, Some(b"30"), Some(b"")]);
1754        let mut columns = [None, None, None, None];
1755
1756        parse_first_four_columns_payload_zerocopy(payload, &mut columns).unwrap();
1757
1758        assert_eq!(columns[0].as_deref(), Some(&b"10"[..]));
1759        assert_eq!(columns[1].as_deref(), None);
1760        assert_eq!(columns[2].as_deref(), Some(&b"30"[..]));
1761        assert_eq!(columns[3].as_deref(), Some(&b""[..]));
1762    }
1763
1764    #[test]
1765    fn parse_first_four_columns_payload_zerocopy_rejects_wrong_arity() {
1766        let payload = build_data_row_payload(&[Some(b"1"), Some(b"2"), Some(b"3")]);
1767        let mut columns = [None, None, None, None];
1768
1769        let err = parse_first_four_columns_payload_zerocopy(payload, &mut columns).unwrap_err();
1770
1771        assert!(
1772            err.to_string()
1773                .contains("fast-path expects exactly 4 columns")
1774        );
1775    }
1776
1777    #[cfg(unix)]
1778    #[tokio::test]
1779    async fn recv_data_zerocopy_rejects_datarow_length_4() {
1780        let mut conn = test_conn();
1781        conn.buffer.extend_from_slice(&[b'D', 0, 0, 0, 4]);
1782
1783        let err = conn.recv_data_zerocopy().await.unwrap_err();
1784
1785        assert!(err.to_string().contains("DataRow payload too short"));
1786        assert!(conn.is_io_desynced());
1787    }
1788
1789    #[cfg(unix)]
1790    #[tokio::test]
1791    async fn recv_data_zerocopy_rejects_datarow_length_5() {
1792        let mut conn = test_conn();
1793        conn.buffer.extend_from_slice(&[b'D', 0, 0, 0, 5, 0]);
1794
1795        let err = conn.recv_data_zerocopy().await.unwrap_err();
1796
1797        assert!(err.to_string().contains("DataRow payload too short"));
1798        assert!(conn.is_io_desynced());
1799    }
1800
1801    #[cfg(unix)]
1802    #[tokio::test]
1803    async fn recv_with_data_fast_desyncs_on_malformed_datarow() {
1804        let mut conn = test_conn();
1805        push_one_column_datarow_without_column_length(&mut conn);
1806
1807        let err = conn.recv_with_data_fast().await.unwrap_err();
1808
1809        assert_protocol_error_contains(err, "DataRow truncated");
1810        assert!(conn.is_io_desynced());
1811    }
1812
1813    #[cfg(unix)]
1814    #[tokio::test]
1815    async fn recv_fill_data_row_fast_desyncs_on_malformed_datarow() {
1816        let mut conn = test_conn();
1817        let mut row = Vec::new();
1818        push_one_column_datarow_without_column_length(&mut conn);
1819
1820        let err = conn.recv_fill_data_row_fast(&mut row).await.unwrap_err();
1821
1822        assert_protocol_error_contains(err, "DataRow truncated");
1823        assert!(conn.is_io_desynced());
1824    }
1825
1826    #[cfg(unix)]
1827    #[tokio::test]
1828    async fn recv_fill_zerocopy_row_fast_desyncs_on_malformed_datarow() {
1829        let mut conn = test_conn();
1830        let mut row = PgBytesRow::default();
1831        push_one_column_datarow_without_column_length(&mut conn);
1832
1833        let err = conn
1834            .recv_fill_zerocopy_row_fast(&mut row)
1835            .await
1836            .unwrap_err();
1837
1838        assert_protocol_error_contains(err, "DataRow truncated");
1839        assert!(conn.is_io_desynced());
1840    }
1841
1842    #[cfg(unix)]
1843    #[tokio::test]
1844    async fn recv_fill_first_column_zerocopy_fast_desyncs_on_malformed_datarow() {
1845        let mut conn = test_conn();
1846        let mut first_column = None;
1847        push_one_column_datarow_without_column_length(&mut conn);
1848
1849        let err = conn
1850            .recv_fill_first_column_zerocopy_fast(&mut first_column)
1851            .await
1852            .unwrap_err();
1853
1854        assert_protocol_error_contains(err, "DataRow truncated");
1855        assert!(conn.is_io_desynced());
1856    }
1857
1858    #[cfg(unix)]
1859    #[tokio::test]
1860    async fn recv_fill_first_four_columns_zerocopy_fast_desyncs_on_malformed_datarow() {
1861        let mut conn = test_conn();
1862        let mut columns = [None, None, None, None];
1863        push_one_column_datarow_without_column_length(&mut conn);
1864
1865        let err = conn
1866            .recv_fill_first_four_columns_zerocopy_fast(&mut columns)
1867            .await
1868            .unwrap_err();
1869
1870        assert_protocol_error_contains(err, "DataRow fast-path expects exactly 4 columns");
1871        assert!(conn.is_io_desynced());
1872    }
1873
1874    #[cfg(unix)]
1875    #[tokio::test]
1876    async fn recv_data_ultra_desyncs_on_malformed_datarow() {
1877        let mut conn = test_conn();
1878        push_one_column_datarow_without_column_length(&mut conn);
1879
1880        let err = conn.recv_data_ultra().await.unwrap_err();
1881
1882        assert_protocol_error_contains(err, "DataRow ultra expects exactly 2 columns");
1883        assert!(conn.is_io_desynced());
1884    }
1885
1886    #[cfg(unix)]
1887    #[tokio::test]
1888    async fn recv_msg_type_fast_rejects_malformed_ready_for_query() {
1889        let mut conn = test_conn();
1890        conn.buffer.extend_from_slice(&[b'Z', 0, 0, 0, 5, b'X']);
1891
1892        let err = conn.recv_msg_type_fast().await.unwrap_err();
1893
1894        assert!(err.to_string().contains("Unknown transaction status"));
1895        assert!(conn.is_io_desynced());
1896    }
1897
1898    #[cfg(unix)]
1899    #[tokio::test]
1900    async fn recv_msg_type_fast_rejects_malformed_command_complete() {
1901        let mut conn = test_conn();
1902        conn.buffer.extend_from_slice(&[
1903            b'C', 0, 0, 0, 12, b'S', b'E', b'L', b'E', b'C', b'T', b' ', b'1',
1904        ]);
1905
1906        let err = conn.recv_msg_type_fast().await.unwrap_err();
1907
1908        assert!(
1909            err.to_string()
1910                .contains("CommandComplete missing null terminator")
1911        );
1912        assert!(conn.is_io_desynced());
1913    }
1914}