Skip to main content

wasi_pg_client/query/
mod.rs

1//! Query protocol implementation — simple and extended.
2//!
3//! This module provides [`Connection`] methods for executing queries via both
4//! the Simple Query Protocol (text-only, no parameters) and the Extended Query
5//! Protocol (parameterized, prepared statements, binary data).
6
7use 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
30// Re-export commonly used types at the `query` level.
31pub use cache::StatementCache;
32pub use cursor::Cursor;
33pub use cursor::CursorStream;
34pub use pipeline::{Pipeline, PipelineResult};
35pub use prepared::PreparedStatement;
36
37// ---------------------------------------------------------------------------
38// Notice
39// ---------------------------------------------------------------------------
40
41/// A notice (non-fatal warning) sent by the PostgreSQL server.
42///
43/// Wraps a [`PgServerError`] which contains all fields from the PostgreSQL
44/// `NoticeResponse` message. Convenience accessor methods are provided for
45/// the most commonly used fields.
46#[derive(Debug, Clone)]
47#[non_exhaustive]
48pub struct Notice {
49    /// The underlying server error/notice with all fields.
50    inner: PgServerError,
51}
52
53/// A callback that is invoked whenever the server sends a [`Notice`].
54pub type NoticeHandler = Box<dyn Fn(&Notice) + Send + Sync>;
55
56impl Notice {
57    /// Parse a [`Notice`] from a [`NoticeResponseBody`](crate::protocol::backend::NoticeResponseBody).
58    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    /// Returns the severity level.
64    ///
65    /// One of: `ERROR`, `FATAL`, `PANIC`, `WARNING`, `NOTICE`, `DEBUG`, `INFO`, `LOG`.
66    pub fn severity(&self) -> &str {
67        &self.inner.severity
68    }
69
70    /// Returns the SQLSTATE error code.
71    pub fn code(&self) -> &str {
72        &self.inner.code
73    }
74
75    /// Returns the primary human-readable message.
76    pub fn message(&self) -> &str {
77        &self.inner.message
78    }
79
80    /// Returns the detailed secondary message, if any.
81    pub fn detail(&self) -> Option<&str> {
82        self.inner.detail.as_deref()
83    }
84
85    /// Returns the suggestion for resolution, if any.
86    pub fn hint(&self) -> Option<&str> {
87        self.inner.hint.as_deref()
88    }
89
90    /// Returns a reference to the underlying [`PgServerError`].
91    ///
92    /// Use this to access all fields (position, schema, table, column,
93    /// constraint, etc.) that are not exposed by the convenience methods.
94    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
116// ---------------------------------------------------------------------------
117// Connection query methods
118// ---------------------------------------------------------------------------
119
120impl Connection {
121    /// Execute a SQL query that returns rows.
122    ///
123    /// This is a convenience method that collects all rows into a
124    /// [`QueryResult`]. For streaming results one row at a time, use
125    /// [`Connection::query_stream`] instead.
126    ///
127    /// # Example
128    /// ```ignore
129    /// let result = conn.query("SELECT id, name FROM users").await?;
130    /// for row in result.iter() {
131    ///     let id: i32 = row.get(0)?;
132    ///     let name: String = row.get(1)?;
133    /// }
134    /// ```
135    #[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    /// Execute a SQL statement that does not return rows.
148    ///
149    /// Returns the number of rows affected where applicable.
150    #[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    /// Execute a query and return at most one row.
157    ///
158    /// Returns `None` if the query returns zero rows.
159    #[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    /// Execute a query, invoking `f` for each row as it arrives.
166    ///
167    /// This avoids buffering all rows in memory, which is useful for large
168    /// result sets.
169    #[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    /// Execute multiple statements separated by semicolons.
185    ///
186    /// Returns a [`QueryResult`] for each statement that produces one.
187    #[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
254// ---------------------------------------------------------------------------
255// Streaming query methods
256// ---------------------------------------------------------------------------
257
258impl Connection {
259    /// Execute a simple query and return a stream of rows.
260    ///
261    /// This is the primary streaming API. Rows are fetched from the server
262    /// one at a time as the consumer calls `next()` on the returned stream.
263    /// Memory usage is O(1) per row regardless of result set size.
264    ///
265    /// # Example
266    ///
267    /// ```ignore
268    /// let mut stream = conn.query_stream("SELECT id, name FROM users").await?;
269    /// while let Some(row) = stream.next().await? {
270    ///     let id: i32 = row.get(0)?;
271    ///     let name: String = row.get(1)?;
272    /// }
273    /// ```
274    #[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        // Ensure connection is in a clean state
282        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    /// Execute a prepared statement and return a stream of rows.
299    #[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        // Bind (unnamed portal, named statement)
317        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        // Describe portal
331        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        // Execute
342        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        // Sync
353        self.codec
354            .encode_and_write(&mut self.transport, &FrontendMessage::Sync)
355            .await?;
356
357        // Flush the entire batch
358        self.transport.flush().await.map_err(PgError::Transport)?;
359
360        // Use the prepared statement's column metadata
361        Ok(stream::RowStream::new_extended_with_columns(
362            self,
363            stmt.columns.clone(),
364        ))
365    }
366
367    /// Execute a query and process rows with an async callback (streaming).
368    ///
369    /// Like `query_each()` but the callback is async, allowing async operations
370    /// (e.g., writing to another connection) per row.
371    #[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
388// ---------------------------------------------------------------------------
389// Internal helpers
390// ---------------------------------------------------------------------------
391
392use crate::connection::ConnectionState;
393
394/// Convert a `RowDescriptionBody` into our `Vec<FieldDescription>`.
395pub(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
414/// Convert a `DataRowBody` into a `Vec<Option<Vec<u8>>>`.
415pub(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// ---------------------------------------------------------------------------
428// Tests
429// ---------------------------------------------------------------------------
430
431#[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        // field count
464        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()); // table_oid
469            body.extend_from_slice(&0i16.to_be_bytes()); // column_id
470            body.extend_from_slice(&type_oid.to_be_bytes()); // type_oid
471            body.extend_from_slice(&(-1i16).to_be_bytes()); // type_size
472            body.extend_from_slice(&(-1i32).to_be_bytes()); // type_modifier
473            body.extend_from_slice(&0i16.to_be_bytes()); // format
474        }
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        // column count
485        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        // ErrorResponse
594        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        // ReadyForQuery
602        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        // First result set
640        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        // Second result set (no rows)
647        data.extend_from_slice(&build_command_complete_msg("INSERT 0 1"));
648        // ReadyForQuery
649        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}