1use std::sync::Arc;
8
9use crate::protocol::{BackendMessage, FrontendMessage, TransactionStatus};
10use fallible_iterator::FallibleIterator;
11
12use crate::connection::Connection;
13use crate::error::{PgError, PgServerError, Result};
14use crate::query::result::{CommandTag, ExecuteResult, QueryResult};
15use crate::query::row::{FieldDescription, Row};
16use crate::transport::AsyncTransport;
17
18#[cfg(feature = "tracing")]
19use crate::tracing_ext::{truncate_str, TARGET_QUERY};
20
21pub mod cache;
22pub mod cursor;
23pub mod params;
24pub mod pipeline;
25pub mod prepared;
26pub mod result;
27pub mod row;
28pub mod stream;
29
30pub use cache::StatementCache;
32pub use cursor::Cursor;
33pub use cursor::CursorStream;
34pub use pipeline::{Pipeline, PipelineResult};
35pub use prepared::PreparedStatement;
36
37#[derive(Debug, Clone)]
47#[non_exhaustive]
48pub struct Notice {
49 inner: PgServerError,
51}
52
53pub type NoticeHandler = Box<dyn Fn(&Notice) + Send + Sync>;
55
56impl Notice {
57 pub fn from_fields(fields: &crate::protocol::backend::NoticeResponseBody) -> Result<Self> {
59 let inner = PgServerError::from_notice_body(fields).map_err(PgError::Io)?;
60 Ok(Self { inner })
61 }
62
63 pub fn severity(&self) -> &str {
67 &self.inner.severity
68 }
69
70 pub fn code(&self) -> &str {
72 &self.inner.code
73 }
74
75 pub fn message(&self) -> &str {
77 &self.inner.message
78 }
79
80 pub fn detail(&self) -> Option<&str> {
82 self.inner.detail.as_deref()
83 }
84
85 pub fn hint(&self) -> Option<&str> {
87 self.inner.hint.as_deref()
88 }
89
90 pub fn as_server_error(&self) -> &PgServerError {
95 &self.inner
96 }
97}
98
99impl std::fmt::Display for Notice {
100 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
101 write!(
102 f,
103 "{}: {} (SQLSTATE {})",
104 self.inner.severity, self.inner.message, self.inner.code
105 )?;
106 if let Some(detail) = &self.inner.detail {
107 write!(f, "\nDETAIL: {}", detail)?;
108 }
109 if let Some(hint) = &self.inner.hint {
110 write!(f, "\nHINT: {}", hint)?;
111 }
112 Ok(())
113 }
114}
115
116impl Connection {
121 #[must_use = "query results should be checked for errors"]
136 pub async fn query(&mut self, sql: &str) -> Result<QueryResult> {
137 let mut stream = self.query_stream(sql).await?;
138 let mut rows = Vec::new();
139 while let Some(row) = stream.next().await? {
140 rows.push(row);
141 }
142 let columns = stream.columns().map(|c| c.to_vec()).unwrap_or_default();
143 let command_tag = stream.command_tag().cloned().unwrap_or_default();
144 Ok(QueryResult::new(rows, command_tag, Arc::new(columns)))
145 }
146
147 #[must_use = "execute results should be checked for errors"]
151 pub async fn execute(&mut self, sql: &str) -> Result<ExecuteResult> {
152 let result = self.query(sql).await?;
153 Ok(ExecuteResult::new(result.command_tag().clone()))
154 }
155
156 #[must_use = "query results should be checked for errors"]
160 pub async fn query_one(&mut self, sql: &str) -> Result<Option<Row>> {
161 let result = self.query(sql).await?;
162 Ok(result.into_rows().into_iter().next())
163 }
164
165 #[must_use = "query results should be checked for errors"]
170 pub async fn query_each<F>(&mut self, sql: &str, mut f: F) -> Result<CommandTag>
171 where
172 F: FnMut(Row) -> Result<()>,
173 {
174 let mut stream = self.query_stream(sql).await?;
175 while let Some(row) = stream.next().await? {
176 f(row)?;
177 }
178 stream
179 .command_tag()
180 .cloned()
181 .ok_or_else(|| PgError::InvalidState("stream ended without command tag".into()))
182 }
183
184 #[must_use = "batch results should be checked for errors"]
188 pub async fn batch_execute(&mut self, sql: &str) -> Result<Vec<QueryResult>> {
189 self.transition(ConnectionState::ActiveSimpleQuery)?;
190
191 self.codec
192 .send(
193 &mut self.transport,
194 &FrontendMessage::Query { sql: sql.into() },
195 )
196 .await?;
197
198 let mut results = Vec::new();
199 let mut current_columns: Option<Arc<Vec<FieldDescription>>> = None;
200 let mut current_rows: Vec<Row> = Vec::new();
201
202 loop {
203 let msg = self.codec.read_message(&mut self.transport).await?;
204 if self.handle_async_message(&msg) {
205 continue;
206 }
207 match msg {
208 BackendMessage::RowDescription(body) => {
209 current_columns = Some(Arc::new(read_row_description(body)?));
210 current_rows.clear();
211 }
212 BackendMessage::DataRow(body) => {
213 let values = read_data_row(body)?;
214 current_rows.push(Row::new(
215 current_columns.clone().unwrap_or_default(),
216 values,
217 ));
218 }
219 BackendMessage::CommandComplete(body) => {
220 let tag = CommandTag::new(body.tag().unwrap_or("").into());
221 results.push(QueryResult::new(
222 std::mem::take(&mut current_rows),
223 tag,
224 current_columns.take().unwrap_or_default(),
225 ));
226 }
227 BackendMessage::EmptyQueryResponse => {
228 results.push(QueryResult::new(
229 Vec::new(),
230 CommandTag::new("".into()),
231 Arc::new(Vec::new()),
232 ));
233 }
234 BackendMessage::ErrorResponse(body) => {
235 let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
236 self.read_until_ready().await?;
237 self.state = ConnectionState::Idle;
238 return Err(PgError::Server(Box::new(server_err)));
239 }
240 BackendMessage::ReadyForQuery(body) => {
241 self.transaction_status = TransactionStatus::from_u8(body.status())
242 .unwrap_or(TransactionStatus::Idle);
243 self.state = ConnectionState::Idle;
244 break;
245 }
246 _ => {}
247 }
248 }
249
250 Ok(results)
251 }
252}
253
254impl Connection {
259 #[must_use = "stream results should be checked for errors"]
275 pub async fn query_stream(&mut self, sql: &str) -> Result<stream::RowStream<'_>> {
276 #[cfg(feature = "tracing")]
277 tracing::debug!(target: TARGET_QUERY, sql_len = sql.len(), sql_truncated = %truncate_str(sql, 200), protocol = "simple", "Executing simple query");
278 #[cfg(feature = "tracing")]
279 tracing::trace!(target: TARGET_QUERY, sql = %sql, "Full SQL text");
280
281 if self.needs_recovery {
283 self.recover().await?;
284 }
285
286 self.transition(ConnectionState::Streaming)?;
287
288 self.codec
289 .send(
290 &mut self.transport,
291 &FrontendMessage::Query { sql: sql.into() },
292 )
293 .await?;
294
295 Ok(stream::RowStream::new_simple(self))
296 }
297
298 #[must_use = "stream results should be checked for errors"]
300 pub async fn query_prepared_stream(
301 &mut self,
302 stmt: &PreparedStatement,
303 params: &[&dyn crate::types::ToSql],
304 ) -> Result<stream::RowStream<'_>> {
305 #[cfg(feature = "tracing")]
306 tracing::debug!(target: TARGET_QUERY, sql_len = stmt.sql.len(), sql_truncated = %truncate_str(&stmt.sql, 200), statement = %stmt.name, "Executing prepared statement");
307
308 if self.needs_recovery {
309 self.recover().await?;
310 }
311
312 self.transition(ConnectionState::Streaming)?;
313
314 let param_values = params::encode_params_binary(params, &stmt.param_types)?;
315
316 self.codec
318 .encode_and_write(
319 &mut self.transport,
320 &FrontendMessage::Bind {
321 portal: String::new(),
322 statement: stmt.name.clone(),
323 param_formats: vec![crate::protocol::FormatCode::Binary],
324 params: param_values,
325 result_formats: vec![crate::protocol::FormatCode::Binary],
326 },
327 )
328 .await?;
329
330 self.codec
332 .encode_and_write(
333 &mut self.transport,
334 &FrontendMessage::Describe {
335 variant: b'P',
336 name: String::new(),
337 },
338 )
339 .await?;
340
341 self.codec
343 .encode_and_write(
344 &mut self.transport,
345 &FrontendMessage::Execute {
346 portal: String::new(),
347 max_rows: 0,
348 },
349 )
350 .await?;
351
352 self.codec
354 .encode_and_write(&mut self.transport, &FrontendMessage::Sync)
355 .await?;
356
357 self.transport.flush().await.map_err(PgError::Transport)?;
359
360 Ok(stream::RowStream::new_extended_with_columns(
362 self,
363 stmt.columns.clone(),
364 ))
365 }
366
367 #[must_use = "query results should be checked for errors"]
372 pub async fn query_each_async<F, Fut>(&mut self, sql: &str, mut f: F) -> Result<CommandTag>
373 where
374 F: FnMut(Row) -> Fut,
375 Fut: std::future::Future<Output = Result<()>>,
376 {
377 let mut stream = self.query_stream(sql).await?;
378 while let Some(row) = stream.next().await? {
379 f(row).await?;
380 }
381 stream
382 .command_tag()
383 .cloned()
384 .ok_or_else(|| PgError::InvalidState("stream ended without command tag".into()))
385 }
386}
387
388use crate::connection::ConnectionState;
393
394pub(crate) fn read_row_description(
396 body: crate::protocol::backend::RowDescriptionBody,
397) -> Result<Vec<FieldDescription>> {
398 let mut fields = Vec::new();
399 let mut iter = body.fields();
400 while let Some(field) = iter.next()? {
401 fields.push(FieldDescription::new(
402 field.name().into(),
403 field.table_oid(),
404 field.column_id(),
405 field.type_oid(),
406 field.type_size(),
407 field.type_modifier(),
408 field.format(),
409 ));
410 }
411 Ok(fields)
412}
413
414pub(crate) fn read_data_row(
416 body: crate::protocol::backend::DataRowBody,
417) -> Result<Vec<Option<Vec<u8>>>> {
418 let buf = body.buffer();
419 let mut values = Vec::new();
420 let mut iter = body.ranges();
421 while let Some(range) = iter.next()? {
422 values.push(range.map(|r| buf[r].to_vec()));
423 }
424 Ok(values)
425}
426
427#[cfg(test)]
432mod tests {
433 use super::*;
434 use crate::auth::{Codec, ServerParams};
435 use crate::config::Config;
436 use crate::connection::ConnectionState;
437 use crate::transport::{BufferedTransport, ClientTransport, MockTransport, PgTransport};
438 use std::collections::VecDeque;
439
440 fn make_connection(read_data: Vec<u8>) -> Connection {
441 let transport = PgTransport::Plain(BufferedTransport::new(ClientTransport::Mock(
442 MockTransport::new(read_data),
443 )));
444 Connection {
445 transport,
446 codec: Codec::new(),
447 server_params: ServerParams::default(),
448 state: ConnectionState::Idle,
449 config: Config::new(),
450 transaction_status: TransactionStatus::Idle,
451 notification_queue: VecDeque::new(),
452 notice_handler: None,
453 statement_counter: 0,
454 needs_recovery: false,
455 health: crate::reconnect::session::ConnectionHealth::new(),
456 session_state: crate::reconnect::session::SessionState::new(),
457 }
458 }
459
460 pub(crate) fn build_row_description_msg(fields: &[(&str, u32)]) -> Vec<u8> {
461 let mut buf = vec![b'T'];
462 let mut body = Vec::new();
463 body.extend_from_slice(&(fields.len() as i16).to_be_bytes());
465 for (name, type_oid) in fields {
466 body.extend_from_slice(name.as_bytes());
467 body.push(0);
468 body.extend_from_slice(&0u32.to_be_bytes()); body.extend_from_slice(&0i16.to_be_bytes()); body.extend_from_slice(&type_oid.to_be_bytes()); body.extend_from_slice(&(-1i16).to_be_bytes()); body.extend_from_slice(&(-1i32).to_be_bytes()); body.extend_from_slice(&0i16.to_be_bytes()); }
475 let len = (body.len() + 4) as i32;
476 buf.extend_from_slice(&len.to_be_bytes());
477 buf.extend_from_slice(&body);
478 buf
479 }
480
481 fn build_data_row_msg(values: &[Option<&str>]) -> Vec<u8> {
482 let mut buf = vec![b'D'];
483 let mut body = Vec::new();
484 body.extend_from_slice(&(values.len() as i16).to_be_bytes());
486 for val in values {
487 match val {
488 Some(v) => {
489 let bytes = v.as_bytes();
490 body.extend_from_slice(&(bytes.len() as i32).to_be_bytes());
491 body.extend_from_slice(bytes);
492 }
493 None => {
494 body.extend_from_slice(&(-1i32).to_be_bytes());
495 }
496 }
497 }
498 let len = (body.len() + 4) as i32;
499 buf.extend_from_slice(&len.to_be_bytes());
500 buf.extend_from_slice(&body);
501 buf
502 }
503
504 fn build_command_complete_msg(tag: &str) -> Vec<u8> {
505 let mut buf = vec![b'C'];
506 let mut body = Vec::new();
507 body.extend_from_slice(tag.as_bytes());
508 body.push(0);
509 let len = (body.len() + 4) as i32;
510 buf.extend_from_slice(&len.to_be_bytes());
511 buf.extend_from_slice(&body);
512 buf
513 }
514
515 fn build_ready_for_query(status: u8) -> Vec<u8> {
516 vec![b'Z', 0, 0, 0, 5, status]
517 }
518
519 #[tokio::test]
520 async fn test_query_basic() {
521 let mut data = Vec::new();
522 data.extend_from_slice(&build_row_description_msg(&[
523 ("id", crate::types::INT4_OID),
524 ("name", crate::types::TEXT_OID),
525 ]));
526 data.extend_from_slice(&build_data_row_msg(&[Some("1"), Some("alice")]));
527 data.extend_from_slice(&build_data_row_msg(&[Some("2"), Some("bob")]));
528 data.extend_from_slice(&build_command_complete_msg("SELECT 2"));
529 data.extend_from_slice(&build_ready_for_query(b'I'));
530
531 let mut conn = make_connection(data);
532 let result = conn.query("SELECT id, name FROM users").await.unwrap();
533 assert_eq!(result.len(), 2);
534 let id: i32 = result.rows()[0].get(0).unwrap();
535 assert_eq!(id, 1);
536 let name: String = result.rows()[0].get(1).unwrap();
537 assert_eq!(name, "alice");
538 }
539
540 #[tokio::test]
541 async fn test_execute_no_rows() {
542 let mut data = Vec::new();
543 data.extend_from_slice(&build_command_complete_msg("INSERT 0 3"));
544 data.extend_from_slice(&build_ready_for_query(b'I'));
545
546 let mut conn = make_connection(data);
547 let result = conn
548 .execute("INSERT INTO users (name) VALUES ('alice')")
549 .await
550 .unwrap();
551 assert_eq!(result.rows_affected(), Some(3));
552 }
553
554 #[tokio::test]
555 async fn test_query_one() {
556 let mut data = Vec::new();
557 data.extend_from_slice(&build_row_description_msg(&[(
558 "id",
559 crate::types::INT4_OID,
560 )]));
561 data.extend_from_slice(&build_data_row_msg(&[Some("42")]));
562 data.extend_from_slice(&build_command_complete_msg("SELECT 1"));
563 data.extend_from_slice(&build_ready_for_query(b'I'));
564
565 let mut conn = make_connection(data);
566 let row = conn.query_one("SELECT 42").await.unwrap();
567 assert!(row.is_some());
568 let id: i32 = row.unwrap().get(0).unwrap();
569 assert_eq!(id, 42);
570 }
571
572 #[tokio::test]
573 async fn test_query_empty() {
574 let mut data = Vec::new();
575 data.extend_from_slice(&build_row_description_msg(&[(
576 "id",
577 crate::types::INT4_OID,
578 )]));
579 data.extend_from_slice(&build_command_complete_msg("SELECT 0"));
580 data.extend_from_slice(&build_ready_for_query(b'I'));
581
582 let mut conn = make_connection(data);
583 let result = conn
584 .query("SELECT id FROM users WHERE false")
585 .await
586 .unwrap();
587 assert!(result.is_empty());
588 }
589
590 #[tokio::test]
591 async fn test_query_error() {
592 let mut data = Vec::new();
593 let mut err = vec![b'E', 0, 0, 0, 26];
595 err.extend_from_slice(b"S");
596 err.extend_from_slice(b"ERROR\0");
597 err.extend_from_slice(b"M");
598 err.extend_from_slice(b"syntax error\0");
599 err.push(0);
600 data.extend_from_slice(&err);
601 data.extend_from_slice(&build_ready_for_query(b'I'));
603
604 let mut conn = make_connection(data);
605 let result = conn.query("BAD SQL").await;
606 assert!(result.is_err());
607 assert!(conn.is_idle());
608 }
609
610 #[tokio::test]
611 async fn test_query_each() {
612 let mut data = Vec::new();
613 data.extend_from_slice(&build_row_description_msg(&[(
614 "val",
615 crate::types::INT4_OID,
616 )]));
617 data.extend_from_slice(&build_data_row_msg(&[Some("10")]));
618 data.extend_from_slice(&build_data_row_msg(&[Some("20")]));
619 data.extend_from_slice(&build_command_complete_msg("SELECT 2"));
620 data.extend_from_slice(&build_ready_for_query(b'I'));
621
622 let mut conn = make_connection(data);
623 let mut sum = 0i32;
624 let tag = conn
625 .query_each("SELECT val FROM nums", |row| {
626 let v: i32 = row.get(0)?;
627 sum += v;
628 Ok(())
629 })
630 .await
631 .unwrap();
632 assert_eq!(sum, 30);
633 assert_eq!(tag.as_str(), "SELECT 2");
634 }
635
636 #[tokio::test]
637 async fn test_batch_execute() {
638 let mut data = Vec::new();
639 data.extend_from_slice(&build_row_description_msg(&[(
641 "id",
642 crate::types::INT4_OID,
643 )]));
644 data.extend_from_slice(&build_data_row_msg(&[Some("1")]));
645 data.extend_from_slice(&build_command_complete_msg("SELECT 1"));
646 data.extend_from_slice(&build_command_complete_msg("INSERT 0 1"));
648 data.extend_from_slice(&build_ready_for_query(b'I'));
650
651 let mut conn = make_connection(data);
652 let results = conn
653 .batch_execute("SELECT 1; INSERT INTO t VALUES (1)")
654 .await
655 .unwrap();
656 assert_eq!(results.len(), 2);
657 assert_eq!(results[0].len(), 1);
658 assert_eq!(results[1].len(), 0);
659 assert_eq!(results[1].rows_affected(), Some(1));
660 }
661
662 #[tokio::test]
663 async fn test_null_handling() {
664 let mut data = Vec::new();
665 data.extend_from_slice(&build_row_description_msg(&[(
666 "val",
667 crate::types::INT4_OID,
668 )]));
669 data.extend_from_slice(&build_data_row_msg(&[None]));
670 data.extend_from_slice(&build_command_complete_msg("SELECT 1"));
671 data.extend_from_slice(&build_ready_for_query(b'I'));
672
673 let mut conn = make_connection(data);
674 let result = conn.query("SELECT NULL").await.unwrap();
675 assert_eq!(result.len(), 1);
676 assert!(result.rows()[0].is_null(0));
677 }
678}