sqlx-turso-driver 0.0.1

An asynchronous SQLx driver for embedded Turso databases
use crate::{
    Turso, TursoColumn, TursoConnection, TursoQueryResult, TursoRow, TursoStatement, TursoTypeInfo,
    error::{engine_error, unsupported},
};
use either::Either;
use futures_core::{future::BoxFuture, stream::BoxStream};
use futures_util::{TryStreamExt, stream};
use sqlx_core::{
    error::Error,
    executor::{Execute, Executor},
    logger::QueryLogger,
    sql_str::SqlStr,
};

impl<'c> Executor<'c> for &'c mut TursoConnection {
    type Database = Turso;
    fn fetch_many<'e, 'q: 'e, E>(
        self,
        query: E,
    ) -> BoxStream<'e, Result<Either<TursoQueryResult, TursoRow>, Error>>
    where
        'c: 'e,
        E: 'q + Execute<'q, Turso>,
    {
        // The stream owns the mutable connection borrow, not a public clone.
        // Turso prepare ignores trailing statements, so reject batches before
        // preparing or executing anything, using its own parser and byte offset.
        Box::pin(stream::try_unfold(
            (self, Some(query), None),
            |(conn, query, active)| async move {
                let (statement, mut rows, names, mut logger) = match active {
                    Some(active) => active,
                    None => {
                        let Some(mut query) = query else {
                            return Ok(None);
                        };
                        conn.ready().await?;
                        let arguments = query
                            .take_arguments()
                            .map_err(Error::Encode)?
                            .unwrap_or_default();
                        let sql = query.sql();
                        let logger = QueryLogger::new(sql.clone(), conn.log_settings.clone());
                        validate_single_statement(
                            sql.as_str(),
                            conn.transaction_state == crate::transaction::TransactionState::Active,
                        )?;
                        let mut statement = conn
                            .inner
                            .prepare(sql.as_str())
                            .await
                            .map_err(engine_error)?;
                        let names = statement.column_names();
                        // Only an internal statement handle is cloned, not a
                        // physical connection for another executor/transaction.
                        // Keep it across every query/step await, including errors.
                        conn.pending_statement = Some(statement.clone());
                        let rows = statement
                            .query(arguments.values)
                            .await
                            .map_err(engine_error)?;
                        (statement, rows, names, logger)
                    }
                };
                if let Some(row) = rows.next().await.map_err(engine_error)? {
                    let row = TursoRow::from_engine(row, &names)?;
                    logger.increment_rows_returned();
                    Ok(Some((
                        Either::Right(row),
                        (conn, None, Some((statement, rows, names, logger))),
                    )))
                } else {
                    // Pinned core Statement::n_change is per-statement, not the
                    // connection's changes(). Read it only after completion, for
                    // writes with RETURNING as well as non-row-returning writes.
                    let rows_affected = statement.n_change();
                    logger.increase_rows_affected(rows_affected);
                    drop(rows);
                    drop(statement);
                    conn.clear_pending_statement()?;
                    Ok(Some((
                        Either::Left(TursoQueryResult { rows_affected }),
                        (conn, None, None),
                    )))
                }
            },
        ))
    }
    fn fetch_optional<'e, 'q: 'e, E>(
        self,
        query: E,
    ) -> BoxFuture<'e, Result<Option<TursoRow>, Error>>
    where
        'c: 'e,
        E: 'q + Execute<'q, Turso>,
    {
        Box::pin(
            self.fetch_many(query)
                .try_fold(None, |first, result| async move {
                    // Drain the engine to completion even after the first row: stopping
                    // early can otherwise roll back a row-returning write in Turso.
                    Ok(match (first, result) {
                        (None, Either::Right(row)) => Some(row),
                        (first, _) => first,
                    })
                }),
        )
    }
    fn prepare_with<'e>(
        self,
        sql: SqlStr,
        parameters: &'e [TursoTypeInfo],
    ) -> BoxFuture<'e, Result<TursoStatement, Error>>
    where
        'c: 'e,
    {
        Box::pin(async move {
            self.ready().await?;
            if !parameters.is_empty() {
                return Err(unsupported("parameter type hints"));
            }
            validate_single_statement(
                sql.as_str(),
                self.transaction_state == crate::transaction::TransactionState::Active,
            )?;
            let prepared = self
                .inner
                .prepare(sql.as_str())
                .await
                .map_err(engine_error)?;
            let columns = prepared
                .columns()
                .into_iter()
                .enumerate()
                .map(|(ordinal, column)| TursoColumn {
                    ordinal,
                    name: column.name().to_owned(),
                    type_info: TursoTypeInfo::from_decl_type(column.decl_type()),
                })
                .collect();
            Ok(TursoStatement { sql, columns })
        })
    }
}

fn validate_single_statement(sql: &str, owned_transaction: bool) -> Result<(), Error> {
    use turso_parser::ast::{Cmd, Stmt};
    let parse_error =
        |error: turso::core::LimboError| engine_error(turso::Error::Error(error.to_string()));
    let (command, end) = turso::core::dialect::sqlite::parse(sql).map_err(parse_error)?;
    // Match actual commands, including END/ROLLBACK TO aliases, not text in
    // comments/literals. Public statements also pass here when re-prepared.
    if owned_transaction
        && matches!(
            command,
            Some(Cmd::Stmt(
                Stmt::Begin { .. }
                    | Stmt::Commit { .. }
                    | Stmt::Rollback { .. }
                    | Stmt::Savepoint { .. }
                    | Stmt::Release { .. }
                    | Stmt::Attach { .. }
                    | Stmt::Detach { .. }
                    | Stmt::Pragma { .. }
                    | Stmt::Vacuum { .. }
            ))
        )
    {
        return Err(unsupported(
            "transaction-escaping SQL inside a SQLx transaction",
        ));
    }
    let (trailing, _) = turso::core::dialect::sqlite::parse(&sql[end..]).map_err(parse_error)?;
    if trailing.is_some() {
        return Err(unsupported("multiple SQL statements"));
    }
    Ok(())
}