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-52f9582423689f098a337f88341239a136ee5132bc97f4b06cc1ee5beaa621e5.json b/backend/.sqlx/query-52f9582423689f098a337f88341239a136ee5132bc97f4b06cc1ee5beaa621e5.json deleted file mode 100644 index cc6a57c005..0000000000 --- a/backend/.sqlx/query-52f9582423689f098a337f88341239a136ee5132bc97f4b06cc1ee5beaa621e5.json +++ /dev/null @@ -1,41 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT c.id IS NOT NULL AS \"completed!\",\n c.result AS \"result: sqlx::types::Json>\",\n c.status = 'success' AS \"success\",\n EXISTS(SELECT 1 FROM v2_job_queue WHERE id = $1) AS \"queued!\"\n FROM (SELECT 1) one\n LEFT JOIN v2_job_completed c ON c.id = $1 AND c.workspace_id = $2", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "completed!", - "type_info": "Bool" - }, - { - "ordinal": 1, - "name": "result: sqlx::types::Json>", - "type_info": "Jsonb" - }, - { - "ordinal": 2, - "name": "success", - "type_info": "Bool" - }, - { - "ordinal": 3, - "name": "queued!", - "type_info": "Bool" - } - ], - "parameters": { - "Left": [ - "Uuid", - "Text" - ] - }, - "nullable": [ - null, - true, - null, - null - ] - }, - "hash": "52f9582423689f098a337f88341239a136ee5132bc97f4b06cc1ee5beaa621e5" -} 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-f5c87b6fc1decd6dc53590fe4d16f59430530b8b7428e42d1777927e63395ff8.json b/backend/.sqlx/query-f5c87b6fc1decd6dc53590fe4d16f59430530b8b7428e42d1777927e63395ff8.json new file mode 100644 index 0000000000..197ecaf3d9 --- /dev/null +++ b/backend/.sqlx/query-f5c87b6fc1decd6dc53590fe4d16f59430530b8b7428e42d1777927e63395ff8.json @@ -0,0 +1,16 @@ +{ + "db_name": "PostgreSQL", + "query": "UPDATE v2_job_queue SET worker = $1, running = true, started_at = now()\n WHERE id = $2 AND tag = ANY($3)", + "describe": { + "columns": [], + "parameters": { + "Left": [ + "Varchar", + "Uuid", + "TextArray" + ] + }, + "nullable": [] + }, + "hash": "f5c87b6fc1decd6dc53590fe4d16f59430530b8b7428e42d1777927e63395ff8" +} diff --git a/backend/tests/ai_tool_queue.rs b/backend/tests/ai_tool_queue.rs new file mode 100644 index 0000000000..4d39176cc9 --- /dev/null +++ b/backend/tests/ai_tool_queue.rs @@ -0,0 +1,220 @@ +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, + unsupported: 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 unsupported && i != 1 { "unserved-tool-tag" } else { "bun" }, + "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","timeout":if unsupported {Some(3)} else {None},"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; + if unsupported { + assert!( + !result.success, + "agent must time out waiting for unsupported tools" + ); + let children: Vec<(String, bool, bool)> = sqlx::query_as( + "SELECT j.runnable_path, c.status = 'success', c.started_at IS NULL + FROM v2_job j JOIN v2_job_completed c USING(id) + WHERE j.parent_job IN (SELECT id FROM v2_job WHERE parent_job = $1) + ORDER BY j.runnable_path", + ) + .bind(id) + .fetch_all(&db) + .await?; + assert_eq!(children.len(), 3); + assert_eq!( + children.iter().map(|c| (c.1, c.2)).collect::>(), + vec![(false, true), (true, false), (false, true)], + "neither the first-tool reservation nor fallback may execute unsupported tags" + ); + stub.abort(); + server.close().await?; + return Ok(()); + } + 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, 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, 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, false).await +} + +#[sqlx::test(fixtures("base"))] +#[serial_test::serial] +async fn parent_leaves_unsupported_first_and_remaining_tools_queued( + db: Pool, +) -> anyhow::Result<()> { + run_batch(db, false, false, true).await +} diff --git a/backend/windmill-common/src/worker.rs b/backend/windmill-common/src/worker.rs index d480a89851..d661483182 100644 --- a/backend/windmill-common/src/worker.rs +++ b/backend/windmill-common/src/worker.rs @@ -864,6 +864,24 @@ pub fn make_pull_query(tags: &[String]) -> String { query } +/// Claim a parent's tool jobs matching its worker tags. 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], tags: &[String]) -> 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(", "); + let tags = tags + .iter() + .map(|tag| format!("'{}'", tag.replace('\'', "''"))) + .join(", "); + format_pull_query(format!( + "SELECT id FROM v2_job_queue + WHERE running = false AND id = ANY(ARRAY[{ids}]::uuid[]) AND scheduled_for <= now() + AND tag = ANY(ARRAY[{tags}]::text[]) + 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 b555e1ea92..c2c42c649c 100644 --- a/backend/windmill-queue/src/jobs.rs +++ b/backend/windmill-queue/src/jobs.rs @@ -4909,6 +4909,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 682714292e..930a329be4 100644 --- a/backend/windmill-worker/src/ai/tools.rs +++ b/backend/windmill-worker/src/ai/tools.rs @@ -7,16 +7,15 @@ use crate::ai::utils::{ use crate::common::OccupancyMetrics; 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, + evaluate_input_transform, raw_script_to_payload, resolve_flow_step_tag, script_to_payload, + JobPayloadWithTag, }; +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 +35,12 @@ 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, WORKER_CONFIG}, }; use windmill_queue::{ - add_completed_job, add_completed_job_error, check_tag_available_for_push, get_mini_pulled_job, - push, resolve_push_tag, MiniCompletedJob, MiniPulledJob, PushArgs, PushIsolationLevel, + append_logs, check_tag_available_for_push, get_mini_pulled_job, pull, push, resolve_push_tag, + tag_reads_flow_expr, try_admit_owned_job, MiniCompletedJob, MiniPulledJob, PushArgs, + PushIsolationLevel, }; /// Shared collection of abort handles for spawned tool tasks. @@ -74,6 +74,8 @@ pub struct ToolExecutionContext<'a> { pub stream_event_processor: Option<&'a StreamEventProcessor>, pub flow_context: &'a mut FlowContext, pub omit_output_from_conversation: bool, + /// The agent's flow `preserve_step_tags`: its tools pick their tag as that flow's steps do. + pub preserve_step_tags: bool, /// The thinking that led to this round's calls, stored on the first tool row written. /// None when the round wrote text, whose row carries it. pub reasoning: Option, @@ -82,6 +84,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 +101,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 +154,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 +326,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<(Uuid, bool), Error> { // 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 +414,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,7 +483,6 @@ async fn execute_windmill_tool( )); } let path = format!("{}/tools/{}", ctx.job.runnable_path(), tool_module.id); - // a nested agent always runs inline on this worker, under this agent's tag JobPayloadWithTag { payload: JobPayload::AIAgent { path }, tag: None, @@ -502,18 +511,34 @@ async fn execute_windmill_tool( let push_args = PushArgs { args: &tool_call_args, extra: None }; - // A tool with a tag of its own naming another queue than this agent's must run on a worker - // that tag selects, and like any tag the caller chose it must pass the workspace's custom tag - // restrictions. It runs inline when this worker serves that queue: the agent holds this - // worker while it waits, so dispatching to a queue it serves could wait on itself. - let mut resolved_tool_tag = None; - if let Some(tag) = job_payload.tag.as_deref() { - resolved_tool_tag = resolve_push_tag(tag, &push_args, &ctx.job.workspace_id, ctx.db) - .await - .filter(|resolved| *resolved != ctx.job.tag); + // The agent's tag stands in for the flow's, as a sub-flow's does for its own steps. + let tool_tag = resolve_flow_step_tag( + false, + &ctx.job.tag, + &ctx.job.workspace_id, + ctx.preserve_step_tags, + job_payload.tag.as_deref(), + ); + if tool_tag.as_deref().is_some_and(tag_reads_flow_expr) { + append_logs( + &ctx.job.id, + &ctx.job.workspace_id, + format!( + "Tool '{}' is tagged with `$flow_expr[...]`, which AI agent tools do not resolve, so it runs on its default tag.\n", + tool_call.function.name + ), + ctx.conn, + ) + .await; } - let tool_tag = match (job_payload.tag.as_deref(), resolved_tool_tag.as_ref()) { - (Some(tag), Some(_)) => { + + // A tool's own tag routes its job, so like any tag the caller chose it must pass the + // workspace's custom tag restrictions. The agent's tag was already checked when it was pushed. + if let Some(tag) = tool_tag.as_deref() { + if resolve_push_tag(tag, &push_args, &ctx.job.workspace_id, ctx.db) + .await + .is_some_and(|resolved| resolved != ctx.job.tag) + { let is_super_admin = windmill_common::auth::is_super_admin_email(ctx.db, email).await?; check_tag_available_for_push( ctx.db, @@ -531,16 +556,8 @@ async fn execute_windmill_tool( )), e => e, })?; - Some(tag.to_string()) } - _ => None, - }; - let run_inline = resolved_tool_tag.is_none_or(|resolved| { - windmill_common::worker::WORKER_CONFIG - .load() - .worker_tags - .contains(&resolved) - }); + } let mut tx = ctx.db.begin().await?; @@ -573,77 +590,71 @@ async fn execute_windmill_tool( false, None, ctx.job.visible_to_owner, - Some(tool_tag.unwrap_or_else(|| ctx.job.tag.clone())), + tool_tag, job_payload.timeout, None, job_priority, job_perms.as_ref(), - run_inline, + false, None, None, None, ) .await?; - tx.commit().await?; - - if !run_inline { - // Like an inline tool's failure, a lost wait reaches the model rather than failing the - // agent. - let (result, success) = - match wait_for_dispatched_tool_job(ctx.db, &uuid, &ctx.job.workspace_id).await { - Ok(outcome) => outcome, - Err(e) => ( - to_raw_value(&serde_json::json!({ "error": { "message": e.to_string() } })), - false, - ), - }; - return report_tool_result( - ctx, - tool_call, - tool_module, - job_id, - &result, - success, - is_ai_agent_tool && success, - messages, - final_events_str, + let mut tx = tx; + let reserved = if reserved { + let tags = local_tool_tags(); + // Reserve a supported first tool before commit so other workers cannot race it. + let claimed = sqlx::query!( + "UPDATE v2_job_queue SET worker = $1, running = true, started_at = now() + WHERE id = $2 AND tag = ANY($3)", + ctx.worker_name, + uuid, + &tags, ) - .await; - } - - 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())); + .execute(&mut *tx) + .await? + .rows_affected() + == 1; + if claimed { + sqlx::query!("UPDATE v2_job_runtime SET ping = now() WHERE id = $1", uuid) + .execute(&mut *tx) + .await?; + } + claimed + } else { + false }; + tx.commit().await?; + Ok((uuid, reserved)) +} - // The tool gets a token of its own, as it would on any worker: the agent's token would - // identify the tool as the agent job (its OIDC path and job id). - let tool_client = AuthedClient::new( - ctx.base_internal_url.to_string(), - tool_job.workspace_id.clone(), - windmill_queue::create_token(ctx.db, &tool_job, None).await, - None, - ); - let tool_job = Arc::new(tool_job); +fn local_tool_tags() -> Vec { + WORKER_CONFIG + .load() + .priority_tags_sorted + .iter() + .flat_map(|group| group.tags.iter().cloned()) + .collect() +} - 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 = tool_client; 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(); @@ -652,147 +663,270 @@ 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?; + } + } + let (id, reserved) = enqueue_windmill_tool(ctx, call, tool, actions, index == 0).await?; + job_ids.push(id); + if reserved { + let job = get_mini_pulled_job(ctx.db, &id) + .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, &local_tool_tags()), + ); + #[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 { @@ -802,200 +936,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(), - )); - }; - - report_tool_result( - ctx, - tool_call, - tool_module, - job_id, - &result, - success, - is_ai_agent_tool && job_success, - messages, - final_events_str, - ) - .await -} - -/// Wait for a tool job another worker runs; returns its result and whether it succeeded. -/// Canceling the agent cancels this job along with it, as a child of the agent's job. -async fn wait_for_dispatched_tool_job( - db: &DB, - id: &Uuid, - w_id: &str, -) -> Result<(Box, bool), Error> { - const MAX_CONSECUTIVE_POLL_ERRORS: u32 = 10; - let mut interval = std::time::Duration::from_millis(50); - let mut poll_errors = 0; - loop { - // One statement reads both tables from one snapshot, so a job completing between two - // reads cannot look like one that vanished. - let state = match sqlx::query!( - "SELECT c.id IS NOT NULL AS \"completed!\", - c.result AS \"result: sqlx::types::Json>\", - c.status = 'success' AS \"success\", - EXISTS(SELECT 1 FROM v2_job_queue WHERE id = $1) AS \"queued!\" - FROM (SELECT 1) one - LEFT JOIN v2_job_completed c ON c.id = $1 AND c.workspace_id = $2", - id, - w_id - ) - .fetch_one(db) - .await - { - Ok(state) => { - poll_errors = 0; - state - } - Err(e) if poll_errors < MAX_CONSECUTIVE_POLL_ERRORS => { - poll_errors += 1; - tracing::warn!("polling tool job {id} failed, retrying: {e}"); - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - continue; - } - Err(e) => return Err(e.into()), - }; - if state.completed { - let result = state - .result - .map(|r| r.0) - .unwrap_or_else(|| to_raw_value(&serde_json::Value::Null)); - return Ok((result, state.success.unwrap_or(false))); - } - if !state.queued { - return Err(Error::internal_err(format!( - "tool job {id} is neither queued nor completed" - ))); - } - tokio::time::sleep(interval).await; - interval = std::cmp::min(interval * 2, std::time::Duration::from_secs(1)); - } -} - -/// Hand a finished tool job's result to the model, the stream, the flow status and the chat. -async fn report_tool_result( - ctx: &mut ToolExecutionContext<'_>, - tool_call: &OpenAIToolCall, - tool_module: &windmill_common::flows::FlowModule, - job_id: Uuid, - result: &RawValue, - success: bool, - is_ai_agent_output: bool, - messages: &mut Vec, - final_events_str: &mut String, -) -> Result<(), Error> { - // 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_output { - 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 @@ -1060,9 +1000,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 d7ded3652f..958caeaefc 100644 --- a/backend/windmill-worker/src/ai_executor.rs +++ b/backend/windmill-worker/src/ai_executor.rs @@ -607,6 +607,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? { @@ -686,6 +687,7 @@ pub async fn handle_ai_agent_job( }; let value = flow_data.value(); + let preserve_step_tags = value.preserve_step_tags; let module = if direct_parent_job_kind == JobKind::AIAgent { let parent_agent_step_id = direct_parent_job_flow_step_id.as_deref().ok_or_else(|| { @@ -1111,8 +1113,10 @@ pub async fn handle_ai_agent_job( has_stream, has_websearch, omit_output_from_conversation, + preserve_step_tags, cancel_rx, tool_abort_handles.clone(), + job_completed_tx, ); let mut occupancy_opt = Some(occupancy_metrics); @@ -1131,13 +1135,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| { @@ -1148,43 +1149,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); - // A tool dispatched to another worker group would otherwise still run later, after - // the agent that wanted its result is gone. - let reason = format!("parent AI agent {} timed out", job.id); - let cb = CanceledBy { username: None, reason: Some(reason) }; - cleanup_orphaned_tool_jobs(db, &job.id, &job.workspace_id, Some(cb)).await; - 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 @@ -1231,12 +1231,14 @@ pub async fn run_agent( has_stream: &mut bool, has_websearch: bool, omit_output_from_conversation: bool, + preserve_step_tags: bool, // cancellation signal from parent cancel_rx: tokio::sync::watch::Receiver, // 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?; @@ -2105,6 +2107,7 @@ pub async fn run_agent( stream_event_processor: stream_event_processor.as_ref(), flow_context: &mut flow_context, omit_output_from_conversation, + preserve_step_tags, reasoning: if structured_output_first { None } else { @@ -2113,6 +2116,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) = @@ -3093,11 +3097,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 f45c295b7a..ce3fe26545 100644 --- a/backend/windmill-worker/src/worker.rs +++ b/backend/windmill-worker/src/worker.rs @@ -4990,6 +4990,7 @@ pub async fn handle_queued_job( hostname, killpill_rx, &mut has_stream, + job_completed_tx.clone(), )) .await } diff --git a/backend/windmill-worker/src/worker_flow.rs b/backend/windmill-worker/src/worker_flow.rs index a3257ace36..7c1ea23de4 100644 --- a/backend/windmill-worker/src/worker_flow.rs +++ b/backend/windmill-worker/src/worker_flow.rs @@ -3239,7 +3239,7 @@ struct PushNextFlowJobRec { /// - when the flow opts into `preserve_step_tags` and the child declares its own non-empty tag, /// that tag is honored instead of being overridden by the flow tag; /// - otherwise the child inherits the parent flow job's tag. -fn resolve_flow_step_tag( +pub(crate) fn resolve_flow_step_tag( is_preprocessor_step: bool, flow_tag: &str, workspace_id: &str,