1#![warn(missing_docs)]
2#![allow(clippy::needless_range_loop)]
3
4mod 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
53fn 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
82fn 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#[derive(Debug)]
101pub struct MySQL {
102 url: String,
103 opts: MySqlConnectOptions,
104 capability: &'static Capability,
105}
106
107impl MySQL {
108 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
159fn 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#[derive(Debug)]
229pub struct Connection {
230 conn: MySqlConnection,
231 valid: Cell<bool>,
234 query_log: QueryLogConfig,
235 capability: &'static Capability,
236}
237
238impl Connection {
239 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 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 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 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}