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 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 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#[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#[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 assert_eq!(Arc::strong_count(&requests.shared), 1);
666 }
667}