diff --git a/Cargo.lock b/Cargo.lock index 9479b1a5b..c2b7cfc8d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -182,7 +182,6 @@ dependencies = [ "thiserror 2.0.18", "tokio", "tokio-stream", - "tokio-util", "tower", "uuid", ] diff --git a/crates/attractor/Cargo.toml b/crates/attractor/Cargo.toml index 6600b9ebb..d6f251557 100644 --- a/crates/attractor/Cargo.toml +++ b/crates/attractor/Cargo.toml @@ -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" diff --git a/crates/attractor/src/cli/backend.rs b/crates/attractor/src/cli/backend.rs index 585351ec5..73b5356f1 100644 --- a/crates/attractor/src/cli/backend.rs +++ b/crates/attractor/src/cli/backend.rs @@ -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 { + fn build_profile(&self) -> Arc { 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 = Arc::from(profile); - + let profile = self.build_profile(); let cwd = std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")); let exec_env: Arc = 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>>) { - let outcome_cell: Arc>> = 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::>() .join(", ") } - -#[cfg(test)] -mod tests { - use super::*; - use tokio_util::sync::CancellationToken; - - fn dummy_env() -> Arc { - 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")); - } -} diff --git a/crates/attractor/src/handler/codergen.rs b/crates/attractor/src/handler/codergen.rs index 15110d732..9a85b8955 100644 --- a/crates/attractor/src/handler/codergen.rs +++ b/crates/attractor/src/handler/codergen.rs @@ -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 { - 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::>() - .join("\n"); + if let Some(ids) = obj.get("suggested_next_ids").and_then(|v| v.as_array()) { + let string_ids: Vec = 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>>, - } - - #[async_trait] - impl CodergenBackend for PromptCapturingBackend { - async fn run( - &self, - _node: &Node, - prompt: &str, - _context: &Context, - _thread_id: Option<&str>, - ) -> Result { - *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]