use std::sync::{
Arc, OnceLock,
atomic::{AtomicUsize, Ordering},
};
use std::fmt;
use saddle_core::{ApplicationId, CallContext, CaptureSite, Diagnostic, DiagnosticCategory, DiagnosticCause,
DiagnosticCode, DiagnosticStage, ErrorKind, Result, SaddleError};
use saddle_observability::EmergencyDiagnosticHandle;
use saddle_observability::root_diagnostic::{OriginalCaptureState, UnrootedCaptureFacts,
runtime_admission_source_description, runtime_unrooted_admission_source_description};
use tokio::sync::Notify;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ApplicationPhase {
Starting = 0,
Ready = 1,
Draining = 2,
Stopped = 3,
}
#[derive(Clone, Debug)]
pub struct ApplicationHealth {
shared: Arc<Shared>,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct ApplicationHealthSnapshot {
phase: ApplicationPhase,
}
impl ApplicationHealthSnapshot {
pub const fn phase(self) -> ApplicationPhase {
self.phase
}
pub const fn is_live(self) -> bool {
!matches!(self.phase, ApplicationPhase::Stopped)
}
pub const fn is_ready(self) -> bool {
matches!(self.phase, ApplicationPhase::Ready)
}
}
impl ApplicationHealth {
pub(crate) fn fail_closed(&self) {
RequestLifecycle { shared: Arc::clone(&self.shared) }.begin_draining();
}
pub fn snapshot(&self) -> ApplicationHealthSnapshot {
ApplicationHealthSnapshot {
phase: phase(self.shared.state.load(Ordering::Acquire)),
}
}
}
const PHASE_SHIFT: u32 = usize::BITS - 2;
const COUNT_MASK: usize = (1 << PHASE_SHIFT) - 1;
#[derive(Debug)]
struct Shared {
state: AtomicUsize,
drained: Notify,
source_binding: OnceLock<AdmissionSourceBinding>,
}
struct AdmissionSourceBinding {
application: ApplicationId,
output: Option<EmergencyDiagnosticHandle>,
}
impl fmt::Debug for AdmissionSourceBinding {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.debug_struct("AdmissionSourceBinding")
.field("application", &self.application)
.field("output_selected", &self.output.is_some()).finish()
}
}
#[derive(Clone, Debug)]
pub struct RequestLifecycle {
shared: Arc<Shared>,
}
impl RequestLifecycle {
pub(crate) fn new() -> Self {
Self {
shared: Arc::new(Shared {
state: AtomicUsize::new(encode(ApplicationPhase::Starting, 0)),
drained: Notify::new(),
source_binding: OnceLock::new(),
}),
}
}
pub fn phase(&self) -> ApplicationPhase {
phase(self.shared.state.load(Ordering::Acquire))
}
pub(crate) fn health(&self) -> ApplicationHealth {
ApplicationHealth {
shared: Arc::clone(&self.shared),
}
}
pub(crate) fn install_admission_source(&self, application: ApplicationId,
output: Option<EmergencyDiagnosticHandle>) -> bool {
self.shared.source_binding.set(AdmissionSourceBinding { application, output }).is_ok()
}
pub fn try_accept(&self) -> Result<RequestGuard> {
self.try_claim().map(RequestClaim::publish)
}
#[doc(hidden)]
pub fn try_accept_recorded(&self, call: &CallContext,
output: Option<&EmergencyDiagnosticHandle>) -> Result<RequestGuard> {
if let Some(binding) = self.shared.source_binding.get()
&& &binding.application != call.application() {
return Err(admission_identity_mismatch(binding, call,
self.shared.state.load(Ordering::Acquire)));
}
self.try_claim_with_source(Some((call, output))).map(RequestClaim::publish)
}
pub(crate) fn try_claim(&self) -> Result<RequestClaim> {
self.try_claim_with_source(None)
}
fn try_claim_with_source(&self,
source: Option<(&CallContext, Option<&EmergencyDiagnosticHandle>)>) -> Result<RequestClaim> {
let mut current = self.shared.state.load(Ordering::Acquire);
loop {
match phase(current) {
ApplicationPhase::Ready => {
if count(current) == COUNT_MASK {
return Err(admission_rejection(source, self.shared.source_binding.get(), current, ErrorKind::Internal,
"runtime.request_count_overflow", "request accounting capacity exhausted"));
}
match self.shared.state.compare_exchange_weak(
current,
current + 1,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
return Ok(RequestClaim {
shared: Some(Arc::clone(&self.shared)),
});
}
Err(observed) => current = observed,
}
}
ApplicationPhase::Starting => {
return Err(admission_rejection(source, self.shared.source_binding.get(), current, ErrorKind::Unavailable,
"runtime.not_ready", "application is not ready"));
}
ApplicationPhase::Draining => {
return Err(admission_rejection(source, self.shared.source_binding.get(), current, ErrorKind::Unavailable,
"runtime.shutting_down", "application is shutting down"));
}
ApplicationPhase::Stopped => {
return Err(admission_rejection(source, self.shared.source_binding.get(), current, ErrorKind::Unavailable,
"runtime.stopped", "application is stopped"));
}
}
}
}
pub(crate) fn mark_ready(&self) {
let result = self.shared.state.compare_exchange(
encode(ApplicationPhase::Starting, 0),
encode(ApplicationPhase::Ready, 0),
Ordering::AcqRel,
Ordering::Acquire,
);
debug_assert!(result.is_ok());
}
pub(crate) fn begin_draining(&self) {
let mut current = self.shared.state.load(Ordering::Acquire);
while matches!(
phase(current),
ApplicationPhase::Starting | ApplicationPhase::Ready
) {
let draining = encode(ApplicationPhase::Draining, count(current));
match self.shared.state.compare_exchange_weak(
current,
draining,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
current = draining;
break;
}
Err(observed) => current = observed,
}
}
if count(current) == 0 {
self.shared.drained.notify_one();
}
}
pub(crate) async fn wait_until_drained(&self) {
loop {
if count(self.shared.state.load(Ordering::Acquire)) == 0 {
return;
}
self.shared.drained.notified().await;
}
}
pub(crate) fn mark_stopped(&self) {
let mut current = self.shared.state.load(Ordering::Acquire);
loop {
debug_assert_eq!(phase(current), ApplicationPhase::Draining);
match self.shared.state.compare_exchange_weak(
current,
encode(ApplicationPhase::Stopped, count(current)),
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return,
Err(observed) => current = observed,
}
}
}
}
#[derive(Debug)]
struct AdmissionRejection {
code: &'static str,
observed_phase: ApplicationPhase,
observed_active_requests: usize,
}
#[derive(Debug)]
struct AdmissionIdentityMismatch<'a> {
bound_application: &'a ApplicationId,
call_application: &'a ApplicationId,
observed_phase: ApplicationPhase,
observed_active_requests: usize,
}
impl fmt::Display for AdmissionIdentityMismatch<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(formatter, "admission application mismatch: bound={:?}, call={:?}, phase={:?}, active_requests={}",
self.bound_application, self.call_application,
self.observed_phase, self.observed_active_requests)
}
}
#[track_caller]
fn admission_identity_mismatch(binding: &AdmissionSourceBinding,
call: &CallContext, state: usize) -> SaddleError {
let original = AdmissionIdentityMismatch {
bound_application: &binding.application,
call_application: call.application(),
observed_phase: phase(state), observed_active_requests: count(state),
};
let code = "runtime.admission_identity_mismatch";
let diagnostic = Diagnostic::capture(
DiagnosticCategory::UnexpectedError, CaptureSite::Origin,
DiagnosticCause::new(DiagnosticStage::RequestAdmission,
DiagnosticCode::new(code).expect("fixed identity code")),
);
let facts = runtime_unrooted_admission_source_description(binding.output.as_ref(),
&diagnostic, &original, Some(&binding.application));
let safe = SaddleError::new(ErrorKind::Internal, code,
"the request admission application does not match its bound lifecycle")
.with_diagnostic(diagnostic);
attach_admission_source(safe, facts)
}
impl fmt::Display for AdmissionRejection {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(formatter, "request admission rejected: code={}, phase={:?}, active_requests={}",
self.code, self.observed_phase, self.observed_active_requests)
}
}
#[track_caller]
fn admission_rejection(
source: Option<(&CallContext, Option<&EmergencyDiagnosticHandle>)>,
binding: Option<&AdmissionSourceBinding>,
state: usize,
kind: ErrorKind,
code: &'static str,
message: &'static str,
) -> SaddleError {
let original = AdmissionRejection {
code, observed_phase: phase(state), observed_active_requests: count(state),
};
let diagnostic = Diagnostic::capture(
DiagnosticCategory::UnexpectedError, CaptureSite::Origin,
DiagnosticCause::new(DiagnosticStage::RequestAdmission,
DiagnosticCode::new(code).expect("fixed admission code")),
);
let facts = match source {
Some((call, output)) => runtime_admission_source_description(output, &diagnostic, &original, call),
None => runtime_unrooted_admission_source_description(
binding.and_then(|binding| binding.output.as_ref()), &diagnostic, &original,
binding.map(|binding| &binding.application)),
};
let safe = SaddleError::new(kind, code, message).with_diagnostic(diagnostic);
attach_admission_source(safe, facts)
}
fn attach_admission_source(safe: SaddleError, facts: UnrootedCaptureFacts) -> SaddleError {
if facts.original_capture() == OriginalCaptureState::CompleteWritten
&& safe.diagnostic().is_some_and(|diagnostic|
facts.occurrence().matches_diagnostic(diagnostic)) {
safe.with_source_receipt(facts)
} else {
safe.with_unconfirmed_source()
}
}
const fn encode(phase: ApplicationPhase, count: usize) -> usize {
((phase as usize) << PHASE_SHIFT) | count
}
const fn phase(state: usize) -> ApplicationPhase {
match state >> PHASE_SHIFT {
0 => ApplicationPhase::Starting,
1 => ApplicationPhase::Ready,
2 => ApplicationPhase::Draining,
3 => ApplicationPhase::Stopped,
_ => unreachable!(),
}
}
const fn count(state: usize) -> usize {
state & COUNT_MASK
}
#[derive(Debug)]
pub(crate) struct RequestClaim {
shared: Option<Arc<Shared>>,
}
impl RequestClaim {
pub(crate) fn publish(mut self) -> RequestGuard {
RequestGuard {
shared: self.shared.take().expect("request claim publishes once"),
}
}
}
impl Drop for RequestClaim {
fn drop(&mut self) {
if let Some(shared) = self.shared.take() {
complete(&shared);
}
}
}
#[derive(Debug)]
pub struct RequestGuard {
shared: Arc<Shared>,
}
impl Drop for RequestGuard {
fn drop(&mut self) {
complete(&self.shared);
}
}
fn complete(shared: &Shared) {
let previous = shared.state.fetch_sub(1, Ordering::AcqRel);
debug_assert!(count(previous) > 0);
if count(previous) == 1 && phase(previous) == ApplicationPhase::Draining {
shared.drained.notify_one();
}
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Barrier};
use super::*;
fn test_runtime() -> tokio::runtime::Runtime {
tokio::runtime::Builder::new_current_thread()
.build()
.expect("test runtime must build")
}
#[test]
fn recorded_admission_rejections_keep_atomic_state_and_call_identity() {
use saddle_observability::{EmergencyDiagnostics, FileLoggingConfig, Observer,
ObserverConfig, Rotation};
use saddle_observability::root_diagnostic::UnrootedCaptureFacts;
let directory = std::env::temp_dir().join(format!("saddle-runtime-admission-{}-{}",
std::process::id(), std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH).unwrap().as_nanos()));
std::fs::create_dir(&directory).unwrap();
let mut writer = EmergencyDiagnostics::start(
&FileLoggingConfig::new(directory.clone(), Rotation::Daily)).unwrap();
let handle = writer.handle();
let observer = Observer::with_writer(ObserverConfig::default(), std::io::sink()).unwrap();
let (call, _) = observer.start_external_call_checked(
"admission-app", "runtime", "admission", "try_accept", None).unwrap();
let lifecycle = RequestLifecycle::new();
let starting = lifecycle.try_accept_recorded(call.context(), Some(&handle)).unwrap_err();
assert_eq!(starting.code(), "runtime.not_ready");
assert!(!starting.source_unavailable());
let receipt = starting.source_receipt::<UnrootedCaptureFacts>().unwrap();
assert_eq!(receipt.original_capture(), OriginalCaptureState::CompleteWritten);
assert!(receipt.occurrence().matches_diagnostic(starting.diagnostic().unwrap()));
lifecycle.mark_ready();
let guard = lifecycle.try_accept_recorded(call.context(), Some(&handle)).unwrap();
lifecycle.begin_draining();
let draining = lifecycle.try_accept_recorded(call.context(), Some(&handle)).unwrap_err();
assert_eq!(draining.code(), "runtime.shutting_down");
assert_ne!(starting.diagnostic().unwrap().id(), draining.diagnostic().unwrap().id());
let rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
.lines().map(|row| serde_json::from_str(row).unwrap()).collect();
for (error, expected) in [(&starting, "phase=Starting, active_requests=0"),
(&draining, "phase=Draining, active_requests=1")] {
let source: Vec<_> = rows.iter().filter(|row| row["occurrence"]["diagnostic_id"]
== error.diagnostic().unwrap().id()).collect();
let header: serde_json::Value = serde_json::from_str(source[0]["payload"].as_str().unwrap()).unwrap();
assert_eq!(header["context"]["application"], "admission-app");
assert_eq!(header["stage"], "runtime_request_admission");
assert!(source.iter().any(|row| row["channel"] == "description"
&& row["payload"].as_str().is_some_and(|text| text.contains(expected))));
assert!(source.iter().any(|row| row["channel"] == "terminal"
&& row["state"] == "description_complete_source_unavailable"));
}
let unconfirmed = lifecycle.try_accept_recorded(call.context(), None).unwrap_err();
assert!(unconfirmed.source_unavailable());
assert!(unconfirmed.source_receipt::<UnrootedCaptureFacts>().is_none());
let unrooted = RequestLifecycle::new();
assert!(unrooted.install_admission_source("unrooted-app".into(), Some(handle.clone())));
let no_call = unrooted.try_claim().unwrap_err();
assert_eq!(no_call.code(), "runtime.not_ready");
assert!(!no_call.source_unavailable());
let no_call_receipt = no_call.source_receipt::<UnrootedCaptureFacts>().unwrap();
assert_eq!(no_call_receipt.original_capture(), OriginalCaptureState::CompleteWritten);
assert!(no_call_receipt.occurrence().matches_diagnostic(no_call.diagnostic().unwrap()));
let unrooted_rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
.lines().map(|row| serde_json::from_str(row).unwrap()).collect();
let unrooted_source: Vec<_> = unrooted_rows.iter().filter(|row|
row["occurrence"]["diagnostic_id"] == no_call.diagnostic().unwrap().id()).collect();
let unrooted_header: serde_json::Value = serde_json::from_str(
unrooted_source[0]["payload"].as_str().unwrap()).unwrap();
assert_eq!(unrooted_header["stage"], "runtime_unrooted_admission");
assert_eq!(unrooted_header["context"]["application"]["value"], "unrooted-app");
assert_eq!(unrooted_header["context"]["trace_id"]["state"], "not_established");
assert!(unrooted_source.iter().any(|row| row["channel"] == "description"
&& row["payload"].as_str().is_some_and(|text|
text.contains("phase=Starting, active_requests=0"))));
assert!(unrooted_source.iter().any(|row| row["channel"] == "terminal"
&& row["state"] == "description_complete_source_unavailable"));
let mut application = crate::Application::new();
application.install_lifecycle_observer_with_output(
observer.clone(), "assembled-app", Some(handle.clone()));
let assembled = application.request_lifecycle().try_claim().unwrap_err();
assert!(!assembled.source_unavailable());
let assembled_rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
.lines().map(|row| serde_json::from_str(row).unwrap()).collect();
let assembled_source: Vec<_> = assembled_rows.iter().filter(|row|
row["occurrence"]["diagnostic_id"] == assembled.diagnostic().unwrap().id()).collect();
let assembled_header: serde_json::Value = serde_json::from_str(
assembled_source[0]["payload"].as_str().unwrap()).unwrap();
assert_eq!(assembled_header["context"]["application"]["value"], "assembled-app");
assert_eq!(assembled_header["stage"], "runtime_unrooted_admission");
let (wrong_call, _) = observer.start_external_call_checked(
"wrong-app", "runtime", "admission", "try_accept", None).unwrap();
let mismatch = application.request_lifecycle().try_accept_recorded(
wrong_call.context(), Some(&handle)).unwrap_err();
assert_eq!(mismatch.code(), "runtime.admission_identity_mismatch");
assert!(!mismatch.source_unavailable());
let mismatch_rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
.lines().map(|row| serde_json::from_str(row).unwrap()).collect();
let mismatch_source: Vec<_> = mismatch_rows.iter().filter(|row|
row["occurrence"]["diagnostic_id"] == mismatch.diagnostic().unwrap().id()).collect();
let mismatch_header: serde_json::Value = serde_json::from_str(
mismatch_source[0]["payload"].as_str().unwrap()).unwrap();
assert_eq!(mismatch_header["context"]["application"]["value"], "assembled-app");
assert!(mismatch_source.iter().any(|row| row["channel"] == "description"
&& row["payload"].as_str().is_some_and(|text|
text.contains("assembled-app") && text.contains("wrong-app"))));
let unrelated = Diagnostic::capture(DiagnosticCategory::UnexpectedError,
CaptureSite::Origin, DiagnosticCause::new(DiagnosticStage::RequestAdmission,
DiagnosticCode::new("runtime.receipt_mismatch_test").unwrap()));
let facts = runtime_unrooted_admission_source_description(Some(&handle), &unrelated,
&AdmissionRejection { code: "runtime.receipt_mismatch_test",
observed_phase: ApplicationPhase::Starting, observed_active_requests: 0 },
Some(&ApplicationId::from("assembled-app")));
assert_eq!(facts.original_capture(), OriginalCaptureState::CompleteWritten);
let safe = SaddleError::new(ErrorKind::Internal,
"runtime.receipt_mismatch_test", "safe").with_diagnostic(
Diagnostic::capture(DiagnosticCategory::UnexpectedError, CaptureSite::Origin,
DiagnosticCause::new(DiagnosticStage::RequestAdmission,
DiagnosticCode::new("runtime.receipt_mismatch_test").unwrap())));
let false_receipt = attach_admission_source(safe, facts);
assert!(false_receipt.source_unavailable());
assert!(false_receipt.source_receipt::<UnrootedCaptureFacts>().is_none());
let no_binding = RequestLifecycle::new().try_claim().unwrap_err();
assert!(no_binding.source_unavailable());
drop(guard);
lifecycle.mark_stopped();
let until = std::time::Instant::now() + std::time::Duration::from_secs(2);
loop {
match writer.shutdown() {
saddle_observability::DiagnosticShutdown::Finished => break,
saddle_observability::DiagnosticShutdown::Pending
if std::time::Instant::now() < until =>
std::thread::sleep(std::time::Duration::from_millis(10)),
other => panic!("admission diagnostic shutdown failed: {other:?}"),
}
}
drop(writer);
std::fs::remove_dir_all(directory).unwrap();
}
#[test]
fn only_ready_applications_accept_requests() {
let requests = RequestLifecycle::new();
assert_eq!(
requests.try_accept().unwrap_err().code(),
"runtime.not_ready"
);
requests.mark_ready();
let request = requests.try_accept().expect("ready request is admitted");
requests.begin_draining();
assert_eq!(requests.phase(), ApplicationPhase::Draining);
assert_eq!(
requests.try_accept().unwrap_err().code(),
"runtime.shutting_down"
);
drop(request);
test_runtime().block_on(requests.wait_until_drained());
requests.mark_stopped();
assert_eq!(requests.phase(), ApplicationPhase::Stopped);
}
#[test]
fn expired_drain_stops_admission_without_refunding_live_request() {
let requests = RequestLifecycle::new();
requests.mark_ready();
let request = requests.try_accept().unwrap();
requests.begin_draining();
requests.mark_stopped();
assert_eq!(requests.phase(), ApplicationPhase::Stopped);
assert_eq!(count(requests.shared.state.load(Ordering::Acquire)), 1);
assert_eq!(requests.try_accept().unwrap_err().code(), "runtime.stopped");
drop(request);
assert_eq!(count(requests.shared.state.load(Ordering::Acquire)), 0);
}
#[test]
fn draining_waits_for_every_admitted_request() {
test_runtime().block_on(async {
let requests = RequestLifecycle::new();
requests.mark_ready();
let first = requests.try_accept().unwrap();
let second = requests.try_accept().unwrap();
requests.begin_draining();
let requests_for_waiter = requests.clone();
let waiter = tokio::spawn(async move {
requests_for_waiter.wait_until_drained().await;
});
tokio::task::yield_now().await;
assert!(!waiter.is_finished());
drop(first);
tokio::task::yield_now().await;
assert!(!waiter.is_finished());
drop(second);
waiter.await.unwrap();
});
}
#[test]
fn admission_racing_with_drain_never_leaks_a_request() {
const WORKERS: usize = 8;
let requests = RequestLifecycle::new();
requests.mark_ready();
let barrier = Arc::new(Barrier::new(WORKERS + 1));
let workers: Vec<_> = (0..WORKERS)
.map(|_| {
let requests = requests.clone();
let barrier = Arc::clone(&barrier);
std::thread::spawn(move || {
let admitted_before_drain = requests.try_accept().unwrap();
barrier.wait();
loop {
match requests.try_accept() {
Ok(request) => drop(request),
Err(error) => {
assert_eq!(error.code(), "runtime.shutting_down");
drop(admitted_before_drain);
break;
}
}
}
})
})
.collect();
barrier.wait();
requests.begin_draining();
for worker in workers {
worker.join().unwrap();
}
test_runtime().block_on(requests.wait_until_drained());
assert_eq!(count(requests.shared.state.load(Ordering::Acquire)), 0);
}
#[test]
fn guard_completion_is_safe_from_another_thread() {
let requests = RequestLifecycle::new();
requests.mark_ready();
let guard = requests.try_accept().unwrap();
requests.begin_draining();
std::thread::spawn(move || drop(guard)).join().unwrap();
test_runtime().block_on(requests.wait_until_drained());
assert_eq!(Arc::strong_count(&requests.shared), 1);
}
}