1use 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; const DEFAULT_READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
15const 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 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 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 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 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 pub async fn recv(&mut self) -> PgResult<BackendMessage> {
771 loop {
772 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 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 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; }
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 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 #[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 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 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 #[inline]
1070 pub fn buffer_bytes(&mut self, bytes: &[u8]) {
1071 self.write_buf.extend_from_slice(bytes);
1072 }
1073
1074 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 #[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 #[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 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 #[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 #[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 #[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 #[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 #[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 if msg_type == b'D' {
1467 let mut msg_bytes = self.buffer.split_to(msg_len + 1);
1469
1470 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 #[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); if msg_bytes.remaining() < 2 {
1585 return self.protocol_desync(
1586 "DataRow ultra: too short for column count".into(),
1587 );
1588 }
1589
1590 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}