use std::collections::BTreeMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use chrono::{DateTime, Utc};
use tokio::sync::mpsc;
use crate::hel_controller::{Controller, WorkerUpgradeOutcome};
use crate::hel_recovery::{backoff_delay, elapsed_at_least};
use crate::hel_session_manager::SessionManagerControl;
use hel::hel_config::HelConfig;
use hel::hel_state::{HelState, RecoveryGate, RecoveryObserver, SessionRecord, SessionState};
use hel::hel_targets::CancellableProcessExecutor;
const WORKER_UPGRADE_RETRY_INTERVAL: Duration = Duration::from_secs(10 * 60);
const MAX_WORKER_UPGRADE_RETRY_INTERVAL: Duration = Duration::from_secs(2 * 60 * 60);
const INSTALLED_BUILD_FRESHNESS: Duration = Duration::from_secs(10 * 60);
const WORKER_UPGRADE_TIMEOUT: Duration = Duration::from_secs(15 * 60);
#[derive(Debug, Clone)]
pub struct WorkerUpgradeObservation {
pub session: SessionRecord,
pub config: HelConfig,
pub worker_build: Option<String>,
pub quiet: bool,
}
#[derive(Clone)]
pub struct WorkerUpgradeObserver {
observations: mpsc::UnboundedSender<WorkerUpgradeObservation>,
}
impl WorkerUpgradeObserver {
pub fn observe(&self, observation: WorkerUpgradeObservation) {
let session_id = observation.session.id.clone();
if let Err(error) = self.observations.send(observation) {
tracing::debug!(
%session_id,
%error,
"worker upgrade observation dropped because the coordinator stopped"
);
}
}
}
#[derive(Debug, Clone)]
pub struct WorkerUpgradeResult {
pub session_id: String,
pub outcome: Result<WorkerUpgradeOutcome, String>,
pub cancelled: bool,
}
pub struct WorkerUpgradeCoordinator {
observer: WorkerUpgradeObserver,
results: mpsc::UnboundedReceiver<WorkerUpgradeResult>,
cancelled: Arc<AtomicBool>,
gate: Arc<RecoveryGate>,
}
impl Drop for WorkerUpgradeCoordinator {
fn drop(&mut self) {
self.cancelled.store(true, Ordering::Release);
self.gate.cancel_all();
}
}
impl WorkerUpgradeCoordinator {
pub fn spawn(session_manager: SessionManagerControl, recovery: &RecoveryObserver) -> Self {
let (observations_tx, mut observations_rx) =
mpsc::unbounded_channel::<WorkerUpgradeObservation>();
let (completed_tx, mut completed_rx) = mpsc::unbounded_channel::<WorkerUpgradeResult>();
let (results_tx, results_rx) = mpsc::unbounded_channel();
let gate = recovery.gate.clone();
let coordinator_gate = gate.clone();
let cancelled = Arc::new(AtomicBool::new(false));
let coordinator_cancelled = cancelled.clone();
tokio::spawn(async move {
let mut policies = BTreeMap::<String, PolicyState>::new();
loop {
tokio::select! {
observed = observations_rx.recv() => {
let Some(observation) = observed else { break };
if coordinator_cancelled.load(Ordering::Acquire) {
break;
}
let session_id = observation.session.id.clone();
let policy = policies.entry(session_id.clone()).or_default();
policy.observe(&observation);
if !policy.due(&observation, Utc::now()) {
continue;
}
let Some(upgrade_cancelled) = coordinator_gate.try_start(&session_id)
else {
continue;
};
policy.attempt_started();
let completed_tx = completed_tx.clone();
let session_manager = session_manager.clone();
let task_cancelled = upgrade_cancelled.clone();
let handle = tokio::runtime::Handle::current();
let task_session_id = session_id.clone();
tokio::spawn(async move {
let joined = tokio::task::spawn_blocking(move || {
let mut state = HelState::default();
state
.sessions
.insert(task_session_id.clone(), observation.session);
let controller = Controller {
config: observation.config,
state,
};
let executor = CancellableProcessExecutor::new(task_cancelled)
.with_deadline(WORKER_UPGRADE_TIMEOUT);
handle
.block_on(controller.upgrade_session_worker(
&task_session_id,
&executor,
&session_manager,
observation.worker_build.as_deref(),
))
.map_err(|error| format!("{error:#}"))
})
.await;
let outcome = match joined {
Ok(outcome) => outcome,
Err(error) => Err(format!("worker upgrade task failed: {error}")),
};
let result = WorkerUpgradeResult {
session_id,
outcome,
cancelled: upgrade_cancelled.load(Ordering::Acquire),
};
let result_session_id = result.session_id.clone();
if let Err(error) = completed_tx.send(result) {
tracing::debug!(
session_id = %result_session_id,
%error,
"worker upgrade result dropped because the coordinator stopped"
);
}
});
}
completed = completed_rx.recv() => {
let Some(result) = completed else { break };
coordinator_gate.finish(&result.session_id);
let policy = policies.entry(result.session_id.clone()).or_default();
policy.record(&result, Utc::now());
let result_session_id = result.session_id.clone();
if let Err(error) = results_tx.send(result) {
tracing::debug!(
session_id = %result_session_id,
%error,
"worker upgrade result dropped because its consumer stopped"
);
}
}
}
}
});
Self {
observer: WorkerUpgradeObserver {
observations: observations_tx,
},
results: results_rx,
cancelled,
gate,
}
}
pub fn observer(&self) -> WorkerUpgradeObserver {
self.observer.clone()
}
pub fn try_result(&mut self) -> Option<WorkerUpgradeResult> {
self.results.try_recv().ok()
}
}
#[derive(Debug, Default, PartialEq, Eq)]
struct PolicyState {
current_build: Option<String>,
current_build_at: Option<DateTime<Utc>>,
attempt_in_flight: bool,
failed_at: Option<DateTime<Utc>>,
consecutive_failures: u32,
}
impl PolicyState {
fn due(&self, observation: &WorkerUpgradeObservation, now: DateTime<Utc>) -> bool {
if !observation.quiet
|| observation.session.state != SessionState::Running
|| self.attempt_in_flight
{
return false;
}
if self.worker_is_known_current(observation.worker_build.as_deref(), now) {
return false;
}
self.failed_at.is_none_or(|failed_at| {
elapsed_at_least(
failed_at,
now,
backoff_delay(
WORKER_UPGRADE_RETRY_INTERVAL,
MAX_WORKER_UPGRADE_RETRY_INTERVAL,
self.consecutive_failures,
),
)
})
}
fn worker_is_known_current(&self, worker_build: Option<&str>, now: DateTime<Utc>) -> bool {
let (Some(observed), Some(current), Some(checked_at)) = (
worker_build,
self.current_build.as_deref(),
self.current_build_at,
) else {
return false;
};
observed == current && !elapsed_at_least(checked_at, now, INSTALLED_BUILD_FRESHNESS)
}
fn observe(&mut self, observation: &WorkerUpgradeObservation) {
if self.current_build.is_some()
&& observation.worker_build.as_deref() == self.current_build.as_deref()
{
self.failed_at = None;
self.consecutive_failures = 0;
}
}
fn attempt_started(&mut self) {
self.attempt_in_flight = true;
}
fn record(&mut self, result: &WorkerUpgradeResult, now: DateTime<Utc>) {
self.attempt_in_flight = false;
if result.cancelled {
return;
}
match &result.outcome {
Ok(WorkerUpgradeOutcome::Deferred) => {
self.failed_at = None;
self.consecutive_failures = 0;
}
Ok(outcome @ WorkerUpgradeOutcome::AlreadyCurrent { .. }) => {
self.failed_at = None;
self.consecutive_failures = 0;
self.current_build = outcome.build().map(str::to_owned);
self.current_build_at = Some(now);
}
Ok(outcome @ WorkerUpgradeOutcome::Upgraded { .. }) => {
self.current_build = outcome.build().map(str::to_owned);
self.current_build_at = Some(now);
self.consecutive_failures = self.consecutive_failures.saturating_add(1);
self.failed_at = Some(now);
}
Err(_) => {
self.consecutive_failures = self.consecutive_failures.saturating_add(1);
self.failed_at = Some(now);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn session_record(state: SessionState) -> SessionRecord {
SessionRecord {
workspace_id: hel::hel_workspace::DEFAULT_WORKSPACE_ID.to_owned(),
archived: false,
container_cpus: None,
container_memory: None,
id: "session-1".to_owned(),
title: "work".into(),
harness_kind: hel::hel_config::HarnessKind::Codex,
last_profile: "codex-1".into(),
bundle_id: "hel".into(),
project_directory: None,
managed_worktree: None,
target_template_id: "podman".into(),
resource_allocation: None,
additional_mounts: Vec::new(),
state,
target: None,
native_session_id: None,
acp_session_title: None,
session_title_override: None,
created_at: "2026-08-09T12:00:00Z".into(),
updated_at: "2026-08-09T12:01:00Z".into(),
viewed_through_event_ordinal: 0,
draft_input: String::new(),
last_error: None,
last_checkpoint_error: None,
checkpoint: None,
}
}
fn observation(worker_build: Option<&str>, quiet: bool) -> WorkerUpgradeObservation {
WorkerUpgradeObservation {
session: session_record(SessionState::Running),
config: HelConfig::default(),
worker_build: worker_build.map(str::to_owned),
quiet,
}
}
fn failure(detail: &str) -> WorkerUpgradeResult {
WorkerUpgradeResult {
session_id: "session-1".into(),
outcome: Err(detail.into()),
cancelled: false,
}
}
fn current(build: &str) -> WorkerUpgradeOutcome {
WorkerUpgradeOutcome::AlreadyCurrent {
build: build.to_owned(),
}
}
fn upgraded(build: &str) -> WorkerUpgradeOutcome {
WorkerUpgradeOutcome::Upgraded {
build: build.to_owned(),
}
}
fn success(outcome: WorkerUpgradeOutcome) -> WorkerUpgradeResult {
WorkerUpgradeResult {
session_id: "session-1".into(),
outcome: Ok(outcome),
cancelled: false,
}
}
#[test]
fn only_a_quiet_session_with_an_unknown_build_is_due() {
let now = Utc::now();
let policy = PolicyState::default();
assert!(policy.due(&observation(Some("build-a"), true), now));
assert!(
!policy.due(&observation(Some("build-a"), false), now),
"a working session must not have its worker killed"
);
assert!(
policy.due(&observation(None, true), now),
"a worker too old to report a build is outdated"
);
}
#[test]
fn a_busy_turn_is_never_upgraded_no_matter_how_long_it_runs() {
let started = Utc::now();
let policy = PolicyState::default();
let two_days_later = started + chrono::Duration::days(2);
assert!(!policy.due(&observation(Some("old-build"), false), two_days_later));
assert!(
policy.due(&observation(Some("old-build"), true), two_days_later),
"the next quiet observation may upgrade without an age-based busy timeout"
);
}
#[test]
fn only_a_running_session_is_due() {
let now = Utc::now();
let policy = PolicyState::default();
for state in [
SessionState::Provisioning,
SessionState::Disconnected,
SessionState::Checkpointing,
SessionState::Closing,
SessionState::Destroying,
SessionState::Stopped,
SessionState::Lost,
SessionState::Error,
SessionState::DestroyedWithDataLoss,
] {
let mut observation = observation(Some("build-a"), true);
observation.session.state = state;
assert!(!policy.due(&observation, now), "{state:?}");
}
}
#[test]
fn an_attempt_in_flight_suppresses_further_observations() {
let now = Utc::now();
let mut policy = PolicyState::default();
policy.attempt_started();
assert!(!policy.due(&observation(Some("build-a"), true), now));
}
#[test]
fn a_worker_proved_current_stops_being_probed_until_the_proof_goes_stale() {
let now = Utc::now();
let mut policy = PolicyState::default();
policy.record(&success(current("build-a")), now);
assert!(!policy.due(&observation(Some("build-a"), true), now));
assert!(
policy.due(&observation(Some("build-b"), true), now),
"a different build is outdated however recently the last one was checked"
);
let stale = now + chrono::Duration::from_std(INSTALLED_BUILD_FRESHNESS).unwrap();
assert!(policy.due(&observation(Some("build-a"), true), stale));
}
#[test]
fn a_failed_upgrade_backs_off_and_widens() {
let now = Utc::now();
let mut policy = PolicyState::default();
policy.attempt_started();
policy.record(&failure("install the current Mjolnir worker binary"), now);
let interval = chrono::Duration::from_std(WORKER_UPGRADE_RETRY_INTERVAL).unwrap();
assert!(!policy.due(&observation(Some("build-a"), true), now));
assert!(!policy.due(
&observation(Some("build-a"), true),
now + interval - chrono::Duration::seconds(1)
));
assert!(policy.due(&observation(Some("build-a"), true), now + interval));
policy.attempt_started();
policy.record(&failure("install the current Mjolnir worker binary"), now);
assert!(!policy.due(
&observation(Some("build-a"), true),
now + interval * 2 - chrono::Duration::seconds(1)
));
assert!(policy.due(&observation(Some("build-a"), true), now + interval * 2));
}
#[test]
fn a_successful_upgrade_stops_the_probing_and_confirming_it_clears_the_backoff() {
let now = Utc::now();
let mut policy = PolicyState::default();
policy.attempt_started();
policy.record(&failure("install the current Mjolnir worker binary"), now);
policy.attempt_started();
policy.record(&success(upgraded("build-b")), now);
let confirmed = observation(Some("build-b"), true);
policy.observe(&confirmed);
assert_eq!(policy.consecutive_failures, 0);
assert_eq!(policy.failed_at, None);
assert!(
!policy.due(&confirmed, now),
"the worker now runs the installed build, so nothing is due"
);
}
#[test]
fn an_upgrade_that_does_not_take_backs_off_instead_of_looping() {
let now = Utc::now();
let mut policy = PolicyState::default();
policy.attempt_started();
policy.record(&success(upgraded("build-b")), now);
let unchanged = observation(Some("build-a"), true);
policy.observe(&unchanged);
let interval = chrono::Duration::from_std(WORKER_UPGRADE_RETRY_INTERVAL).unwrap();
assert!(!policy.due(&unchanged, now));
assert!(policy.due(&unchanged, now + interval));
policy.attempt_started();
policy.record(&success(upgraded("build-b")), now + interval);
policy.observe(&unchanged);
assert!(!policy.due(&unchanged, now + interval * 2));
assert!(policy.due(&unchanged, now + interval * 3));
}
#[test]
fn a_deferred_upgrade_is_retried_at_the_next_quiet_observation() {
let now = Utc::now();
let mut policy = PolicyState::default();
policy.attempt_started();
policy.record(&success(WorkerUpgradeOutcome::Deferred), now);
assert!(policy.due(&observation(Some("build-a"), true), now));
}
#[test]
fn a_preempted_attempt_is_neither_a_success_nor_a_failure() {
let now = Utc::now();
let mut policy = PolicyState::default();
policy.attempt_started();
policy.record(
&WorkerUpgradeResult {
session_id: "session-1".into(),
outcome: Err("operation cancelled".into()),
cancelled: true,
},
now,
);
assert_eq!(policy.consecutive_failures, 0);
assert!(policy.due(&observation(Some("build-a"), true), now));
}
#[test]
fn the_shared_gate_keeps_an_upgrade_and_a_recovery_copy_apart() {
let gate = Arc::new(RecoveryGate::default());
let recovery_copy = gate.try_start("session-1").expect("the slot starts free");
assert!(gate.try_start("session-1").is_none());
gate.finish("session-1");
assert!(gate.try_start("session-1").is_some());
drop(recovery_copy);
}
#[test]
fn observing_hands_off_without_waiting() {
let (observations, mut queued) = mpsc::unbounded_channel();
let observer = WorkerUpgradeObserver { observations };
for _ in 0..64 {
observer.observe(observation(Some("build-a"), true));
}
let received = std::iter::from_fn(|| queued.try_recv().ok()).count();
assert_eq!(received, 64);
}
#[test]
fn observing_a_stopped_coordinator_is_a_no_op() {
let (observations, queued) = mpsc::unbounded_channel();
let observer = WorkerUpgradeObserver { observations };
drop(queued);
observer.observe(observation(Some("build-a"), true));
}
}