fix(workflow): flush agent events before stage completion

This commit is contained in:
Bryan Helmkamp 2026-07-23 15:52:25 -04:00
parent 7ff153d222
commit c9b5303128
No known key found for this signature in database
2 changed files with 81 additions and 7 deletions

View file

@ -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<Emitter>,
file_tracking: Arc<Mutex<FileTracking>>,
) {
) -> 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();

View file

@ -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"
);
}
// ---------------------------------------------------------------------------