wasi_pg_client/query/
prepared.rs1use 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#[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 pub fn name(&self) -> &str {
37 &self.name
38 }
39
40 pub fn sql(&self) -> &str {
42 &self.sql
43 }
44
45 pub fn param_types(&self) -> &[crate::types::Type] {
47 &self.param_types
48 }
49
50 pub fn columns(&self) -> &[FieldDescription] {
52 &self.columns
53 }
54}
55
56impl Connection {
61 fn next_statement_name(&mut self) -> String {
63 self.statement_counter += 1;
64 format!("__pg_stmt_{}", self.statement_counter)
65 }
66
67 #[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 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![], },
93 )
94 .await?;
95
96 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 self.codec
109 .encode_and_write(&mut self.transport, &FrontendMessage::Sync)
110 .await?;
111
112 self.transport.flush().await.map_err(PgError::Transport)?;
114
115 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 }
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 #[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 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#[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 data.extend_from_slice(&build_parameter_description(&[23, 25]));
279 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 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]); 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}