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();
}
}
}
#[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;
}
}
struct TransactionTask<O: StaticWriteOperation> {
request: ManagedDatabaseRequestCapability,
context: CallContext,
event: EventContext,
state: Arc<RequestState>,
cancel_state: Arc<RequestState>,
proof: TransactionOperationProof<O>,
parameters: O::Parameters,
decision: TransactionDecision,
sender: tokio::sync::oneshot::Sender<Result<ManagedWriteResult, TransactionExecutionError>>,
}
impl<O: StaticWriteOperation> TransactionTask<O> {
async fn run(self) {
let Self {
request,
context,
event,
state,
cancel_state,
proof,
parameters,
decision,
sender,
} = self;
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();
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 transaction_memory_layout<O: StaticWriteOperation, C>() -> [std::alloc::Layout; 4] {
[
future_layout(TransactionTask::<O>::run),
std::alloc::Layout::new::<Result<ManagedWriteResult, TransactionExecutionError>>(),
std::alloc::Layout::new::<RequestState>(),
std::alloc::Layout::new::<(Result<ManagedWriteResult, TransactionExecutionError>, C)>(),
]
}
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(
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())
}
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(
TransactionTask {
request,
context,
event,
state,
cancel_state,
proof,
parameters,
decision,
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::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",
)
}