Revert "Replace text-scanning with report_outcome tool for routing"

This reverts commit f166bb4959.
This commit is contained in:
Bryan Helmkamp 2026-02-24 12:10:39 -05:00
parent c9b7bee2da
commit 0431fdb242
4 changed files with 157 additions and 325 deletions

1
Cargo.lock generated
View file

@ -182,7 +182,6 @@ dependencies = [
"thiserror 2.0.18",
"tokio",
"tokio-stream",
"tokio-util",
"tower",
"uuid",
]

View file

@ -45,7 +45,6 @@ tokio-stream = { workspace = true, optional = true, features = ["sync"] }
[dev-dependencies]
tokio = { workspace = true, features = ["test-util", "macros"] }
tokio-util.workspace = true
tempfile = "3"
axum = "0.8"
tower = "0.5"

View file

@ -1,5 +1,5 @@
use std::path::PathBuf;
use std::sync::{Arc, Mutex};
use std::sync::Arc;
use async_trait::async_trait;
@ -8,16 +8,14 @@ use agent::{
ExecutionEnvironment, GeminiProfile, LocalExecutionEnvironment, OpenAiProfile, ProviderProfile,
Session, SessionConfig, Turn,
};
use agent::tool_registry::RegisteredTool;
use llm::client::Client;
use llm::types::ToolDefinition;
use terminal::Styles;
use crate::context::Context;
use crate::error::AttractorError;
use crate::graph::Node;
use crate::handler::codergen::{CodergenBackend, CodergenResult};
use crate::outcome::{Outcome, StageStatus, StageUsage};
use crate::outcome::StageUsage;
/// LLM backend that delegates to an `agent` Session per invocation.
pub struct AgentBackend {
@ -46,12 +44,12 @@ impl AgentBackend {
}
}
fn build_profile(&self) -> Box<dyn ProviderProfile> {
fn build_profile(&self) -> Arc<dyn ProviderProfile> {
let provider = self.provider.as_deref().unwrap_or("anthropic");
match provider {
"openai" => Box::new(OpenAiProfile::new(&self.model)),
"gemini" => Box::new(GeminiProfile::new(&self.model)),
_ => Box::new(AnthropicProfile::new(&self.model)),
"openai" => Arc::new(OpenAiProfile::new(&self.model)),
"gemini" => Arc::new(GeminiProfile::new(&self.model)),
_ => Arc::new(AnthropicProfile::new(&self.model)),
}
}
}
@ -69,11 +67,7 @@ impl CodergenBackend for AgentBackend {
.await
.map_err(|e| AttractorError::Handler(format!("Failed to create LLM client: {e}")))?;
let mut profile = self.build_profile();
let (tool, outcome_cell) = make_report_outcome_tool();
profile.tool_registry_mut().register(tool);
let profile: Arc<dyn ProviderProfile> = Arc::from(profile);
let profile = self.build_profile();
let cwd = std::env::current_dir().unwrap_or_else(|_| PathBuf::from("."));
let exec_env: Arc<dyn ExecutionEnvironment> = if self.docker {
@ -196,13 +190,6 @@ impl CodergenBackend for AgentBackend {
);
}
// If the LLM called report_outcome, return a Full outcome.
let tool_outcome = outcome_cell.lock().unwrap().take();
if let Some(mut outcome) = tool_outcome {
outcome.usage = Some(stage_usage);
return Ok(CodergenResult::Full(outcome));
}
// Extract last assistant response from the session history.
let response = session
.history()
@ -223,94 +210,6 @@ impl CodergenBackend for AgentBackend {
}
}
/// Creates a `report_outcome` tool that the LLM calls to declare routing decisions.
///
/// Returns the `RegisteredTool` and a shared cell where the tool stores the latest `Outcome`.
///
/// # Panics
///
/// The tool executor panics if the internal mutex is poisoned.
#[must_use]
pub fn make_report_outcome_tool() -> (RegisteredTool, Arc<Mutex<Option<Outcome>>>) {
let outcome_cell: Arc<Mutex<Option<Outcome>>> = Arc::new(Mutex::new(None));
let cell = outcome_cell.clone();
let definition = ToolDefinition {
name: "report_outcome".to_string(),
description: "Report the outcome of this task, including routing preference and context updates. \
Call this tool when you have completed your work to declare the result status \
and optionally indicate which next step should be taken."
.to_string(),
parameters: serde_json::json!({
"type": "object",
"required": ["status"],
"properties": {
"status": {
"type": "string",
"enum": ["success", "fail", "partial_success", "retry", "skipped"],
"description": "The result status of this task."
},
"preferred_next_label": {
"type": "string",
"description": "The label of the preferred next edge to follow."
},
"context_updates": {
"type": "object",
"description": "Key-value pairs to merge into the pipeline context."
},
"notes": {
"type": "string",
"description": "Optional notes about the outcome."
},
"failure_reason": {
"type": "string",
"description": "Reason for failure (when status is 'fail')."
}
}
}),
};
let executor: agent::tool_registry::ToolExecutor = Arc::new(move |args, _env, _cancel| {
let cell = cell.clone();
Box::pin(async move {
let status_str = args
.get("status")
.and_then(|v| v.as_str())
.ok_or_else(|| "missing required field: status".to_string())?;
let status: StageStatus = status_str
.parse()
.map_err(|e: String| e)?;
let mut outcome = Outcome {
status,
preferred_label: args.get("preferred_next_label").and_then(|v| v.as_str()).map(String::from),
suggested_next_ids: Vec::new(),
context_updates: std::collections::HashMap::new(),
notes: args.get("notes").and_then(|v| v.as_str()).map(String::from),
failure_reason: args.get("failure_reason").and_then(|v| v.as_str()).map(String::from),
usage: None,
};
if let Some(updates) = args.get("context_updates").and_then(|v| v.as_object()) {
for (key, val) in updates {
outcome.context_updates.insert(key.clone(), val.clone());
}
}
*cell.lock().unwrap() = Some(outcome);
Ok("Outcome recorded.".to_string())
})
});
let tool = RegisteredTool {
definition,
executor,
};
(tool, outcome_cell)
}
fn format_tool_args(args: &serde_json::Value) -> String {
let Some(obj) = args.as_object() else {
return args.to_string();
@ -330,77 +229,3 @@ fn format_tool_args(args: &serde_json::Value) -> String {
.collect::<Vec<_>>()
.join(", ")
}
#[cfg(test)]
mod tests {
use super::*;
use tokio_util::sync::CancellationToken;
fn dummy_env() -> Arc<dyn ExecutionEnvironment> {
Arc::new(LocalExecutionEnvironment::new(std::path::PathBuf::from(".")))
}
#[tokio::test]
async fn report_outcome_tool_captures_outcome() {
let (tool, cell) = make_report_outcome_tool();
let args = serde_json::json!({
"status": "success",
"preferred_next_label": "Fix"
});
let result = (tool.executor)(args, dummy_env(), CancellationToken::new()).await;
assert!(result.is_ok());
let outcome = cell.lock().unwrap().clone().unwrap();
assert_eq!(outcome.status, StageStatus::Success);
assert_eq!(outcome.preferred_label.as_deref(), Some("Fix"));
}
#[tokio::test]
async fn report_outcome_tool_last_call_wins() {
let (tool, cell) = make_report_outcome_tool();
let _ = (tool.executor)(
serde_json::json!({"status": "success", "notes": "first"}),
dummy_env(),
CancellationToken::new(),
)
.await;
let _ = (tool.executor)(
serde_json::json!({"status": "fail", "failure_reason": "oops"}),
dummy_env(),
CancellationToken::new(),
)
.await;
let outcome = cell.lock().unwrap().clone().unwrap();
assert_eq!(outcome.status, StageStatus::Fail);
assert_eq!(outcome.failure_reason.as_deref(), Some("oops"));
}
#[tokio::test]
async fn report_outcome_tool_parses_context_updates() {
let (tool, cell) = make_report_outcome_tool();
let args = serde_json::json!({
"status": "success",
"context_updates": {"k": "v"}
});
let result = (tool.executor)(args, dummy_env(), CancellationToken::new()).await;
assert!(result.is_ok());
let outcome = cell.lock().unwrap().clone().unwrap();
assert_eq!(
outcome.context_updates.get("k"),
Some(&serde_json::json!("v"))
);
}
#[tokio::test]
async fn report_outcome_tool_invalid_status_errors() {
let (tool, _cell) = make_report_outcome_tool();
let args = serde_json::json!({"status": "bogus"});
let result = (tool.executor)(args, dummy_env(), CancellationToken::new()).await;
assert!(result.is_err());
assert!(result.unwrap_err().contains("unknown stage status"));
}
}

View file

@ -47,32 +47,93 @@ fn expand_variables(text: &str, graph: &Graph) -> String {
text.replace("$goal", graph.goal())
}
/// Build a routing preamble for the prompt when the node has 2+ labeled unconditional edges.
///
/// Returns `None` if there are fewer than 2 eligible edges (unconditional + labeled).
fn build_routing_preamble(node_id: &str, graph: &Graph) -> Option<String> {
let labels: Vec<&str> = graph
.outgoing_edges(node_id)
.into_iter()
.filter(|e| e.condition().is_none())
.filter_map(|e| e.label())
.collect();
/// Status fields that indicate a JSON object contains routing directives.
const STATUS_FIELDS: &[&str] = &[
"preferred_next_label",
"outcome",
"suggested_next_ids",
"context_updates",
];
if labels.len() < 2 {
return None;
/// Find all balanced `{...}` JSON object substrings in the text.
fn find_json_objects(text: &str) -> Vec<&str> {
let mut results = Vec::new();
let bytes = text.as_bytes();
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'{' {
let start = i;
let mut depth = 0;
let mut in_string = false;
let mut escape = false;
let mut j = i;
while j < bytes.len() {
let c = bytes[j];
if escape {
escape = false;
} else if c == b'\\' && in_string {
escape = true;
} else if c == b'"' {
in_string = !in_string;
} else if !in_string {
if c == b'{' {
depth += 1;
} else if c == b'}' {
depth -= 1;
if depth == 0 {
results.push(&text[start..=j]);
break;
}
}
}
j += 1;
}
}
i += 1;
}
results
}
/// Extract routing directives from LLM response text.
///
/// Searches for the last JSON object in the response that contains at least
/// one status field (`preferred_next_label`, `outcome`, `suggested_next_ids`,
/// `context_updates`). Merges extracted fields into the outcome.
fn extract_status_fields(text: &str, outcome: &mut Outcome) {
let candidates = find_json_objects(text);
let parsed = candidates.iter().rev().find_map(|candidate| {
let value: serde_json::Value = serde_json::from_str(candidate).ok()?;
if let Some(obj) = value.as_object() {
if STATUS_FIELDS.iter().any(|f| obj.contains_key(*f)) {
return Some(value);
}
}
None
});
let Some(value) = parsed else { return };
let Some(obj) = value.as_object() else { return };
if let Some(label) = obj.get("preferred_next_label").and_then(|v| v.as_str()) {
outcome.preferred_label = Some(label.to_string());
}
let label_list: String = labels
.iter()
.map(|l| format!("- {l}"))
.collect::<Vec<_>>()
.join("\n");
if let Some(ids) = obj.get("suggested_next_ids").and_then(|v| v.as_array()) {
let string_ids: Vec<String> = ids
.iter()
.filter_map(|v| v.as_str().map(String::from))
.collect();
if !string_ids.is_empty() {
outcome.suggested_next_ids = string_ids;
}
}
Some(format!(
"\n\n---\nWhen you have completed your work, call the `report_outcome` tool to declare the result.\n\
Use the `preferred_next_label` parameter to indicate which path to take next.\n\
Available next steps:\n{label_list}"
))
if let Some(updates) = obj.get("context_updates").and_then(|v| v.as_object()) {
for (key, val) in updates {
outcome.context_updates.insert(key.clone(), val.clone());
}
}
}
/// Truncate a string to at most `max_chars` characters.
@ -121,12 +182,7 @@ impl Handler for CodergenHandler {
.prompt()
.filter(|p| !p.is_empty())
.unwrap_or_else(|| node.label());
let mut prompt = expand_variables(raw_prompt, graph);
// 1b. Append routing preamble when the node has routing choices
if let Some(preamble) = build_routing_preamble(&node.id, graph) {
prompt.push_str(&preamble);
}
let prompt = expand_variables(raw_prompt, graph);
// 2. Write prompt to logs
let stage_dir = logs_root.join(&node.id);
@ -195,6 +251,8 @@ impl Handler for CodergenHandler {
serde_json::json!(&response_text),
);
// 7b. Parse routing directives from response text
extract_status_fields(&response_text, &mut outcome);
outcome.usage = stage_usage;
let status_json = serde_json::to_string_pretty(&outcome)
@ -565,135 +623,86 @@ mod tests {
assert!(err.to_string().contains("Request timed out"));
}
// --- build_routing_preamble tests ---
#[test]
fn build_routing_preamble_none_for_no_edges() {
let mut graph = Graph::new("test");
graph.nodes.insert("work".to_string(), Node::new("work"));
assert!(build_routing_preamble("work", &graph).is_none());
fn extract_status_fields_from_fenced_code_block() {
let text = r#"Here is my analysis of the code.
```json
{"preferred_next_label": "fix", "outcome": "success"}
```
That's it."#;
let mut outcome = Outcome::success();
extract_status_fields(text, &mut outcome);
assert_eq!(outcome.preferred_label.as_deref(), Some("fix"));
}
#[test]
fn build_routing_preamble_none_for_single_edge() {
let mut graph = Graph::new("test");
graph.nodes.insert("work".to_string(), Node::new("work"));
let mut edge = crate::graph::Edge::new("work", "next");
edge.attrs.insert("label".to_string(), AttrValue::String("Go".to_string()));
graph.edges.push(edge);
assert!(build_routing_preamble("work", &graph).is_none());
fn extract_status_fields_from_bare_json() {
let text = r#"I recommend routing to fix.
{"preferred_next_label": "fix_batch"}"#;
let mut outcome = Outcome::success();
extract_status_fields(text, &mut outcome);
assert_eq!(outcome.preferred_label.as_deref(), Some("fix_batch"));
}
#[test]
fn build_routing_preamble_includes_labels() {
let mut graph = Graph::new("test");
graph.nodes.insert("work".to_string(), Node::new("work"));
let mut e1 = crate::graph::Edge::new("work", "fix");
e1.attrs.insert("label".to_string(), AttrValue::String("Fix".to_string()));
let mut e2 = crate::graph::Edge::new("work", "review");
e2.attrs.insert("label".to_string(), AttrValue::String("Review".to_string()));
graph.edges.push(e1);
graph.edges.push(e2);
let preamble = build_routing_preamble("work", &graph).unwrap();
assert!(preamble.contains("Fix"));
assert!(preamble.contains("Review"));
assert!(preamble.contains("report_outcome"));
fn extract_status_fields_no_json() {
let text = "Just some plain text response with no JSON at all.";
let mut outcome = Outcome::success();
extract_status_fields(text, &mut outcome);
assert!(outcome.preferred_label.is_none());
assert!(outcome.suggested_next_ids.is_empty());
}
#[test]
fn build_routing_preamble_none_for_unlabeled_edges() {
let mut graph = Graph::new("test");
graph.nodes.insert("work".to_string(), Node::new("work"));
graph.edges.push(crate::graph::Edge::new("work", "a"));
graph.edges.push(crate::graph::Edge::new("work", "b"));
assert!(build_routing_preamble("work", &graph).is_none());
fn extract_status_fields_json_without_status_fields() {
let text = r#"Here is some data: {"name": "test", "count": 42}"#;
let mut outcome = Outcome::success();
extract_status_fields(text, &mut outcome);
assert!(outcome.preferred_label.is_none());
assert!(outcome.suggested_next_ids.is_empty());
}
#[test]
fn build_routing_preamble_none_for_conditional_edges() {
let mut graph = Graph::new("test");
graph.nodes.insert("work".to_string(), Node::new("work"));
let mut e1 = crate::graph::Edge::new("work", "a");
e1.attrs.insert("label".to_string(), AttrValue::String("A".to_string()));
e1.attrs.insert("condition".to_string(), AttrValue::String("outcome=success".to_string()));
let mut e2 = crate::graph::Edge::new("work", "b");
e2.attrs.insert("label".to_string(), AttrValue::String("B".to_string()));
e2.attrs.insert("condition".to_string(), AttrValue::String("outcome=fail".to_string()));
graph.edges.push(e1);
graph.edges.push(e2);
assert!(build_routing_preamble("work", &graph).is_none());
fn extract_status_fields_context_updates_and_suggested_ids() {
let text = r#"```json
{
"preferred_next_label": "review",
"suggested_next_ids": ["node_a", "node_b"],
"context_updates": {"fix.files_changed": 3, "fix.summary": "patched"}
}
```"#;
let mut outcome = Outcome::success();
outcome
.context_updates
.insert("existing_key".to_string(), serde_json::json!("keep"));
extract_status_fields(text, &mut outcome);
assert_eq!(outcome.preferred_label.as_deref(), Some("review"));
assert_eq!(outcome.suggested_next_ids, vec!["node_a", "node_b"]);
assert_eq!(
outcome.context_updates.get("fix.files_changed"),
Some(&serde_json::json!(3))
);
assert_eq!(
outcome.context_updates.get("fix.summary"),
Some(&serde_json::json!("patched"))
);
// Existing keys preserved
assert_eq!(
outcome.context_updates.get("existing_key"),
Some(&serde_json::json!("keep"))
);
}
#[test]
fn build_routing_preamble_mixed_conditional() {
let mut graph = Graph::new("test");
graph.nodes.insert("work".to_string(), Node::new("work"));
// 2 conditional labeled edges + 1 unconditional labeled edge = only 1 unconditional
let mut e1 = crate::graph::Edge::new("work", "a");
e1.attrs.insert("label".to_string(), AttrValue::String("A".to_string()));
e1.attrs.insert("condition".to_string(), AttrValue::String("outcome=success".to_string()));
let mut e2 = crate::graph::Edge::new("work", "b");
e2.attrs.insert("label".to_string(), AttrValue::String("B".to_string()));
e2.attrs.insert("condition".to_string(), AttrValue::String("outcome=fail".to_string()));
let mut e3 = crate::graph::Edge::new("work", "c");
e3.attrs.insert("label".to_string(), AttrValue::String("C".to_string()));
graph.edges.push(e1);
graph.edges.push(e2);
graph.edges.push(e3);
// Only 1 unconditional labeled edge → None
assert!(build_routing_preamble("work", &graph).is_none());
}
#[tokio::test]
async fn codergen_handler_appends_routing_preamble() {
use std::sync::{Arc, Mutex};
struct PromptCapturingBackend {
captured_prompt: Arc<Mutex<Option<String>>>,
}
#[async_trait]
impl CodergenBackend for PromptCapturingBackend {
async fn run(
&self,
_node: &Node,
prompt: &str,
_context: &Context,
_thread_id: Option<&str>,
) -> Result<CodergenResult, AttractorError> {
*self.captured_prompt.lock().unwrap() = Some(prompt.to_string());
Ok(CodergenResult::Text { text: "ok".to_string(), usage: None })
}
}
let captured = Arc::new(Mutex::new(None));
let backend = PromptCapturingBackend {
captured_prompt: captured.clone(),
};
let handler = CodergenHandler::new(Some(Box::new(backend)));
let node = Node::new("work");
let context = Context::new();
let mut graph = Graph::new("test");
graph.nodes.insert("work".to_string(), Node::new("work"));
let mut e1 = crate::graph::Edge::new("work", "fix");
e1.attrs.insert("label".to_string(), AttrValue::String("Fix".to_string()));
let mut e2 = crate::graph::Edge::new("work", "review");
e2.attrs.insert("label".to_string(), AttrValue::String("Review".to_string()));
graph.edges.push(e1);
graph.edges.push(e2);
let tmp = TempDir::new().unwrap();
handler
.execute(&node, &context, &graph, tmp.path(), &make_services())
.await
.unwrap();
let prompt = captured.lock().unwrap().clone().unwrap();
assert!(prompt.contains("report_outcome"), "prompt should mention report_outcome tool");
assert!(prompt.contains("Fix"), "prompt should contain edge label Fix");
assert!(prompt.contains("Review"), "prompt should contain edge label Review");
fn extract_status_fields_uses_last_match() {
let text = r#"{"preferred_next_label": "first"}
Some text in between.
{"preferred_next_label": "second"}"#;
let mut outcome = Outcome::success();
extract_status_fields(text, &mut outcome);
assert_eq!(outcome.preferred_label.as_deref(), Some("second"));
}
#[tokio::test]