diff --git a/lib/crates/fabro-workflow/src/handler/llm/api.rs b/lib/crates/fabro-workflow/src/handler/llm/api.rs index 2b0ffc1df..93ee9fed8 100644 --- a/lib/crates/fabro-workflow/src/handler/llm/api.rs +++ b/lib/crates/fabro-workflow/src/handler/llm/api.rs @@ -24,7 +24,7 @@ use fabro_model::{AgentProfileKind, Catalog, FallbackTarget, ModelRef, ProviderI use fabro_types::settings::run::RunModelControls; use fabro_types::{PermissionLevel, RunId, SessionCapability, StageId, StageTiming}; use serde::de::DeserializeOwned; -use tokio::sync::Mutex as TokioMutex; +use tokio::sync::{Mutex as TokioMutex, mpsc}; use tokio::task::JoinHandle; use tokio_util::sync::CancellationToken; @@ -532,16 +532,50 @@ fn emit_agent_tools_available( /// Spawn a task that subscribes to session events and: /// 1. Tracks file changes (write_file/edit_file tool calls) into shared state. /// 2. Forwards non-streaming agent events to the pipeline emitter. +/// +/// The returned handle exposes a per-input barrier. A successful +/// `process_input_with_runtime` emits `ProcessingEnd` after all events for +/// that input, so waiting for the barrier keeps terminal stage events from +/// overtaking queued agent events. +struct EventForwarder { + processing_end_rx: mpsc::UnboundedReceiver<()>, + task: JoinHandle<()>, +} + +impl EventForwarder { + async fn wait_for_processing_end(&mut self) { + if self.processing_end_rx.recv().await.is_none() { + tracing::warn!("Agent event forwarder stopped before processing input events"); + } + } + + fn abort(&self) { + self.task.abort(); + } +} + +impl Drop for EventForwarder { + fn drop(&mut self) { + self.task.abort(); + } +} + fn spawn_event_forwarder( session: &Session, node_id: String, scope: StageScope, emitter: Arc, file_tracking: Arc>, -) { +) -> EventForwarder { let mut rx = session.subscribe(); - tokio::spawn(async move { + let root_session_id = session.id().to_string(); + let (processing_end_tx, processing_end_rx) = mpsc::unbounded_channel(); + let task = tokio::spawn(async move { while let Ok(event) = rx.recv().await { + let is_root_processing_end = event.session_id == root_session_id + && event.parent_session_id.is_none() + && matches!(&event.event, AgentEvent::ProcessingEnd); + // Reset watchdog on every event, including streaming deltas emitter.touch(); @@ -573,8 +607,17 @@ fn spawn_event_forwarder( &scope, ); } + + if is_root_processing_end { + let _ = processing_end_tx.send(()); + } } }); + + EventForwarder { + processing_end_rx, + task, + } } /// LLM backend that delegates to an `agent` Session per invocation. @@ -1213,7 +1256,7 @@ impl CodergenBackend for AgentApiBackend { let stage_scope = StageScope::for_handler(context, &node.id); // Subscribe to session events: forward to pipeline emitter + track files. - spawn_event_forwarder( + let mut event_forwarder = spawn_event_forwarder( &session, node.id.clone(), stage_scope.clone(), @@ -1284,6 +1327,7 @@ impl CodergenBackend for AgentApiBackend { inference_duration = inference_duration.saturating_add(timing.inference); tool_duration = tool_duration.saturating_add(timing.tool); if process_result.is_ok() { + event_forwarder.wait_for_processing_end().await; total_usage += session.last_input_usage(); UsdMicros::accumulate(&mut total_cost, session.last_input_cost()); } @@ -1315,6 +1359,7 @@ impl CodergenBackend for AgentApiBackend { let mut succeeded = false; bridge.abort(); + event_forwarder.abort(); discard_session(&mut session, &mut lease, emitter); for (index, target) in self.fallback_chain.iter().enumerate() { @@ -1368,7 +1413,7 @@ impl CodergenBackend for AgentApiBackend { bridge.replace(cancel_token.clone(), &session); // Re-subscribe to forward events + track files from the new session - spawn_event_forwarder( + event_forwarder = spawn_event_forwarder( &session, node.id.clone(), stage_scope.clone(), @@ -1420,6 +1465,7 @@ impl CodergenBackend for AgentApiBackend { tool_duration = tool_duration.saturating_add(timing.tool); match process_result { Ok(()) => { + event_forwarder.wait_for_processing_end().await; total_usage += session.last_input_usage(); UsdMicros::accumulate(&mut total_cost, session.last_input_cost()); succeeded = true; @@ -1492,6 +1538,7 @@ impl CodergenBackend for AgentApiBackend { tool_duration = tool_duration.saturating_add(timing.tool); match repair_result { Ok(()) => { + event_forwarder.wait_for_processing_end().await; total_usage += session.last_input_usage(); UsdMicros::accumulate(&mut total_cost, session.last_input_cost()); repair_attempts += 1; @@ -1534,6 +1581,7 @@ impl CodergenBackend for AgentApiBackend { // Collect files_touched from the shared tracking state. let (files_touched, last_file_touched) = file_tracking_snapshot(&file_tracking); + drop(event_forwarder); if let Some(lease) = lease.take() { lease.release(); diff --git a/lib/crates/fabro-workflow/tests/it/integration.rs b/lib/crates/fabro-workflow/tests/it/integration.rs index 2b817434a..71a1e1103 100644 --- a/lib/crates/fabro-workflow/tests/it/integration.rs +++ b/lib/crates/fabro-workflow/tests/it/integration.rs @@ -2319,7 +2319,7 @@ reasoning = false ); } -#[tokio::test] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn workflow_persists_authoritative_openrouter_cost_for_agent_stage() { use fabro_auth::EnvCredentialSource; use fabro_workflow::steering_hub::SteeringHub; @@ -2398,8 +2398,18 @@ base_url = "{}" registry.register("start", Box::new(StartHandler)); registry.register("exit", Box::new(ExitHandler)); + let events = Arc::new(std::sync::Mutex::new(Vec::new())); + let events_for_listener = Arc::clone(&events); + let emitter = Arc::new(Emitter::default()); + emitter.on_event(move |event| { + if event.event_name() == "agent.message" { + std::thread::sleep(Duration::from_millis(50)); + } + events_for_listener.lock().unwrap().push(event.clone()); + }); + let dir = tempfile::tempdir().unwrap(); - let engine = WorkflowRunner::new(registry, Arc::new(Emitter::default()), local_env()); + let engine = WorkflowRunner::new(registry, emitter, local_env()); let run_options = RunOptions { settings: WorkflowSettings::default(), run_dir: dir.path().to_path_buf(), @@ -2432,6 +2442,22 @@ base_url = "{}" Some(AUTHORITATIVE_COST_USD_MICROS), "provider-reported usage.cost should override the catalog estimate" ); + + let events = events.lock().unwrap(); + let agent_message = events + .iter() + .position(|event| event.event_name() == "agent.message") + .expect("agent message should be emitted"); + let stage_completed = events + .iter() + .position(|event| { + event.event_name() == "stage.completed" && event.node_id.as_deref() == Some("work") + }) + .expect("work stage completion should be emitted"); + assert!( + agent_message < stage_completed, + "agent messages must be forwarded before terminal stage events" + ); } // ---------------------------------------------------------------------------