use std::{
future::Future,
pin::{Pin, pin},
task::{Context, Poll},
};
use saddle_admission::{AdmissionError, DbFinalizerRoleClaim, DbRequestPermit, ManagedResponse};
#[doc(hidden)]
pub struct DbFinalizingOutput<T, F> {
value: T,
return_to_pool: F,
claim: DbFinalizerRoleClaim,
}
impl<F> DbFinalizingOutput<ManagedResponse, F>
where
F: Future<Output = ()> + Send + 'static,
{
#[doc(hidden)]
pub fn begin(
response: ManagedResponse,
permit: DbRequestPermit,
return_to_pool: F,
) -> Result<Self, AdmissionError> {
let claim = permit.begin_finalizing()?;
Ok(Self {
value: response,
return_to_pool,
claim,
})
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[doc(hidden)]
pub enum DbTransitionRequest {
Cancel,
Shutdown,
}
#[doc(hidden)]
pub enum DbQueryPoll<T> {
Ready(T),
Transition(DbTransitionRequest),
}
#[must_use = "the DB transition owner must be consumed into Finalizing"]
#[doc(hidden)]
pub struct DbQueryTransition<C, S> {
cancel: C,
shutdown: S,
selected: Option<DbTransitionRequest>,
active: bool,
}
impl<C, S> DbQueryTransition<C, S>
where
C: Future + Unpin,
S: Future + Unpin,
{
#[doc(hidden)]
pub fn poll_query<Q>(
&mut self,
permit: &DbRequestPermit,
query: Pin<&mut Q>,
context: &mut Context<'_>,
) -> Poll<DbQueryPoll<Q::Output>>
where
Q: Future,
{
if let Poll::Ready(output) = permit.poll_query(query, context) {
return Poll::Ready(DbQueryPoll::Ready(output));
}
self.poll_cancel_or_shutdown(context)
.map(DbQueryPoll::Transition)
}
fn poll_cancel_or_shutdown(&mut self, context: &mut Context<'_>) -> Poll<DbTransitionRequest> {
if let Some(selected) = self.selected {
return Poll::Ready(selected);
}
if Pin::new(&mut self.shutdown).poll(context).is_ready() {
self.selected = Some(DbTransitionRequest::Shutdown);
return Poll::Ready(DbTransitionRequest::Shutdown);
}
if Pin::new(&mut self.cancel).poll(context).is_ready() {
self.selected = Some(DbTransitionRequest::Cancel);
return Poll::Ready(DbTransitionRequest::Cancel);
}
Poll::Pending
}
#[doc(hidden)]
pub fn begin_finalizing<T, F>(
mut self,
value: T,
permit: DbRequestPermit,
return_to_pool: F,
) -> Result<DbFinalizingOutput<T, F>, AdmissionError>
where
T: Send + 'static,
F: Future<Output = ()> + Send + 'static,
{
let claim = permit.begin_finalizing()?;
self.active = false;
Ok(DbFinalizingOutput {
value,
return_to_pool,
claim,
})
}
}
impl<C, S> Drop for DbQueryTransition<C, S> {
fn drop(&mut self) {
if self.active {
std::process::abort();
}
}
}
struct FinalizationGuard {
complete: bool,
}
impl Drop for FinalizationGuard {
fn drop(&mut self) {
if !self.complete {
std::process::abort();
}
}
}
#[doc(hidden)]
pub async fn drive_db_finalizer<T, Q, F>(query: Q) -> Result<T, AdmissionError>
where
T: Send + 'static,
Q: Future<Output = Result<DbFinalizingOutput<T, F>, AdmissionError>> + Send + 'static,
F: Future<Output = ()> + Send + 'static,
{
let mut guard = FinalizationGuard { complete: false };
let output = query.await?;
let mut return_to_pool = pin!(output.return_to_pool);
std::future::poll_fn(|context| {
output
.claim
.poll_connection_return(return_to_pool.as_mut(), context)
})
.await;
output.claim.complete_after_connection_return()?;
guard.complete = true;
Ok(output.value)
}
#[doc(hidden)]
pub async fn drive_db_finalizer_with_transition<T, C, S, B, Q, F>(
cancel: C,
shutdown: S,
factory: B,
) -> Result<T, AdmissionError>
where
T: Send + 'static,
C: Future + Unpin + Send + 'static,
S: Future + Unpin + Send + 'static,
B: FnOnce(DbQueryTransition<C, S>) -> Q,
Q: Future<Output = Result<DbFinalizingOutput<T, F>, AdmissionError>> + Send + 'static,
F: Future<Output = ()> + Send + 'static,
{
let transition = DbQueryTransition {
cancel,
shutdown,
selected: None,
active: true,
};
drive_db_finalizer(factory(transition)).await
}
#[doc(hidden)]
pub async fn drive_db_query<Q>(permit: &DbRequestPermit, query: Q) -> Q::Output
where
Q: Future,
{
let mut query = pin!(query);
std::future::poll_fn(|context| permit.poll_query(query.as_mut(), context)).await
}
#[cfg(test)]
mod tests {
use std::{
env,
future::{Future, Pending, Ready},
os::unix::process::ExitStatusExt,
pin::pin,
process::{Command, Stdio},
task::{Context, Poll, Waker},
time::{Duration, Instant},
};
use super::*;
async fn pending_query()
-> Result<DbFinalizingOutput<ManagedResponse, Ready<()>>, AdmissionError> {
std::future::pending().await
}
async fn pending_transition_query(
_transition: DbQueryTransition<Pending<()>, Pending<()>>,
) -> Result<DbFinalizingOutput<ManagedResponse, Ready<()>>, AdmissionError> {
std::future::pending().await
}
#[test]
fn lost_db_role_is_finite_fail_closed_in_subprocess() {
const CHILD: &str = "db_finalizer::tests::lost_db_role_child";
for mode in [
"direct",
"during-unwind",
"transition-direct",
"transition-during-unwind",
] {
let mut child = Command::new(env::current_exe().unwrap())
.args(["--exact", CHILD, "--nocapture"])
.env("SADDLE_DB_FINALIZER_LOST_ROLE_CHILD", mode)
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.unwrap();
let started = Instant::now();
let status = loop {
if let Some(status) = child.try_wait().unwrap() {
break status;
}
if started.elapsed() >= Duration::from_secs(5) {
child.kill().unwrap();
child.wait().unwrap();
panic!("lost DB role ({mode}) exceeded finite deadline");
}
std::thread::sleep(Duration::from_millis(10));
};
assert_eq!(status.signal(), Some(6));
}
}
#[test]
fn lost_db_role_child() {
let Some(mode) = env::var_os("SADDLE_DB_FINALIZER_LOST_ROLE_CHILD") else {
return;
};
let transition = mode.to_string_lossy().starts_with("transition-");
let mut future = pin!(async {
if transition {
drive_db_finalizer_with_transition(
std::future::pending(),
std::future::pending(),
pending_transition_query,
)
.await
} else {
drive_db_finalizer(pending_query()).await
}
});
let waker = Waker::noop();
let mut context = Context::from_waker(waker);
assert!(matches!(future.as_mut().poll(&mut context), Poll::Pending));
if mode.to_string_lossy().ends_with("during-unwind") {
let _future = future;
panic!("existing unwind must not discard the DB finalizer role");
}
}
}