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, Transaction,
Type, TypeInfo,
};
use super::catalog::{Catalog, CatalogEntry, CatalogKind, MAX_ENTRIES};
use super::config::{ConnectionConfig, Engine};
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::{mysql, postgres, sqlite, statement, typed_placeholder};
pub(crate) const POOL_SIZE: u32 = 5;
#[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,
}
impl Connection {
pub async fn open(config: ConnectionConfig, password: Option<String>) -> Result<Self> {
let password = password.as_deref();
let pool = match config.engine {
Engine::Postgres => Pool::Postgres(postgres::connect(&config, password).await?),
Engine::MySql => Pool::MySql(mysql::connect(&config, password).await?),
Engine::Sqlite => Pool::Sqlite(sqlite::connect(&config).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(),
})
}
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();
Self::open(config, self.password.clone()).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| {
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 })
})
.collect())
}
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| {
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,
})
})
.collect())
}
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 })
}
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
}
pub async fn run_query_as_user(&self, sql: &str) -> Result<QueryResult> {
self.run_query_tagged(sql, Vec::new(), QuerySource::User)
.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
}
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> {
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
}
}
}
pub async fn fetch_binary(&self, sql: &str, params: Vec<Cell>) -> Result<Option<Vec<u8>>> {
self.refuse_write(sql)?;
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,
}
}
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).len() != 1 {
anyhow::bail!("Select one statement to explain.");
}
let header = statement::explained(statement);
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) {
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 async fn run_script(&self, sql: &str) -> Result<Vec<QueryResult>> {
let statements = statement::split(sql);
for statement in &statements {
self.refuse_write(&statement.text)?;
}
let started = Instant::now();
let result = match &self.pool {
Pool::Postgres(pool) => {
let scale = self.money_scale;
run_script_transactional(
pool,
&statements,
move |row, index| postgres::cell(row, index, scale),
postgres::rows_affected,
)
.await
}
Pool::Sqlite(pool) => {
run_script_transactional(pool, &statements, sqlite::cell, sqlite::rows_affected)
.await
}
Pool::MySql(pool) => {
run_script_untransacted(pool, &statements, mysql::cell, mysql::rows_affected).await
}
};
match &result {
Ok(results) => {
for (statement, result) in statements.iter().zip(results) {
self.log(
&statement.text,
QuerySource::User,
result.elapsed,
query_outcome(result),
);
}
}
Err(error) => self.log(
sql,
QuerySource::User,
started.elapsed(),
QueryOutcome::Error(format!("{error:#}")),
),
}
result
}
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 = match &self.pool {
Pool::Postgres(pool) => import::Session::postgres(pool).await?,
Pool::MySql(pool) => import::Session::mysql(pool).await?,
Pool::Sqlite(pool) => import::Session::sqlite(pool).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
}
fn refuse_write(&self, sql: &str) -> Result<()> {
if !self.config.safety.is_read_only() {
return Ok(());
}
if let Some(word) = statement::first_write(sql) {
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,
};
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;
self.log_script(statements, started.elapsed(), &result);
result
}
Pool::Sqlite(pool) => {
let started = Instant::now();
let result = execute_script_transactional(pool, statements).await;
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,
}
}
}
async fn run_script_untransacted<DB, F>(
pool: &sqlx::Pool<DB>,
statements: &[statement::Statement],
cell: F,
rows_affected: fn(&DB::QueryResult) -> u64,
) -> Result<Vec<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 mut results = Vec::new();
for (position, statement) in statements.iter().enumerate() {
let result = fetch_all(pool, &statement.text, Vec::new(), &cell, rows_affected)
.await
.with_context(|| format!("statement {}", position + 1))?;
results.push(result);
}
Ok(results)
}
async fn run_script_transactional<DB, F>(
pool: &sqlx::Pool<DB>,
statements: &[statement::Statement],
cell: F,
rows_affected: fn(&DB::QueryResult) -> u64,
) -> Result<Vec<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 mut tx = pool
.begin()
.await
.context("could not start the script's transaction")?;
let mut results = Vec::new();
for (position, statement) in statements.iter().enumerate() {
match fetch_in_transaction(&mut tx, &statement.text, &cell, rows_affected).await {
Ok(result) => results.push(result),
Err(error) => {
tx.rollback().await.ok();
return Err(error.context(format!("statement {}", position + 1)));
}
}
}
tx.commit().await.context("could not commit the script")?;
Ok(results)
}
async fn fetch_in_transaction<DB, F>(
tx: &mut Transaction<'_, DB>,
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 **tx);
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 **tx, 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)))
}
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(())
}