diff --git a/src/openhuman/socket/medulla/mod.rs b/src/openhuman/socket/medulla/mod.rs index 4d902e9ca..1a961e4c4 100644 --- a/src/openhuman/socket/medulla/mod.rs +++ b/src/openhuman/socket/medulla/mod.rs @@ -27,8 +27,9 @@ use std::sync::atomic::{AtomicI64, Ordering}; use std::sync::{Arc, OnceLock}; use parking_lot::Mutex; -use tokio::sync::{mpsc, Notify}; -use tokio::time::Duration; +use tokio::sync::mpsc; +use tokio::time::{Duration, Instant}; +use tokio_util::sync::CancellationToken; use crate::openhuman::agent::progress::AgentProgress; use crate::openhuman::agent::turn_origin::{with_origin, AgentTurnOrigin}; @@ -60,8 +61,12 @@ pub fn manager() -> &'static Arc { /// One in-flight task: a cooperative abort signal and a steering input channel. struct RunningTask { - /// Fired by `medulla:task_abort` to cancel the in-flight turn. - abort: Arc, + /// Fired by `medulla:task_abort` to cancel the task. A + /// [`CancellationToken`] *latches*: an abort that arrives before the run + /// path starts awaiting it (e.g. while the driver is still building the + /// agent) is still observed when the turn checks it, so no cancellation is + /// lost. + abort: CancellationToken, /// Mid-task steering input (`medulla:task_send`) delivered as follow-up /// turns on the same agent session. steer_tx: mpsc::UnboundedSender, @@ -93,12 +98,12 @@ impl MedullaTaskManager { return; } - let abort = Arc::new(Notify::new()); + let abort = CancellationToken::new(); let (steer_tx, steer_rx) = mpsc::unbounded_channel::(); self.tasks.lock().insert( task_id.clone(), RunningTask { - abort: Arc::clone(&abort), + abort: abort.clone(), steer_tx, }, ); @@ -132,7 +137,7 @@ impl MedullaTaskManager { match self.tasks.lock().get(&abort.task_id) { Some(task) => { log::info!("[medulla] aborting task_id={}", abort.task_id); - task.abort.notify_waiters(); + task.abort.cancel(); } None => log::warn!( "[medulla] task_abort for unknown task_id={} — dropping", @@ -146,7 +151,7 @@ impl MedullaTaskManager { let tasks = self.tasks.lock(); for (task_id, task) in tasks.iter() { log::debug!("[medulla] socket down — aborting task_id={task_id}"); - task.abort.notify_waiters(); + task.abort.cancel(); } } @@ -161,7 +166,7 @@ impl MedullaTaskManager { async fn drive( &self, run: payloads::TaskRun, - abort: Arc, + abort: CancellationToken, mut steer_rx: mpsc::UnboundedReceiver, ) { let task_id = run.task_id.clone(); @@ -174,7 +179,7 @@ impl MedullaTaskManager { .unwrap_or_else(|| DEFAULT_AGENT_ID.to_string()); let seq = Arc::new(AtomicI64::new(0)); - let mut agent = match build_agent(&agent_id, &task_id).await { + let mut agent = match build_agent(&agent_id, &task_id, &session_id).await { Ok(agent) => agent, Err(err) => { log::error!("[medulla] task_id={task_id} failed to build agent: {err}"); @@ -194,11 +199,25 @@ impl MedullaTaskManager { } }; - let deadline = (run.timeout_ms > 0).then(|| Duration::from_millis(run.timeout_ms)); + // `timeout_ms` is a hard wall-clock budget for the WHOLE task, not + // per turn: anchor a single deadline now and charge every turn + // (initial + steering follow-ups) against the time remaining until it. + let deadline = + (run.timeout_ms > 0).then(|| Instant::now() + Duration::from_millis(run.timeout_ms)); let mut next_input = run.instruction.clone(); let result; 'outer: loop { + // Charge this turn against the remaining task budget; if it's + // already spent, settle the whole task as timed out. + let remaining = match remaining_budget(deadline, Instant::now()) { + Ok(remaining) => remaining, + Err(()) => { + result = timeout_result(&task_id, &session_id, &seq); + break 'outer; + } + }; + let (progress_tx, progress_rx) = mpsc::channel::(256); agent.set_on_progress(Some(progress_tx)); let forwarder = spawn_forwarder( @@ -216,7 +235,7 @@ impl MedullaTaskManager { }; let turn = Box::pin(with_origin(origin, agent.run_single(&next_input))); - let turn_result = run_with_optional_timeout(deadline, &abort, turn).await; + let turn_result = run_with_optional_timeout(remaining, &abort, turn).await; // The forwarder ends when `progress_tx` drops; make sure it's flushed. agent.set_on_progress(None); let _ = forwarder.await; @@ -233,22 +252,7 @@ impl MedullaTaskManager { break 'outer; } TurnOutcome::TimedOut => { - emit_envelope( - &task_id, - envelope::error_envelope( - &session_id, - next_seq(&seq), - "task timed out", - true, - ), - ); - result = TaskResult { - task_id: task_id.clone(), - ok: false, - reply: String::new(), - usage: None, - error: Some("timeout".to_string()), - }; + result = timeout_result(&task_id, &session_id, &seq); break 'outer; } TurnOutcome::Errored(err) => { @@ -302,17 +306,22 @@ enum TurnOutcome { TimedOut, } -/// Race the agent turn against the cooperative abort signal and an optional -/// wall-clock deadline. +/// Race the agent turn against the cooperative abort signal and the remaining +/// slice of the task's wall-clock budget. +/// +/// `abort` is a [`CancellationToken`], so a cancellation that landed *before* +/// this call starts polling is not lost: `cancelled()` resolves immediately for +/// an already-cancelled token, and the `biased` select settles the turn as +/// [`TurnOutcome::Aborted`] before the agent future is ever polled. async fn run_with_optional_timeout( - deadline: Option, - abort: &Arc, + remaining: Option, + abort: &CancellationToken, turn: std::pin::Pin>>>, ) -> TurnOutcome { let run = async { tokio::select! { biased; - _ = abort.notified() => TurnOutcome::Aborted, + _ = abort.cancelled() => TurnOutcome::Aborted, res = turn => match res { Ok(reply) => TurnOutcome::Completed(reply), Err(err) => TurnOutcome::Errored(err.to_string()), @@ -320,7 +329,7 @@ async fn run_with_optional_timeout( } }; - match deadline { + match remaining { Some(d) => match tokio::time::timeout(d, run).await { Ok(outcome) => outcome, Err(_) => TurnOutcome::TimedOut, @@ -329,6 +338,35 @@ async fn run_with_optional_timeout( } } +/// Remaining wall-clock budget until the task `deadline`. +/// +/// `Ok(None)` = no deadline configured (run unbounded); `Ok(Some(d))` = `d` +/// left before the deadline; `Err(())` = the deadline has already passed, so +/// the caller must settle the task as timed out instead of starting a turn. +fn remaining_budget(deadline: Option, now: Instant) -> Result, ()> { + match deadline { + Some(d) if now >= d => Err(()), + Some(d) => Ok(Some(d - now)), + None => Ok(None), + } +} + +/// Build the terminal timeout [`TaskResult`], emitting the fatal `error` +/// envelope that bookends a timed-out task's stream. +fn timeout_result(task_id: &str, session_id: &str, seq: &AtomicI64) -> TaskResult { + emit_envelope( + task_id, + envelope::error_envelope(session_id, next_seq(seq), "task timed out", true), + ); + TaskResult { + task_id: task_id.to_string(), + ok: false, + reply: String::new(), + usage: None, + error: Some("timeout".to_string()), + } +} + /// Return queued steering input (if any) after briefly waiting for input that /// was sent while the turn was still in flight. async fn drain_steer(steer_rx: &mut mpsc::UnboundedReceiver) -> Option { @@ -342,18 +380,38 @@ async fn drain_steer(steer_rx: &mut mpsc::UnboundedReceiver) -> Option Result { +/// +/// The transcript identity is scoped by `session_id` (which the caller has +/// already resolved to the medulla `sessionId`, falling back to the `taskId`), +/// not the bare `agent_definition_name`. Without this, two `medulla:task_run`s +/// on the same `agentId` would collide onto one shared transcript and the +/// second would resume the first's history. +async fn build_agent(agent_id: &str, task_id: &str, session_id: &str) -> Result { let config = crate::openhuman::config::rpc::load_config_with_timeout().await?; crate::openhuman::agent::harness::AgentDefinitionRegistry::init_global(&config.workspace_dir) .map_err(|err| format!("failed to init agent definition registry: {err}"))?; let mut agent = Agent::from_config_for_agent(&config, agent_id) .map_err(|err| format!("failed to build agent `{agent_id}`: {err}"))?; agent.set_event_context(format!("medulla:{task_id}"), "medulla_harness"); + // Scope the transcript/session key per medulla session so distinct sessions + // on the same agent get isolated history (mirrors the web channel's + // per-thread `set_agent_definition_name`). + agent.set_agent_definition_name(medulla_session_key(agent_id, session_id)); agent.fetch_connected_integrations().await; let _ = agent.refresh_delegation_tools(); Ok(agent) } +/// Derive a per-session agent-definition (transcript) key from the medulla +/// `session_id`, namespaced by `agent_id`. The session id is truncated on a +/// char boundary to keep transcript filenames bounded; the underlying +/// [`Agent::set_agent_definition_name`] sanitizes any remaining non-filename +/// characters. +fn medulla_session_key(agent_id: &str, session_id: &str) -> String { + let short: String = session_id.chars().take(32).collect(); + format!("{agent_id}_{short}") +} + /// Spawn the per-turn progress → `medulla:task_envelope` forwarder. Returns its /// join handle so the driver can flush it before the next turn. fn spawn_forwarder( @@ -473,7 +531,7 @@ mod tests { let mgr = Arc::new(MedullaTaskManager::new()); // Manually seed a running task to simulate an in-flight run, then prove // a second registration under the same id is ignored. - let abort = Arc::new(Notify::new()); + let abort = CancellationToken::new(); let (steer_tx, _rx) = mpsc::unbounded_channel(); mgr.tasks .lock() @@ -490,4 +548,38 @@ mod tests { }); assert_eq!(mgr.tasks.lock().len(), 1); } + + #[test] + fn session_key_scopes_transcript_by_session_id() { + // Same agent id + distinct session ids => distinct transcript keys, so + // two medulla sessions can't collide onto one shared transcript. + let a = medulla_session_key("orchestrator", "sess-abc"); + let b = medulla_session_key("orchestrator", "sess-xyz"); + assert_ne!(a, b); + assert_eq!(a, "orchestrator_sess-abc"); + assert!(a.starts_with("orchestrator_")); + // Overlong session ids are truncated on a char boundary. + let long = "x".repeat(100); + let key = medulla_session_key("orchestrator", &long); + assert_eq!(key, format!("orchestrator_{}", "x".repeat(32))); + } + + #[test] + fn remaining_budget_reports_time_left_and_exhaustion() { + let now = Instant::now(); + // No deadline configured => unbounded. + assert_eq!(remaining_budget(None, now), Ok(None)); + // Deadline in the future => remaining time until it. + let future = now + Duration::from_secs(10); + match remaining_budget(Some(future), now) { + Ok(Some(d)) => assert!(d <= Duration::from_secs(10) && d > Duration::from_secs(9)), + other => panic!("expected some remaining budget, got {other:?}"), + } + // Deadline already reached / passed => exhausted. + assert_eq!(remaining_budget(Some(now), now), Err(())); + assert_eq!( + remaining_budget(Some(now), now + Duration::from_secs(1)), + Err(()) + ); + } }