mirror of
https://github.com/openai/codex.git
synced 2026-09-09 15:58:47 +00:00
Parallelize CSV agent job startup
This commit is contained in:
@@ -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;
|
||||
|
||||
430
codex-rs/core/src/tools/handlers/agent_jobs_startup.rs
Normal file
430
codex-rs/core/src/tools/handlers/agent_jobs_startup.rs
Normal 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"]);
|
||||
}
|
||||
}
|
||||
@@ -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(())
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user