From ab68fc27d10412e9d827d07a62374ccc3cd12f12 Mon Sep 17 00:00:00 2001 From: Bryan Helmkamp <19+brynary@users.noreply.github.com> Date: Tue, 26 May 2026 21:19:53 -0400 Subject: [PATCH] fix(workflow): preserve usage across session compaction (#420) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Summary Fixes shared-thread workflow stages that compact their session before routing/audit bookkeeping finishes. The workflow backend now records token usage from each `Session::process_input` call as it happens, instead of slicing assistant turns out of the final session history after the session may have been compacted or replaced. ## Changes - Track per-input token usage inside `fabro-agent::Session` alongside the existing timing data. - Use the recorded per-input usage in the workflow LLM backend for initial prompts, retry-after-compaction prompts, and schema repair prompts. - Keep the invariant panic message for inconsistent session history explicit with `expect(...)`. - Add a black-box workflow integration test that drives a shared-thread audit through pre-routing compaction and asserts the audit still succeeds. ## Verification - `ulimit -n 4096 && cargo nextest run --workspace` - `cargo +nightly-2026-04-14 fmt --check --all` - `cargo +nightly-2026-04-14 clippy --workspace --all-targets -- -D warnings` - `cd apps/fabro-web && bun test --isolate` - `cd apps/fabro-web && bun run typecheck` - `cd lib/packages/fabro-api-client && bun run typecheck` - After rebasing onto current `origin/main`: `ulimit -n 4096 && cargo nextest run -p fabro-workflow --test it integration::shared_thread_compaction_before_routing_audit_succeeds` --- [![Compound Engineering](https://img.shields.io/badge/Compound_Engineering-6366f1)](https://github.com/EveryInc/compound-engineering-plugin) 🤖 Generated with GPT-5 via [Codex](https://openai.com/codex) --- lib/crates/fabro-agent/src/session.rs | 18 +- .../fabro-workflow/src/handler/llm/api.rs | 18 +- .../fabro-workflow/tests/it/integration.rs | 235 +++++++++++++++++- 3 files changed, 255 insertions(+), 16 deletions(-) diff --git a/lib/crates/fabro-agent/src/session.rs b/lib/crates/fabro-agent/src/session.rs index 38d14a63c..2bd840bda 100644 --- a/lib/crates/fabro-agent/src/session.rs +++ b/lib/crates/fabro-agent/src/session.rs @@ -9,7 +9,7 @@ use fabro_llm::generate::StreamAccumulator; use fabro_llm::provider::StreamEventStream; use fabro_llm::types::{ ContentPart, Message as LlmMessage, ReasoningEffort, Request, RetryPolicy, StreamEvent, - ToolChoice, + TokenCounts, ToolChoice, }; use fabro_llm::{Error as LlmError, retry}; use fabro_mcp::config::{McpServerSettings, McpTransport}; @@ -353,6 +353,7 @@ pub struct Session { subagent_manager: Option>>, completion_coordinator: Option>, last_input_timing: SessionInputTiming, + last_input_usage: TokenCounts, } impl Session { @@ -391,6 +392,7 @@ impl Session { subagent_manager, completion_coordinator: None, last_input_timing: SessionInputTiming::default(), + last_input_usage: TokenCounts::default(), } } @@ -1215,6 +1217,11 @@ impl Session { self.last_input_timing } + #[must_use] + pub fn last_input_usage(&self) -> TokenCounts { + self.last_input_usage.clone() + } + /// Process an input. The inference/tool timing accumulated during the call /// is available via [`Self::last_input_timing`] after this returns, even on /// error. @@ -1224,7 +1231,9 @@ impl Session { agent_tool_runtime: AgentToolRuntime, ) -> Result<(), Error> { let mut timing = SessionInputTiming::default(); + let mut usage = TokenCounts::default(); self.last_input_timing = timing; + self.last_input_usage = TokenCounts::default(); if self.state == SessionState::Closed { return Err(Error::SessionClosed); } @@ -1249,7 +1258,7 @@ impl Session { // Process the initial input, then drain any followups let mut result = self - .run_single_input(input, &agent_tool_runtime, &mut timing) + .run_single_input(input, &agent_tool_runtime, &mut timing, &mut usage) .await; if result.is_ok() { @@ -1261,7 +1270,7 @@ impl Session { .pop_front(); let Some(followup) = followup else { break }; result = self - .run_single_input(&followup, &agent_tool_runtime, &mut timing) + .run_single_input(&followup, &agent_tool_runtime, &mut timing, &mut usage) .await; if result.is_err() { break; @@ -1280,6 +1289,7 @@ impl Session { } self.last_input_timing = timing; + self.last_input_usage = usage; result } @@ -1288,6 +1298,7 @@ impl Session { input: &str, agent_tool_runtime: &AgentToolRuntime, timing: &mut SessionInputTiming, + usage_accumulator: &mut TokenCounts, ) -> Result<(), Error> { const STREAM_CONSUME_RETRIES: usize = 3; @@ -1692,6 +1703,7 @@ impl Session { &local_context_window, &usage, )); + *usage_accumulator += usage.clone(); self.history.push(Message::Assistant { content: text.clone(), diff --git a/lib/crates/fabro-workflow/src/handler/llm/api.rs b/lib/crates/fabro-workflow/src/handler/llm/api.rs index e00a7c255..91753f4c9 100644 --- a/lib/crates/fabro-workflow/src/handler/llm/api.rs +++ b/lib/crates/fabro-workflow/src/handler/llm/api.rs @@ -1213,8 +1213,7 @@ impl CodergenBackend for AgentApiBackend { Arc::clone(&file_tracking), ); - // Record turn count before processing so we only aggregate new usage. - let mut turns_before = session.history().turns().len(); + let mut total_usage = TokenCounts::default(); let mut inference_duration = Duration::ZERO; let mut tool_duration = Duration::ZERO; @@ -1275,6 +1274,9 @@ impl CodergenBackend for AgentApiBackend { let timing = session.last_input_timing(); inference_duration = inference_duration.saturating_add(timing.inference); tool_duration = tool_duration.saturating_add(timing.tool); + if process_result.is_ok() { + total_usage += session.last_input_usage(); + } process_result } Err(err) => Err(err), @@ -1354,7 +1356,6 @@ impl CodergenBackend for AgentApiBackend { }; session = new_session; bridge.replace(cancel_token.clone(), &session); - turns_before = session.history().turns().len(); // Re-subscribe to forward events + track files from the new session spawn_event_forwarder( @@ -1409,6 +1410,7 @@ impl CodergenBackend for AgentApiBackend { tool_duration = tool_duration.saturating_add(timing.tool); match process_result { Ok(()) => { + total_usage += session.last_input_usage(); succeeded = true; break; } @@ -1479,6 +1481,7 @@ impl CodergenBackend for AgentApiBackend { tool_duration = tool_duration.saturating_add(timing.tool); match repair_result { Ok(()) => { + total_usage += session.last_input_usage(); repair_attempts += 1; response = last_assistant_response(&session); } @@ -1505,15 +1508,6 @@ impl CodergenBackend for AgentApiBackend { } } - // Aggregate token usage only from new turns (prevents double-counting on - // reuse), including any output-schema repair turns. - let mut total_usage = TokenCounts::default(); - for turn in &session.history().turns()[turns_before..] { - if let AgentMessage::Assistant { usage, .. } = turn { - total_usage += *usage.clone(); - } - } - let billing_controls = self.resolve_effective_request_controls(node)?; let stage_usage = billed_model_usage_from_llm( self.catalog.as_ref(), diff --git a/lib/crates/fabro-workflow/tests/it/integration.rs b/lib/crates/fabro-workflow/tests/it/integration.rs index 27409d025..4ddf6b75c 100644 --- a/lib/crates/fabro-workflow/tests/it/integration.rs +++ b/lib/crates/fabro-workflow/tests/it/integration.rs @@ -30,8 +30,8 @@ use fabro_interview::{ Answer, AnswerValue, AutoApproveInterviewer, CallbackInterviewer, Interviewer, QueueInterviewer, RecordingInterviewer, }; -use fabro_model::Catalog; use fabro_model::catalog::{LlmCatalogSettings, ProviderCatalogSettings}; +use fabro_model::{Catalog, ProviderId}; use fabro_store::{ArtifactKey, ArtifactStore, Database}; use fabro_types::{RunEvent, RunId, StageId, WorkflowSettings, parse_blob_ref}; use fabro_validate::{Severity, validate, validate_or_raise}; @@ -45,6 +45,7 @@ use fabro_workflow::handler::command::CommandHandler; use fabro_workflow::handler::conditional::ConditionalHandler; use fabro_workflow::handler::exit::ExitHandler; use fabro_workflow::handler::human::HumanHandler; +use fabro_workflow::handler::llm::AgentApiBackend; use fabro_workflow::handler::manager_loop::SubWorkflowHandler; use fabro_workflow::handler::start::StartHandler; use fabro_workflow::handler::wait::WaitHandler; @@ -2086,6 +2087,238 @@ async fn smoke_test_with_mock_codergen_backend() { ); } +#[tokio::test] +async fn shared_thread_compaction_before_routing_audit_succeeds() { + use fabro_auth::EnvCredentialSource; + use fabro_workflow::steering_hub::SteeringHub; + use httpmock::Method::POST; + use httpmock::MockServer; + + fn chat_completion_stream(text: &str, input_tokens: i64, output_tokens: i64) -> String { + let text_chunk = serde_json::json!({ + "id": uuid::Uuid::new_v4().to_string(), + "model": "compact-model", + "choices": [{ + "delta": {"content": text}, + "finish_reason": null + }] + }); + let usage_chunk = serde_json::json!({ + "id": uuid::Uuid::new_v4().to_string(), + "model": "compact-model", + "choices": [], + "usage": { + "prompt_tokens": input_tokens, + "completion_tokens": output_tokens, + "total_tokens": input_tokens + output_tokens + } + }); + format!("data: {text_chunk}\n\ndata: {usage_chunk}\n\ndata: [DONE]\n\n") + } + + fn chat_completion_response(text: &str) -> serde_json::Value { + serde_json::json!({ + "id": uuid::Uuid::new_v4().to_string(), + "model": "compact-model", + "choices": [{ + "message": {"content": text}, + "finish_reason": "stop" + }], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 1, + "total_tokens": 11 + } + }) + } + + let server = MockServer::start_async().await; + let warmup_count = 10; + + for index in 1..=warmup_count { + let prompt = format!("Warmup {index}"); + let next_prompt = if index == warmup_count { + "Audit shared-thread work".to_string() + } else { + format!("Warmup {}", index + 1) + }; + let response = chat_completion_stream(r#"{"outcome":"succeeded"}"#, 1, 1); + server + .mock_async(move |when, then| { + when.method(POST) + .path("/chat/completions") + .body_includes(r#""stream":true"#) + .body_includes(prompt) + .body_excludes(next_prompt); + then.status(200) + .header("content-type", "text/event-stream") + .body(response); + }) + .await; + } + + let audit_stream = chat_completion_stream( + r#"{"outcome":"succeeded","preferred_next_label":"Done"}"#, + 1_000_000, + 1, + ); + let audit_mock = server + .mock_async(|when, then| { + when.method(POST) + .path("/chat/completions") + .body_includes(r#""stream":true"#) + .body_includes("Audit shared-thread work"); + then.status(200) + .header("content-type", "text/event-stream") + .body(audit_stream); + }) + .await; + + let compaction_mock = server + .mock_async(|when, then| { + when.method(POST) + .path("/chat/completions") + .body_excludes(r#""stream":true"#); + then.status(200) + .header("content-type", "application/json") + .json_body(chat_completion_response( + "Previous work completed and the audit can finish.", + )); + }) + .await; + + let settings: LlmCatalogSettings = toml::from_str(&format!( + r#" +[providers.compact] +adapter = "openai_compatible" +agent_profile = "openai" +base_url = "{}" + +[providers.compact.auth] +credentials = ["env:COMPACT_API_KEY"] + +[models.compact-model] +provider = "compact" +display_name = "Compact Model" +family = "mock" +default = true + +[models.compact-model.limits] +context_window = 100000 +max_output = 1024 + +[models.compact-model.features] +tools = true +vision = false +reasoning = false +"#, + server.base_url() + )) + .expect("test catalog should parse"); + let catalog = Arc::new(Catalog::from_builtin_with_overrides(&settings).unwrap()); + let source = Arc::new(EnvCredentialSource::with_env_lookup(Arc::new(|name| { + (name == "COMPACT_API_KEY").then(|| "sk-test".to_string()) + }))); + let backend = AgentApiBackend::new_with_catalog( + "compact-model".to_string(), + ProviderId::from("compact"), + Vec::new(), + source, + Arc::new(SteeringHub::new(Arc::new(Emitter::default()))), + catalog, + ); + + let mut graph = make_graph_with_start_exit("SharedThreadCompactionAudit"); + graph.attrs.insert( + "default_fidelity".to_string(), + AttrValue::String("full".to_string()), + ); + let mut previous = "start".to_string(); + for index in 1..=warmup_count { + let node_id = format!("warmup_{index}"); + let mut node = Node::new(&node_id); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String(format!("Warmup {index}")), + ); + node.attrs.insert( + "thread_id".to_string(), + AttrValue::String("shared-audit-thread".to_string()), + ); + graph.nodes.insert(node_id.clone(), node); + graph.edges.push(Edge::new(&previous, &node_id)); + previous = node_id; + } + + let mut audit = Node::new("audit"); + audit.attrs.insert( + "prompt".to_string(), + AttrValue::String("Audit shared-thread work".to_string()), + ); + audit.attrs.insert( + "thread_id".to_string(), + AttrValue::String("shared-audit-thread".to_string()), + ); + audit.attrs.insert( + "output_schema".to_string(), + AttrValue::String("routing".to_string()), + ); + graph.nodes.insert("audit".to_string(), audit); + graph.edges.push(Edge::new(&previous, "audit")); + graph.edges.push(Edge::new("audit", "exit")); + + let mut registry = HandlerRegistry::new(Box::new(AgentHandler::new(Some(Box::new(backend))))); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + + let dir = tempfile::tempdir().unwrap(); + let engine = WorkflowRunner::new(registry, Arc::new(Emitter::default()), local_env()); + let run_options = RunOptions { + settings: WorkflowSettings::default(), + run_dir: dir.path().to_path_buf(), + cancel_token: CancellationToken::new(), + run_id: test_run_id("shared-thread-compaction-audit"), + labels: std::collections::HashMap::new(), + workflow_slug: None, + github_app: None, + base_branch: None, + display_base_sha: None, + pre_run_git: None, + fork_source_ref: None, + git: None, + }; + + let (_outcome, state) = engine + .run_with_state(&graph, &run_options) + .await + .expect("workflow execution should complete"); + + assert_eq!( + audit_mock.calls_async().await, + 1, + "audit should use the high-usage response that triggers compaction" + ); + assert_eq!( + compaction_mock.calls_async().await, + 1, + "audit response should trigger context compaction before routing finishes" + ); + let checkpoint = state + .current_checkpoint() + .cloned() + .expect("checkpoint should be captured"); + let audit_outcome = checkpoint + .node_outcomes + .get("audit") + .expect("audit outcome should be captured"); + assert_eq!( + audit_outcome.status, + StageOutcome::Succeeded, + "audit should succeed after compaction, got failure: {:?}", + audit_outcome.failure + ); +} + // --------------------------------------------------------------------------- // 12. Parallel fan-out / fan-in integration test (Gap #14) // ---------------------------------------------------------------------------