use std::sync::{Arc, Mutex};
use saddle_core::CallContext;
use saddle_db::internal::{
ManagedDatabaseRequestCapability, ManagedOptionalRow, ManagedWriteResult,
QueryOptionalExecutionError, QueryOptionalOperationProof, StaticQueryOptionalOperation,
StaticWriteOperation, 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};
#[doc(hidden)]
pub trait GeneratedNamedQuery: StaticQueryOptionalOperation {
type NamedParameters: Send;
type NamedRow: Send;
const REQUIRES_TRANSACTION: bool;
fn parameters(value: Self::NamedParameters) -> Self::Parameters;
fn row(value: ManagedOptionalRow<Self::Row>) -> Option<Self::NamedRow>;
}
#[doc(hidden)]
pub trait GeneratedNamedWrite: StaticWriteOperation {
type NamedParameters: Send;
fn parameters(value: Self::NamedParameters) -> Self::Parameters;
}
enum RequestTerminal {
Untouched(ProfuseGwManagedDispatch),
InFlight,
PostDatabase(ProfuseGwPostDatabaseRequestTerminal),
Serial(SerialScope),
Finished,
}
struct RequestState {
terminal: Mutex<RequestTerminal>,
cancel: Arc<Notify>,
completed: Notify,
}
struct AbortOnLostRequestOwner(bool);
type SerialScope =
saddle_runtime::profusegw::ProfuseGwSerialScope<tokio::sync::futures::OwnedNotified>;
#[doc(hidden)]
pub type TransactionSession<'session, 'request> = saddle_db::internal::DatabaseTransactionSession<
'session,
'request,
tokio::sync::futures::OwnedNotified,
>;
impl Drop for AbortOnLostRequestOwner {
fn drop(&mut self) {
if self.0 {
std::process::abort();
}
}
}
#[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>,
}
struct QueryTask<O: StaticQueryOptionalOperation> {
request: ManagedDatabaseRequestCapability,
context: CallContext,
event: EventContext,
state: Arc<RequestState>,
cancel_state: Arc<RequestState>,
proof: QueryOptionalOperationProof<O>,
parameters: O::Parameters,
sender: tokio::sync::oneshot::Sender<
Result<ManagedOptionalRow<O::Row>, QueryOptionalExecutionError>,
>,
}
impl<O: StaticQueryOptionalOperation> QueryTask<O> {
async fn run(self) {
let Self {
request,
context,
event,
state,
cancel_state,
proof,
parameters,
sender,
} = self;
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();
DatabaseRequest::publish(&state, terminal, value, sender);
owner_guard.0 = false;
}
}
struct WriteTask<O: StaticWriteOperation> {
request: ManagedDatabaseRequestCapability,
context: CallContext,
event: EventContext,
state: Arc<RequestState>,
cancel_state: Arc<RequestState>,
proof: WriteOperationProof<O>,
parameters: O::Parameters,
sender: tokio::sync::oneshot::Sender<Result<ManagedWriteResult, WriteExecutionError>>,
}
impl<O: StaticWriteOperation> WriteTask<O> {
async fn run(self) {
let Self {
request,
context,
event,
state,
cancel_state,
proof,
parameters,
sender,
} = self;
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();
DatabaseRequest::publish(&state, terminal, value, sender);
owner_guard.0 = false;
}
}
fn future_layout<A, F: std::future::Future>(_: fn(A) -> F) -> std::alloc::Layout {
std::alloc::Layout::new::<F>()
}
pub(crate) fn query_memory_layout<O: StaticQueryOptionalOperation, C>() -> [std::alloc::Layout; 4] {
[
future_layout(QueryTask::<O>::run),
std::alloc::Layout::new::<Result<ManagedOptionalRow<O::Row>, QueryOptionalExecutionError>>(
),
std::alloc::Layout::new::<RequestState>(),
std::alloc::Layout::new::<(
Result<ManagedOptionalRow<O::Row>, QueryOptionalExecutionError>,
C,
)>(),
]
}
pub(crate) fn write_memory_layout<O: StaticWriteOperation, C>() -> [std::alloc::Layout; 4] {
[
future_layout(WriteTask::<O>::run),
std::alloc::Layout::new::<Result<ManagedWriteResult, WriteExecutionError>>(),
std::alloc::Layout::new::<RequestState>(),
std::alloc::Layout::new::<(Result<ManagedWriteResult, WriteExecutionError>, C)>(),
]
}
pub(crate) fn scoped_memory_layout<C>() -> [std::alloc::Layout; 4] {
[
std::alloc::Layout::new::<SerialScope>(),
std::alloc::Layout::new::<RequestState>(),
std::alloc::Layout::new::<DatabaseRequest>(),
std::alloc::Layout::new::<C>(),
]
}
impl DatabaseRequest {
pub(crate) fn supervise_external<T: Send + 'static>(
&self,
future: std::pin::Pin<
Box<
dyn std::future::Future<Output = crate::programming::ExternalFunctionResult<T>>
+ Send,
>,
>,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = crate::programming::ExternalFunctionResult<T>> + Send>,
> {
let state = Arc::clone(&self.state);
Box::pin(async move {
use crate::programming::{
ExecutionCertainty, ExternalFunctionResult, TechnicalFailure, TechnicalFailureCode,
};
let stopped = |code, certainty| {
ExternalFunctionResult::TechnicalFailure(TechnicalFailure::from_framework(
code, certainty,
))
};
let scope = {
let mut terminal = state.terminal.lock().unwrap();
match std::mem::replace(&mut *terminal, RequestTerminal::InFlight) {
RequestTerminal::Serial(scope) => Some(scope),
other => {
let untouched = matches!(other, RequestTerminal::Untouched(_));
*terminal = other;
if !untouched {
return stopped(
TechnicalFailureCode::InternalFailure,
ExecutionCertainty::NotExecuted,
);
}
None
}
}
};
let Some(mut scope) = scope else {
return future.await;
};
let (sender, receiver) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let mut guard = AbortOnLostRequestOwner(true);
let result = scope.supervise_between(future).await;
*state.terminal.lock().unwrap() = RequestTerminal::Serial(scope);
state.completed.notify_one();
let _ = sender.send(result);
guard.0 = false;
});
match receiver.await {
Ok(Ok(result)) => result,
Ok(Err(failure)) => {
use saddle_runtime::profusegw::{ProfuseGwScopeFailure, ProfuseGwScopeStop};
let code = match failure {
ProfuseGwScopeFailure::Stopped(ProfuseGwScopeStop::TimedOut) => {
TechnicalFailureCode::DeadlineExceeded
}
ProfuseGwScopeFailure::Stopped(ProfuseGwScopeStop::Cancelled) => {
TechnicalFailureCode::DependencyUnavailable
}
ProfuseGwScopeFailure::Panicked
| ProfuseGwScopeFailure::AlreadySupervised
| ProfuseGwScopeFailure::ScopeAlreadyEntered => {
TechnicalFailureCode::InternalFailure
}
};
stopped(code, ExecutionCertainty::MayHaveExecuted)
}
Err(_) => std::process::abort(),
}
})
}
pub async fn scoped_transaction<T, E, B>(
&mut self,
body: B,
) -> saddle_db::internal::ScopeTransactionOutcome<T, E>
where
T: Send + 'static,
E: Send + 'static,
B: for<'tx, 'session, 'request> FnOnce(
&'tx mut TransactionSession<'session, 'request>,
)
-> saddle_db::internal::ScopeTransactionFuture<
'tx,
T,
E,
> + Send
+ 'static,
{
use saddle_db::internal::{
ScopeDatabaseError, ScopeTransactionAbort, ScopeTransactionOutcome,
};
let rejected = || {
ScopeTransactionOutcome::Rejected(ScopeTransactionAbort::Technical(
ScopeDatabaseError::Unavailable,
))
};
let Some(process) = self.process.as_ref().cloned() else {
return rejected();
};
let scope = {
let mut state = self.state.terminal.lock().unwrap();
match std::mem::replace(&mut *state, RequestTerminal::InFlight) {
RequestTerminal::Untouched(dispatch) => dispatch
.into_database_request()
.into_serial_scope(Arc::clone(&self.state.cancel).notified_owned()),
RequestTerminal::Serial(scope) => scope,
other => {
*state = other;
return rejected();
}
}
};
let state = Arc::clone(&self.state);
let context = self.context.clone();
let event = self.event.clone();
let (sender, receiver) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let mut guard = AbortOnLostRequestOwner(true);
let mut suspended = process
.transaction_scope(scope, &context, event, body)
.await;
let (terminal, result) = match suspended.resume().await {
Ok((next, result)) => (RequestTerminal::Serial(next), result),
Err(_) => {
let (result, terminal) = suspended
.into_response_parts()
.unwrap_or_else(|_| std::process::abort());
(RequestTerminal::PostDatabase(terminal), result)
}
};
*state.terminal.lock().unwrap() = terminal;
state.completed.notify_one();
let _ = sender.send(result.outcome);
guard.0 = false;
});
receiver.await.unwrap_or_else(|_| std::process::abort())
}
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: Arc::new(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(
QueryTask {
request,
context,
event,
state,
cancel_state,
proof,
parameters,
sender,
}
.run(),
);
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(
WriteTask {
request,
context,
event,
state,
cancel_state,
proof,
parameters,
sender,
}
.run(),
);
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::Serial(scope)) => {
if scope.finish_unentered_response(()).is_err() {
std::process::abort();
}
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",
)
}