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 { .. }))); +}