Skip to main content

mj_controller/
recovery_gate.rs

1//! Per-session coordination between foreground operations and background recovery.
2use 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};
7/// Validate a queued placement after admission. A lifecycle may have finished
8/// between the observation and this attempt, so the old target cannot be used.
9pub(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/// Reports session activity to the recovery coordinator.
23///
24/// Reporting retains only the newest observation per session. Completed-turn
25/// frontiers are merged monotonically, so coalescing does not lose the boundary
26/// that makes a recovery copy due.
27///
28/// A caller that must know no copy can start uses [`RecoveryObserver::reserve`]
29/// rather than the queue: the reservation blocks a copy from starting whether
30/// or not queued observations have been read yet.
31#[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    // A busy-to-idle edge releases a deferral even when both observations are
41    // folded into one pending entry before the policy consumes them.
42    pub observed_wait: bool,
43}
44
45/// A per-session reservation held by a foreground lifecycle operation. The
46/// coordinator cannot start another recovery copy until this value is dropped.
47pub 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
58/// The one slot per session that background work has to hold.
59///
60/// It is shared rather than per-coordinator: a recovery copy and a worker
61/// upgrade both act on a session's live worker, so only one of them may run at
62/// a time, and a foreground lifecycle operation preempts whichever it is.
63pub struct RecoveryGate {
64    state: Mutex<RecoveryGateState>,
65    closed: Notify,
66    /// Which sessions are busy, for waiters. Published from inside the gate so
67    /// every holder updates it, whatever started the work.
68    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    /// In-flight copies, each with the cancel flag its executor watches, so a
85    /// foreground lifecycle operation can preempt one instead of waiting.
86    busy: BTreeMap<String, Arc<AtomicBool>>,
87    reservations: BTreeMap<String, usize>,
88}
89
90impl RecoveryGate {
91    /// Run disposable background work under the same admission as recovery
92    /// and worker replacement. Lifecycle reservations preempt it before they
93    /// touch the worker. Dropping the future also releases admission.
94    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    /// Claims the session with an identity-bearing guard whose cancellation
134    /// flag the executor watches, or `None` when work or a reservation already
135    /// holds it, or a daemon upgrade is waiting.
136    ///
137    /// The work does not hold up a daemon upgrade: it is cancelled when the
138    /// daemon exits, and the next daemon starts it again. Starting it while a
139    /// handoff waits would only waste it.
140    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    /// Asks the in-flight copy for this session, if any, to stop.
213    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    /// Close admission and cancel every admitted attempt in one transition.
226    /// Existing guards retain ownership until their executors have settled.
227    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/// The cancel flag also identifies this particular admission. Only its owner
251/// can release the slot, including when an executor or coordinator unwinds.
252#[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    /// Both the waiter and blocking executor own admission. Aborting the waiter
267    /// cannot release a running executor; a panicking executor cannot release
268    /// its slot before the supervisor applies the failure outcome.
269    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/// Per-session latest-state mailbox shared by the two background policies.
293/// A capacity-one wake channel never queues an observation history.
294#[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        // Full means a wake is already pending. Closure discards disposable observations.
346        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    /// Queues one observation for the coordinator. Returns as soon as the
373    /// observation is queued; a stopped coordinator makes this a no-op.
374    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    /// Holds off any recovery copy for this session until the returned
397    /// reservation is dropped. This, not the observation queue, is what a
398    /// lifecycle operation relies on: queued observations may still be
399    /// unread, and the coordinator refuses to start a copy for a reserved
400    /// session whenever it reads them.
401    pub fn reserve(&self, session_id: &str) -> RecoveryReservation {
402        self.gate.reserve(session_id)
403    }
404
405    /// Asks an in-flight recovery copy for this session to stop. A foreground
406    /// lifecycle operation calls this after reserving so it preempts the copy
407    /// instead of waiting behind it.
408    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}