use crate::db::plan::{self, InsertPlan, PkSource};
use crate::db::platform::Platform;
use crate::db::reference::TableRef;
use crate::db::{bind_all, text_col, value};
use crate::world::World;
use sqlx::AnyPool;
use sqlx::any::AnyArguments;
use sqlx::{Any, query::Query};
use std::future::Future;
use std::sync::Arc;
use std::time::Instant;
pub fn log_sql(w: &World, sql: &str, binds: &[Option<String>], logs: &[String]) {
if w.debug {
eprintln!("SQL: {sql}");
eprintln!("PARAMETERS: {binds:?}");
for l in logs {
eprintln!(" auto: {l}");
}
}
}
async fn timed<T, E>(w: &World, fut: impl Future<Output = Result<T, E>>) -> Result<T, E> {
let start = Instant::now();
let result = fut.await;
if w.debug {
eprintln!("TIME: {:.2} ms", start.elapsed().as_secs_f64() * 1000.0);
}
result
}
pub async fn resolve<'a>(
w: &'a World,
raw_table: &str,
) -> Result<
(
&'a AnyPool,
&'static dyn Platform,
Arc<plan::TableSchema>,
TableRef,
),
String,
> {
let tref = TableRef::parse(raw_table)?;
let conn = tref
.conn
.clone()
.unwrap_or_else(|| w.db.current().to_string());
let db = w.db.resources()?;
let state = db.connection(&conn)?;
let schema = db.schema(&conn, &tref).await?;
Ok((&state.pool, state.platform, schema, tref))
}
pub async fn insert(
w: &mut World,
raw_table: &str,
values: &[(String, Option<String>)],
index: Option<usize>,
) -> Result<(), String> {
let (pool, platform, schema, tref) = resolve(w, raw_table).await?;
let InsertPlan {
sql,
binds,
logs,
has_returning,
pk_vars,
} = plan::build_insert(
platform,
&schema,
&tref.sql_name(),
&tref.table,
values,
index,
)?;
log_sql(w, &sql, &binds, &logs);
if pk_vars.is_empty() {
timed(w, bind_all(sqlx::query(&sql), &binds).execute(pool))
.await
.map_err(|e| format!("INSERT into {}: {e}", tref.sql_name()))?;
return Ok(());
}
let assignments: Vec<(String, String)> = if has_returning {
let row = timed(w, bind_all(sqlx::query(&sql), &binds).fetch_one(pool))
.await
.map_err(|e| format!("INSERT into {}: {e}", tref.sql_name()))?;
let mut assignments = Vec::new();
for (i, (name, _)) in pk_vars.iter().enumerate() {
let v = text_col(&row, i)
.map_err(|e| format!("reading RETURNING: {e}"))?
.ok_or_else(|| "reading RETURNING: unexpected NULL".to_string())?;
assignments.push((name.clone(), v));
}
assignments
} else {
let result = timed(w, bind_all(sqlx::query(&sql), &binds).execute(pool))
.await
.map_err(|e| format!("INSERT into {}: {e}", tref.sql_name()))?;
let mut assignments = Vec::new();
for (name, source) in &pk_vars {
let v = match source {
PkSource::Known(v) => v.clone(),
PkSource::AutoIncrement => result
.last_insert_id()
.filter(|id| *id > 0)
.ok_or_else(|| {
format!(
"INSERT into {}: expected an auto-increment id in the INSERT result, got none",
tref.sql_name()
)
})?
.to_string(),
PkSource::Unknown(col) => unreachable!(
"build_insert refuses PkSource::Unknown ({col}) before returning a plan when has_returning is false"
),
};
assignments.push((name.clone(), v));
}
assignments
};
for (name, v) in assignments {
w.vars.set(&name, v);
}
Ok(())
}
pub async fn update(w: &mut World, raw_table: &str, set: &str, where_: &str) -> Result<(), String> {
let (pool, platform, schema, tref) = resolve(w, raw_table).await?;
let set_pairs = value::parse_oneliner(set)?;
let where_pairs = value::parse_oneliner(where_)?;
let (sql, binds) = plan::build_update(
platform,
&schema,
&tref.sql_name(),
&set_pairs,
&where_pairs,
)?;
log_sql(w, &sql, &binds, &[]);
let done = timed(w, bind_all(sqlx::query(&sql), &binds).execute(pool))
.await
.map_err(|e| format!("UPDATE {}: {e}", tref.sql_name()))?;
let table = tref.table.clone();
w.vars.set(
&format!("updated_{table}"),
done.rows_affected().to_string(),
);
Ok(())
}
pub async fn delete(w: &mut World, raw_table: &str, where_: &str) -> Result<(), String> {
let (pool, platform, schema, tref) = resolve(w, raw_table).await?;
let where_pairs = value::parse_oneliner(where_)?;
let (sql, binds) = plan::build_delete(platform, &schema, &tref.sql_name(), &where_pairs)?;
log_sql(w, &sql, &binds, &[]);
let done = timed(w, bind_all(sqlx::query(&sql), &binds).execute(pool))
.await
.map_err(|e| format!("DELETE {}: {e}", tref.sql_name()))?;
let table = tref.table.clone();
w.vars.set(
&format!("deleted_{table}"),
done.rows_affected().to_string(),
);
Ok(())
}
pub async fn exists(
w: &mut World,
raw_table: &str,
where_pairs: &[(String, Option<String>)],
) -> Result<bool, String> {
let (pool, platform, schema, tref) = resolve(w, raw_table).await?;
let (sql, binds) = plan::build_exists(platform, &schema, &tref.sql_name(), where_pairs)?;
log_sql(w, &sql, &binds, &[]);
Ok(
timed(w, bind_all(sqlx::query(&sql), &binds).fetch_optional(pool))
.await
.map_err(|e| format!("checking presence in {}: {e}", tref.sql_name()))?
.is_some(),
)
}
pub async fn extract(
w: &mut World,
column: &str,
raw_table: &str,
where_str: &str,
var: &str,
) -> Result<(), String> {
let (pool, platform, schema, tref) = resolve(w, raw_table).await?;
plan::plain_column(column)?;
let col = schema
.col(column)
.ok_or_else(|| format!("column {column:?} is missing from {}", tref.sql_name()))?;
platform.check_bindable(col)?;
let where_pairs = value::parse_oneliner(where_str)?;
let (where_sql, binds) = plan::build_where(platform, &schema, &where_pairs, 1)?;
let sql = format!(
"SELECT {} FROM {} WHERE {where_sql} LIMIT 1",
platform.cast_text(column),
tref.sql_name()
);
log_sql(w, &sql, &binds, &[]);
let row = timed(w, bind_all(sqlx::query(&sql), &binds).fetch_optional(pool))
.await
.map_err(|e| format!("SELECT {}: {e}", tref.sql_name()))?
.ok_or_else(|| format!("no row in {} matched the condition", tref.sql_name()))?;
let v = text_col(&row, 0).map_err(|e| format!("reading value: {e}"))?;
let value = v.unwrap_or_default();
w.vars.set(var, value);
Ok(())
}
pub async fn delete_all(w: &mut World, raw_table: &str) -> Result<(), String> {
let (pool, _platform, _schema, tref) = resolve(w, raw_table).await?;
let sql = plan::build_delete_all(&tref.sql_name());
log_sql(w, &sql, &[], &[]);
let done = timed(w, sqlx::query(&sql).execute(pool))
.await
.map_err(|e| format!("DELETE ALL {}: {e}", tref.sql_name()))?;
let table = tref.table.clone();
w.vars.set(
&format!("deleted_{table}"),
done.rows_affected().to_string(),
);
Ok(())
}
fn bind_args<'q>(
mut q: Query<'q, Any, AnyArguments<'q>>,
args: &'q [value::Arg],
) -> Query<'q, Any, AnyArguments<'q>> {
for a in args {
q = match a {
value::Arg::Null => q.bind(Option::<String>::None),
value::Arg::Int(i) => q.bind(*i),
value::Arg::Float(f) => q.bind(*f),
value::Arg::Bool(b) => q.bind(*b),
value::Arg::Text(t) => q.bind(t.as_str()),
};
}
q
}
fn placeholder_list(p: &dyn Platform, n: usize) -> String {
(1..=n)
.map(|i| p.placeholder(i))
.collect::<Vec<_>>()
.join(", ")
}
pub async fn call_procedure(w: &mut World, name: &str, args_str: &str) -> Result<(), String> {
let args = value::parse_args(args_str)?;
let db = w.db.resources()?;
let state = db.connection(w.db.current())?;
let sql = format!(
"CALL {name}({})",
placeholder_list(state.platform, args.len())
);
if w.debug {
eprintln!("SQL: {sql}\nARGUMENTS: {args:?}");
}
timed(w, bind_args(sqlx::query(&sql), &args).execute(&state.pool))
.await
.map_err(|e| format!("CALL {name}: {e}"))?;
Ok(())
}
pub async fn call_function(
w: &mut World,
name: &str,
args_str: &str,
var: &str,
) -> Result<(), String> {
let args = value::parse_args(args_str)?;
let db = w.db.resources()?;
let state = db.connection(w.db.current())?;
let call = format!("{name}({})", placeholder_list(state.platform, args.len()));
let sql = state.platform.cast_text(&call);
let sql = format!("SELECT {sql}");
if w.debug {
eprintln!("SQL: {sql}\nARGUMENTS: {args:?}");
}
let row = timed(
w,
bind_args(sqlx::query(&sql), &args).fetch_one(&state.pool),
)
.await
.map_err(|e| format!("SELECT {name}(...): {e}"))?;
let v = text_col(&row, 0).map_err(|e| format!("reading function result: {e}"))?;
let value = v.unwrap_or_default();
w.vars.set(var, value);
Ok(())
}
pub async fn next_sequence(w: &mut World, seq: &str, var: &str) -> Result<(), String> {
let db = w.db.resources()?;
let state = db.connection(w.db.current())?;
let (sql, binds) = state
.platform
.next_sequence(seq)
.ok_or_else(|| format!("sequences are not supported on {}", state.platform.name()))?;
if w.debug {
eprintln!("SQL: {sql} [{seq}]");
}
let row = timed(
w,
bind_all(sqlx::query(&sql), &binds).fetch_one(&state.pool),
)
.await
.map_err(|e| format!("next value of sequence {seq:?}: {e}"))?;
let v = text_col(&row, 0)
.map_err(|e| format!("reading sequence: {e}"))?
.unwrap_or_default();
w.vars.set(var, v);
Ok(())
}