From c15b532bcfc1b84549c5b5b6b73b271ca0ea16c9 Mon Sep 17 00:00:00 2001 From: Bryan Helmkamp Date: Sun, 22 Feb 2026 12:07:44 -0400 Subject: [PATCH] Implement three missing spec features: SubPipelineHandler, GraphMergeTransform, HTTP Server Add SubPipelineHandler (handler/sub_pipeline.rs) that inline-executes a parsed sub-graph within the same engine, reading DOT source from node attributes and propagating context diffs back to the parent pipeline. Add GraphMergeTransform (transform.rs) that merges nodes and edges from secondary graphs into a primary graph with namespace-prefixed IDs to avoid collisions. Add WebInterviewer (interviewer/web.rs) backed by oneshot channels for async question/answer flow, and HTTP server (server.rs) with 8 axum endpoints behind a "server" feature flag for pipeline management and human-in-the-loop via web. Fix tempfile dev-dependency usage in server production code by using std::env::temp_dir. Co-Authored-By: Claude Opus 4.6 --- Cargo.lock | 82 + crates/attractor/Cargo.toml | 10 + crates/attractor/src/handler/mod.rs | 1 + crates/attractor/src/handler/sub_pipeline.rs | 371 +++++ crates/attractor/src/interviewer/mod.rs | 1 + crates/attractor/src/interviewer/web.rs | 256 ++++ crates/attractor/src/lib.rs | 2 + crates/attractor/src/server.rs | 727 +++++++++ crates/attractor/src/transform.rs | 227 ++- crates/attractor/tests/integration.rs | 1413 +++++++++++++++++- 10 files changed, 3085 insertions(+), 5 deletions(-) create mode 100644 crates/attractor/src/handler/sub_pipeline.rs create mode 100644 crates/attractor/src/interviewer/web.rs create mode 100644 crates/attractor/src/server.rs diff --git a/Cargo.lock b/Cargo.lock index dd1c48513..32efc9625 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -133,6 +133,7 @@ name = "attractor" version = "0.1.0" dependencies = [ "async-trait", + "axum", "chrono", "coding-agent-loop", "futures", @@ -143,6 +144,8 @@ dependencies = [ "tempfile", "thiserror 2.0.18", "tokio", + "tokio-stream", + "tower", "unified-llm", "uuid", ] @@ -175,6 +178,58 @@ dependencies = [ "fs_extra", ] +[[package]] +name = "axum" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b52af3cb4058c895d37317bb27508dccc8e5f2d39454016b297bf4a400597b8" +dependencies = [ + "axum-core", + "bytes", + "form_urlencoded", + "futures-util", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-util", + "itoa", + "matchit", + "memchr", + "mime", + "percent-encoding", + "pin-project-lite", + "serde_core", + "serde_json", + "serde_path_to_error", + "serde_urlencoded", + "sync_wrapper", + "tokio", + "tower", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "axum-core" +version = "0.5.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08c78f31d7b1291f7ee735c1c6780ccde7785daae9a9206026862dab7d8792d1" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "http-body-util", + "mime", + "pin-project-lite", + "sync_wrapper", + "tower-layer", + "tower-service", + "tracing", +] + [[package]] name = "base64" version = "0.22.1" @@ -770,6 +825,12 @@ version = "1.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87" +[[package]] +name = "httpdate" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" + [[package]] name = "hyper" version = "1.8.1" @@ -784,6 +845,7 @@ dependencies = [ "http", "http-body", "httparse", + "httpdate", "itoa", "pin-project-lite", "pin-utils", @@ -1137,6 +1199,12 @@ version = "0.4.29" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" +[[package]] +name = "matchit" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" + [[package]] name = "memchr" version = "2.8.0" @@ -1863,6 +1931,17 @@ dependencies = [ "zmij", ] +[[package]] +name = "serde_path_to_error" +version = "0.1.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10a9ff822e371bb5403e391ecd83e182e0e77ba7f6fe0160b795797109d1b457" +dependencies = [ + "itoa", + "serde", + "serde_core", +] + [[package]] name = "serde_urlencoded" version = "0.7.1" @@ -2109,6 +2188,7 @@ dependencies = [ "futures-core", "pin-project-lite", "tokio", + "tokio-util", ] [[package]] @@ -2137,6 +2217,7 @@ dependencies = [ "tokio", "tower-layer", "tower-service", + "tracing", ] [[package]] @@ -2175,6 +2256,7 @@ version = "0.1.44" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" dependencies = [ + "log", "pin-project-lite", "tracing-core", ] diff --git a/crates/attractor/Cargo.toml b/crates/attractor/Cargo.toml index a42c38bcc..d67ea6476 100644 --- a/crates/attractor/Cargo.toml +++ b/crates/attractor/Cargo.toml @@ -9,6 +9,10 @@ keywords = ["llm", "ai", "pipeline", "workflow", "dot"] categories = ["development-tools"] readme = "README.md" +[features] +default = [] +server = ["axum", "tower", "tokio-stream"] + [dependencies] coding-agent-loop = { path = "../coding-agent-loop" } unified-llm = { path = "../unified-llm" } @@ -22,10 +26,16 @@ async-trait.workspace = true futures.workspace = true chrono = { workspace = true, features = ["serde"] } nom = "7" +axum = { version = "0.8", optional = true } +tower = { version = "0.5", optional = true } +tokio-stream = { workspace = true, optional = true, features = ["sync"] } [dev-dependencies] tokio = { workspace = true, features = ["test-util", "macros"] } tempfile = "3" +axum = "0.8" +tower = "0.5" +tokio-stream = { workspace = true, features = ["sync"] } [lints] workspace = true diff --git a/crates/attractor/src/handler/mod.rs b/crates/attractor/src/handler/mod.rs index 152bd238a..f0261c5a5 100644 --- a/crates/attractor/src/handler/mod.rs +++ b/crates/attractor/src/handler/mod.rs @@ -5,6 +5,7 @@ pub mod fan_in; pub mod manager_loop; pub mod parallel; pub mod start; +pub mod sub_pipeline; pub mod tool; pub mod wait_human; diff --git a/crates/attractor/src/handler/sub_pipeline.rs b/crates/attractor/src/handler/sub_pipeline.rs new file mode 100644 index 000000000..1cb8cc802 --- /dev/null +++ b/crates/attractor/src/handler/sub_pipeline.rs @@ -0,0 +1,371 @@ +use std::path::Path; +use std::sync::Arc; + +use async_trait::async_trait; + +use crate::context::Context; +use crate::engine::select_edge; +use crate::error::AttractorError; +use crate::event::EventEmitter; +use crate::graph::{Graph, Node}; +use crate::outcome::Outcome; +use crate::pipeline::prepare_pipeline; + +use super::{Handler, HandlerRegistry}; + +/// Executes a sub-pipeline defined by inline DOT source in a node attribute. +/// The sub-pipeline runs with a cloned context; context updates propagate back. +pub struct SubPipelineHandler { + registry: Arc, + _emitter: Arc, +} + +impl SubPipelineHandler { + #[must_use] + pub fn new(registry: Arc, emitter: Arc) -> Self { + Self { + registry, + _emitter: emitter, + } + } +} + +/// Check whether a node is a terminal (exit) node. +fn is_terminal(node: &Node) -> bool { + node.shape() == "Msquare" || node.handler_type() == Some("exit") +} + +#[async_trait] +impl Handler for SubPipelineHandler { + async fn execute( + &self, + node: &Node, + context: &Context, + _graph: &Graph, + logs_root: &Path, + ) -> Result { + // 1. Get DOT source from node attribute + let dot_source = match node.attrs.get("sub_pipeline.dot_source").and_then(|v| v.as_str()) { + Some(s) if !s.is_empty() => s, + _ => return Ok(Outcome::fail("No sub_pipeline.dot_source attribute specified")), + }; + + // 2. Parse the sub-pipeline DOT + let sub_graph = match prepare_pipeline(dot_source) { + Ok(g) => g, + Err(e) => return Ok(Outcome::fail(format!("Failed to parse sub-pipeline: {e}"))), + }; + + // 3. Find start node + let start_node = match sub_graph.find_start_node() { + Some(n) => n.id.clone(), + None => return Ok(Outcome::fail("Sub-pipeline has no start node")), + }; + + // 4. Clone parent context for isolation + let sub_context = context.clone_context(); + let before_snapshot = context.snapshot(); + + // 5. Walk the sub-graph + let sub_logs_root = logs_root.join(&node.id); + let mut current_node_id = start_node; + let mut last_outcome = Outcome::success(); + + let max_steps: usize = 1000; + let mut steps: usize = 0; + + while steps < max_steps { + steps += 1; + + let sub_node = match sub_graph.nodes.get(¤t_node_id) { + Some(n) => n, + None => { + return Ok(Outcome::fail(format!( + "Sub-pipeline node not found: {current_node_id}" + ))); + } + }; + + // Check for terminal node + if is_terminal(sub_node) { + break; + } + + // Execute the node handler + let handler = self.registry.resolve(sub_node); + last_outcome = handler + .execute(sub_node, &sub_context, &sub_graph, &sub_logs_root) + .await?; + + // Apply context updates from the outcome + sub_context.apply_updates(&last_outcome.context_updates); + sub_context.set("outcome", serde_json::json!(last_outcome.status.to_string())); + + // Select next edge + match select_edge(¤t_node_id, &last_outcome, &sub_context, &sub_graph) { + Some(edge) => { + current_node_id.clone_from(&edge.to); + } + None => break, + } + } + + // 6. Compute context diff (sub_context changes vs parent's original snapshot) + let after_snapshot = sub_context.snapshot(); + let mut context_updates = std::collections::HashMap::new(); + for (key, value) in &after_snapshot { + match before_snapshot.get(key) { + Some(old_value) if old_value == value => {} + _ => { + context_updates.insert(key.clone(), value.clone()); + } + } + } + + // 7. Return the last outcome with context updates propagated + let mut result = last_outcome; + result.context_updates.extend(context_updates); + Ok(result) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::AttrValue; + use crate::handler::exit::ExitHandler; + use crate::handler::start::StartHandler; + use crate::outcome::StageStatus; + + fn make_registry() -> Arc { + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + Arc::new(registry) + } + + fn make_emitter() -> Arc { + Arc::new(EventEmitter::new()) + } + + #[tokio::test] + async fn executes_simple_sub_pipeline() { + let registry = make_registry(); + let handler = SubPipelineHandler::new(registry, make_emitter()); + + let mut node = Node::new("sub"); + node.attrs.insert( + "sub_pipeline.dot_source".to_string(), + AttrValue::String( + r#"digraph Sub { + start [shape=Mdiamond] + exit [shape=Msquare] + start -> exit + }"# + .to_string(), + ), + ); + + let context = Context::new(); + let graph = Graph::new("parent"); + let tmp = tempfile::tempdir().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Success); + } + + #[tokio::test] + async fn parent_context_available_in_sub_pipeline() { + let registry = make_registry(); + let handler = SubPipelineHandler::new(registry, make_emitter()); + + let mut node = Node::new("sub"); + node.attrs.insert( + "sub_pipeline.dot_source".to_string(), + AttrValue::String( + r#"digraph Sub { + start [shape=Mdiamond] + exit [shape=Msquare] + start -> exit + }"# + .to_string(), + ), + ); + + let context = Context::new(); + context.set("parent.value", serde_json::json!("hello")); + let graph = Graph::new("parent"); + let tmp = tempfile::tempdir().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Success); + // The sub-pipeline clones the context, so the parent value should be + // available during sub-execution. After execution, any sub-pipeline + // context updates should be in the outcome's context_updates. + } + + #[tokio::test] + async fn context_updates_propagate_back() { + // Use a handler that sets a context value, register it in the sub-pipeline registry + struct ContextSettingHandler; + + #[async_trait] + impl Handler for ContextSettingHandler { + async fn execute( + &self, + _node: &Node, + context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + context.set("sub.result", serde_json::json!("from_sub")); + Ok(Outcome::success()) + } + } + + let mut registry = HandlerRegistry::new(Box::new(ContextSettingHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + let registry = Arc::new(registry); + + let handler = SubPipelineHandler::new(registry, make_emitter()); + + let mut node = Node::new("sub"); + node.attrs.insert( + "sub_pipeline.dot_source".to_string(), + AttrValue::String( + r#"digraph Sub { + start [shape=Mdiamond] + work [shape=box] + exit [shape=Msquare] + start -> work -> exit + }"# + .to_string(), + ), + ); + + let context = Context::new(); + let graph = Graph::new("parent"); + let tmp = tempfile::tempdir().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Success); + + // Context updates from the sub-pipeline should be in the outcome + assert!( + outcome.context_updates.contains_key("sub.result"), + "sub-pipeline context updates should propagate back" + ); + assert_eq!( + outcome.context_updates.get("sub.result"), + Some(&serde_json::json!("from_sub")) + ); + } + + #[tokio::test] + async fn failing_sub_pipeline_returns_fail() { + struct AlwaysFailHandler; + + #[async_trait] + impl Handler for AlwaysFailHandler { + async fn execute( + &self, + _node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + Ok(Outcome::fail("sub-pipeline failure")) + } + } + + let mut registry = HandlerRegistry::new(Box::new(AlwaysFailHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + let registry = Arc::new(registry); + + let handler = SubPipelineHandler::new(registry, make_emitter()); + + let mut node = Node::new("sub"); + // Sub-pipeline where the work node fails and there's a fail edge to exit + node.attrs.insert( + "sub_pipeline.dot_source".to_string(), + AttrValue::String( + r#"digraph Sub { + start [shape=Mdiamond] + work [shape=box, max_retries="0"] + exit [shape=Msquare] + start -> work + work -> exit [condition="outcome=fail"] + }"# + .to_string(), + ), + ); + + let context = Context::new(); + let graph = Graph::new("parent"); + let tmp = tempfile::tempdir().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Fail); + } + + #[tokio::test] + async fn missing_dot_source_returns_fail() { + let registry = make_registry(); + let handler = SubPipelineHandler::new(registry, make_emitter()); + + let node = Node::new("sub"); + let context = Context::new(); + let graph = Graph::new("parent"); + let tmp = tempfile::tempdir().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Fail); + assert!( + outcome + .failure_reason + .as_deref() + .unwrap() + .contains("sub_pipeline.dot_source"), + "should mention the missing attribute" + ); + } + + #[tokio::test] + async fn invalid_dot_source_returns_fail() { + let registry = make_registry(); + let handler = SubPipelineHandler::new(registry, make_emitter()); + + let mut node = Node::new("sub"); + node.attrs.insert( + "sub_pipeline.dot_source".to_string(), + AttrValue::String("not valid dot".to_string()), + ); + + let context = Context::new(); + let graph = Graph::new("parent"); + let tmp = tempfile::tempdir().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Fail); + } +} diff --git a/crates/attractor/src/interviewer/mod.rs b/crates/attractor/src/interviewer/mod.rs index b6fcab581..64558e370 100644 --- a/crates/attractor/src/interviewer/mod.rs +++ b/crates/attractor/src/interviewer/mod.rs @@ -3,6 +3,7 @@ pub mod callback; pub mod console; pub mod queue; pub mod recording; +pub mod web; use std::collections::HashMap; diff --git a/crates/attractor/src/interviewer/web.rs b/crates/attractor/src/interviewer/web.rs new file mode 100644 index 000000000..65420fadc --- /dev/null +++ b/crates/attractor/src/interviewer/web.rs @@ -0,0 +1,256 @@ +use std::collections::HashMap; +use std::sync::{Arc, Mutex}; + +use async_trait::async_trait; +use tokio::sync::oneshot; + +use super::{Answer, Interviewer, Question}; + +/// A pending question waiting for an answer from an external source (e.g., HTTP endpoint). +#[derive(Debug)] +pub struct PendingQuestion { + pub id: String, + pub question: Question, +} + +/// Internal state: maps question ID to its oneshot sender. +struct WebInterviewerInner { + pending: HashMap>, + questions: Vec, + next_id: u64, +} + +/// An interviewer that holds questions until answers are submitted externally. +/// +/// When `ask()` is called, the question is enqueued with a unique ID and the call +/// blocks until `submit_answer()` is called with the matching ID. +pub struct WebInterviewer { + inner: Arc>, +} + +impl WebInterviewer { + #[must_use] + pub fn new() -> Self { + Self { + inner: Arc::new(Mutex::new(WebInterviewerInner { + pending: HashMap::new(), + questions: Vec::new(), + next_id: 1, + })), + } + } + + /// Returns a snapshot of currently pending questions. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + #[must_use] + pub fn pending_questions(&self) -> Vec { + let inner = self.inner.lock().expect("web interviewer lock poisoned"); + inner + .questions + .iter() + .map(|pq| PendingQuestion { + id: pq.id.clone(), + question: pq.question.clone(), + }) + .collect() + } + + /// Submit an answer for a pending question by ID. + /// Returns `true` if the question was found and the answer was delivered, + /// `false` if no such question was pending. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + #[must_use] + pub fn submit_answer(&self, question_id: &str, answer: Answer) -> bool { + let sender = { + let mut inner = self.inner.lock().expect("web interviewer lock poisoned"); + let sender = inner.pending.remove(question_id); + if sender.is_some() { + inner.questions.retain(|pq| pq.id != question_id); + } + sender + }; + sender.is_some_and(|tx| tx.send(answer).is_ok()) + } +} + +impl Default for WebInterviewer { + fn default() -> Self { + Self::new() + } +} + +#[async_trait] +impl Interviewer for WebInterviewer { + async fn ask(&self, question: Question) -> Answer { + let (tx, rx) = oneshot::channel(); + + { + let mut inner = self.inner.lock().expect("web interviewer lock poisoned"); + let id = format!("q-{}", inner.next_id); + inner.next_id += 1; + inner.pending.insert(id.clone(), tx); + inner.questions.push(PendingQuestion { + id, + question: question.clone(), + }); + } + + // Block until answer arrives or sender is dropped + rx.await.unwrap_or_else(|_| Answer::skipped()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::interviewer::{AnswerValue, QuestionType}; + use std::sync::Arc; + + #[tokio::test] + async fn ask_blocks_until_answer_submitted() { + let interviewer = Arc::new(WebInterviewer::new()); + let interviewer_clone = Arc::clone(&interviewer); + + let ask_handle = tokio::spawn(async move { + let q = Question::new("approve?", QuestionType::YesNo); + interviewer_clone.ask(q).await + }); + + // Give the ask task a moment to register the question + tokio::time::sleep(std::time::Duration::from_millis(50)).await; + + // Question should be pending + let pending = interviewer.pending_questions(); + assert_eq!(pending.len(), 1); + assert_eq!(pending[0].question.text, "approve?"); + + // Submit answer + let submitted = interviewer.submit_answer(&pending[0].id, Answer::yes()); + assert!(submitted); + + // ask() should now return + let answer = ask_handle.await.expect("task should complete"); + assert_eq!(answer.value, AnswerValue::Yes); + } + + #[tokio::test] + async fn submit_answer_unblocks_ask() { + let interviewer = Arc::new(WebInterviewer::new()); + let interviewer_clone = Arc::clone(&interviewer); + + let ask_handle = tokio::spawn(async move { + let q = Question::new("name?", QuestionType::Freeform); + interviewer_clone.ask(q).await + }); + + tokio::time::sleep(std::time::Duration::from_millis(50)).await; + + let pending = interviewer.pending_questions(); + assert_eq!(pending.len(), 1); + + let _ = interviewer.submit_answer(&pending[0].id, Answer::text("Alice")); + + let answer = ask_handle.await.expect("task should complete"); + assert_eq!(answer.value, AnswerValue::Text("Alice".to_string())); + assert_eq!(answer.text, Some("Alice".to_string())); + } + + #[tokio::test] + async fn timeout_returns_default_or_timeout_answer() { + let interviewer = Arc::new(WebInterviewer::new()); + + let mut q = Question::new("approve?", QuestionType::YesNo); + q.timeout_seconds = Some(0.05); + + // Use ask_with_timeout from the parent module + let answer = crate::interviewer::ask_with_timeout(interviewer.as_ref(), q).await; + assert_eq!(answer.value, AnswerValue::Timeout); + } + + #[tokio::test] + async fn question_id_correlation() { + let interviewer = Arc::new(WebInterviewer::new()); + let i1 = Arc::clone(&interviewer); + let i2 = Arc::clone(&interviewer); + + // Spawn two concurrent asks + let handle1 = tokio::spawn(async move { + let q = Question::new("first?", QuestionType::YesNo); + i1.ask(q).await + }); + + let handle2 = tokio::spawn(async move { + let q = Question::new("second?", QuestionType::YesNo); + i2.ask(q).await + }); + + tokio::time::sleep(std::time::Duration::from_millis(50)).await; + + let pending = interviewer.pending_questions(); + assert_eq!(pending.len(), 2); + + // Find which ID corresponds to which question + let first_id = pending + .iter() + .find(|pq| pq.question.text == "first?") + .expect("first question should be pending") + .id + .clone(); + let second_id = pending + .iter() + .find(|pq| pq.question.text == "second?") + .expect("second question should be pending") + .id + .clone(); + + // Answer them in reverse order + let _ = interviewer.submit_answer(&second_id, Answer::no()); + let _ = interviewer.submit_answer(&first_id, Answer::yes()); + + let answer1 = handle1.await.expect("task should complete"); + let answer2 = handle2.await.expect("task should complete"); + + assert_eq!(answer1.value, AnswerValue::Yes); + assert_eq!(answer2.value, AnswerValue::No); + } + + #[test] + fn submit_answer_for_unknown_id_returns_false() { + let interviewer = WebInterviewer::new(); + let result = interviewer.submit_answer("nonexistent", Answer::yes()); + assert!(!result); + } + + #[tokio::test] + async fn pending_questions_empty_initially() { + let interviewer = WebInterviewer::new(); + assert!(interviewer.pending_questions().is_empty()); + } + + #[tokio::test] + async fn pending_questions_cleared_after_answer() { + let interviewer = Arc::new(WebInterviewer::new()); + let i_clone = Arc::clone(&interviewer); + + let handle = tokio::spawn(async move { + let q = Question::new("q?", QuestionType::YesNo); + i_clone.ask(q).await + }); + + tokio::time::sleep(std::time::Duration::from_millis(50)).await; + + let pending = interviewer.pending_questions(); + assert_eq!(pending.len(), 1); + + let _ = interviewer.submit_answer(&pending[0].id, Answer::yes()); + handle.await.expect("task should complete"); + + assert!(interviewer.pending_questions().is_empty()); + } +} diff --git a/crates/attractor/src/lib.rs b/crates/attractor/src/lib.rs index 8ed07702a..6fe2d8eef 100644 --- a/crates/attractor/src/lib.rs +++ b/crates/attractor/src/lib.rs @@ -11,6 +11,8 @@ pub mod interviewer; pub mod outcome; pub mod parser; pub mod pipeline; +#[cfg(any(feature = "server", test))] +pub mod server; pub mod stylesheet; pub mod transform; pub mod validation; diff --git a/crates/attractor/src/server.rs b/crates/attractor/src/server.rs new file mode 100644 index 000000000..e7e3def07 --- /dev/null +++ b/crates/attractor/src/server.rs @@ -0,0 +1,727 @@ +use std::collections::HashMap; +use std::sync::{Arc, Mutex}; + +use axum::extract::{Path, State}; +use axum::http::StatusCode; +use axum::response::sse::{Event, Sse}; +use axum::response::{IntoResponse, Response}; +use axum::routing::{get, post}; +use axum::{Json, Router}; +use serde::{Deserialize, Serialize}; +use tokio::sync::broadcast; +use tokio_stream::wrappers::BroadcastStream; +use tokio_stream::StreamExt; + +use crate::checkpoint::Checkpoint; +use crate::context::Context; +use crate::engine::{PipelineEngine, RunConfig}; +use crate::event::{EventEmitter, PipelineEvent}; +use crate::handler::HandlerRegistry; +use crate::interviewer::web::WebInterviewer; +use crate::interviewer::{Answer, AnswerValue}; + +/// Status of a managed pipeline. +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub enum PipelineStatus { + Running, + Completed, + Failed, + Cancelled, +} + +/// A pending question exposed via the API. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ApiQuestion { + pub id: String, + pub text: String, + pub question_type: String, +} + +/// Snapshot of a managed pipeline. +struct ManagedPipeline { + status: PipelineStatus, + error: Option, + interviewer: Arc, + event_tx: broadcast::Sender, + context: Option, + checkpoint: Option, + cancel_tx: Option>, +} + +/// Shared application state for the server. +pub struct AppState { + pipelines: Mutex>, + registry_factory: Box HandlerRegistry + Send + Sync>, +} + +/// Request body for POST /pipelines. +#[derive(Debug, Deserialize)] +pub struct StartPipelineRequest { + pub dot_source: String, +} + +/// Response body for POST /pipelines. +#[derive(Debug, Serialize)] +pub struct StartPipelineResponse { + pub id: String, +} + +/// Response body for GET /pipelines/{id}. +#[derive(Debug, Serialize)] +pub struct PipelineStatusResponse { + pub id: String, + pub status: PipelineStatus, + #[serde(skip_serializing_if = "Option::is_none")] + pub error: Option, +} + +/// Request body for POST /pipelines/{id}/questions/{qid}/answer. +#[derive(Debug, Deserialize)] +pub struct SubmitAnswerRequest { + pub value: String, +} + +/// Response for answer submission. +#[derive(Debug, Serialize)] +pub struct SubmitAnswerResponse { + pub accepted: bool, +} + +/// Build the axum Router with all pipeline endpoints. +pub fn build_router(state: Arc) -> Router { + Router::new() + .route("/pipelines", post(start_pipeline)) + .route("/pipelines/{id}", get(get_pipeline_status)) + .route("/pipelines/{id}/questions", get(get_questions)) + .route( + "/pipelines/{id}/questions/{qid}/answer", + post(submit_answer), + ) + .route("/pipelines/{id}/events", get(get_events)) + .route("/pipelines/{id}/checkpoint", get(get_checkpoint)) + .route("/pipelines/{id}/context", get(get_context)) + .route("/pipelines/{id}/cancel", post(cancel_pipeline)) + .with_state(state) +} + +/// Create an `AppState` with the given registry factory. +pub fn create_app_state( + registry_factory: impl Fn() -> HandlerRegistry + Send + Sync + 'static, +) -> Arc { + Arc::new(AppState { + pipelines: Mutex::new(HashMap::new()), + registry_factory: Box::new(registry_factory), + }) +} + +async fn start_pipeline( + State(state): State>, + Json(req): Json, +) -> Response { + // Parse the DOT source + let graph = match crate::pipeline::prepare_pipeline(&req.dot_source) { + Ok(g) => g, + Err(e) => { + return (StatusCode::BAD_REQUEST, Json(serde_json::json!({"error": e.to_string()}))) + .into_response(); + } + }; + + let pipeline_id = uuid::Uuid::new_v4().to_string(); + let interviewer = Arc::new(WebInterviewer::new()); + let (event_tx, _) = broadcast::channel(256); + let (cancel_tx, cancel_rx) = tokio::sync::oneshot::channel::<()>(); + + let context = Context::new(); + + // Set up event emitter that broadcasts to the channel + let mut emitter = EventEmitter::new(); + let tx_clone = event_tx.clone(); + emitter.on_event(move |event| { + let _ = tx_clone.send(event.clone()); + }); + + let registry = (state.registry_factory)(); + let engine = PipelineEngine::new(registry, emitter); + + { + let mut pipelines = state.pipelines.lock().expect("pipelines lock poisoned"); + pipelines.insert( + pipeline_id.clone(), + ManagedPipeline { + status: PipelineStatus::Running, + error: None, + interviewer: Arc::clone(&interviewer), + event_tx: event_tx.clone(), + context: Some(context.clone()), + checkpoint: None, + cancel_tx: Some(cancel_tx), + }, + ); + } + + // Spawn pipeline execution + let state_clone = Arc::clone(&state); + let id_clone = pipeline_id.clone(); + tokio::spawn(async move { + let logs_root = std::env::temp_dir().join(format!("attractor-{}", uuid::Uuid::new_v4())); + std::fs::create_dir_all(&logs_root).expect("failed to create logs directory"); + let config = RunConfig { logs_root }; + + let result = tokio::select! { + result = engine.run(&graph, &config) => result, + _ = cancel_rx => { + let mut pipelines = state_clone.pipelines.lock().expect("pipelines lock poisoned"); + if let Some(pipeline) = pipelines.get_mut(&id_clone) { + pipeline.status = PipelineStatus::Cancelled; + } + return; + } + }; + + // Save final checkpoint + let checkpoint = Checkpoint::load(&config.logs_root.join("checkpoint.json")).ok(); + + let mut pipelines = state_clone.pipelines.lock().expect("pipelines lock poisoned"); + if let Some(pipeline) = pipelines.get_mut(&id_clone) { + match result { + Ok(_) => { + pipeline.status = PipelineStatus::Completed; + } + Err(e) => { + pipeline.status = PipelineStatus::Failed; + pipeline.error = Some(e.to_string()); + } + } + pipeline.checkpoint = checkpoint; + } + }); + + ( + StatusCode::CREATED, + Json(StartPipelineResponse { id: pipeline_id }), + ) + .into_response() +} + +async fn get_pipeline_status( + State(state): State>, + Path(id): Path, +) -> Response { + let pipelines = state.pipelines.lock().expect("pipelines lock poisoned"); + match pipelines.get(&id) { + Some(pipeline) => ( + StatusCode::OK, + Json(PipelineStatusResponse { + id: id.clone(), + status: pipeline.status.clone(), + error: pipeline.error.clone(), + }), + ) + .into_response(), + None => StatusCode::NOT_FOUND.into_response(), + } +} + +async fn get_questions( + State(state): State>, + Path(id): Path, +) -> Response { + let pipelines = state.pipelines.lock().expect("pipelines lock poisoned"); + match pipelines.get(&id) { + Some(pipeline) => { + let pending = pipeline.interviewer.pending_questions(); + let questions: Vec = pending + .into_iter() + .map(|pq| ApiQuestion { + id: pq.id, + text: pq.question.text, + question_type: format!("{:?}", pq.question.question_type), + }) + .collect(); + (StatusCode::OK, Json(questions)).into_response() + } + None => StatusCode::NOT_FOUND.into_response(), + } +} + +async fn submit_answer( + State(state): State>, + Path((id, qid)): Path<(String, String)>, + Json(req): Json, +) -> Response { + let pipelines = state.pipelines.lock().expect("pipelines lock poisoned"); + match pipelines.get(&id) { + Some(pipeline) => { + let answer = Answer { + value: AnswerValue::Text(req.value.clone()), + selected_option: None, + text: Some(req.value), + }; + let accepted = pipeline.interviewer.submit_answer(&qid, answer); + (StatusCode::OK, Json(SubmitAnswerResponse { accepted })).into_response() + } + None => StatusCode::NOT_FOUND.into_response(), + } +} + +async fn get_events( + State(state): State>, + Path(id): Path, +) -> Response { + let rx = { + let pipelines = state.pipelines.lock().expect("pipelines lock poisoned"); + match pipelines.get(&id) { + Some(pipeline) => pipeline.event_tx.subscribe(), + None => return StatusCode::NOT_FOUND.into_response(), + } + }; + + let stream = BroadcastStream::new(rx).filter_map(|result| match result { + Ok(event) => { + let data = serde_json::to_string(&event).unwrap_or_default(); + Some(Ok::( + Event::default().data(data), + )) + } + Err(_) => None, + }); + + Sse::new(stream).into_response() +} + +async fn get_checkpoint( + State(state): State>, + Path(id): Path, +) -> Response { + let pipelines = state.pipelines.lock().expect("pipelines lock poisoned"); + match pipelines.get(&id) { + Some(pipeline) => match &pipeline.checkpoint { + Some(cp) => (StatusCode::OK, Json(cp.clone())).into_response(), + None => (StatusCode::OK, Json(serde_json::json!(null))).into_response(), + }, + None => StatusCode::NOT_FOUND.into_response(), + } +} + +async fn get_context( + State(state): State>, + Path(id): Path, +) -> Response { + let pipelines = state.pipelines.lock().expect("pipelines lock poisoned"); + match pipelines.get(&id) { + Some(pipeline) => match &pipeline.context { + Some(ctx) => (StatusCode::OK, Json(ctx.snapshot())).into_response(), + None => (StatusCode::OK, Json(serde_json::json!({}))).into_response(), + }, + None => StatusCode::NOT_FOUND.into_response(), + } +} + +async fn cancel_pipeline( + State(state): State>, + Path(id): Path, +) -> Response { + let mut pipelines = state.pipelines.lock().expect("pipelines lock poisoned"); + match pipelines.get_mut(&id) { + Some(pipeline) => { + if pipeline.status != PipelineStatus::Running { + return ( + StatusCode::CONFLICT, + Json(serde_json::json!({"error": "pipeline is not running"})), + ) + .into_response(); + } + if let Some(cancel_tx) = pipeline.cancel_tx.take() { + let _ = cancel_tx.send(()); + } + pipeline.status = PipelineStatus::Cancelled; + (StatusCode::OK, Json(serde_json::json!({"cancelled": true}))).into_response() + } + None => StatusCode::NOT_FOUND.into_response(), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use axum::body::Body; + use axum::http::Request; + use tower::ServiceExt; + + use crate::handler::exit::ExitHandler; + use crate::handler::start::StartHandler; + + const MINIMAL_DOT: &str = r#"digraph Test { + graph [goal="Test"] + start [shape=Mdiamond] + exit [shape=Msquare] + start -> exit + }"#; + + fn test_registry() -> HandlerRegistry { + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry + } + + fn test_app() -> Router { + let state = create_app_state(test_registry); + build_router(state) + } + + async fn body_json(body: Body) -> serde_json::Value { + let bytes = axum::body::to_bytes(body, usize::MAX).await.unwrap(); + serde_json::from_slice(&bytes).unwrap() + } + + #[tokio::test] + async fn post_pipelines_starts_pipeline_and_returns_id() { + let app = test_app(); + + let req = Request::builder() + .method("POST") + .uri("/pipelines") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(), + )) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::CREATED); + + let body = body_json(response.into_body()).await; + assert!(body["id"].is_string()); + assert!(!body["id"].as_str().unwrap().is_empty()); + } + + #[tokio::test] + async fn post_pipelines_invalid_dot_returns_bad_request() { + let app = test_app(); + + let req = Request::builder() + .method("POST") + .uri("/pipelines") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"dot_source": "not a graph"})).unwrap(), + )) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + } + + #[tokio::test] + async fn get_pipeline_status_returns_status() { + let state = create_app_state(test_registry); + let app = build_router(Arc::clone(&state)); + + // Start a pipeline + let req = Request::builder() + .method("POST") + .uri("/pipelines") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(), + )) + .unwrap(); + + let response = app.clone().oneshot(req).await.unwrap(); + let body = body_json(response.into_body()).await; + let pipeline_id = body["id"].as_str().unwrap().to_string(); + + // Give pipeline a moment to run + tokio::time::sleep(std::time::Duration::from_millis(200)).await; + + // Check status + let req = Request::builder() + .method("GET") + .uri(format!("/pipelines/{pipeline_id}")) + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + + let body = body_json(response.into_body()).await; + assert_eq!(body["id"].as_str().unwrap(), pipeline_id); + // Status should be either "running" or "completed" + let status = body["status"].as_str().unwrap(); + assert!( + status == "running" || status == "completed", + "unexpected status: {status}" + ); + } + + #[tokio::test] + async fn get_pipeline_status_not_found() { + let app = test_app(); + + let req = Request::builder() + .method("GET") + .uri("/pipelines/nonexistent") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::NOT_FOUND); + } + + #[tokio::test] + async fn get_questions_returns_empty_list() { + let state = create_app_state(test_registry); + let app = build_router(Arc::clone(&state)); + + // Start a pipeline + let req = Request::builder() + .method("POST") + .uri("/pipelines") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(), + )) + .unwrap(); + + let response = app.clone().oneshot(req).await.unwrap(); + let body = body_json(response.into_body()).await; + let pipeline_id = body["id"].as_str().unwrap().to_string(); + + // Get questions (should be empty for a pipeline without wait.human nodes) + let req = Request::builder() + .method("GET") + .uri(format!("/pipelines/{pipeline_id}/questions")) + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + + let body = body_json(response.into_body()).await; + assert!(body.is_array()); + } + + #[tokio::test] + async fn submit_answer_not_found_pipeline() { + let app = test_app(); + + let req = Request::builder() + .method("POST") + .uri("/pipelines/nonexistent/questions/q1/answer") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"value": "yes"})).unwrap(), + )) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::NOT_FOUND); + } + + #[tokio::test] + async fn get_events_not_found() { + let app = test_app(); + + let req = Request::builder() + .method("GET") + .uri("/pipelines/nonexistent/events") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::NOT_FOUND); + } + + #[tokio::test] + async fn get_checkpoint_returns_null_initially() { + let state = create_app_state(test_registry); + let app = build_router(Arc::clone(&state)); + + // Start a pipeline + let req = Request::builder() + .method("POST") + .uri("/pipelines") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(), + )) + .unwrap(); + + let response = app.clone().oneshot(req).await.unwrap(); + let body = body_json(response.into_body()).await; + let pipeline_id = body["id"].as_str().unwrap().to_string(); + + // Get checkpoint immediately (before pipeline completes, may be null) + let req = Request::builder() + .method("GET") + .uri(format!("/pipelines/{pipeline_id}/checkpoint")) + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn get_context_returns_map() { + let state = create_app_state(test_registry); + let app = build_router(Arc::clone(&state)); + + // Start a pipeline + let req = Request::builder() + .method("POST") + .uri("/pipelines") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(), + )) + .unwrap(); + + let response = app.clone().oneshot(req).await.unwrap(); + let body = body_json(response.into_body()).await; + let pipeline_id = body["id"].as_str().unwrap().to_string(); + + // Get context + let req = Request::builder() + .method("GET") + .uri(format!("/pipelines/{pipeline_id}/context")) + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + + let body = body_json(response.into_body()).await; + assert!(body.is_object()); + } + + #[tokio::test] + async fn cancel_pipeline_succeeds() { + let state = create_app_state(test_registry); + let app = build_router(Arc::clone(&state)); + + // Start a pipeline + let req = Request::builder() + .method("POST") + .uri("/pipelines") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(), + )) + .unwrap(); + + let response = app.clone().oneshot(req).await.unwrap(); + let body = body_json(response.into_body()).await; + let pipeline_id = body["id"].as_str().unwrap().to_string(); + + // Cancel it + let req = Request::builder() + .method("POST") + .uri(format!("/pipelines/{pipeline_id}/cancel")) + .body(Body::empty()) + .unwrap(); + + let response = app.clone().oneshot(req).await.unwrap(); + // Could be OK (cancelled) or CONFLICT (already completed) + let status = response.status(); + assert!( + status == StatusCode::OK || status == StatusCode::CONFLICT, + "unexpected status: {status}" + ); + } + + #[tokio::test] + async fn cancel_nonexistent_pipeline_returns_not_found() { + let app = test_app(); + + let req = Request::builder() + .method("POST") + .uri("/pipelines/nonexistent/cancel") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::NOT_FOUND); + } + + #[tokio::test] + async fn get_events_returns_sse_stream() { + let state = create_app_state(test_registry); + let app = build_router(Arc::clone(&state)); + + // Start a pipeline + let req = Request::builder() + .method("POST") + .uri("/pipelines") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(), + )) + .unwrap(); + + let response = app.clone().oneshot(req).await.unwrap(); + let body = body_json(response.into_body()).await; + let pipeline_id = body["id"].as_str().unwrap().to_string(); + + // Request the SSE stream + let req = Request::builder() + .method("GET") + .uri(format!("/pipelines/{pipeline_id}/events")) + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + + // Check content-type is text/event-stream + let content_type = response + .headers() + .get("content-type") + .expect("content-type header should be present") + .to_str() + .unwrap(); + assert!( + content_type.contains("text/event-stream"), + "expected text/event-stream, got: {content_type}" + ); + } + + #[tokio::test] + async fn pipeline_completes_and_status_is_completed() { + let state = create_app_state(test_registry); + let app = build_router(Arc::clone(&state)); + + // Start a pipeline + let req = Request::builder() + .method("POST") + .uri("/pipelines") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(), + )) + .unwrap(); + + let response = app.clone().oneshot(req).await.unwrap(); + let body = body_json(response.into_body()).await; + let pipeline_id = body["id"].as_str().unwrap().to_string(); + + // Wait for pipeline to complete + tokio::time::sleep(std::time::Duration::from_millis(500)).await; + + // Check status + let req = Request::builder() + .method("GET") + .uri(format!("/pipelines/{pipeline_id}")) + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + + let body = body_json(response.into_body()).await; + assert_eq!(body["status"].as_str().unwrap(), "completed"); + } +} diff --git a/crates/attractor/src/transform.rs b/crates/attractor/src/transform.rs index 1e9bcf564..03efff66d 100644 --- a/crates/attractor/src/transform.rs +++ b/crates/attractor/src/transform.rs @@ -1,4 +1,4 @@ -use crate::graph::{AttrValue, Graph}; +use crate::graph::{AttrValue, Edge, Graph, Node}; use crate::stylesheet::{apply_stylesheet, parse_stylesheet}; /// A transform that modifies the pipeline graph after parsing and before validation. @@ -48,6 +48,43 @@ impl Transform for PreambleTransform { } } +/// Merges nodes and edges from secondary graphs into the primary graph. +/// Node IDs from secondary graphs are prefixed with a namespace to avoid collisions. +pub struct GraphMergeTransform { + secondary_graphs: Vec, +} + +impl GraphMergeTransform { + pub fn new(secondary_graphs: Vec) -> Self { + Self { secondary_graphs } + } +} + +impl Transform for GraphMergeTransform { + fn apply(&self, graph: &mut Graph) { + for secondary in &self.secondary_graphs { + let prefix = &secondary.name; + + for (id, node) in &secondary.nodes { + let prefixed_id = format!("{prefix}.{id}"); + let mut merged_node = Node::new(&prefixed_id); + merged_node.attrs = node.attrs.clone(); + merged_node.classes = node.classes.clone(); + graph.nodes.insert(prefixed_id, merged_node); + } + + for edge in &secondary.edges { + let mut merged_edge = Edge::new( + format!("{prefix}.{}", edge.from), + format!("{prefix}.{}", edge.to), + ); + merged_edge.attrs = edge.attrs.clone(); + graph.edges.push(merged_edge); + } + } + } +} + /// Applies the `model_stylesheet` graph attribute to resolve LLM properties for each node. pub struct StylesheetApplicationTransform; @@ -67,7 +104,6 @@ impl Transform for StylesheetApplicationTransform { #[cfg(test)] mod tests { use super::*; - use crate::graph::Node; #[test] fn variable_expansion_replaces_goal() { @@ -251,4 +287,191 @@ mod tests { assert!(graph.nodes["work"].attrs.get("prompt").is_none()); } + + // ----------------------------------------------------------------------- + // GraphMergeTransform tests + // ----------------------------------------------------------------------- + + #[test] + fn graph_merge_combines_nodes_and_edges() { + let mut primary = Graph::new("primary"); + primary.nodes.insert("a".to_string(), Node::new("a")); + primary.nodes.insert("b".to_string(), Node::new("b")); + primary.edges.push(Edge::new("a", "b")); + + let mut secondary = Graph::new("secondary"); + secondary.nodes.insert("x".to_string(), Node::new("x")); + secondary.nodes.insert("y".to_string(), Node::new("y")); + secondary.edges.push(Edge::new("x", "y")); + + let transform = GraphMergeTransform::new(vec![secondary]); + transform.apply(&mut primary); + + // Primary should now have 4 nodes: a, b, secondary.x, secondary.y + assert_eq!(primary.nodes.len(), 4); + assert!(primary.nodes.contains_key("secondary.x")); + assert!(primary.nodes.contains_key("secondary.y")); + // Should have 2 edges: a->b and secondary.x->secondary.y + assert_eq!(primary.edges.len(), 2); + } + + #[test] + fn graph_merge_prefixes_node_ids_to_avoid_collisions() { + let mut primary = Graph::new("primary"); + primary.nodes.insert("work".to_string(), Node::new("work")); + + let mut secondary = Graph::new("sub"); + secondary.nodes.insert("work".to_string(), Node::new("work")); + + let transform = GraphMergeTransform::new(vec![secondary]); + transform.apply(&mut primary); + + // Primary "work" is preserved, secondary "work" becomes "sub.work" + assert!(primary.nodes.contains_key("work")); + assert!(primary.nodes.contains_key("sub.work")); + assert_eq!(primary.nodes.len(), 2); + } + + #[test] + fn graph_merge_remaps_edges_to_prefixed_ids() { + let mut primary = Graph::new("primary"); + primary.nodes.insert("a".to_string(), Node::new("a")); + + let mut secondary = Graph::new("sub"); + secondary.nodes.insert("x".to_string(), Node::new("x")); + secondary.nodes.insert("y".to_string(), Node::new("y")); + secondary.edges.push(Edge::new("x", "y")); + + let transform = GraphMergeTransform::new(vec![secondary]); + transform.apply(&mut primary); + + // The edge from secondary should be remapped to sub.x -> sub.y + let merged_edge = primary + .edges + .iter() + .find(|e| e.from == "sub.x") + .expect("should have edge from sub.x"); + assert_eq!(merged_edge.to, "sub.y"); + } + + #[test] + fn graph_merge_preserves_primary_attributes() { + let mut primary = Graph::new("primary"); + primary.attrs.insert( + "goal".to_string(), + AttrValue::String("Build feature".to_string()), + ); + primary.attrs.insert( + "model_stylesheet".to_string(), + AttrValue::String("* { llm_model: sonnet; }".to_string()), + ); + + let mut secondary = Graph::new("sub"); + secondary.attrs.insert( + "goal".to_string(), + AttrValue::String("Sub goal".to_string()), + ); + secondary.nodes.insert("x".to_string(), Node::new("x")); + + let transform = GraphMergeTransform::new(vec![secondary]); + transform.apply(&mut primary); + + assert_eq!(primary.goal(), "Build feature"); + assert_eq!(primary.model_stylesheet(), "* { llm_model: sonnet; }"); + } + + #[test] + fn graph_merge_empty_secondary_is_noop() { + let mut primary = Graph::new("primary"); + primary.nodes.insert("a".to_string(), Node::new("a")); + primary.edges.push(Edge::new("a", "a")); + + let secondary = Graph::new("empty"); + + let transform = GraphMergeTransform::new(vec![secondary]); + transform.apply(&mut primary); + + assert_eq!(primary.nodes.len(), 1); + assert_eq!(primary.edges.len(), 1); + } + + #[test] + fn graph_merge_multiple_secondary_graphs() { + let mut primary = Graph::new("primary"); + primary.nodes.insert("a".to_string(), Node::new("a")); + + let mut sub1 = Graph::new("sub1"); + sub1.nodes.insert("n1".to_string(), Node::new("n1")); + + let mut sub2 = Graph::new("sub2"); + sub2.nodes.insert("n2".to_string(), Node::new("n2")); + + let transform = GraphMergeTransform::new(vec![sub1, sub2]); + transform.apply(&mut primary); + + assert_eq!(primary.nodes.len(), 3); + assert!(primary.nodes.contains_key("a")); + assert!(primary.nodes.contains_key("sub1.n1")); + assert!(primary.nodes.contains_key("sub2.n2")); + } + + #[test] + fn graph_merge_preserves_node_attributes() { + let mut primary = Graph::new("primary"); + + let mut secondary = Graph::new("sub"); + let mut node = Node::new("worker"); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Do the work".to_string()), + ); + node.attrs.insert( + "shape".to_string(), + AttrValue::String("box".to_string()), + ); + secondary.nodes.insert("worker".to_string(), node); + + let transform = GraphMergeTransform::new(vec![secondary]); + transform.apply(&mut primary); + + let merged = &primary.nodes["sub.worker"]; + assert_eq!(merged.id, "sub.worker"); + assert_eq!( + merged.attrs.get("prompt").and_then(AttrValue::as_str), + Some("Do the work") + ); + assert_eq!( + merged.attrs.get("shape").and_then(AttrValue::as_str), + Some("box") + ); + } + + #[test] + fn graph_merge_preserves_edge_attributes() { + let mut primary = Graph::new("primary"); + + let mut secondary = Graph::new("sub"); + secondary.nodes.insert("x".to_string(), Node::new("x")); + secondary.nodes.insert("y".to_string(), Node::new("y")); + let mut edge = Edge::new("x", "y"); + edge.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + secondary.edges.push(edge); + + let transform = GraphMergeTransform::new(vec![secondary]); + transform.apply(&mut primary); + + let merged_edge = primary + .edges + .iter() + .find(|e| e.from == "sub.x") + .expect("should have merged edge"); + assert_eq!(merged_edge.to, "sub.y"); + assert_eq!( + merged_edge.attrs.get("condition").and_then(AttrValue::as_str), + Some("outcome=success") + ); + } } diff --git a/crates/attractor/tests/integration.rs b/crates/attractor/tests/integration.rs index 94538f4e9..cfc86732d 100644 --- a/crates/attractor/tests/integration.rs +++ b/crates/attractor/tests/integration.rs @@ -6,21 +6,25 @@ use attractor::checkpoint::Checkpoint; use attractor::context::Context; use attractor::engine::{PipelineEngine, RunConfig}; use attractor::error::AttractorError; -use attractor::event::EventEmitter; +use attractor::event::{EventEmitter, PipelineEvent}; use attractor::graph::{AttrValue, Edge, Graph, Node}; use attractor::handler::codergen::{CodergenBackend, CodergenHandler, CodergenResult}; use attractor::handler::conditional::ConditionalHandler; use attractor::handler::exit::ExitHandler; +use attractor::handler::manager_loop::ManagerLoopHandler; use attractor::handler::start::StartHandler; +use attractor::handler::tool::ToolHandler; use attractor::handler::wait_human::WaitHumanHandler; use attractor::handler::{Handler, HandlerRegistry}; +use attractor::interviewer::auto_approve::AutoApproveInterviewer; use attractor::interviewer::queue::QueueInterviewer; -use attractor::interviewer::{Answer, AnswerValue}; +use attractor::interviewer::recording::RecordingInterviewer; +use attractor::interviewer::{Answer, AnswerValue, Interviewer}; use attractor::outcome::{Outcome, StageStatus}; use attractor::parser::parse; use attractor::stylesheet::{apply_stylesheet, parse_stylesheet}; use attractor::transform::{StylesheetApplicationTransform, Transform, VariableExpansionTransform}; -use attractor::validation::validate_or_raise; +use attractor::validation::{validate, validate_or_raise, Severity}; // --------------------------------------------------------------------------- // 1. Parse and validate all 3 spec examples (Section 2.13) @@ -1041,6 +1045,97 @@ impl CodergenBackend for MockCodergenBackend { } } +// --------------------------------------------------------------------------- +// Helpers for parity tests +// --------------------------------------------------------------------------- + +/// A handler backed by a shared AtomicU32 counter. +/// Returns Fail on call 0, Success on call >= 1. +struct CounterHandler { + call_count: Arc, +} + +#[async_trait::async_trait] +impl Handler for CounterHandler { + async fn execute( + &self, + _node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + let count = self + .call_count + .fetch_add(1, std::sync::atomic::Ordering::SeqCst); + if count == 0 { + Ok(Outcome::fail("first call fails")) + } else { + Ok(Outcome::success()) + } + } +} + +/// A handler that sets context_updates = {"my_flag": "set"}. +struct ContextSetterHandler; + +#[async_trait::async_trait] +impl Handler for ContextSetterHandler { + async fn execute( + &self, + _node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + let mut outcome = Outcome::success(); + outcome + .context_updates + .insert("my_flag".to_string(), serde_json::json!("set")); + Ok(outcome) + } +} + +fn collect_events(emitter: &mut EventEmitter) -> Arc>> { + let events = Arc::new(std::sync::Mutex::new(Vec::new())); + let events_clone = Arc::clone(&events); + emitter.on_event(move |event| { + events_clone.lock().unwrap().push(event.clone()); + }); + events +} + +fn make_full_registry(interviewer: Arc) -> HandlerRegistry { + let mut registry = HandlerRegistry::new(Box::new(CodergenHandler::new(None))); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register("codergen", Box::new(CodergenHandler::new(None))); + registry.register("conditional", Box::new(ConditionalHandler)); + registry.register("tool", Box::new(ToolHandler)); + registry.register( + "wait.human", + Box::new(WaitHumanHandler::new(interviewer)), + ); + registry.register( + "stack.manager_loop", + Box::new(ManagerLoopHandler::new(None)), + ); + registry +} + +fn make_graph_with_start_exit(name: &str) -> Graph { + let mut graph = Graph::new(name); + 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); + graph +} + #[tokio::test] async fn smoke_test_with_mock_codergen_backend() { // Pipeline: @@ -1445,3 +1540,1315 @@ async fn resume_from_checkpoint_preserves_goal_gate_outcomes() { .expect("resume with goal gate should succeed"); assert_eq!(outcome.status, StageStatus::Success); } + +// =========================================================================== +// Parity tests — P1: Core pipeline behaviors +// =========================================================================== + +#[tokio::test] +async fn graph_goal_in_context() { + let input = r#"digraph GoalTest { + graph [goal="Ship the widget"] + start [shape=Mdiamond] + exit [shape=Msquare] + work [shape=box, prompt="Build it"] + start -> work -> exit + }"#; + let graph = parse(input).expect("parse"); + let dir = tempfile::tempdir().unwrap(); + let engine = PipelineEngine::new(make_linear_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&graph, &config).await.expect("run"); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert_eq!( + cp.context_values.get("graph.goal"), + Some(&serde_json::json!("Ship the widget")) + ); +} + +#[tokio::test] +async fn event_streaming_lifecycle() { + let input = r#"digraph EventTest { + start [shape=Mdiamond] + exit [shape=Msquare] + task [shape=box, prompt="Do something"] + start -> task -> exit + }"#; + let graph = parse(input).expect("parse"); + let dir = tempfile::tempdir().unwrap(); + let mut emitter = EventEmitter::new(); + let events = collect_events(&mut emitter); + let engine = PipelineEngine::new(make_linear_registry(), emitter); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&graph, &config).await.expect("run"); + + let collected = events.lock().unwrap(); + assert!(collected + .iter() + .any(|e| matches!(e, PipelineEvent::PipelineStarted { .. }))); + assert!(collected + .iter() + .any(|e| matches!(e, PipelineEvent::StageStarted { name, .. } if name == "start"))); + assert!(collected + .iter() + .any(|e| matches!(e, PipelineEvent::StageCompleted { name, .. } if name == "start"))); + assert!(collected + .iter() + .any(|e| matches!(e, PipelineEvent::StageStarted { name, .. } if name == "task"))); + assert!(collected + .iter() + .any(|e| matches!(e, PipelineEvent::StageCompleted { name, .. } if name == "task"))); + assert!(collected + .iter() + .any(|e| matches!(e, PipelineEvent::CheckpointSaved { .. }))); + assert!(collected + .iter() + .any(|e| matches!(e, PipelineEvent::PipelineCompleted { .. }))); + // PipelineStarted first, PipelineCompleted last + assert!(matches!( + collected.first().unwrap(), + PipelineEvent::PipelineStarted { .. } + )); + assert!(matches!( + collected.last().unwrap(), + PipelineEvent::PipelineCompleted { .. } + )); +} + +#[tokio::test] +async fn context_flow_between_stages() { + let mut graph = make_graph_with_start_exit("ContextFlowTest"); + let mut step_a = Node::new("step_a"); + step_a + .attrs + .insert("shape".to_string(), AttrValue::String("box".to_string())); + step_a.attrs.insert( + "prompt".to_string(), + AttrValue::String("Step A work".to_string()), + ); + graph.nodes.insert("step_a".to_string(), step_a); + let mut step_b = Node::new("step_b"); + step_b + .attrs + .insert("shape".to_string(), AttrValue::String("box".to_string())); + step_b.attrs.insert( + "prompt".to_string(), + AttrValue::String("Step B work".to_string()), + ); + graph.nodes.insert("step_b".to_string(), step_b); + graph.edges.push(Edge::new("start", "step_a")); + graph.edges.push(Edge::new("step_a", "step_b")); + graph.edges.push(Edge::new("step_b", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let engine = PipelineEngine::new(make_linear_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&graph, &config).await.expect("run"); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert_eq!( + cp.context_values.get("last_stage"), + Some(&serde_json::json!("step_b")) + ); + let last_response = cp + .context_values + .get("last_response") + .unwrap() + .as_str() + .unwrap(); + assert!(last_response.contains("[Simulated]")); +} + +#[tokio::test] +async fn tool_handler_e2e() { + let mut graph = make_graph_with_start_exit("ToolTest"); + let mut echo_task = Node::new("echo_task"); + echo_task.attrs.insert( + "shape".to_string(), + AttrValue::String("parallelogram".to_string()), + ); + echo_task.attrs.insert( + "tool_command".to_string(), + AttrValue::String("echo hello-from-tool".to_string()), + ); + graph.nodes.insert("echo_task".to_string(), echo_task); + graph.edges.push(Edge::new("start", "echo_task")); + graph.edges.push(Edge::new("echo_task", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let interviewer = Arc::new(AutoApproveInterviewer); + let engine = PipelineEngine::new(make_full_registry(interviewer), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + let tool_output = cp + .context_values + .get("tool.output") + .expect("tool.output should exist"); + assert!(tool_output.as_str().unwrap().contains("hello-from-tool")); +} + +#[tokio::test] +async fn auto_approve_interviewer_e2e() { + let mut graph = make_graph_with_start_exit("AutoApproveTest"); + let mut gate = Node::new("gate"); + gate.attrs.insert( + "shape".to_string(), + AttrValue::String("hexagon".to_string()), + ); + gate.attrs.insert( + "type".to_string(), + AttrValue::String("wait.human".to_string()), + ); + gate.attrs.insert( + "label".to_string(), + AttrValue::String("Review".to_string()), + ); + graph.nodes.insert("gate".to_string(), gate); + graph + .nodes + .insert("approve".to_string(), Node::new("approve")); + graph + .nodes + .insert("reject".to_string(), Node::new("reject")); + graph.edges.push(Edge::new("start", "gate")); + let mut e_approve = Edge::new("gate", "approve"); + e_approve.attrs.insert( + "label".to_string(), + AttrValue::String("[A] Approve".to_string()), + ); + graph.edges.push(e_approve); + let mut e_reject = Edge::new("gate", "reject"); + e_reject.attrs.insert( + "label".to_string(), + AttrValue::String("[R] Reject".to_string()), + ); + graph.edges.push(e_reject); + graph.edges.push(Edge::new("approve", "exit")); + graph.edges.push(Edge::new("reject", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let interviewer = Arc::new(AutoApproveInterviewer); + let engine = PipelineEngine::new(make_full_registry(interviewer), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!(cp.completed_nodes.contains(&"approve".to_string())); + assert!(!cp.completed_nodes.contains(&"reject".to_string())); +} + +#[tokio::test] +async fn codergen_without_backend_simulated() { + let input = r#"digraph SimTest { + start [shape=Mdiamond] + exit [shape=Msquare] + code [shape=box, prompt="Write the code"] + start -> code -> exit + }"#; + let graph = parse(input).expect("parse"); + let dir = tempfile::tempdir().unwrap(); + let engine = PipelineEngine::new(make_linear_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&graph, &config).await.expect("run"); + + let response = + std::fs::read_to_string(dir.path().join("code").join("response.md")).unwrap(); + assert!(response.contains("[Simulated]")); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + let last_response = cp + .context_values + .get("last_response") + .unwrap() + .as_str() + .unwrap(); + assert!(last_response.contains("[Simulated]")); +} + +// =========================================================================== +// Parity tests — P2: Complex scenarios +// =========================================================================== + +#[tokio::test] +async fn branching_loop_back_on_failure() { + struct FailThenSucceedHandler { + call_count: std::sync::atomic::AtomicU32, + } + + #[async_trait::async_trait] + impl Handler for FailThenSucceedHandler { + async fn execute( + &self, + _node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + let count = self + .call_count + .fetch_add(1, std::sync::atomic::Ordering::SeqCst); + if count == 0 { + Ok(Outcome::fail("first attempt fails")) + } else { + Ok(Outcome::success()) + } + } + } + + let mut graph = make_graph_with_start_exit("LoopTest"); + let mut implement = Node::new("implement"); + implement + .attrs + .insert("shape".to_string(), AttrValue::String("box".to_string())); + implement.attrs.insert( + "prompt".to_string(), + AttrValue::String("Implement".to_string()), + ); + graph.nodes.insert("implement".to_string(), implement); + let mut validate_node = Node::new("validate"); + validate_node.attrs.insert( + "type".to_string(), + AttrValue::String("fail_then_succeed".to_string()), + ); + graph + .nodes + .insert("validate".to_string(), validate_node); + + graph.edges.push(Edge::new("start", "implement")); + graph.edges.push(Edge::new("implement", "validate")); + let mut e_success = Edge::new("validate", "exit"); + e_success.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + graph.edges.push(e_success); + let mut e_fail = Edge::new("validate", "implement"); + e_fail.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=fail".to_string()), + ); + graph.edges.push(e_fail); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(CodergenHandler::new(None))); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register("codergen", Box::new(CodergenHandler::new(None))); + registry.register( + "fail_then_succeed", + Box::new(FailThenSucceedHandler { + call_count: std::sync::atomic::AtomicU32::new(0), + }), + ); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + let implement_count = cp + .completed_nodes + .iter() + .filter(|n| *n == "implement") + .count(); + assert!( + implement_count >= 2, + "implement should appear at least 2x, got {implement_count}" + ); +} + +#[tokio::test] +async fn human_gate_loops_back() { + let mut graph = make_graph_with_start_exit("HumanLoopTest"); + let mut gate = Node::new("gate"); + gate.attrs.insert( + "shape".to_string(), + AttrValue::String("hexagon".to_string()), + ); + gate.attrs.insert( + "type".to_string(), + AttrValue::String("wait.human".to_string()), + ); + gate.attrs.insert( + "label".to_string(), + AttrValue::String("Review".to_string()), + ); + graph.nodes.insert("gate".to_string(), gate); + graph + .nodes + .insert("approve".to_string(), Node::new("approve")); + graph.nodes.insert("fix".to_string(), Node::new("fix")); + + graph.edges.push(Edge::new("start", "gate")); + let mut e_approve = Edge::new("gate", "approve"); + e_approve.attrs.insert( + "label".to_string(), + AttrValue::String("[A] Approve".to_string()), + ); + graph.edges.push(e_approve); + let mut e_fix = Edge::new("gate", "fix"); + e_fix.attrs.insert( + "label".to_string(), + AttrValue::String("[F] Fix".to_string()), + ); + graph.edges.push(e_fix); + graph.edges.push(Edge::new("fix", "gate")); + graph.edges.push(Edge::new("approve", "exit")); + + let answers = VecDeque::from([ + Answer { + value: AnswerValue::Selected("F".to_string()), + selected_option: None, + text: None, + }, + Answer { + value: AnswerValue::Selected("A".to_string()), + selected_option: None, + text: None, + }, + ]); + let interviewer = Arc::new(QueueInterviewer::new(answers)); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register( + "wait.human", + Box::new(WaitHumanHandler::new(interviewer)), + ); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + let gate_count = cp + .completed_nodes + .iter() + .filter(|n| *n == "gate") + .count(); + assert!( + gate_count >= 2, + "gate should appear at least 2x, got {gate_count}" + ); + assert!(cp.completed_nodes.contains(&"approve".to_string())); +} + +#[tokio::test] +async fn scenario_ship_a_feature() { + let dot = r#"digraph ShipFeature { + graph [goal="Ship the widget"] + rankdir=LR + start [shape=Mdiamond] + exit [shape=Msquare] + plan [shape=box, prompt="Plan to achieve: $goal"] + implement [shape=box, prompt="Implement the plan"] + test [shape=parallelogram, tool_command="echo PASS"] + review [shape=hexagon, label="Review Changes"] + start -> plan -> implement -> test -> review + review -> exit [label="[A] Approve"] + review -> implement [label="[F] Fix"] + }"#; + let mut graph = parse(dot).expect("parse"); + validate_or_raise(&graph, &[]).expect("validate"); + VariableExpansionTransform.apply(&mut graph); + assert_eq!( + graph.nodes["plan"].prompt().unwrap(), + "Plan to achieve: Ship the widget" + ); + + let interviewer = Arc::new(AutoApproveInterviewer); + let dir = tempfile::tempdir().unwrap(); + let mut emitter = EventEmitter::new(); + let events = collect_events(&mut emitter); + let engine = PipelineEngine::new(make_full_registry(interviewer), emitter); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + let tool_output = cp.context_values.get("tool.output").expect("tool.output"); + assert!(tool_output.as_str().unwrap().contains("PASS")); + assert!(cp.completed_nodes.contains(&"plan".to_string())); + assert!(cp.completed_nodes.contains(&"implement".to_string())); + assert!(cp.completed_nodes.contains(&"test".to_string())); + assert!(cp.completed_nodes.contains(&"review".to_string())); + + let collected = events.lock().unwrap(); + assert!(collected + .iter() + .any(|e| matches!(e, PipelineEvent::PipelineStarted { .. }))); + assert!(collected + .iter() + .any(|e| matches!(e, PipelineEvent::PipelineCompleted { .. }))); +} + +#[tokio::test] +async fn scenario_parallel_expert_review() { + use attractor::handler::fan_in::FanInHandler; + use attractor::handler::parallel::ParallelHandler; + + let input = r#"digraph ParallelReview { + start [shape=Mdiamond] + fan_out [shape=component] + expert_a [shape=box, prompt="Expert A review"] + expert_b [shape=box, prompt="Expert B review"] + expert_c [shape=box, prompt="Expert C review"] + fan_in_node [shape=tripleoctagon] + review [shape=hexagon, label="Final Review"] + exit [shape=Msquare] + start -> fan_out + fan_out -> expert_a + fan_out -> expert_b + fan_out -> expert_c + expert_a -> fan_in_node + expert_b -> fan_in_node + expert_c -> fan_in_node + fan_in_node -> review + review -> exit [label="[A] Approve"] + review -> fan_out [label="[F] Redo"] + }"#; + let graph = parse(input).expect("parse"); + validate_or_raise(&graph, &[]).expect("validate"); + + let recorder = Arc::new(RecordingInterviewer::new(Box::new(AutoApproveInterviewer))); + let dir = tempfile::tempdir().unwrap(); + + let mut base_registry = HandlerRegistry::new(Box::new(CodergenHandler::new(Some( + Box::new(MockCodergenBackend), + )))); + base_registry.register("start", Box::new(StartHandler)); + base_registry.register("exit", Box::new(ExitHandler)); + base_registry.register( + "codergen", + Box::new(CodergenHandler::new(Some(Box::new(MockCodergenBackend)))), + ); + let base_registry = Arc::new(base_registry); + let emitter = Arc::new(EventEmitter::new()); + let parallel_handler = + ParallelHandler::new(Arc::clone(&base_registry), Arc::clone(&emitter)); + let fan_in_handler = FanInHandler::new(Some(Box::new(MockCodergenBackend))); + + let interviewer: Arc = recorder.clone(); + let mut full_registry = HandlerRegistry::new(Box::new(CodergenHandler::new(Some( + Box::new(MockCodergenBackend), + )))); + full_registry.register("start", Box::new(StartHandler)); + full_registry.register("exit", Box::new(ExitHandler)); + full_registry.register( + "codergen", + Box::new(CodergenHandler::new(Some(Box::new(MockCodergenBackend)))), + ); + full_registry.register("parallel", Box::new(parallel_handler)); + full_registry.register("parallel.fan_in", Box::new(fan_in_handler)); + full_registry.register( + "wait.human", + Box::new(WaitHumanHandler::new(interviewer)), + ); + + let engine = PipelineEngine::new(full_registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + let results = cp + .context_values + .get("parallel.results") + .expect("parallel.results"); + assert_eq!(results.as_array().unwrap().len(), 3); + + let recordings = recorder.recordings(); + assert_eq!(recordings.len(), 1, "should have 1 interview recording"); + assert!(cp.completed_nodes.contains(&"review".to_string())); +} + +#[tokio::test] +async fn scenario_node_retries_on_retry_status() { + struct RetryHandler { + call_count: std::sync::atomic::AtomicU32, + } + + #[async_trait::async_trait] + impl Handler for RetryHandler { + async fn execute( + &self, + _node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + let count = self + .call_count + .fetch_add(1, std::sync::atomic::Ordering::SeqCst); + if count == 0 { + Ok(Outcome::retry("transient failure")) + } else { + Ok(Outcome::success()) + } + } + } + + let mut graph = make_graph_with_start_exit("RetryScenarioTest"); + let mut flaky = Node::new("flaky"); + flaky.attrs.insert( + "type".to_string(), + AttrValue::String("retry_handler".to_string()), + ); + flaky + .attrs + .insert("max_retries".to_string(), AttrValue::Integer(2)); + graph.nodes.insert("flaky".to_string(), flaky); + graph.edges.push(Edge::new("start", "flaky")); + graph.edges.push(Edge::new("flaky", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register( + "retry_handler", + Box::new(RetryHandler { + call_count: std::sync::atomic::AtomicU32::new(0), + }), + ); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + let retry_count = cp + .node_retries + .get("flaky") + .expect("flaky should have retries"); + assert_eq!(*retry_count, 2, "should have been called 2x"); +} + +#[tokio::test] +async fn scenario_loop_restart_resets_context() { + let mut graph = make_graph_with_start_exit("LoopRestartTest"); + let mut work = Node::new("work"); + work.attrs.insert( + "type".to_string(), + AttrValue::String("counter".to_string()), + ); + graph.nodes.insert("work".to_string(), work); + + graph.edges.push(Edge::new("start", "work")); + let mut success_edge = Edge::new("work", "exit"); + success_edge.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + graph.edges.push(success_edge); + let mut fail_edge = Edge::new("work", "start"); + fail_edge.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=fail".to_string()), + ); + fail_edge + .attrs + .insert("loop_restart".to_string(), AttrValue::Boolean(true)); + graph.edges.push(fail_edge); + + let call_count = Arc::new(std::sync::atomic::AtomicU32::new(0)); + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register( + "counter", + Box::new(CounterHandler { + call_count: Arc::clone(&call_count), + }), + ); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + assert!(call_count.load(std::sync::atomic::Ordering::SeqCst) >= 2); +} + +#[tokio::test] +async fn scenario_bug_triage_router() { + let mut graph = make_graph_with_start_exit("TriageTest"); + let mut triage = Node::new("triage"); + triage.attrs.insert( + "shape".to_string(), + AttrValue::String("diamond".to_string()), + ); + graph.nodes.insert("triage".to_string(), triage); + graph + .nodes + .insert("critical".to_string(), Node::new("critical")); + graph + .nodes + .insert("normal".to_string(), Node::new("normal")); + graph + .nodes + .insert("wontfix".to_string(), Node::new("wontfix")); + + graph.edges.push(Edge::new("start", "triage")); + let mut e_critical = Edge::new("triage", "critical"); + e_critical.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + e_critical + .attrs + .insert("weight".to_string(), AttrValue::Integer(10)); + graph.edges.push(e_critical); + let mut e_normal = Edge::new("triage", "normal"); + e_normal.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + e_normal + .attrs + .insert("weight".to_string(), AttrValue::Integer(5)); + graph.edges.push(e_normal); + graph.edges.push(Edge::new("triage", "wontfix")); + graph.edges.push(Edge::new("critical", "exit")); + graph.edges.push(Edge::new("normal", "exit")); + graph.edges.push(Edge::new("wontfix", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register("conditional", Box::new(ConditionalHandler)); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!( + cp.completed_nodes.contains(&"critical".to_string()), + "critical should be selected (highest weight)" + ); + assert!(!cp.completed_nodes.contains(&"normal".to_string())); + assert!(!cp.completed_nodes.contains(&"wontfix".to_string())); +} + +#[tokio::test] +async fn scenario_crash_recovery() { + let mut graph = make_graph_with_start_exit("CrashRecoveryTest"); + graph.nodes.insert("a".to_string(), Node::new("a")); + graph.nodes.insert("b".to_string(), Node::new("b")); + graph.nodes.insert("c".to_string(), Node::new("c")); + graph.edges.push(Edge::new("start", "a")); + graph.edges.push(Edge::new("a", "b")); + graph.edges.push(Edge::new("b", "c")); + graph.edges.push(Edge::new("c", "exit")); + + let ctx = Context::new(); + ctx.set("outcome", serde_json::json!("success")); + let mut outcomes = std::collections::HashMap::new(); + outcomes.insert("start".to_string(), Outcome::success()); + outcomes.insert("a".to_string(), Outcome::success()); + let checkpoint = Checkpoint::from_context( + &ctx, + "a", + vec!["start".to_string(), "a".to_string()], + std::collections::HashMap::new(), + outcomes, + Some("b".to_string()), + ); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine + .run_from_checkpoint(&graph, &config, &checkpoint) + .await + .expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!(cp.completed_nodes.contains(&"b".to_string())); + assert!(cp.completed_nodes.contains(&"c".to_string())); + assert!(cp.completed_nodes.contains(&"a".to_string())); + let a_count = cp.completed_nodes.iter().filter(|n| *n == "a").count(); + assert_eq!(a_count, 1, "a should not be re-executed"); +} + +#[tokio::test] +async fn manager_loop_stop_condition_satisfied_e2e() { + struct DoneSetterHandler; + + #[async_trait::async_trait] + impl Handler for DoneSetterHandler { + async fn execute( + &self, + _node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + let mut outcome = Outcome::success(); + outcome + .context_updates + .insert("done".to_string(), serde_json::json!("true")); + Ok(outcome) + } + } + + let mut graph = make_graph_with_start_exit("ManagerStopTest"); + let mut setter = Node::new("setter"); + setter.attrs.insert( + "type".to_string(), + AttrValue::String("done_setter".to_string()), + ); + graph.nodes.insert("setter".to_string(), setter); + let mut manager = Node::new("manager"); + manager.attrs.insert( + "type".to_string(), + AttrValue::String("stack.manager_loop".to_string()), + ); + manager.attrs.insert( + "manager.stop_condition".to_string(), + AttrValue::String("context.done=true".to_string()), + ); + manager + .attrs + .insert("manager.max_cycles".to_string(), AttrValue::Integer(10)); + manager.attrs.insert( + "manager.poll_interval".to_string(), + AttrValue::Duration(std::time::Duration::from_millis(1)), + ); + graph.nodes.insert("manager".to_string(), manager); + graph.edges.push(Edge::new("start", "setter")); + graph.edges.push(Edge::new("setter", "manager")); + graph.edges.push(Edge::new("manager", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register("done_setter", Box::new(DoneSetterHandler)); + registry.register( + "stack.manager_loop", + Box::new(ManagerLoopHandler::new(None)), + ); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + let manager_outcome = cp.node_outcomes.get("manager").expect("manager outcome"); + assert_eq!(manager_outcome.status, StageStatus::Success); + assert!(manager_outcome + .notes + .as_deref() + .unwrap() + .contains("Stop condition satisfied")); + // Overall pipeline succeeds because manager succeeded + assert_eq!(outcome.status, StageStatus::Success); +} + +#[tokio::test] +async fn manager_loop_max_cycles_exceeded_e2e() { + let mut graph = make_graph_with_start_exit("ManagerMaxCyclesTest"); + let mut manager = Node::new("manager"); + manager.attrs.insert( + "type".to_string(), + AttrValue::String("stack.manager_loop".to_string()), + ); + manager + .attrs + .insert("manager.max_cycles".to_string(), AttrValue::Integer(2)); + manager.attrs.insert( + "manager.poll_interval".to_string(), + AttrValue::Duration(std::time::Duration::from_millis(1)), + ); + graph.nodes.insert("manager".to_string(), manager); + graph.edges.push(Edge::new("start", "manager")); + graph.edges.push(Edge::new("manager", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register( + "stack.manager_loop", + Box::new(ManagerLoopHandler::new(None)), + ); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + let manager_outcome = cp.node_outcomes.get("manager").expect("manager outcome"); + assert_eq!(manager_outcome.status, StageStatus::Fail); + assert!(manager_outcome + .failure_reason + .as_deref() + .unwrap() + .contains("Max cycles")); + // Overall pipeline outcome is from last completed node (manager) = Fail + assert_eq!(outcome.status, StageStatus::Fail); +} + +// =========================================================================== +// Parity tests — P3: Validation +// =========================================================================== + +#[test] +fn validation_missing_start_node() { + let mut graph = Graph::new("NoStartTest"); + 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 diagnostics = validate(&graph, &[]); + let start_errors: Vec<_> = diagnostics + .iter() + .filter(|d| d.severity == Severity::Error && d.rule == "start_node") + .collect(); + assert!( + !start_errors.is_empty(), + "should have start_node error diagnostic" + ); +} + +#[test] +fn validation_missing_exit_node() { + let mut graph = Graph::new("NoExitTest"); + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + graph.nodes.insert("start".to_string(), start); + graph + .nodes + .insert("work".to_string(), Node::new("work")); + graph.edges.push(Edge::new("start", "work")); + + let diagnostics = validate(&graph, &[]); + let exit_errors: Vec<_> = diagnostics + .iter() + .filter(|d| d.severity == Severity::Error && d.rule == "terminal_node") + .collect(); + assert!( + !exit_errors.is_empty(), + "should have terminal_node error diagnostic" + ); +} + +#[test] +fn validation_orphan_unreachable_node() { + let mut graph = make_graph_with_start_exit("OrphanTest"); + graph + .nodes + .insert("orphan".to_string(), Node::new("orphan")); + graph.edges.push(Edge::new("start", "exit")); + + let diagnostics = validate(&graph, &[]); + let reachability_errors: Vec<_> = diagnostics + .iter() + .filter(|d| d.rule == "reachability") + .collect(); + assert!( + !reachability_errors.is_empty(), + "should have reachability diagnostic for orphan node" + ); +} + +// =========================================================================== +// Parity tests — P4: Edge selection and cross-feature +// =========================================================================== + +#[tokio::test] +async fn conditional_branching_success_fail_paths() { + let mut graph = make_graph_with_start_exit("CondBranchTest"); + let mut work = Node::new("work"); + work.attrs.insert( + "type".to_string(), + AttrValue::String("always_fail".to_string()), + ); + graph.nodes.insert("work".to_string(), work); + graph + .nodes + .insert("success_path".to_string(), Node::new("success_path")); + graph + .nodes + .insert("fail_path".to_string(), Node::new("fail_path")); + + graph.edges.push(Edge::new("start", "work")); + let mut e_success = Edge::new("work", "success_path"); + e_success.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + graph.edges.push(e_success); + let mut e_fail = Edge::new("work", "fail_path"); + e_fail.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=fail".to_string()), + ); + graph.edges.push(e_fail); + graph.edges.push(Edge::new("success_path", "exit")); + graph.edges.push(Edge::new("fail_path", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register("always_fail", Box::new(AlwaysFailHandler)); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!(cp.completed_nodes.contains(&"fail_path".to_string())); + assert!(!cp.completed_nodes.contains(&"success_path".to_string())); +} + +#[tokio::test] +async fn edge_selection_condition_match_wins_over_weight() { + let mut graph = make_graph_with_start_exit("CondVsWeightTest"); + graph.nodes.insert("a".to_string(), Node::new("a")); + graph + .nodes + .insert("cond_target".to_string(), Node::new("cond_target")); + graph.nodes.insert( + "weighted_target".to_string(), + Node::new("weighted_target"), + ); + + graph.edges.push(Edge::new("start", "a")); + let mut e_cond = Edge::new("a", "cond_target"); + e_cond.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + graph.edges.push(e_cond); + let mut e_weight = Edge::new("a", "weighted_target"); + e_weight + .attrs + .insert("weight".to_string(), AttrValue::Integer(100)); + graph.edges.push(e_weight); + graph.edges.push(Edge::new("cond_target", "exit")); + graph.edges.push(Edge::new("weighted_target", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&graph, &config).await.expect("run"); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!(cp.completed_nodes.contains(&"cond_target".to_string())); + assert!(!cp + .completed_nodes + .contains(&"weighted_target".to_string())); +} + +#[tokio::test] +async fn edge_selection_weight_breaks_ties() { + let mut graph = make_graph_with_start_exit("WeightTiesTest"); + graph.nodes.insert("a".to_string(), Node::new("a")); + graph.nodes.insert("low".to_string(), Node::new("low")); + graph.nodes.insert("high".to_string(), Node::new("high")); + + graph.edges.push(Edge::new("start", "a")); + let mut e_low = Edge::new("a", "low"); + e_low + .attrs + .insert("weight".to_string(), AttrValue::Integer(1)); + graph.edges.push(e_low); + let mut e_high = Edge::new("a", "high"); + e_high + .attrs + .insert("weight".to_string(), AttrValue::Integer(10)); + graph.edges.push(e_high); + graph.edges.push(Edge::new("low", "exit")); + graph.edges.push(Edge::new("high", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&graph, &config).await.expect("run"); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!(cp.completed_nodes.contains(&"high".to_string())); + assert!(!cp.completed_nodes.contains(&"low".to_string())); +} + +#[tokio::test] +async fn edge_selection_lexical_tiebreak() { + let mut graph = make_graph_with_start_exit("LexicalTieTest"); + graph.nodes.insert("a".to_string(), Node::new("a")); + graph.nodes.insert("beta".to_string(), Node::new("beta")); + graph + .nodes + .insert("alpha".to_string(), Node::new("alpha")); + + graph.edges.push(Edge::new("start", "a")); + graph.edges.push(Edge::new("a", "beta")); + graph.edges.push(Edge::new("a", "alpha")); + graph.edges.push(Edge::new("beta", "exit")); + graph.edges.push(Edge::new("alpha", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&graph, &config).await.expect("run"); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!(cp.completed_nodes.contains(&"alpha".to_string())); + assert!(!cp.completed_nodes.contains(&"beta".to_string())); +} + +#[tokio::test] +async fn context_updates_visible_across_nodes() { + let mut graph = make_graph_with_start_exit("ContextVisibilityTest"); + let mut setter = Node::new("setter"); + setter.attrs.insert( + "type".to_string(), + AttrValue::String("context_setter".to_string()), + ); + graph.nodes.insert("setter".to_string(), setter); + let mut gate = Node::new("gate"); + gate.attrs.insert( + "shape".to_string(), + AttrValue::String("diamond".to_string()), + ); + graph.nodes.insert("gate".to_string(), gate); + graph.nodes.insert("yes".to_string(), Node::new("yes")); + graph.nodes.insert("no".to_string(), Node::new("no")); + + graph.edges.push(Edge::new("start", "setter")); + graph.edges.push(Edge::new("setter", "gate")); + let mut e_yes = Edge::new("gate", "yes"); + e_yes.attrs.insert( + "condition".to_string(), + AttrValue::String("context.my_flag=set".to_string()), + ); + graph.edges.push(e_yes); + graph.edges.push(Edge::new("gate", "no")); + graph.edges.push(Edge::new("yes", "exit")); + graph.edges.push(Edge::new("no", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register("conditional", Box::new(ConditionalHandler)); + registry.register("context_setter", Box::new(ContextSetterHandler)); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&graph, &config).await.expect("run"); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!(cp.completed_nodes.contains(&"yes".to_string())); + assert!(!cp.completed_nodes.contains(&"no".to_string())); +} + +#[tokio::test] +async fn stylesheet_applies_model_override() { + let input = r#"digraph StylesheetTest { + graph [ + goal="Test stylesheet", + model_stylesheet="* { llm_model: custom-model; }" + ] + start [shape=Mdiamond] + exit [shape=Msquare] + work [shape=box, prompt="Do work"] + start -> work -> exit + }"#; + let mut graph = parse(input).expect("parse"); + validate_or_raise(&graph, &[]).expect("validate"); + StylesheetApplicationTransform.apply(&mut graph); + assert_eq!(graph.nodes["work"].llm_model(), Some("custom-model")); + + let dir = tempfile::tempdir().unwrap(); + let engine = PipelineEngine::new(make_linear_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); +} + +#[tokio::test] +async fn custom_handler_registration_and_execution() { + struct CustomHandler; + + #[async_trait::async_trait] + impl Handler for CustomHandler { + async fn execute( + &self, + _node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + let mut outcome = Outcome::success(); + outcome + .context_updates + .insert("custom.ran".to_string(), serde_json::json!("true")); + Ok(outcome) + } + } + + let mut graph = make_graph_with_start_exit("CustomHandlerTest"); + let mut custom = Node::new("custom"); + custom.attrs.insert( + "type".to_string(), + AttrValue::String("my_custom".to_string()), + ); + graph.nodes.insert("custom".to_string(), custom); + graph.edges.push(Edge::new("start", "custom")); + graph.edges.push(Edge::new("custom", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register("my_custom", Box::new(CustomHandler)); + let engine = PipelineEngine::new(registry, EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&graph, &config).await.expect("run"); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert_eq!( + cp.context_values.get("custom.ran"), + Some(&serde_json::json!("true")) + ); +} + +#[tokio::test] +async fn integration_smoke_plan_implement_review_done() { + let dot = r#"digraph SmokeIntegration { + graph [ + goal="Build the feature", + model_stylesheet="* { llm_model: test-model; }" + ] + rankdir=LR + start [shape=Mdiamond] + exit [shape=Msquare] + plan [shape=box, prompt="Plan: $goal"] + implement [shape=box, prompt="Implement"] + review [shape=hexagon, label="Review"] + start -> plan -> implement -> review + review -> exit [label="[A] Approve"] + review -> implement [label="[F] Fix"] + }"#; + + // Parse and validate + let mut graph = parse(dot).expect("parse"); + let diagnostics = validate_or_raise(&graph, &[]).expect("validate"); + let errors: Vec<_> = diagnostics + .iter() + .filter(|d| d.severity == Severity::Error) + .collect(); + assert!(errors.is_empty()); + + // Apply transforms + VariableExpansionTransform.apply(&mut graph); + StylesheetApplicationTransform.apply(&mut graph); + + // Verify transforms applied + assert_eq!( + graph.nodes["plan"].prompt().unwrap(), + "Plan: Build the feature" + ); + assert_eq!(graph.nodes["plan"].llm_model(), Some("test-model")); + + // Run pipeline + let interviewer = Arc::new(AutoApproveInterviewer); + let dir = tempfile::tempdir().unwrap(); + let mut emitter = EventEmitter::new(); + let events = collect_events(&mut emitter); + let engine = PipelineEngine::new(make_full_registry(interviewer), emitter); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&graph, &config).await.expect("run"); + assert_eq!(outcome.status, StageStatus::Success); + + // Verify all nodes completed + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!(cp.completed_nodes.contains(&"plan".to_string())); + assert!(cp.completed_nodes.contains(&"implement".to_string())); + assert!(cp.completed_nodes.contains(&"review".to_string())); + + // Verify prompt.md and response.md exist + assert!(dir.path().join("plan").join("prompt.md").exists()); + assert!(dir.path().join("plan").join("response.md").exists()); + + // Verify events + let collected = events.lock().unwrap(); + assert!(collected + .iter() + .any(|e| matches!(e, PipelineEvent::PipelineStarted { .. }))); + assert!(collected + .iter() + .any(|e| matches!(e, PipelineEvent::PipelineCompleted { .. }))); +}