mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-06 02:48:25 +00:00
fix(workflow): preserve usage across session compaction (#420)
## 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` --- [](https://github.com/EveryInc/compound-engineering-plugin) 🤖 Generated with GPT-5 via [Codex](https://openai.com/codex)
This commit is contained in:
parent
535cbda355
commit
ab68fc27d1
3 changed files with 255 additions and 16 deletions
|
|
@ -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<Arc<AsyncMutex<SubAgentManager>>>,
|
||||
completion_coordinator: Option<Arc<dyn CompletionCoordinator>>,
|
||||
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(),
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue