From 245bbd5d27aa7a2ee26ee9e6ff0591d8042d4bbf Mon Sep 17 00:00:00 2001 From: hugocasa Date: Fri, 18 Sep 2026 17:36:21 +0200 Subject: [PATCH] feat: dispatch ai agent tools through the worker queue --- ...9e7801129570106b5acd49c0f708398606771.json | 23 + ...ca7cb85bdfc66110b6534260f8642f91f006a.json | 22 + ...5a1754ccb8ddeb08c23795dfe14e9f332e9dc.json | 35 ++ ...26d32b476a43ed391c09d7678215195c4283a.json | 15 + backend/tests/ai_tool_queue.rs | 181 ++++++ backend/windmill-common/src/worker.rs | 13 + backend/windmill-queue/src/jobs.rs | 39 ++ backend/windmill-worker/src/ai/tools.rs | 578 +++++++++--------- backend/windmill-worker/src/ai_executor.rs | 81 +-- backend/windmill-worker/src/worker.rs | 1 + 10 files changed, 670 insertions(+), 318 deletions(-) create mode 100644 backend/.sqlx/query-0c2bf6925de6dd4d9f8d47d7fdd9e7801129570106b5acd49c0f708398606771.json create mode 100644 backend/.sqlx/query-451f303f4a24d848ba4ccdc2441ca7cb85bdfc66110b6534260f8642f91f006a.json create mode 100644 backend/.sqlx/query-789c3f6d29f46fcb17a22fe97405a1754ccb8ddeb08c23795dfe14e9f332e9dc.json create mode 100644 backend/.sqlx/query-e459c277c0bc27293d71972d12d26d32b476a43ed391c09d7678215195c4283a.json create mode 100644 backend/tests/ai_tool_queue.rs diff --git a/backend/.sqlx/query-0c2bf6925de6dd4d9f8d47d7fdd9e7801129570106b5acd49c0f708398606771.json b/backend/.sqlx/query-0c2bf6925de6dd4d9f8d47d7fdd9e7801129570106b5acd49c0f708398606771.json new file mode 100644 index 0000000000..a7b61d7282 --- /dev/null +++ b/backend/.sqlx/query-0c2bf6925de6dd4d9f8d47d7fdd9e7801129570106b5acd49c0f708398606771.json @@ -0,0 +1,23 @@ +{ + "db_name": "PostgreSQL", + "query": "WITH RECURSIVE descendants AS (\n SELECT id FROM v2_job WHERE parent_job = $1 AND workspace_id = $2\n UNION ALL\n SELECT j.id FROM v2_job j JOIN descendants d ON j.parent_job = d.id\n WHERE j.workspace_id = $2\n ) SELECT d.id AS \"id!\" FROM descendants d JOIN v2_job_queue q ON q.id = d.id", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "id!", + "type_info": "Uuid" + } + ], + "parameters": { + "Left": [ + "Uuid", + "Text" + ] + }, + "nullable": [ + null + ] + }, + "hash": "0c2bf6925de6dd4d9f8d47d7fdd9e7801129570106b5acd49c0f708398606771" +} diff --git a/backend/.sqlx/query-451f303f4a24d848ba4ccdc2441ca7cb85bdfc66110b6534260f8642f91f006a.json b/backend/.sqlx/query-451f303f4a24d848ba4ccdc2441ca7cb85bdfc66110b6534260f8642f91f006a.json new file mode 100644 index 0000000000..8209f67c75 --- /dev/null +++ b/backend/.sqlx/query-451f303f4a24d848ba4ccdc2441ca7cb85bdfc66110b6534260f8642f91f006a.json @@ -0,0 +1,22 @@ +{ + "db_name": "PostgreSQL", + "query": "UPDATE v2_job_queue SET started_at = now() WHERE id = $1 RETURNING started_at", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "started_at", + "type_info": "Timestamptz" + } + ], + "parameters": { + "Left": [ + "Uuid" + ] + }, + "nullable": [ + true + ] + }, + "hash": "451f303f4a24d848ba4ccdc2441ca7cb85bdfc66110b6534260f8642f91f006a" +} diff --git a/backend/.sqlx/query-789c3f6d29f46fcb17a22fe97405a1754ccb8ddeb08c23795dfe14e9f332e9dc.json b/backend/.sqlx/query-789c3f6d29f46fcb17a22fe97405a1754ccb8ddeb08c23795dfe14e9f332e9dc.json new file mode 100644 index 0000000000..9702cbf61d --- /dev/null +++ b/backend/.sqlx/query-789c3f6d29f46fcb17a22fe97405a1754ccb8ddeb08c23795dfe14e9f332e9dc.json @@ -0,0 +1,35 @@ +{ + "db_name": "PostgreSQL", + "query": "SELECT id, status = 'success' AS \"success!\", result AS \"result: Json>\"\n FROM v2_job_completed WHERE workspace_id = $1 AND id = ANY($2)", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "id", + "type_info": "Uuid" + }, + { + "ordinal": 1, + "name": "success!", + "type_info": "Bool" + }, + { + "ordinal": 2, + "name": "result: Json>", + "type_info": "Jsonb" + } + ], + "parameters": { + "Left": [ + "Text", + "UuidArray" + ] + }, + "nullable": [ + false, + null, + true + ] + }, + "hash": "789c3f6d29f46fcb17a22fe97405a1754ccb8ddeb08c23795dfe14e9f332e9dc" +} diff --git a/backend/.sqlx/query-e459c277c0bc27293d71972d12d26d32b476a43ed391c09d7678215195c4283a.json b/backend/.sqlx/query-e459c277c0bc27293d71972d12d26d32b476a43ed391c09d7678215195c4283a.json new file mode 100644 index 0000000000..3385b83f74 --- /dev/null +++ b/backend/.sqlx/query-e459c277c0bc27293d71972d12d26d32b476a43ed391c09d7678215195c4283a.json @@ -0,0 +1,15 @@ +{ + "db_name": "PostgreSQL", + "query": "UPDATE v2_job_queue SET worker = $1 WHERE id = $2", + "describe": { + "columns": [], + "parameters": { + "Left": [ + "Varchar", + "Uuid" + ] + }, + "nullable": [] + }, + "hash": "e459c277c0bc27293d71972d12d26d32b476a43ed391c09d7678215195c4283a" +} diff --git a/backend/tests/ai_tool_queue.rs b/backend/tests/ai_tool_queue.rs new file mode 100644 index 0000000000..878231f23c --- /dev/null +++ b/backend/tests/ai_tool_queue.rs @@ -0,0 +1,181 @@ +use axum::{ + extract::State, + routing::{get, post}, + Json, Router, +}; +use serde_json::{json, Value}; +use sqlx::{Pool, Postgres}; +use std::{sync::Arc, time::Duration}; +use tokio::sync::Notify; +use windmill_common::{flows::FlowValue, jobs::JobPayload}; +use windmill_test_utils::{ + completed_job, in_test_worker, listen_for_completed_jobs, ApiServer, RunJob, StreamFind, +}; + +async fn model(Json(body): Json) -> ([(&'static str, &'static str); 1], String) { + let finished = body["messages"] + .as_array() + .unwrap() + .iter() + .any(|m| m["role"] == "tool"); + let delta = if finished { + json!({"role":"assistant", "content":"done"}) + } else { + json!({"role":"assistant", "tool_calls": (0..3).map(|i| json!({ + "index":i, "id":format!("call_{i}"), "type":"function", + "function":{"name":format!("tool_{i}"), "arguments":"{}"} + })).collect::>()}) + }; + let event = json!({"choices":[{"index":0, "delta":delta, "finish_reason":null}]}); + let end = json!({"choices":[{"index":0, "delta":{}, "finish_reason":if finished {"stop"} else {"tool_calls"}}]}); + ( + [("content-type", "text/event-stream")], + format!("data: {event}\n\ndata: {end}\n\ndata: [DONE]\n\n"), + ) +} + +async fn run_batch(db: Pool, parallel: bool, limited: bool) -> anyhow::Result<()> { + std::env::set_var("ALLOW_PRIVATE_AI_BASE_URLS", "true"); + let server = ApiServer::start(db.clone()).await?; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?; + let base = format!("http://{}", listener.local_addr()?); + let gate = Arc::new(Notify::new()); + let router = Router::new() + .route("/v1/chat/completions", post(model)) + .route( + "/first", + get(|State(gate): State>| async move { + gate.notified().await; + "ok" + }), + ) + .route( + "/second", + get(|State(gate): State>| async move { + gate.notify_one(); + "ok" + }), + ) + .with_state(gate); + let stub = tokio::spawn(async move { + axum::serve(listener, router).await.unwrap(); + }); + let tools: Vec = (0..3).map(|i| { + let wait = if parallel && !limited && i < 2 { + format!("await fetch('{base}/{}');", if i == 0 {"first"} else {"second"}) + } else { String::new() }; + let finish = if i == 2 { "throw new Error('expected tool failure');".to_string() } + else { format!("return {i};") }; + json!({"id":format!("t{i}"),"summary":format!("tool_{i}"),"value":{ + "type":"rawscript", "language":"bun", "input_transforms":{}, + "tag": if parallel { "bun" } else { "unserved-tool-tag" }, + "concurrent_limit": if limited { Some(1) } else { None }, + "custom_concurrency_key": if limited { Some("agent-tool-test") } else { None }, + "content":format!("export async function main() {{ {wait} await Bun.sleep({}); {finish} }}", if i == 0 {200} else {0}) + }}) + }).collect(); + let flow: FlowValue = serde_json::from_value(json!({"modules":[{"id":"agent","value":{ + "type":"aiagent", "tools":tools, "input_transforms":{ + "provider":{"type":"static","value":{"kind":"customai","model":"queue-test","resource":{"base_url":format!("{base}/v1")}}}, + "user_message":{"type":"static","value":"run the tools"}, + "max_iterations":{"type":"static","value":3} + } + }}]}))?; + let id = RunJob::from(JobPayload::RawFlow { + value: flow, + path: Some("u/test/agent_queue".into()), + restarted_from: None, + }) + .push(&db) + .await; + let notifications = listen_for_completed_jobs(&db).await; + let wait = async { + if parallel { + in_test_worker(&db, notifications.find(&id), server.addr.port()).await; + } else { + notifications.find(&id).await; + } + }; + tokio::time::timeout( + Duration::from_secs(45), + in_test_worker(&db, wait, server.addr.port()), + ) + .await?; + let result = completed_job(id, &db).await; + assert!(result.success, "{:?}", result.result); + let result = result.json_result().expect("agent result"); + let messages: Vec<_> = result["messages"] + .as_array() + .unwrap() + .iter() + .filter(|m| m["role"] == "tool") + .collect(); + assert_eq!(messages.len(), 3); + for (index, message) in messages.iter().enumerate() { + assert_eq!(message["tool_call_id"], format!("call_{index}")); + } + assert_eq!(messages[0]["content"], "0"); + assert_eq!(messages[1]["content"], "1"); + assert!(messages[2]["content"] + .as_str() + .unwrap() + .contains("expected tool failure")); + let (parent_id, parent_worker): (uuid::Uuid, String) = sqlx::query_as( + "SELECT j.id, c.worker FROM v2_job j JOIN v2_job_completed c USING(id) WHERE j.parent_job = $1" + ).bind(id).fetch_one(&db).await?; + let children: Vec<(String, String, bool)> = sqlx::query_as( + "SELECT j.runnable_path, c.worker, c.status = 'success' FROM v2_job j JOIN v2_job_completed c USING(id) WHERE j.parent_job = $1 ORDER BY j.runnable_path" + ).bind(parent_id).fetch_all(&db).await?; + assert_eq!(children.len(), 3); + assert_eq!( + children[0].1, parent_worker, + "first tool must stay on the parent worker" + ); + assert_eq!( + children.iter().map(|c| c.2).collect::>(), + vec![true, true, false] + ); + if limited { + let overlapping: i64 = sqlx::query_scalar( + "SELECT count(*) FROM v2_job j1 JOIN v2_job_completed c1 ON c1.id = j1.id + JOIN v2_job j2 ON j2.parent_job = j1.parent_job AND j2.id > j1.id + JOIN v2_job_completed c2 ON c2.id = j2.id + WHERE j1.parent_job = $1 AND c1.started_at < c2.completed_at AND c2.started_at < c1.completed_at" + ).bind(parent_id).fetch_one(&db).await?; + assert_eq!( + overlapping, 0, + "the reserved first job must count toward the shared limit" + ); + } else if parallel { + assert_ne!( + children[1].1, parent_worker, + "second tool must unblock the first from another worker" + ); + } else { + assert!(children.iter().all(|child| child.1 == parent_worker)); + } + stub.abort(); + server.close().await?; + Ok(()) +} + +#[sqlx::test(fixtures("base"))] +#[serial_test::serial] +async fn parent_drains_tools_without_another_worker(db: Pool) -> anyhow::Result<()> { + run_batch(db, false, false).await +} + +#[sqlx::test(fixtures("base"))] +#[serial_test::serial] +async fn parent_keeps_first_tool_while_another_worker_runs_siblings( + db: Pool, +) -> anyhow::Result<()> { + run_batch(db, true, false).await +} + +#[cfg(all(feature = "enterprise", feature = "private"))] +#[sqlx::test(fixtures("base"))] +#[serial_test::serial] +async fn reserved_tool_obeys_shared_concurrency_limit(db: Pool) -> anyhow::Result<()> { + run_batch(db, true, true).await +} diff --git a/backend/windmill-common/src/worker.rs b/backend/windmill-common/src/worker.rs index c742a3efcb..56452602cd 100644 --- a/backend/windmill-common/src/worker.rs +++ b/backend/windmill-common/src/worker.rs @@ -775,6 +775,19 @@ pub fn make_pull_query(tags: &[String]) -> String { query } +/// Claim a parent's tool jobs without tag filtering. The caller must supply only child IDs +/// it owns; this query is an internal scheduling primitive and does not authorize job access. +pub fn make_tool_job_pull_query(job_ids: &[uuid::Uuid]) -> String { + // pull() binds only the worker name. These literals come from typed UUIDs, never input SQL. + let ids = job_ids.iter().map(|id| format!("'{id}'::uuid")).join(", "); + format_pull_query(format!( + "SELECT id FROM v2_job_queue + WHERE running = false AND id = ANY(ARRAY[{ids}]::uuid[]) AND scheduled_for <= now() + ORDER BY priority DESC NULLS LAST, scheduled_for + FOR UPDATE SKIP LOCKED LIMIT 1" + )) +} + // Variant of `make_pull_query` that additionally excludes jobs whose workspace_id is in the // overloaded-list bind parameter ($2::text[]). Built as a separate string (rather than reusing // `make_pull_query` with an always-bound array) so the planner can keep using the same indexes diff --git a/backend/windmill-queue/src/jobs.rs b/backend/windmill-queue/src/jobs.rs index 429701f7e0..372e947ddc 100644 --- a/backend/windmill-queue/src/jobs.rs +++ b/backend/windmill-queue/src/jobs.rs @@ -4399,6 +4399,45 @@ pub fn has_active_concurrency_limit(concurrent_limit: Option) -> bool { concurrent_limit.is_some_and(|n| n > 0) } +/// Admit a job already owned by a worker without releasing its queue reservation. +/// The caller must maintain its heartbeat while waiting and complete it on failure. +pub async fn try_admit_owned_job(db: &DB, job: &MiniPulledJob) -> error::Result { + #[cfg(all(feature = "private", feature = "enterprise"))] + { + let settings = windmill_common::runnable_settings::prefetch_cached_from_handle( + job.runnable_settings_handle, + db, + ) + .await? + .1 + .maybe_fallback(None, job.concurrent_limit, job.concurrency_time_window_s); + if has_active_concurrency_limit(settings.concurrent_limit) + && !*DISABLE_CONCURRENCY_LIMIT + && job.canceled_by.is_none() + { + let key = concurrency_key(db, &job.id).await?.ok_or_else(|| { + Error::internal_err(format!("No concurrency key found for job {}", job.id)) + })?; + if !key.is_empty() { + return Ok(crate::jobs_ee::update_concurrency_counter( + db, + &job.id, + key, + serde_json::json!({ job.id.to_string(): {} }), + job.id.to_string(), + settings.concurrency_time_window_s.unwrap_or(0), + settings.concurrent_limit.unwrap_or_default(), + ) + .await? + .0); + } + } + } + #[cfg(not(all(feature = "private", feature = "enterprise")))] + let _ = (db, job); + Ok(true) +} + pub async fn custom_concurrency_key( db: &Pool, job_id: &Uuid, diff --git a/backend/windmill-worker/src/ai/tools.rs b/backend/windmill-worker/src/ai/tools.rs index 7e775f34fe..e32bb456bc 100644 --- a/backend/windmill-worker/src/ai/tools.rs +++ b/backend/windmill-worker/src/ai/tools.rs @@ -9,14 +9,12 @@ use crate::result_processor::handle_non_flow_job_error; use crate::worker_flow::{ evaluate_input_transform, raw_script_to_payload, script_to_payload, JobPayloadWithTag, }; -use crate::{ - create_job_dir, handle_queued_job, JobCompletedReceiver, JobCompletedSender, SendResult, - SendResultPayload, -}; +use crate::{create_job_dir, handle_queued_job, JobCompletedSender}; use anyhow::Context; use mappable_rc::Marc; use serde_json::value::RawValue; -use std::{collections::HashMap, sync::Arc}; +use sqlx::types::Json; +use std::{collections::HashMap, sync::Arc, time::Duration}; use uuid::Uuid; use windmill_ai::{ai_types::OpenAIToolCall, query_builder::StreamEventSink, types::*}; use windmill_common::jobs::JobPayload; @@ -36,11 +34,11 @@ use windmill_common::{ flow_conversations::{MessageExtras, MessageType}, flow_status::AgentAction, flows::FlowModuleValue, - worker::{to_raw_value, Connection}, + worker::{make_tool_job_pull_query, to_raw_value, Connection}, }; use windmill_queue::{ - add_completed_job, add_completed_job_error, get_mini_pulled_job, push, MiniCompletedJob, - MiniPulledJob, PushArgs, PushIsolationLevel, + get_mini_pulled_job, pull, push, try_admit_owned_job, MiniCompletedJob, MiniPulledJob, + PushArgs, PushIsolationLevel, }; /// Shared collection of abort handles for spawned tool tasks. @@ -82,6 +80,7 @@ pub struct ToolExecutionContext<'a> { // Abort handles for spawned tool tasks (used for force-cancel cleanup) pub tool_abort_handles: ToolAbortHandles, + pub job_completed_tx: JobCompletedSender, } /// Execute all tool calls from an AI response @@ -98,7 +97,8 @@ pub async fn execute_tool_calls( let mut used_structured_output_tool = false; let mut final_content = None; - for tool_call in tool_calls.iter() { + let mut calls = tool_calls.iter().peekable(); + while let Some(tool_call) = calls.next() { // Stream tool call progress if let Some(stream_event_processor) = ctx.stream_event_processor { let event = StreamingEvent::ToolExecution { @@ -150,15 +150,24 @@ pub async fn execute_tool_calls( ) .await?; } else if tool.module.is_some() { - execute_windmill_tool( - &mut ctx, - tool_call, - tool, - actions, - &mut messages, - final_events_str, - ) - .await?; + let mut batch = vec![(tool_call, tool)]; + while let Some(next_call) = calls.peek() { + if structured_output_tool_name.as_deref() + == Some(next_call.function.name.as_str()) + { + break; + } + let Some(next_tool) = tools.iter().find(|t| { + t.def.function.name == next_call.function.name + && t.mcp_source.is_none() + && t.module.is_some() + }) else { + break; + }; + batch.push((calls.next().unwrap(), next_tool)); + } + execute_windmill_tools(&mut ctx, &batch, actions, &mut messages, final_events_str) + .await?; } else { return Err(Error::internal_err(format!( "Tool type not supported: {}", @@ -313,14 +322,13 @@ async fn execute_mcp_tool_call( } /// Execute a Windmill tool (script or flow) -async fn execute_windmill_tool( - ctx: &mut ToolExecutionContext<'_>, +async fn enqueue_windmill_tool( + ctx: &ToolExecutionContext<'_>, tool_call: &OpenAIToolCall, tool: &Tool, actions: &mut Vec, - messages: &mut Vec, - final_events_str: &mut String, -) -> Result<(), Error> { + reserved: bool, +) -> Result { // Regular Windmill tools must have a module let tool_module = tool.module.as_ref().ok_or_else(|| { Error::internal_err(format!("Tool {} has no module", tool_call.function.name)) @@ -402,8 +410,6 @@ async fn execute_windmill_tool( tool_call_args.insert(key.clone(), result); } - let is_ai_agent_tool = matches!(tool_value, FlowModuleValue::AIAgent { .. }); - let job_payload = match tool_value { FlowModuleValue::Script { path: script_path, hash: script_hash, tag_override, .. } => { script_to_payload( @@ -473,8 +479,6 @@ async fn execute_windmill_tool( )); } let path = format!("{}/tools/{}", ctx.job.runnable_path(), tool_module.id); - // tool jobs are pushed with the parent agent job's tag and executed inline on the - // same worker, so a tag override on a nested agent tool does not apply here JobPayloadWithTag { payload: JobPayload::AIAgent { path }, tag: None, @@ -532,44 +536,52 @@ async fn execute_windmill_tool( false, None, ctx.job.visible_to_owner, - Some(ctx.job.tag.clone()), + job_payload.tag, job_payload.timeout, None, job_priority, job_perms.as_ref(), - true, + reserved, None, None, None, ) .await?; + let mut tx = tx; + if reserved { + // Running ownership reserves the first child; its normal tag allows zombie recovery. + sqlx::query!( + "UPDATE v2_job_queue SET worker = $1 WHERE id = $2", + ctx.worker_name, + uuid, + ) + .execute(&mut *tx) + .await?; + sqlx::query!("UPDATE v2_job_runtime SET ping = now() WHERE id = $1", uuid) + .execute(&mut *tx) + .await?; + } tx.commit().await?; + Ok(uuid) +} - let tool_job = get_mini_pulled_job(ctx.db, &uuid).await?; - - let Some(tool_job) = tool_job else { - return Err(Error::internal_err("Tool job not found".to_string())); - }; - - let tool_job = Arc::new(tool_job); - - let (inner_job_completed_tx, inner_job_completed_rx) = JobCompletedSender::new(ctx.conn, 1); - - let inner_job_completed_rx = inner_job_completed_rx.expect( - "inner_job_completed_tx should be set as agent jobs are not supported on agent workers", - ); +fn spawn_local_tool( + ctx: &ToolExecutionContext<'_>, + tool_job: MiniPulledJob, + reserved: bool, +) -> tokio::task::JoinHandle> { + let mut tool_job = Arc::new(tool_job); // Spawn handle_queued_job on separate task to prevent tokio stack overflow // Clone everything needed for the spawned task - let tool_job_spawn = tool_job.clone(); + let db = ctx.db.clone(); let conn_spawn = ctx.conn.clone(); - let client_spawn = ctx.client.clone(); let hostname_spawn = ctx.hostname.to_string(); let worker_name_spawn = ctx.worker_name.to_string(); let worker_dir_spawn = ctx.worker_dir.to_string(); let base_internal_url_spawn = ctx.base_internal_url.to_string(); - let inner_job_completed_tx_spawn = inner_job_completed_tx.clone(); + let job_completed_tx = ctx.job_completed_tx.clone(); let mut occupancy_metrics_spawn = ctx.occupancy_metrics.clone(); let mut killpill_rx_spawn = ctx.killpill_rx.resubscribe(); @@ -578,147 +590,266 @@ async fn execute_windmill_tool( #[cfg(feature = "benchmark")] let mut bench_spawn = windmill_common::bench::BenchmarkIter::new(); - let job_dir = create_job_dir(&worker_dir_spawn, tool_job_spawn.id).await; + let result = async { + if reserved { + loop { + let queued = get_mini_pulled_job(&db, &tool_job.id) + .await? + .ok_or_else(|| Error::AlreadyCompleted("Tool job already completed".to_string()))?; + tool_job = Arc::new(queued); + sqlx::query!( + "UPDATE v2_job_runtime SET ping = now() WHERE id = $1", + tool_job.id, + ) + .execute(&db) + .await?; + if try_admit_owned_job(&db, &tool_job).await? { + let started_at = sqlx::query_scalar!( + "UPDATE v2_job_queue SET started_at = now() WHERE id = $1 RETURNING started_at", + tool_job.id, + ) + .fetch_optional(&db) + .await? + .flatten(); + Arc::make_mut(&mut tool_job).started_at = started_at; + break; + } + tokio::time::sleep(Duration::from_secs(1)).await; + } + } + let perms = + windmill_common::auth::get_job_perms(&db, &tool_job.id, &tool_job.workspace_id).await?; + let token = windmill_queue::create_token(&db, &tool_job, perms).await; + let client_spawn = AuthedClient::new( + base_internal_url_spawn.clone(), + tool_job.workspace_id.clone(), + token, + None, + ); + let job_dir = create_job_dir(&worker_dir_spawn, tool_job.id).await; - let result = handle_queued_job( - tool_job_spawn, - None, - None, - None, - None, - &conn_spawn, - &client_spawn, - &hostname_spawn, - &worker_name_spawn, - &worker_dir_spawn, - &job_dir, - None, - &base_internal_url_spawn, - inner_job_completed_tx_spawn, - &mut occupancy_metrics_spawn, - &mut killpill_rx_spawn, - None, - None, - #[cfg(feature = "benchmark")] - &mut bench_spawn, - ) + handle_queued_job( + tool_job.clone(), + None, + None, + None, + None, + &conn_spawn, + &client_spawn, + &hostname_spawn, + &worker_name_spawn, + &worker_dir_spawn, + &job_dir, + None, + &base_internal_url_spawn, + job_completed_tx, + &mut occupancy_metrics_spawn, + &mut killpill_rx_spawn, + None, + None, + #[cfg(feature = "benchmark")] + &mut bench_spawn, + ) + .await + } .await; - // Return both result and updated metrics - (result, occupancy_metrics_spawn) + match result { + Err(err) => { + let err_string = format!("{}: {}", err.name(), err); + handle_non_flow_job_error( + &db, + &MiniCompletedJob::from(tool_job), + 0, + None, + err_string, + windmill_common::worker::error_to_value(&err), + &worker_name_spawn, + ) + .await?; + } + Ok(_) => {} + } + Ok(occupancy_metrics_spawn) }); // Register abort handle so the task can be killed on force-cancel let abort_handle = join_handle.abort_handle(); // unwrap safe: lock is only held briefly for push/drain, no panic possible inside ctx.tool_abort_handles.lock().unwrap().push(abort_handle); - - // Await the spawned task - let (handle_result, updated_occupancy) = join_handle.await.map_err(|e| { - if e.is_cancelled() { - Error::ExecutionErr("Tool execution task was cancelled".to_string()) - } else { - Error::internal_err(format!("Tool execution task failed: {}", e)) - } - })?; - - // Merge occupancy metrics back - ctx.occupancy_metrics.total_duration_of_running_jobs = - updated_occupancy.total_duration_of_running_jobs; - - match handle_result { - Err(err) => { - handle_tool_execution_error( - ctx, - tool_call, - tool_module, - &MiniCompletedJob::from(tool_job), - job_id, - err, - messages, - final_events_str, - ) - .await?; - } - Ok(outcome) => { - handle_tool_execution_success( - ctx, - tool_call, - tool_module, - job_id, - outcome.is_success(), - is_ai_agent_tool, - inner_job_completed_rx, - messages, - final_events_str, - ) - .await?; - } - } - - Ok(()) + join_handle } -/// Handle tool execution error -async fn handle_tool_execution_error( +async fn execute_windmill_tools( ctx: &mut ToolExecutionContext<'_>, - tool_call: &OpenAIToolCall, - tool_module: &windmill_common::flows::FlowModule, - tool_job: &MiniCompletedJob, - job_id: Uuid, - err: Error, + batch: &[(&OpenAIToolCall, &Tool)], + actions: &mut Vec, messages: &mut Vec, final_events_str: &mut String, ) -> Result<(), Error> { - let err_string = format!("{}: {}", err.name(), err.to_string()); - let err_json = windmill_common::worker::error_to_value(&err); - let _ = handle_non_flow_job_error( - ctx.db, - tool_job, - 0, - None, - err_string.clone(), - err_json, - ctx.worker_name, - ) - .await; - - let error_message = format!("Error running tool: {}", err_string); - messages.push(OpenAIMessage { - role: "tool".to_string(), - content: Some(OpenAIContent::Text(error_message.clone())), - tool_call_id: Some(tool_call.id.clone()), - agent_action: Some(AgentAction::ToolCall { - job_id, - function_name: tool_call.function.name.clone(), - module_id: tool_module.id.clone(), - }), - ..Default::default() - }); - - // Stream tool result (error case) - if let Some(stream_event_processor) = ctx.stream_event_processor { - let tool_result_event = StreamingEvent::ToolResult { - call_id: tool_call.id.clone(), - function_name: tool_call.function.name.clone(), - result: error_message.clone(), - success: false, - }; - stream_event_processor - .send(tool_result_event, final_events_str) - .await?; + let mut job_ids = Vec::with_capacity(batch.len()); + let mut local = None; + for (index, (call, tool)) in batch.iter().enumerate() { + if index > 0 { + if let Some(processor) = ctx.stream_event_processor { + processor + .send( + StreamingEvent::ToolExecution { + call_id: call.id.clone(), + function_name: call.function.name.clone(), + }, + final_events_str, + ) + .await?; + } + } + job_ids.push(enqueue_windmill_tool(ctx, call, tool, actions, index == 0).await?); + if index == 0 { + let job = get_mini_pulled_job(ctx.db, &job_ids[0]) + .await? + .ok_or_else(|| Error::internal_err("Reserved tool job not found".to_string()))?; + local = Some(spawn_local_tool(ctx, job, true)); + } } - if let Some(parent_job) = ctx.parent_job { - update_flow_status_module_with_actions_success(ctx.db, parent_job, false).await?; + let mut results: HashMap = HashMap::new(); + let mut next_result = 0; + let mut poll = tokio::time::interval(Duration::from_millis(100)); + poll.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + + while next_result < batch.len() || local.is_some() { + tokio::select! { + result = async { local.as_mut().unwrap().await }, if local.is_some() => { + local = None; + let metrics = result.map_err(|e| Error::internal_err(format!("Tool task failed: {e}")))??; + ctx.occupancy_metrics.total_duration_of_running_jobs = metrics.total_duration_of_running_jobs; + ctx.tool_abort_handles.lock().unwrap().retain(|handle| !handle.is_finished()); + } + _ = poll.tick() => {} + } + + let pending: Vec = job_ids[next_result..] + .iter() + .filter(|id| !results.contains_key(id)) + .copied() + .collect(); + if !pending.is_empty() { + let completed = sqlx::query!( + "SELECT id, status = 'success' AS \"success!\", result AS \"result: Json>\" + FROM v2_job_completed WHERE workspace_id = $1 AND id = ANY($2)", + ctx.job.workspace_id, + &pending, + ).fetch_all(ctx.db).await?; + for completed in completed { + let index = job_ids + .iter() + .position(|id| *id == completed.id) + .ok_or_else(|| Error::internal_err("Unexpected tool completion".to_string()))?; + let (call, tool) = batch[index]; + let result = completed + .result + .map(|value| value.0) + .unwrap_or_else(|| to_raw_value(&serde_json::Value::Null)); + let is_agent = tool.module.as_ref().is_some_and(|module| { + matches!(module.get_value(), Ok(FlowModuleValue::AIAgent { .. })) + }); + let content = if is_agent && completed.success { + extract_ai_agent_output(&result).unwrap_or_else(|| result.get().to_string()) + } else { + result.get().to_string() + }; + if let Some(processor) = ctx.stream_event_processor { + processor + .send( + StreamingEvent::ToolResult { + call_id: call.id.clone(), + function_name: call.function.name.clone(), + result: content.clone(), + success: completed.success, + }, + final_events_str, + ) + .await?; + } + results.insert(completed.id, (completed.success, content)); + } + } + + // Transcript rows, model messages, and positional action statuses share call order. + while next_result < batch.len() { + let job_id = job_ids[next_result]; + let Some((success, content)) = results.remove(&job_id) else { + break; + }; + let (call, tool) = batch[next_result]; + let module = tool + .module + .as_ref() + .ok_or_else(|| Error::internal_err("Windmill tool has no module".to_string()))?; + messages.push(OpenAIMessage { + role: "tool".to_string(), + content: Some(OpenAIContent::Text(content.clone())), + tool_call_id: Some(call.id.clone()), + agent_action: Some(AgentAction::ToolCall { + job_id, + function_name: call.function.name.clone(), + module_id: module.id.clone(), + }), + ..Default::default() + }); + if let Some(parent) = ctx.parent_job { + update_flow_status_module_with_actions_success(ctx.db, parent, success).await?; + } + let (content, extras) = windmill_tool_row(call, success, &content); + add_tool_message_to_chat(ctx, Some(job_id), &content, success, Some(extras)).await; + next_result += 1; + } + + if local.is_none() && next_result < batch.len() { + let pending: Vec = job_ids[next_result..] + .iter() + .filter(|id| !results.contains_key(id)) + .copied() + .collect(); + if !pending.is_empty() { + local = claim_local_tool(ctx, &pending).await?; + } + } } - - let (content, extras) = windmill_tool_row(tool_call, false, &error_message); - add_tool_message_to_chat(ctx, Some(job_id), &content, false, Some(extras)).await; - Ok(()) } +async fn claim_local_tool( + ctx: &ToolExecutionContext<'_>, + pending: &[Uuid], +) -> Result>>, Error> { + let query = (String::new(), make_tool_job_pull_query(pending)); + #[cfg(feature = "benchmark")] + let mut bench = windmill_common::bench::BenchmarkIter::new(); + let mut pulled = pull( + ctx.db, + false, + ctx.worker_name, + Some(&query), + #[cfg(feature = "benchmark")] + &mut bench, + ) + .await?; + if let Err(err) = pulled.maybe_apply_debouncing(ctx.db).await { + pulled.error_while_preprocessing = Some(err.to_string()); + } + match pulled.to_pulled_job() { + Ok(job) => Ok(job.map(|job| spawn_local_tool(ctx, job.job, false))), + Err( + windmill_queue::PulledJobResultToJobErr::MissingConcurrencyKey(job) + | windmill_queue::PulledJobResultToJobErr::ErrorWhilePreprocessing(job), + ) => { + ctx.job_completed_tx.send_job(job, true).await?; + Ok(None) + } + } +} + /// Extract the `output` field of an `AIAgentResult` envelope, serialized back to JSON. /// Returns `None` if the result is not a JSON object carrying an `output` field. fn extract_ai_agent_output(result: &RawValue) -> Option { @@ -728,119 +859,6 @@ fn extract_ai_agent_output(result: &RawValue) -> Option { .map(|output| output.get().to_string()) } -/// Handle tool execution success -async fn handle_tool_execution_success( - ctx: &mut ToolExecutionContext<'_>, - tool_call: &OpenAIToolCall, - tool_module: &windmill_common::flows::FlowModule, - job_id: Uuid, - success: bool, - is_ai_agent_tool: bool, - inner_job_completed_rx: JobCompletedReceiver, - messages: &mut Vec, - final_events_str: &mut String, -) -> Result<(), Error> { - let send_result = inner_job_completed_rx.bounded_rx.try_recv().ok(); - - let (result, job_success) = if let Some(SendResult { - result: SendResultPayload::JobCompleted(ref jc), - .. - }) = send_result - { - let result = jc.result.clone(); - // Write tool completion to the DB inline instead of forwarding through - // the parent channel. Forwarding would deadlock for nested agents: the - // sub-tool result would fill the parent's bounded(1) channel, leaving - // no room for the agent's own completion from process_result. - if jc.success { - add_completed_job( - ctx.db, - &jc.job, - true, - false, - sqlx::types::Json(&*jc.result), - jc.result_columns.clone(), - jc.mem_peak, - jc.canceled_by.clone(), - false, - jc.duration, - jc.from_cache.unwrap_or(false), - ) - .await - .map_err(|e| Error::internal_err(format!("Failed to add completed job: {e}")))?; - } else { - let error_value: serde_json::Value = - serde_json::from_str(jc.result.get()).unwrap_or_else(|_| { - serde_json::json!({ "message": format!("Non serializable error: {}", jc.result.get()) }) - }); - add_completed_job_error( - ctx.db, - &jc.job, - jc.mem_peak, - jc.canceled_by.clone(), - error_value, - ctx.worker_name, - false, - jc.duration, - ) - .await - .map_err(|e| Error::internal_err(format!("Failed to add completed job error: {e}")))?; - } - (result, jc.success) - } else { - return Err(Error::internal_err( - "Tool job completed but no result".to_string(), - )); - }; - - // A nested agent returns the whole `AIAgentResult` envelope: on top of `output` it carries - // the child's entire message history, stream log and token usage. Feeding that back would - // grow the caller's context by the child's full transcript on every call, so the caller only - // sees `output`. The envelope stays intact in the tool job's completed row. - let tool_result = if is_ai_agent_tool && job_success { - extract_ai_agent_output(&result).unwrap_or_else(|| result.get().to_string()) - } else { - result.get().to_string() - }; - - messages.push(OpenAIMessage { - role: "tool".to_string(), - content: Some(OpenAIContent::Text(tool_result.clone())), - tool_call_id: Some(tool_call.id.clone()), - agent_action: Some(AgentAction::ToolCall { - job_id, - function_name: tool_call.function.name.clone(), - module_id: tool_module.id.clone(), - }), - ..Default::default() - }); - - let (content, extras) = windmill_tool_row(tool_call, success, &tool_result); - - // The job ran; whether it ran successfully is `success`, and the row stored below is - // worded from it. The stream has to carry the same value, or the card the reader watches - // and the row that replaces it describe the same call differently. - if let Some(stream_event_processor) = ctx.stream_event_processor { - let tool_result_event = StreamingEvent::ToolResult { - call_id: tool_call.id.clone(), - function_name: tool_call.function.name.clone(), - result: tool_result, - success, - }; - stream_event_processor - .send(tool_result_event, final_events_str) - .await?; - } - - if let Some(parent_job) = ctx.parent_job { - update_flow_status_module_with_actions_success(ctx.db, parent_job, success).await?; - } - - add_tool_message_to_chat(ctx, Some(job_id), &content, success, Some(extras)).await; - - Ok(()) -} - /// A Windmill tool's conversation row: worded from the tool, carrying the model's call and /// the exact text the model got back, the same text agent memory keeps for that tool /// message, so a card needs no job fetch. The call is the model's arguments, not the job's @@ -905,9 +923,7 @@ async fn add_tool_message_to_chat( .or(ctx.job.flow_step_id.as_deref()); let step_name = get_step_name_from_flow(ctx.summary.as_deref(), effective_step_id); - // Awaited, not spawned: `created_seq` is the transcript's order, so a round's rows - // must commit in the order of its calls. Calls run one after another; running them - // in parallel would need their rows written in call order all the same. + // created_seq defines transcript order, so these writes must stay sequential. if let Err(e) = add_message_to_conversation( ctx.db, &memory_id, diff --git a/backend/windmill-worker/src/ai_executor.rs b/backend/windmill-worker/src/ai_executor.rs index 3c3a26875b..87c8bdabcf 100644 --- a/backend/windmill-worker/src/ai_executor.rs +++ b/backend/windmill-worker/src/ai_executor.rs @@ -504,6 +504,7 @@ pub async fn handle_ai_agent_job( hostname: &str, killpill_rx: &mut tokio::sync::broadcast::Receiver<()>, has_stream: &mut bool, + job_completed_tx: crate::JobCompletedSender, ) -> Result, Error> { // build_args_map returns None if no $res:/$var: transforms needed, in which case use original args let local_args = match build_args_map(job, client, conn).await? { @@ -1010,6 +1011,7 @@ pub async fn handle_ai_agent_job( omit_output_from_conversation, cancel_rx, tool_abort_handles.clone(), + job_completed_tx, ); let mut occupancy_opt = Some(occupancy_metrics); @@ -1028,13 +1030,10 @@ pub async fn handle_ai_agent_job( cancel_tx, CANCEL_GRACE_PERIOD, ) - .await? + .await }; // agent_fut and update_job are now dropped — borrows on mcp_clients and canceled_by released - // Cleanup MCP clients - cleanup_mcp_clients(mcp_clients).await; - let format_cancel_info = |cb: &Option| { cb.as_ref() .map_or(("unknown".to_string(), "unknown".to_string()), |x| { @@ -1045,38 +1044,42 @@ pub async fn handle_ai_agent_job( }) }; - match outcome { - GracefulPollOutcome::Ok(result) => Ok(result), - GracefulPollOutcome::Timeout(ms) => { - tracing::error!("AI agent timeout after {}s", ms / 1000); - Err(Error::ExecutionErr(format!( - "AI agent timeout after (>{}s)", - ms / 1000 - ))) - } - GracefulPollOutcome::Cancelled { canceled_by: cb } => { - let (by, reason) = format_cancel_info(&cb); - Err(Error::ExecutionErr(format!( - "Job cancelled by {by} (reason: {reason})" - ))) - } - GracefulPollOutcome::CancelledTimeout { canceled_by: cb } => { - let (by, reason) = format_cancel_info(&cb); - // Abort any still-running spawned tool tasks - // unwrap safe: lock is only held briefly for push/drain, no panic possible inside - for handle in tool_abort_handles.lock().unwrap().drain(..) { - handle.abort(); + let result = match outcome { + Err(error) => Err(error), + Ok(outcome) => match outcome { + GracefulPollOutcome::Ok(result) => Ok(result), + GracefulPollOutcome::Timeout(ms) => { + tracing::error!("AI agent timeout after {}s", ms / 1000); + Err(Error::ExecutionErr(format!( + "AI agent timeout after (>{}s)", + ms / 1000 + ))) } - // Hard timeout: clean up orphaned jobs still stuck in v2_job_queue - cleanup_orphaned_tool_jobs(db, &job.id, &job.workspace_id, cb).await; - Err(Error::ExecutionErr(format!( - "Job cancelled by {by} (reason: {reason}, timed out waiting for tool calls)" - ))) - } - GracefulPollOutcome::AlreadyCompleted => { - Err(Error::AlreadyCompleted("Job already completed".to_string())) + GracefulPollOutcome::Cancelled { canceled_by: cb } => { + let (by, reason) = format_cancel_info(&cb); + Err(Error::ExecutionErr(format!( + "Job cancelled by {by} (reason: {reason})" + ))) + } + GracefulPollOutcome::CancelledTimeout { canceled_by: cb } => { + let (by, reason) = format_cancel_info(&cb); + Err(Error::ExecutionErr(format!( + "Job cancelled by {by} (reason: {reason}, timed out waiting for tool calls)" + ))) + } + GracefulPollOutcome::AlreadyCompleted => { + Err(Error::AlreadyCompleted("Job already completed".to_string())) + } + }, + }; + if result.is_err() { + for handle in tool_abort_handles.lock().unwrap().drain(..) { + handle.abort(); } + cleanup_orphaned_tool_jobs(db, &job.id, &job.workspace_id, canceled_by.clone()).await; } + cleanup_mcp_clients(mcp_clients).await; + result } /// OpenAI rejects a `prompt_cache_key` over 64 characters @@ -1129,6 +1132,7 @@ pub async fn run_agent( // abort handles for spawned tool tasks tool_abort_handles: ToolAbortHandles, + job_completed_tx: crate::JobCompletedSender, ) -> error::Result> { let output_type = args.output_type.as_ref().unwrap_or(&OutputType::Text); let credentials = args.provider.to_provider_credentials(db).await?; @@ -1859,6 +1863,7 @@ pub async fn run_agent( previous_result: &previous_result, id_context: &id_context, tool_abort_handles: tool_abort_handles.clone(), + job_completed_tx: job_completed_tx.clone(), }; let (tool_messages, tool_content, tool_used_structured_output) = @@ -2758,11 +2763,13 @@ async fn cleanup_orphaned_tool_jobs( ) }); - // Find direct child jobs still in v2_job_queue (agent tool jobs are always direct children) let orphaned_ids: Vec = match sqlx::query_scalar!( - r#"SELECT j.id FROM v2_job j - JOIN v2_job_queue q ON q.id = j.id - WHERE j.parent_job = $1 AND j.workspace_id = $2"#, + r#"WITH RECURSIVE descendants AS ( + SELECT id FROM v2_job WHERE parent_job = $1 AND workspace_id = $2 + UNION ALL + SELECT j.id FROM v2_job j JOIN descendants d ON j.parent_job = d.id + WHERE j.workspace_id = $2 + ) SELECT d.id AS "id!" FROM descendants d JOIN v2_job_queue q ON q.id = d.id"#, parent_job_id, w_id, ) diff --git a/backend/windmill-worker/src/worker.rs b/backend/windmill-worker/src/worker.rs index 55a698a9dd..3cc1264970 100644 --- a/backend/windmill-worker/src/worker.rs +++ b/backend/windmill-worker/src/worker.rs @@ -4973,6 +4973,7 @@ pub async fn handle_queued_job( hostname, killpill_rx, &mut has_stream, + job_completed_tx.clone(), )) .await }