use std::sync::Arc;
use std::time::{Duration, Instant};
use anyhow::{Context as _, Result};
use futures::StreamExt as _;
use serde::{Deserialize, Serialize};
use sqlx::Either;
use sqlx::{
AssertSqlSafe, Column, Database, Encode, Executor, IntoArguments, Row, SqlSafeStr, Type,
TypeInfo,
};
use super::catalog::{Catalog, CatalogEntry, CatalogKind, MAX_ENTRIES};
use super::config::{ConnectionConfig, Engine};
use super::dedicated::Dedicated;
use super::import::{self, Dialect, ImportProgress, ImportRequest, ImportSummary};
use super::plan::{self, Explained, Plan};
use super::query::{Cell, QueryResult};
use super::query_log::{LoggedQuery, QueryLog, QueryOutcome, QuerySource};
use super::tunnel::Tunnel;
use super::{
health, mysql, postgres, quote_identifier, quote_literal, sqlite, statement, typed_placeholder,
};
pub(crate) const POOL_SIZE: u32 = 5;
pub(crate) const PINNED_MAX: u32 = 8;
pub(crate) const POOL_MAX: u32 = POOL_SIZE + PINNED_MAX;
pub(crate) fn pool_options<DB: Database>() -> sqlx::pool::PoolOptions<DB> {
sqlx::pool::PoolOptions::new()
.max_connections(POOL_MAX)
.acquire_timeout(health::ACQUIRE_TIMEOUT)
}
#[derive(Debug)]
enum Pool {
Postgres(sqlx::PgPool),
MySql(sqlx::MySqlPool),
Sqlite(sqlx::SqlitePool),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct DatabaseObject {
pub schema: Option<String>,
pub name: String,
pub kind: ObjectKind,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum ObjectKind {
Table,
View,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StoredObject {
pub schema: Option<String>,
pub name: String,
pub arguments: Option<String>,
pub kind: StoredKind,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StoredKind {
Function,
Procedure,
Sequence,
}
impl StoredObject {
pub fn label(&self) -> String {
let name = qualified(self.schema.as_deref(), &self.name);
match &self.arguments {
Some(arguments) => format!("{name}({arguments})"),
None => name,
}
}
}
impl DatabaseObject {
pub fn label(&self) -> String {
qualified(self.schema.as_deref(), &self.name)
}
}
fn qualified(schema: Option<&str>, name: &str) -> String {
match schema {
Some(schema) => format!("{schema}.{name}"),
None => name.to_string(),
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RowKey {
Columns(Vec<String>),
RowId(&'static str),
Unavailable(&'static str),
}
impl RowKey {
pub fn is_available(&self) -> bool {
!matches!(self, RowKey::Unavailable(_))
}
}
#[derive(Debug)]
pub enum QueryDigest {
Available(QueryResult),
Unavailable(String),
}
#[derive(Debug)]
pub struct Connection {
pub config: ConnectionConfig,
password: Option<String>,
pool: Pool,
money_scale: i64,
query_log: QueryLog,
pinned: Arc<tokio::sync::Semaphore>,
tunnel: Option<Arc<Tunnel>>,
ssh_secret: Option<String>,
}
#[derive(Clone, Default)]
pub struct Credentials {
pub password: Option<String>,
pub ssh: Option<String>,
}
impl Connection {
#[cfg(test)]
pub async fn open(config: ConnectionConfig, password: Option<String>) -> Result<Self> {
Self::open_with(
config,
Credentials {
password,
ssh: None,
},
)
.await
}
pub async fn open_with(config: ConnectionConfig, credentials: Credentials) -> Result<Self> {
let tunnel = Self::tunnel(&config, credentials.ssh.as_deref()).await?;
Self::open_through(config, credentials, tunnel).await
}
async fn tunnel(
config: &ConnectionConfig,
secret: Option<&str>,
) -> Result<Option<Arc<Tunnel>>> {
if !config.ssh.enabled || config.engine.is_file_based() {
return Ok(None);
}
let tunnel = Tunnel::open(&config.ssh, secret, &config.host, config.port).await?;
Ok(Some(Arc::new(tunnel)))
}
async fn open_through(
config: ConnectionConfig,
credentials: Credentials,
tunnel: Option<Arc<Tunnel>>,
) -> Result<Self> {
let Credentials {
password,
ssh: ssh_secret,
} = credentials;
let mut target = config.clone();
if let Some(tunnel) = &tunnel {
target.host = "127.0.0.1".into();
target.port = tunnel.local_port();
}
let password = password.as_deref();
let pool = match config.engine {
Engine::Postgres => Pool::Postgres(postgres::connect(&target, password).await?),
Engine::MySql => Pool::MySql(mysql::connect(&target, password).await?),
Engine::Sqlite => Pool::Sqlite(sqlite::connect(&target).await?),
};
let money_scale = match &pool {
Pool::Postgres(pool) => postgres::money_scale(pool).await,
Pool::MySql(_) | Pool::Sqlite(_) => postgres::DEFAULT_MONEY_SCALE,
};
Ok(Self {
config,
password: password.map(str::to_string),
pool,
money_scale,
query_log: QueryLog::default(),
pinned: Arc::new(tokio::sync::Semaphore::new(PINNED_MAX as usize)),
tunnel,
ssh_secret,
})
}
pub(crate) fn pinned_permit(&self) -> Result<tokio::sync::OwnedSemaphorePermit> {
self.pinned.clone().try_acquire_owned().map_err(|_| {
anyhow::anyhow!(
"too many query tabs are holding connections ({PINNED_MAX}). Close one to \
open another."
)
})
}
pub(crate) async fn cancel_backend(&self, backend: u64) -> Result<()> {
match &self.pool {
Pool::Postgres(pool) => {
sqlx::query("SELECT pg_cancel_backend($1)")
.bind(backend as i64)
.execute(pool)
.await?;
}
Pool::MySql(pool) => {
sqlx::raw_sql(AssertSqlSafe(format!("KILL QUERY {backend}")))
.execute(pool)
.await?;
}
Pool::Sqlite(_) => {}
}
Ok(())
}
pub fn query_log(&self) -> &QueryLog {
&self.query_log
}
pub fn database(&self) -> &str {
&self.config.database
}
pub async fn with_database(&self, database: &str) -> Result<Self> {
let mut config = self.config.clone();
config.database = database.to_string();
let tunnel = match &self.tunnel {
Some(tunnel) if !tunnel.is_closed() => Some(tunnel.clone()),
Some(_) => Self::tunnel(&config, self.ssh_secret.as_deref()).await?,
None => None,
};
let credentials = Credentials {
password: self.password.clone(),
ssh: self.ssh_secret.clone(),
};
Self::open_through(config, credentials, tunnel).await
}
pub async fn databases(&self) -> Result<Vec<String>> {
let sql = match self.config.engine {
Engine::Postgres => postgres::DATABASES_SQL,
Engine::MySql => mysql::DATABASES_SQL,
Engine::Sqlite => sqlite::DATABASES_SQL,
};
let result = self.run_query(sql).await?;
Ok(result
.rows
.iter()
.filter_map(|row| row.first().cloned().flatten())
.collect())
}
pub async fn objects(&self) -> Result<Vec<DatabaseObject>> {
let sql = match self.config.engine {
Engine::Postgres => postgres::OBJECTS_SQL,
Engine::MySql => mysql::OBJECTS_SQL,
Engine::Sqlite => sqlite::OBJECTS_SQL,
};
let result = self.run_query(sql).await?;
Ok(result
.rows
.iter()
.filter_map(|row| self.object_row(row))
.collect())
}
fn object_row(&self, row: &[Cell]) -> Option<DatabaseObject> {
let name = row.get(1)?.clone()?;
let kind = match row.get(2).and_then(|cell| cell.as_deref()) {
Some(kind) if kind.eq_ignore_ascii_case("VIEW") => ObjectKind::View,
_ => ObjectKind::Table,
};
let schema = self.shown_schema(row.first().cloned().flatten());
Some(DatabaseObject { schema, name, kind })
}
pub async fn stored_objects(&self) -> Result<Vec<StoredObject>> {
let sql = match self.config.engine {
Engine::Postgres => postgres::ROUTINES_SQL,
Engine::MySql => mysql::ROUTINES_SQL,
Engine::Sqlite => return Ok(Vec::new()),
};
let result = self.run_query(sql).await?;
Ok(result
.rows
.iter()
.filter_map(|row| self.stored_row(row))
.collect())
}
fn stored_row(&self, row: &[Cell]) -> Option<StoredObject> {
let name = row.get(1)?.clone()?;
let kind = match row.get(2).and_then(|cell| cell.as_deref())? {
kind if kind.eq_ignore_ascii_case("PROCEDURE") => StoredKind::Procedure,
kind if kind.eq_ignore_ascii_case("SEQUENCE") => StoredKind::Sequence,
_ => StoredKind::Function,
};
let schema = self.shown_schema(row.first().cloned().flatten());
let arguments = match kind {
StoredKind::Sequence => None,
_ => Some(row.get(3).cloned().flatten().unwrap_or_default()),
};
Some(StoredObject {
schema,
name,
arguments,
kind,
})
}
pub async fn catalog(&self) -> Result<Catalog> {
let (columns, indexes, triggers) = match self.config.engine {
Engine::Postgres => (
postgres::CATALOG_COLUMNS_SQL,
postgres::CATALOG_INDEXES_SQL,
postgres::CATALOG_TRIGGERS_SQL,
),
Engine::MySql => (
mysql::CATALOG_COLUMNS_SQL,
mysql::CATALOG_INDEXES_SQL,
mysql::CATALOG_TRIGGERS_SQL,
),
Engine::Sqlite => (
sqlite::CATALOG_COLUMNS_SQL,
sqlite::CATALOG_INDEXES_SQL,
sqlite::CATALOG_TRIGGERS_SQL,
),
};
let objects = self.objects().await?;
let stored = self.stored_objects().await?;
let columns = self.run_query(columns).await?;
let indexes = self.run_query(indexes).await?;
let triggers = self.run_query(triggers).await?;
let total = objects.len()
+ stored.len()
+ columns.rows.len()
+ indexes.rows.len()
+ triggers.rows.len();
let mut entries: Vec<CatalogEntry> = Vec::with_capacity(total.min(MAX_ENTRIES));
entries.extend(objects.into_iter().map(CatalogEntry::object));
entries.extend(stored.into_iter().map(CatalogEntry::routine));
let members = [
(&columns.rows, CatalogKind::Column),
(&indexes.rows, CatalogKind::Index),
(&triggers.rows, CatalogKind::Trigger),
]
.into_iter()
.flat_map(|(rows, kind)| rows.iter().map(move |row| (kind, row)));
for (kind, row) in members {
if entries.len() >= MAX_ENTRIES {
break;
}
let (owner, name, detail) = self.catalog_row(row);
entries.push(CatalogEntry::member(kind, owner, name, detail));
}
entries.truncate(MAX_ENTRIES);
Ok(Catalog { entries, total })
}
pub async fn catalog_for_database(&self, database: &str) -> Result<Catalog> {
let (objects, stored, columns) = match self.config.engine {
Engine::MySql => {
let param = vec![Some(database.to_string())];
(
self.run_query_with(mysql::OBJECTS_FOR_DB_SQL, param.clone())
.await?,
self.run_query_with(mysql::ROUTINES_FOR_DB_SQL, param.clone())
.await?,
self.run_query_with(mysql::COLUMNS_FOR_DB_SQL, param)
.await?,
)
}
Engine::Sqlite => {
if database.eq_ignore_ascii_case("main") {
return self.catalog().await;
}
let db = quote_identifier(database, Engine::Sqlite);
let name = quote_literal(database);
let objects = self
.run_query(&format!(
"SELECT NULL AS table_schema, name, \
CASE type WHEN 'view' THEN 'VIEW' ELSE 'BASE TABLE' END AS table_type \
FROM {db}.sqlite_master \
WHERE type IN ('table', 'view') AND name NOT LIKE 'sqlite_%' \
ORDER BY name"
))
.await?;
let columns = self
.run_query(&format!(
"SELECT NULL, m.name, m.type, ti.name, ti.type \
FROM {db}.sqlite_master m \
JOIN pragma_table_info(m.name, {name}) ti \
WHERE m.type IN ('table', 'view') AND m.name NOT LIKE 'sqlite_%' \
ORDER BY m.name, ti.cid"
))
.await?;
(objects, QueryResult::default(), columns)
}
Engine::Postgres => {
anyhow::bail!("cross-database completion is not supported on PostgreSQL")
}
};
let total = objects.rows.len() + stored.rows.len() + columns.rows.len();
let mut entries: Vec<CatalogEntry> = Vec::with_capacity(total.min(MAX_ENTRIES));
entries.extend(
objects
.rows
.iter()
.filter_map(|row| self.object_row(row))
.map(CatalogEntry::object),
);
entries.extend(
stored
.rows
.iter()
.filter_map(|row| self.stored_row(row))
.map(CatalogEntry::routine),
);
for row in &columns.rows {
if entries.len() >= MAX_ENTRIES {
break;
}
let (owner, name, detail) = self.catalog_row(row);
entries.push(CatalogEntry::member(
CatalogKind::Column,
owner,
name,
detail,
));
}
entries.truncate(MAX_ENTRIES);
Ok(Catalog { entries, total })
}
fn shown_schema(&self, schema: Option<String>) -> Option<String> {
match self.config.engine {
Engine::Postgres => schema.filter(|schema| schema != "public"),
Engine::MySql | Engine::Sqlite => None,
}
}
fn catalog_row(&self, row: &[Cell]) -> (DatabaseObject, String, String) {
let text = |index: usize| {
row.get(index)
.and_then(|cell| cell.clone())
.unwrap_or_default()
};
let schema = self.shown_schema(Some(text(0)));
let kind = match text(2) {
kind if kind.eq_ignore_ascii_case("VIEW") => ObjectKind::View,
_ => ObjectKind::Table,
};
(
DatabaseObject {
schema,
name: text(1),
kind,
},
text(3),
text(4),
)
}
pub async fn run_query(&self, sql: &str) -> Result<QueryResult> {
self.run_query_with(sql, Vec::new()).await
}
pub async fn run_query_with(&self, sql: &str, params: Vec<Cell>) -> Result<QueryResult> {
self.run_query_tagged(sql, params, QuerySource::Internal)
.await
}
async fn run_query_tagged(
&self,
sql: &str,
params: Vec<Cell>,
source: QuerySource,
) -> Result<QueryResult> {
self.refuse_write(sql)?;
let started = Instant::now();
let result = self.fetch(sql, params).await;
let outcome = match &result {
Ok(result) => query_outcome(result),
Err(error) => QueryOutcome::Error(format!("{error:#}")),
};
self.log(sql, source, started.elapsed(), outcome);
result
}
pub(crate) fn log(
&self,
sql: &str,
source: QuerySource,
elapsed: Duration,
outcome: QueryOutcome,
) {
self.query_log.record(LoggedQuery {
sql: sql.to_string(),
source,
outcome,
elapsed,
at: std::time::SystemTime::now(),
});
}
async fn fetch(&self, sql: &str, params: Vec<Cell>) -> Result<QueryResult> {
let result = match &self.pool {
Pool::Postgres(pool) => {
let scale = self.money_scale;
fetch_all(
pool,
sql,
params,
move |row, index| postgres::cell(row, index, scale),
postgres::rows_affected,
)
.await
}
Pool::MySql(pool) => {
fetch_all(pool, sql, params, mysql::cell, mysql::rows_affected).await
}
Pool::Sqlite(pool) => {
fetch_all(pool, sql, params, sqlite::cell, sqlite::rows_affected).await
}
};
result.map_err(health::plain)
}
pub async fn fetch_binary(&self, sql: &str, params: Vec<Cell>) -> Result<Option<Vec<u8>>> {
self.refuse_write(sql)?;
let result = match &self.pool {
Pool::Postgres(pool) => {
fetch_binary_column(pool, sql, params, postgres::raw_bytes).await
}
Pool::MySql(pool) => fetch_binary_column(pool, sql, params, mysql::raw_bytes).await,
Pool::Sqlite(pool) => fetch_binary_column(pool, sql, params, sqlite::raw_bytes).await,
};
result.map_err(health::plain)
}
pub async fn explain(&self, sql: &str, analyze: bool) -> Result<Explained> {
let started = Instant::now();
let result = self.explain_inner(sql, analyze).await;
let outcome = match &result {
Ok(_) => QueryOutcome::Ran,
Err(error) => QueryOutcome::Error(format!("{error:#}")),
};
self.log(sql, QuerySource::User, started.elapsed(), outcome);
result
}
async fn explain_inner(&self, sql: &str, analyze: bool) -> Result<Explained> {
let statement = sql.trim().trim_end_matches(';').trim();
if statement.is_empty() || statement::split(statement, self.config.engine).len() != 1 {
anyhow::bail!("Select one statement to explain.");
}
let header = statement::explained(statement, self.config.engine);
let (inner, runs) = match &header {
Some((inner, analyzes)) => (inner.clone(), analyze || *analyzes),
None => (statement.to_string(), analyze),
};
if runs && let Some(word) = statement::first_write(&inner, self.config.engine) {
anyhow::bail!("EXPLAIN ANALYZE would run this statement; it changes data ({word}).");
}
let already = header.is_some();
match self.config.engine {
Engine::Postgres => {
let sql = match (already, analyze) {
(true, _) => statement.to_string(),
(false, true) => {
format!("EXPLAIN (ANALYZE, BUFFERS, FORMAT JSON) {inner}")
}
(false, false) => format!("EXPLAIN (FORMAT JSON) {inner}"),
};
let result = self.fetch(&sql, Vec::new()).await?;
let text = result
.rows
.first()
.and_then(|row| row.first())
.and_then(|cell| cell.clone())
.unwrap_or_else(|| "(no plan returned)".to_string());
let mut plan = plan::postgres(&text, analyze).unwrap_or_else(|| Plan::raw(&text));
plan.elapsed = result.elapsed;
Ok(Explained::Plan(plan))
}
Engine::MySql => {
let tree = match (already, analyze) {
(true, _) => statement.to_string(),
(false, true) => format!("EXPLAIN ANALYZE {inner}"),
(false, false) => format!("EXPLAIN FORMAT=TREE {inner}"),
};
if let Ok(result) = self.fetch(&tree, Vec::new()).await {
let text = result
.rows
.first()
.and_then(|row| row.first())
.and_then(|cell| cell.clone())
.unwrap_or_default();
if let Some(mut plan) = plan::mysql(&text, analyze) {
plan.elapsed = result.elapsed;
return Ok(Explained::Plan(plan));
}
}
let classic = format!("EXPLAIN {inner}");
let result = self.fetch(&classic, Vec::new()).await?;
Ok(Explained::Rows(result))
}
Engine::Sqlite => {
if runs {
anyhow::bail!(
"SQLite has no EXPLAIN ANALYZE, so actual times are not available. \
Use Explain to see the query plan."
);
}
let sql = if already {
statement.to_string()
} else {
format!("EXPLAIN QUERY PLAN {inner}")
};
let result = self.fetch(&sql, Vec::new()).await?;
let mut plan = plan::sqlite(&result);
plan.elapsed = result.elapsed;
Ok(Explained::Plan(plan))
}
}
}
pub(crate) async fn dedicated(&self) -> Result<Dedicated> {
let dedicated = match &self.pool {
Pool::Postgres(pool) => Dedicated::postgres(pool, self.money_scale).await,
Pool::MySql(pool) => Dedicated::mysql(pool).await,
Pool::Sqlite(pool) => Dedicated::sqlite(pool).await,
};
dedicated.map_err(health::plain)
}
pub async fn import_dump(
&self,
request: ImportRequest,
progress: tokio::sync::mpsc::UnboundedSender<ImportProgress>,
) -> Result<ImportSummary> {
if self.config.safety.is_read_only() {
anyhow::bail!("this connection is read-only, so a dump cannot be imported");
}
let dialect = match self.config.engine {
Engine::Postgres => Dialect::Postgres,
Engine::MySql => Dialect::MySql,
Engine::Sqlite => Dialect::Sqlite,
};
let session = self.dedicated().await?;
import::run(session, dialect, &request, progress).await
}
pub async fn row_key(&self, object: &DatabaseObject) -> Result<RowKey> {
if object.kind == ObjectKind::View {
return Ok(RowKey::Unavailable("a view cannot be edited"));
}
let engine = self.config.engine;
let sql = match engine {
Engine::Postgres => postgres::primary_key_sql(
object.schema.as_deref().unwrap_or("public"),
&object.name,
),
Engine::MySql => mysql::primary_key_sql(&object.name),
Engine::Sqlite => sqlite::primary_key_sql(&object.name),
};
let result = self.run_query(&sql).await?;
let columns: Vec<String> = result
.rows
.iter()
.filter_map(|row| row.first().cloned().flatten())
.collect();
if !columns.is_empty() {
return Ok(RowKey::Columns(columns));
}
Ok(match engine {
Engine::Postgres => RowKey::RowId("ctid"),
Engine::Sqlite => RowKey::RowId("rowid"),
Engine::MySql => RowKey::Unavailable("a table without a primary key cannot be edited"),
})
}
pub async fn processes(&self) -> Result<QueryResult> {
let sql = match self.config.engine {
Engine::Postgres => postgres::PROCESSES_SQL,
Engine::MySql => mysql::PROCESSES_SQL,
Engine::Sqlite => anyhow::bail!("SQLite has no server processes to list"),
};
self.run_query(sql).await
}
pub async fn kill_process(&self, id: &str) -> Result<()> {
if self.config.safety.is_read_only() {
anyhow::bail!("this connection is read-only");
}
match self.config.engine {
Engine::Postgres => {
let sql = format!(
"select pg_terminate_backend({})",
typed_placeholder(Engine::Postgres, 1, "INT4")
);
self.execute(&sql, vec![Some(id.to_string())]).await?;
}
Engine::MySql => {
let pid: u64 = id
.parse()
.map_err(|_| anyhow::anyhow!("not a process id: {id}"))?;
self.execute(&format!("KILL {pid}"), vec![]).await?;
}
Engine::Sqlite => anyhow::bail!("SQLite has no server processes to end"),
}
Ok(())
}
pub async fn server_variables(&self) -> Result<QueryResult> {
let sql = match self.config.engine {
Engine::Postgres => postgres::VARIABLES_SQL,
Engine::MySql => mysql::VARIABLES_SQL,
Engine::Sqlite => anyhow::bail!("SQLite has no server variables to list"),
};
self.run_query(sql).await
}
pub async fn query_digest(&self) -> Result<QueryDigest> {
match self.config.engine {
Engine::Postgres => {
let available = self.run_query(postgres::DIGEST_AVAILABLE_SQL).await?;
if available.rows.is_empty() {
return Ok(QueryDigest::Unavailable(
"pg_stat_statements is not installed on this server. Enable it with \
`CREATE EXTENSION pg_stat_statements;` after adding it to \
shared_preload_libraries and restarting the server."
.to_string(),
));
}
let result = self.run_query(postgres::DIGEST_SQL).await?;
Ok(QueryDigest::Available(result))
}
Engine::MySql => {
let available = self.run_query(mysql::DIGEST_AVAILABLE_SQL).await?;
let on = available
.rows
.first()
.and_then(|row| row.get(1))
.and_then(|cell| cell.as_deref())
.is_some_and(|value| value.eq_ignore_ascii_case("ON"));
if !on {
return Ok(QueryDigest::Unavailable(
"performance_schema is off on this server. It cannot be turned on \
while the server is running — set performance_schema=ON in its \
configuration and restart."
.to_string(),
));
}
let result = self.run_query(mysql::DIGEST_SQL).await?;
Ok(QueryDigest::Available(result))
}
Engine::Sqlite => anyhow::bail!("SQLite has no query digest to show"),
}
}
pub async fn run_maintenance(&self, task: sqlite::Maintenance) -> Result<QueryResult> {
if self.config.engine != Engine::Sqlite {
anyhow::bail!("maintenance commands are only offered for SQLite");
}
if task.writes() && self.config.safety.is_read_only() {
anyhow::bail!(
"this connection is read-only, so {} was not run",
task.label()
);
}
self.run_query(task.sql()).await
}
pub(crate) fn refuse_write(&self, sql: &str) -> Result<()> {
if !self.config.safety.is_read_only() {
return Ok(());
}
if let Some(word) = statement::first_write(sql, self.config.engine) {
anyhow::bail!("this connection is read-only, so the {word} statement was not run");
}
Ok(())
}
pub async fn execute(&self, sql: &str, params: Vec<Cell>) -> Result<u64> {
if self.config.safety.is_read_only() {
anyhow::bail!("this connection is read-only");
}
let started = Instant::now();
let result = match &self.pool {
Pool::Postgres(pool) => execute_with(pool, sql, params, postgres::rows_affected).await,
Pool::MySql(pool) => execute_with(pool, sql, params, mysql::rows_affected).await,
Pool::Sqlite(pool) => execute_with(pool, sql, params, sqlite::rows_affected).await,
}
.map_err(health::plain);
let outcome = match &result {
Ok(affected) => QueryOutcome::Affected(*affected),
Err(error) => QueryOutcome::Error(format!("{error:#}")),
};
self.log(sql, QuerySource::Internal, started.elapsed(), outcome);
result
}
pub async fn execute_script(&self, statements: &[String]) -> Result<()> {
if self.config.safety.is_read_only() {
anyhow::bail!("this connection is read-only");
}
match &self.pool {
Pool::Postgres(pool) => {
let started = Instant::now();
let result = execute_script_transactional(pool, statements)
.await
.map_err(health::plain);
self.log_script(statements, started.elapsed(), &result);
result
}
Pool::Sqlite(pool) => {
let started = Instant::now();
let result = execute_script_transactional(pool, statements)
.await
.map_err(health::plain);
self.log_script(statements, started.elapsed(), &result);
result
}
Pool::MySql(_) => {
for (position, statement) in statements.iter().enumerate() {
self.execute(statement, Vec::new()).await.with_context(|| {
format!("statement {} of {}", position + 1, statements.len())
})?;
}
Ok(())
}
}
}
fn log_script(&self, statements: &[String], elapsed: Duration, result: &Result<()>) {
let outcome = match result {
Ok(()) => QueryOutcome::Ran,
Err(error) => QueryOutcome::Error(format!("{error:#}")),
};
self.log(
&statements.join("; "),
QuerySource::Internal,
elapsed,
outcome,
);
}
pub async fn rebuild_table(&self, statements: Vec<String>) -> Result<()> {
if self.config.safety.is_read_only() {
anyhow::bail!("this connection is read-only");
}
match &self.pool {
Pool::Sqlite(pool) => sqlite::rebuild(pool, &statements).await,
Pool::Postgres(_) | Pool::MySql(_) => {
anyhow::bail!("a table rebuild is only used on SQLite")
}
}
}
pub async fn close(&self) {
match &self.pool {
Pool::Postgres(pool) => pool.close().await,
Pool::MySql(pool) => pool.close().await,
Pool::Sqlite(pool) => pool.close().await,
}
}
}
pub(crate) async fn fetch_on<DB, F>(
connection: &mut DB::Connection,
sql: &str,
cell: F,
rows_affected: fn(&DB::QueryResult) -> u64,
) -> Result<QueryResult>
where
DB: Database,
for<'c> &'c mut DB::Connection: Executor<'c, Database = DB>,
<DB as Database>::Arguments: IntoArguments<DB>,
F: Fn(&DB::Row, usize) -> Cell,
{
let statement = AssertSqlSafe(sql.to_string()).into_sql_str();
let started = Instant::now();
let query = sqlx::query(statement.clone());
#[allow(deprecated)]
let mut results = query.fetch_many(&mut *connection);
let mut rows = Vec::new();
let mut affected: Option<u64> = None;
while let Some(result) = results.next().await {
match result? {
Either::Left(done) => {
*affected.get_or_insert(0) += rows_affected(&done);
}
Either::Right(row) => rows.push(row),
}
}
drop(results);
let elapsed = started.elapsed();
let (columns, column_types): (Vec<String>, Vec<String>) = match rows.first() {
Some(row) => describe_columns(row.columns()),
None => Executor::describe(&mut *connection, statement)
.await
.map(|described| describe_columns(described.columns()))
.unwrap_or_default(),
};
let rows = rows
.iter()
.map(|row| (0..columns.len()).map(|index| cell(row, index)).collect())
.collect();
Ok(QueryResult {
columns,
column_types,
rows,
elapsed,
affected,
})
}
async fn fetch_all<DB, F>(
pool: &sqlx::Pool<DB>,
sql: &str,
params: Vec<Cell>,
cell: F,
rows_affected: fn(&DB::QueryResult) -> u64,
) -> Result<QueryResult>
where
DB: Database,
for<'c> &'c sqlx::Pool<DB>: Executor<'c, Database = DB>,
<DB as Database>::Arguments: IntoArguments<DB>,
for<'q> Option<String>: Encode<'q, DB>,
String: Type<DB>,
F: Fn(&DB::Row, usize) -> Cell,
{
let statement = AssertSqlSafe(sql.to_string()).into_sql_str();
let started = Instant::now();
let mut query = sqlx::query(statement.clone());
for param in params {
query = query.bind(param);
}
#[allow(deprecated)]
let mut results = query.fetch_many(pool);
let mut rows = Vec::new();
let mut affected: Option<u64> = None;
while let Some(result) = results.next().await {
match result? {
Either::Left(done) => {
*affected.get_or_insert(0) += rows_affected(&done);
}
Either::Right(row) => rows.push(row),
}
}
drop(results);
let elapsed = started.elapsed();
let (columns, column_types): (Vec<String>, Vec<String>) = match rows.first() {
Some(row) => describe_columns(row.columns()),
None => Executor::describe(pool, statement)
.await
.map(|described| describe_columns(described.columns()))
.unwrap_or_default(),
};
let rows = rows
.iter()
.map(|row| (0..columns.len()).map(|index| cell(row, index)).collect())
.collect();
Ok(QueryResult {
columns,
column_types,
rows,
elapsed,
affected,
})
}
async fn fetch_binary_column<DB>(
pool: &sqlx::Pool<DB>,
sql: &str,
params: Vec<Cell>,
raw_bytes: fn(&DB::Row, usize) -> Option<Vec<u8>>,
) -> Result<Option<Vec<u8>>>
where
DB: Database,
for<'c> &'c sqlx::Pool<DB>: Executor<'c, Database = DB>,
<DB as Database>::Arguments: IntoArguments<DB>,
for<'q> Option<String>: Encode<'q, DB>,
String: Type<DB>,
{
let statement = AssertSqlSafe(sql.to_string()).into_sql_str();
let mut query = sqlx::query(statement);
for param in params {
query = query.bind(param);
}
let row = query.fetch_optional(pool).await?;
Ok(row.and_then(|row| raw_bytes(&row, 0)))
}
pub(crate) fn query_outcome(result: &QueryResult) -> QueryOutcome {
if result.rows.is_empty() && result.affected.is_some() {
QueryOutcome::Affected(result.affected.unwrap_or_default())
} else {
QueryOutcome::Rows(result.row_count())
}
}
fn describe_columns<C: Column>(columns: &[C]) -> (Vec<String>, Vec<String>) {
columns
.iter()
.map(|column| {
(
column.name().to_string(),
column.type_info().name().to_string(),
)
})
.unzip()
}
async fn execute_with<DB>(
pool: &sqlx::Pool<DB>,
sql: &str,
params: Vec<Cell>,
rows_affected: fn(&DB::QueryResult) -> u64,
) -> Result<u64>
where
DB: Database,
for<'c> &'c sqlx::Pool<DB>: Executor<'c, Database = DB>,
<DB as Database>::Arguments: IntoArguments<DB>,
for<'q> Option<String>: Encode<'q, DB>,
String: Type<DB>,
{
let statement = AssertSqlSafe(sql.to_string()).into_sql_str();
let mut query = sqlx::query(statement);
for param in params {
query = query.bind(param);
}
let result = query.execute(pool).await?;
Ok(rows_affected(&result))
}
async fn execute_script_transactional<DB>(
pool: &sqlx::Pool<DB>,
statements: &[String],
) -> Result<()>
where
DB: Database,
for<'c> &'c mut DB::Connection: Executor<'c, Database = DB>,
<DB as Database>::Arguments: IntoArguments<DB>,
{
let mut tx = pool
.begin()
.await
.context("could not start the change's transaction")?;
for (position, statement) in statements.iter().enumerate() {
let sql = AssertSqlSafe(statement.clone()).into_sql_str();
if let Err(error) = sqlx::query(sql).execute(&mut *tx).await {
tx.rollback().await.ok();
return Err(anyhow::Error::from(error))
.with_context(|| format!("statement {} of {}", position + 1, statements.len()));
}
}
tx.commit().await.context("could not commit the change")?;
Ok(())
}