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