1use std::collections::BTreeMap;
13use std::sync::Arc;
14use std::sync::atomic::{AtomicBool, Ordering};
15use std::time::Duration;
16
17use chrono::{DateTime, Utc};
18use tokio::sync::mpsc;
19
20use crate::controller::{Controller, WorkerUpgradeOutcome};
21use crate::recovery::{backoff_delay, elapsed_at_least};
22use crate::recovery_gate::{RecoveryGate, RecoveryObserver};
23use crate::session_manager::SessionManagerControl;
24use crate::targets::CancellableProcessExecutor;
25use mj_core::config::Config;
26use mj_core::state::{SessionRecord, SessionState, State};
27
28const WORKER_UPGRADE_RETRY_INTERVAL: Duration = Duration::from_secs(10 * 60);
32
33const MAX_WORKER_UPGRADE_RETRY_INTERVAL: Duration = Duration::from_secs(2 * 60 * 60);
36
37const WORKER_UPGRADE_TIMEOUT: Duration = Duration::from_secs(15 * 60);
41
42#[derive(Debug, Clone)]
44pub struct WorkerUpgradeObservation {
45 pub session: SessionRecord,
46 pub config: Config,
47 pub worker_build: Option<String>,
50 pub quiet: bool,
57}
58
59#[derive(Clone)]
64pub struct WorkerUpgradeObserver {
65 observations: mpsc::UnboundedSender<WorkerUpgradeObservation>,
66}
67
68impl WorkerUpgradeObserver {
69 pub fn observe(&self, observation: WorkerUpgradeObservation) {
70 let session_id = observation.session.id.clone();
71 if let Err(error) = self.observations.send(observation) {
72 tracing::debug!(
73 %session_id,
74 %error,
75 "worker upgrade observation dropped because the coordinator stopped"
76 );
77 }
78 }
79}
80
81#[derive(Debug, Clone)]
82pub struct WorkerUpgradeResult {
83 pub session_id: String,
84 pub outcome: Result<WorkerUpgradeOutcome, String>,
85 pub cancelled: bool,
88}
89
90pub struct WorkerUpgradeCoordinator {
91 observer: WorkerUpgradeObserver,
92 results: mpsc::UnboundedReceiver<WorkerUpgradeResult>,
93 cancelled: Arc<AtomicBool>,
94 gate: Arc<RecoveryGate>,
95}
96
97impl Drop for WorkerUpgradeCoordinator {
98 fn drop(&mut self) {
99 self.cancelled.store(true, Ordering::Release);
104 self.gate.cancel_all();
105 }
106}
107
108impl WorkerUpgradeCoordinator {
109 pub fn spawn(session_manager: SessionManagerControl, recovery: &RecoveryObserver) -> Self {
112 let (observations_tx, mut observations_rx) =
113 mpsc::unbounded_channel::<WorkerUpgradeObservation>();
114 let (completed_tx, mut completed_rx) = mpsc::unbounded_channel::<WorkerUpgradeResult>();
115 let (results_tx, results_rx) = mpsc::unbounded_channel();
116 let gate = recovery.gate.clone();
117 let coordinator_gate = gate.clone();
118 let cancelled = Arc::new(AtomicBool::new(false));
119 let coordinator_cancelled = cancelled.clone();
120 tokio::spawn(async move {
121 let mut policies = BTreeMap::<String, PolicyState>::new();
122 loop {
123 tokio::select! {
124 observed = observations_rx.recv() => {
125 let Some(observation) = observed else { break };
126 if coordinator_cancelled.load(Ordering::Acquire) {
127 break;
128 }
129 let session_id = observation.session.id.clone();
130 let policy = policies.entry(session_id.clone()).or_default();
131 policy.observe(&observation);
132 if !policy.due(&observation, Utc::now()) {
133 continue;
134 }
135 let Some(upgrade_cancelled) = coordinator_gate.try_start(&session_id)
136 else {
137 continue;
138 };
139 policy.attempt_started();
140 let completed_tx = completed_tx.clone();
141 let session_manager = session_manager.clone();
142 let task_cancelled = upgrade_cancelled.clone();
143 let task_session_id = session_id.clone();
144 tokio::spawn(async move {
145 let joined = tokio::task::spawn_blocking(move || {
146 let mut state = State::default();
147 state
148 .sessions
149 .insert(task_session_id.clone(), observation.session);
150 let controller = Controller {
151 config: observation.config,
152 state,
153 };
154 let executor = CancellableProcessExecutor::new(task_cancelled)
155 .with_deadline(WORKER_UPGRADE_TIMEOUT);
156 mj_core::runtime::block_on(controller.upgrade_session_worker(
157 &task_session_id,
158 &executor,
159 &session_manager,
160 observation.worker_build.as_deref(),
161 ))
162 .and_then(|result| result)
163 .map_err(|error| format!("{error:#}"))
164 })
165 .await;
166 let outcome = match joined {
167 Ok(outcome) => outcome,
168 Err(error) => Err(format!("worker upgrade task failed: {error}")),
169 };
170 let result = WorkerUpgradeResult {
171 session_id,
172 outcome,
173 cancelled: upgrade_cancelled.load(Ordering::Acquire),
174 };
175 let result_session_id = result.session_id.clone();
176 if let Err(error) = completed_tx.send(result) {
177 tracing::debug!(
178 session_id = %result_session_id,
179 %error,
180 "worker upgrade result dropped because the coordinator stopped"
181 );
182 }
183 });
184 }
185 completed = completed_rx.recv() => {
186 let Some(result) = completed else { break };
187 coordinator_gate.finish(&result.session_id);
188 let policy = policies.entry(result.session_id.clone()).or_default();
189 policy.record(&result, Utc::now());
190 let result_session_id = result.session_id.clone();
191 if let Err(error) = results_tx.send(result) {
192 tracing::debug!(
193 session_id = %result_session_id,
194 %error,
195 "worker upgrade result dropped because its consumer stopped"
196 );
197 }
198 }
199 }
200 }
201 });
202 Self {
203 observer: WorkerUpgradeObserver {
204 observations: observations_tx,
205 },
206 results: results_rx,
207 cancelled,
208 gate,
209 }
210 }
211
212 pub fn observer(&self) -> WorkerUpgradeObserver {
213 self.observer.clone()
214 }
215
216 pub fn try_result(&mut self) -> Option<WorkerUpgradeResult> {
217 self.results.try_recv().ok()
218 }
219}
220
221#[derive(Debug, Default, PartialEq, Eq)]
223struct PolicyState {
224 current_build: Option<String>,
227 attempt_in_flight: bool,
230 failed_at: Option<DateTime<Utc>>,
231 consecutive_failures: u32,
232}
233
234impl PolicyState {
235 fn due(&self, observation: &WorkerUpgradeObservation, now: DateTime<Utc>) -> bool {
237 if !observation.quiet
240 || observation.session.state != SessionState::Running
241 || self.attempt_in_flight
242 {
243 return false;
244 }
245 if self.worker_is_known_current(observation.worker_build.as_deref()) {
246 return false;
247 }
248 self.failed_at.is_none_or(|failed_at| {
249 elapsed_at_least(
250 failed_at,
251 now,
252 backoff_delay(
253 WORKER_UPGRADE_RETRY_INTERVAL,
254 MAX_WORKER_UPGRADE_RETRY_INTERVAL,
255 self.consecutive_failures,
256 ),
257 )
258 })
259 }
260
261 fn worker_is_known_current(&self, worker_build: Option<&str>) -> bool {
265 let (Some(observed), Some(current)) = (worker_build, self.current_build.as_deref()) else {
266 return false;
267 };
268 observed == current
269 }
270
271 fn observe(&mut self, observation: &WorkerUpgradeObservation) {
277 if self.current_build.is_some()
278 && observation.worker_build.as_deref() == self.current_build.as_deref()
279 {
280 self.failed_at = None;
281 self.consecutive_failures = 0;
282 }
283 }
284
285 fn attempt_started(&mut self) {
286 self.attempt_in_flight = true;
287 }
288
289 fn record(&mut self, result: &WorkerUpgradeResult, now: DateTime<Utc>) {
290 self.attempt_in_flight = false;
291 if result.cancelled {
292 return;
295 }
296 match &result.outcome {
297 Ok(WorkerUpgradeOutcome::Deferred) => {
298 self.failed_at = None;
302 self.consecutive_failures = 0;
303 }
304 Ok(outcome @ WorkerUpgradeOutcome::AlreadyCurrent { .. }) => {
305 self.failed_at = None;
306 self.consecutive_failures = 0;
307 self.current_build = outcome.build().map(str::to_owned);
308 }
309 Ok(outcome @ WorkerUpgradeOutcome::Upgraded { .. }) => {
310 self.current_build = outcome.build().map(str::to_owned);
314 self.consecutive_failures = self.consecutive_failures.saturating_add(1);
320 self.failed_at = Some(now);
321 }
322 Err(_) => {
323 self.consecutive_failures = self.consecutive_failures.saturating_add(1);
324 self.failed_at = Some(now);
325 }
326 }
327 }
328}
329
330#[cfg(test)]
331mod tests {
332 use super::*;
333
334 fn session_record(state: SessionState) -> SessionRecord {
335 SessionRecord {
336 target_runtime: None,
337 launch_base: None,
338 launch_branch: None,
339 publication: None,
340 build_cache: None,
341 container_workspace: None,
342 mjolnir_subagents: None,
343 create_managed_worktree: None,
344 workspace_id: mj_core::workspace::DEFAULT_WORKSPACE_ID.to_owned(),
345 archived: false,
346 container_cpus: None,
347 container_memory: None,
348 id: "session-1".to_owned(),
349 title: "work".into(),
350 harness_kind: mj_core::config::HarnessKind::Codex,
351 last_profile: "codex-1".into(),
352 bundle_id: "hel".into(),
353 project_directory: None,
354 managed_worktree: None,
355 target_template_id: "podman".into(),
356 resource_allocation: None,
357 additional_mounts: Vec::new(),
358 state,
359 target: None,
360 native_session_id: None,
361 acp_session_title: None,
362 session_title_override: None,
363 created_at: "2026-08-09T12:00:00Z".into(),
364 updated_at: "2026-08-09T12:01:00Z".into(),
365 viewed_through_event_ordinal: 0,
366 draft_input: String::new(),
367 last_error: None,
368 last_checkpoint_error: None,
369 checkpoint: None,
370 }
371 }
372
373 fn observation(worker_build: Option<&str>, quiet: bool) -> WorkerUpgradeObservation {
374 WorkerUpgradeObservation {
375 session: session_record(SessionState::Running),
376 config: Config::default(),
377 worker_build: worker_build.map(str::to_owned),
378 quiet,
379 }
380 }
381
382 fn failure(detail: &str) -> WorkerUpgradeResult {
383 WorkerUpgradeResult {
384 session_id: "session-1".into(),
385 outcome: Err(detail.into()),
386 cancelled: false,
387 }
388 }
389
390 fn current(build: &str) -> WorkerUpgradeOutcome {
391 WorkerUpgradeOutcome::AlreadyCurrent {
392 build: build.to_owned(),
393 }
394 }
395
396 fn upgraded(build: &str) -> WorkerUpgradeOutcome {
397 WorkerUpgradeOutcome::Upgraded {
398 build: build.to_owned(),
399 }
400 }
401
402 fn success(outcome: WorkerUpgradeOutcome) -> WorkerUpgradeResult {
403 WorkerUpgradeResult {
404 session_id: "session-1".into(),
405 outcome: Ok(outcome),
406 cancelled: false,
407 }
408 }
409
410 #[test]
412 fn only_a_quiet_session_with_an_unknown_build_is_due() {
413 let now = Utc::now();
414 let policy = PolicyState::default();
415
416 assert!(policy.due(&observation(Some("build-a"), true), now));
417 assert!(
418 !policy.due(&observation(Some("build-a"), false), now),
419 "a working session must not have its worker killed"
420 );
421 assert!(
422 policy.due(&observation(None, true), now),
423 "a worker too old to report a build is outdated"
424 );
425 }
426
427 #[test]
428 fn a_busy_turn_is_never_upgraded_no_matter_how_long_it_runs() {
429 let started = Utc::now();
430 let policy = PolicyState::default();
431 let two_days_later = started + chrono::Duration::days(2);
432
433 assert!(!policy.due(&observation(Some("old-build"), false), two_days_later));
434 assert!(
435 policy.due(&observation(Some("old-build"), true), two_days_later),
436 "the next quiet observation may upgrade without an age-based busy timeout"
437 );
438 }
439
440 #[test]
443 fn only_a_running_session_is_due() {
444 let now = Utc::now();
445 let policy = PolicyState::default();
446 for state in [
447 SessionState::Provisioning,
448 SessionState::Disconnected,
449 SessionState::Checkpointing,
450 SessionState::Closing,
451 SessionState::Destroying,
452 SessionState::Stopped,
453 SessionState::Lost,
454 SessionState::Error,
455 SessionState::DestroyedWithDataLoss,
456 ] {
457 let mut observation = observation(Some("build-a"), true);
458 observation.session.state = state;
459 assert!(!policy.due(&observation, now), "{state:?}");
460 }
461 }
462
463 #[test]
466 fn an_attempt_in_flight_suppresses_further_observations() {
467 let now = Utc::now();
468 let mut policy = PolicyState::default();
469 policy.attempt_started();
470
471 assert!(!policy.due(&observation(Some("build-a"), true), now));
472 }
473
474 #[test]
477 fn a_worker_proved_current_stays_trusted_for_coordinator_lifetime() {
478 let now = Utc::now();
479 let mut policy = PolicyState::default();
480 policy.record(&success(current("build-a")), now);
481
482 assert!(!policy.due(&observation(Some("build-a"), true), now));
483 assert!(
484 !policy.due(
485 &observation(Some("build-a"), true),
486 now + chrono::Duration::days(2)
487 ),
488 "the launched build remains trusted for the coordinator lifetime"
489 );
490 assert!(
491 policy.due(&observation(Some("build-b"), true), now),
492 "a different build is outdated however recently the last one was checked"
493 );
494 }
495
496 #[test]
499 fn a_failed_upgrade_backs_off_and_widens() {
500 let now = Utc::now();
501 let mut policy = PolicyState::default();
502 policy.attempt_started();
503 policy.record(&failure("install the current Mjolnir worker binary"), now);
504
505 let interval = chrono::Duration::from_std(WORKER_UPGRADE_RETRY_INTERVAL).unwrap();
506 assert!(!policy.due(&observation(Some("build-a"), true), now));
507 assert!(!policy.due(
508 &observation(Some("build-a"), true),
509 now + interval - chrono::Duration::seconds(1)
510 ));
511 assert!(policy.due(&observation(Some("build-a"), true), now + interval));
512
513 policy.attempt_started();
514 policy.record(&failure("install the current Mjolnir worker binary"), now);
515 assert!(!policy.due(
516 &observation(Some("build-a"), true),
517 now + interval * 2 - chrono::Duration::seconds(1)
518 ));
519 assert!(policy.due(&observation(Some("build-a"), true), now + interval * 2));
520 }
521
522 #[test]
526 fn a_successful_upgrade_stops_the_probing_and_confirming_it_clears_the_backoff() {
527 let now = Utc::now();
528 let mut policy = PolicyState::default();
529 policy.attempt_started();
530 policy.record(&failure("install the current Mjolnir worker binary"), now);
531 policy.attempt_started();
532 policy.record(&success(upgraded("build-b")), now);
533
534 let confirmed = observation(Some("build-b"), true);
535 policy.observe(&confirmed);
536 assert_eq!(policy.consecutive_failures, 0);
537 assert_eq!(policy.failed_at, None);
538 assert!(
539 !policy.due(&confirmed, now),
540 "the worker now runs the installed build, so nothing is due"
541 );
542 }
543
544 #[test]
548 fn an_upgrade_that_does_not_take_backs_off_instead_of_looping() {
549 let now = Utc::now();
550 let mut policy = PolicyState::default();
551 policy.attempt_started();
552 policy.record(&success(upgraded("build-b")), now);
553
554 let unchanged = observation(Some("build-a"), true);
557 policy.observe(&unchanged);
558 let interval = chrono::Duration::from_std(WORKER_UPGRADE_RETRY_INTERVAL).unwrap();
559 assert!(!policy.due(&unchanged, now));
560 assert!(policy.due(&unchanged, now + interval));
561
562 policy.attempt_started();
563 policy.record(&success(upgraded("build-b")), now + interval);
564 policy.observe(&unchanged);
565 assert!(!policy.due(&unchanged, now + interval * 2));
566 assert!(policy.due(&unchanged, now + interval * 3));
567 }
568
569 #[test]
572 fn a_deferred_upgrade_is_retried_at_the_next_quiet_observation() {
573 let now = Utc::now();
574 let mut policy = PolicyState::default();
575 policy.attempt_started();
576 policy.record(&success(WorkerUpgradeOutcome::Deferred), now);
577
578 assert!(policy.due(&observation(Some("build-a"), true), now));
579 }
580
581 #[test]
584 fn a_preempted_attempt_is_neither_a_success_nor_a_failure() {
585 let now = Utc::now();
586 let mut policy = PolicyState::default();
587 policy.attempt_started();
588 policy.record(
589 &WorkerUpgradeResult {
590 session_id: "session-1".into(),
591 outcome: Err("operation cancelled".into()),
592 cancelled: true,
593 },
594 now,
595 );
596
597 assert_eq!(policy.consecutive_failures, 0);
598 assert!(policy.due(&observation(Some("build-a"), true), now));
599 }
600
601 #[test]
604 fn the_shared_gate_keeps_an_upgrade_and_a_recovery_copy_apart() {
605 let gate = Arc::new(RecoveryGate::default());
606 let recovery_copy = gate.try_start("session-1").expect("the slot starts free");
607
608 assert!(gate.try_start("session-1").is_none());
609
610 gate.finish("session-1");
611 assert!(gate.try_start("session-1").is_some());
612 drop(recovery_copy);
613 }
614
615 #[test]
618 fn observing_hands_off_without_waiting() {
619 let (observations, mut queued) = mpsc::unbounded_channel();
620 let observer = WorkerUpgradeObserver { observations };
621
622 for _ in 0..64 {
623 observer.observe(observation(Some("build-a"), true));
624 }
625
626 let received = std::iter::from_fn(|| queued.try_recv().ok()).count();
627 assert_eq!(received, 64);
628 }
629
630 #[test]
632 fn observing_a_stopped_coordinator_is_a_no_op() {
633 let (observations, queued) = mpsc::unbounded_channel();
634 let observer = WorkerUpgradeObserver { observations };
635 drop(queued);
636
637 observer.observe(observation(Some("build-a"), true));
638 }
639}