use anyhow::Result;
use sqlx::mysql::{MySqlConnectOptions, MySqlPool, MySqlQueryResult, MySqlRow, MySqlSslMode};
use sqlx::{AssertSqlSafe, Executor, Row, TypeInfo, ValueRef};
use super::config::Engine;
use super::query::{self, Cell};
use super::sql::quote_literal_for;
use super::{ConnectionConfig, SslConfig, SslMode, decode, pool_options};
pub(crate) const DATABASES_SQL: &str = "SELECT schema_name FROM information_schema.schemata \
WHERE schema_name NOT IN ('information_schema', 'performance_schema', 'mysql', 'sys') \
ORDER BY schema_name";
pub(crate) const OBJECTS_SQL: &str = "SELECT table_schema, table_name, table_type \
FROM information_schema.tables WHERE table_schema = DATABASE() ORDER BY table_name";
pub(crate) const OBJECTS_FOR_DB_SQL: &str = "SELECT table_schema, table_name, table_type \
FROM information_schema.tables WHERE table_schema = ? ORDER BY table_name";
pub(crate) const ROUTINES_FOR_DB_SQL: &str = "SELECT r.routine_schema, r.routine_name, r.routine_type, \
(SELECT GROUP_CONCAT(p.dtd_identifier ORDER BY p.ordinal_position SEPARATOR ', ') \
FROM information_schema.parameters p \
WHERE p.specific_schema = r.routine_schema \
AND p.specific_name = r.specific_name \
AND p.ordinal_position > 0) \
FROM information_schema.routines r \
WHERE r.routine_schema = ? ORDER BY r.routine_name";
pub(crate) const COLUMNS_FOR_DB_SQL: &str = "SELECT c.table_schema, c.table_name, t.table_type, \
c.column_name, c.column_type \
FROM information_schema.columns c \
JOIN information_schema.tables t \
ON t.table_schema = c.table_schema AND t.table_name = c.table_name \
WHERE c.table_schema = ? \
ORDER BY c.table_name, c.ordinal_position";
pub(crate) const PROCESSES_SQL: &str = "SELECT id, user, host, db, command, time, state, info \
FROM information_schema.processlist \
WHERE id <> CONNECTION_ID() \
ORDER BY time DESC";
pub(crate) const VARIABLES_SQL: &str = "SELECT v.VARIABLE_NAME, v.VARIABLE_VALUE, i.VARIABLE_SOURCE, \
CASE WHEN i.VARIABLE_SOURCE = 'COMPILED' THEN '' ELSE 'yes' END AS changed \
FROM performance_schema.global_variables v \
JOIN performance_schema.variables_info i ON v.VARIABLE_NAME = i.VARIABLE_NAME \
ORDER BY v.VARIABLE_NAME";
pub(crate) const DIGEST_AVAILABLE_SQL: &str = "SHOW VARIABLES LIKE 'performance_schema'";
pub(crate) const DIGEST_SQL: &str = "SELECT digest_text, count_star, \
round(sum_timer_wait / 1000000000.0, 2) AS total_time_ms, \
round(avg_timer_wait / 1000000000.0, 2) AS avg_time_ms, sum_rows_examined \
FROM performance_schema.events_statements_summary_by_digest \
WHERE digest_text IS NOT NULL \
ORDER BY avg_timer_wait DESC LIMIT 200";
pub(crate) const ROUTINES_SQL: &str = "SELECT r.routine_schema, r.routine_name, r.routine_type, \
(SELECT GROUP_CONCAT(p.dtd_identifier ORDER BY p.ordinal_position SEPARATOR ', ') \
FROM information_schema.parameters p \
WHERE p.specific_schema = r.routine_schema \
AND p.specific_name = r.specific_name \
AND p.ordinal_position > 0) \
FROM information_schema.routines r \
WHERE r.routine_schema = DATABASE() ORDER BY r.routine_name";
pub(crate) const CATALOG_COLUMNS_SQL: &str = "SELECT c.table_schema, c.table_name, t.table_type, \
c.column_name, c.column_type \
FROM information_schema.columns c \
JOIN information_schema.tables t \
ON t.table_schema = c.table_schema AND t.table_name = c.table_name \
WHERE c.table_schema = DATABASE() \
ORDER BY c.table_name, c.ordinal_position";
pub(crate) const CATALOG_INDEXES_SQL: &str = "SELECT table_schema, table_name, 'BASE TABLE', index_name, \
GROUP_CONCAT(column_name ORDER BY seq_in_index SEPARATOR ', ') \
FROM information_schema.statistics \
WHERE table_schema = DATABASE() \
GROUP BY table_schema, table_name, index_name \
ORDER BY table_name, index_name";
pub(crate) const CATALOG_TRIGGERS_SQL: &str = "SELECT trigger_schema, event_object_table, 'BASE TABLE', trigger_name, \
CONCAT(action_timing, ' ', event_manipulation) \
FROM information_schema.triggers \
WHERE trigger_schema = DATABASE() \
ORDER BY event_object_table, trigger_name";
pub(crate) async fn connect(
config: &ConnectionConfig,
password: Option<&str>,
) -> Result<MySqlPool> {
let options = options(config, password);
let mut pool_options = pool_options();
let read_only = config.safety.is_read_only();
let timeout = config.statement_timeout.filter(|&seconds| seconds > 0);
if read_only || timeout.is_some() {
pool_options =
pool_options.after_connect(move |connection: &mut sqlx::MySqlConnection, _| {
Box::pin(async move {
if read_only {
connection
.execute("SET SESSION TRANSACTION READ ONLY")
.await?;
}
if let Some(seconds) = timeout {
set_statement_timeout(connection, seconds).await;
}
Ok(())
})
});
}
let pool = pool_options.connect_with(options).await?;
Ok(pool)
}
fn ssl_mode(mode: SslMode) -> MySqlSslMode {
match mode {
SslMode::Disable => MySqlSslMode::Disabled,
SslMode::Prefer => MySqlSslMode::Preferred,
SslMode::Require => MySqlSslMode::Required,
SslMode::VerifyCa => MySqlSslMode::VerifyCa,
SslMode::VerifyFull => MySqlSslMode::VerifyIdentity,
}
}
fn options(config: &ConnectionConfig, password: Option<&str>) -> MySqlConnectOptions {
let mut options = MySqlConnectOptions::new()
.host(&config.host)
.port(config.port)
.username(&config.username);
if !config.database.is_empty() {
options = options.database(&config.database);
}
if let Some(password) = password {
options = options.password(password);
}
let ssl = &config.ssl;
options = options.ssl_mode(ssl_mode(ssl.mode));
if ssl.mode.uses_files() {
if let Some(path) = SslConfig::path(&ssl.ca_cert) {
options = options.ssl_ca(path);
}
if let Some(path) = SslConfig::path(&ssl.client_cert) {
options = options.ssl_client_cert(path);
}
if let Some(path) = SslConfig::path(&ssl.client_key) {
options = options.ssl_client_key(path);
}
}
options
}
async fn set_statement_timeout(connection: &mut sqlx::MySqlConnection, seconds: u32) {
let mysql = format!(
"SET SESSION max_execution_time = {}",
u64::from(seconds) * 1000
);
if sqlx::raw_sql(AssertSqlSafe(mysql))
.execute(&mut *connection)
.await
.is_ok()
{
return;
}
let mariadb = format!("SET SESSION max_statement_time = {seconds}");
let _ = sqlx::raw_sql(AssertSqlSafe(mariadb))
.execute(&mut *connection)
.await;
}
pub(crate) fn primary_key_sql(table: &str) -> String {
format!(
"SELECT kcu.column_name FROM information_schema.table_constraints tc \
JOIN information_schema.key_column_usage kcu \
ON kcu.constraint_name = tc.constraint_name \
AND kcu.table_schema = tc.table_schema \
AND kcu.table_name = tc.table_name \
WHERE tc.constraint_type = 'PRIMARY KEY' \
AND tc.table_schema = DATABASE() AND tc.table_name = {} \
ORDER BY kcu.ordinal_position",
quote_literal_for(Engine::MySql, table)
)
}
pub(crate) fn columns_sql(table: &str) -> String {
format!(
"SELECT column_name, column_type, (is_nullable = 'YES'), column_default, \
extra, collation_name, column_comment, generation_expression \
FROM information_schema.columns \
WHERE table_schema = DATABASE() AND table_name = {} \
ORDER BY ordinal_position",
quote_literal_for(Engine::MySql, table)
)
}
pub(crate) fn indexes_sql(table: &str) -> String {
format!(
"SELECT index_name, \
GROUP_CONCAT(column_name ORDER BY seq_in_index SEPARATOR ','), \
(non_unique = 0), (index_name = 'PRIMARY') \
FROM information_schema.statistics \
WHERE table_schema = DATABASE() AND table_name = {} \
GROUP BY index_name, non_unique \
ORDER BY index_name",
quote_literal_for(Engine::MySql, table)
)
}
pub(crate) fn foreign_keys_sql(table: &str) -> String {
format!(
"SELECT kcu.constraint_name, \
GROUP_CONCAT(kcu.column_name ORDER BY kcu.ordinal_position SEPARATOR ','), \
kcu.referenced_table_schema, kcu.referenced_table_name, \
GROUP_CONCAT(kcu.referenced_column_name ORDER BY kcu.ordinal_position SEPARATOR ','), \
rc.delete_rule, rc.update_rule \
FROM information_schema.key_column_usage kcu \
JOIN information_schema.referential_constraints rc \
ON rc.constraint_name = kcu.constraint_name \
AND rc.constraint_schema = kcu.constraint_schema \
WHERE kcu.table_schema = DATABASE() AND kcu.table_name = {} \
AND kcu.referenced_table_name IS NOT NULL \
GROUP BY kcu.constraint_name, kcu.referenced_table_schema, \
kcu.referenced_table_name, rc.delete_rule, rc.update_rule \
ORDER BY kcu.constraint_name",
quote_literal_for(Engine::MySql, table)
)
}
pub(crate) fn rows_affected(result: &MySqlQueryResult) -> u64 {
result.rows_affected()
}
pub(crate) fn unpreparable(error: &anyhow::Error) -> bool {
error
.downcast_ref::<sqlx::Error>()
.and_then(|error| error.as_database_error())
.and_then(|error| error.try_downcast_ref::<sqlx::mysql::MySqlDatabaseError>())
.is_some_and(|error| error.number() == 1295)
}
pub(crate) fn cell(row: &MySqlRow, index: usize) -> Cell {
let Ok(raw) = row.try_get_raw(index) else {
return None;
};
if raw.is_null() {
return None;
}
let type_name = raw.type_info().name().to_string();
let value = match type_name.as_str() {
"BOOLEAN" | "TINYINT" => decode!(row, index, i8).or_else(|| decode!(row, index, u8)),
"TINYINT UNSIGNED" => decode!(row, index, u8),
"SMALLINT" => decode!(row, index, i16),
"SMALLINT UNSIGNED" => decode!(row, index, u16),
"INT" | "MEDIUMINT" => decode!(row, index, i32),
"INT UNSIGNED" | "MEDIUMINT UNSIGNED" => decode!(row, index, u32),
"BIGINT" => decode!(row, index, i64),
"BIGINT UNSIGNED" => decode!(row, index, u64),
"FLOAT" => decode!(row, index, f32),
"DOUBLE" => decode!(row, index, f64),
"DECIMAL" => decode!(row, index, sqlx::types::BigDecimal),
"JSON" => decode!(row, index, sqlx::types::JsonValue),
"DATE" => decode!(row, index, chrono::NaiveDate),
"YEAR" => decode!(row, index, u16),
"BIT" => row
.try_get::<u64, _>(index)
.ok()
.map(|bits| format!("{bits:b}")),
"DATETIME" => decode!(row, index, chrono::NaiveDateTime),
"TIMESTAMP" => row
.try_get::<chrono::DateTime<chrono::Utc>, _>(index)
.ok()
.map(|value| value.naive_utc().to_string()),
"TIME" => decode!(row, index, chrono::NaiveTime)
.or_else(|| decode!(row, index, sqlx::mysql::types::MySqlTime)),
"BLOB" | "TINYBLOB" | "MEDIUMBLOB" | "LONGBLOB" | "BINARY" | "VARBINARY" => {
decode!(row, index, String).or_else(|| {
row.try_get::<Vec<u8>, _>(index)
.ok()
.and_then(|bytes| query::blob(&bytes))
})
}
_ => decode!(row, index, String),
};
value.or_else(|| query::unsupported(&type_name))
}
pub(crate) fn raw_bytes(row: &MySqlRow, index: usize) -> Option<Vec<u8>> {
row.try_get::<Option<Vec<u8>>, _>(index).ok().flatten()
}
#[cfg(test)]
mod tests {
use super::*;
fn ssl(mode: SslMode) -> ConnectionConfig {
ConnectionConfig {
ssl: SslConfig {
mode,
ca_cert: "/etc/zippa/ca.pem".into(),
client_cert: "/etc/zippa/client.pem".into(),
client_key: "/etc/zippa/client.key".into(),
},
..ConnectionConfig::new(Engine::MySql)
}
}
#[test]
fn every_ssl_mode_reaches_the_driver() {
for (mode, expected) in [
(SslMode::Disable, "Disabled"),
(SslMode::Prefer, "Preferred"),
(SslMode::Require, "Required"),
(SslMode::VerifyCa, "VerifyCa"),
(SslMode::VerifyFull, "VerifyIdentity"),
] {
let options = options(&ssl(mode), None);
assert_eq!(format!("{:?}", options.get_ssl_mode()), expected);
}
}
#[test]
fn certificate_files_are_passed_only_when_encrypting() {
let verify = format!("{:?}", options(&ssl(SslMode::VerifyCa), None));
for file in ["ca.pem", "client.pem", "client.key"] {
assert!(verify.contains(file), "{file} missing from {verify}");
}
for mode in [SslMode::Disable, SslMode::Prefer] {
let plain = format!("{:?}", options(&ssl(mode), None));
assert!(!plain.contains("/etc/zippa"), "{plain}");
}
}
}