mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-07 03:00:29 +00:00
fix(workflow): flush agent events before stage completion
This commit is contained in:
parent
7ff153d222
commit
c9b5303128
2 changed files with 81 additions and 7 deletions
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue