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#[derive(Clone, Copy, Debug, Eq, PartialEq)]
16pub enum ApplicationPhase {
17 Starting = 0,
18 Ready = 1,
19 Draining = 2,
20 Stopped = 3,
21}
22
23#[derive(Clone, Debug)]
26pub struct ApplicationHealth {
27 shared: Arc<Shared>,
28}
29
30#[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 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 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#[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 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 pub fn try_accept(&self) -> Result<RequestGuard> {
135 self.try_claim().map(RequestClaim::publish)
136 }
137
138 #[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 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 record_lifecycle(&self.shared, "claim_cutoff", draining);
225 current = draining;
226 break;
227 }
228 Err(observed) => current = observed,
229 }
230 }
231 if count(current) == 0 {
232 self.shared.drained.notify_one();
233 }
234 }
235
236 pub(crate) fn observe_shutdown_decision(&self, succeeded: bool) {
237 record_lifecycle(&self.shared,
238 if succeeded { "shutdown_future_completed" } else { "shutdown_future_failed" },
239 self.shared.state.load(Ordering::Acquire));
240 }
241
242 pub(crate) fn shutdown_signal_source(&self) -> Option<(saddle_core::ContextLabel, EmergencyDiagnosticHandle)> {
243 let binding = self.shared.source_binding.get()?;
244 Some((saddle_core::ContextLabel::checked(binding.application.as_str()).ok()?,
245 binding.output.as_ref()?.clone()))
246 }
247
248 pub(crate) async fn wait_until_drained(&self) {
249 loop {
250 if count(self.shared.state.load(Ordering::Acquire)) == 0 {
251 return;
252 }
253 self.shared.drained.notified().await;
254 }
255 }
256
257 pub(crate) fn mark_stopped(&self) {
258 let mut current = self.shared.state.load(Ordering::Acquire);
262 loop {
263 debug_assert_eq!(phase(current), ApplicationPhase::Draining);
264 match self.shared.state.compare_exchange_weak(
265 current,
266 encode(ApplicationPhase::Stopped, count(current)),
267 Ordering::AcqRel,
268 Ordering::Acquire,
269 ) {
270 Ok(_) => return,
271 Err(observed) => current = observed,
272 }
273 }
274 }
275}
276
277#[derive(Debug)]
278struct AdmissionRejection {
279 code: &'static str,
280 observed_phase: ApplicationPhase,
281 observed_active_requests: usize,
282}
283
284#[derive(Debug)]
285struct AdmissionIdentityMismatch<'a> {
286 bound_application: &'a ApplicationId,
287 call_application: &'a ApplicationId,
288 observed_phase: ApplicationPhase,
289 observed_active_requests: usize,
290}
291
292impl fmt::Display for AdmissionIdentityMismatch<'_> {
293 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
294 write!(formatter, "admission application mismatch: bound={:?}, call={:?}, phase={:?}, active_requests={}",
295 self.bound_application, self.call_application,
296 self.observed_phase, self.observed_active_requests)
297 }
298}
299
300#[track_caller]
301fn admission_identity_mismatch(binding: &AdmissionSourceBinding,
302 call: &CallContext, state: usize) -> SaddleError {
303 let original = AdmissionIdentityMismatch {
304 bound_application: &binding.application,
305 call_application: call.application(),
306 observed_phase: phase(state), observed_active_requests: count(state),
307 };
308 let code = "runtime.admission_identity_mismatch";
309 let diagnostic = Diagnostic::capture(
310 DiagnosticCategory::UnexpectedError, CaptureSite::Origin,
311 DiagnosticCause::new(DiagnosticStage::RequestAdmission,
312 DiagnosticCode::new(code).expect("fixed identity code")),
313 );
314 let facts = runtime_unrooted_admission_source_description(binding.output.as_ref(),
315 &diagnostic, &original, Some(&binding.application));
316 let safe = SaddleError::new(ErrorKind::Internal, code,
317 "the request admission application does not match its bound lifecycle")
318 .with_diagnostic(diagnostic);
319 attach_admission_source(safe, facts)
320}
321
322impl fmt::Display for AdmissionRejection {
323 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
324 write!(formatter, "request admission rejected: code={}, phase={:?}, active_requests={}",
325 self.code, self.observed_phase, self.observed_active_requests)
326 }
327}
328
329#[track_caller]
330fn admission_rejection(
331 source: Option<(&CallContext, Option<&EmergencyDiagnosticHandle>)>,
332 binding: Option<&AdmissionSourceBinding>,
333 state: usize,
334 kind: ErrorKind,
335 code: &'static str,
336 message: &'static str,
337) -> SaddleError {
338 let original = AdmissionRejection {
339 code, observed_phase: phase(state), observed_active_requests: count(state),
340 };
341 let diagnostic = Diagnostic::capture(
342 DiagnosticCategory::UnexpectedError, CaptureSite::Origin,
343 DiagnosticCause::new(DiagnosticStage::RequestAdmission,
344 DiagnosticCode::new(code).expect("fixed admission code")),
345 );
346 let facts = match source {
347 Some((call, output)) => runtime_admission_source_description(output, &diagnostic, &original, call),
348 None => runtime_unrooted_admission_source_description(
349 binding.and_then(|binding| binding.output.as_ref()), &diagnostic, &original,
350 binding.map(|binding| &binding.application)),
351 };
352 let safe = SaddleError::new(kind, code, message).with_diagnostic(diagnostic);
353 attach_admission_source(safe, facts)
354}
355
356fn attach_admission_source(safe: SaddleError, facts: UnrootedCaptureFacts) -> SaddleError {
357 if facts.original_capture() == OriginalCaptureState::CompleteWritten
358 && safe.diagnostic().is_some_and(|diagnostic|
359 facts.occurrence().matches_diagnostic(diagnostic)) {
360 safe.with_source_receipt(facts)
361 } else {
362 safe.with_unconfirmed_source()
363 }
364}
365
366const fn encode(phase: ApplicationPhase, count: usize) -> usize {
367 ((phase as usize) << PHASE_SHIFT) | count
368}
369
370const fn phase(state: usize) -> ApplicationPhase {
371 match state >> PHASE_SHIFT {
372 0 => ApplicationPhase::Starting,
373 1 => ApplicationPhase::Ready,
374 2 => ApplicationPhase::Draining,
375 3 => ApplicationPhase::Stopped,
376 _ => unreachable!(),
377 }
378}
379
380const fn count(state: usize) -> usize {
381 state & COUNT_MASK
382}
383
384#[derive(Debug)]
386pub(crate) struct RequestClaim {
387 shared: Option<Arc<Shared>>,
388}
389
390impl RequestClaim {
391 pub(crate) fn publish(mut self) -> RequestGuard {
392 let guard = RequestGuard {
393 shared: self.shared.take().expect("request claim publishes once"),
394 };
395 record_lifecycle(&guard.shared, "claim_published",
396 guard.shared.state.load(Ordering::Acquire));
397 guard
398 }
399}
400
401fn record_lifecycle(shared: &Shared, event: &'static str, observed: usize) {
402 if let Some(binding) = shared.source_binding.get()
403 && let Some(output) = binding.output.as_ref() {
404 output.record_request_lifecycle(binding.application.as_str(), event,
405 phase(observed) as u8, count(observed));
406 }
407}
408
409impl Drop for RequestClaim {
410 fn drop(&mut self) {
411 if let Some(shared) = self.shared.take() {
412 complete(&shared);
413 }
414 }
415}
416
417#[derive(Debug)]
419pub struct RequestGuard {
420 shared: Arc<Shared>,
421}
422
423impl Drop for RequestGuard {
424 fn drop(&mut self) {
425 complete(&self.shared);
426 }
427}
428
429fn complete(shared: &Shared) {
430 let previous = shared.state.fetch_sub(1, Ordering::AcqRel);
431 debug_assert!(count(previous) > 0);
432 if count(previous) == 1 && phase(previous) == ApplicationPhase::Draining {
433 shared.drained.notify_one();
434 }
435}
436
437#[cfg(test)]
438mod tests {
439 use std::sync::{Arc, Barrier};
440
441 use super::*;
442
443 fn test_runtime() -> tokio::runtime::Runtime {
444 tokio::runtime::Builder::new_current_thread()
445 .build()
446 .expect("test runtime must build")
447 }
448
449 #[test]
450 fn notice27_stop_new_claim_preserves_old_claim_completion_under_cas_race() {
451 for _ in 0..32 {
452 let requests = Arc::new(RequestLifecycle::new());
453 requests.mark_ready();
454 let old = requests.try_claim().expect("pre-cutoff linear claim");
455 let barrier = Arc::new(Barrier::new(2));
456 let peer_requests = Arc::clone(&requests);
457 let peer_barrier = Arc::clone(&barrier);
458 let contender = std::thread::spawn(move || {
459 peer_barrier.wait();
460 peer_requests.try_claim()
461 });
462 barrier.wait();
463 requests.begin_draining();
464 assert_eq!(requests.try_claim().unwrap_err().code(), "runtime.shutting_down");
465 if let Ok(claim) = contender.join().unwrap() { drop(claim.publish()); }
468 let guard = old.publish();
469 assert_eq!(requests.phase(), ApplicationPhase::Draining);
470 assert_eq!(count(requests.shared.state.load(Ordering::Acquire)), 1);
471 drop(guard);
472 assert_eq!(count(requests.shared.state.load(Ordering::Acquire)), 0);
473 test_runtime().block_on(requests.wait_until_drained());
474 requests.mark_stopped();
475 }
476 }
477
478 #[test]
479 fn recorded_admission_rejections_keep_atomic_state_and_call_identity() {
480 use saddle_observability::{EmergencyDiagnostics, FileLoggingConfig, Observer,
481 ObserverConfig, Rotation};
482 use saddle_observability::root_diagnostic::UnrootedCaptureFacts;
483 let directory = std::env::temp_dir().join(format!("saddle-runtime-admission-{}-{}",
484 std::process::id(), std::time::SystemTime::now()
485 .duration_since(std::time::UNIX_EPOCH).unwrap().as_nanos()));
486 std::fs::create_dir(&directory).unwrap();
487 let mut writer = EmergencyDiagnostics::start(
488 &FileLoggingConfig::new(directory.clone(), Rotation::Daily)).unwrap();
489 let handle = writer.handle();
490 let observer = Observer::with_writer(ObserverConfig::default(), std::io::sink()).unwrap();
491 let (call, _) = observer.start_external_call_checked(
492 "admission-app", "runtime", "admission", "try_accept", None).unwrap();
493 let lifecycle = RequestLifecycle::new();
494 let starting = lifecycle.try_accept_recorded(call.context(), Some(&handle)).unwrap_err();
495 assert_eq!(starting.code(), "runtime.not_ready");
496 assert!(!starting.source_unavailable());
497 let receipt = starting.source_receipt::<UnrootedCaptureFacts>().unwrap();
498 assert_eq!(receipt.original_capture(), OriginalCaptureState::CompleteWritten);
499 assert!(receipt.occurrence().matches_diagnostic(starting.diagnostic().unwrap()));
500 lifecycle.mark_ready();
501 let guard = lifecycle.try_accept_recorded(call.context(), Some(&handle)).unwrap();
502 lifecycle.begin_draining();
503 let draining = lifecycle.try_accept_recorded(call.context(), Some(&handle)).unwrap_err();
504 assert_eq!(draining.code(), "runtime.shutting_down");
505 assert_ne!(starting.diagnostic().unwrap().id(), draining.diagnostic().unwrap().id());
506 let rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
507 .lines().map(|row| serde_json::from_str(row).unwrap()).collect();
508 for (error, expected) in [(&starting, "phase=Starting, active_requests=0"),
509 (&draining, "phase=Draining, active_requests=1")] {
510 let source: Vec<_> = rows.iter().filter(|row| row["occurrence"]["diagnostic_id"]
511 == error.diagnostic().unwrap().id()).collect();
512 let header: serde_json::Value = serde_json::from_str(source[0]["payload"].as_str().unwrap()).unwrap();
513 assert_eq!(header["context"]["application"], "admission-app");
514 assert_eq!(header["stage"], "runtime_request_admission");
515 assert!(source.iter().any(|row| row["channel"] == "description"
516 && row["payload"].as_str().is_some_and(|text| text.contains(expected))));
517 assert!(source.iter().any(|row| row["channel"] == "terminal"
518 && row["state"] == "description_complete_source_unavailable"));
519 }
520 let unconfirmed = lifecycle.try_accept_recorded(call.context(), None).unwrap_err();
521 assert!(unconfirmed.source_unavailable());
522 assert!(unconfirmed.source_receipt::<UnrootedCaptureFacts>().is_none());
523 let unrooted = RequestLifecycle::new();
524 assert!(unrooted.install_admission_source("unrooted-app".into(), Some(handle.clone())));
525 let no_call = unrooted.try_claim().unwrap_err();
526 assert_eq!(no_call.code(), "runtime.not_ready");
527 assert!(!no_call.source_unavailable());
528 let no_call_receipt = no_call.source_receipt::<UnrootedCaptureFacts>().unwrap();
529 assert_eq!(no_call_receipt.original_capture(), OriginalCaptureState::CompleteWritten);
530 assert!(no_call_receipt.occurrence().matches_diagnostic(no_call.diagnostic().unwrap()));
531 let unrooted_rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
532 .lines().map(|row| serde_json::from_str(row).unwrap()).collect();
533 let unrooted_source: Vec<_> = unrooted_rows.iter().filter(|row|
534 row["occurrence"]["diagnostic_id"] == no_call.diagnostic().unwrap().id()).collect();
535 let unrooted_header: serde_json::Value = serde_json::from_str(
536 unrooted_source[0]["payload"].as_str().unwrap()).unwrap();
537 assert_eq!(unrooted_header["stage"], "runtime_unrooted_admission");
538 assert_eq!(unrooted_header["context"]["application"]["value"], "unrooted-app");
539 assert_eq!(unrooted_header["context"]["trace_id"]["state"], "not_established");
540 assert!(unrooted_source.iter().any(|row| row["channel"] == "description"
541 && row["payload"].as_str().is_some_and(|text|
542 text.contains("phase=Starting, active_requests=0"))));
543 assert!(unrooted_source.iter().any(|row| row["channel"] == "terminal"
544 && row["state"] == "description_complete_source_unavailable"));
545 let mut application = crate::Application::new();
546 application.install_lifecycle_observer_with_output(
547 observer.clone(), "assembled-app", Some(handle.clone()));
548 let assembled = application.request_lifecycle().try_claim().unwrap_err();
549 assert!(!assembled.source_unavailable());
550 let assembled_rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
551 .lines().map(|row| serde_json::from_str(row).unwrap()).collect();
552 let assembled_source: Vec<_> = assembled_rows.iter().filter(|row|
553 row["occurrence"]["diagnostic_id"] == assembled.diagnostic().unwrap().id()).collect();
554 let assembled_header: serde_json::Value = serde_json::from_str(
555 assembled_source[0]["payload"].as_str().unwrap()).unwrap();
556 assert_eq!(assembled_header["context"]["application"]["value"], "assembled-app");
557 assert_eq!(assembled_header["stage"], "runtime_unrooted_admission");
558 let (wrong_call, _) = observer.start_external_call_checked(
559 "wrong-app", "runtime", "admission", "try_accept", None).unwrap();
560 let mismatch = application.request_lifecycle().try_accept_recorded(
561 wrong_call.context(), Some(&handle)).unwrap_err();
562 assert_eq!(mismatch.code(), "runtime.admission_identity_mismatch");
563 assert!(!mismatch.source_unavailable());
564 let mismatch_rows: Vec<serde_json::Value> = std::fs::read_to_string(writer.target()).unwrap()
565 .lines().map(|row| serde_json::from_str(row).unwrap()).collect();
566 let mismatch_source: Vec<_> = mismatch_rows.iter().filter(|row|
567 row["occurrence"]["diagnostic_id"] == mismatch.diagnostic().unwrap().id()).collect();
568 let mismatch_header: serde_json::Value = serde_json::from_str(
569 mismatch_source[0]["payload"].as_str().unwrap()).unwrap();
570 assert_eq!(mismatch_header["context"]["application"]["value"], "assembled-app");
571 assert!(mismatch_source.iter().any(|row| row["channel"] == "description"
572 && row["payload"].as_str().is_some_and(|text|
573 text.contains("assembled-app") && text.contains("wrong-app"))));
574 let unrelated = Diagnostic::capture(DiagnosticCategory::UnexpectedError,
575 CaptureSite::Origin, DiagnosticCause::new(DiagnosticStage::RequestAdmission,
576 DiagnosticCode::new("runtime.receipt_mismatch_test").unwrap()));
577 let facts = runtime_unrooted_admission_source_description(Some(&handle), &unrelated,
578 &AdmissionRejection { code: "runtime.receipt_mismatch_test",
579 observed_phase: ApplicationPhase::Starting, observed_active_requests: 0 },
580 Some(&ApplicationId::from("assembled-app")));
581 assert_eq!(facts.original_capture(), OriginalCaptureState::CompleteWritten);
582 let safe = SaddleError::new(ErrorKind::Internal,
583 "runtime.receipt_mismatch_test", "safe").with_diagnostic(
584 Diagnostic::capture(DiagnosticCategory::UnexpectedError, CaptureSite::Origin,
585 DiagnosticCause::new(DiagnosticStage::RequestAdmission,
586 DiagnosticCode::new("runtime.receipt_mismatch_test").unwrap())));
587 let false_receipt = attach_admission_source(safe, facts);
588 assert!(false_receipt.source_unavailable());
589 assert!(false_receipt.source_receipt::<UnrootedCaptureFacts>().is_none());
590 let no_binding = RequestLifecycle::new().try_claim().unwrap_err();
591 assert!(no_binding.source_unavailable());
592 drop(guard);
593 lifecycle.mark_stopped();
594 let until = std::time::Instant::now() + std::time::Duration::from_secs(2);
595 loop {
596 match writer.shutdown() {
597 saddle_observability::DiagnosticShutdown::Finished => break,
598 saddle_observability::DiagnosticShutdown::Pending
599 if std::time::Instant::now() < until =>
600 std::thread::sleep(std::time::Duration::from_millis(10)),
601 other => panic!("admission diagnostic shutdown failed: {other:?}"),
602 }
603 }
604 drop(writer);
605 std::fs::remove_dir_all(directory).unwrap();
606 }
607
608 #[test]
609 fn only_ready_applications_accept_requests() {
610 let requests = RequestLifecycle::new();
611 assert_eq!(
612 requests.try_accept().unwrap_err().code(),
613 "runtime.not_ready"
614 );
615
616 requests.mark_ready();
617 let request = requests.try_accept().expect("ready request is admitted");
618 requests.begin_draining();
619
620 assert_eq!(requests.phase(), ApplicationPhase::Draining);
621 assert_eq!(
622 requests.try_accept().unwrap_err().code(),
623 "runtime.shutting_down"
624 );
625
626 drop(request);
627 test_runtime().block_on(requests.wait_until_drained());
628 requests.mark_stopped();
629 assert_eq!(requests.phase(), ApplicationPhase::Stopped);
630 }
631
632 #[test]
633 fn expired_drain_stops_admission_without_refunding_live_request() {
634 let requests = RequestLifecycle::new();
635 requests.mark_ready();
636 let request = requests.try_accept().unwrap();
637 requests.begin_draining();
638 requests.mark_stopped();
639 assert_eq!(requests.phase(), ApplicationPhase::Stopped);
640 assert_eq!(count(requests.shared.state.load(Ordering::Acquire)), 1);
641 assert_eq!(requests.try_accept().unwrap_err().code(), "runtime.stopped");
642 drop(request);
643 assert_eq!(count(requests.shared.state.load(Ordering::Acquire)), 0);
644 }
645
646 #[test]
647 fn draining_waits_for_every_admitted_request() {
648 test_runtime().block_on(async {
649 let requests = RequestLifecycle::new();
650 requests.mark_ready();
651 let first = requests.try_accept().unwrap();
652 let second = requests.try_accept().unwrap();
653 requests.begin_draining();
654
655 let requests_for_waiter = requests.clone();
656 let waiter = tokio::spawn(async move {
657 requests_for_waiter.wait_until_drained().await;
658 });
659 tokio::task::yield_now().await;
660 assert!(!waiter.is_finished());
661
662 drop(first);
663 tokio::task::yield_now().await;
664 assert!(!waiter.is_finished());
665
666 drop(second);
667 waiter.await.unwrap();
668 });
669 }
670
671 #[test]
672 fn admission_racing_with_drain_never_leaks_a_request() {
673 const WORKERS: usize = 8;
674 let requests = RequestLifecycle::new();
675 requests.mark_ready();
676 let barrier = Arc::new(Barrier::new(WORKERS + 1));
677 let workers: Vec<_> = (0..WORKERS)
678 .map(|_| {
679 let requests = requests.clone();
680 let barrier = Arc::clone(&barrier);
681 std::thread::spawn(move || {
682 let admitted_before_drain = requests.try_accept().unwrap();
683 barrier.wait();
684 loop {
685 match requests.try_accept() {
686 Ok(request) => drop(request),
687 Err(error) => {
688 assert_eq!(error.code(), "runtime.shutting_down");
689 drop(admitted_before_drain);
690 break;
691 }
692 }
693 }
694 })
695 })
696 .collect();
697
698 barrier.wait();
699 requests.begin_draining();
700 for worker in workers {
701 worker.join().unwrap();
702 }
703 test_runtime().block_on(requests.wait_until_drained());
704 assert_eq!(count(requests.shared.state.load(Ordering::Acquire)), 0);
705 }
706
707 #[test]
708 fn guard_completion_is_safe_from_another_thread() {
709 let requests = RequestLifecycle::new();
710 requests.mark_ready();
711 let guard = requests.try_accept().unwrap();
712 requests.begin_draining();
713
714 std::thread::spawn(move || drop(guard)).join().unwrap();
715 test_runtime().block_on(requests.wait_until_drained());
716
717 assert_eq!(Arc::strong_count(&requests.shared), 1);
720 }
721}