diff --git a/crates/attractor/src/cli/backend.rs b/crates/attractor/src/cli/backend.rs index f25196589..5d3f125c2 100644 --- a/crates/attractor/src/cli/backend.rs +++ b/crates/attractor/src/cli/backend.rs @@ -57,6 +57,60 @@ impl AgentBackend { #[async_trait] impl CodergenBackend for AgentBackend { + async fn one_shot( + &self, + node: &Node, + prompt: &str, + ) -> Result { + let client = Client::from_env() + .await + .map_err(|e| AttractorError::Handler(format!("Failed to create LLM client: {e}")))?; + + let model = node.llm_model().unwrap_or(&self.model); + let provider = node + .llm_provider() + .or(self.provider.as_deref()) + .map(String::from); + + let request = llm::types::Request { + model: model.to_string(), + messages: vec![llm::types::Message::user(prompt)], + provider, + reasoning_effort: Some(node.reasoning_effort().to_string()), + tools: None, + tool_choice: None, + response_format: None, + temperature: None, + top_p: None, + max_tokens: None, + stop_sequences: None, + metadata: None, + provider_options: None, + }; + + let response = client + .complete(&request) + .await + .map_err(|e| AttractorError::Handler(format!("one_shot LLM call failed: {e}")))?; + + let mut stage_usage = StageUsage { + model: model.to_string(), + input_tokens: response.usage.input_tokens, + output_tokens: response.usage.output_tokens, + cache_read_tokens: response.usage.cache_read_tokens, + cache_write_tokens: response.usage.cache_write_tokens, + reasoning_tokens: response.usage.reasoning_tokens, + cost: None, + }; + stage_usage.cost = super::compute_stage_cost(&stage_usage); + + Ok(CodergenResult::Text { + text: response.text(), + usage: Some(stage_usage), + files_touched: Vec::new(), + }) + } + async fn run( &self, node: &Node, diff --git a/crates/attractor/src/graph/types.rs b/crates/attractor/src/graph/types.rs index bec13b866..ffb3aeb39 100644 --- a/crates/attractor/src/graph/types.rs +++ b/crates/attractor/src/graph/types.rs @@ -3,6 +3,27 @@ use std::time::Duration; use serde::{Deserialize, Serialize}; +use crate::error::AttractorError; + +/// Whether a codergen node runs as a multi-turn agent loop or a single LLM call. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum CodergenMode { + AgentLoop, + OneShot, +} + +impl CodergenMode { + pub fn parse(s: &str) -> Result { + match s { + "agent_loop" => Ok(Self::AgentLoop), + "one_shot" => Ok(Self::OneShot), + other => Err(AttractorError::Validation(format!( + "invalid codergen_mode: {other:?} (expected \"agent_loop\" or \"one_shot\")" + ))), + } + } +} + /// Typed attribute values for nodes, edges, and graph-level attributes. #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub enum AttrValue { @@ -204,6 +225,14 @@ impl Node { self.str_attr("retry_policy") } + /// Returns the codergen mode for this node. Defaults to `AgentLoop` when absent. + pub fn codergen_mode(&self) -> Result { + match self.str_attr("codergen_mode") { + Some(s) => CodergenMode::parse(s), + None => Ok(CodergenMode::AgentLoop), + } + } + /// Resolve the handler type for this node using explicit type or shape mapping. #[must_use] pub fn handler_type(&self) -> Option<&str> { @@ -653,4 +682,46 @@ mod tests { .insert("max_node_visits".to_string(), AttrValue::Integer(10)); assert_eq!(g.max_node_visits(), 10); } + + #[test] + fn codergen_mode_parse_agent_loop() { + assert_eq!(CodergenMode::parse("agent_loop").unwrap(), CodergenMode::AgentLoop); + } + + #[test] + fn codergen_mode_parse_one_shot() { + assert_eq!(CodergenMode::parse("one_shot").unwrap(), CodergenMode::OneShot); + } + + #[test] + fn codergen_mode_parse_invalid() { + let err = CodergenMode::parse("bogus").unwrap_err(); + assert!(err.to_string().contains("bogus")); + } + + #[test] + fn node_codergen_mode_defaults_to_agent_loop() { + let node = Node::new("test"); + assert_eq!(node.codergen_mode().unwrap(), CodergenMode::AgentLoop); + } + + #[test] + fn node_codergen_mode_one_shot() { + let mut node = Node::new("test"); + node.attrs.insert( + "codergen_mode".to_string(), + AttrValue::String("one_shot".to_string()), + ); + assert_eq!(node.codergen_mode().unwrap(), CodergenMode::OneShot); + } + + #[test] + fn node_codergen_mode_invalid_value() { + let mut node = Node::new("test"); + node.attrs.insert( + "codergen_mode".to_string(), + AttrValue::String("invalid".to_string()), + ); + assert!(node.codergen_mode().is_err()); + } } diff --git a/crates/attractor/src/handler/codergen.rs b/crates/attractor/src/handler/codergen.rs index 31cccd0a7..e90ed0856 100644 --- a/crates/attractor/src/handler/codergen.rs +++ b/crates/attractor/src/handler/codergen.rs @@ -6,7 +6,7 @@ use async_trait::async_trait; use crate::context::Context; use crate::error::AttractorError; use crate::event::EventEmitter; -use crate::graph::{Graph, Node}; +use crate::graph::{CodergenMode, Graph, Node}; use crate::outcome::{Outcome, StageUsage}; use super::{EngineServices, Handler}; @@ -24,6 +24,7 @@ pub enum CodergenResult { /// Backend interface for LLM execution in codergen nodes. #[async_trait] pub trait CodergenBackend: Send + Sync { + /// Run a multi-turn agent loop (the default codergen mode). async fn run( &self, node: &Node, @@ -32,6 +33,17 @@ pub trait CodergenBackend: Send + Sync { thread_id: Option<&str>, emitter: &Arc, ) -> Result; + + /// Run a single LLM call with no tools (one_shot mode). + async fn one_shot( + &self, + _node: &Node, + _prompt: &str, + ) -> Result { + Err(AttractorError::Validation( + "one_shot mode not supported by this backend".into(), + )) + } } /// The default handler for LLM task nodes. @@ -203,11 +215,18 @@ impl Handler for CodergenHandler { } // 4. Call LLM backend + let mode = node.codergen_mode()?; let thread_id = context .get("internal.thread_id") .and_then(|v| v.as_str().map(String::from)); let (response_text, stage_usage, backend_files_touched) = if let Some(backend) = &self.backend { - match backend.run(node, &prompt, context, thread_id.as_deref(), &services.emitter).await { + let result = match mode { + CodergenMode::AgentLoop => { + backend.run(node, &prompt, context, thread_id.as_deref(), &services.emitter).await + } + CodergenMode::OneShot => backend.one_shot(node, &prompt).await, + }; + match result { Ok(CodergenResult::Full(outcome)) => { let status_json = serde_json::to_string_pretty(&outcome) .unwrap_or_else(|_| "{}".to_string()); @@ -713,6 +732,107 @@ Some text in between. assert_eq!(outcome.preferred_label.as_deref(), Some("second")); } + #[tokio::test] + async fn codergen_handler_one_shot_dispatches_to_backend() { + struct OneShotBackend; + + #[async_trait] + impl CodergenBackend for OneShotBackend { + async fn run( + &self, + _node: &Node, + _prompt: &str, + _context: &Context, + _thread_id: Option<&str>, + _emitter: &Arc, + ) -> Result { + panic!("run() should not be called in one_shot mode"); + } + + async fn one_shot( + &self, + _node: &Node, + _prompt: &str, + ) -> Result { + Ok(CodergenResult::Text { + text: "one-shot response".to_string(), + usage: None, + files_touched: Vec::new(), + }) + } + } + + let handler = CodergenHandler::new(Some(Box::new(OneShotBackend))); + let mut node = Node::new("classify"); + node.attrs.insert( + "codergen_mode".to_string(), + AttrValue::String("one_shot".to_string()), + ); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Classify this".to_string()), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let tmp = TempDir::new().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path(), &make_services()) + .await + .unwrap(); + assert_eq!(outcome.status, crate::outcome::StageStatus::Success); + + let prompt_content = + std::fs::read_to_string(tmp.path().join("classify").join("prompt.md")).unwrap(); + assert_eq!(prompt_content, "Classify this"); + + let response_content = + std::fs::read_to_string(tmp.path().join("classify").join("response.md")).unwrap(); + assert_eq!(response_content, "one-shot response"); + } + + #[tokio::test] + async fn codergen_handler_one_shot_simulation_mode() { + let handler = CodergenHandler::new(None); + let mut node = Node::new("classify"); + node.attrs.insert( + "codergen_mode".to_string(), + AttrValue::String("one_shot".to_string()), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let tmp = TempDir::new().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path(), &make_services()) + .await + .unwrap(); + assert_eq!(outcome.status, crate::outcome::StageStatus::Success); + + let response_content = + std::fs::read_to_string(tmp.path().join("classify").join("response.md")).unwrap(); + assert!(response_content.contains("[Simulated]")); + } + + #[tokio::test] + async fn codergen_handler_invalid_mode_returns_error() { + let handler = CodergenHandler::new(None); + let mut node = Node::new("step"); + node.attrs.insert( + "codergen_mode".to_string(), + AttrValue::String("bogus".to_string()), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let tmp = TempDir::new().unwrap(); + + let result = handler + .execute(&node, &context, &graph, tmp.path(), &make_services()) + .await; + assert!(result.is_err()); + assert!(result.unwrap_err().to_string().contains("bogus")); + } + #[tokio::test] async fn codergen_handler_returns_fail_outcome_for_non_retryable_backend_error() { struct ValidationFailBackend; diff --git a/crates/attractor/tests/integration.rs b/crates/attractor/tests/integration.rs index 2e1b39361..c80351abc 100644 --- a/crates/attractor/tests/integration.rs +++ b/crates/attractor/tests/integration.rs @@ -5480,6 +5480,92 @@ mod real_llm { "revise should NOT be traversed with auto-approve" ); } + + #[tokio::test] + #[ignore] + async fn real_llm_one_shot_pipeline() { + let client = if let Some(c) = make_llm_client().await { + c + } else { + eprintln!("Skipping: ANTHROPIC_API_KEY not set"); + return; + }; + + let mut graph = Graph::new("RealLLMOneShot"); + graph.attrs.insert( + "goal".to_string(), + AttrValue::String("Classify a fruit".to_string()), + ); + + let mut start = Node::new("start"); + start.attrs.insert( + "shape".to_string(), + AttrValue::String("Mdiamond".to_string()), + ); + graph.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs.insert( + "shape".to_string(), + AttrValue::String("Msquare".to_string()), + ); + graph.nodes.insert("exit".to_string(), exit); + + let mut classify = Node::new("classify"); + classify.attrs.insert( + "shape".to_string(), + AttrValue::String("box".to_string()), + ); + classify.attrs.insert( + "prompt".to_string(), + AttrValue::String("Reply with exactly one word: is an apple a fruit or vegetable?".to_string()), + ); + classify.attrs.insert( + "codergen_mode".to_string(), + AttrValue::String("one_shot".to_string()), + ); + classify.attrs.insert( + "llm_model".to_string(), + AttrValue::String("claude-haiku-4-5-20251001".to_string()), + ); + graph.nodes.insert("classify".to_string(), classify); + + graph.edges.push(Edge::new("start", "classify")); + graph.edges.push(Edge::new("classify", "exit")); + + let dir = tempfile::tempdir().unwrap(); + + let mut registry = HandlerRegistry::new(Box::new(CodergenHandler::new(Some( + make_llm_backend(Arc::clone(&client)), + )))); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register( + "codergen", + Box::new(CodergenHandler::new(Some(make_llm_backend(client)))), + ); + + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + cancel_token: None, + dry_run: false, + }; + + let outcome = tokio::time::timeout( + std::time::Duration::from_secs(30), + engine.run(&graph, &config), + ) + .await + .expect("should not timeout") + .expect("one_shot pipeline should succeed"); + + assert_eq!(outcome.status, StageStatus::Success); + + let response_path = dir.path().join("classify").join("response.md"); + let response = std::fs::read_to_string(&response_path).unwrap(); + assert!(!response.is_empty(), "response.md should be non-empty"); + } } // ---------------------------------------------------------------------------