saddle-framework 0.3.9

The single business-facing facade for Saddle applications
use std::sync::{Arc, Mutex};

use saddle_core::CallContext;
use saddle_db::internal::{
    ManagedDatabaseRequestCapability, ManagedOptionalRow, ManagedWriteResult,
    QueryOptionalExecutionError, QueryOptionalOperationProof, StaticQueryOptionalOperation,
    StaticWriteOperation, TransactionDecision, TransactionExecutionError,
    TransactionOperationProof, WriteExecutionError, WriteOperationProof,
};
use saddle_observability::EventContext;
use saddle_runtime::profusegw::{
    ProfuseGwManagedDispatch, ProfuseGwPostDatabaseRequestTerminal,
    finish_profusegw_after_database, finish_profusegw_without_database,
};
use tokio::sync::Notify;

use crate::{ErrorKind, SaddleError};

enum RequestTerminal {
    Untouched(ProfuseGwManagedDispatch),
    InFlight,
    PostDatabase(ProfuseGwPostDatabaseRequestTerminal),
    Finished,
}

struct RequestState {
    terminal: Mutex<RequestTerminal>,
    cancel: Notify,
    completed: Notify,
}

struct AbortOnLostRequestOwner(bool);

impl Drop for AbortOnLostRequestOwner {
    fn drop(&mut self) {
        if self.0 {
            std::process::abort();
        }
    }
}

/// Framework-owned request database carrier. Generated capabilities retain
/// only a private reference to this state; business code cannot construct or
/// recover Runtime or Database owners from it.
#[doc(hidden)]
pub struct DatabaseRequest {
    process: Option<Arc<saddle_db::internal::StartupManagedDatabaseProcessCapability>>,
    context: CallContext,
    event: EventContext,
    state: Arc<RequestState>,
}

#[doc(hidden)]
pub struct DatabaseRequestCompletion {
    state: Arc<RequestState>,
}

impl DatabaseRequest {
    pub(crate) fn new(
        process: Option<Arc<saddle_db::internal::StartupManagedDatabaseProcessCapability>>,
        dispatch: ProfuseGwManagedDispatch,
        context: CallContext,
        event: EventContext,
    ) -> (Self, DatabaseRequestCompletion) {
        let state = Arc::new(RequestState {
            terminal: Mutex::new(RequestTerminal::Untouched(dispatch)),
            cancel: Notify::new(),
            completed: Notify::new(),
        });
        (
            Self {
                process,
                context,
                event,
                state: Arc::clone(&state),
            },
            DatabaseRequestCompletion { state },
        )
    }

    fn begin(&self) -> Result<ManagedDatabaseRequestCapability, SaddleError> {
        let process = self.process.as_ref().ok_or_else(database_unavailable)?;
        let dispatch = {
            let mut terminal = self.state.terminal.lock().unwrap();
            match std::mem::replace(&mut *terminal, RequestTerminal::InFlight) {
                RequestTerminal::Untouched(dispatch) => dispatch,
                other => {
                    *terminal = other;
                    return Err(database_already_used());
                }
            }
        };
        Ok(process.bind_profusegw_request(dispatch.into_database_request()))
    }

    fn publish<T>(
        state: &RequestState,
        terminal: ProfuseGwPostDatabaseRequestTerminal,
        value: T,
        sender: tokio::sync::oneshot::Sender<T>,
    ) {
        *state.terminal.lock().unwrap() = RequestTerminal::PostDatabase(terminal);
        state.completed.notify_waiters();
        let _ = sender.send(value);
    }

    pub async fn query_optional<O>(
        self,
        parameters: O::Parameters,
    ) -> Result<ManagedOptionalRow<O::Row>, QueryOptionalExecutionError>
    where
        O: StaticQueryOptionalOperation,
    {
        let proof = QueryOptionalOperationProof::<O>::bind()
            .map_err(|_| QueryOptionalExecutionError::MappingUnavailable)?;
        let request = self
            .begin()
            .map_err(|_| QueryOptionalExecutionError::Shutdown)?;
        let context = self.context;
        let event = self.event;
        let state = Arc::clone(&self.state);
        let cancel_state = Arc::clone(&state);
        let (sender, receiver) = tokio::sync::oneshot::channel();
        tokio::spawn(async move {
            let mut owner_guard = AbortOnLostRequestOwner(true);
            let terminal = request
                .query_optional(&context, event, proof.invocation(parameters), async move {
                    cancel_state.cancel.notified().await;
                })
                .await;
            let (_, owner) = terminal.into_outcome();
            let (value, terminal) = owner.into_response_parts();
            Self::publish(&state, terminal, value, sender);
            owner_guard.0 = false;
        });
        receiver.await.unwrap_or_else(|_| std::process::abort())
    }

    pub async fn write<O>(
        self,
        parameters: O::Parameters,
    ) -> Result<ManagedWriteResult, WriteExecutionError>
    where
        O: StaticWriteOperation,
    {
        let proof = WriteOperationProof::<O>::bind()
            .map_err(|_| WriteExecutionError::MappingUnavailable)?;
        let request = self.begin().map_err(|_| WriteExecutionError::Shutdown)?;
        let context = self.context;
        let event = self.event;
        let state = Arc::clone(&self.state);
        let cancel_state = Arc::clone(&state);
        let (sender, receiver) = tokio::sync::oneshot::channel();
        tokio::spawn(async move {
            let mut owner_guard = AbortOnLostRequestOwner(true);
            let terminal = request
                .write(&context, event, proof.invocation(parameters), async move {
                    cancel_state.cancel.notified().await;
                })
                .await;
            let (_, owner) = terminal.into_outcome();
            let (value, terminal) = owner.into_response_parts();
            Self::publish(&state, terminal, value, sender);
            owner_guard.0 = false;
        });
        receiver.await.unwrap_or_else(|_| std::process::abort())
    }

    pub async fn transaction<O>(
        self,
        parameters: O::Parameters,
        decision: TransactionDecision,
    ) -> Result<ManagedWriteResult, TransactionExecutionError>
    where
        O: StaticWriteOperation,
    {
        let proof = TransactionOperationProof::<O>::bind()
            .map_err(|_| TransactionExecutionError::MappingUnavailable)?;
        let request = self
            .begin()
            .map_err(|_| TransactionExecutionError::Shutdown)?;
        let context = self.context;
        let event = self.event;
        let state = Arc::clone(&self.state);
        let cancel_state = Arc::clone(&state);
        let (sender, receiver) = tokio::sync::oneshot::channel();
        tokio::spawn(async move {
            let mut owner_guard = AbortOnLostRequestOwner(true);
            let terminal = request
                .transaction(
                    &context,
                    event,
                    proof.invocation(parameters, decision),
                    async move { cancel_state.cancel.notified().await },
                )
                .await;
            let (_, owner) = terminal.into_outcome();
            let (value, terminal) = owner.into_response_parts();
            Self::publish(&state, terminal, value, sender);
            owner_guard.0 = false;
        });
        receiver.await.unwrap_or_else(|_| std::process::abort())
    }
}

impl DatabaseRequestCompletion {
    pub(crate) async fn finish(self, cancelled: bool) {
        if cancelled {
            self.state.cancel.notify_one();
        }
        loop {
            let terminal = {
                let mut state = self.state.terminal.lock().unwrap();
                match std::mem::replace(&mut *state, RequestTerminal::Finished) {
                    RequestTerminal::InFlight => {
                        *state = RequestTerminal::InFlight;
                        None
                    }
                    terminal => Some(terminal),
                }
            };
            match terminal {
                Some(RequestTerminal::Untouched(dispatch)) => {
                    if cancelled {
                        dispatch.cancel();
                    } else if finish_profusegw_without_database(dispatch, ()).is_err() {
                        std::process::abort();
                    }
                    return;
                }
                Some(RequestTerminal::PostDatabase(terminal)) => {
                    finish_profusegw_after_database(terminal);
                    return;
                }
                Some(RequestTerminal::Finished) => return,
                Some(RequestTerminal::InFlight) | None => self.state.completed.notified().await,
            }
        }
    }
}

fn database_unavailable() -> SaddleError {
    SaddleError::new(
        ErrorKind::Unavailable,
        "saddle.database.not_configured",
        "database is not configured for this process",
    )
}

fn database_already_used() -> SaddleError {
    SaddleError::new(
        ErrorKind::Conflict,
        "saddle.database.operation_already_used",
        "the route database operation was already consumed",
    )
}