toasty 0.11.0

An async ORM for Rust supporting SQL and NoSQL databases
Documentation
use toasty_core::{
    driver::{ExecResponse, Rows, operation},
    stmt,
};

use crate::{
    Result,
    engine::{eval, exec::Exec, mir},
};

/// How to interpret a statement's output rows.
///
/// A conditional write (the SQL `#[version]` / OCC path compiled as a single
/// CTE statement) prefixes its result with two probe columns: the number of
/// rows matching the filter and, of those, the number satisfying the condition.
/// The write applied only when the two agree; a mismatch is a condition
/// failure, and zero matched rows means the record no longer exists.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ConditionalOutput {
    /// Not a conditional write. Output is passed through unchanged.
    None,

    /// Conditional write with no `RETURNING`. The two probe columns are the
    /// only output; the action reports the matched-row count.
    Count,

    /// Conditional write with a `RETURNING`. The two probe columns are followed
    /// by the changed rows' columns, which become the action's output.
    Returning,
}

/// Configuration for pagination at the execution level.
#[derive(Debug, Clone)]
pub(crate) struct PaginationConfig {
    /// Number of items per page
    pub page_size: i64,
    /// Whether the query starts after a cursor and can have a previous page.
    pub has_previous_page: bool,
    /// Function to extract cursor from a row (SQL only).
    /// For NoSQL drivers, this is None (driver provides cursor).
    pub extract_cursor: Option<eval::Func>,
}

/// Information about a MySQL UPDATE with RETURNING that needs special handling.
///
/// MySQL doesn't support `RETURNING` on `UPDATE`. The workaround is to strip
/// the returning, run the UPDATE, then run a follow-up `SELECT` over the same
/// table and filter to fetch the post-update column values. The two
/// statements are not atomic relative to concurrent writers — see #881 for
/// the broader design discussion.
#[derive(Debug)]
pub(super) struct MySQLUpdateReturning {
    /// The `SELECT` statement that returns the post-update values. Carries
    /// the same filter as the original `UPDATE` plus the projected
    /// returning expression.
    select_stmt: stmt::Statement,
}

impl Exec<'_> {
    pub(super) async fn exec_statement(
        &mut self,
        action: &mir::ExecStatement,
    ) -> Result<ExecResponse> {
        // Databases always return rows as a vec of values; this specifies the
        // type of each value. `None` means the statement returns only a count.
        let output_ty = mir::row_field_types(&action.ty);

        let mut stmt = action.stmt.clone();

        // Collect input values and substitute into the statement
        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:#?}"
        );

        // MySQL does not support `RETURNING` on `UPDATE`. Strip the returning
        // and capture an equivalent `SELECT` to run after the UPDATE.
        let mysql_update_returning = self.process_stmt_update_with_returning_on_mysql(&mut stmt);

        // Short circuit if we can statically determine there are no results
        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));
        }

        // Legalize the statement for the target backend and extract bind
        // parameters (SQL drivers only; key-value drivers read values
        // directly from the statement).
        let params = self.engine.prepare_for_driver(&mut stmt);

        let ret = match action.conditional {
            // A conditional write prefixes its result with two `I64` probe
            // counts; the `Returning` variant follows them with the changed
            // rows' columns.
            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() => {
                // The UPDATE has had its RETURNING stripped; the driver runs
                // a plain UPDATE that returns no rows. The follow-up SELECT
                // below produces the returning values.
                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])?;

                // A conditional write targets a row the caller holds an
                // instance of: zero matched rows means it has since been
                // deleted.
                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),
                    _ => {
                        // The probe locked the matched rows, so the write
                        // applied to exactly those rows and every result row is
                        // a real changed row — strip the two leading probe
                        // columns.
                        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)
                    }
                };
            }
        }

        // Apply pagination if configured
        if let Some(pagination) = &action.pagination {
            assert!(res.is_unpaginated());
            res.values.buffer().await?;
            self.apply_sql_pagination(&mut res, pagination)?;
        }

        Ok(res)
    }

    /// Apply SQL pagination by extracting cursor from last row.
    /// If we got a full page (page_size rows), extract cursor for potential next page.
    /// The client will naturally discover there's no more data when the next request returns empty.
    ///
    /// The response values must already be buffered (via `Rows::buffer()`).
    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;

        // Extract cursors for potential next/prev pages
        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 {
            // Got fewer than page_size rows, no more data
            None
        };

        // Extract a previous cursor only when this query started after another page.
        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<'_> {
    /// Detects an UPDATE with a non-empty `RETURNING` on a MySQL backend
    /// and rewrites the statement for the workaround path:
    ///
    /// - The returning clause is stripped from the UPDATE so the SQL
    ///   serializer doesn't reject it.
    /// - An equivalent `SELECT` over the same table + filter is captured,
    ///   carrying the original returning expression as its projection.
    ///
    /// Returns `None` when the backend supports `RETURNING` natively (PG,
    /// SQLite) or when the statement is not an UPDATE with a returning
    /// project. The two-statement path is not atomic relative to concurrent
    /// writers — see #881.
    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 })
    }

    /// Runs the follow-up `SELECT` for a MySQL UPDATE with stripped
    /// `RETURNING`. The driver receives a plain query whose result rows
    /// take the place of the original RETURNING output.
    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
    }
}

/// Collects a conditional write's result rows. The probe (a `COUNT` aggregate)
/// always yields at least one row, so an empty result is a driver bug.
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)
}

/// Reads the two leading probe counts (`matched`, `conditioned`) from a
/// conditional write's result row.
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:?}"
        ))),
    }
}