use toasty_core::{
driver::{ExecResponse, Rows, operation},
stmt,
};
use crate::{
Result,
engine::{eval, exec::Exec, mir},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ConditionalOutput {
None,
Count,
Returning,
}
#[derive(Debug, Clone)]
pub(crate) struct PaginationConfig {
pub page_size: i64,
pub has_previous_page: bool,
pub extract_cursor: Option<eval::Func>,
}
#[derive(Debug)]
pub(super) struct MySQLUpdateReturning {
select_stmt: stmt::Statement,
}
impl Exec<'_> {
pub(super) async fn exec_statement(
&mut self,
action: &mir::ExecStatement,
) -> Result<ExecResponse> {
let output_ty = mir::row_field_types(&action.ty);
let mut stmt = action.stmt.clone();
if !action.inputs.is_empty() {
let input_values = self.collect_input(action.inputs.iter().copied()).await?;
stmt.substitute(&input_values);
self.engine.simplify_stmt(&mut stmt);
}
debug_assert!(
stmt.returning()
.and_then(|returning| returning.as_project())
.map(|expr| expr.is_record())
.unwrap_or(true),
"stmt={stmt:#?}"
);
let mysql_update_returning = self.process_stmt_update_with_returning_on_mysql(&mut stmt);
if let stmt::Statement::Query(query) = &stmt
&& let stmt::ExprSet::Values(values) = &query.body
&& values.is_empty()
{
assert_eq!(action.conditional, ConditionalOutput::None);
let rows = if output_ty.is_some() {
Rows::Stream(stmt::ValueStream::default())
} else {
Rows::Count(0)
};
return Ok(ExecResponse::from_rows(rows));
}
let params = self.engine.prepare_for_driver(&mut stmt);
let ret = match action.conditional {
ConditionalOutput::Count => Some(vec![stmt::Type::I64, stmt::Type::I64]),
ConditionalOutput::Returning => {
let mut tys = vec![stmt::Type::I64, stmt::Type::I64];
tys.extend(
output_ty
.clone()
.expect("conditional write with RETURNING has output columns"),
);
Some(tys)
}
ConditionalOutput::None if mysql_update_returning.is_some() => {
None
}
ConditionalOutput::None => output_ty.clone(),
};
let op: toasty_core::driver::Operation = if stmt.is_insert() {
operation::Insert { stmt, params, ret }.into()
} else {
operation::QuerySql { stmt, params, ret }.into()
};
let mut res = self.connection.exec(&self.engine.schema, op).await?;
match action.conditional {
ConditionalOutput::None => {
if let Some(mysql_update) = mysql_update_returning {
res = self
.run_mysql_update_returning_select(mysql_update, output_ty.clone())
.await?;
}
}
ConditionalOutput::Count | ConditionalOutput::Returning => {
let rows = collect_conditional_probe(res.values).await?;
let (matched, conditioned) = conditional_probe_counts(&rows[0])?;
if matched == 0 {
return Err(toasty_core::Error::record_not_found(
"conditional write matched no rows",
));
}
if matched != conditioned {
return Err(toasty_core::Error::condition_failed(
"write condition did not match",
));
}
res.values = match action.conditional {
ConditionalOutput::Count => Rows::Count(matched as u64),
_ => {
let changed = rows
.into_iter()
.map(|row| {
let stmt::Value::Record(record) = row else {
return Err(toasty_core::Error::invalid_result(
"conditional write expected Record",
));
};
Ok(stmt::Value::record_from_vec(
record.fields.into_iter().skip(2).collect(),
))
})
.collect::<Result<Vec<_>>>()?;
Rows::value_stream(changed)
}
};
}
}
if let Some(pagination) = &action.pagination {
assert!(res.is_unpaginated());
res.values.buffer().await?;
self.apply_sql_pagination(&mut res, pagination)?;
}
Ok(res)
}
fn apply_sql_pagination(
&mut self,
res: &mut ExecResponse,
pagination: &PaginationConfig,
) -> Result<()> {
let Some(extract_cursor) = &pagination.extract_cursor else {
return Ok(());
};
let Rows::Value(stmt::Value::List(ref row_vec)) = res.values else {
return Ok(());
};
let page_size = pagination.page_size as usize;
res.next_cursor = if row_vec.len() == page_size {
let cursor_row = &row_vec[page_size - 1];
Some(Box::new(extract_cursor.eval(
&self.engine.schema,
std::slice::from_ref(cursor_row),
)?))
} else {
None
};
res.prev_cursor = if pagination.has_previous_page
&& !row_vec.is_empty()
&& self.engine.capability().backward_pagination
{
let cursor_row = &row_vec[0];
Some(Box::new(extract_cursor.eval(
&self.engine.schema,
std::slice::from_ref(cursor_row),
)?))
} else {
None
};
Ok(())
}
}
impl Exec<'_> {
pub(super) fn process_stmt_update_with_returning_on_mysql(
&self,
stmt: &mut stmt::Statement,
) -> Option<MySQLUpdateReturning> {
if self.engine.capability().returning_from_update || !self.engine.capability().sql() {
return None;
}
let stmt::Statement::Update(update) = stmt else {
return None;
};
let table_id = match &update.target {
stmt::UpdateTarget::Table(table_id) => *table_id,
_ => return None,
};
let returning = update.returning.take()?;
let select = stmt::Select {
returning,
source: stmt::Source::table(table_id),
filter: update.filter.clone(),
distinct: false,
};
let select_stmt =
stmt::Statement::Query(stmt::Query::new(stmt::ExprSet::Select(Box::new(select))));
Some(MySQLUpdateReturning { select_stmt })
}
pub(super) async fn run_mysql_update_returning_select(
&mut self,
mysql_update: MySQLUpdateReturning,
ret_ty: Option<Vec<stmt::Type>>,
) -> Result<toasty_core::driver::ExecResponse> {
let mut select_stmt = mysql_update.select_stmt;
let select_params = self.engine.prepare_for_driver(&mut select_stmt);
let op = operation::QuerySql {
stmt: select_stmt,
params: select_params,
ret: ret_ty,
};
self.connection.exec(&self.engine.schema, op.into()).await
}
}
async fn collect_conditional_probe(rows: Rows) -> Result<Vec<stmt::Value>> {
let Rows::Stream(rows) = rows else {
return Err(toasty_core::Error::invalid_result(format!(
"conditional write expected Stream, got {rows:?}"
)));
};
let rows = rows.collect().await?;
if rows.is_empty() {
return Err(toasty_core::Error::invalid_result(
"conditional write probe returned no rows",
));
}
Ok(rows)
}
fn conditional_probe_counts(row: &stmt::Value) -> Result<(i64, i64)> {
let stmt::Value::Record(record) = row else {
return Err(toasty_core::Error::invalid_result(format!(
"conditional write expected Record, got {row:?}"
)));
};
match (record.fields.first(), record.fields.get(1)) {
(Some(stmt::Value::I64(matched)), Some(stmt::Value::I64(conditioned))) => {
Ok((*matched, *conditioned))
}
_ => Err(toasty_core::Error::invalid_result(format!(
"conditional write probe columns are not I64; row={row:?}"
))),
}
}