diff --git a/codex-rs/core/src/tools/handlers/agent_jobs.rs b/codex-rs/core/src/tools/handlers/agent_jobs.rs index c9475eeb39..835783e272 100644 --- a/codex-rs/core/src/tools/handlers/agent_jobs.rs +++ b/codex-rs/core/src/tools/handlers/agent_jobs.rs @@ -177,6 +177,9 @@ impl JobProgressEmitter { } } +#[path = "agent_jobs_startup.rs"] +mod startup; + impl ToolHandler for BatchJobHandler { type Output = FunctionToolOutput; @@ -577,6 +580,7 @@ async fn run_agent_job_loop( .ok_or_else(|| anyhow::anyhow!("agent job {job_id} was not found"))?; let runtime_timeout = job_runtime_timeout(&job); let mut active_items: HashMap = HashMap::new(); + let mut starting_items = startup::StartupTasks::default(); let mut progress_emitter = JobProgressEmitter::new(); recover_running_items( session.clone(), @@ -611,85 +615,55 @@ async fn run_agent_job_loop( .await; } - if !cancel_requested && active_items.len() < options.max_concurrency { - let slots = options.max_concurrency - active_items.len(); - let pending_items = db - .list_agent_job_items( - job_id.as_str(), - Some(codex_state::AgentJobItemStatus::Pending), - Some(slots), - ) - .await?; - for item in pending_items { - let prompt = build_worker_prompt(&job, &item)?; - let items = vec![UserInput::Text { - text: prompt, - text_elements: Vec::new(), - }]; - let thread_id = match session - .services - .agent_control - .spawn_agent( - options.spawn_config.clone(), - items.into(), - Some(SessionSource::SubAgent(SubAgentSource::Other(format!( - "agent_job:{job_id}" - )))), - ) - .await - { - Ok(thread_id) => thread_id, - Err(CodexErr::AgentLimitReached { .. }) => { - db.mark_agent_job_item_pending( - job_id.as_str(), - item.item_id.as_str(), - /*error_message*/ None, - ) - .await?; - break; - } - Err(err) => { - let error_message = format!("failed to spawn worker: {err}"); - db.mark_agent_job_item_failed( - job_id.as_str(), - item.item_id.as_str(), - error_message.as_str(), - ) - .await?; - progressed = true; - continue; - } - }; - let assigned = db - .mark_agent_job_item_running_with_thread( - job_id.as_str(), - item.item_id.as_str(), - thread_id.to_string().as_str(), - ) - .await?; - if !assigned { - let _ = session - .services - .agent_control - .shutdown_live_agent(thread_id) - .await; - continue; - } - active_items.insert( - thread_id, - ActiveJobItem { - item_id: item.item_id.clone(), - started_at: Instant::now(), - status_rx: session - .services - .agent_control - .subscribe_status(thread_id) - .await - .ok(), - }, - ); - progressed = true; - } + if startup::drain_ready_startups( + session.clone(), + db.clone(), + job_id.as_str(), + &mut active_items, + &mut starting_items, + ) + .await? + { + progressed = true; + } + + if !cancel_requested + && active_items.len() + starting_items.len() < options.max_concurrency + && startup::launch_pending_items( + session.clone(), + db.clone(), + &job, + job_id.as_str(), + &options, + active_items.len(), + &mut starting_items, + ) + .await? + { + progressed = true; + } + + if startup::drain_ready_startups( + session.clone(), + db.clone(), + job_id.as_str(), + &mut active_items, + &mut starting_items, + ) + .await? + { + progressed = true; + } + + if startup::reap_stale_startups( + db.clone(), + job_id.as_str(), + &mut starting_items, + runtime_timeout, + ) + .await? + { + progressed = true; } if reap_stale_active_items( @@ -708,17 +682,28 @@ async fn run_agent_job_loop( if finished.is_empty() { let progress = db.get_agent_job_progress(job_id.as_str()).await?; if cancel_requested { - if progress.running_items == 0 && active_items.is_empty() { + if progress.running_items == 0 + && active_items.is_empty() + && starting_items.is_empty() + { break; } } else if progress.pending_items == 0 && progress.running_items == 0 && active_items.is_empty() + && starting_items.is_empty() { break; } if !progressed { - wait_for_status_change(&active_items).await; + startup::wait_for_startup_or_status_change( + session.clone(), + db.clone(), + job_id.as_str(), + &mut active_items, + &mut starting_items, + ) + .await?; } continue; } @@ -837,10 +822,10 @@ async fn recover_running_items( continue; } let Some(assigned_thread_id) = item.assigned_thread_id.clone() else { - db.mark_agent_job_item_failed( + db.mark_agent_job_item_pending( job_id, item.item_id.as_str(), - "running item is missing assigned_thread_id", + Some("worker startup was interrupted before a thread was assigned"), ) .await?; continue; diff --git a/codex-rs/core/src/tools/handlers/agent_jobs_startup.rs b/codex-rs/core/src/tools/handlers/agent_jobs_startup.rs new file mode 100644 index 0000000000..83bf3d3636 --- /dev/null +++ b/codex-rs/core/src/tools/handlers/agent_jobs_startup.rs @@ -0,0 +1,430 @@ +use super::*; +use std::collections::HashMap; +use std::future::Future; +use tokio::task::AbortHandle; +use tokio::task::Id as TaskId; +use tokio::task::JoinError; +use tokio::task::JoinSet; + +#[derive(Debug)] +pub(super) struct WorkerStartup { + pub(super) item_id: String, + pub(super) started_at: Instant, + pub(super) spawn_latency: Duration, + pub(super) result: Result, +} + +#[derive(Debug)] +pub(super) struct LaunchingJobItem { + item_id: String, + started_at: Instant, + abort_handle: AbortHandle, +} + +#[derive(Debug, Default)] +pub(super) struct StartupTasks { + starting_items: JoinSet, + launching_items: HashMap, +} + +impl StartupTasks { + pub(super) fn len(&self) -> usize { + self.starting_items.len() + } + + pub(super) fn is_empty(&self) -> bool { + self.starting_items.is_empty() + } +} + +fn spawn_tracked_startup_task( + startup_tasks: &mut StartupTasks, + item_id: String, + started_at: Instant, + task: F, +) where + F: Future + Send + 'static, +{ + let abort_handle = startup_tasks.starting_items.spawn(task); + startup_tasks.launching_items.insert( + abort_handle.id(), + LaunchingJobItem { + item_id, + started_at, + abort_handle, + }, + ); +} + +pub(super) async fn launch_pending_items( + session: Arc, + db: Arc, + job: &codex_state::AgentJob, + job_id: &str, + options: &JobRunnerOptions, + active_items_len: usize, + startup_tasks: &mut StartupTasks, +) -> anyhow::Result { + let slots = options + .max_concurrency + .saturating_sub(active_items_len + startup_tasks.len()); + if slots == 0 { + return Ok(false); + } + + let pending_items = db + .list_agent_job_items( + job_id, + Some(codex_state::AgentJobItemStatus::Pending), + Some(slots), + ) + .await?; + + let mut launched = 0usize; + let mut progressed = false; + for item in pending_items { + let claimed = db + .mark_agent_job_item_running(job_id, item.item_id.as_str()) + .await?; + if !claimed { + continue; + } + + let prompt = match build_worker_prompt(job, &item) { + Ok(prompt) => prompt, + Err(err) => { + let error_message = format!("failed to build worker prompt: {err}"); + db.mark_agent_job_item_failed( + job_id, + item.item_id.as_str(), + error_message.as_str(), + ) + .await?; + progressed = true; + continue; + } + }; + + let item_id = item.item_id.clone(); + let session = session.clone(); + let spawn_config = options.spawn_config.clone(); + let session_source = + SessionSource::SubAgent(SubAgentSource::Other(format!("agent_job:{job_id}"))); + let started_at = Instant::now(); + spawn_tracked_startup_task(startup_tasks, item_id.clone(), started_at, async move { + let items = vec![UserInput::Text { + text: prompt, + text_elements: Vec::new(), + }]; + let result = session + .services + .agent_control + .spawn_agent(spawn_config, items.into(), Some(session_source)) + .await; + WorkerStartup { + item_id, + started_at, + spawn_latency: started_at.elapsed(), + result, + } + }); + launched = launched.saturating_add(1); + progressed = true; + } + + if launched > 0 { + tracing::info!( + job_id, + launched, + active_items = active_items_len, + starting_items = startup_tasks.len(), + target_concurrency = options.max_concurrency, + "agent job queued worker startups" + ); + } + Ok(progressed) +} + +pub(super) async fn drain_ready_startups( + session: Arc, + db: Arc, + job_id: &str, + active_items: &mut HashMap, + startup_tasks: &mut StartupTasks, +) -> anyhow::Result { + let mut progressed = false; + while let Some(result) = startup_tasks.starting_items.try_join_next_with_id() { + let starting_items_len = startup_tasks.starting_items.len(); + handle_worker_startup_result( + session.clone(), + db.clone(), + job_id, + active_items, + startup_tasks, + result, + starting_items_len, + ) + .await?; + progressed = true; + } + Ok(progressed) +} + +pub(super) async fn wait_for_startup_or_status_change( + session: Arc, + db: Arc, + job_id: &str, + active_items: &mut HashMap, + startup_tasks: &mut StartupTasks, +) -> anyhow::Result<()> { + if startup_tasks.is_empty() { + wait_for_status_change(active_items).await; + return Ok(()); + } + + let active_items_ref = &*active_items; + if active_items_ref.is_empty() { + if let Some(result) = startup_tasks.starting_items.join_next_with_id().await { + let starting_items_len = startup_tasks.starting_items.len(); + handle_worker_startup_result( + session, + db, + job_id, + active_items, + startup_tasks, + result, + starting_items_len, + ) + .await?; + } + return Ok(()); + } + + tokio::select! { + startup = startup_tasks.starting_items.join_next_with_id() => { + if let Some(result) = startup { + let starting_items_len = startup_tasks.starting_items.len(); + handle_worker_startup_result( + session, + db, + job_id, + active_items, + startup_tasks, + result, + starting_items_len, + ) + .await?; + } + } + _ = wait_for_status_change(active_items_ref) => {} + } + Ok(()) +} + +pub(super) async fn reap_stale_startups( + db: Arc, + job_id: &str, + startup_tasks: &mut StartupTasks, + runtime_timeout: Duration, +) -> anyhow::Result { + let stale_task_ids: Vec<_> = startup_tasks + .launching_items + .iter() + .filter_map(|(task_id, item)| { + (item.started_at.elapsed() >= runtime_timeout).then_some(*task_id) + }) + .collect(); + if stale_task_ids.is_empty() { + return Ok(false); + } + + for task_id in stale_task_ids { + let Some(item) = startup_tasks.launching_items.remove(&task_id) else { + continue; + }; + item.abort_handle.abort(); + let error_message = + format!("worker exceeded max runtime of {runtime_timeout:?} before startup completed"); + db.mark_agent_job_item_failed(job_id, item.item_id.as_str(), error_message.as_str()) + .await?; + tracing::warn!( + job_id, + item_id = item.item_id, + ?task_id, + "agent job worker startup timed out" + ); + } + Ok(true) +} + +async fn handle_worker_startup_result( + session: Arc, + db: Arc, + job_id: &str, + active_items: &mut HashMap, + startup_tasks: &mut StartupTasks, + result: Result<(TaskId, WorkerStartup), JoinError>, + starting_items_len: usize, +) -> anyhow::Result<()> { + match result { + Ok((task_id, startup)) => { + startup_tasks.launching_items.remove(&task_id); + match startup.result { + Ok(thread_id) => { + let thread_id_str = thread_id.to_string(); + let assigned = db + .set_agent_job_item_thread( + job_id, + startup.item_id.as_str(), + thread_id_str.as_str(), + ) + .await?; + if !assigned { + let _ = session + .services + .agent_control + .shutdown_live_agent(thread_id) + .await; + tracing::debug!( + job_id, + item_id = startup.item_id, + thread_id = %thread_id, + "agent job worker startup finished after item left running state" + ); + return Ok(()); + } + + let item_id = startup.item_id; + active_items.insert( + thread_id, + ActiveJobItem { + item_id: item_id.clone(), + started_at: startup.started_at, + status_rx: session + .services + .agent_control + .subscribe_status(thread_id) + .await + .ok(), + }, + ); + tracing::info!( + job_id, + item_id, + thread_id = %thread_id, + spawn_latency_ms = startup.spawn_latency.as_millis() as u64, + active_items = active_items.len(), + starting_items = starting_items_len, + "agent job worker startup completed" + ); + } + Err(CodexErr::AgentLimitReached { .. }) => { + let _ = db + .mark_agent_job_item_pending( + job_id, + startup.item_id.as_str(), + /*error_message*/ None, + ) + .await?; + tracing::debug!( + job_id, + item_id = startup.item_id, + starting_items = starting_items_len, + "agent job worker startup hit agent limit" + ); + } + Err(err) => { + let error_message = format!("failed to spawn worker: {err}"); + let _ = db + .mark_agent_job_item_failed( + job_id, + startup.item_id.as_str(), + error_message.as_str(), + ) + .await?; + tracing::warn!( + job_id, + item_id = startup.item_id, + error = %err, + "agent job worker startup failed" + ); + } + } + } + Err(join_error) => { + let task_id = join_error.id(); + let Some(item) = startup_tasks.launching_items.remove(&task_id) else { + return Ok(()); + }; + let error_message = format!("worker startup task failed: {join_error}"); + let _ = db + .mark_agent_job_item_failed(job_id, item.item_id.as_str(), error_message.as_str()) + .await?; + tracing::warn!( + job_id, + item_id = item.item_id, + error = %join_error, + "agent job worker startup task exited unexpectedly" + ); + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use pretty_assertions::assert_eq; + use std::sync::Arc; + use std::sync::atomic::AtomicUsize; + use std::sync::atomic::Ordering; + use tokio::sync::Barrier; + use tokio::time::timeout; + + #[tokio::test] + async fn spawn_tracked_startup_task_starts_multiple_workers_without_serial_waiting() { + let mut startup_tasks = StartupTasks::default(); + let started = Arc::new(AtomicUsize::new(0)); + let barrier = Arc::new(Barrier::new(4)); + + for idx in 0..3usize { + let started = Arc::clone(&started); + let barrier = Arc::clone(&barrier); + spawn_tracked_startup_task( + &mut startup_tasks, + format!("item-{idx}"), + Instant::now(), + async move { + started.fetch_add(1, Ordering::SeqCst); + barrier.wait().await; + WorkerStartup { + item_id: format!("item-{idx}"), + started_at: Instant::now(), + spawn_latency: Duration::ZERO, + result: Err(CodexErr::ThreadNotFound(ThreadId::new())), + } + }, + ); + } + + timeout(Duration::from_secs(1), async { + while started.load(Ordering::SeqCst) < 3 { + tokio::task::yield_now().await; + } + }) + .await + .expect("all startup tasks should begin running"); + + assert_eq!(startup_tasks.len(), 3); + assert_eq!(startup_tasks.launching_items.len(), 3); + + barrier.wait().await; + + let mut outputs = Vec::new(); + while let Some(result) = startup_tasks.starting_items.join_next().await { + outputs.push(result.expect("startup task should complete").item_id); + } + outputs.sort(); + assert_eq!(outputs, vec!["item-0", "item-1", "item-2"]); + } +} diff --git a/codex-rs/state/src/runtime/agent_jobs.rs b/codex-rs/state/src/runtime/agent_jobs.rs index 3f5526c58d..bbc7dc4d1e 100644 --- a/codex-rs/state/src/runtime/agent_jobs.rs +++ b/codex-rs/state/src/runtime/agent_jobs.rs @@ -319,15 +319,14 @@ WHERE id = ? let now = Utc::now().timestamp(); let result = sqlx::query( r#" -UPDATE agent_job_items -SET - status = ?, - assigned_thread_id = NULL, - attempt_count = attempt_count + 1, - updated_at = ?, - last_error = NULL -WHERE job_id = ? AND item_id = ? AND status = ? - "#, + UPDATE agent_job_items + SET + status = ?, + assigned_thread_id = NULL, + updated_at = ?, + last_error = NULL + WHERE job_id = ? AND item_id = ? AND status = ? + "#, ) .bind(AgentJobItemStatus::Running.as_str()) .bind(now) @@ -407,10 +406,10 @@ WHERE job_id = ? AND item_id = ? AND status = ? let now = Utc::now().timestamp(); let result = sqlx::query( r#" -UPDATE agent_job_items -SET assigned_thread_id = ?, updated_at = ? -WHERE job_id = ? AND item_id = ? AND status = ? - "#, + UPDATE agent_job_items + SET assigned_thread_id = ?, attempt_count = attempt_count + 1, updated_at = ? + WHERE job_id = ? AND item_id = ? AND status = ? + "#, ) .bind(thread_id) .bind(now) @@ -681,4 +680,60 @@ mod tests { assert_eq!(item.last_error, Some("missing report".to_string())); Ok(()) } + + #[tokio::test] + async fn set_agent_job_item_thread_increments_attempt_count_after_claim() -> anyhow::Result<()> + { + let codex_home = unique_temp_dir(); + let runtime = StateRuntime::init(codex_home, "test-provider".to_string()).await?; + let job_id = "job-1".to_string(); + let item_id = "item-1".to_string(); + runtime + .create_agent_job( + &AgentJobCreateParams { + id: job_id.clone(), + name: "test-job".to_string(), + instruction: "Return a result".to_string(), + auto_export: true, + max_runtime_seconds: None, + output_schema_json: None, + input_headers: vec!["path".to_string()], + input_csv_path: "/tmp/in.csv".to_string(), + output_csv_path: "/tmp/out.csv".to_string(), + }, + &[AgentJobItemCreateParams { + item_id: item_id.clone(), + row_index: 0, + source_id: None, + row_json: json!({"path":"file-1"}), + }], + ) + .await?; + runtime.mark_agent_job_running(job_id.as_str()).await?; + + let claimed = runtime + .mark_agent_job_item_running(job_id.as_str(), item_id.as_str()) + .await?; + assert!(claimed); + + let item = runtime + .get_agent_job_item(job_id.as_str(), item_id.as_str()) + .await? + .expect("job item should exist"); + assert_eq!(item.attempt_count, 0); + assert_eq!(item.assigned_thread_id, None); + + let assigned = runtime + .set_agent_job_item_thread(job_id.as_str(), item_id.as_str(), "thread-1") + .await?; + assert!(assigned); + + let item = runtime + .get_agent_job_item(job_id.as_str(), item_id.as_str()) + .await? + .expect("job item should exist"); + assert_eq!(item.attempt_count, 1); + assert_eq!(item.assigned_thread_id, Some("thread-1".to_string())); + Ok(()) + } }