Skip to main content

mysql_async/queryable/
stmt.rs

1// Copyright (c) 2017 Anatoly Ikorsky
2//
3// Licensed under the Apache License, Version 2.0
4// <LICENSE-APACHE or http://www.apache.org/licenses/LICENSE-2.0> or the MIT
5// license <LICENSE-MIT or http://opensource.org/licenses/MIT>, at your
6// option. All files in the project carrying such notice may not be copied,
7// modified, or distributed except according to those terms.
8
9use arc_swap::ArcSwapOption;
10use futures_util::FutureExt;
11use mysql_common::{
12    io::ParseBuf,
13    named_params::ParsedNamedParams,
14    packets::{ComStmtClose, StmtPacket},
15};
16
17use std::{borrow::Cow, fmt, sync::Arc};
18
19use crate::{
20    conn::routines::{ExecBulkRoutine, ExecRoutine, PrepareRoutine},
21    consts::CapabilityFlags,
22    error::*,
23    Column, Params,
24};
25
26use super::AsQuery;
27
28/// Result of a `StatementLike::to_statement` call.
29pub enum ToStatementResult<'a> {
30    /// Statement is immediately available.
31    Immediate(Statement),
32    /// We need some time to get a statement and the operation itself may fail.
33    Mediate(crate::BoxFuture<'a, Statement>),
34}
35
36pub trait StatementLike: Send + Sync {
37    /// Returns a statement.
38    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
39    where
40        Self: 'a;
41}
42
43fn to_statement_move<'a, T: AsQuery + 'a>(
44    stmt: T,
45    conn: &'a mut crate::Conn,
46) -> ToStatementResult<'a> {
47    let fut = async move {
48        let query = stmt.as_query();
49        let parsed = ParsedNamedParams::parse(query.as_ref())?;
50        let inner_stmt = match conn.get_cached_stmt(parsed.query()) {
51            Some(inner_stmt) => inner_stmt,
52            None => {
53                conn.prepare_statement(Cow::Borrowed(parsed.query()))
54                    .await?
55            }
56        };
57        Ok(Statement::new(
58            inner_stmt,
59            parsed
60                .params()
61                .iter()
62                .map(|x| x.as_ref().to_vec())
63                .collect::<Vec<_>>(),
64        ))
65    }
66    .boxed();
67    ToStatementResult::Mediate(fut)
68}
69
70impl StatementLike for Cow<'_, str> {
71    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
72    where
73        Self: 'a,
74    {
75        to_statement_move(self, conn)
76    }
77}
78
79impl StatementLike for &'_ str {
80    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
81    where
82        Self: 'a,
83    {
84        to_statement_move(self, conn)
85    }
86}
87
88impl StatementLike for String {
89    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
90    where
91        Self: 'a,
92    {
93        to_statement_move(self, conn)
94    }
95}
96
97impl StatementLike for Box<str> {
98    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
99    where
100        Self: 'a,
101    {
102        to_statement_move(self, conn)
103    }
104}
105
106impl StatementLike for Arc<str> {
107    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
108    where
109        Self: 'a,
110    {
111        to_statement_move(self, conn)
112    }
113}
114
115impl StatementLike for Cow<'_, [u8]> {
116    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
117    where
118        Self: 'a,
119    {
120        to_statement_move(self, conn)
121    }
122}
123
124impl StatementLike for &'_ [u8] {
125    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
126    where
127        Self: 'a,
128    {
129        to_statement_move(self, conn)
130    }
131}
132
133impl StatementLike for Vec<u8> {
134    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
135    where
136        Self: 'a,
137    {
138        to_statement_move(self, conn)
139    }
140}
141
142impl StatementLike for Box<[u8]> {
143    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
144    where
145        Self: 'a,
146    {
147        to_statement_move(self, conn)
148    }
149}
150
151impl StatementLike for Arc<[u8]> {
152    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
153    where
154        Self: 'a,
155    {
156        to_statement_move(self, conn)
157    }
158}
159
160impl StatementLike for Statement {
161    fn to_statement<'a>(self, _conn: &'a mut crate::Conn) -> ToStatementResult<'static>
162    where
163        Self: 'a,
164    {
165        ToStatementResult::Immediate(self.clone())
166    }
167}
168
169impl<T: StatementLike + Clone> StatementLike for &'_ T {
170    fn to_statement<'a>(self, conn: &'a mut crate::Conn) -> ToStatementResult<'a>
171    where
172        Self: 'a,
173    {
174        self.clone().to_statement(conn)
175    }
176}
177
178/// Statement data.
179#[derive(Debug, Eq, PartialEq)]
180pub struct StmtInner {
181    pub(crate) raw_query: Arc<[u8]>,
182    columns: Option<Arc<[Column]>>,
183    /// This cached value overrides the column metadata stored in the `inner` field.
184    ///
185    /// See MARIADB_CLIENT_CACHE_METADATA capability.
186    columns_cache: ColumnCache,
187    params: Option<Box<[Column]>>,
188    stmt_packet: StmtPacket,
189    connection_id: u32,
190}
191
192impl StmtInner {
193    pub(crate) fn from_payload(
194        pld: &[u8],
195        connection_id: u32,
196        raw_query: Arc<[u8]>,
197    ) -> std::io::Result<Self> {
198        let stmt_packet = ParseBuf(pld).parse(())?;
199
200        Ok(Self {
201            raw_query,
202            columns: None,
203            columns_cache: ColumnCache::new(),
204            params: None,
205            stmt_packet,
206            connection_id,
207        })
208    }
209
210    pub(crate) fn with_params(mut self, params: Vec<Column>) -> Self {
211        self.params = if params.is_empty() {
212            None
213        } else {
214            Some(params.into_boxed_slice())
215        };
216        self
217    }
218
219    pub(crate) fn with_columns(mut self, columns: Vec<Column>) -> Self {
220        self.columns = if columns.is_empty() {
221            None
222        } else {
223            Some(Arc::from(columns))
224        };
225        self
226    }
227
228    pub(crate) fn columns(&self) -> Arc<[Column]> {
229        self.columns_cache
230            .get_columns()
231            .or_else(|| self.columns.clone())
232            .unwrap_or_default()
233    }
234
235    pub fn update_columns_metadata(&self, columns: Vec<Column>) {
236        self.columns_cache.set_columns(columns);
237    }
238
239    pub(crate) fn params(&self) -> &[Column] {
240        self.params.as_ref().map(AsRef::as_ref).unwrap_or(&[])
241    }
242
243    pub(crate) fn id(&self) -> u32 {
244        self.stmt_packet.statement_id()
245    }
246
247    pub(crate) const fn connection_id(&self) -> u32 {
248        self.connection_id
249    }
250
251    pub(crate) fn num_params(&self) -> u16 {
252        self.stmt_packet.num_params()
253    }
254
255    pub(crate) fn num_columns(&self) -> u16 {
256        self.stmt_packet.num_columns()
257    }
258}
259
260/// Prepared statement.
261///
262/// Statement is only valid for connection with id `Statement::connection_id()`.
263#[derive(Debug, Clone, Eq, PartialEq)]
264pub struct Statement {
265    pub(crate) inner: Arc<StmtInner>,
266    /// An empty vector in case of no named params.
267    pub(crate) named_params: Vec<Vec<u8>>,
268}
269
270impl Statement {
271    pub(crate) fn new(inner: Arc<StmtInner>, named_params: Vec<Vec<u8>>) -> Self {
272        Self {
273            inner,
274            named_params,
275        }
276    }
277
278    /// Returned columns.
279    pub fn columns(&self) -> Arc<[Column]> {
280        self.inner.columns()
281    }
282
283    /// Overrides columns metadata for this statement.
284    ///
285    /// See MARIADB_CLIENT_CACHE_METADATA capability.
286    pub(crate) fn update_columns_metadata(&self, columns: Vec<Column>) {
287        self.inner.update_columns_metadata(columns);
288    }
289
290    /// Required parameters.
291    pub fn params(&self) -> &[Column] {
292        self.inner.params()
293    }
294
295    /// MySql statement identifier.
296    pub fn id(&self) -> u32 {
297        self.inner.id()
298    }
299
300    /// MySql connection identifier.
301    pub fn connection_id(&self) -> u32 {
302        self.inner.connection_id()
303    }
304
305    /// Number of parameters.
306    pub fn num_params(&self) -> u16 {
307        self.inner.num_params()
308    }
309
310    /// Number of columns.
311    pub fn num_columns(&self) -> u16 {
312        self.inner.num_columns()
313    }
314}
315
316impl crate::Conn {
317    /// Low-level helpers, that reads the given number of column packets terminated by EOF packet.
318    ///
319    /// Requires `num > 0`.
320    pub(crate) async fn read_column_defs<U>(&mut self, num: U) -> Result<Vec<Column>>
321    where
322        U: Into<usize>,
323    {
324        let num = num.into();
325        debug_assert!(num > 0);
326        let packets = self.read_packets(num).await?;
327        let defs = packets
328            .into_iter()
329            .map(|x| ParseBuf(&x).parse(()))
330            .collect::<std::result::Result<Vec<Column>, _>>()
331            .map_err(Error::from)?;
332
333        if !self.has_capabilities(CapabilityFlags::CLIENT_DEPRECATE_EOF) {
334            self.read_packet().await?;
335        }
336
337        Ok(defs)
338    }
339
340    /// Helper, that retrieves `Statement` from `StatementLike`.
341    pub(crate) async fn get_statement<U>(&mut self, stmt_like: U) -> Result<Statement>
342    where
343        U: StatementLike,
344    {
345        match stmt_like.to_statement(self) {
346            ToStatementResult::Immediate(statement) => Ok(statement),
347            ToStatementResult::Mediate(statement) => statement.await,
348        }
349    }
350
351    /// Low-level helper, that prepares the given statement.
352    ///
353    /// `raw_query` is a query with `?` placeholders (if any).
354    async fn prepare_statement(&mut self, raw_query: Cow<'_, [u8]>) -> Result<Arc<StmtInner>> {
355        let inner_stmt = self.routine(PrepareRoutine::new(raw_query)).await?;
356
357        if let Some(old_stmt) = self.cache_stmt(&inner_stmt) {
358            self.close_statement(old_stmt.id()).await?;
359        }
360
361        Ok(inner_stmt)
362    }
363
364    /// Helper, that executes the given statement with the given params.
365    pub(crate) async fn execute_statement<P>(
366        &mut self,
367        statement: &Statement,
368        params: P,
369    ) -> Result<()>
370    where
371        P: Into<Params>,
372    {
373        self.routine(ExecRoutine::new(statement, params.into()))
374            .await?;
375        Ok(())
376    }
377
378    /// Helper, that executes the given statement as a bulk operation
379    ///
380    /// Available in MariaDb with MARIADB_CLIENT_STMT_BULK_OPERATIONS capability.
381    pub(crate) async fn execute_bulk<P, I>(
382        &mut self,
383        statement: &Statement,
384        params: I,
385    ) -> Result<()>
386    where
387        P: Into<Params> + Send,
388        I: IntoIterator<Item = P> + Send,
389        I::IntoIter: Send,
390    {
391        self.routine(ExecBulkRoutine::new(statement, params))
392            .await?;
393        Ok(())
394    }
395
396    /// Helper, that closes statement with the given id.
397    pub(crate) async fn close_statement(&mut self, id: u32) -> Result<()> {
398        self.stmt_cache_mut().remove(id);
399        self.write_command(&ComStmtClose::new(id)).await
400    }
401}
402
403struct ColumnCache {
404    columns: ArcSwapOption<Arc<[Column]>>,
405}
406
407impl ColumnCache {
408    fn new() -> Self {
409        Self {
410            columns: ArcSwapOption::const_empty(),
411        }
412    }
413
414    fn get_columns(&self) -> Option<Arc<[Column]>> {
415        self.columns.load_full().map(|x| (*x).clone())
416    }
417
418    fn set_columns(&self, new_columns: Vec<Column>) {
419        let new_columns: Arc<[Column]> = new_columns.into();
420        self.columns.store(Some(Arc::new(new_columns)));
421    }
422}
423
424impl fmt::Debug for ColumnCache {
425    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
426        f.debug_struct("ColumnCache")
427            .field("columns", &self.get_columns())
428            .finish()
429    }
430}
431
432impl PartialEq for ColumnCache {
433    fn eq(&self, other: &Self) -> bool {
434        self.get_columns() == other.get_columns()
435    }
436}
437
438impl Eq for ColumnCache {}
439
440#[cfg(test)]
441mod tests {
442    use super::ColumnCache;
443    use crate::Column;
444    use mysql_common::constants::ColumnType;
445    use std::{sync::Arc, thread};
446
447    const ROUNDS: usize = if cfg!(miri) { 8 } else { 10_000 };
448
449    fn columns(table: &str, len: usize) -> Vec<Column> {
450        (0..len)
451            .map(|_| Column::new(ColumnType::MYSQL_TYPE_LONG).with_table(table.as_bytes()))
452            .collect()
453    }
454
455    #[test]
456    fn reader_keeps_columns_alive_while_writer_replaces_them() {
457        let cache = Arc::new(ColumnCache::new());
458        cache.set_columns(columns("t", 42));
459
460        let writer = {
461            let cache = Arc::clone(&cache);
462            thread::spawn(move || {
463                for _ in 0..ROUNDS {
464                    cache.set_columns(columns("t", 46));
465                    cache.set_columns(columns("t", 42));
466                }
467            })
468        };
469
470        let readers: Vec<_> = (0..4)
471            .map(|_| {
472                let cache = Arc::clone(&cache);
473                thread::spawn(move || {
474                    for _ in 0..ROUNDS {
475                        let Some(columns) = cache.get_columns() else {
476                            continue;
477                        };
478                        for column in columns.iter() {
479                            assert_eq!(column.table_str(), "t");
480                        }
481                    }
482                })
483            })
484            .collect();
485
486        writer.join().unwrap();
487        for reader in readers {
488            reader.join().unwrap();
489        }
490    }
491}