Skip to main content

toasty_driver_mysql/
lib.rs

1#![warn(missing_docs)]
2#![allow(clippy::needless_range_loop)]
3
4//! Toasty drivers for [MySQL](https://www.mysql.com/) and
5//! [MariaDB](https://mariadb.org/) 11.8 and later, using [SQLx](https://docs.rs/sqlx).
6//!
7//! # Examples
8//!
9//! ```no_run
10//! use toasty_driver_mysql::{MySQL, MariaDB};
11//!
12//! let driver = MySQL::new("mysql://localhost/mydb").unwrap();
13//! let driver = MariaDB::new("mariadb://localhost/mydb").unwrap();
14//! ```
15
16mod mariadb;
17pub use mariadb::MariaDB;
18
19mod value;
20pub(crate) use value::Value;
21
22use async_trait::async_trait;
23use sqlx_core::{
24    connection::{ConnectOptions as _, Connection as SqlxConnection},
25    row::Row,
26    sql_str::AssertSqlSafe,
27};
28use sqlx_mysql::{MySqlArguments, MySqlConnectOptions, MySqlConnection, MySqlDatabaseError};
29use std::{borrow::Cow, cell::Cell, sync::Arc};
30use toasty_core::{
31    Result, Schema,
32    driver::{
33        Capability, ConnectContext, ConnectionUrl, Dialect, Driver, ExecResponse, Operation,
34        QueryLogConfig,
35        log::QueryLog,
36        operation::{RawSqlRet, Transaction, TransactionMode},
37    },
38    schema::{
39        db::{self, Migration, Table},
40        diff,
41    },
42    stmt::{self, ValueRecord},
43};
44use toasty_sql::{self as sql};
45
46enum SqlReturn {
47    Count,
48    LastInsertId(stmt::Type),
49    Infer,
50    Types(Vec<stmt::Type>),
51}
52
53/// Classifies a SQLx MySQL error into a Toasty error.
54///
55/// Transport and protocol failures become `ConnectionLost`. MySQL
56/// errors with numbers that Toasty understands become typed errors.
57/// Everything else becomes `DriverOperationFailed`.
58fn classify_mysql_error(e: sqlx_core::Error) -> toasty_core::Error {
59    match e {
60        error @ (sqlx_core::Error::Io(_)
61        | sqlx_core::Error::Protocol(_)
62        | sqlx_core::Error::WorkerCrashed) => toasty_core::Error::connection_lost(error),
63        sqlx_core::Error::Database(database_error) => {
64            let mysql_error = database_error
65                .try_downcast_ref::<MySqlDatabaseError>()
66                .expect("SQLx returned a non-MySQL database error from a MySQL connection");
67            let number = mysql_error.number();
68            let message = mysql_error.message().to_owned();
69
70            match number {
71                1213 => toasty_core::Error::serialization_failure(message),
72                1792 => toasty_core::Error::read_only_transaction(message),
73                _ => toasty_core::Error::driver_operation_failed(sqlx_core::Error::Database(
74                    database_error,
75                )),
76            }
77        }
78        other => toasty_core::Error::driver_operation_failed(other),
79    }
80}
81
82/// Classifies a SQLx error and records whether the connection is still usable.
83fn record_mysql_err(valid: &Cell<bool>, e: sqlx_core::Error) -> toasty_core::Error {
84    let err = classify_mysql_error(e);
85    if err.is_connection_lost() {
86        valid.set(false);
87    }
88    err
89}
90
91/// A MySQL [`Driver`] that connects through SQLx.
92///
93/// # Examples
94///
95/// ```no_run
96/// use toasty_driver_mysql::MySQL;
97///
98/// let driver = MySQL::new("mysql://localhost/mydb").unwrap();
99/// ```
100#[derive(Debug)]
101pub struct MySQL {
102    url: String,
103    opts: MySqlConnectOptions,
104    capability: &'static Capability,
105}
106
107impl MySQL {
108    /// Creates a MySQL driver from a SQLx connection URL.
109    ///
110    /// The URL must use the `mysql` scheme and include a database path, such as
111    /// `mysql://user:pass@host:3306/dbname`.
112    pub fn new(url: impl Into<String>) -> Result<Self> {
113        Self::with_capability(url, "mysql", &Capability::MYSQL)
114    }
115
116    fn with_capability(
117        url: impl Into<String>,
118        scheme: &str,
119        capability: &'static Capability,
120    ) -> Result<Self> {
121        let url_str = url.into();
122        let url = ConnectionUrl::parse(&url_str)?;
123
124        if !url.has_scheme(scheme) {
125            return Err(toasty_core::Error::invalid_connection_url(format!(
126                "connection url does not have a `{scheme}` scheme; url={}",
127                url.as_str()
128            )));
129        }
130
131        url.host()?.ok_or_else(|| {
132            toasty_core::Error::invalid_connection_url(format!(
133                "missing host in connection URL; url={}",
134                url.as_str()
135            ))
136        })?;
137
138        if url.path().is_empty() {
139            return Err(toasty_core::Error::invalid_connection_url(format!(
140                "no database specified - missing path in connection URL; url={}",
141                url.as_str()
142            )));
143        }
144
145        let opts = url
146            .as_str()
147            .parse::<MySqlConnectOptions>()
148            .map_err(toasty_core::Error::driver_operation_failed)?
149            .disable_statement_logging();
150
151        Ok(Self {
152            url: url_str,
153            opts,
154            capability,
155        })
156    }
157}
158
159/// The serializer for the dialect `capability` names.
160fn serializer<'a>(capability: &Capability, schema: &'a db::Schema) -> sql::Serializer<'a> {
161    match capability.sql {
162        Some(Dialect::MariaDb) => sql::Serializer::mariadb(schema),
163        _ => sql::Serializer::mysql(schema),
164    }
165}
166
167#[async_trait]
168impl Driver for MySQL {
169    fn url(&self) -> Cow<'_, str> {
170        Cow::Borrowed(&self.url)
171    }
172
173    fn capability(&self) -> &'static Capability {
174        self.capability
175    }
176
177    async fn connect(
178        &self,
179        cx: &ConnectContext,
180    ) -> Result<Box<dyn toasty_core::driver::Connection>> {
181        let conn = MySqlConnection::connect_with(&self.opts)
182            .await
183            .map_err(classify_mysql_error)?;
184        let mut connection = Connection::with_capability(conn, self.capability);
185        connection.query_log = cx.query_log;
186        Ok(Box::new(connection))
187    }
188
189    fn generate_migration(&self, schema_diff: &diff::Schema<'_>) -> Migration {
190        let statements = sql::MigrationStatement::from_diff(schema_diff, self.capability);
191
192        let sql_strings: Vec<String> = statements
193            .iter()
194            .map(|stmt| serializer(self.capability, stmt.schema()).serialize(stmt.statement()))
195            .collect();
196
197        Migration::new_sql_with_breakpoints(&sql_strings)
198    }
199
200    async fn reset_db(&self) -> Result<()> {
201        let mut conn = MySqlConnection::connect_with(&self.opts)
202            .await
203            .map_err(classify_mysql_error)?;
204        let dbname = self.opts.get_database().ok_or_else(|| {
205            toasty_core::Error::invalid_connection_url("no database name configured")
206        })?;
207
208        let dbname = format!("`{}`", dbname.replace('`', "``"));
209
210        sqlx_core::raw_sql::raw_sql(AssertSqlSafe(format!("DROP DATABASE IF EXISTS {dbname}")))
211            .execute(&mut conn)
212            .await
213            .map_err(classify_mysql_error)?;
214        sqlx_core::raw_sql::raw_sql(AssertSqlSafe(format!("CREATE DATABASE {dbname}")))
215            .execute(&mut conn)
216            .await
217            .map_err(classify_mysql_error)?;
218        sqlx_core::raw_sql::raw_sql(AssertSqlSafe(format!("USE {dbname}")))
219            .execute(&mut conn)
220            .await
221            .map_err(classify_mysql_error)?;
222
223        Ok(())
224    }
225}
226
227/// An open connection to a MySQL database.
228#[derive(Debug)]
229pub struct Connection {
230    conn: MySqlConnection,
231    /// Set to `false` after a connection-level failure. SQLx does not expose a
232    /// passive validity flag, so the driver records one for [`is_valid`].
233    valid: Cell<bool>,
234    query_log: QueryLogConfig,
235    capability: &'static Capability,
236}
237
238impl Connection {
239    /// Wraps an existing SQLx [`MySqlConnection`] as a Toasty connection.
240    pub fn new(conn: MySqlConnection) -> Self {
241        Self::with_capability(conn, &Capability::MYSQL)
242    }
243
244    fn with_capability(conn: MySqlConnection, capability: &'static Capability) -> Self {
245        Self {
246            conn,
247            valid: Cell::new(true),
248            query_log: QueryLogConfig::default(),
249            capability,
250        }
251    }
252
253    async fn exec_sql(
254        &mut self,
255        sql_as_str: &str,
256        args: MySqlArguments,
257        ret: SqlReturn,
258        log: &mut QueryLog<'_>,
259    ) -> Result<ExecResponse> {
260        if matches!(ret, SqlReturn::Count | SqlReturn::LastInsertId(_)) {
261            let result = sqlx_core::query::query_with(AssertSqlSafe(sql_as_str), args)
262                .execute(&mut self.conn)
263                .await
264                .map_err(|e| record_mysql_err(&self.valid, e))?;
265
266            if let SqlReturn::LastInsertId(ty) = ret {
267                let id = ty.cast(&(), stmt::Value::U64(result.last_insert_id()))?;
268                log.rows(1);
269                return Ok(ExecResponse::value_stream(stmt::ValueStream::from_vec(
270                    vec![ValueRecord::from_vec(vec![id]).into()],
271                )));
272            }
273
274            return Ok(ExecResponse::count(result.rows_affected()));
275        }
276
277        let rows = sqlx_core::query::query_with(AssertSqlSafe(sql_as_str), args)
278            .fetch_all(&mut self.conn)
279            .await
280            .map_err(|e| record_mysql_err(&self.valid, e))?;
281
282        log.rows(rows.len() as u64);
283        let mut records = Vec::with_capacity(rows.len());
284
285        for row in rows {
286            let mut values = Vec::with_capacity(row.len());
287
288            match &ret {
289                SqlReturn::Count | SqlReturn::LastInsertId(_) => unreachable!(),
290                SqlReturn::Infer => {
291                    for i in 0..row.len() {
292                        let column = row.column(i);
293                        values.push(
294                            Value::from_sql_infer(i, &row, column)
295                                .map_err(|e| record_mysql_err(&self.valid, e))?
296                                .into_inner(),
297                        );
298                    }
299                }
300                SqlReturn::Types(returning) => {
301                    assert_eq!(
302                        row.len(),
303                        returning.len(),
304                        "row={row:#?}; returning={returning:#?}"
305                    );
306
307                    for i in 0..row.len() {
308                        let column = row.column(i);
309                        values.push(
310                            Value::from_sql(i, &row, column, &returning[i])
311                                .map_err(|e| record_mysql_err(&self.valid, e))?
312                                .into_inner(),
313                        );
314                    }
315                }
316            }
317
318            records.push(Ok(ValueRecord::from_vec(values)));
319        }
320
321        Ok(ExecResponse::value_stream(stmt::ValueStream::from_iter(
322            records.into_iter(),
323        )))
324    }
325
326    /// Creates a table and its indices from a schema definition.
327    pub async fn create_table(&mut self, schema: &db::Schema, table: &Table) -> Result<()> {
328        let serializer = serializer(self.capability, schema);
329        let statement = serializer.serialize(&sql::Statement::create_table(table, self.capability));
330
331        sqlx_core::query::query(AssertSqlSafe(statement))
332            .execute(&mut self.conn)
333            .await
334            .map_err(|e| record_mysql_err(&self.valid, e))?;
335
336        for index in &table.indices {
337            if index.primary_key {
338                continue;
339            }
340
341            let statement = serializer.serialize(&sql::Statement::create_index(index));
342            sqlx_core::query::query(AssertSqlSafe(statement))
343                .execute(&mut self.conn)
344                .await
345                .map_err(|e| record_mysql_err(&self.valid, e))?;
346        }
347
348        Ok(())
349    }
350}
351
352impl From<MySqlConnection> for Connection {
353    fn from(conn: MySqlConnection) -> Self {
354        Self::new(conn)
355    }
356}
357
358#[async_trait]
359impl toasty_core::driver::Connection for Connection {
360    async fn exec(&mut self, schema: &Arc<Schema>, op: Operation) -> Result<ExecResponse> {
361        let driver = match self.capability.sql {
362            Some(Dialect::MariaDb) => "mariadb",
363            _ => "mysql",
364        };
365        tracing::trace!(driver, op = %op.name(), "driver exec");
366
367        let (sql, typed_params, ret) = match op {
368            Operation::Insert(op) => {
369                // A surviving RETURNING clause means MariaDB, and a result
370                // set to read. For MySQL the planner strips it and asks for
371                // the auto-increment key alone, via LAST_INSERT_ID().
372                let returns_rows = op.stmt.returning().is_some();
373
374                let ret = match op.ret {
375                    Some(types) if returns_rows => SqlReturn::Types(types),
376                    Some(types) => {
377                        let [ty] = &types[..] else {
378                            return Err(toasty_core::Error::invalid_result(format!(
379                                "MySQL insert ID result requires one type, got {types:?}"
380                            )));
381                        };
382                        SqlReturn::LastInsertId(ty.clone())
383                    }
384                    None => SqlReturn::Count,
385                };
386                (sql::Statement::from(op.stmt), op.params, ret)
387            }
388            Operation::QuerySql(op) => {
389                let ret = match op.ret {
390                    Some(types) => SqlReturn::Types(types),
391                    None => SqlReturn::Count,
392                };
393                (sql::Statement::from(op.stmt), op.params, ret)
394            }
395            Operation::RawSql(op) => {
396                let ret = match op.ret {
397                    RawSqlRet::None => SqlReturn::Count,
398                    RawSqlRet::Infer => SqlReturn::Infer,
399                    RawSqlRet::Types(types) => SqlReturn::Types(types),
400                };
401                let mut log = QueryLog::sql(
402                    &self.query_log,
403                    driver,
404                    &op.sql,
405                    op.params.iter().map(|tv| &tv.value),
406                );
407                let mut args = MySqlArguments::default();
408                for param in op.params {
409                    Value::from(param.value)
410                        .add_to(&mut args)
411                        .map_err(|e| record_mysql_err(&self.valid, e))?;
412                }
413                let result = self.exec_sql(&op.sql, args, ret, &mut log).await;
414                log.finish(&result);
415                return result;
416            }
417            Operation::Transaction(op) => {
418                if let Transaction::Start {
419                    mode: mode @ (TransactionMode::Immediate | TransactionMode::Exclusive),
420                    ..
421                } = &op
422                {
423                    return Err(toasty_core::Error::unsupported_feature(format!(
424                        "{} does not support TransactionMode::{mode:?}",
425                        self.capability.driver_name
426                    )));
427                }
428                let statement = serializer(self.capability, &schema.db).serialize_transaction(&op);
429                sqlx_core::raw_sql::raw_sql(AssertSqlSafe(statement))
430                    .execute(&mut self.conn)
431                    .await
432                    .map_err(|e| record_mysql_err(&self.valid, e))?;
433                return Ok(ExecResponse::count(0));
434            }
435            op => todo!("op={op:#?}"),
436        };
437
438        let (sql_as_str, arg_order) =
439            serializer(self.capability, &schema.db).serialize_with_arg_order(&sql);
440
441        let mut log = QueryLog::sql(
442            &self.query_log,
443            driver,
444            &sql_as_str,
445            arg_order.iter().map(|&pos| &typed_params[pos].value),
446        );
447
448        // MySQL uses positional `?` placeholders, so parameters must follow the
449        // order in which their `Expr::Arg(n)` values appear in serialized SQL.
450        let mut remaining = vec![0usize; typed_params.len()];
451        for &pos in &arg_order {
452            remaining[pos] += 1;
453        }
454        let mut values = typed_params
455            .into_iter()
456            .map(|param| Some(param.value))
457            .collect::<Vec<_>>();
458        let mut args = MySqlArguments::default();
459        for pos in arg_order {
460            remaining[pos] -= 1;
461            let value = if remaining[pos] == 0 {
462                values[pos].take().expect("MySQL parameter already moved")
463            } else {
464                values[pos]
465                    .as_ref()
466                    .expect("MySQL parameter missing")
467                    .clone()
468            };
469            Value::from(value)
470                .add_to(&mut args)
471                .map_err(|e| record_mysql_err(&self.valid, e))?;
472        }
473
474        let result = self.exec_sql(&sql_as_str, args, ret, &mut log).await;
475        log.finish(&result);
476        result
477    }
478
479    async fn push_schema(&mut self, schema: &Schema) -> Result<()> {
480        for table in &schema.db.tables {
481            tracing::debug!(table = %table.name, "creating table");
482            self.create_table(&schema.db, table).await?;
483        }
484        Ok(())
485    }
486
487    async fn applied_migrations(
488        &mut self,
489    ) -> Result<Vec<toasty_core::schema::db::AppliedMigration>> {
490        sqlx_core::query::query(
491            "CREATE TABLE IF NOT EXISTS __toasty_migrations (
492                id BIGINT UNSIGNED PRIMARY KEY,
493                name TEXT NOT NULL,
494                applied_at TIMESTAMP NOT NULL
495            )",
496        )
497        .execute(&mut self.conn)
498        .await
499        .map_err(|e| record_mysql_err(&self.valid, e))?;
500
501        let ids = sqlx_core::query_scalar::query_scalar::<sqlx_mysql::MySql, u64>(
502            "SELECT id FROM __toasty_migrations ORDER BY applied_at",
503        )
504        .fetch_all(&mut self.conn)
505        .await
506        .map_err(|e| record_mysql_err(&self.valid, e))?;
507
508        Ok(ids
509            .into_iter()
510            .map(toasty_core::schema::db::AppliedMigration::new)
511            .collect())
512    }
513
514    async fn apply_migration(
515        &mut self,
516        id: u64,
517        name: &str,
518        migration: &toasty_core::schema::db::Migration,
519    ) -> Result<()> {
520        tracing::info!(id, name, "applying migration");
521        sqlx_core::query::query(
522            "CREATE TABLE IF NOT EXISTS __toasty_migrations (
523                id BIGINT UNSIGNED PRIMARY KEY,
524                name TEXT NOT NULL,
525                applied_at TIMESTAMP NOT NULL
526            )",
527        )
528        .execute(&mut self.conn)
529        .await
530        .map_err(|e| record_mysql_err(&self.valid, e))?;
531
532        let mut transaction = self
533            .conn
534            .begin()
535            .await
536            .map_err(|e| record_mysql_err(&self.valid, e))?;
537
538        for statement in migration.statements() {
539            if let Err(error) = sqlx_core::raw_sql::raw_sql(AssertSqlSafe(statement))
540                .execute(&mut *transaction)
541                .await
542            {
543                let error = record_mysql_err(&self.valid, error);
544                transaction
545                    .rollback()
546                    .await
547                    .map_err(|e| record_mysql_err(&self.valid, e))?;
548                return Err(error);
549            }
550        }
551
552        if let Err(error) = sqlx_core::query::query(
553            "INSERT INTO __toasty_migrations (id, name, applied_at) VALUES (?, ?, NOW())",
554        )
555        .bind(id)
556        .bind(name)
557        .execute(&mut *transaction)
558        .await
559        {
560            let error = record_mysql_err(&self.valid, error);
561            transaction
562                .rollback()
563                .await
564                .map_err(|e| record_mysql_err(&self.valid, e))?;
565            return Err(error);
566        }
567
568        transaction
569            .commit()
570            .await
571            .map_err(|e| record_mysql_err(&self.valid, e))?;
572        Ok(())
573    }
574
575    fn is_valid(&self) -> bool {
576        self.valid.get()
577    }
578
579    async fn ping(&mut self) -> Result<()> {
580        match self.conn.ping().await {
581            Ok(()) => Ok(()),
582            Err(error) => {
583                self.valid.set(false);
584                Err(toasty_core::Error::connection_lost(error))
585            }
586        }
587    }
588}