Skip to main content

wasi_pg_client/query/
prepared.rs

1//! Prepared statement management.
2//!
3//! This module provides [`PreparedStatement`] and the [`Connection`] methods
4//! for creating and closing prepared statements via the Extended Query Protocol.
5
6use std::sync::Arc;
7
8use crate::protocol::{BackendMessage, FrontendMessage};
9use fallible_iterator::FallibleIterator;
10
11use crate::connection::{Connection, ConnectionState};
12use crate::error::{PgError, PgServerError, Result};
13use crate::query::read_row_description;
14use crate::query::row::FieldDescription;
15use crate::transport::AsyncTransport;
16
17// ---------------------------------------------------------------------------
18// PreparedStatement
19// ---------------------------------------------------------------------------
20
21/// A server-side prepared statement.
22///
23/// Created via [`Connection::prepare`], a prepared statement can be executed
24/// repeatedly with different parameters via [`Connection::query_prepared`].
25#[derive(Debug, Clone)]
26#[non_exhaustive]
27pub struct PreparedStatement {
28    pub(crate) name: String,
29    pub(crate) sql: String,
30    pub(crate) param_types: Vec<crate::types::Type>,
31    pub(crate) columns: Arc<Vec<FieldDescription>>,
32}
33
34impl PreparedStatement {
35    /// The server-side name of this prepared statement.
36    pub fn name(&self) -> &str {
37        &self.name
38    }
39
40    /// The SQL text used to create this statement.
41    pub fn sql(&self) -> &str {
42        &self.sql
43    }
44
45    /// The parameter types inferred by the server.
46    pub fn param_types(&self) -> &[crate::types::Type] {
47        &self.param_types
48    }
49
50    /// The result column descriptions (empty for non-SELECT statements).
51    pub fn columns(&self) -> &[FieldDescription] {
52        &self.columns
53    }
54}
55
56// ---------------------------------------------------------------------------
57// Connection methods
58// ---------------------------------------------------------------------------
59
60impl Connection {
61    /// Generate a unique statement name.
62    fn next_statement_name(&mut self) -> String {
63        self.statement_counter += 1;
64        format!("__pg_stmt_{}", self.statement_counter)
65    }
66
67    /// Prepare a statement for repeated execution.
68    ///
69    /// The server parses the SQL, infers parameter types, and returns result
70    /// column metadata. The returned [`PreparedStatement`] can be passed to
71    /// [`Connection::query_prepared`] to execute with parameters.
72    ///
73    /// # Example
74    /// ```ignore
75    /// let stmt = conn.prepare("SELECT * FROM users WHERE id = $1").await?;
76    /// let rows = conn.query_prepared(&stmt, &[&42i32]).await?;
77    /// ```
78    #[must_use = "prepare errors should be checked"]
79    pub async fn prepare(&mut self, sql: &str) -> Result<PreparedStatement> {
80        self.transition(ConnectionState::ActiveExtendedQuery)?;
81
82        let name = self.next_statement_name();
83
84        // Parse
85        self.codec
86            .encode_and_write(
87                &mut self.transport,
88                &FrontendMessage::Parse {
89                    name: name.clone(),
90                    sql: sql.to_string(),
91                    param_types: vec![], // let server infer
92                },
93            )
94            .await?;
95
96        // Describe (to get param types and result columns)
97        self.codec
98            .encode_and_write(
99                &mut self.transport,
100                &FrontendMessage::Describe {
101                    variant: b'S',
102                    name: name.clone(),
103                },
104            )
105            .await?;
106
107        // Sync
108        self.codec
109            .encode_and_write(&mut self.transport, &FrontendMessage::Sync)
110            .await?;
111
112        // Flush the batch
113        self.transport.flush().await.map_err(PgError::Transport)?;
114
115        // Read responses
116        let mut param_types = Vec::new();
117        let mut columns = Vec::new();
118
119        loop {
120            let msg = self.codec.read_message(&mut self.transport).await?;
121            if self.handle_async_message(&msg) {
122                continue;
123            }
124            match msg {
125                BackendMessage::ParseComplete => {}
126                BackendMessage::ParameterDescription(body) => {
127                    let mut iter = body.parameters();
128                    while let Some(oid) = iter.next()? {
129                        if let Some(ty) = crate::types::type_from_oid(oid) {
130                            param_types.push(ty);
131                        } else {
132                            param_types.push(crate::types::Type::UNKNOWN);
133                        }
134                    }
135                }
136                BackendMessage::RowDescription(body) => {
137                    columns = read_row_description(body)?;
138                }
139                BackendMessage::NoData => {
140                    // Statement doesn't return rows (INSERT, UPDATE, etc.)
141                }
142                BackendMessage::ReadyForQuery(body) => {
143                    self.transaction_status =
144                        crate::protocol::TransactionStatus::from_u8(body.status())
145                            .unwrap_or(crate::protocol::TransactionStatus::Idle);
146                    self.state = ConnectionState::Idle;
147                    break;
148                }
149                BackendMessage::ErrorResponse(body) => {
150                    let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
151                    self.read_until_ready().await?;
152                    return Err(PgError::Server(Box::new(server_err)));
153                }
154                _ => {}
155            }
156        }
157
158        Ok(PreparedStatement {
159            name,
160            sql: sql.to_string(),
161            param_types,
162            columns: Arc::new(columns),
163        })
164    }
165
166    /// Deallocate a prepared statement on the server.
167    #[must_use = "close errors should be checked"]
168    pub async fn close_statement(&mut self, stmt: &PreparedStatement) -> Result<()> {
169        self.transition(ConnectionState::ActiveExtendedQuery)?;
170
171        self.codec
172            .encode_and_write(
173                &mut self.transport,
174                &FrontendMessage::Close {
175                    variant: b'S',
176                    name: stmt.name.clone(),
177                },
178            )
179            .await?;
180
181        self.codec
182            .encode_and_write(&mut self.transport, &FrontendMessage::Sync)
183            .await?;
184
185        // Flush the batch
186        self.transport.flush().await.map_err(PgError::Transport)?;
187
188        loop {
189            let msg = self.codec.read_message(&mut self.transport).await?;
190            if self.handle_async_message(&msg) {
191                continue;
192            }
193            match msg {
194                BackendMessage::CloseComplete => {}
195                BackendMessage::ReadyForQuery(body) => {
196                    self.transaction_status =
197                        crate::protocol::TransactionStatus::from_u8(body.status())
198                            .unwrap_or(crate::protocol::TransactionStatus::Idle);
199                    self.state = ConnectionState::Idle;
200                    break;
201                }
202                BackendMessage::ErrorResponse(body) => {
203                    let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
204                    self.read_until_ready().await?;
205                    return Err(PgError::Server(Box::new(server_err)));
206                }
207                _ => {}
208            }
209        }
210
211        Ok(())
212    }
213}
214
215// ---------------------------------------------------------------------------
216// Tests
217// ---------------------------------------------------------------------------
218
219#[cfg(test)]
220mod tests {
221    use super::*;
222    use crate::auth::{Codec, ServerParams};
223    use crate::config::Config;
224    use crate::connection::ConnectionState;
225    use crate::transport::{BufferedTransport, ClientTransport, MockTransport, PgTransport};
226    use std::collections::VecDeque;
227
228    fn make_connection(read_data: Vec<u8>) -> Connection {
229        let transport = PgTransport::Plain(BufferedTransport::new(ClientTransport::Mock(
230            MockTransport::new(read_data),
231        )));
232        Connection {
233            transport,
234            codec: Codec::new(),
235            server_params: ServerParams::default(),
236            state: ConnectionState::Idle,
237            config: Config::new(),
238            transaction_status: crate::protocol::TransactionStatus::Idle,
239            notification_queue: VecDeque::new(),
240            notice_handler: None,
241            statement_counter: 0,
242            needs_recovery: false,
243            health: crate::reconnect::session::ConnectionHealth::new(),
244            session_state: crate::reconnect::session::SessionState::new(),
245        }
246    }
247
248    fn build_parse_complete() -> Vec<u8> {
249        vec![b'1', 0, 0, 0, 4]
250    }
251
252    fn build_parameter_description(oids: &[u32]) -> Vec<u8> {
253        let mut buf = vec![b't'];
254        let mut body = Vec::new();
255        body.extend_from_slice(&(oids.len() as i16).to_be_bytes());
256        for oid in oids {
257            body.extend_from_slice(&oid.to_be_bytes());
258        }
259        let len = (body.len() + 4) as i32;
260        buf.extend_from_slice(&len.to_be_bytes());
261        buf.extend_from_slice(&body);
262        buf
263    }
264
265    fn build_no_data() -> Vec<u8> {
266        vec![b'n', 0, 0, 0, 4]
267    }
268
269    fn build_ready_for_query(status: u8) -> Vec<u8> {
270        vec![b'Z', 0, 0, 0, 5, status]
271    }
272
273    #[tokio::test]
274    async fn test_prepare_select() {
275        let mut data = Vec::new();
276        data.extend_from_slice(&build_parse_complete());
277        // ParameterDescription: 2 params (INT4=23, TEXT=25)
278        data.extend_from_slice(&build_parameter_description(&[23, 25]));
279        // RowDescription: id INT4, name TEXT
280        data.extend_from_slice(&super::super::tests::build_row_description_msg(&[
281            ("id", crate::types::INT4_OID),
282            ("name", crate::types::TEXT_OID),
283        ]));
284        data.extend_from_slice(&build_ready_for_query(b'I'));
285
286        let mut conn = make_connection(data);
287        let stmt = conn
288            .prepare("SELECT * FROM users WHERE id = $1 AND name = $2")
289            .await
290            .unwrap();
291
292        assert_eq!(stmt.name(), "__pg_stmt_1");
293        assert_eq!(
294            stmt.sql(),
295            "SELECT * FROM users WHERE id = $1 AND name = $2"
296        );
297        assert_eq!(stmt.param_types().len(), 2);
298        assert_eq!(stmt.param_types()[0], crate::types::Type::INT4);
299        assert_eq!(stmt.param_types()[1], crate::types::Type::TEXT);
300        assert_eq!(stmt.columns().len(), 2);
301        assert_eq!(stmt.columns()[0].name(), "id");
302        assert_eq!(stmt.columns()[1].name(), "name");
303    }
304
305    #[tokio::test]
306    async fn test_prepare_insert() {
307        let mut data = Vec::new();
308        data.extend_from_slice(&build_parse_complete());
309        // ParameterDescription: 1 param (INT4=23)
310        data.extend_from_slice(&build_parameter_description(&[23]));
311        data.extend_from_slice(&build_no_data());
312        data.extend_from_slice(&build_ready_for_query(b'I'));
313
314        let mut conn = make_connection(data);
315        let stmt = conn
316            .prepare("INSERT INTO users (id) VALUES ($1)")
317            .await
318            .unwrap();
319
320        assert_eq!(stmt.param_types().len(), 1);
321        assert!(stmt.columns().is_empty());
322    }
323
324    #[tokio::test]
325    async fn test_close_statement() {
326        let mut data = Vec::new();
327        data.extend_from_slice(&[b'3', 0, 0, 0, 4]); // CloseComplete
328        data.extend_from_slice(&build_ready_for_query(b'I'));
329
330        let mut conn = make_connection(data);
331        let stmt = PreparedStatement {
332            name: "__pg_stmt_1".into(),
333            sql: "SELECT 1".into(),
334            param_types: vec![],
335            columns: Arc::new(vec![]),
336        };
337        conn.close_statement(&stmt).await.unwrap();
338        assert!(conn.is_idle());
339    }
340}