Parallelize CSV agent job startup

This commit is contained in:
Dave Aitel
2026-04-08 21:54:24 -04:00
parent aac1e74cd5
commit da391e4bd2
3 changed files with 566 additions and 96 deletions

View File

@@ -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<ThreadId, ActiveJobItem> = 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;

View File

@@ -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<ThreadId, CodexErr>,
}
#[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<WorkerStartup>,
launching_items: HashMap<TaskId, LaunchingJobItem>,
}
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<F>(
startup_tasks: &mut StartupTasks,
item_id: String,
started_at: Instant,
task: F,
) where
F: Future<Output = WorkerStartup> + 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<Session>,
db: Arc<codex_state::StateRuntime>,
job: &codex_state::AgentJob,
job_id: &str,
options: &JobRunnerOptions,
active_items_len: usize,
startup_tasks: &mut StartupTasks,
) -> anyhow::Result<bool> {
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<Session>,
db: Arc<codex_state::StateRuntime>,
job_id: &str,
active_items: &mut HashMap<ThreadId, ActiveJobItem>,
startup_tasks: &mut StartupTasks,
) -> anyhow::Result<bool> {
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<Session>,
db: Arc<codex_state::StateRuntime>,
job_id: &str,
active_items: &mut HashMap<ThreadId, ActiveJobItem>,
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<codex_state::StateRuntime>,
job_id: &str,
startup_tasks: &mut StartupTasks,
runtime_timeout: Duration,
) -> anyhow::Result<bool> {
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<Session>,
db: Arc<codex_state::StateRuntime>,
job_id: &str,
active_items: &mut HashMap<ThreadId, ActiveJobItem>,
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"]);
}
}

View File

@@ -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(())
}
}