saddle-db 0.1.0

Saddle managed asynchronous database access and transactions
Documentation
use std::{future::Future, pin::Pin};

use futures_util::TryStreamExt;
use saddle_observability::{ActiveCall, CallKind, Observer};
use sqlx::MySql;
use tokio::sync::mpsc;

use crate::{
    CallContext, DbRow, MAX_QUERY_ROWS, MAX_RESULT_BYTES, Result, Statement, WriteResult,
    cleanup::RollbackJob,
    error::{
        map_operation_error, result_limit_exceeded, transaction_commit_failed,
        transaction_rollback_failed,
    },
    row::row_payload_bytes,
};

pub type TransactionFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T>> + Send + 'a>>;

/// A single-level managed transaction.
///
/// There is intentionally no API for beginning another transaction, creating
/// savepoints, committing, or rolling back. The surrounding `Database` owns
/// that boundary.
pub struct Transaction {
    inner: Option<sqlx::Transaction<'static, MySql>>,
    observer: Observer,
    context: CallContext,
    boundary: Option<ActiveCall>,
    cleanup: Option<mpsc::UnboundedSender<RollbackJob>>,
}

impl Transaction {
    pub(crate) fn new(
        inner: sqlx::Transaction<'static, MySql>,
        observer: Observer,
        boundary: ActiveCall,
        cleanup: mpsc::UnboundedSender<RollbackJob>,
    ) -> Self {
        let context = boundary.context().clone();
        Self {
            inner: Some(inner),
            observer,
            context,
            boundary: Some(boundary),
            cleanup: Some(cleanup),
        }
    }

    fn inner(&mut self) -> &mut sqlx::Transaction<'static, MySql> {
        self.inner
            .as_mut()
            .expect("active transaction has an inner connection")
    }

    pub async fn query_all(&mut self, statement: Statement) -> Result<Vec<DbRow>> {
        statement.validate()?;
        let operation = statement.operation().to_owned();
        let call = self.observer.start_child_call(
            &self.context,
            CallKind::Database,
            "database",
            "database",
            operation,
        );
        let result = async {
            let mut stream = statement.query().fetch(&mut **self.inner());
            let mut rows = Vec::new();
            let mut result_bytes = 0_usize;
            while let Some(row) = stream.try_next().await.map_err(map_operation_error)? {
                if rows.len() == MAX_QUERY_ROWS {
                    return Err(result_limit_exceeded());
                }
                result_bytes = result_bytes.saturating_add(row_payload_bytes(&row)?);
                if result_bytes > MAX_RESULT_BYTES {
                    return Err(result_limit_exceeded());
                }
                rows.push(DbRow(row));
            }
            Ok(rows)
        }
        .await;
        finish_call(call, &result);
        result
    }

    pub async fn query_optional(&mut self, statement: Statement) -> Result<Option<DbRow>> {
        statement.validate()?;
        let operation = statement.operation().to_owned();
        let call = self.observer.start_child_call(
            &self.context,
            CallKind::Database,
            "database",
            "database",
            operation,
        );
        let result = statement
            .query()
            .fetch_optional(&mut **self.inner())
            .await
            .map_err(map_operation_error)
            .and_then(|row| {
                row.map(|row| {
                    row_payload_bytes(&row)?;
                    Ok(DbRow(row))
                })
                .transpose()
            });
        finish_call(call, &result);
        result
    }

    pub async fn write(&mut self, statement: Statement) -> Result<WriteResult> {
        statement.validate()?;
        let operation = statement.operation().to_owned();
        let call = self.observer.start_child_call(
            &self.context,
            CallKind::Database,
            "database",
            "database",
            operation,
        );
        let result = statement
            .query()
            .execute(&mut **self.inner())
            .await
            .map(|result| WriteResult::new(result.rows_affected(), result.last_insert_id()))
            .map_err(map_operation_error);
        finish_call(call, &result);
        result
    }

    pub(crate) async fn commit(mut self) -> Result<()> {
        let inner = self
            .inner
            .take()
            .expect("active transaction has an inner connection");
        let boundary = self
            .boundary
            .take()
            .expect("active transaction has a boundary");
        let _cleanup = self.cleanup.take();
        let commit = self.observer.start_child_call(
            boundary.context(),
            CallKind::Transaction,
            "database",
            "database",
            "commit",
        );
        match inner.commit().await {
            Ok(()) => {
                commit.succeed();
                boundary.succeed();
                Ok(())
            }
            Err(_) => {
                let error = transaction_commit_failed();
                commit.fail(&error);
                boundary.fail(&error);
                Err(error)
            }
        }
    }

    pub(crate) async fn rollback(mut self, work_error: crate::SaddleError) -> crate::SaddleError {
        let inner = self
            .inner
            .take()
            .expect("active transaction has an inner connection");
        let boundary = self
            .boundary
            .take()
            .expect("active transaction has a boundary");
        let _cleanup = self.cleanup.take();
        let rollback = self.observer.start_child_call(
            boundary.context(),
            CallKind::Transaction,
            "database",
            "database",
            "rollback",
        );
        match inner.rollback().await {
            Ok(()) => {
                rollback.succeed();
                boundary.fail(&work_error);
                work_error
            }
            Err(_) => {
                let error = transaction_rollback_failed();
                rollback.fail(&error);
                boundary.fail(&error);
                error
            }
        }
    }
}

impl Drop for Transaction {
    fn drop(&mut self) {
        let (Some(inner), Some(boundary), Some(cleanup)) =
            (self.inner.take(), self.boundary.take(), self.cleanup.take())
        else {
            return;
        };
        let rollback = self.observer.start_child_call(
            boundary.context(),
            CallKind::Transaction,
            "database",
            "database",
            "rollback",
        );
        let job = RollbackJob {
            inner,
            boundary,
            rollback,
        };
        if let Err(error) = cleanup.send(job) {
            error.0.fail_without_worker();
        }
    }
}

fn finish_call<T>(call: saddle_observability::ActiveCall, result: &Result<T>) {
    match result {
        Ok(_) => call.succeed(),
        Err(error) => call.fail(error),
    }
}