Skip to main content

saddle_runtime/
request.rs

1use std::sync::{
2    Arc, OnceLock,
3    atomic::{AtomicUsize, Ordering},
4};
5use std::fmt;
6
7use saddle_core::{ApplicationId, CallContext, CaptureSite, Diagnostic, DiagnosticCategory, DiagnosticCause,
8    DiagnosticCode, DiagnosticStage, ErrorKind, Result, SaddleError};
9use saddle_observability::EmergencyDiagnosticHandle;
10use saddle_observability::root_diagnostic::{OriginalCaptureState, UnrootedCaptureFacts,
11    runtime_admission_source_description, runtime_unrooted_admission_source_description};
12use tokio::sync::Notify;
13
14/// The externally visible phase of the Saddle application lifecycle.
15#[derive(Clone, Copy, Debug, Eq, PartialEq)]
16pub enum ApplicationPhase {
17    Starting = 0,
18    Ready = 1,
19    Draining = 2,
20    Stopped = 3,
21}
22
23/// Read-only process health observer backed by the application's unique
24/// request-lifecycle state. It carries no request-admission authority.
25#[derive(Clone, Debug)]
26pub struct ApplicationHealth {
27    shared: Arc<Shared>,
28}
29
30/// One atomic view of process liveness and readiness.
31///
32/// Liveness remains true while the managed runtime can complete startup or
33/// drain admitted work. Readiness is true only after all components have
34/// started and becomes false before shutdown begins stopping components.
35#[derive(Clone, Copy, Debug, Eq, PartialEq)]
36pub struct ApplicationHealthSnapshot {
37    phase: ApplicationPhase,
38}
39
40impl ApplicationHealthSnapshot {
41    pub const fn phase(self) -> ApplicationPhase {
42        self.phase
43    }
44
45    pub const fn is_live(self) -> bool {
46        !matches!(self.phase, ApplicationPhase::Stopped)
47    }
48
49    pub const fn is_ready(self) -> bool {
50        matches!(self.phase, ApplicationPhase::Ready)
51    }
52}
53
54impl ApplicationHealth {
55    pub(crate) fn fail_closed(&self) {
56        RequestLifecycle { shared: Arc::clone(&self.shared) }.begin_draining();
57    }
58    /// Reads phase, liveness and readiness from one atomic lifecycle value.
59    pub fn snapshot(&self) -> ApplicationHealthSnapshot {
60        ApplicationHealthSnapshot {
61            phase: phase(self.shared.state.load(Ordering::Acquire)),
62        }
63    }
64}
65
66const PHASE_SHIFT: u32 = usize::BITS - 2;
67const COUNT_MASK: usize = (1 << PHASE_SHIFT) - 1;
68
69#[derive(Debug)]
70struct Shared {
71    /// The top two bits hold `ApplicationPhase`; the remaining bits hold the
72    /// number of admitted requests. Updating both in one CAS closes the race
73    /// between request admission and the transition to draining.
74    state: AtomicUsize,
75    drained: Notify,
76    source_binding: OnceLock<AdmissionSourceBinding>,
77}
78
79struct AdmissionSourceBinding {
80    application: ApplicationId,
81    output: Option<EmergencyDiagnosticHandle>,
82}
83
84impl fmt::Debug for AdmissionSourceBinding {
85    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
86        formatter.debug_struct("AdmissionSourceBinding")
87            .field("application", &self.application)
88            .field("output_selected", &self.output.is_some()).finish()
89    }
90}
91
92/// Coordinates request admission with graceful application shutdown.
93///
94/// Service adapters must acquire a [`RequestGuard`] before dispatching an
95/// accepted request. Once shutdown begins, new acquisitions fail while guards
96/// already issued remain valid until the corresponding request completes.
97#[derive(Clone, Debug)]
98pub struct RequestLifecycle {
99    shared: Arc<Shared>,
100}
101
102impl RequestLifecycle {
103    pub(crate) fn new() -> Self {
104        Self {
105            shared: Arc::new(Shared {
106                state: AtomicUsize::new(encode(ApplicationPhase::Starting, 0)),
107                drained: Notify::new(),
108                source_binding: OnceLock::new(),
109            }),
110        }
111    }
112
113    /// Returns the application's current lifecycle phase.
114    pub fn phase(&self) -> ApplicationPhase {
115        phase(self.shared.state.load(Ordering::Acquire))
116    }
117
118    pub(crate) fn health(&self) -> ApplicationHealth {
119        ApplicationHealth {
120            shared: Arc::clone(&self.shared),
121        }
122    }
123
124    pub(crate) fn install_admission_source(&self, application: ApplicationId,
125        output: Option<EmergencyDiagnosticHandle>) -> bool {
126        self.shared.source_binding.set(AdmissionSourceBinding { application, output }).is_ok()
127    }
128
129    /// Admits one request if the application is ready and accepting work.
130    ///
131    /// The returned guard must be held for the complete request execution. Its
132    /// `Drop` implementation records completion even when the request future is
133    /// cancelled or unwinds.
134    pub fn try_accept(&self) -> Result<RequestGuard> {
135        self.try_claim().map(RequestClaim::publish)
136    }
137
138    /// The Service entry supplies the established call and selected file so
139    /// a rejected atomic admission can be recorded at its source.
140    #[doc(hidden)]
141    pub fn try_accept_recorded(&self, call: &CallContext,
142        output: Option<&EmergencyDiagnosticHandle>) -> Result<RequestGuard> {
143        if let Some(binding) = self.shared.source_binding.get()
144            && &binding.application != call.application() {
145            return Err(admission_identity_mismatch(binding, call,
146                self.shared.state.load(Ordering::Acquire)));
147        }
148        self.try_claim_with_source(Some((call, output))).map(RequestClaim::publish)
149    }
150
151    /// Claims one request slot without publishing executable work.
152    ///
153    /// Runtime holds this linear claim across Admission commitment and turns it
154    /// into a [`RequestGuard`] only when the owned envelope is ready to submit.
155    pub(crate) fn try_claim(&self) -> Result<RequestClaim> {
156        self.try_claim_with_source(None)
157    }
158
159    fn try_claim_with_source(&self,
160        source: Option<(&CallContext, Option<&EmergencyDiagnosticHandle>)>) -> Result<RequestClaim> {
161        let mut current = self.shared.state.load(Ordering::Acquire);
162        loop {
163            match phase(current) {
164                ApplicationPhase::Ready => {
165                    if count(current) == COUNT_MASK {
166                        return Err(admission_rejection(source, self.shared.source_binding.get(), current, ErrorKind::Internal,
167                            "runtime.request_count_overflow", "request accounting capacity exhausted"));
168                    }
169                    match self.shared.state.compare_exchange_weak(
170                        current,
171                        current + 1,
172                        Ordering::AcqRel,
173                        Ordering::Acquire,
174                    ) {
175                        Ok(_) => {
176                            return Ok(RequestClaim {
177                                shared: Some(Arc::clone(&self.shared)),
178                            });
179                        }
180                        Err(observed) => current = observed,
181                    }
182                }
183                ApplicationPhase::Starting => {
184                    return Err(admission_rejection(source, self.shared.source_binding.get(), current, ErrorKind::Unavailable,
185                        "runtime.not_ready", "application is not ready"));
186                }
187                ApplicationPhase::Draining => {
188                    return Err(admission_rejection(source, self.shared.source_binding.get(), current, ErrorKind::Unavailable,
189                        "runtime.shutting_down", "application is shutting down"));
190                }
191                ApplicationPhase::Stopped => {
192                    return Err(admission_rejection(source, self.shared.source_binding.get(), current, ErrorKind::Unavailable,
193                        "runtime.stopped", "application is stopped"));
194                }
195            }
196        }
197    }
198
199    pub(crate) fn mark_ready(&self) {
200        let result = self.shared.state.compare_exchange(
201            encode(ApplicationPhase::Starting, 0),
202            encode(ApplicationPhase::Ready, 0),
203            Ordering::AcqRel,
204            Ordering::Acquire,
205        );
206        debug_assert!(result.is_ok());
207    }
208
209    pub(crate) fn begin_draining(&self) {
210        let mut current = self.shared.state.load(Ordering::Acquire);
211        while matches!(
212            phase(current),
213            ApplicationPhase::Starting | ApplicationPhase::Ready
214        ) {
215            let draining = encode(ApplicationPhase::Draining, count(current));
216            match self.shared.state.compare_exchange_weak(
217                current,
218                draining,
219                Ordering::AcqRel,
220                Ordering::Acquire,
221            ) {
222                Ok(_) => {
223                    current = draining;
224                    break;
225                }
226                Err(observed) => current = observed,
227            }
228        }
229        if count(current) == 0 {
230            self.shared.drained.notify_one();
231        }
232    }
233
234    pub(crate) async fn wait_until_drained(&self) {
235        loop {
236            if count(self.shared.state.load(Ordering::Acquire)) == 0 {
237                return;
238            }
239            self.shared.drained.notified().await;
240        }
241    }
242
243    pub(crate) fn mark_stopped(&self) {
244        // A drain deadline can expire while admitted guards still exist.
245        // Stop admission and preserve their physical count until each guard
246        // actually drops; never report an early resource refund.
247        let mut current = self.shared.state.load(Ordering::Acquire);
248        loop {
249            debug_assert_eq!(phase(current), ApplicationPhase::Draining);
250            match self.shared.state.compare_exchange_weak(
251                current,
252                encode(ApplicationPhase::Stopped, count(current)),
253                Ordering::AcqRel,
254                Ordering::Acquire,
255            ) {
256                Ok(_) => return,
257                Err(observed) => current = observed,
258            }
259        }
260    }
261}
262
263#[derive(Debug)]
264struct AdmissionRejection {
265    code: &'static str,
266    observed_phase: ApplicationPhase,
267    observed_active_requests: usize,
268}
269
270#[derive(Debug)]
271struct AdmissionIdentityMismatch<'a> {
272    bound_application: &'a ApplicationId,
273    call_application: &'a ApplicationId,
274    observed_phase: ApplicationPhase,
275    observed_active_requests: usize,
276}
277
278impl fmt::Display for AdmissionIdentityMismatch<'_> {
279    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
280        write!(formatter, "admission application mismatch: bound={:?}, call={:?}, phase={:?}, active_requests={}",
281            self.bound_application, self.call_application,
282            self.observed_phase, self.observed_active_requests)
283    }
284}
285
286#[track_caller]
287fn admission_identity_mismatch(binding: &AdmissionSourceBinding,
288    call: &CallContext, state: usize) -> SaddleError {
289    let original = AdmissionIdentityMismatch {
290        bound_application: &binding.application,
291        call_application: call.application(),
292        observed_phase: phase(state), observed_active_requests: count(state),
293    };
294    let code = "runtime.admission_identity_mismatch";
295    let diagnostic = Diagnostic::capture(
296        DiagnosticCategory::UnexpectedError, CaptureSite::Origin,
297        DiagnosticCause::new(DiagnosticStage::RequestAdmission,
298            DiagnosticCode::new(code).expect("fixed identity code")),
299    );
300    let facts = runtime_unrooted_admission_source_description(binding.output.as_ref(),
301        &diagnostic, &original, Some(&binding.application));
302    let safe = SaddleError::new(ErrorKind::Internal, code,
303        "the request admission application does not match its bound lifecycle")
304        .with_diagnostic(diagnostic);
305    attach_admission_source(safe, facts)
306}
307
308impl fmt::Display for AdmissionRejection {
309    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
310        write!(formatter, "request admission rejected: code={}, phase={:?}, active_requests={}",
311            self.code, self.observed_phase, self.observed_active_requests)
312    }
313}
314
315#[track_caller]
316fn admission_rejection(
317    source: Option<(&CallContext, Option<&EmergencyDiagnosticHandle>)>,
318    binding: Option<&AdmissionSourceBinding>,
319    state: usize,
320    kind: ErrorKind,
321    code: &'static str,
322    message: &'static str,
323) -> SaddleError {
324    let original = AdmissionRejection {
325        code, observed_phase: phase(state), observed_active_requests: count(state),
326    };
327    let diagnostic = Diagnostic::capture(
328        DiagnosticCategory::UnexpectedError, CaptureSite::Origin,
329        DiagnosticCause::new(DiagnosticStage::RequestAdmission,
330            DiagnosticCode::new(code).expect("fixed admission code")),
331    );
332    let facts = match source {
333        Some((call, output)) => runtime_admission_source_description(output, &diagnostic, &original, call),
334        None => runtime_unrooted_admission_source_description(
335            binding.and_then(|binding| binding.output.as_ref()), &diagnostic, &original,
336            binding.map(|binding| &binding.application)),
337    };
338    let safe = SaddleError::new(kind, code, message).with_diagnostic(diagnostic);
339    attach_admission_source(safe, facts)
340}
341
342fn attach_admission_source(safe: SaddleError, facts: UnrootedCaptureFacts) -> SaddleError {
343    if facts.original_capture() == OriginalCaptureState::CompleteWritten
344        && safe.diagnostic().is_some_and(|diagnostic|
345            facts.occurrence().matches_diagnostic(diagnostic)) {
346        safe.with_source_receipt(facts)
347    } else {
348        safe.with_unconfirmed_source()
349    }
350}
351
352const fn encode(phase: ApplicationPhase, count: usize) -> usize {
353    ((phase as usize) << PHASE_SHIFT) | count
354}
355
356const fn phase(state: usize) -> ApplicationPhase {
357    match state >> PHASE_SHIFT {
358        0 => ApplicationPhase::Starting,
359        1 => ApplicationPhase::Ready,
360        2 => ApplicationPhase::Draining,
361        3 => ApplicationPhase::Stopped,
362        _ => unreachable!(),
363    }
364}
365
366const fn count(state: usize) -> usize {
367    state & COUNT_MASK
368}
369
370/// An unpublished lifecycle claim held only by Runtime.
371#[derive(Debug)]
372pub(crate) struct RequestClaim {
373    shared: Option<Arc<Shared>>,
374}
375
376impl RequestClaim {
377    pub(crate) fn publish(mut self) -> RequestGuard {
378        RequestGuard {
379            shared: self.shared.take().expect("request claim publishes once"),
380        }
381    }
382}
383
384impl Drop for RequestClaim {
385    fn drop(&mut self) {
386        if let Some(shared) = self.shared.take() {
387            complete(&shared);
388        }
389    }
390}
391
392/// Proof that one request was admitted by Saddle before shutdown began.
393#[derive(Debug)]
394pub struct RequestGuard {
395    shared: Arc<Shared>,
396}
397
398impl Drop for RequestGuard {
399    fn drop(&mut self) {
400        complete(&self.shared);
401    }
402}
403
404fn complete(shared: &Shared) {
405    let previous = shared.state.fetch_sub(1, Ordering::AcqRel);
406    debug_assert!(count(previous) > 0);
407    if count(previous) == 1 && phase(previous) == ApplicationPhase::Draining {
408        shared.drained.notify_one();
409    }
410}
411
412#[cfg(test)]
413mod tests {
414    use std::sync::{Arc, Barrier};
415
416    use super::*;
417
418    fn test_runtime() -> tokio::runtime::Runtime {
419        tokio::runtime::Builder::new_current_thread()
420            .build()
421            .expect("test runtime must build")
422    }
423
424    #[test]
425    fn recorded_admission_rejections_keep_atomic_state_and_call_identity() {
426        use saddle_observability::{EmergencyDiagnostics, FileLoggingConfig, Observer,
427            ObserverConfig, Rotation};
428        use saddle_observability::root_diagnostic::UnrootedCaptureFacts;
429        let directory = std::env::temp_dir().join(format!("saddle-runtime-admission-{}-{}",
430            std::process::id(), std::time::SystemTime::now()
431                .duration_since(std::time::UNIX_EPOCH).unwrap().as_nanos()));
432        std::fs::create_dir(&directory).unwrap();
433        let mut writer = EmergencyDiagnostics::start(
434            &FileLoggingConfig::new(directory.clone(), Rotation::Daily)).unwrap();
435        let handle = writer.handle();
436        let observer = Observer::with_writer(ObserverConfig::default(), std::io::sink()).unwrap();
437        let (call, _) = observer.start_external_call_checked(
438            "admission-app", "runtime", "admission", "try_accept", None).unwrap();
439        let lifecycle = RequestLifecycle::new();
440        let starting = lifecycle.try_accept_recorded(call.context(), Some(&handle)).unwrap_err();
441        assert_eq!(starting.code(), "runtime.not_ready");
442        assert!(!starting.source_unavailable());
443        let receipt = starting.source_receipt::<UnrootedCaptureFacts>().unwrap();
444        assert_eq!(receipt.original_capture(), OriginalCaptureState::CompleteWritten);
445        assert!(receipt.occurrence().matches_diagnostic(starting.diagnostic().unwrap()));
446        lifecycle.mark_ready();
447        let guard = lifecycle.try_accept_recorded(call.context(), Some(&handle)).unwrap();
448        lifecycle.begin_draining();
449        let draining = lifecycle.try_accept_recorded(call.context(), Some(&handle)).unwrap_err();
450        assert_eq!(draining.code(), "runtime.shutting_down");
451        assert_ne!(starting.diagnostic().unwrap().id(), draining.diagnostic().unwrap().id());
452        let rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
453            .lines().map(|row| serde_json::from_str(row).unwrap()).collect();
454        for (error, expected) in [(&starting, "phase=Starting, active_requests=0"),
455            (&draining, "phase=Draining, active_requests=1")] {
456            let source: Vec<_> = rows.iter().filter(|row| row["occurrence"]["diagnostic_id"]
457                == error.diagnostic().unwrap().id()).collect();
458            let header: serde_json::Value = serde_json::from_str(source[0]["payload"].as_str().unwrap()).unwrap();
459            assert_eq!(header["context"]["application"], "admission-app");
460            assert_eq!(header["stage"], "runtime_request_admission");
461            assert!(source.iter().any(|row| row["channel"] == "description"
462                && row["payload"].as_str().is_some_and(|text| text.contains(expected))));
463            assert!(source.iter().any(|row| row["channel"] == "terminal"
464                && row["state"] == "description_complete_source_unavailable"));
465        }
466        let unconfirmed = lifecycle.try_accept_recorded(call.context(), None).unwrap_err();
467        assert!(unconfirmed.source_unavailable());
468        assert!(unconfirmed.source_receipt::<UnrootedCaptureFacts>().is_none());
469        let unrooted = RequestLifecycle::new();
470        assert!(unrooted.install_admission_source("unrooted-app".into(), Some(handle.clone())));
471        let no_call = unrooted.try_claim().unwrap_err();
472        assert_eq!(no_call.code(), "runtime.not_ready");
473        assert!(!no_call.source_unavailable());
474        let no_call_receipt = no_call.source_receipt::<UnrootedCaptureFacts>().unwrap();
475        assert_eq!(no_call_receipt.original_capture(), OriginalCaptureState::CompleteWritten);
476        assert!(no_call_receipt.occurrence().matches_diagnostic(no_call.diagnostic().unwrap()));
477        let unrooted_rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
478            .lines().map(|row| serde_json::from_str(row).unwrap()).collect();
479        let unrooted_source: Vec<_> = unrooted_rows.iter().filter(|row|
480            row["occurrence"]["diagnostic_id"] == no_call.diagnostic().unwrap().id()).collect();
481        let unrooted_header: serde_json::Value = serde_json::from_str(
482            unrooted_source[0]["payload"].as_str().unwrap()).unwrap();
483        assert_eq!(unrooted_header["stage"], "runtime_unrooted_admission");
484        assert_eq!(unrooted_header["context"]["application"]["value"], "unrooted-app");
485        assert_eq!(unrooted_header["context"]["trace_id"]["state"], "not_established");
486        assert!(unrooted_source.iter().any(|row| row["channel"] == "description"
487            && row["payload"].as_str().is_some_and(|text|
488                text.contains("phase=Starting, active_requests=0"))));
489        assert!(unrooted_source.iter().any(|row| row["channel"] == "terminal"
490            && row["state"] == "description_complete_source_unavailable"));
491        let mut application = crate::Application::new();
492        application.install_lifecycle_observer_with_output(
493            observer.clone(), "assembled-app", Some(handle.clone()));
494        let assembled = application.request_lifecycle().try_claim().unwrap_err();
495        assert!(!assembled.source_unavailable());
496        let assembled_rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
497            .lines().map(|row| serde_json::from_str(row).unwrap()).collect();
498        let assembled_source: Vec<_> = assembled_rows.iter().filter(|row|
499            row["occurrence"]["diagnostic_id"] == assembled.diagnostic().unwrap().id()).collect();
500        let assembled_header: serde_json::Value = serde_json::from_str(
501            assembled_source[0]["payload"].as_str().unwrap()).unwrap();
502        assert_eq!(assembled_header["context"]["application"]["value"], "assembled-app");
503        assert_eq!(assembled_header["stage"], "runtime_unrooted_admission");
504        let (wrong_call, _) = observer.start_external_call_checked(
505            "wrong-app", "runtime", "admission", "try_accept", None).unwrap();
506        let mismatch = application.request_lifecycle().try_accept_recorded(
507            wrong_call.context(), Some(&handle)).unwrap_err();
508        assert_eq!(mismatch.code(), "runtime.admission_identity_mismatch");
509        assert!(!mismatch.source_unavailable());
510        let mismatch_rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
511            .lines().map(|row| serde_json::from_str(row).unwrap()).collect();
512        let mismatch_source: Vec<_> = mismatch_rows.iter().filter(|row|
513            row["occurrence"]["diagnostic_id"] == mismatch.diagnostic().unwrap().id()).collect();
514        let mismatch_header: serde_json::Value = serde_json::from_str(
515            mismatch_source[0]["payload"].as_str().unwrap()).unwrap();
516        assert_eq!(mismatch_header["context"]["application"]["value"], "assembled-app");
517        assert!(mismatch_source.iter().any(|row| row["channel"] == "description"
518            && row["payload"].as_str().is_some_and(|text|
519                text.contains("assembled-app") && text.contains("wrong-app"))));
520        let unrelated = Diagnostic::capture(DiagnosticCategory::UnexpectedError,
521            CaptureSite::Origin, DiagnosticCause::new(DiagnosticStage::RequestAdmission,
522                DiagnosticCode::new("runtime.receipt_mismatch_test").unwrap()));
523        let facts = runtime_unrooted_admission_source_description(Some(&handle), &unrelated,
524            &AdmissionRejection { code: "runtime.receipt_mismatch_test",
525                observed_phase: ApplicationPhase::Starting, observed_active_requests: 0 },
526            Some(&ApplicationId::from("assembled-app")));
527        assert_eq!(facts.original_capture(), OriginalCaptureState::CompleteWritten);
528        let safe = SaddleError::new(ErrorKind::Internal,
529            "runtime.receipt_mismatch_test", "safe").with_diagnostic(
530                Diagnostic::capture(DiagnosticCategory::UnexpectedError, CaptureSite::Origin,
531                    DiagnosticCause::new(DiagnosticStage::RequestAdmission,
532                        DiagnosticCode::new("runtime.receipt_mismatch_test").unwrap())));
533        let false_receipt = attach_admission_source(safe, facts);
534        assert!(false_receipt.source_unavailable());
535        assert!(false_receipt.source_receipt::<UnrootedCaptureFacts>().is_none());
536        let no_binding = RequestLifecycle::new().try_claim().unwrap_err();
537        assert!(no_binding.source_unavailable());
538        drop(guard);
539        lifecycle.mark_stopped();
540        let until = std::time::Instant::now() + std::time::Duration::from_secs(2);
541        loop {
542            match writer.shutdown() {
543                saddle_observability::DiagnosticShutdown::Finished => break,
544                saddle_observability::DiagnosticShutdown::Pending
545                    if std::time::Instant::now() < until =>
546                    std::thread::sleep(std::time::Duration::from_millis(10)),
547                other => panic!("admission diagnostic shutdown failed: {other:?}"),
548            }
549        }
550        drop(writer);
551        std::fs::remove_dir_all(directory).unwrap();
552    }
553
554    #[test]
555    fn only_ready_applications_accept_requests() {
556        let requests = RequestLifecycle::new();
557        assert_eq!(
558            requests.try_accept().unwrap_err().code(),
559            "runtime.not_ready"
560        );
561
562        requests.mark_ready();
563        let request = requests.try_accept().expect("ready request is admitted");
564        requests.begin_draining();
565
566        assert_eq!(requests.phase(), ApplicationPhase::Draining);
567        assert_eq!(
568            requests.try_accept().unwrap_err().code(),
569            "runtime.shutting_down"
570        );
571
572        drop(request);
573        test_runtime().block_on(requests.wait_until_drained());
574        requests.mark_stopped();
575        assert_eq!(requests.phase(), ApplicationPhase::Stopped);
576    }
577
578    #[test]
579    fn expired_drain_stops_admission_without_refunding_live_request() {
580        let requests = RequestLifecycle::new();
581        requests.mark_ready();
582        let request = requests.try_accept().unwrap();
583        requests.begin_draining();
584        requests.mark_stopped();
585        assert_eq!(requests.phase(), ApplicationPhase::Stopped);
586        assert_eq!(count(requests.shared.state.load(Ordering::Acquire)), 1);
587        assert_eq!(requests.try_accept().unwrap_err().code(), "runtime.stopped");
588        drop(request);
589        assert_eq!(count(requests.shared.state.load(Ordering::Acquire)), 0);
590    }
591
592    #[test]
593    fn draining_waits_for_every_admitted_request() {
594        test_runtime().block_on(async {
595            let requests = RequestLifecycle::new();
596            requests.mark_ready();
597            let first = requests.try_accept().unwrap();
598            let second = requests.try_accept().unwrap();
599            requests.begin_draining();
600
601            let requests_for_waiter = requests.clone();
602            let waiter = tokio::spawn(async move {
603                requests_for_waiter.wait_until_drained().await;
604            });
605            tokio::task::yield_now().await;
606            assert!(!waiter.is_finished());
607
608            drop(first);
609            tokio::task::yield_now().await;
610            assert!(!waiter.is_finished());
611
612            drop(second);
613            waiter.await.unwrap();
614        });
615    }
616
617    #[test]
618    fn admission_racing_with_drain_never_leaks_a_request() {
619        const WORKERS: usize = 8;
620        let requests = RequestLifecycle::new();
621        requests.mark_ready();
622        let barrier = Arc::new(Barrier::new(WORKERS + 1));
623        let workers: Vec<_> = (0..WORKERS)
624            .map(|_| {
625                let requests = requests.clone();
626                let barrier = Arc::clone(&barrier);
627                std::thread::spawn(move || {
628                    let admitted_before_drain = requests.try_accept().unwrap();
629                    barrier.wait();
630                    loop {
631                        match requests.try_accept() {
632                            Ok(request) => drop(request),
633                            Err(error) => {
634                                assert_eq!(error.code(), "runtime.shutting_down");
635                                drop(admitted_before_drain);
636                                break;
637                            }
638                        }
639                    }
640                })
641            })
642            .collect();
643
644        barrier.wait();
645        requests.begin_draining();
646        for worker in workers {
647            worker.join().unwrap();
648        }
649        test_runtime().block_on(requests.wait_until_drained());
650        assert_eq!(count(requests.shared.state.load(Ordering::Acquire)), 0);
651    }
652
653    #[test]
654    fn guard_completion_is_safe_from_another_thread() {
655        let requests = RequestLifecycle::new();
656        requests.mark_ready();
657        let guard = requests.try_accept().unwrap();
658        requests.begin_draining();
659
660        std::thread::spawn(move || drop(guard)).join().unwrap();
661        test_runtime().block_on(requests.wait_until_drained());
662
663        // Keep this assertion explicit: the thread above is runtime-internal
664        // test coverage, not a business-facing execution API.
665        assert_eq!(Arc::strong_count(&requests.shared), 1);
666    }
667}