1use mj_core::state::RecoveryObservation;
3use std::collections::{BTreeMap, BTreeSet, VecDeque};
4use std::sync::atomic::{AtomicBool, Ordering};
5use std::sync::{Arc, Mutex};
6use tokio::sync::{Notify, mpsc, watch};
7pub(crate) fn current_background_session(
10 observed: &mj_core::state::SessionRecord,
11) -> anyhow::Result<Option<mj_core::state::SessionRecord>> {
12 let current = crate::database::read_durable_session_record(&observed.id)?;
13 Ok(current.filter(|current| {
14 current.state == mj_core::state::SessionState::Running
15 && current.target == observed.target
16 && current.native_session_id == observed.native_session_id
17 && current.harness_kind == observed.harness_kind
18 && current.last_profile == observed.last_profile
19 }))
20}
21
22#[derive(Clone)]
32pub struct RecoveryObserver {
33 pub(crate) observations: ObservationSender<PendingRecoveryObservation>,
34 pub gate: Arc<RecoveryGate>,
35}
36
37#[derive(Clone)]
38pub(crate) struct PendingRecoveryObservation {
39 pub observation: RecoveryObservation,
40 pub observed_wait: bool,
43}
44
45pub struct RecoveryReservation {
48 session_id: String,
49 gate: Arc<RecoveryGate>,
50}
51
52impl Drop for RecoveryReservation {
53 fn drop(&mut self) {
54 self.gate.release(&self.session_id);
55 }
56}
57
58pub struct RecoveryGate {
64 state: Mutex<RecoveryGateState>,
65 closed: Notify,
66 busy: watch::Sender<BTreeSet<String>>,
69}
70
71impl Default for RecoveryGate {
72 fn default() -> Self {
73 Self {
74 state: Mutex::default(),
75 closed: Notify::new(),
76 busy: watch::channel(BTreeSet::new()).0,
77 }
78 }
79}
80
81#[derive(Default)]
82struct RecoveryGateState {
83 closed: bool,
84 busy: BTreeMap<String, Arc<AtomicBool>>,
87 reservations: BTreeMap<String, usize>,
88}
89
90impl RecoveryGate {
91 pub async fn run_background<T>(
95 self: &Arc<Self>,
96 session_id: &str,
97 work: impl std::future::Future<Output = T>,
98 ) -> Option<T> {
99 let admission = self.try_start(session_id)?;
100 let cancelled = admission.cancellation();
101 let cancellation = async {
102 while !cancelled.load(Ordering::Acquire) {
103 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
104 }
105 };
106 tokio::select! {
107 biased;
108 _ = cancellation => None,
109 result = work => Some(result),
110 }
111 }
112
113 pub fn reserve(self: &Arc<Self>, session_id: &str) -> RecoveryReservation {
114 let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
115 *state.reservations.entry(session_id.to_owned()).or_default() += 1;
116 RecoveryReservation {
117 session_id: session_id.to_owned(),
118 gate: self.clone(),
119 }
120 }
121
122 fn release(&self, session_id: &str) {
123 let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
124 let Some(count) = state.reservations.get_mut(session_id) else {
125 return;
126 };
127 *count -= 1;
128 if *count == 0 {
129 state.reservations.remove(session_id);
130 }
131 }
132
133 pub fn try_start(self: &Arc<Self>, session_id: &str) -> Option<RecoveryAttempt> {
141 if crate::upgrade::is_draining() {
142 return None;
143 }
144 let mut state = self
145 .state
146 .lock()
147 .unwrap_or_else(std::sync::PoisonError::into_inner);
148 if state.closed
149 || state.busy.contains_key(session_id)
150 || state.reservations.contains_key(session_id)
151 {
152 return None;
153 }
154 let cancelled = Arc::new(AtomicBool::new(false));
155 state.busy.insert(session_id.to_owned(), cancelled.clone());
156 self.publish_busy(&state);
157 Some(RecoveryAttempt(Arc::new(RecoveryAdmission {
158 gate: self.clone(),
159 session_id: session_id.to_owned(),
160 cancelled,
161 })))
162 }
163
164 fn finish(&self, session_id: &str, identity: &Arc<AtomicBool>) {
165 let mut state = self
166 .state
167 .lock()
168 .unwrap_or_else(std::sync::PoisonError::into_inner);
169 if state
170 .busy
171 .get(session_id)
172 .is_some_and(|current| Arc::ptr_eq(current, identity))
173 {
174 state.busy.remove(session_id);
175 self.publish_busy(&state);
176 }
177 }
178
179 fn publish_busy(&self, state: &RecoveryGateState) {
180 self.busy.send_replace(state.busy.keys().cloned().collect());
181 }
182
183 pub async fn closed(&self) {
184 loop {
185 let notified = self.closed.notified();
186 tokio::pin!(notified);
187 notified.as_mut().enable();
188 if self
189 .state
190 .lock()
191 .unwrap_or_else(std::sync::PoisonError::into_inner)
192 .closed
193 {
194 return;
195 }
196 notified.await;
197 }
198 }
199
200 pub fn subscribe(&self) -> watch::Receiver<BTreeSet<String>> {
201 self.busy.subscribe()
202 }
203
204 pub fn is_busy(&self, session_id: &str) -> bool {
205 self.state
206 .lock()
207 .unwrap_or_else(|error| error.into_inner())
208 .busy
209 .contains_key(session_id)
210 }
211
212 pub fn cancel_busy(&self, session_id: &str) {
214 if let Some(cancelled) = self
215 .state
216 .lock()
217 .unwrap_or_else(|error| error.into_inner())
218 .busy
219 .get(session_id)
220 {
221 cancelled.store(true, Ordering::Release);
222 }
223 }
224
225 pub fn close(&self) {
228 let mut state = self
229 .state
230 .lock()
231 .unwrap_or_else(std::sync::PoisonError::into_inner);
232 state.closed = true;
233 for cancelled in state.busy.values() {
234 cancelled.store(true, Ordering::Release);
235 }
236 self.closed.notify_waiters();
237 }
238
239 pub fn busy_sessions(&self) -> BTreeSet<String> {
240 self.state
241 .lock()
242 .unwrap_or_else(|error| error.into_inner())
243 .busy
244 .keys()
245 .cloned()
246 .collect()
247 }
248}
249
250#[derive(Clone)]
253pub struct RecoveryAttempt(Arc<RecoveryAdmission>);
254
255struct RecoveryAdmission {
256 gate: Arc<RecoveryGate>,
257 session_id: String,
258 cancelled: Arc<AtomicBool>,
259}
260
261impl RecoveryAttempt {
262 pub fn cancellation(&self) -> Arc<AtomicBool> {
263 self.0.cancelled.clone()
264 }
265
266 pub(crate) async fn run_blocking<T: Send + 'static>(
270 self,
271 work: impl FnOnce(Arc<AtomicBool>) -> T + Send + 'static,
272 ) -> (Result<T, tokio::task::JoinError>, Self) {
273 let executing = self.clone();
274 let result = tokio::task::spawn_blocking(move || work(executing.cancellation())).await;
275 (result, self)
276 }
277}
278
279impl std::ops::Deref for RecoveryAttempt {
280 type Target = AtomicBool;
281 fn deref(&self) -> &AtomicBool {
282 &self.0.cancelled
283 }
284}
285
286impl Drop for RecoveryAdmission {
287 fn drop(&mut self) {
288 self.gate.finish(&self.session_id, &self.cancelled);
289 }
290}
291
292#[derive(Clone)]
295pub(crate) struct ObservationSender<T> {
296 pending: Arc<Mutex<PendingObservations<T>>>,
297 wake: mpsc::Sender<()>,
298}
299
300pub(crate) struct ObservationReceiver<T> {
301 pending: Arc<Mutex<PendingObservations<T>>>,
302 wake: mpsc::Receiver<()>,
303}
304
305struct PendingObservations<T> {
306 values: BTreeMap<String, T>,
307 ready: VecDeque<String>,
308}
309
310pub(crate) fn observation_channel<T>() -> (ObservationSender<T>, ObservationReceiver<T>) {
311 let pending = Arc::new(Mutex::new(PendingObservations {
312 values: BTreeMap::new(),
313 ready: VecDeque::new(),
314 }));
315 let (tx, rx) = mpsc::channel(1);
316 (
317 ObservationSender {
318 pending: pending.clone(),
319 wake: tx,
320 },
321 ObservationReceiver { pending, wake: rx },
322 )
323}
324
325impl<T> ObservationSender<T> {
326 pub(crate) fn send(
327 &self,
328 session_id: String,
329 mut observation: T,
330 merge: impl FnOnce(&T, &mut T),
331 ) {
332 if self.wake.is_closed() {
333 return;
334 }
335 let mut pending = self
336 .pending
337 .lock()
338 .unwrap_or_else(std::sync::PoisonError::into_inner);
339 if let Some(previous) = pending.values.get(&session_id) {
340 merge(previous, &mut observation);
341 } else {
342 pending.ready.push_back(session_id.clone());
343 }
344 pending.values.insert(session_id, observation);
345 let _ = self.wake.try_send(());
347 }
348}
349
350impl<T> ObservationReceiver<T> {
351 pub(crate) fn try_recv(&mut self) -> Option<T> {
352 let mut pending = self
353 .pending
354 .lock()
355 .unwrap_or_else(std::sync::PoisonError::into_inner);
356 let session = pending.ready.pop_front()?;
357 pending.values.remove(&session)
358 }
359
360 pub(crate) async fn recv(&mut self) -> Option<T> {
361 tokio::task::yield_now().await;
362 loop {
363 if let Some(observation) = self.try_recv() {
364 return Some(observation);
365 }
366 self.wake.recv().await?;
367 }
368 }
369}
370
371impl RecoveryObserver {
372 pub fn observe(&self, observation: RecoveryObservation) {
375 let pending = PendingRecoveryObservation {
376 observed_wait: observation.checkpoint_wait.is_some(),
377 observation,
378 };
379 self.observations.send(
380 pending.observation.session.id.clone(),
381 pending,
382 |previous, next| {
383 next.observed_wait |= previous.observed_wait;
384 next.observation.latest_completed_turn_ordinal = next
385 .observation
386 .latest_completed_turn_ordinal
387 .max(previous.observation.latest_completed_turn_ordinal);
388 },
389 );
390 }
391
392 pub fn is_busy(&self, session_id: &str) -> bool {
393 self.gate.is_busy(session_id)
394 }
395
396 pub fn reserve(&self, session_id: &str) -> RecoveryReservation {
402 self.gate.reserve(session_id)
403 }
404
405 pub fn cancel_busy(&self, session_id: &str) {
409 self.gate.cancel_busy(session_id);
410 }
411}
412
413#[cfg(test)]
414mod tests {
415 use super::*;
416
417 #[tokio::test]
418 async fn lifecycle_reservation_preempts_background_work_before_releasing_admission() {
419 let gate = Arc::new(RecoveryGate::default());
420 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
421 let task = tokio::spawn({
422 let gate = gate.clone();
423 async move {
424 gate.run_background("child", async {
425 started_tx.send(()).unwrap();
426 std::future::pending::<()>().await;
427 })
428 .await
429 }
430 });
431 started_rx.await.unwrap();
432 let reservation = gate.reserve("child");
433 assert!(gate.is_busy("child"));
434 gate.cancel_busy("child");
435 assert!(
436 tokio::time::timeout(std::time::Duration::from_secs(1), task)
437 .await
438 .unwrap()
439 .unwrap()
440 .is_none()
441 );
442 assert!(!gate.is_busy("child"));
443 assert!(
444 gate.run_background("child", async { panic!("reserved worker accessed") })
445 .await
446 .is_none()
447 );
448 assert_eq!(gate.run_background("other", async { 7 }).await, Some(7));
449 drop(reservation);
450 assert_eq!(gate.run_background("child", async { 9 }).await, Some(9));
451 }
452
453 #[tokio::test]
454 async fn aborting_background_task_releases_worker_admission() {
455 let gate = Arc::new(RecoveryGate::default());
456 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
457 let task = tokio::spawn({
458 let gate = gate.clone();
459 async move {
460 gate.run_background("child", async {
461 started_tx.send(()).unwrap();
462 std::future::pending::<()>().await;
463 })
464 .await
465 }
466 });
467 started_rx.await.unwrap();
468 assert!(gate.run_background("child", async { 1 }).await.is_none());
469 task.abort();
470 assert!(task.await.unwrap_err().is_cancelled());
471 assert_eq!(gate.run_background("child", async { 2 }).await, Some(2));
472 }
473 #[test]
474 fn closing_gate_cancels_admitted_work_and_refuses_late_admission() {
475 for _ in 0..100 {
476 let gate = Arc::new(RecoveryGate::default());
477 let worker_gate = gate.clone();
478 let worker = std::thread::spawn(move || worker_gate.try_start("session"));
479 gate.close();
480 if let Some(attempt) = worker.join().unwrap() {
481 assert!(attempt.load(Ordering::Acquire));
482 assert!(gate.is_busy("session"));
483 drop(attempt);
484 }
485 assert!(gate.try_start("late").is_none());
486 assert!(gate.busy_sessions().is_empty());
487 }
488 }
489
490 #[test]
491 fn old_attempt_identity_cannot_release_a_replacement() {
492 let gate = Arc::new(RecoveryGate::default());
493 let old = gate.try_start("session").unwrap();
494 let identity = old.cancellation();
495 drop(old);
496 let replacement = gate.try_start("session").unwrap();
497 gate.finish("session", &identity);
498 assert!(gate.is_busy("session"));
499 drop(replacement);
500 assert!(!gate.is_busy("session"));
501 }
502
503 #[test]
504 fn busy_publication_tracks_concurrent_session_transitions() {
505 let gate = Arc::new(RecoveryGate::default());
506 let view = gate.subscribe();
507 std::thread::scope(|scope| {
508 for session in 0..8 {
509 let gate = gate.clone();
510 scope.spawn(move || {
511 for _ in 0..1000 {
512 drop(gate.try_start(&session.to_string()).unwrap());
513 }
514 });
515 }
516 });
517 assert_eq!(*view.borrow(), gate.busy_sessions());
518 assert!(view.borrow().is_empty());
519 }
520
521 #[tokio::test]
522 async fn aborting_a_blocking_waiter_keeps_admission_until_executor_settles() {
523 let gate = Arc::new(RecoveryGate::default());
524 let attempt = gate.try_start("session").unwrap();
525 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
526 let (finish_tx, finish_rx) = std::sync::mpsc::channel();
527 let waiter = tokio::spawn(attempt.run_blocking(move |_| {
528 started_tx.send(()).unwrap();
529 finish_rx.recv().unwrap();
530 }));
531 started_rx.await.unwrap();
532 waiter.abort();
533 assert!(matches!(waiter.await, Err(error) if error.is_cancelled()));
534 assert!(gate.is_busy("session"));
535 gate.close();
536 finish_tx.send(()).unwrap();
537 let mut busy = gate.subscribe();
538 tokio::time::timeout(
539 std::time::Duration::from_secs(5),
540 busy.wait_for(|sessions| sessions.is_empty()),
541 )
542 .await
543 .unwrap()
544 .unwrap();
545 }
546
547 #[tokio::test]
548 async fn panicking_executor_retains_admission_until_failure_is_settled() {
549 let gate = Arc::new(RecoveryGate::default());
550 let attempt = gate.try_start("session").unwrap();
551 let (result, settlement) = attempt
552 .run_blocking::<()>(|_| panic!("executor failed"))
553 .await;
554 assert!(result.unwrap_err().is_panic());
555 assert!(gate.try_start("session").is_none());
556 drop(settlement);
557 assert!(gate.try_start("session").is_some());
558 }
559
560 #[tokio::test]
561 async fn coalesced_observations_keep_independent_sessions_fair() {
562 let (sender, mut receiver) = observation_channel();
563 for value in 0..100_000 {
564 sender.send("a".into(), value, |_, _| {});
565 }
566 sender.send("b".into(), 7, |_, _| {});
567 assert_eq!(receiver.recv().await, Some(99_999));
568 sender.send("a".into(), 100_000, |_, _| {});
569 assert_eq!(receiver.recv().await, Some(7));
570 assert_eq!(receiver.recv().await, Some(100_000));
571 assert!(receiver.try_recv().is_none());
572 drop(sender);
573 assert_eq!(receiver.recv().await, None);
574 }
575}