1use 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
28pub enum ToStatementResult<'a> {
30 Immediate(Statement),
32 Mediate(crate::BoxFuture<'a, Statement>),
34}
35
36pub trait StatementLike: Send + Sync {
37 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#[derive(Debug, Eq, PartialEq)]
180pub struct StmtInner {
181 pub(crate) raw_query: Arc<[u8]>,
182 columns: Option<Arc<[Column]>>,
183 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#[derive(Debug, Clone, Eq, PartialEq)]
264pub struct Statement {
265 pub(crate) inner: Arc<StmtInner>,
266 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 pub fn columns(&self) -> Arc<[Column]> {
280 self.inner.columns()
281 }
282
283 pub(crate) fn update_columns_metadata(&self, columns: Vec<Column>) {
287 self.inner.update_columns_metadata(columns);
288 }
289
290 pub fn params(&self) -> &[Column] {
292 self.inner.params()
293 }
294
295 pub fn id(&self) -> u32 {
297 self.inner.id()
298 }
299
300 pub fn connection_id(&self) -> u32 {
302 self.inner.connection_id()
303 }
304
305 pub fn num_params(&self) -> u16 {
307 self.inner.num_params()
308 }
309
310 pub fn num_columns(&self) -> u16 {
312 self.inner.num_columns()
313 }
314}
315
316impl crate::Conn {
317 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 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 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 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 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 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}