diff --git a/Cargo.lock b/Cargo.lock index dd84cda23..dd1c48513 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -135,6 +135,7 @@ dependencies = [ "async-trait", "chrono", "coding-agent-loop", + "futures", "nom", "rand", "serde", diff --git a/crates/attractor/Cargo.toml b/crates/attractor/Cargo.toml new file mode 100644 index 000000000..a42c38bcc --- /dev/null +++ b/crates/attractor/Cargo.toml @@ -0,0 +1,31 @@ +[package] +name = "attractor" +edition.workspace = true +version.workspace = true +license.workspace = true +description = "A DOT-based pipeline runner for multi-stage AI workflows" +repository = "https://github.com/brynary/attractor-rust" +keywords = ["llm", "ai", "pipeline", "workflow", "dot"] +categories = ["development-tools"] +readme = "README.md" + +[dependencies] +coding-agent-loop = { path = "../coding-agent-loop" } +unified-llm = { path = "../unified-llm" } +thiserror.workspace = true +serde.workspace = true +serde_json.workspace = true +tokio.workspace = true +uuid.workspace = true +rand.workspace = true +async-trait.workspace = true +futures.workspace = true +chrono = { workspace = true, features = ["serde"] } +nom = "7" + +[dev-dependencies] +tokio = { workspace = true, features = ["test-util", "macros"] } +tempfile = "3" + +[lints] +workspace = true diff --git a/crates/attractor/README.md b/crates/attractor/README.md new file mode 100644 index 000000000..b96a20a68 --- /dev/null +++ b/crates/attractor/README.md @@ -0,0 +1,3 @@ +# attractor + +A DOT-based pipeline runner for multi-stage AI workflows. diff --git a/crates/attractor/src/artifact.rs b/crates/attractor/src/artifact.rs new file mode 100644 index 000000000..f2078dbbc --- /dev/null +++ b/crates/attractor/src/artifact.rs @@ -0,0 +1,300 @@ +use std::collections::HashMap; +use std::path::PathBuf; +use std::sync::RwLock; + +use chrono::{DateTime, Utc}; +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +use crate::error::{AttractorError, Result}; + +/// Threshold above which artifacts are stored on disk instead of in memory (100KB). +const FILE_BACKING_THRESHOLD: usize = 100 * 1024; + +/// Metadata about a stored artifact. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ArtifactInfo { + pub id: String, + pub name: String, + pub size_bytes: usize, + pub stored_at: DateTime, + pub is_file_backed: bool, +} + +/// Storage for artifacts, either held in memory or backed by files on disk. +enum StoredData { + InMemory(Value), + FileBacked(PathBuf), +} + +/// Named, typed storage for large stage outputs. +pub struct ArtifactStore { + base_dir: Option, + artifacts: RwLock>, +} + +impl std::fmt::Debug for ArtifactStore { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ArtifactStore") + .field("base_dir", &self.base_dir) + .finish_non_exhaustive() + } +} + +impl ArtifactStore { + #[must_use] + pub fn new(base_dir: Option) -> Self { + Self { + base_dir, + artifacts: RwLock::new(HashMap::new()), + } + } + + /// Store an artifact. Large artifacts with a configured `base_dir` are written to disk. + /// + /// # Errors + /// + /// Returns an error if serialization fails or the file cannot be written. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + pub fn store(&self, id: impl Into, name: impl Into, data: Value) -> Result { + let id = id.into(); + let name = name.into(); + let serialized = serde_json::to_string(&data) + .map_err(|e| AttractorError::Engine(format!("artifact serialize failed: {e}")))?; + let size_bytes = serialized.len(); + + let is_file_backed = size_bytes > FILE_BACKING_THRESHOLD && self.base_dir.is_some(); + + let stored = if is_file_backed { + let base = self.base_dir.as_ref().expect("base_dir checked above"); + let artifacts_dir = base.join("artifacts"); + std::fs::create_dir_all(&artifacts_dir)?; + let file_path = artifacts_dir.join(format!("{id}.json")); + std::fs::write(&file_path, &serialized)?; + StoredData::FileBacked(file_path) + } else { + StoredData::InMemory(data) + }; + + let info = ArtifactInfo { + id: id.clone(), + name, + size_bytes, + stored_at: Utc::now(), + is_file_backed, + }; + + self.artifacts + .write() + .expect("artifact lock poisoned") + .insert(id, (info.clone(), stored)); + + Ok(info) + } + + /// Retrieve an artifact's data by ID. + /// + /// # Errors + /// + /// Returns an error if the artifact is not found or cannot be read from disk. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + pub fn retrieve(&self, id: &str) -> Result { + let guard = self.artifacts.read().expect("artifact lock poisoned"); + let (_, stored) = guard + .get(id) + .ok_or_else(|| AttractorError::Engine(format!("artifact not found: {id}")))?; + + match stored { + StoredData::InMemory(v) => Ok(v.clone()), + StoredData::FileBacked(path) => { + let path = path.clone(); + drop(guard); + let data = std::fs::read_to_string(&path).map_err(|e| { + AttractorError::Engine(format!( + "failed to read file-backed artifact {id}: {e}" + )) + })?; + serde_json::from_str(&data).map_err(|e| { + AttractorError::Engine(format!( + "failed to deserialize file-backed artifact {id}: {e}" + )) + }) + } + } + } + + /// Check if an artifact exists. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + pub fn has(&self, id: &str) -> bool { + self.artifacts + .read() + .expect("artifact lock poisoned") + .contains_key(id) + } + + /// List all artifact metadata. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + #[must_use] + pub fn list(&self) -> Vec { + self.artifacts + .read() + .expect("artifact lock poisoned") + .values() + .map(|(info, _)| info.clone()) + .collect() + } + + /// Remove an artifact by ID. Also deletes file-backed data from disk. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + pub fn remove(&self, id: &str) { + let mut guard = self.artifacts.write().expect("artifact lock poisoned"); + if let Some((_, StoredData::FileBacked(path))) = guard.remove(id) { + let _ = std::fs::remove_file(path); + } + } + + /// Remove all artifacts. Also deletes file-backed data from disk. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + pub fn clear(&self) { + let mut guard = self.artifacts.write().expect("artifact lock poisoned"); + for (_, stored) in guard.values() { + if let StoredData::FileBacked(path) = stored { + let _ = std::fs::remove_file(path); + } + } + guard.clear(); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn store_and_retrieve_small_artifact() { + let store = ArtifactStore::new(None); + let data = serde_json::json!({"result": "ok"}); + let info = store.store("art1", "test artifact", data.clone()).unwrap(); + + assert_eq!(info.id, "art1"); + assert_eq!(info.name, "test artifact"); + assert!(!info.is_file_backed); + assert!(info.size_bytes > 0); + + let retrieved = store.retrieve("art1").unwrap(); + assert_eq!(retrieved, data); + } + + #[test] + fn retrieve_nonexistent() { + let store = ArtifactStore::new(None); + assert!(store.retrieve("missing").is_err()); + } + + #[test] + fn has_artifact() { + let store = ArtifactStore::new(None); + assert!(!store.has("x")); + store.store("x", "x", serde_json::json!(1)).unwrap(); + assert!(store.has("x")); + } + + #[test] + fn list_artifacts() { + let store = ArtifactStore::new(None); + store.store("a", "alpha", serde_json::json!(1)).unwrap(); + store.store("b", "beta", serde_json::json!(2)).unwrap(); + let list = store.list(); + assert_eq!(list.len(), 2); + } + + #[test] + fn remove_artifact() { + let store = ArtifactStore::new(None); + store.store("r", "remove me", serde_json::json!(1)).unwrap(); + assert!(store.has("r")); + store.remove("r"); + assert!(!store.has("r")); + } + + #[test] + fn clear_artifacts() { + let store = ArtifactStore::new(None); + store.store("a", "a", serde_json::json!(1)).unwrap(); + store.store("b", "b", serde_json::json!(2)).unwrap(); + assert_eq!(store.list().len(), 2); + store.clear(); + assert!(store.list().is_empty()); + } + + #[test] + fn file_backed_storage() { + let dir = tempfile::tempdir().unwrap(); + let store = ArtifactStore::new(Some(dir.path().to_path_buf())); + + // Create data larger than the 100KB threshold + let large_string = "x".repeat(FILE_BACKING_THRESHOLD + 1); + let data = serde_json::json!(large_string); + + let info = store.store("big", "large artifact", data.clone()).unwrap(); + assert!(info.is_file_backed); + assert!(info.size_bytes > FILE_BACKING_THRESHOLD); + + let retrieved = store.retrieve("big").unwrap(); + assert_eq!(retrieved, data); + } + + #[test] + fn file_backed_remove_deletes_file() { + let dir = tempfile::tempdir().unwrap(); + let store = ArtifactStore::new(Some(dir.path().to_path_buf())); + + let large_string = "x".repeat(FILE_BACKING_THRESHOLD + 1); + let data = serde_json::json!(large_string); + store.store("big", "large", data).unwrap(); + + let file_path = dir.path().join("artifacts").join("big.json"); + assert!(file_path.exists()); + + store.remove("big"); + assert!(!file_path.exists()); + } + + #[test] + fn small_artifact_stays_in_memory_even_with_base_dir() { + let dir = tempfile::tempdir().unwrap(); + let store = ArtifactStore::new(Some(dir.path().to_path_buf())); + + let data = serde_json::json!({"small": true}); + let info = store.store("small", "tiny", data).unwrap(); + assert!(!info.is_file_backed); + } + + #[test] + fn no_file_backing_without_base_dir() { + let store = ArtifactStore::new(None); + + let large_string = "x".repeat(FILE_BACKING_THRESHOLD + 1); + let data = serde_json::json!(large_string); + let info = store.store("big", "large", data).unwrap(); + assert!(!info.is_file_backed); + } +} diff --git a/crates/attractor/src/checkpoint.rs b/crates/attractor/src/checkpoint.rs new file mode 100644 index 000000000..1dc790026 --- /dev/null +++ b/crates/attractor/src/checkpoint.rs @@ -0,0 +1,143 @@ +use std::collections::HashMap; +use std::path::Path; + +use chrono::{DateTime, Utc}; +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +use crate::context::Context; +use crate::error::{AttractorError, Result}; + +/// Serializable snapshot of execution state for crash recovery and resume. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Checkpoint { + pub timestamp: DateTime, + pub current_node: String, + pub completed_nodes: Vec, + pub node_retries: HashMap, + pub context_values: HashMap, + pub logs: Vec, +} + +impl Checkpoint { + /// Create a checkpoint from the current execution state. + pub fn from_context( + context: &Context, + current_node: impl Into, + completed_nodes: Vec, + ) -> Self { + Self { + timestamp: Utc::now(), + current_node: current_node.into(), + completed_nodes, + node_retries: HashMap::new(), + context_values: context.snapshot(), + logs: context.logs_snapshot(), + } + } + + /// Save the checkpoint as JSON to a file. + /// + /// # Errors + /// + /// Returns an error if serialization or file writing fails. + pub fn save(&self, path: &Path) -> Result<()> { + let json = serde_json::to_string_pretty(self) + .map_err(|e| AttractorError::Checkpoint(format!("serialize failed: {e}")))?; + std::fs::write(path, json)?; + Ok(()) + } + + /// Load a checkpoint from a JSON file. + /// + /// # Errors + /// + /// Returns an error if the file cannot be read or deserialization fails. + pub fn load(path: &Path) -> Result { + let data = std::fs::read_to_string(path)?; + let checkpoint: Self = serde_json::from_str(&data) + .map_err(|e| AttractorError::Checkpoint(format!("deserialize failed: {e}")))?; + Ok(checkpoint) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn from_context_captures_state() { + let ctx = Context::new(); + ctx.set("key", serde_json::json!("value")); + ctx.append_log("started"); + + let cp = Checkpoint::from_context( + &ctx, + "node_a", + vec!["start".to_string(), "node_a".to_string()], + ); + + assert_eq!(cp.current_node, "node_a"); + assert_eq!(cp.completed_nodes.len(), 2); + assert_eq!(cp.completed_nodes[0], "start"); + assert_eq!(cp.completed_nodes[1], "node_a"); + assert_eq!( + cp.context_values.get("key"), + Some(&serde_json::json!("value")) + ); + assert_eq!(cp.logs.len(), 1); + assert_eq!(cp.logs[0], "started"); + assert!(cp.node_retries.is_empty()); + } + + #[test] + fn save_and_load_roundtrip() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("checkpoint.json"); + + let ctx = Context::new(); + ctx.set("goal", serde_json::json!("test")); + ctx.append_log("log entry"); + + let mut cp = Checkpoint::from_context(&ctx, "work", vec!["start".to_string()]); + cp.node_retries.insert("work".to_string(), 2); + + cp.save(&path).unwrap(); + let loaded = Checkpoint::load(&path).unwrap(); + + assert_eq!(loaded.current_node, "work"); + assert_eq!(loaded.completed_nodes, vec!["start"]); + assert_eq!(loaded.node_retries.get("work"), Some(&2)); + assert_eq!( + loaded.context_values.get("goal"), + Some(&serde_json::json!("test")) + ); + assert_eq!(loaded.logs, vec!["log entry"]); + } + + #[test] + fn load_nonexistent_file() { + let result = Checkpoint::load(Path::new("/nonexistent/checkpoint.json")); + assert!(result.is_err()); + } + + #[test] + fn load_invalid_json() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("bad.json"); + std::fs::write(&path, "not json").unwrap(); + + let result = Checkpoint::load(&path); + assert!(result.is_err()); + } + + #[test] + fn serialization_roundtrip() { + let ctx = Context::new(); + let cp = Checkpoint::from_context(&ctx, "n1", vec![]); + + let json = serde_json::to_string(&cp).unwrap(); + let deserialized: Checkpoint = serde_json::from_str(&json).unwrap(); + assert_eq!(deserialized.current_node, "n1"); + } +} diff --git a/crates/attractor/src/condition.rs b/crates/attractor/src/condition.rs new file mode 100644 index 000000000..7705bddc3 --- /dev/null +++ b/crates/attractor/src/condition.rs @@ -0,0 +1,324 @@ +/// Condition expression evaluator for edge guards (spec Section 10). +/// +/// Grammar: `ConditionExpr ::= Clause ('&&' Clause)*`, `Clause ::= Key Op Literal`, +/// `Op ::= '=' | '!='`. +use crate::context::Context; +use crate::error::AttractorError; +use crate::outcome::Outcome; + +#[derive(Debug, Clone, PartialEq, Eq)] +struct Clause { + key: String, + op: Op, + value: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +enum Op { + Eq, + NotEq, + Truthy, +} + +fn parse_clauses(expr: &str) -> Result, AttractorError> { + let expr = expr.trim(); + if expr.is_empty() { + return Ok(Vec::new()); + } + + expr.split("&&") + .filter(|part| !part.trim().is_empty()) + .map(|part| { + let part = part.trim(); + if let Some(pos) = part.find("!=") { + let key = part[..pos].trim().to_string(); + let value = part[pos + 2..].trim().to_string(); + if key.is_empty() { + return Err(AttractorError::Parse(format!( + "empty key in condition clause: {part:?}" + ))); + } + Ok(Clause { + key, + op: Op::NotEq, + value, + }) + } else if let Some(pos) = part.find('=') { + let key = part[..pos].trim().to_string(); + let value = part[pos + 1..].trim().to_string(); + if key.is_empty() { + return Err(AttractorError::Parse(format!( + "empty key in condition clause: {part:?}" + ))); + } + Ok(Clause { + key, + op: Op::Eq, + value, + }) + } else { + // Bare key: truthiness check + let key = part.to_string(); + if key.is_empty() { + return Err(AttractorError::Parse(format!( + "empty key in condition clause: {part:?}" + ))); + } + Ok(Clause { + key, + op: Op::Truthy, + value: String::new(), + }) + } + }) + .collect() +} + +/// Parse and validate a condition expression. +/// +/// # Errors +/// +/// Returns an error if the expression contains invalid syntax. +pub fn parse_condition(expr: &str) -> Result<(), AttractorError> { + parse_clauses(expr)?; + Ok(()) +} + +fn resolve_key(key: &str, outcome: &Outcome, context: &Context) -> String { + if key == "outcome" { + return outcome.status.to_string(); + } + if key == "preferred_label" { + return outcome + .preferred_label + .as_deref() + .unwrap_or("") + .to_string(); + } + if let Some(path) = key.strip_prefix("context.") { + if let Some(val) = context.get(key) { + return json_value_to_string(&val); + } + if let Some(val) = context.get(path) { + return json_value_to_string(&val); + } + return String::new(); + } + context + .get(key) + .map_or_else(String::new, |val| json_value_to_string(&val)) +} + +fn json_value_to_string(val: &serde_json::Value) -> String { + match val { + serde_json::Value::String(s) => s.clone(), + serde_json::Value::Bool(b) => b.to_string(), + serde_json::Value::Number(n) => n.to_string(), + serde_json::Value::Null => String::new(), + other => other.to_string(), + } +} + +/// Evaluate a condition expression against an outcome and context. +/// Empty conditions always return true. +#[must_use] +pub fn evaluate_condition(expr: &str, outcome: &Outcome, context: &Context) -> bool { + let Ok(clauses) = parse_clauses(expr) else { + return false; + }; + + if clauses.is_empty() { + return true; + } + + clauses.iter().all(|clause| { + let resolved = resolve_key(&clause.key, outcome, context); + match clause.op { + Op::Eq => resolved == clause.value, + Op::NotEq => resolved != clause.value, + Op::Truthy => !resolved.is_empty() && resolved != "false" && resolved != "0", + } + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::outcome::StageStatus; + + fn make_outcome(status: StageStatus) -> Outcome { + Outcome { + status, + preferred_label: None, + suggested_next_ids: Vec::new(), + context_updates: std::collections::HashMap::new(), + notes: None, + failure_reason: None, + } + } + + #[test] + fn empty_condition_is_true() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + assert!(evaluate_condition("", &outcome, &context)); + assert!(evaluate_condition(" ", &outcome, &context)); + } + + #[test] + fn outcome_equals_success() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + assert!(evaluate_condition("outcome=success", &outcome, &context)); + assert!(!evaluate_condition("outcome=fail", &outcome, &context)); + } + + #[test] + fn outcome_not_equals() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + assert!(evaluate_condition("outcome!=fail", &outcome, &context)); + assert!(!evaluate_condition( + "outcome!=success", + &outcome, + &context + )); + } + + #[test] + fn preferred_label_match() { + let mut outcome = make_outcome(StageStatus::Success); + outcome.preferred_label = Some("Fix".to_string()); + let context = Context::new(); + assert!(evaluate_condition( + "preferred_label=Fix", + &outcome, + &context + )); + assert!(!evaluate_condition( + "preferred_label=Approve", + &outcome, + &context + )); + } + + #[test] + fn context_key_with_prefix() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + context.set("tests_passed", serde_json::json!("true")); + assert!(evaluate_condition( + "context.tests_passed=true", + &outcome, + &context + )); + } + + #[test] + fn bare_key_context_lookup() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + context.set("custom_key", serde_json::json!("custom_value")); + assert!(evaluate_condition( + "custom_key=custom_value", + &outcome, + &context + )); + } + + #[test] + fn missing_key_compares_as_empty() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + assert!(!evaluate_condition( + "missing_key=something", + &outcome, + &context + )); + assert!(evaluate_condition("missing_key=", &outcome, &context)); + } + + #[test] + fn multiple_clauses_and() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + context.set("tests_passed", serde_json::json!("true")); + assert!(evaluate_condition( + "outcome=success && context.tests_passed=true", + &outcome, + &context + )); + assert!(!evaluate_condition( + "outcome=fail && context.tests_passed=true", + &outcome, + &context + )); + } + + #[test] + fn parse_condition_validates() { + assert!(parse_condition("outcome=success").is_ok()); + assert!(parse_condition("outcome=success && context.x=y").is_ok()); + assert!(parse_condition("").is_ok()); + } + + #[test] + fn parse_condition_accepts_bare_key() { + assert!(parse_condition("some_flag").is_ok()); + } + + #[test] + fn context_dotted_fallback() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + context.set("loop_state", serde_json::json!("exhausted")); + assert!(evaluate_condition( + "context.loop_state=exhausted", + &outcome, + &context + )); + } + + #[test] + fn bare_key_truthy_when_non_empty() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + context.set("my_flag", serde_json::json!("yes")); + assert!(evaluate_condition("my_flag", &outcome, &context)); + } + + #[test] + fn bare_key_falsy_when_empty() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + assert!(!evaluate_condition("missing_key", &outcome, &context)); + } + + #[test] + fn bare_key_falsy_when_false_string() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + context.set("my_flag", serde_json::json!("false")); + assert!(!evaluate_condition("my_flag", &outcome, &context)); + } + + #[test] + fn bare_key_falsy_when_zero_string() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + context.set("my_flag", serde_json::json!("0")); + assert!(!evaluate_condition("my_flag", &outcome, &context)); + } + + #[test] + fn bare_key_with_and_clause() { + let outcome = make_outcome(StageStatus::Success); + let context = Context::new(); + context.set("flag", serde_json::json!("yes")); + assert!(evaluate_condition( + "outcome=success && flag", + &outcome, + &context + )); + } +} diff --git a/crates/attractor/src/context.rs b/crates/attractor/src/context.rs new file mode 100644 index 000000000..76c6cc309 --- /dev/null +++ b/crates/attractor/src/context.rs @@ -0,0 +1,225 @@ +use std::collections::HashMap; +use std::sync::{Arc, RwLock}; + +use serde_json::Value; + +/// Thread-safe key-value context shared across pipeline stages. +#[derive(Debug, Clone)] +pub struct Context { + values: Arc>>, + logs: Arc>>, +} + +impl Default for Context { + fn default() -> Self { + Self::new() + } +} + +impl Context { + #[must_use] + pub fn new() -> Self { + Self { + values: Arc::new(RwLock::new(HashMap::new())), + logs: Arc::new(RwLock::new(Vec::new())), + } + } + + /// Set a key-value pair in the context. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + pub fn set(&self, key: impl Into, value: Value) { + self.values + .write() + .expect("context lock poisoned") + .insert(key.into(), value); + } + + /// Get a value by key, returning None if not present. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + #[must_use] + pub fn get(&self, key: &str) -> Option { + self.values + .read() + .expect("context lock poisoned") + .get(key) + .cloned() + } + + /// Get a value as a string, returning the default if not present or not a string. + #[must_use] + pub fn get_string(&self, key: &str, default: &str) -> String { + self.get(key) + .and_then(|v| v.as_str().map(String::from)) + .unwrap_or_else(|| default.to_string()) + } + + /// Append a log entry. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + pub fn append_log(&self, entry: impl Into) { + self.logs + .write() + .expect("context lock poisoned") + .push(entry.into()); + } + + /// Return a snapshot (clone) of all current context values. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + #[must_use] + pub fn snapshot(&self) -> HashMap { + self.values + .read() + .expect("context lock poisoned") + .clone() + } + + /// Return a snapshot of the logs. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + #[must_use] + pub fn logs_snapshot(&self) -> Vec { + self.logs.read().expect("context lock poisoned").clone() + } + + /// Deep copy for parallel branch isolation. + #[must_use] + pub fn clone_context(&self) -> Self { + let values = self.snapshot(); + let logs = self.logs_snapshot(); + Self { + values: Arc::new(RwLock::new(values)), + logs: Arc::new(RwLock::new(logs)), + } + } + + /// Merge a map of updates into the context. + /// + /// # Panics + /// + /// Panics if the internal lock is poisoned. + pub fn apply_updates(&self, updates: &HashMap) { + let mut values = self.values.write().expect("context lock poisoned"); + for (key, value) in updates { + values.insert(key.clone(), value.clone()); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn new_context_is_empty() { + let ctx = Context::new(); + assert!(ctx.snapshot().is_empty()); + assert!(ctx.logs_snapshot().is_empty()); + } + + #[test] + fn set_and_get() { + let ctx = Context::new(); + ctx.set("key", serde_json::json!("value")); + assert_eq!(ctx.get("key"), Some(serde_json::json!("value"))); + } + + #[test] + fn get_missing_key() { + let ctx = Context::new(); + assert_eq!(ctx.get("missing"), None); + } + + #[test] + fn get_string_with_value() { + let ctx = Context::new(); + ctx.set("name", serde_json::json!("alice")); + assert_eq!(ctx.get_string("name", "default"), "alice"); + } + + #[test] + fn get_string_missing_key() { + let ctx = Context::new(); + assert_eq!(ctx.get_string("missing", "fallback"), "fallback"); + } + + #[test] + fn get_string_non_string_value() { + let ctx = Context::new(); + ctx.set("num", serde_json::json!(42)); + assert_eq!(ctx.get_string("num", "default"), "default"); + } + + #[test] + fn append_and_snapshot_logs() { + let ctx = Context::new(); + ctx.append_log("first entry"); + ctx.append_log("second entry"); + let logs = ctx.logs_snapshot(); + assert_eq!(logs.len(), 2); + assert_eq!(logs[0], "first entry"); + assert_eq!(logs[1], "second entry"); + } + + #[test] + fn snapshot_is_independent() { + let ctx = Context::new(); + ctx.set("a", serde_json::json!(1)); + let snap = ctx.snapshot(); + ctx.set("b", serde_json::json!(2)); + // snapshot should not contain "b" + assert!(snap.contains_key("a")); + assert!(!snap.contains_key("b")); + } + + #[test] + fn clone_context_is_independent() { + let ctx = Context::new(); + ctx.set("shared", serde_json::json!("original")); + ctx.append_log("log1"); + + let cloned = ctx.clone_context(); + cloned.set("shared", serde_json::json!("modified")); + cloned.append_log("log2"); + + // original should be unchanged + assert_eq!(ctx.get("shared"), Some(serde_json::json!("original"))); + assert_eq!(ctx.logs_snapshot().len(), 1); + + // cloned has the modification + assert_eq!(cloned.get("shared"), Some(serde_json::json!("modified"))); + assert_eq!(cloned.logs_snapshot().len(), 2); + } + + #[test] + fn apply_updates() { + let ctx = Context::new(); + ctx.set("existing", serde_json::json!("old")); + + let mut updates = HashMap::new(); + updates.insert("existing".to_string(), serde_json::json!("new")); + updates.insert("added".to_string(), serde_json::json!(true)); + ctx.apply_updates(&updates); + + assert_eq!(ctx.get("existing"), Some(serde_json::json!("new"))); + assert_eq!(ctx.get("added"), Some(serde_json::json!(true))); + } + + #[test] + fn default_creates_empty_context() { + let ctx = Context::default(); + assert!(ctx.snapshot().is_empty()); + } +} diff --git a/crates/attractor/src/engine.rs b/crates/attractor/src/engine.rs new file mode 100644 index 000000000..672a7e93f --- /dev/null +++ b/crates/attractor/src/engine.rs @@ -0,0 +1,1570 @@ +use std::collections::HashMap; +use std::panic::AssertUnwindSafe; +use std::path::{Path, PathBuf}; +use std::time::Instant; + +use chrono::Utc; +use futures::FutureExt; +use rand::Rng; + +use crate::checkpoint::Checkpoint; +use crate::condition::evaluate_condition; +use crate::context::Context; +use crate::error::{AttractorError, Result}; +use crate::event::{EventEmitter, PipelineEvent}; +use crate::graph::{Edge, Graph, Node}; +use crate::handler::HandlerRegistry; +use crate::outcome::{Outcome, StageStatus}; + +/// Convert a Duration's milliseconds to u64, saturating on overflow. +fn millis_u64(d: std::time::Duration) -> u64 { + u64::try_from(d.as_millis()).unwrap_or(u64::MAX) +} + +// --- Retry policy types --- + +/// Configuration for exponential backoff between retry attempts. +#[derive(Debug, Clone)] +pub struct BackoffConfig { + pub initial_delay_ms: u64, + pub backoff_factor: f64, + pub max_delay_ms: u64, + pub jitter: bool, +} + +impl Default for BackoffConfig { + fn default() -> Self { + Self { + initial_delay_ms: 200, + backoff_factor: 2.0, + max_delay_ms: 60_000, + jitter: true, + } + } +} + +impl BackoffConfig { + /// Calculate delay for a given attempt (1-indexed). + #[must_use] + #[allow(clippy::missing_panics_doc)] + pub fn delay_for_attempt(&self, attempt: u32) -> std::time::Duration { + let exponent = attempt.saturating_sub(1); + let initial = f64::from(u32::try_from(self.initial_delay_ms).unwrap_or(u32::MAX)); + let max = f64::from(u32::try_from(self.max_delay_ms).unwrap_or(u32::MAX)); + let exp_i32 = i32::try_from(exponent).unwrap_or(i32::MAX); + let delay_f64 = initial * self.backoff_factor.powi(exp_i32); + let capped = delay_f64.min(max); + let final_ms = if self.jitter { + let mut rng = rand::thread_rng(); + let jitter_factor: f64 = rng.gen_range(0.5..1.5); + capped * jitter_factor + } else { + capped + }; + // f64 -> u64: clamp to non-negative, truncate via string-free path + let ms = if final_ms <= 0.0 { + 0u64 + } else if final_ms >= f64::from(u32::MAX) { + u64::from(u32::MAX) + } else { + // Safe: final_ms is in [0, u32::MAX] so the truncated integer fits in u64 + #[allow(clippy::cast_sign_loss, clippy::cast_possible_truncation)] + { final_ms as u64 } + }; + std::time::Duration::from_millis(ms) + } +} + +/// Predicate that determines whether an error is retryable. +/// Returns `true` if the error should be retried, `false` to fail immediately. +pub type ShouldRetryFn = std::sync::Arc bool + Send + Sync>; + +/// Retry policy for node execution. +#[derive(Clone)] +pub struct RetryPolicy { + pub max_attempts: u32, + pub backoff: BackoffConfig, + pub should_retry: ShouldRetryFn, +} + +impl std::fmt::Debug for RetryPolicy { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RetryPolicy") + .field("max_attempts", &self.max_attempts) + .field("backoff", &self.backoff) + .field("should_retry", &"") + .finish() + } +} + +/// Default should_retry predicate: retries all errors. +fn default_should_retry() -> ShouldRetryFn { + std::sync::Arc::new(|_| true) +} + +impl RetryPolicy { + /// No retries -- fail immediately. + #[must_use] + pub fn none() -> Self { + Self { + max_attempts: 1, + backoff: BackoffConfig::default(), + should_retry: default_should_retry(), + } + } + + /// Standard retry policy: 5 attempts, 200ms initial, 2x factor. + #[must_use] + pub fn standard() -> Self { + Self { + max_attempts: 5, + backoff: BackoffConfig { + initial_delay_ms: 200, + backoff_factor: 2.0, + max_delay_ms: 60_000, + jitter: true, + }, + should_retry: default_should_retry(), + } + } + + /// Aggressive retry: 5 attempts, 500ms initial, 2x factor. + #[must_use] + pub fn aggressive() -> Self { + Self { + max_attempts: 5, + backoff: BackoffConfig { + initial_delay_ms: 500, + backoff_factor: 2.0, + max_delay_ms: 60_000, + jitter: true, + }, + should_retry: default_should_retry(), + } + } + + /// Linear retry: 3 attempts, 500ms fixed delay. + #[must_use] + pub fn linear() -> Self { + Self { + max_attempts: 3, + backoff: BackoffConfig { + initial_delay_ms: 500, + backoff_factor: 1.0, + max_delay_ms: 60_000, + jitter: true, + }, + should_retry: default_should_retry(), + } + } + + /// Patient retry: 3 attempts, 2000ms initial, 3x factor. + #[must_use] + pub fn patient() -> Self { + Self { + max_attempts: 3, + backoff: BackoffConfig { + initial_delay_ms: 2000, + backoff_factor: 3.0, + max_delay_ms: 60_000, + jitter: true, + }, + should_retry: default_should_retry(), + } + } +} + +/// Build a retry policy from node and graph attributes. +fn build_retry_policy(node: &Node, graph: &Graph) -> RetryPolicy { + let max_retries = node + .max_retries() + .unwrap_or_else(|| graph.default_max_retry()); + // max_retries=0 means 1 attempt (no retries) + let max_attempts = u32::try_from(max_retries + 1).unwrap_or(1).max(1); + RetryPolicy { + max_attempts, + backoff: BackoffConfig::default(), + should_retry: default_should_retry(), + } +} + +// --- Fidelity resolution (spec 5.4) --- + +/// Resolve the context fidelity for a node, following the precedence: +/// 1. Incoming edge `fidelity` attribute +/// 2. Target node `fidelity` attribute +/// 3. Graph `default_fidelity` attribute +/// 4. Default: "compact" +#[must_use] +pub fn resolve_fidelity(incoming_edge: Option<&Edge>, node: &Node, graph: &Graph) -> String { + if let Some(edge) = incoming_edge { + if let Some(f) = edge.fidelity() { + return f.to_string(); + } + } + if let Some(f) = node.fidelity() { + return f.to_string(); + } + if let Some(f) = graph.default_fidelity() { + return f.to_string(); + } + "compact".to_string() +} + +// --- Run directory helpers (spec 5.6) --- + +/// Write manifest.json at the start of a pipeline run. +fn write_manifest(logs_root: &Path, graph: &Graph) { + let pipeline_name = if graph.name.is_empty() { + "unnamed" + } else { + &graph.name + }; + let manifest = serde_json::json!({ + "pipeline_name": pipeline_name, + "start_time": Utc::now().to_rfc3339(), + "node_count": graph.nodes.len(), + "edge_count": graph.edges.len(), + }); + if let Ok(json) = serde_json::to_string_pretty(&manifest) { + let _ = std::fs::create_dir_all(logs_root); + let _ = std::fs::write(logs_root.join("manifest.json"), json); + } +} + +/// Write status.json for a completed node into {logs_root}/{node_id}/status.json. +fn write_node_status(logs_root: &Path, node_id: &str, outcome: &Outcome) { + let node_dir = logs_root.join(node_id); + let _ = std::fs::create_dir_all(&node_dir); + let status = serde_json::json!({ + "status": outcome.status.to_string(), + "notes": outcome.notes, + "failure_reason": outcome.failure_reason, + "timestamp": Utc::now().to_rfc3339(), + }); + if let Ok(json) = serde_json::to_string_pretty(&status) { + let _ = std::fs::write(node_dir.join("status.json"), json); + } +} + +// --- Edge selection --- + +/// Normalize a label for comparison: lowercase, trim, strip accelerator prefixes. +/// Patterns: "[Y] ", "Y) ", "Y - " +fn normalize_label(label: &str) -> String { + let s = label.trim().to_lowercase(); + // Strip "[X] " prefix + if s.starts_with('[') { + if let Some(rest) = s.strip_prefix('[').and_then(|s| { + s.find(']') + .map(|i| s[i + 1..].trim_start().to_string()) + }) { + return rest; + } + } + // Strip "X) " prefix + if s.len() >= 2 { + let bytes = s.as_bytes(); + if bytes.get(1) == Some(&b')') { + return s[2..].trim_start().to_string(); + } + } + // Strip "X - " prefix + if s.len() >= 3 { + if let Some(rest) = s.get(1..).and_then(|r| r.strip_prefix(" - ")) { + return rest.to_string(); + } + } + s +} + +/// Pick the best edge by highest weight, then lexical target node ID tiebreak. +fn best_by_weight_then_lexical<'a>(edges: &[&'a Edge]) -> Option<&'a Edge> { + if edges.is_empty() { + return None; + } + let mut best = edges[0]; + for &edge in &edges[1..] { + if edge.weight() > best.weight() + || (edge.weight() == best.weight() && edge.to < best.to) + { + best = edge; + } + } + Some(best) +} + +/// Select the next edge from a node's outgoing edges (spec Section 3.3). +#[must_use] +pub fn select_edge<'a>( + node_id: &str, + outcome: &Outcome, + context: &Context, + graph: &'a Graph, +) -> Option<&'a Edge> { + let edges = graph.outgoing_edges(node_id); + if edges.is_empty() { + return None; + } + + // Step 1: Condition matching + let condition_matched: Vec<&Edge> = edges + .iter() + .filter(|e| { + e.condition() + .is_some_and(|c| !c.is_empty() && evaluate_condition(c, outcome, context)) + }) + .copied() + .collect(); + if !condition_matched.is_empty() { + return best_by_weight_then_lexical(&condition_matched); + } + + // Step 2: Preferred label match + if let Some(pref) = &outcome.preferred_label { + let normalized_pref = normalize_label(pref); + for edge in &edges { + if let Some(label) = edge.label() { + if normalize_label(label) == normalized_pref { + return Some(edge); + } + } + } + } + + // Step 3: Suggested next IDs + for suggested_id in &outcome.suggested_next_ids { + for edge in &edges { + if edge.to == *suggested_id { + return Some(edge); + } + } + } + + // Step 4 & 5: Weight with lexical tiebreak (unconditional edges only) + let unconditional: Vec<&Edge> = edges + .iter() + .filter(|e| e.condition().is_none_or(str::is_empty)) + .copied() + .collect(); + if !unconditional.is_empty() { + return best_by_weight_then_lexical(&unconditional); + } + + // Fallback: any edge + best_by_weight_then_lexical(&edges) +} + +// --- Goal gate enforcement --- + +/// Check if all goal gates have been satisfied. +/// Returns Ok(()) if all gates passed, or Err with the failed node ID. +fn check_goal_gates( + graph: &Graph, + node_outcomes: &HashMap, +) -> std::result::Result<(), String> { + for (node_id, outcome) in node_outcomes { + if let Some(node) = graph.nodes.get(node_id) { + if node.goal_gate() + && outcome.status != StageStatus::Success + && outcome.status != StageStatus::PartialSuccess + { + return Err(node_id.clone()); + } + } + } + Ok(()) +} + +/// Resolve the retry target for a failed goal gate node. +fn get_retry_target(failed_node_id: &str, graph: &Graph) -> Option { + if let Some(node) = graph.nodes.get(failed_node_id) { + // Node-level retry_target + if let Some(target) = node.retry_target() { + if graph.nodes.contains_key(target) { + return Some(target.to_string()); + } + } + // Node-level fallback_retry_target + if let Some(target) = node.fallback_retry_target() { + if graph.nodes.contains_key(target) { + return Some(target.to_string()); + } + } + } + // Graph-level retry_target + if let Some(target) = graph.retry_target() { + if graph.nodes.contains_key(target) { + return Some(target.to_string()); + } + } + // Graph-level fallback_retry_target + if let Some(target) = graph.fallback_retry_target() { + if graph.nodes.contains_key(target) { + return Some(target.to_string()); + } + } + None +} + +/// Check whether a node is a terminal (exit) node. +fn is_terminal(node: &Node) -> bool { + node.shape() == "Msquare" + || node.handler_type() == Some("exit") +} + +// --- Pipeline engine --- + +/// Configuration for a pipeline run. +pub struct RunConfig { + pub logs_root: PathBuf, +} + +/// The pipeline execution engine. +pub struct PipelineEngine { + pub registry: HandlerRegistry, + pub emitter: EventEmitter, +} + +impl PipelineEngine { + #[must_use] + pub const fn new(registry: HandlerRegistry, emitter: EventEmitter) -> Self { + Self { registry, emitter } + } + + /// Mirror graph-level attributes into the context. + fn mirror_graph_attributes(graph: &Graph, context: &Context) { + if !graph.goal().is_empty() { + context.set("graph.goal", serde_json::json!(graph.goal())); + } + for (key, val) in &graph.attrs { + context.set( + format!("graph.{key}"), + serde_json::json!(val.to_string_value()), + ); + } + } + + /// Execute a node handler with retry policy. + async fn execute_with_retry( + &self, + node: &Node, + context: &Context, + graph: &Graph, + logs_root: &Path, + policy: &RetryPolicy, + stage_index: usize, + ) -> Result { + let handler = self.registry.resolve(node); + + for attempt in 1..=policy.max_attempts { + // Gap #11: Panic safety -- catch panics from handler execution + let result = { + let future = handler.execute(node, context, graph, logs_root); + match AssertUnwindSafe(future).catch_unwind().await { + Ok(r) => r, + Err(panic_payload) => { + let msg = if let Some(s) = panic_payload.downcast_ref::<&str>() { + format!("handler panicked: {s}") + } else if let Some(s) = panic_payload.downcast_ref::() { + format!("handler panicked: {s}") + } else { + "handler panicked".to_string() + }; + Err(AttractorError::Handler(msg)) + } + } + }; + + let outcome = match result { + Ok(o) => o, + Err(e) => { + // Gap #7: Check should_retry predicate before retrying + if attempt < policy.max_attempts && (policy.should_retry)(&e) { + let delay = policy.backoff.delay_for_attempt(attempt); + self.emitter.emit(&PipelineEvent::StageFailed { + name: node.label().to_string(), + index: stage_index, + error: e.to_string(), + will_retry: true, + }); + self.emitter.emit(&PipelineEvent::StageRetrying { + name: node.label().to_string(), + index: stage_index, + attempt: usize::try_from(attempt).unwrap_or(usize::MAX), + delay_ms: millis_u64(delay), + }); + tokio::time::sleep(delay).await; + continue; + } + return Ok(Outcome::fail(e.to_string())); + } + }; + + match outcome.status { + StageStatus::Success + | StageStatus::PartialSuccess + | StageStatus::Fail + | StageStatus::Skipped => { + return Ok(outcome); + } + StageStatus::Retry => { + if attempt < policy.max_attempts { + let delay = policy.backoff.delay_for_attempt(attempt); + self.emitter.emit(&PipelineEvent::StageRetrying { + name: node.label().to_string(), + index: stage_index, + attempt: usize::try_from(attempt).unwrap_or(usize::MAX), + delay_ms: millis_u64(delay), + }); + tokio::time::sleep(delay).await; + continue; + } + if node.allow_partial() { + return Ok(Outcome { + status: StageStatus::PartialSuccess, + notes: Some("retries exhausted, partial accepted".to_string()), + ..Outcome::success() + }); + } + return Ok(Outcome::fail("max retries exceeded")); + } + } + } + + Ok(Outcome::fail("max retries exceeded")) + } + + /// Run the pipeline. Returns the final outcome. + /// + /// # Errors + /// + /// Returns an error if no start node is found, a node is missing, or a goal gate fails + /// without a retry target. + pub async fn run(&self, graph: &Graph, config: &RunConfig) -> Result { + self.run_internal(graph, config, None, None).await + } + + /// Resume from a checkpoint. Restores context, completed nodes, and continues + /// execution from the node after the checkpoint's current_node. + /// + /// # Errors + /// + /// Returns an error if the checkpoint's current node is not found or execution fails. + pub async fn run_from_checkpoint( + &self, + graph: &Graph, + config: &RunConfig, + checkpoint: &Checkpoint, + ) -> Result { + self.run_internal(graph, config, Some(checkpoint), None).await + } + + /// Internal run implementation supporting optional checkpoint resume and start_at override. + #[allow(clippy::too_many_lines)] + async fn run_internal( + &self, + graph: &Graph, + config: &RunConfig, + resume_checkpoint: Option<&Checkpoint>, + start_at: Option<&str>, + ) -> Result { + let run_start = Instant::now(); + let run_id = uuid::Uuid::new_v4().to_string(); + + self.emitter.emit(&PipelineEvent::PipelineStarted { + name: graph.name.clone(), + id: run_id, + }); + + // Write manifest.json (spec 5.6) + write_manifest(&config.logs_root, graph); + + // Gap #4: Initialize from checkpoint, start_at, or fresh + let context; + let mut completed_nodes: Vec; + let mut node_outcomes: HashMap = HashMap::new(); + let mut stage_index: usize; + let mut current_node_id: String; + let mut incoming_edge: Option<&Edge> = None; + + if let Some(cp) = resume_checkpoint { + // Restore context from checkpoint + context = Context::new(); + for (key, value) in &cp.context_values { + context.set(key.clone(), value.clone()); + } + for log_entry in &cp.logs { + context.append_log(log_entry.clone()); + } + completed_nodes = cp.completed_nodes.clone(); + stage_index = completed_nodes.len(); + // Resume from the node after the checkpoint's current_node + let edges = graph.outgoing_edges(&cp.current_node); + if let Some(edge) = edges.first() { + current_node_id = edge.to.clone(); + } else { + current_node_id = cp.current_node.clone(); + } + } else if let Some(start) = start_at { + context = Context::new(); + Self::mirror_graph_attributes(graph, &context); + completed_nodes = Vec::new(); + stage_index = 0; + current_node_id = start.to_string(); + } else { + context = Context::new(); + Self::mirror_graph_attributes(graph, &context); + completed_nodes = Vec::new(); + stage_index = 0; + + let start_node = graph + .find_start_node() + .ok_or_else(|| AttractorError::Engine("no start node found".to_string()))?; + current_node_id = start_node.id.clone(); + } + + loop { + let node = graph.nodes.get(¤t_node_id).ok_or_else(|| { + AttractorError::Engine(format!("node not found: {current_node_id}")) + })?; + + // Step 1: Check for terminal node + if is_terminal(node) { + match check_goal_gates(graph, &node_outcomes) { + Ok(()) => break, + Err(failed_node_id) => { + if let Some(retry_target) = + get_retry_target(&failed_node_id, graph) + { + current_node_id = retry_target; + continue; + } + let duration_ms = millis_u64(run_start.elapsed()); + let error_msg = + format!("goal gate unsatisfied for node {failed_node_id} and no retry target"); + self.emitter.emit(&PipelineEvent::PipelineFailed { + error: error_msg.clone(), + duration_ms, + }); + return Err(AttractorError::Engine(error_msg)); + } + } + } + + // Resolve fidelity (spec 5.4) and store in context + let fidelity = resolve_fidelity(incoming_edge, node, graph); + context.set("internal.fidelity", serde_json::json!(&fidelity)); + + // Thread context sharing: store thread association + if let Some(tid) = node.thread_id() { + context.set( + format!("thread.{tid}.current_node"), + serde_json::json!(&node.id), + ); + } + + // Step 2: Execute node handler with retry policy + context.set("current_node", serde_json::json!(&node.id)); + let retry_policy = build_retry_policy(node, graph); + + self.emitter.emit(&PipelineEvent::StageStarted { + name: node.label().to_string(), + index: stage_index, + }); + let stage_start = Instant::now(); + + let outcome = self + .execute_with_retry(node, &context, graph, &config.logs_root, &retry_policy, stage_index) + .await?; + + let stage_duration_ms = millis_u64(stage_start.elapsed()); + + if outcome.status == StageStatus::Fail { + self.emitter.emit(&PipelineEvent::StageFailed { + name: node.label().to_string(), + index: stage_index, + error: outcome + .failure_reason + .as_deref() + .unwrap_or("unknown") + .to_string(), + will_retry: false, + }); + } else { + self.emitter.emit(&PipelineEvent::StageCompleted { + name: node.label().to_string(), + index: stage_index, + duration_ms: stage_duration_ms, + }); + } + + // Write per-node status.json (spec 5.6) + write_node_status(&config.logs_root, &node.id, &outcome); + + // Step 3: Record completion + completed_nodes.push(node.id.clone()); + node_outcomes.insert(node.id.clone(), outcome.clone()); + stage_index += 1; + + // Step 4: Apply context updates from outcome + context.apply_updates(&outcome.context_updates); + context.set("outcome", serde_json::json!(outcome.status.to_string())); + if let Some(ref pref) = outcome.preferred_label { + context.set("preferred_label", serde_json::json!(pref)); + } + + // Step 5: Save checkpoint + let checkpoint = Checkpoint::from_context( + &context, + &node.id, + completed_nodes.clone(), + ); + let checkpoint_path = config.logs_root.join("checkpoint.json"); + if let Err(e) = checkpoint.save(&checkpoint_path) { + context.append_log(format!("checkpoint save failed: {e}")); + } else { + self.emitter.emit(&PipelineEvent::CheckpointSaved { + node_id: node.id.clone(), + }); + } + + // Step 6: Select next edge + let next_edge = select_edge(&node.id, &outcome, &context, graph); + match next_edge { + None => { + // Gap #1: Failure routing -- when FAIL and no matching edge, + // check retry_target / fallback_retry_target before terminating + if outcome.status == StageStatus::Fail { + if let Some(retry_target) = get_retry_target(&node.id, graph) { + current_node_id = retry_target; + continue; + } + let duration_ms = millis_u64(run_start.elapsed()); + let error_msg = format!( + "stage {} failed with no outgoing fail edge", + node.id + ); + self.emitter.emit(&PipelineEvent::PipelineFailed { + error: error_msg.clone(), + duration_ms, + }); + return Err(AttractorError::Engine(error_msg)); + } + break; + } + Some(edge) => { + // Track incoming edge for fidelity resolution on the next node + incoming_edge = Some(edge); + // Gap #6: Handle loop_restart by recursively running from the target + if edge.loop_restart() { + return Box::pin(self.run_internal( + graph, + config, + None, + Some(&edge.to), + )).await; + } + current_node_id.clone_from(&edge.to); + } + } + } + + let duration_ms = millis_u64(run_start.elapsed()); + self.emitter.emit(&PipelineEvent::PipelineCompleted { + duration_ms, + artifact_count: 0, + }); + + // Return last outcome, or success if no outcomes recorded + let last_outcome = node_outcomes + .get(completed_nodes.last().unwrap_or(&String::new())) + .cloned() + .unwrap_or_else(Outcome::success); + Ok(last_outcome) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::AttrValue; + use crate::handler::start::StartHandler; + + // --- BackoffConfig tests --- + + #[test] + fn backoff_no_jitter_first_attempt() { + let config = BackoffConfig { + initial_delay_ms: 200, + backoff_factor: 2.0, + max_delay_ms: 60_000, + jitter: false, + }; + let delay = config.delay_for_attempt(1); + assert_eq!(delay.as_millis(), 200); + } + + #[test] + fn backoff_no_jitter_second_attempt() { + let config = BackoffConfig { + initial_delay_ms: 200, + backoff_factor: 2.0, + max_delay_ms: 60_000, + jitter: false, + }; + let delay = config.delay_for_attempt(2); + assert_eq!(delay.as_millis(), 400); + } + + #[test] + fn backoff_no_jitter_third_attempt() { + let config = BackoffConfig { + initial_delay_ms: 200, + backoff_factor: 2.0, + max_delay_ms: 60_000, + jitter: false, + }; + let delay = config.delay_for_attempt(3); + assert_eq!(delay.as_millis(), 800); + } + + #[test] + fn backoff_respects_max_delay() { + let config = BackoffConfig { + initial_delay_ms: 10_000, + backoff_factor: 10.0, + max_delay_ms: 30_000, + jitter: false, + }; + let delay = config.delay_for_attempt(5); + assert_eq!(delay.as_millis(), 30_000); + } + + #[test] + fn backoff_with_jitter_is_in_range() { + let config = BackoffConfig { + initial_delay_ms: 1000, + backoff_factor: 1.0, + max_delay_ms: 60_000, + jitter: true, + }; + let delay = config.delay_for_attempt(1); + // With jitter factor 0.5..1.5, delay should be 500..1500 + assert!(delay.as_millis() >= 500); + assert!(delay.as_millis() <= 1500); + } + + #[test] + fn backoff_linear_factor() { + let config = BackoffConfig { + initial_delay_ms: 500, + backoff_factor: 1.0, + max_delay_ms: 60_000, + jitter: false, + }; + assert_eq!(config.delay_for_attempt(1).as_millis(), 500); + assert_eq!(config.delay_for_attempt(2).as_millis(), 500); + assert_eq!(config.delay_for_attempt(3).as_millis(), 500); + } + + // --- RetryPolicy preset tests --- + + #[test] + fn retry_policy_none() { + let policy = RetryPolicy::none(); + assert_eq!(policy.max_attempts, 1); + } + + #[test] + fn retry_policy_standard() { + let policy = RetryPolicy::standard(); + assert_eq!(policy.max_attempts, 5); + assert_eq!(policy.backoff.initial_delay_ms, 200); + } + + #[test] + fn retry_policy_aggressive() { + let policy = RetryPolicy::aggressive(); + assert_eq!(policy.max_attempts, 5); + assert_eq!(policy.backoff.initial_delay_ms, 500); + } + + #[test] + fn retry_policy_linear() { + let policy = RetryPolicy::linear(); + assert_eq!(policy.max_attempts, 3); + assert_eq!(policy.backoff.backoff_factor, 1.0); + } + + #[test] + fn retry_policy_patient() { + let policy = RetryPolicy::patient(); + assert_eq!(policy.max_attempts, 3); + assert_eq!(policy.backoff.initial_delay_ms, 2000); + } + + // --- build_retry_policy tests --- + + #[test] + fn build_retry_policy_from_node() { + let mut node = Node::new("n"); + node.attrs + .insert("max_retries".to_string(), AttrValue::Integer(3)); + let graph = Graph::new("test"); + let policy = build_retry_policy(&node, &graph); + assert_eq!(policy.max_attempts, 4); // 3 retries + 1 initial + } + + #[test] + fn build_retry_policy_from_graph_default() { + let node = Node::new("n"); + let mut graph = Graph::new("test"); + graph + .attrs + .insert("default_max_retry".to_string(), AttrValue::Integer(2)); + let policy = build_retry_policy(&node, &graph); + assert_eq!(policy.max_attempts, 3); // 2 retries + 1 initial + } + + #[test] + fn build_retry_policy_no_attrs_uses_graph_default_50() { + let node = Node::new("n"); + let graph = Graph::new("test"); + let policy = build_retry_policy(&node, &graph); + assert_eq!(policy.max_attempts, 51); // default_max_retry=50 + 1 + } + + // --- normalize_label tests --- + + #[test] + fn normalize_label_lowercase_and_trim() { + assert_eq!(normalize_label(" Yes "), "yes"); + } + + #[test] + fn normalize_label_strip_bracket_prefix() { + assert_eq!(normalize_label("[A] Approve"), "approve"); + assert_eq!(normalize_label("[F] Fix"), "fix"); + } + + #[test] + fn normalize_label_strip_paren_prefix() { + assert_eq!(normalize_label("Y) Yes"), "yes"); + } + + #[test] + fn normalize_label_strip_dash_prefix() { + assert_eq!(normalize_label("Y - Yes"), "yes"); + } + + #[test] + fn normalize_label_plain() { + assert_eq!(normalize_label("next"), "next"); + } + + // --- best_by_weight_then_lexical tests --- + + #[test] + fn best_by_weight_highest_wins() { + let e1 = Edge::new("a", "x"); + let mut e2 = Edge::new("a", "y"); + e2.attrs + .insert("weight".to_string(), AttrValue::Integer(5)); + let result = best_by_weight_then_lexical(&[&e1, &e2]).unwrap(); + assert_eq!(result.to, "y"); + } + + #[test] + fn best_by_weight_lexical_tiebreak() { + let e1 = Edge::new("a", "beta"); + let e2 = Edge::new("a", "alpha"); + let result = best_by_weight_then_lexical(&[&e1, &e2]).unwrap(); + assert_eq!(result.to, "alpha"); + } + + #[test] + fn best_by_weight_empty_returns_none() { + let result = best_by_weight_then_lexical(&[]); + assert!(result.is_none()); + } + + // --- select_edge tests --- + + fn make_graph_with_edges(edges: Vec) -> Graph { + let mut g = Graph::new("test"); + for edge in &edges { + if !g.nodes.contains_key(&edge.from) { + g.nodes.insert(edge.from.clone(), Node::new(&edge.from)); + } + if !g.nodes.contains_key(&edge.to) { + g.nodes.insert(edge.to.clone(), Node::new(&edge.to)); + } + } + g.edges = edges; + g + } + + #[test] + fn select_edge_no_edges() { + let g = Graph::new("test"); + let outcome = Outcome::success(); + let context = Context::new(); + assert!(select_edge("a", &outcome, &context, &g).is_none()); + } + + #[test] + fn select_edge_single_unconditional() { + let g = make_graph_with_edges(vec![Edge::new("a", "b")]); + let outcome = Outcome::success(); + let context = Context::new(); + let edge = select_edge("a", &outcome, &context, &g).unwrap(); + assert_eq!(edge.to, "b"); + } + + #[test] + fn select_edge_condition_match() { + let mut e1 = Edge::new("a", "fail_path"); + e1.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=fail".to_string()), + ); + let mut e2 = Edge::new("a", "success_path"); + e2.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + let g = make_graph_with_edges(vec![e1, e2]); + let outcome = Outcome::success(); + let context = Context::new(); + let edge = select_edge("a", &outcome, &context, &g).unwrap(); + assert_eq!(edge.to, "success_path"); + } + + #[test] + fn select_edge_preferred_label() { + let mut e1 = Edge::new("a", "approve"); + e1.attrs.insert( + "label".to_string(), + AttrValue::String("[A] Approve".to_string()), + ); + let mut e2 = Edge::new("a", "fix"); + e2.attrs.insert( + "label".to_string(), + AttrValue::String("[F] Fix".to_string()), + ); + let g = make_graph_with_edges(vec![e1, e2]); + let mut outcome = Outcome::success(); + outcome.preferred_label = Some("Fix".to_string()); + let context = Context::new(); + let edge = select_edge("a", &outcome, &context, &g).unwrap(); + assert_eq!(edge.to, "fix"); + } + + #[test] + fn select_edge_suggested_next_ids() { + let e1 = Edge::new("a", "path1"); + let e2 = Edge::new("a", "path2"); + let g = make_graph_with_edges(vec![e1, e2]); + let mut outcome = Outcome::success(); + outcome.suggested_next_ids = vec!["path2".to_string()]; + let context = Context::new(); + let edge = select_edge("a", &outcome, &context, &g).unwrap(); + assert_eq!(edge.to, "path2"); + } + + #[test] + fn select_edge_weight_tiebreak() { + let mut e1 = Edge::new("a", "low"); + e1.attrs + .insert("weight".to_string(), AttrValue::Integer(1)); + let mut e2 = Edge::new("a", "high"); + e2.attrs + .insert("weight".to_string(), AttrValue::Integer(10)); + let g = make_graph_with_edges(vec![e1, e2]); + let outcome = Outcome::success(); + let context = Context::new(); + let edge = select_edge("a", &outcome, &context, &g).unwrap(); + assert_eq!(edge.to, "high"); + } + + #[test] + fn select_edge_lexical_tiebreak() { + let e1 = Edge::new("a", "charlie"); + let e2 = Edge::new("a", "alpha"); + let g = make_graph_with_edges(vec![e1, e2]); + let outcome = Outcome::success(); + let context = Context::new(); + let edge = select_edge("a", &outcome, &context, &g).unwrap(); + assert_eq!(edge.to, "alpha"); + } + + #[test] + fn select_edge_condition_beats_unconditional() { + let mut e_cond = Edge::new("a", "cond_path"); + e_cond.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + let e_uncond = Edge::new("a", "uncond_path"); + let g = make_graph_with_edges(vec![e_cond, e_uncond]); + let outcome = Outcome::success(); + let context = Context::new(); + let edge = select_edge("a", &outcome, &context, &g).unwrap(); + assert_eq!(edge.to, "cond_path"); + } + + // --- check_goal_gates tests --- + + #[test] + fn goal_gates_all_satisfied() { + let mut g = Graph::new("test"); + let mut n = Node::new("work"); + n.attrs + .insert("goal_gate".to_string(), AttrValue::Boolean(true)); + g.nodes.insert("work".to_string(), n); + + let mut outcomes = HashMap::new(); + outcomes.insert("work".to_string(), Outcome::success()); + + assert!(check_goal_gates(&g, &outcomes).is_ok()); + } + + #[test] + fn goal_gates_partial_success_counts() { + let mut g = Graph::new("test"); + let mut n = Node::new("work"); + n.attrs + .insert("goal_gate".to_string(), AttrValue::Boolean(true)); + g.nodes.insert("work".to_string(), n); + + let mut outcomes = HashMap::new(); + let mut o = Outcome::success(); + o.status = StageStatus::PartialSuccess; + outcomes.insert("work".to_string(), o); + + assert!(check_goal_gates(&g, &outcomes).is_ok()); + } + + #[test] + fn goal_gates_failed_returns_node_id() { + let mut g = Graph::new("test"); + let mut n = Node::new("work"); + n.attrs + .insert("goal_gate".to_string(), AttrValue::Boolean(true)); + g.nodes.insert("work".to_string(), n); + + let mut outcomes = HashMap::new(); + outcomes.insert("work".to_string(), Outcome::fail("test")); + + assert_eq!(check_goal_gates(&g, &outcomes), Err("work".to_string())); + } + + #[test] + fn goal_gates_non_gate_nodes_ignored() { + let mut g = Graph::new("test"); + g.nodes.insert("work".to_string(), Node::new("work")); + + let mut outcomes = HashMap::new(); + outcomes.insert("work".to_string(), Outcome::fail("test")); + + assert!(check_goal_gates(&g, &outcomes).is_ok()); + } + + // --- get_retry_target tests --- + + #[test] + fn retry_target_from_node() { + let mut g = Graph::new("test"); + let mut n = Node::new("work"); + n.attrs.insert( + "retry_target".to_string(), + AttrValue::String("plan".to_string()), + ); + g.nodes.insert("work".to_string(), n); + g.nodes.insert("plan".to_string(), Node::new("plan")); + + assert_eq!( + get_retry_target("work", &g), + Some("plan".to_string()) + ); + } + + #[test] + fn retry_target_from_fallback() { + let mut g = Graph::new("test"); + let mut n = Node::new("work"); + n.attrs.insert( + "fallback_retry_target".to_string(), + AttrValue::String("plan".to_string()), + ); + g.nodes.insert("work".to_string(), n); + g.nodes.insert("plan".to_string(), Node::new("plan")); + + assert_eq!( + get_retry_target("work", &g), + Some("plan".to_string()) + ); + } + + #[test] + fn retry_target_from_graph() { + let mut g = Graph::new("test"); + g.nodes.insert("work".to_string(), Node::new("work")); + g.nodes.insert("plan".to_string(), Node::new("plan")); + g.attrs.insert( + "retry_target".to_string(), + AttrValue::String("plan".to_string()), + ); + + assert_eq!( + get_retry_target("work", &g), + Some("plan".to_string()) + ); + } + + #[test] + fn retry_target_none_when_missing() { + let mut g = Graph::new("test"); + g.nodes.insert("work".to_string(), Node::new("work")); + assert!(get_retry_target("work", &g).is_none()); + } + + #[test] + fn retry_target_skips_nonexistent_node() { + let mut g = Graph::new("test"); + let mut n = Node::new("work"); + n.attrs.insert( + "retry_target".to_string(), + AttrValue::String("nonexistent".to_string()), + ); + g.nodes.insert("work".to_string(), n); + // No "nonexistent" node -- should fall through to graph-level + assert!(get_retry_target("work", &g).is_none()); + } + + // --- is_terminal tests --- + + #[test] + fn terminal_by_shape() { + let mut n = Node::new("exit"); + n.attrs.insert( + "shape".to_string(), + AttrValue::String("Msquare".to_string()), + ); + assert!(is_terminal(&n)); + } + + #[test] + fn terminal_by_type() { + let mut n = Node::new("end"); + n.attrs.insert( + "type".to_string(), + AttrValue::String("exit".to_string()), + ); + assert!(is_terminal(&n)); + } + + #[test] + fn non_terminal_node() { + let n = Node::new("work"); + assert!(!is_terminal(&n)); + } + + // --- PipelineEngine integration tests --- + + fn simple_graph() -> Graph { + let mut g = Graph::new("test_pipeline"); + g.attrs.insert( + "goal".to_string(), + AttrValue::String("Run tests".to_string()), + ); + + let mut start = Node::new("start"); + start.attrs.insert( + "shape".to_string(), + AttrValue::String("Mdiamond".to_string()), + ); + g.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs.insert( + "shape".to_string(), + AttrValue::String("Msquare".to_string()), + ); + g.nodes.insert("exit".to_string(), exit); + + g.edges.push(Edge::new("start", "exit")); + g + } + + fn make_registry() -> HandlerRegistry { + use crate::handler::exit::ExitHandler; + let mut registry = HandlerRegistry::new(Box::new(StartHandler)); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry + } + + #[tokio::test] + async fn engine_runs_simple_pipeline() { + let dir = tempfile::tempdir().unwrap(); + let g = simple_graph(); + let engine = PipelineEngine::new(make_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&g, &config).await.unwrap(); + assert_eq!(outcome.status, StageStatus::Success); + } + + #[tokio::test] + async fn engine_saves_checkpoint() { + let dir = tempfile::tempdir().unwrap(); + let g = simple_graph(); + let engine = PipelineEngine::new(make_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&g, &config).await.unwrap(); + let checkpoint_path = dir.path().join("checkpoint.json"); + assert!(checkpoint_path.exists()); + } + + #[tokio::test] + async fn engine_emits_events() { + let dir = tempfile::tempdir().unwrap(); + let g = simple_graph(); + + let events = std::sync::Arc::new(std::sync::Mutex::new(Vec::new())); + let events_clone = events.clone(); + let mut emitter = EventEmitter::new(); + emitter.on_event(move |event| { + events_clone.lock().unwrap().push(format!("{event:?}")); + }); + + let engine = PipelineEngine::new(make_registry(), emitter); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&g, &config).await.unwrap(); + + let collected = events.lock().unwrap(); + // Should have: PipelineStarted, StageStarted (start), StageCompleted (start), + // CheckpointSaved, PipelineCompleted + assert!(collected.len() >= 4); + } + + #[tokio::test] + async fn engine_error_when_no_start_node() { + let dir = tempfile::tempdir().unwrap(); + let g = Graph::new("empty"); + let engine = PipelineEngine::new(make_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let result = engine.run(&g, &config).await; + assert!(result.is_err()); + } + + #[tokio::test] + async fn engine_mirrors_graph_goal_to_context() { + let dir = tempfile::tempdir().unwrap(); + let g = simple_graph(); + let engine = PipelineEngine::new(make_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&g, &config).await.unwrap(); + + // Verify checkpoint has graph.goal mirrored + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert_eq!( + cp.context_values.get("graph.goal"), + Some(&serde_json::json!("Run tests")) + ); + } + + #[tokio::test] + async fn engine_multi_node_pipeline() { + let dir = tempfile::tempdir().unwrap(); + let mut g = simple_graph(); + // Insert a work node between start and exit + let work = Node::new("work"); + g.nodes.insert("work".to_string(), work); + g.edges.clear(); + g.edges.push(Edge::new("start", "work")); + g.edges.push(Edge::new("work", "exit")); + + let engine = PipelineEngine::new(make_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + let outcome = engine.run(&g, &config).await.unwrap(); + assert_eq!(outcome.status, StageStatus::Success); + + // Checkpoint should show work was completed + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!(cp.completed_nodes.contains(&"start".to_string())); + assert!(cp.completed_nodes.contains(&"work".to_string())); + } + + #[tokio::test] + async fn engine_conditional_routing() { + let dir = tempfile::tempdir().unwrap(); + let mut g = Graph::new("cond_test"); + + let mut start = Node::new("start"); + start.attrs.insert( + "shape".to_string(), + AttrValue::String("Mdiamond".to_string()), + ); + g.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs.insert( + "shape".to_string(), + AttrValue::String("Msquare".to_string()), + ); + g.nodes.insert("exit".to_string(), exit); + + g.nodes + .insert("path_a".to_string(), Node::new("path_a")); + g.nodes + .insert("path_b".to_string(), Node::new("path_b")); + + // start -> path_a (condition: outcome=fail) + let mut e1 = Edge::new("start", "path_a"); + e1.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=fail".to_string()), + ); + g.edges.push(e1); + + // start -> path_b (unconditional, should be taken since start returns success) + g.edges.push(Edge::new("start", "path_b")); + + g.edges.push(Edge::new("path_a", "exit")); + g.edges.push(Edge::new("path_b", "exit")); + + let engine = PipelineEngine::new(make_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&g, &config).await.unwrap(); + + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + // Should have gone through path_b (unconditional) not path_a (condition=fail) + assert!(cp.completed_nodes.contains(&"path_b".to_string())); + assert!(!cp.completed_nodes.contains(&"path_a".to_string())); + } + + // --- resolve_fidelity tests --- + + #[test] + fn fidelity_defaults_to_compact() { + let node = Node::new("work"); + let graph = Graph::new("test"); + assert_eq!(resolve_fidelity(None, &node, &graph), "compact"); + } + + #[test] + fn fidelity_from_graph_default() { + let node = Node::new("work"); + let mut graph = Graph::new("test"); + graph.attrs.insert( + "default_fidelity".to_string(), + AttrValue::String("truncate".to_string()), + ); + assert_eq!(resolve_fidelity(None, &node, &graph), "truncate"); + } + + #[test] + fn fidelity_from_node_overrides_graph() { + let mut node = Node::new("work"); + node.attrs.insert( + "fidelity".to_string(), + AttrValue::String("full".to_string()), + ); + let mut graph = Graph::new("test"); + graph.attrs.insert( + "default_fidelity".to_string(), + AttrValue::String("truncate".to_string()), + ); + assert_eq!(resolve_fidelity(None, &node, &graph), "full"); + } + + #[test] + fn fidelity_from_edge_overrides_node() { + let mut node = Node::new("work"); + node.attrs.insert( + "fidelity".to_string(), + AttrValue::String("full".to_string()), + ); + let mut edge = Edge::new("a", "work"); + edge.attrs.insert( + "fidelity".to_string(), + AttrValue::String("summary:high".to_string()), + ); + let graph = Graph::new("test"); + assert_eq!(resolve_fidelity(Some(&edge), &node, &graph), "summary:high"); + } + + // --- manifest.json and node status tests --- + + #[tokio::test] + async fn engine_writes_manifest_json() { + let dir = tempfile::tempdir().unwrap(); + let g = simple_graph(); + let engine = PipelineEngine::new(make_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&g, &config).await.unwrap(); + + let manifest_path = dir.path().join("manifest.json"); + assert!(manifest_path.exists()); + let manifest: serde_json::Value = + serde_json::from_str(&std::fs::read_to_string(&manifest_path).unwrap()).unwrap(); + assert_eq!(manifest["pipeline_name"], "test_pipeline"); + assert!(manifest["start_time"].is_string()); + assert!(manifest["node_count"].is_number()); + assert!(manifest["edge_count"].is_number()); + } + + #[tokio::test] + async fn engine_writes_node_status_json() { + let dir = tempfile::tempdir().unwrap(); + let g = simple_graph(); + let engine = PipelineEngine::new(make_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&g, &config).await.unwrap(); + + // start node should have status.json + let status_path = dir.path().join("start").join("status.json"); + assert!(status_path.exists()); + let status: serde_json::Value = + serde_json::from_str(&std::fs::read_to_string(&status_path).unwrap()).unwrap(); + assert_eq!(status["status"], "success"); + } + + #[tokio::test] + async fn engine_stores_fidelity_in_context() { + let dir = tempfile::tempdir().unwrap(); + let g = simple_graph(); + let engine = PipelineEngine::new(make_registry(), EventEmitter::new()); + let config = RunConfig { + logs_root: dir.path().to_path_buf(), + }; + engine.run(&g, &config).await.unwrap(); + + // The checkpoint context should contain internal.fidelity + let cp = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert_eq!( + cp.context_values.get("internal.fidelity"), + Some(&serde_json::json!("compact")) + ); + } +} diff --git a/crates/attractor/src/error.rs b/crates/attractor/src/error.rs new file mode 100644 index 000000000..b19e835d3 --- /dev/null +++ b/crates/attractor/src/error.rs @@ -0,0 +1,91 @@ +use thiserror::Error; + +#[derive(Error, Debug, Clone)] +pub enum AttractorError { + #[error("Parse error: {0}")] + Parse(String), + + #[error("Validation error: {0}")] + Validation(String), + + #[error("Engine error: {0}")] + Engine(String), + + #[error("Handler error: {0}")] + Handler(String), + + #[error("Checkpoint error: {0}")] + Checkpoint(String), + + #[error("Stylesheet error: {0}")] + Stylesheet(String), + + #[error("I/O error: {0}")] + Io(String), +} + +impl From for AttractorError { + fn from(err: std::io::Error) -> Self { + Self::Io(err.to_string()) + } +} + +pub type Result = std::result::Result; + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parse_error_display() { + let err = AttractorError::Parse("unexpected token".to_string()); + assert_eq!(err.to_string(), "Parse error: unexpected token"); + } + + #[test] + fn validation_error_display() { + let err = AttractorError::Validation("missing start node".to_string()); + assert_eq!(err.to_string(), "Validation error: missing start node"); + } + + #[test] + fn engine_error_display() { + let err = AttractorError::Engine("no outgoing edge".to_string()); + assert_eq!(err.to_string(), "Engine error: no outgoing edge"); + } + + #[test] + fn handler_error_display() { + let err = AttractorError::Handler("LLM call failed".to_string()); + assert_eq!(err.to_string(), "Handler error: LLM call failed"); + } + + #[test] + fn checkpoint_error_display() { + let err = AttractorError::Checkpoint("file not found".to_string()); + assert_eq!(err.to_string(), "Checkpoint error: file not found"); + } + + #[test] + fn io_error_display() { + let err = AttractorError::Io("permission denied".to_string()); + assert_eq!(err.to_string(), "I/O error: permission denied"); + } + + #[test] + fn io_error_from_std() { + let io_err = std::io::Error::new(std::io::ErrorKind::NotFound, "not found"); + let err = AttractorError::from(io_err); + assert!(matches!(err, AttractorError::Io(_))); + assert!(err.to_string().contains("not found")); + } + + #[test] + fn result_type_alias_works() { + let ok: Result = Ok(42); + assert_eq!(ok.unwrap(), 42); + + let err: Result = Err(AttractorError::Parse("bad".to_string())); + assert!(err.is_err()); + } +} diff --git a/crates/attractor/src/event.rs b/crates/attractor/src/event.rs new file mode 100644 index 000000000..c0451d979 --- /dev/null +++ b/crates/attractor/src/event.rs @@ -0,0 +1,165 @@ +use serde::{Deserialize, Serialize}; + +/// Events emitted during pipeline execution for observability. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub enum PipelineEvent { + PipelineStarted { + name: String, + id: String, + }, + PipelineCompleted { + duration_ms: u64, + artifact_count: usize, + }, + PipelineFailed { + error: String, + duration_ms: u64, + }, + StageStarted { + name: String, + index: usize, + }, + StageCompleted { + name: String, + index: usize, + duration_ms: u64, + }, + StageFailed { + name: String, + index: usize, + error: String, + will_retry: bool, + }, + StageRetrying { + name: String, + index: usize, + attempt: usize, + delay_ms: u64, + }, + ParallelStarted { + branch_count: usize, + }, + ParallelBranchStarted { + branch: String, + index: usize, + }, + ParallelBranchCompleted { + branch: String, + index: usize, + duration_ms: u64, + success: bool, + }, + ParallelCompleted { + duration_ms: u64, + success_count: usize, + failure_count: usize, + }, + InterviewStarted { + question: String, + stage: String, + }, + InterviewCompleted { + question: String, + answer: String, + duration_ms: u64, + }, + InterviewTimeout { + question: String, + stage: String, + duration_ms: u64, + }, + CheckpointSaved { + node_id: String, + }, +} + +/// Listener callback type for pipeline events. +type EventListener = Box; + +/// Callback-based event emitter for pipeline events. +pub struct EventEmitter { + listeners: Vec, +} + +impl std::fmt::Debug for EventEmitter { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("EventEmitter") + .field("listener_count", &self.listeners.len()) + .finish() + } +} + +impl Default for EventEmitter { + fn default() -> Self { + Self::new() + } +} + +impl EventEmitter { + #[must_use] + pub fn new() -> Self { + Self { + listeners: Vec::new(), + } + } + + pub fn on_event(&mut self, listener: impl Fn(&PipelineEvent) + Send + Sync + 'static) { + self.listeners.push(Box::new(listener)); + } + + pub fn emit(&self, event: &PipelineEvent) { + for listener in &self.listeners { + listener(event); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::{Arc, Mutex}; + + #[test] + fn event_emitter_new_has_no_listeners() { + let emitter = EventEmitter::new(); + assert_eq!(emitter.listeners.len(), 0); + } + + #[test] + fn event_emitter_calls_listener() { + let mut emitter = EventEmitter::new(); + let received = Arc::new(Mutex::new(Vec::new())); + let received_clone = Arc::clone(&received); + emitter.on_event(move |event| { + let name = match event { + PipelineEvent::PipelineStarted { name, .. } => name.clone(), + _ => "other".to_string(), + }; + received_clone.lock().unwrap().push(name); + }); + emitter.emit(&PipelineEvent::PipelineStarted { + name: "test".to_string(), + id: "1".to_string(), + }); + let events = received.lock().unwrap(); + assert_eq!(events.len(), 1); + assert_eq!(events[0], "test"); + } + + #[test] + fn pipeline_event_serialization() { + let event = PipelineEvent::StageStarted { + name: "plan".to_string(), + index: 0, + }; + let json = serde_json::to_string(&event).unwrap(); + assert!(json.contains("StageStarted")); + assert!(json.contains("plan")); + } + + #[test] + fn event_emitter_default() { + let emitter = EventEmitter::default(); + assert_eq!(emitter.listeners.len(), 0); + } +} diff --git a/crates/attractor/src/graph/mod.rs b/crates/attractor/src/graph/mod.rs new file mode 100644 index 000000000..26535e727 --- /dev/null +++ b/crates/attractor/src/graph/mod.rs @@ -0,0 +1,3 @@ +pub mod types; + +pub use types::*; diff --git a/crates/attractor/src/graph/types.rs b/crates/attractor/src/graph/types.rs new file mode 100644 index 000000000..f6c255067 --- /dev/null +++ b/crates/attractor/src/graph/types.rs @@ -0,0 +1,620 @@ +use std::collections::HashMap; +use std::time::Duration; + +use serde::{Deserialize, Serialize}; + +/// Typed attribute values for nodes, edges, and graph-level attributes. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub enum AttrValue { + String(String), + Integer(i64), + Float(f64), + Boolean(bool), + Duration(Duration), +} + +impl AttrValue { + #[must_use] + pub fn as_str(&self) -> Option<&str> { + match self { + Self::String(s) => Some(s), + _ => None, + } + } + + #[must_use] + pub const fn as_i64(&self) -> Option { + match self { + Self::Integer(n) => Some(*n), + _ => None, + } + } + + #[must_use] + pub const fn as_f64(&self) -> Option { + match self { + Self::Float(n) => Some(*n), + _ => None, + } + } + + #[must_use] + pub const fn as_bool(&self) -> Option { + match self { + Self::Boolean(b) => Some(*b), + _ => None, + } + } + + #[must_use] + pub const fn as_duration(&self) -> Option { + match self { + Self::Duration(d) => Some(*d), + _ => None, + } + } + + /// Convert any variant to its string representation. + #[must_use] + pub fn to_string_value(&self) -> String { + match self { + Self::String(s) => s.clone(), + Self::Integer(n) => n.to_string(), + Self::Float(f) => f.to_string(), + Self::Boolean(b) => b.to_string(), + Self::Duration(d) => format!("{}ms", d.as_millis()), + } + } +} + +/// Maps Graphviz shapes to handler type strings (Section 2.8). +#[must_use] +pub fn shape_to_handler_type(shape: &str) -> Option<&'static str> { + match shape { + "Mdiamond" => Some("start"), + "Msquare" => Some("exit"), + "box" => Some("codergen"), + "hexagon" => Some("wait.human"), + "diamond" => Some("conditional"), + "component" => Some("parallel"), + "tripleoctagon" => Some("parallel.fan_in"), + "parallelogram" => Some("tool"), + "house" => Some("stack.manager_loop"), + _ => None, + } +} + +/// A node in the workflow graph. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Node { + pub id: String, + pub attrs: HashMap, + /// CSS-like classes for model stylesheet targeting (from `class` attr and subgraph derivation). + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub classes: Vec, +} + +impl Node { + pub fn new(id: impl Into) -> Self { + Self { + id: id.into(), + attrs: HashMap::new(), + classes: Vec::new(), + } + } + + fn str_attr(&self, key: &str) -> Option<&str> { + self.attrs.get(key).and_then(AttrValue::as_str) + } + + fn bool_attr(&self, key: &str) -> Option { + self.attrs.get(key).and_then(AttrValue::as_bool) + } + + fn int_attr(&self, key: &str) -> Option { + self.attrs.get(key).and_then(AttrValue::as_i64) + } + + #[must_use] + pub fn label(&self) -> &str { + self.str_attr("label").unwrap_or(&self.id) + } + + #[must_use] + pub fn shape(&self) -> &str { + self.str_attr("shape").unwrap_or("box") + } + + #[must_use] + pub fn node_type(&self) -> Option<&str> { + self.str_attr("type") + } + + #[must_use] + pub fn prompt(&self) -> Option<&str> { + self.str_attr("prompt") + } + + #[must_use] + pub fn max_retries(&self) -> Option { + self.int_attr("max_retries") + } + + #[must_use] + pub fn goal_gate(&self) -> bool { + self.bool_attr("goal_gate").unwrap_or(false) + } + + #[must_use] + pub fn retry_target(&self) -> Option<&str> { + self.str_attr("retry_target") + } + + #[must_use] + pub fn fallback_retry_target(&self) -> Option<&str> { + self.str_attr("fallback_retry_target") + } + + #[must_use] + pub fn fidelity(&self) -> Option<&str> { + self.str_attr("fidelity") + } + + #[must_use] + pub fn thread_id(&self) -> Option<&str> { + self.str_attr("thread_id") + } + + #[must_use] + pub fn class(&self) -> Option<&str> { + self.str_attr("class") + } + + pub fn timeout(&self) -> Option { + self.attrs.get("timeout").and_then(AttrValue::as_duration) + } + + #[must_use] + pub fn llm_model(&self) -> Option<&str> { + self.str_attr("llm_model") + } + + #[must_use] + pub fn llm_provider(&self) -> Option<&str> { + self.str_attr("llm_provider") + } + + #[must_use] + pub fn reasoning_effort(&self) -> &str { + self.str_attr("reasoning_effort").unwrap_or("high") + } + + #[must_use] + pub fn auto_status(&self) -> bool { + self.bool_attr("auto_status").unwrap_or(false) + } + + #[must_use] + pub fn allow_partial(&self) -> bool { + self.bool_attr("allow_partial").unwrap_or(false) + } + + /// Resolve the handler type for this node using explicit type or shape mapping. + #[must_use] + pub fn handler_type(&self) -> Option<&str> { + if let Some(t) = self.node_type() { + return Some(t); + } + shape_to_handler_type(self.shape()) + } +} + +/// An edge connecting two nodes in the workflow graph. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Edge { + pub from: String, + pub to: String, + pub attrs: HashMap, +} + +impl Edge { + pub fn new(from: impl Into, to: impl Into) -> Self { + Self { + from: from.into(), + to: to.into(), + attrs: HashMap::new(), + } + } + + fn str_attr(&self, key: &str) -> Option<&str> { + self.attrs.get(key).and_then(AttrValue::as_str) + } + + fn bool_attr(&self, key: &str) -> Option { + self.attrs.get(key).and_then(AttrValue::as_bool) + } + + fn int_attr(&self, key: &str) -> Option { + self.attrs.get(key).and_then(AttrValue::as_i64) + } + + #[must_use] + pub fn label(&self) -> Option<&str> { + self.str_attr("label") + } + + #[must_use] + pub fn condition(&self) -> Option<&str> { + self.str_attr("condition") + } + + #[must_use] + pub fn weight(&self) -> i64 { + self.int_attr("weight").unwrap_or(0) + } + + #[must_use] + pub fn fidelity(&self) -> Option<&str> { + self.str_attr("fidelity") + } + + #[must_use] + pub fn thread_id(&self) -> Option<&str> { + self.str_attr("thread_id") + } + + #[must_use] + pub fn loop_restart(&self) -> bool { + self.bool_attr("loop_restart").unwrap_or(false) + } + + #[must_use] + pub fn freeform(&self) -> bool { + self.bool_attr("freeform").unwrap_or(false) + } +} + +/// The parsed workflow graph containing nodes, edges, and graph-level attributes. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Graph { + pub name: String, + pub nodes: HashMap, + pub edges: Vec, + pub attrs: HashMap, +} + +impl Graph { + pub fn new(name: impl Into) -> Self { + Self { + name: name.into(), + nodes: HashMap::new(), + edges: Vec::new(), + attrs: HashMap::new(), + } + } + + /// Returns all outgoing edges from the given node. + #[must_use] + pub fn outgoing_edges(&self, node_id: &str) -> Vec<&Edge> { + self.edges.iter().filter(|e| e.from == node_id).collect() + } + + /// Returns all incoming edges to the given node. + #[must_use] + pub fn incoming_edges(&self, node_id: &str) -> Vec<&Edge> { + self.edges.iter().filter(|e| e.to == node_id).collect() + } + + /// Find the start node: shape=Mdiamond, or id "start"/"Start". + #[must_use] + pub fn find_start_node(&self) -> Option<&Node> { + // First: look for shape=Mdiamond + let by_shape = self.nodes.values().find(|n| n.shape() == "Mdiamond"); + if by_shape.is_some() { + return by_shape; + } + // Second: look for id "start" or "Start" + self.nodes + .get("start") + .or_else(|| self.nodes.get("Start")) + } + + /// Find the exit node: shape=Msquare, or id "exit"/"Exit". + #[must_use] + pub fn find_exit_node(&self) -> Option<&Node> { + let by_shape = self.nodes.values().find(|n| n.shape() == "Msquare"); + if by_shape.is_some() { + return by_shape; + } + self.nodes.get("exit").or_else(|| self.nodes.get("Exit")) + } + + /// Graph-level goal attribute. + pub fn goal(&self) -> &str { + self.attrs + .get("goal") + .and_then(AttrValue::as_str) + .unwrap_or("") + } + + /// Graph-level model stylesheet attribute. + pub fn model_stylesheet(&self) -> &str { + self.attrs + .get("model_stylesheet") + .and_then(AttrValue::as_str) + .unwrap_or("") + } + + /// Graph-level `default_max_retry` (default 50). + pub fn default_max_retry(&self) -> i64 { + self.attrs + .get("default_max_retry") + .and_then(AttrValue::as_i64) + .unwrap_or(50) + } + + /// Graph-level `retry_target`. + pub fn retry_target(&self) -> Option<&str> { + self.attrs.get("retry_target").and_then(AttrValue::as_str) + } + + /// Graph-level `fallback_retry_target`. + pub fn fallback_retry_target(&self) -> Option<&str> { + self.attrs + .get("fallback_retry_target") + .and_then(AttrValue::as_str) + } + + /// Graph-level `default_fidelity`. + pub fn default_fidelity(&self) -> Option<&str> { + self.attrs + .get("default_fidelity") + .and_then(AttrValue::as_str) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn attr_value_as_str() { + let val = AttrValue::String("hello".to_string()); + assert_eq!(val.as_str(), Some("hello")); + assert_eq!(AttrValue::Integer(1).as_str(), None); + } + + #[test] + fn attr_value_as_i64() { + assert_eq!(AttrValue::Integer(42).as_i64(), Some(42)); + assert_eq!(AttrValue::String("x".to_string()).as_i64(), None); + } + + #[test] + fn attr_value_as_f64() { + assert_eq!(AttrValue::Float(3.14).as_f64(), Some(3.14)); + assert_eq!(AttrValue::Integer(1).as_f64(), None); + } + + #[test] + fn attr_value_as_bool() { + assert_eq!(AttrValue::Boolean(true).as_bool(), Some(true)); + assert_eq!(AttrValue::String("true".to_string()).as_bool(), None); + } + + #[test] + fn attr_value_as_duration() { + let d = Duration::from_secs(10); + assert_eq!(AttrValue::Duration(d).as_duration(), Some(d)); + assert_eq!(AttrValue::Integer(10).as_duration(), None); + } + + #[test] + fn shape_to_handler_type_mappings() { + assert_eq!(shape_to_handler_type("Mdiamond"), Some("start")); + assert_eq!(shape_to_handler_type("Msquare"), Some("exit")); + assert_eq!(shape_to_handler_type("box"), Some("codergen")); + assert_eq!(shape_to_handler_type("hexagon"), Some("wait.human")); + assert_eq!(shape_to_handler_type("diamond"), Some("conditional")); + assert_eq!(shape_to_handler_type("component"), Some("parallel")); + assert_eq!( + shape_to_handler_type("tripleoctagon"), + Some("parallel.fan_in") + ); + assert_eq!(shape_to_handler_type("parallelogram"), Some("tool")); + assert_eq!( + shape_to_handler_type("house"), + Some("stack.manager_loop") + ); + assert_eq!(shape_to_handler_type("unknown"), None); + } + + #[test] + fn node_defaults() { + let node = Node::new("test"); + assert_eq!(node.id, "test"); + assert_eq!(node.label(), "test"); + assert_eq!(node.shape(), "box"); + assert_eq!(node.node_type(), None); + assert_eq!(node.prompt(), None); + assert_eq!(node.max_retries(), None); + assert!(!node.goal_gate()); + assert_eq!(node.retry_target(), None); + assert_eq!(node.fallback_retry_target(), None); + assert_eq!(node.fidelity(), None); + assert_eq!(node.thread_id(), None); + assert_eq!(node.class(), None); + assert_eq!(node.timeout(), None); + assert_eq!(node.llm_model(), None); + assert_eq!(node.llm_provider(), None); + assert_eq!(node.reasoning_effort(), "high"); + assert!(!node.auto_status()); + assert!(!node.allow_partial()); + } + + #[test] + fn node_with_attrs() { + let mut node = Node::new("plan"); + node.attrs + .insert("label".to_string(), AttrValue::String("Plan step".to_string())); + node.attrs + .insert("shape".to_string(), AttrValue::String("diamond".to_string())); + node.attrs + .insert("goal_gate".to_string(), AttrValue::Boolean(true)); + node.attrs + .insert("max_retries".to_string(), AttrValue::Integer(3)); + + assert_eq!(node.label(), "Plan step"); + assert_eq!(node.shape(), "diamond"); + assert!(node.goal_gate()); + assert_eq!(node.max_retries(), Some(3)); + } + + #[test] + fn node_handler_type_explicit() { + let mut node = Node::new("gate"); + node.attrs + .insert("type".to_string(), AttrValue::String("wait.human".to_string())); + assert_eq!(node.handler_type(), Some("wait.human")); + } + + #[test] + fn node_handler_type_from_shape() { + let mut node = Node::new("entry"); + node.attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + assert_eq!(node.handler_type(), Some("start")); + } + + #[test] + fn edge_defaults() { + let edge = Edge::new("a", "b"); + assert_eq!(edge.from, "a"); + assert_eq!(edge.to, "b"); + assert_eq!(edge.label(), None); + assert_eq!(edge.condition(), None); + assert_eq!(edge.weight(), 0); + assert_eq!(edge.fidelity(), None); + assert_eq!(edge.thread_id(), None); + assert!(!edge.loop_restart()); + assert!(!edge.freeform()); + } + + #[test] + fn edge_with_attrs() { + let mut edge = Edge::new("a", "b"); + edge.attrs + .insert("label".to_string(), AttrValue::String("next".to_string())); + edge.attrs + .insert("condition".to_string(), AttrValue::String("outcome=success".to_string())); + edge.attrs + .insert("weight".to_string(), AttrValue::Integer(5)); + edge.attrs + .insert("loop_restart".to_string(), AttrValue::Boolean(true)); + edge.attrs + .insert("freeform".to_string(), AttrValue::Boolean(true)); + + assert_eq!(edge.label(), Some("next")); + assert_eq!(edge.condition(), Some("outcome=success")); + assert_eq!(edge.weight(), 5); + assert!(edge.loop_restart()); + assert!(edge.freeform()); + } + + fn sample_graph() -> Graph { + let mut g = Graph::new("test_pipeline"); + + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + g.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs + .insert("shape".to_string(), AttrValue::String("Msquare".to_string())); + g.nodes.insert("exit".to_string(), exit); + + let work = Node::new("work"); + g.nodes.insert("work".to_string(), work); + + g.edges.push(Edge::new("start", "work")); + g.edges.push(Edge::new("work", "exit")); + + g.attrs.insert( + "goal".to_string(), + AttrValue::String("Run tests".to_string()), + ); + + g + } + + #[test] + fn graph_find_start_node() { + let g = sample_graph(); + let start = g.find_start_node().unwrap(); + assert_eq!(start.id, "start"); + } + + #[test] + fn graph_find_exit_node() { + let g = sample_graph(); + let exit = g.find_exit_node().unwrap(); + assert_eq!(exit.id, "exit"); + } + + #[test] + fn graph_outgoing_edges() { + let g = sample_graph(); + let edges = g.outgoing_edges("start"); + assert_eq!(edges.len(), 1); + assert_eq!(edges[0].to, "work"); + } + + #[test] + fn graph_incoming_edges() { + let g = sample_graph(); + let edges = g.incoming_edges("exit"); + assert_eq!(edges.len(), 1); + assert_eq!(edges[0].from, "work"); + } + + #[test] + fn graph_goal() { + let g = sample_graph(); + assert_eq!(g.goal(), "Run tests"); + } + + #[test] + fn graph_goal_default() { + let g = Graph::new("empty"); + assert_eq!(g.goal(), ""); + } + + #[test] + fn graph_model_stylesheet_default() { + let g = Graph::new("empty"); + assert_eq!(g.model_stylesheet(), ""); + } + + #[test] + fn graph_default_max_retry() { + let g = Graph::new("empty"); + assert_eq!(g.default_max_retry(), 50); + } + + #[test] + fn graph_find_start_by_id_fallback() { + let mut g = Graph::new("test"); + // No Mdiamond shape, but id is "start" + let node = Node::new("start"); + g.nodes.insert("start".to_string(), node); + assert!(g.find_start_node().is_some()); + } + + #[test] + fn graph_no_start_node() { + let g = Graph::new("empty"); + assert!(g.find_start_node().is_none()); + } +} diff --git a/crates/attractor/src/handler/codergen.rs b/crates/attractor/src/handler/codergen.rs new file mode 100644 index 000000000..fe4cb6da1 --- /dev/null +++ b/crates/attractor/src/handler/codergen.rs @@ -0,0 +1,383 @@ +use std::path::Path; + +use async_trait::async_trait; + +use crate::context::Context; +use crate::error::AttractorError; +use crate::graph::{Graph, Node}; +use crate::outcome::Outcome; + +use super::Handler; + +/// Result from a `CodergenBackend` invocation. +pub enum CodergenResult { + Text(String), + Full(Outcome), +} + +/// Backend interface for LLM execution in codergen nodes. +#[async_trait] +pub trait CodergenBackend: Send + Sync { + async fn run( + &self, + node: &Node, + prompt: &str, + context: &Context, + ) -> Result; +} + +/// The default handler for LLM task nodes. +pub struct CodergenHandler { + backend: Option>, +} + +impl CodergenHandler { + #[must_use] + pub fn new(backend: Option>) -> Self { + Self { backend } + } +} + +/// Expand `$goal` in text using the graph goal. +fn expand_variables(text: &str, graph: &Graph) -> String { + text.replace("$goal", graph.goal()) +} + +/// Truncate a string to at most `max_chars` characters. +fn truncate(s: &str, max_chars: usize) -> &str { + if s.len() <= max_chars { + s + } else { + &s[..max_chars] + } +} + +/// Resolve a tool hook command from node attributes, falling back to graph attributes. +fn resolve_hook(node: &Node, graph: &Graph, key: &str) -> Option { + node.attrs + .get(key) + .and_then(|v| v.as_str()) + .or_else(|| graph.attrs.get(key).and_then(|v| v.as_str())) + .map(String::from) +} + +/// Execute a tool hook shell command. Returns true if the command succeeded (exit 0). +fn run_hook(command: &str, node_id: &str) -> bool { + match std::process::Command::new("sh") + .arg("-c") + .arg(command) + .env("ATTRACTOR_NODE_ID", node_id) + .output() + { + Ok(output) => output.status.success(), + Err(_) => false, + } +} + +#[async_trait] +impl Handler for CodergenHandler { + async fn execute( + &self, + node: &Node, + context: &Context, + graph: &Graph, + logs_root: &Path, + ) -> Result { + // 1. Build prompt + let raw_prompt = node + .prompt() + .filter(|p| !p.is_empty()) + .unwrap_or_else(|| node.label()); + let prompt = expand_variables(raw_prompt, graph); + + // 2. Write prompt to logs + let stage_dir = logs_root.join(&node.id); + tokio::fs::create_dir_all(&stage_dir).await?; + tokio::fs::write(stage_dir.join("prompt.md"), &prompt).await?; + + // 3. Execute pre-hook (spec 9.7) + if let Some(pre_hook) = resolve_hook(node, graph, "tool_hooks.pre") { + if !run_hook(&pre_hook, &node.id) { + return Ok(Outcome::fail("pre-hook failed, skipping LLM call")); + } + } + + // 4. Call LLM backend + let response_text = if let Some(backend) = &self.backend { + match backend.run(node, &prompt, context).await { + Ok(CodergenResult::Full(outcome)) => { + let status_json = serde_json::to_string_pretty(&outcome) + .unwrap_or_else(|_| "{}".to_string()); + tokio::fs::write(stage_dir.join("status.json"), &status_json).await?; + return Ok(outcome); + } + Ok(CodergenResult::Text(text)) => text, + Err(e) => { + return Ok(Outcome::fail(e.to_string())); + } + } + } else { + format!("[Simulated] Response for stage: {}", node.id) + }; + + // 5. Execute post-hook (spec 9.7) + if let Some(post_hook) = resolve_hook(node, graph, "tool_hooks.post") { + if !run_hook(&post_hook, &node.id) { + context.append_log(format!( + "post-hook failed for node {}, continuing", + node.id + )); + } + } + + // 6. Write response to logs + tokio::fs::write(stage_dir.join("response.md"), &response_text).await?; + + // 7. Build and write status + let mut outcome = Outcome::success(); + outcome.notes = Some(format!("Stage completed: {}", node.id)); + outcome.context_updates.insert( + "last_stage".to_string(), + serde_json::json!(node.id), + ); + outcome.context_updates.insert( + "last_response".to_string(), + serde_json::json!(truncate(&response_text, 200)), + ); + + let status_json = serde_json::to_string_pretty(&outcome) + .unwrap_or_else(|_| "{}".to_string()); + tokio::fs::write(stage_dir.join("status.json"), &status_json).await?; + + Ok(outcome) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::AttrValue; + use tempfile::TempDir; + + #[tokio::test] + async fn codergen_handler_simulation_mode() { + let handler = CodergenHandler::new(None); + let mut node = Node::new("plan"); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Plan the implementation".to_string()), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let tmp = TempDir::new().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + assert_eq!(outcome.status, crate::outcome::StageStatus::Success); + assert_eq!(outcome.notes.as_deref(), Some("Stage completed: plan")); + + // Check files were written + let prompt_path = tmp.path().join("plan").join("prompt.md"); + assert!(prompt_path.exists()); + let prompt_content = std::fs::read_to_string(&prompt_path).unwrap(); + assert_eq!(prompt_content, "Plan the implementation"); + + let response_path = tmp.path().join("plan").join("response.md"); + assert!(response_path.exists()); + let response_content = std::fs::read_to_string(&response_path).unwrap(); + assert!(response_content.contains("[Simulated]")); + + let status_path = tmp.path().join("plan").join("status.json"); + assert!(status_path.exists()); + } + + #[tokio::test] + async fn codergen_handler_variable_expansion() { + let handler = CodergenHandler::new(None); + let mut node = Node::new("plan"); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Achieve: $goal".to_string()), + ); + let context = Context::new(); + let mut graph = Graph::new("test"); + graph.attrs.insert( + "goal".to_string(), + AttrValue::String("Build a feature".to_string()), + ); + let tmp = TempDir::new().unwrap(); + + handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + + let prompt_content = + std::fs::read_to_string(tmp.path().join("plan").join("prompt.md")).unwrap(); + assert_eq!(prompt_content, "Achieve: Build a feature"); + } + + #[tokio::test] + async fn codergen_handler_falls_back_to_label() { + let handler = CodergenHandler::new(None); + let mut node = Node::new("work"); + node.attrs.insert( + "label".to_string(), + AttrValue::String("Do work".to_string()), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let tmp = TempDir::new().unwrap(); + + handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + + let prompt_content = + std::fs::read_to_string(tmp.path().join("work").join("prompt.md")).unwrap(); + assert_eq!(prompt_content, "Do work"); + } + + #[tokio::test] + async fn codergen_handler_context_updates() { + let handler = CodergenHandler::new(None); + let node = Node::new("step"); + let context = Context::new(); + let graph = Graph::new("test"); + let tmp = TempDir::new().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + + assert_eq!( + outcome.context_updates.get("last_stage"), + Some(&serde_json::json!("step")) + ); + assert!(outcome.context_updates.contains_key("last_response")); + } + + #[test] + fn expand_variables_replaces_goal() { + let mut graph = Graph::new("test"); + graph.attrs.insert( + "goal".to_string(), + AttrValue::String("Fix bugs".to_string()), + ); + let result = expand_variables("Goal is: $goal, do it", &graph); + assert_eq!(result, "Goal is: Fix bugs, do it"); + } + + #[test] + fn truncate_short_string() { + assert_eq!(truncate("hello", 200), "hello"); + } + + #[test] + fn truncate_long_string() { + let long = "a".repeat(300); + assert_eq!(truncate(&long, 200).len(), 200); + } + + #[tokio::test] + async fn codergen_handler_pre_hook_failure_skips_backend() { + let handler = CodergenHandler::new(None); + let mut node = Node::new("step"); + node.attrs.insert( + "tool_hooks.pre".to_string(), + AttrValue::String("exit 1".to_string()), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let tmp = TempDir::new().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + assert_eq!(outcome.status, crate::outcome::StageStatus::Fail); + assert!(outcome + .failure_reason + .as_deref() + .unwrap() + .contains("pre-hook")); + } + + #[tokio::test] + async fn codergen_handler_pre_hook_success_continues() { + let handler = CodergenHandler::new(None); + let mut node = Node::new("step"); + node.attrs.insert( + "tool_hooks.pre".to_string(), + AttrValue::String("exit 0".to_string()), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let tmp = TempDir::new().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + assert_eq!(outcome.status, crate::outcome::StageStatus::Success); + } + + #[tokio::test] + async fn codergen_handler_post_hook_failure_logs_warning() { + let handler = CodergenHandler::new(None); + let mut node = Node::new("step"); + node.attrs.insert( + "tool_hooks.post".to_string(), + AttrValue::String("exit 1".to_string()), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let tmp = TempDir::new().unwrap(); + + let outcome = handler + .execute(&node, &context, &graph, tmp.path()) + .await + .unwrap(); + // Post-hook failure should not fail the node + assert_eq!(outcome.status, crate::outcome::StageStatus::Success); + } + + #[test] + fn resolve_hook_from_node_attr() { + let mut node = Node::new("step"); + node.attrs.insert( + "tool_hooks.pre".to_string(), + AttrValue::String("echo node".to_string()), + ); + let graph = Graph::new("test"); + assert_eq!( + resolve_hook(&node, &graph, "tool_hooks.pre"), + Some("echo node".to_string()) + ); + } + + #[test] + fn resolve_hook_falls_back_to_graph() { + let node = Node::new("step"); + let mut graph = Graph::new("test"); + graph.attrs.insert( + "tool_hooks.pre".to_string(), + AttrValue::String("echo graph".to_string()), + ); + assert_eq!( + resolve_hook(&node, &graph, "tool_hooks.pre"), + Some("echo graph".to_string()) + ); + } + + #[test] + fn resolve_hook_none_when_missing() { + let node = Node::new("step"); + let graph = Graph::new("test"); + assert_eq!(resolve_hook(&node, &graph, "tool_hooks.pre"), None); + } +} diff --git a/crates/attractor/src/handler/conditional.rs b/crates/attractor/src/handler/conditional.rs new file mode 100644 index 000000000..4e1251655 --- /dev/null +++ b/crates/attractor/src/handler/conditional.rs @@ -0,0 +1,52 @@ +use std::path::Path; + +use async_trait::async_trait; + +use crate::context::Context; +use crate::error::AttractorError; +use crate::graph::{Graph, Node}; +use crate::outcome::Outcome; + +use super::Handler; + +/// Conditional routing handler. Returns SUCCESS with a note; actual routing +/// is handled by the engine's edge selection algorithm. +pub struct ConditionalHandler; + +#[async_trait] +impl Handler for ConditionalHandler { + async fn execute( + &self, + node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + let mut outcome = Outcome::success(); + outcome.notes = Some(format!("Conditional node evaluated: {}", node.id)); + Ok(outcome) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn conditional_handler_returns_success_with_note() { + let handler = ConditionalHandler; + let node = Node::new("gate"); + let context = Context::new(); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, crate::outcome::StageStatus::Success); + assert_eq!( + outcome.notes.as_deref(), + Some("Conditional node evaluated: gate") + ); + } +} diff --git a/crates/attractor/src/handler/exit.rs b/crates/attractor/src/handler/exit.rs new file mode 100644 index 000000000..97ac5fd88 --- /dev/null +++ b/crates/attractor/src/handler/exit.rs @@ -0,0 +1,45 @@ +use std::path::Path; + +use async_trait::async_trait; + +use crate::context::Context; +use crate::error::AttractorError; +use crate::graph::{Graph, Node}; +use crate::outcome::Outcome; + +use super::Handler; + +/// No-op handler for pipeline exit point. Returns SUCCESS immediately. +pub struct ExitHandler; + +#[async_trait] +impl Handler for ExitHandler { + async fn execute( + &self, + _node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + Ok(Outcome::success()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn exit_handler_returns_success() { + let handler = ExitHandler; + let node = Node::new("exit"); + let context = Context::new(); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, crate::outcome::StageStatus::Success); + } +} diff --git a/crates/attractor/src/handler/fan_in.rs b/crates/attractor/src/handler/fan_in.rs new file mode 100644 index 000000000..a8dff44d3 --- /dev/null +++ b/crates/attractor/src/handler/fan_in.rs @@ -0,0 +1,341 @@ +use std::path::Path; + +use async_trait::async_trait; + +use crate::context::Context; +use crate::error::AttractorError; +use crate::graph::{Graph, Node}; +use crate::outcome::Outcome; + +use super::codergen::{CodergenBackend, CodergenResult}; +use super::Handler; + +/// Consolidates results from a preceding parallel node and selects the best candidate. +pub struct FanInHandler { + backend: Option>, +} + +impl FanInHandler { + #[must_use] + pub fn new(backend: Option>) -> Self { + Self { backend } + } +} + +#[async_trait] +impl Handler for FanInHandler { + async fn execute( + &self, + node: &Node, + context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + let results = context.get("parallel.results"); + let Some(results) = results else { + return Ok(Outcome::fail("No parallel results to evaluate")); + }; + + let prompt = node.prompt().filter(|p| !p.is_empty()); + + let best = if let (Some(prompt_text), Some(backend)) = (prompt, &self.backend) { + llm_evaluate(backend.as_ref(), prompt_text, &results, context).await? + } else { + heuristic_select(&results) + }; + + let mut outcome = Outcome::success(); + outcome.context_updates.insert( + "parallel.fan_in.best_id".to_string(), + serde_json::json!(best.id), + ); + outcome.context_updates.insert( + "parallel.fan_in.best_outcome".to_string(), + serde_json::json!(best.status), + ); + outcome.notes = Some(format!("Selected best candidate: {}", best.id)); + + Ok(outcome) + } +} + +struct Candidate { + id: String, + status: String, +} + +fn status_rank(status: &str) -> u32 { + match status { + "success" => 0, + "partial_success" => 1, + "retry" => 2, + "fail" => 3, + _ => 4, + } +} + +fn heuristic_select(results: &serde_json::Value) -> Candidate { + let empty_vec = vec![]; + let arr = results.as_array().unwrap_or(&empty_vec); + if arr.is_empty() { + return Candidate { + id: "unknown".to_string(), + status: "fail".to_string(), + }; + } + + let mut candidates: Vec = arr + .iter() + .map(|v| Candidate { + id: v + .get("id") + .and_then(|v| v.as_str()) + .unwrap_or("unknown") + .to_string(), + status: v + .get("status") + .and_then(|v| v.as_str()) + .unwrap_or("fail") + .to_string(), + }) + .collect(); + + candidates.sort_by(|a, b| { + let rank_cmp = status_rank(&a.status).cmp(&status_rank(&b.status)); + if rank_cmp != std::cmp::Ordering::Equal { + return rank_cmp; + } + a.id.cmp(&b.id) + }); + + candidates.into_iter().next().unwrap_or_else(|| Candidate { + id: "unknown".to_string(), + status: "fail".to_string(), + }) +} + +/// Use an LLM backend to evaluate and rank parallel branch results. +async fn llm_evaluate( + backend: &dyn CodergenBackend, + prompt: &str, + results: &serde_json::Value, + context: &Context, +) -> Result { + let results_text = serde_json::to_string_pretty(results) + .unwrap_or_else(|_| results.to_string()); + + let full_prompt = format!( + "{prompt}\n\nParallel branch results:\n{results_text}\n\n\ + Respond with the ID of the best candidate." + ); + + // Build a synthetic node for the backend call + let eval_node = Node::new("fan_in_eval"); + + match backend.run(&eval_node, &full_prompt, context).await { + Ok(CodergenResult::Full(outcome)) => { + // If the backend returned a full Outcome, extract best_id from context_updates + let best_id = outcome + .context_updates + .get("parallel.fan_in.best_id") + .and_then(|v| v.as_str()) + .map(String::from) + .or_else(|| outcome.notes.clone()) + .unwrap_or_else(|| "unknown".to_string()); + Ok(Candidate { + id: best_id, + status: outcome.status.to_string(), + }) + } + Ok(CodergenResult::Text(text)) => { + // The LLM responded with text; try to find a matching candidate ID + let text = text.trim().to_string(); + let empty_vec = vec![]; + let arr = results.as_array().unwrap_or(&empty_vec); + + // Check if the response text matches any candidate ID + for v in arr { + if let Some(id) = v.get("id").and_then(|v| v.as_str()) { + if text.contains(id) { + let status = v + .get("status") + .and_then(|v| v.as_str()) + .unwrap_or("success") + .to_string(); + return Ok(Candidate { + id: id.to_string(), + status, + }); + } + } + } + + // No match found; fall back to heuristic + Ok(heuristic_select(results)) + } + Err(_) => { + // LLM call failed; fall back to heuristic + Ok(heuristic_select(results)) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::outcome::StageStatus; + + #[tokio::test] + async fn fan_in_no_results() { + let handler = FanInHandler::new(None); + let node = Node::new("fan_in"); + let context = Context::new(); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Fail); + } + + #[tokio::test] + async fn fan_in_selects_best() { + let handler = FanInHandler::new(None); + let node = Node::new("fan_in"); + let context = Context::new(); + context.set( + "parallel.results", + serde_json::json!([ + {"id": "branch_a", "status": "fail"}, + {"id": "branch_b", "status": "success"}, + ]), + ); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Success); + assert_eq!( + outcome.context_updates.get("parallel.fan_in.best_id"), + Some(&serde_json::json!("branch_b")) + ); + } + + #[tokio::test] + async fn fan_in_lexical_tiebreak() { + let handler = FanInHandler::new(None); + let node = Node::new("fan_in"); + let context = Context::new(); + context.set( + "parallel.results", + serde_json::json!([ + {"id": "c", "status": "success"}, + {"id": "a", "status": "success"}, + {"id": "b", "status": "success"}, + ]), + ); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!( + outcome.context_updates.get("parallel.fan_in.best_id"), + Some(&serde_json::json!("a")) + ); + } + + #[test] + fn status_rank_ordering() { + assert!(status_rank("success") < status_rank("partial_success")); + assert!(status_rank("partial_success") < status_rank("retry")); + assert!(status_rank("retry") < status_rank("fail")); + } + + #[tokio::test] + async fn fan_in_no_backend_ignores_prompt() { + // When there's a prompt but no backend, it should fall back to heuristic + let handler = FanInHandler::new(None); + let mut node = Node::new("fan_in"); + node.attrs.insert( + "prompt".to_string(), + crate::graph::AttrValue::String("Pick the best branch".to_string()), + ); + let context = Context::new(); + context.set( + "parallel.results", + serde_json::json!([ + {"id": "branch_a", "status": "success"}, + {"id": "branch_b", "status": "fail"}, + ]), + ); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Success); + // Should still pick branch_a via heuristic (success beats fail) + assert_eq!( + outcome.context_updates.get("parallel.fan_in.best_id"), + Some(&serde_json::json!("branch_a")) + ); + } + + #[tokio::test] + async fn fan_in_with_backend_llm_eval() { + use crate::handler::codergen::CodergenBackend; + + struct MockBackend; + + #[async_trait] + impl CodergenBackend for MockBackend { + async fn run( + &self, + _node: &Node, + _prompt: &str, + _context: &Context, + ) -> Result { + // Return text that contains the ID "branch_b" + Ok(CodergenResult::Text("The best candidate is branch_b".to_string())) + } + } + + let handler = FanInHandler::new(Some(Box::new(MockBackend))); + let mut node = Node::new("fan_in"); + node.attrs.insert( + "prompt".to_string(), + crate::graph::AttrValue::String("Pick the best branch".to_string()), + ); + let context = Context::new(); + context.set( + "parallel.results", + serde_json::json!([ + {"id": "branch_a", "status": "success"}, + {"id": "branch_b", "status": "success"}, + ]), + ); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Success); + // LLM chose branch_b + assert_eq!( + outcome.context_updates.get("parallel.fan_in.best_id"), + Some(&serde_json::json!("branch_b")) + ); + } +} diff --git a/crates/attractor/src/handler/manager_loop.rs b/crates/attractor/src/handler/manager_loop.rs new file mode 100644 index 000000000..94092db1d --- /dev/null +++ b/crates/attractor/src/handler/manager_loop.rs @@ -0,0 +1,309 @@ +use std::path::Path; +use std::time::Duration; + +use async_trait::async_trait; + +use crate::condition::evaluate_condition; +use crate::context::Context; +use crate::error::AttractorError; +use crate::graph::{Graph, Node}; +use crate::outcome::{Outcome, StageStatus}; + +use super::Handler; + +/// Trait for observing child pipeline state during the manager loop. +#[async_trait] +pub trait ChildObserver: Send + Sync { + /// Ingest child telemetry into the context. + async fn observe(&self, context: &Context) -> Result<(), AttractorError>; + + /// Optionally steer the child pipeline (e.g., write intervention instructions). + async fn steer(&self, context: &Context, node: &Node) -> Result<(), AttractorError>; +} + +/// Orchestrates observe/steer/wait cycles over a child pipeline. +pub struct ManagerLoopHandler { + observer: Option>, +} + +impl ManagerLoopHandler { + #[must_use] + pub fn new(observer: Option>) -> Self { + Self { observer } + } +} + +/// Parse a duration string like "45s", "200ms", "5m" into a Duration. +/// Falls back to 45 seconds on parse failure. +fn parse_duration_str(s: &str) -> Duration { + let s = s.trim(); + if let Some(secs) = s.strip_suffix('s') { + if let Some(ms) = secs.strip_suffix('m') { + // "ms" suffix + if let Ok(val) = ms.parse::() { + return Duration::from_millis(val); + } + } else if let Ok(val) = secs.parse::() { + return Duration::from_secs(val); + } + } + if let Some(mins) = s.strip_suffix('m') { + if let Ok(val) = mins.parse::() { + return Duration::from_secs(val * 60); + } + } + Duration::from_secs(45) +} + +#[async_trait] +impl Handler for ManagerLoopHandler { + async fn execute( + &self, + node: &Node, + context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + let poll_interval = node + .attrs + .get("manager.poll_interval") + .and_then(|v| v.as_duration()) + .unwrap_or_else(|| { + let raw = node + .attrs + .get("manager.poll_interval") + .and_then(|v| v.as_str()) + .unwrap_or("45s"); + parse_duration_str(raw) + }); + + let max_cycles = node + .attrs + .get("manager.max_cycles") + .and_then(|v| v.as_i64()) + .unwrap_or(1000); + let max_cycles = u64::try_from(max_cycles).unwrap_or(1000).max(1); + + let stop_condition = node + .attrs + .get("manager.stop_condition") + .and_then(|v| v.as_str()) + .unwrap_or(""); + + let actions_str = node + .attrs + .get("manager.actions") + .and_then(|v| v.as_str()) + .unwrap_or("observe,wait"); + let do_observe = actions_str.contains("observe"); + let do_steer = actions_str.contains("steer"); + let do_wait = actions_str.contains("wait"); + + // Observation loop + for cycle in 1..=max_cycles { + // Observe + if do_observe { + if let Some(ref observer) = self.observer { + observer.observe(context).await?; + } + } + + // Steer + if do_steer { + if let Some(ref observer) = self.observer { + observer.steer(context, node).await?; + } + } + + // Check child status from context + let child_status = context.get_string("context.stack.child.status", ""); + if child_status == "completed" || child_status == "failed" { + let child_outcome = context.get_string("context.stack.child.outcome", ""); + if child_outcome == "success" { + return Ok(Outcome { + status: StageStatus::Success, + notes: Some(format!("Child completed at cycle {cycle}")), + ..Outcome::success() + }); + } + if child_status == "failed" { + return Ok(Outcome::fail(format!("Child failed at cycle {cycle}"))); + } + } + + // Evaluate stop condition + if !stop_condition.is_empty() { + let dummy_outcome = Outcome::success(); + if evaluate_condition(stop_condition, &dummy_outcome, context) { + return Ok(Outcome { + status: StageStatus::Success, + notes: Some(format!("Stop condition satisfied at cycle {cycle}")), + ..Outcome::success() + }); + } + } + + // Wait + if do_wait { + tokio::time::sleep(poll_interval).await; + } + } + + Ok(Outcome::fail(format!( + "Max cycles ({max_cycles}) exceeded for manager loop node: {}", + node.id + ))) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::AttrValue; + + #[tokio::test] + async fn manager_loop_max_cycles_exceeded() { + let handler = ManagerLoopHandler::new(None); + let mut node = Node::new("manager"); + node.attrs.insert( + "manager.max_cycles".to_string(), + AttrValue::Integer(2), + ); + node.attrs.insert( + "manager.poll_interval".to_string(), + AttrValue::Duration(Duration::from_millis(1)), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Fail); + assert!(outcome + .failure_reason + .as_deref() + .unwrap() + .contains("Max cycles")); + } + + #[tokio::test] + async fn manager_loop_child_completed_success() { + let handler = ManagerLoopHandler::new(None); + let mut node = Node::new("manager"); + node.attrs.insert( + "manager.max_cycles".to_string(), + AttrValue::Integer(10), + ); + node.attrs.insert( + "manager.poll_interval".to_string(), + AttrValue::Duration(Duration::from_millis(1)), + ); + // Pre-set child status to "completed" and outcome to "success" + let context = Context::new(); + context.set( + "context.stack.child.status", + serde_json::json!("completed"), + ); + context.set( + "context.stack.child.outcome", + serde_json::json!("success"), + ); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Success); + assert!(outcome.notes.as_deref().unwrap().contains("Child completed")); + } + + #[tokio::test] + async fn manager_loop_child_failed() { + let handler = ManagerLoopHandler::new(None); + let mut node = Node::new("manager"); + node.attrs.insert( + "manager.max_cycles".to_string(), + AttrValue::Integer(10), + ); + node.attrs.insert( + "manager.poll_interval".to_string(), + AttrValue::Duration(Duration::from_millis(1)), + ); + let context = Context::new(); + context.set( + "context.stack.child.status", + serde_json::json!("failed"), + ); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Fail); + assert!(outcome + .failure_reason + .as_deref() + .unwrap() + .contains("Child failed")); + } + + #[tokio::test] + async fn manager_loop_stop_condition_satisfied() { + let handler = ManagerLoopHandler::new(None); + let mut node = Node::new("manager"); + node.attrs.insert( + "manager.max_cycles".to_string(), + AttrValue::Integer(10), + ); + node.attrs.insert( + "manager.poll_interval".to_string(), + AttrValue::Duration(Duration::from_millis(1)), + ); + node.attrs.insert( + "manager.stop_condition".to_string(), + AttrValue::String("context.done=true".to_string()), + ); + let context = Context::new(); + context.set("done", serde_json::json!("true")); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Success); + assert!(outcome + .notes + .as_deref() + .unwrap() + .contains("Stop condition satisfied")); + } + + #[test] + fn parse_duration_str_seconds() { + assert_eq!(parse_duration_str("45s"), Duration::from_secs(45)); + } + + #[test] + fn parse_duration_str_milliseconds() { + assert_eq!(parse_duration_str("200ms"), Duration::from_millis(200)); + } + + #[test] + fn parse_duration_str_minutes() { + assert_eq!(parse_duration_str("5m"), Duration::from_secs(300)); + } + + #[test] + fn parse_duration_str_invalid_fallback() { + assert_eq!(parse_duration_str("bad"), Duration::from_secs(45)); + } +} diff --git a/crates/attractor/src/handler/mod.rs b/crates/attractor/src/handler/mod.rs new file mode 100644 index 000000000..152bd238a --- /dev/null +++ b/crates/attractor/src/handler/mod.rs @@ -0,0 +1,178 @@ +pub mod codergen; +pub mod conditional; +pub mod exit; +pub mod fan_in; +pub mod manager_loop; +pub mod parallel; +pub mod start; +pub mod tool; +pub mod wait_human; + +use std::collections::HashMap; +use std::path::Path; + +use async_trait::async_trait; + +use crate::context::Context; +use crate::error::AttractorError; +use crate::graph::{shape_to_handler_type, Graph, Node}; +use crate::outcome::Outcome; + +/// The handler interface for node execution. +#[async_trait] +pub trait Handler: Send + Sync { + async fn execute( + &self, + node: &Node, + context: &Context, + graph: &Graph, + logs_root: &Path, + ) -> Result; +} + +/// Maps handler type strings to handler implementations. +pub struct HandlerRegistry { + handlers: HashMap>, + default_handler: Box, +} + +impl HandlerRegistry { + #[must_use] + pub fn new(default_handler: Box) -> Self { + Self { + handlers: HashMap::new(), + default_handler, + } + } + + /// Register a handler for a given type string. + pub fn register(&mut self, type_string: impl Into, handler: Box) { + self.handlers.insert(type_string.into(), handler); + } + + /// Resolve which handler should execute for a given node. + /// Priority: explicit type -> shape-based -> default. + #[must_use] + pub fn resolve(&self, node: &Node) -> &dyn Handler { + // 1. Explicit type attribute + if let Some(node_type) = node.node_type() { + if let Some(handler) = self.handlers.get(node_type) { + return handler.as_ref(); + } + } + + // 2. Shape-based resolution + if let Some(handler_type) = shape_to_handler_type(node.shape()) { + if let Some(handler) = self.handlers.get(handler_type) { + return handler.as_ref(); + } + } + + // 3. Default + self.default_handler.as_ref() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::AttrValue; + + struct TestHandler { + _name: String, + } + + #[async_trait] + impl Handler for TestHandler { + async fn execute( + &self, + _node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + Ok(Outcome::success()) + } + } + + #[test] + fn resolve_by_explicit_type() { + let mut registry = HandlerRegistry::new(Box::new(TestHandler { + _name: "default".to_string(), + })); + registry.register( + "wait.human", + Box::new(TestHandler { + _name: "human".to_string(), + }), + ); + + let mut node = Node::new("gate"); + node.attrs.insert( + "type".to_string(), + AttrValue::String("wait.human".to_string()), + ); + let handler = registry.resolve(&node); + // We can verify it returns the right handler by checking it doesn't panic + // and returns a valid reference + let _ = handler; + } + + #[test] + fn resolve_by_shape() { + let mut registry = HandlerRegistry::new(Box::new(TestHandler { + _name: "default".to_string(), + })); + registry.register( + "start", + Box::new(TestHandler { + _name: "start".to_string(), + }), + ); + + let mut node = Node::new("entry"); + node.attrs.insert( + "shape".to_string(), + AttrValue::String("Mdiamond".to_string()), + ); + let handler = registry.resolve(&node); + let _ = handler; + } + + #[test] + fn resolve_falls_back_to_default() { + let registry = HandlerRegistry::new(Box::new(TestHandler { + _name: "default".to_string(), + })); + let node = Node::new("work"); + let handler = registry.resolve(&node); + let _ = handler; + } + + #[test] + fn register_replaces_existing() { + let mut registry = HandlerRegistry::new(Box::new(TestHandler { + _name: "default".to_string(), + })); + registry.register( + "start", + Box::new(TestHandler { + _name: "first".to_string(), + }), + ); + registry.register( + "start", + Box::new(TestHandler { + _name: "second".to_string(), + }), + ); + // Should not panic + let mut node = Node::new("s"); + node.attrs.insert( + "shape".to_string(), + AttrValue::String("Mdiamond".to_string()), + ); + let handler = registry.resolve(&node); + let _ = handler; + } +} diff --git a/crates/attractor/src/handler/parallel.rs b/crates/attractor/src/handler/parallel.rs new file mode 100644 index 000000000..78ec1dc2a --- /dev/null +++ b/crates/attractor/src/handler/parallel.rs @@ -0,0 +1,462 @@ +use std::path::Path; +use std::sync::Arc; +use std::time::Instant; + +use async_trait::async_trait; +use tokio::sync::Semaphore; + +use crate::context::Context; +use crate::error::AttractorError; +use crate::event::{EventEmitter, PipelineEvent}; +use crate::graph::{Graph, Node}; +use crate::outcome::{Outcome, StageStatus}; + +use super::{Handler, HandlerRegistry}; + +/// Convert a Duration's milliseconds to u64, saturating on overflow. +fn millis_u64(d: std::time::Duration) -> u64 { + u64::try_from(d.as_millis()).unwrap_or(u64::MAX) +} + +/// Fans out execution to multiple branches concurrently. +/// Each branch gets an isolated context clone and runs independently. +pub struct ParallelHandler { + registry: Arc, + emitter: Arc, +} + +impl ParallelHandler { + #[must_use] + pub fn new(registry: Arc, emitter: Arc) -> Self { + Self { registry, emitter } + } +} + +/// Parse join policy from node attributes. +#[derive(Debug, Clone)] +enum JoinPolicy { + WaitAll, + FirstSuccess, + KOfN(usize), + Quorum(f64), +} + +fn parse_join_policy(raw: &str) -> JoinPolicy { + if raw == "first_success" { + return JoinPolicy::FirstSuccess; + } + if let Some(inner) = raw.strip_prefix("k_of_n(").and_then(|s| s.strip_suffix(')')) { + if let Ok(k) = inner.trim().parse::() { + return JoinPolicy::KOfN(k); + } + } + if let Some(inner) = raw.strip_prefix("quorum(").and_then(|s| s.strip_suffix(')')) { + if let Ok(frac) = inner.trim().parse::() { + return JoinPolicy::Quorum(frac); + } + } + JoinPolicy::WaitAll +} + +/// Parse error policy from node attributes. +#[derive(Debug, Clone, PartialEq, Eq)] +enum ErrorPolicy { + Continue, + FailFast, + Ignore, +} + +fn parse_error_policy(raw: &str) -> ErrorPolicy { + match raw { + "fail_fast" => ErrorPolicy::FailFast, + "ignore" => ErrorPolicy::Ignore, + _ => ErrorPolicy::Continue, + } +} + +struct BranchResult { + id: String, + outcome: Outcome, +} + +#[async_trait] +impl Handler for ParallelHandler { + async fn execute( + &self, + node: &Node, + context: &Context, + graph: &Graph, + logs_root: &Path, + ) -> Result { + let parallel_start = Instant::now(); + let branches = graph.outgoing_edges(&node.id); + if branches.is_empty() { + return Ok(Outcome::fail("No branches for parallel node")); + } + + self.emitter.emit(&PipelineEvent::ParallelStarted { + branch_count: branches.len(), + }); + + let join_policy = parse_join_policy( + node.attrs + .get("join_policy") + .and_then(|v| v.as_str()) + .unwrap_or("wait_all"), + ); + let error_policy = parse_error_policy( + node.attrs + .get("error_policy") + .and_then(|v| v.as_str()) + .unwrap_or("continue"), + ); + let max_parallel = node + .attrs + .get("max_parallel") + .and_then(|v| v.as_i64()) + .unwrap_or(4); + let max_parallel = usize::try_from(max_parallel).unwrap_or(4).max(1); + + let semaphore = Arc::new(Semaphore::new(max_parallel)); + + // Build branch tasks + let mut handles = Vec::new(); + for (branch_index, edge) in branches.iter().enumerate() { + let target_id = edge.to.clone(); + let branch_context = context.clone_context(); + let registry = Arc::clone(&self.registry); + let emitter = Arc::clone(&self.emitter); + let graph = graph.clone(); + let logs_root = logs_root.to_path_buf(); + let sem = Arc::clone(&semaphore); + + let handle = tokio::spawn(async move { + let _permit = sem.acquire().await.map_err(|e| { + AttractorError::Handler(format!("semaphore error: {e}")) + })?; + + emitter.emit(&PipelineEvent::ParallelBranchStarted { + branch: target_id.clone(), + index: branch_index, + }); + let branch_start = Instant::now(); + + let Some(target_node) = graph.nodes.get(&target_id) else { + let outcome = Outcome::fail(format!("branch target node not found: {target_id}")); + emitter.emit(&PipelineEvent::ParallelBranchCompleted { + branch: target_id.clone(), + index: branch_index, + duration_ms: millis_u64(branch_start.elapsed()), + success: false, + }); + return Ok(BranchResult { + id: target_id.clone(), + outcome, + }); + }; + + let handler = registry.resolve(target_node); + let outcome = handler + .execute(target_node, &branch_context, &graph, &logs_root) + .await?; + + let success = outcome.status == StageStatus::Success + || outcome.status == StageStatus::PartialSuccess; + emitter.emit(&PipelineEvent::ParallelBranchCompleted { + branch: target_id.clone(), + index: branch_index, + duration_ms: millis_u64(branch_start.elapsed()), + success, + }); + + Ok::(BranchResult { + id: target_id, + outcome, + }) + }); + handles.push(handle); + } + + // Collect results + let mut results: Vec = Vec::new(); + for handle in handles { + match handle.await { + Ok(Ok(result)) => { + if error_policy == ErrorPolicy::FailFast + && result.outcome.status == StageStatus::Fail + { + results.push(result); + break; + } + results.push(result); + } + Ok(Err(e)) => { + let result = BranchResult { + id: String::new(), + outcome: Outcome::fail(e.to_string()), + }; + if error_policy == ErrorPolicy::FailFast { + results.push(result); + break; + } + results.push(result); + } + Err(join_err) => { + let result = BranchResult { + id: String::new(), + outcome: Outcome::fail(format!("task join error: {join_err}")), + }; + if error_policy == ErrorPolicy::FailFast { + results.push(result); + break; + } + results.push(result); + } + } + } + + // Count successes and failures + let success_count = results + .iter() + .filter(|r| r.outcome.status == StageStatus::Success) + .count(); + let fail_count = results + .iter() + .filter(|r| r.outcome.status == StageStatus::Fail) + .count(); + let total = results.len(); + + // Store results as JSON in context for downstream fan-in + let results_json: Vec = results + .iter() + .map(|r| { + serde_json::json!({ + "id": r.id, + "status": r.outcome.status.to_string(), + }) + }) + .collect(); + context.set("parallel.results", serde_json::json!(results_json)); + context.set("parallel.branch_count", serde_json::json!(total)); + + self.emitter.emit(&PipelineEvent::ParallelCompleted { + duration_ms: millis_u64(parallel_start.elapsed()), + success_count, + failure_count: fail_count, + }); + + // Evaluate join policy + let status = match join_policy { + JoinPolicy::WaitAll => { + if fail_count == 0 { + StageStatus::Success + } else if error_policy == ErrorPolicy::Ignore { + StageStatus::Success + } else { + StageStatus::PartialSuccess + } + } + JoinPolicy::FirstSuccess => { + if success_count > 0 { + StageStatus::Success + } else { + StageStatus::Fail + } + } + JoinPolicy::KOfN(k) => { + if success_count >= k { + StageStatus::Success + } else { + StageStatus::Fail + } + } + JoinPolicy::Quorum(fraction) => { + #[allow(clippy::cast_precision_loss)] + let total_f64 = total as f64; + let threshold_f64 = (fraction * total_f64).ceil(); + // Safe: threshold_f64 is non-negative and bounded by total + #[allow(clippy::cast_sign_loss, clippy::cast_possible_truncation)] + let threshold = threshold_f64 as usize; + if success_count >= threshold { + StageStatus::Success + } else { + StageStatus::Fail + } + } + }; + + // Build suggested_next_ids from successful branch targets + let branch_ids: Vec = results.iter().map(|r| r.id.clone()).collect(); + + let is_fail = status == StageStatus::Fail; + let mut outcome = Outcome { + status, + preferred_label: None, + suggested_next_ids: branch_ids, + context_updates: std::collections::HashMap::new(), + notes: Some(format!( + "Parallel node dispatched {total} branches ({success_count} succeeded, {fail_count} failed)" + )), + failure_reason: if is_fail { + Some(format!("Join policy not satisfied: {success_count}/{total} succeeded")) + } else { + None + }, + }; + + if is_fail { + outcome.suggested_next_ids.clear(); + } + + Ok(outcome) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::{AttrValue, Edge}; + use crate::handler::start::StartHandler; + + fn make_registry() -> Arc { + let registry = HandlerRegistry::new(Box::new(StartHandler)); + Arc::new(registry) + } + + fn make_emitter() -> Arc { + Arc::new(EventEmitter::new()) + } + + #[tokio::test] + async fn parallel_handler_no_branches() { + let registry = make_registry(); + let handler = ParallelHandler::new(registry, make_emitter()); + let node = Node::new("par"); + let context = Context::new(); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Fail); + } + + #[tokio::test] + async fn parallel_handler_with_branches() { + let registry = make_registry(); + let handler = ParallelHandler::new(registry, make_emitter()); + let mut node = Node::new("par"); + node.attrs.insert( + "shape".to_string(), + AttrValue::String("component".to_string()), + ); + let context = Context::new(); + let mut graph = Graph::new("test"); + graph.nodes.insert("par".to_string(), node.clone()); + graph + .nodes + .insert("branch_a".to_string(), Node::new("branch_a")); + graph + .nodes + .insert("branch_b".to_string(), Node::new("branch_b")); + graph.edges.push(Edge::new("par", "branch_a")); + graph.edges.push(Edge::new("par", "branch_b")); + + let logs_root = Path::new("/tmp/test"); + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + + assert_eq!(outcome.status, StageStatus::Success); + assert!(outcome.notes.as_deref().unwrap().contains("2 branches")); + + // Check context was set + let results = context.get("parallel.results"); + assert!(results.is_some()); + } + + #[tokio::test] + async fn parallel_handler_first_success_policy() { + let registry = make_registry(); + let handler = ParallelHandler::new(registry, make_emitter()); + let mut node = Node::new("par"); + node.attrs.insert( + "join_policy".to_string(), + AttrValue::String("first_success".to_string()), + ); + let context = Context::new(); + let mut graph = Graph::new("test"); + graph.nodes.insert("par".to_string(), node.clone()); + graph + .nodes + .insert("branch_a".to_string(), Node::new("branch_a")); + graph.edges.push(Edge::new("par", "branch_a")); + + let logs_root = Path::new("/tmp/test"); + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + + assert_eq!(outcome.status, StageStatus::Success); + } + + #[tokio::test] + async fn parallel_handler_k_of_n_policy() { + let registry = make_registry(); + let handler = ParallelHandler::new(registry, make_emitter()); + let mut node = Node::new("par"); + node.attrs.insert( + "join_policy".to_string(), + AttrValue::String("k_of_n(2)".to_string()), + ); + let context = Context::new(); + let mut graph = Graph::new("test"); + graph.nodes.insert("par".to_string(), node.clone()); + graph + .nodes + .insert("branch_a".to_string(), Node::new("branch_a")); + graph + .nodes + .insert("branch_b".to_string(), Node::new("branch_b")); + graph + .nodes + .insert("branch_c".to_string(), Node::new("branch_c")); + graph.edges.push(Edge::new("par", "branch_a")); + graph.edges.push(Edge::new("par", "branch_b")); + graph.edges.push(Edge::new("par", "branch_c")); + + let logs_root = Path::new("/tmp/test"); + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + + // All 3 succeed (default StartHandler returns success), need 2 + assert_eq!(outcome.status, StageStatus::Success); + } + + #[test] + fn parse_join_policy_variants() { + assert!(matches!(parse_join_policy("wait_all"), JoinPolicy::WaitAll)); + assert!(matches!( + parse_join_policy("first_success"), + JoinPolicy::FirstSuccess + )); + assert!(matches!(parse_join_policy("k_of_n(3)"), JoinPolicy::KOfN(3))); + assert!(matches!(parse_join_policy("quorum(0.5)"), JoinPolicy::Quorum(_))); + // Invalid falls back to WaitAll + assert!(matches!(parse_join_policy("invalid"), JoinPolicy::WaitAll)); + } + + #[test] + fn parse_error_policy_variants() { + assert_eq!(parse_error_policy("continue"), ErrorPolicy::Continue); + assert_eq!(parse_error_policy("fail_fast"), ErrorPolicy::FailFast); + assert_eq!(parse_error_policy("ignore"), ErrorPolicy::Ignore); + assert_eq!(parse_error_policy("unknown"), ErrorPolicy::Continue); + } +} diff --git a/crates/attractor/src/handler/start.rs b/crates/attractor/src/handler/start.rs new file mode 100644 index 000000000..46ca51bff --- /dev/null +++ b/crates/attractor/src/handler/start.rs @@ -0,0 +1,45 @@ +use std::path::Path; + +use async_trait::async_trait; + +use crate::context::Context; +use crate::error::AttractorError; +use crate::graph::{Graph, Node}; +use crate::outcome::Outcome; + +use super::Handler; + +/// No-op handler for pipeline entry point. Returns SUCCESS immediately. +pub struct StartHandler; + +#[async_trait] +impl Handler for StartHandler { + async fn execute( + &self, + _node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + Ok(Outcome::success()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn start_handler_returns_success() { + let handler = StartHandler; + let node = Node::new("start"); + let context = Context::new(); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, crate::outcome::StageStatus::Success); + } +} diff --git a/crates/attractor/src/handler/tool.rs b/crates/attractor/src/handler/tool.rs new file mode 100644 index 000000000..0864cb57e --- /dev/null +++ b/crates/attractor/src/handler/tool.rs @@ -0,0 +1,185 @@ +use std::path::Path; + +use async_trait::async_trait; + +use crate::context::Context; +use crate::error::AttractorError; +use crate::graph::{Graph, Node}; +use crate::outcome::Outcome; + +use super::Handler; + +/// Executes an external tool (shell command) configured via node attributes. +pub struct ToolHandler; + +fn process_output( + output: std::process::Output, + command: &str, +) -> Result { + let stdout = String::from_utf8_lossy(&output.stdout).to_string(); + let stderr = String::from_utf8_lossy(&output.stderr).to_string(); + + if output.status.success() { + let mut outcome = Outcome::success(); + outcome + .context_updates + .insert("tool.output".to_string(), serde_json::json!(stdout)); + outcome.notes = Some(format!("Tool completed: {command}")); + Ok(outcome) + } else { + let reason = if stderr.is_empty() { + format!( + "Tool failed with exit code: {}", + output.status.code().unwrap_or(-1) + ) + } else { + format!("Tool failed: {}", stderr.trim()) + }; + Ok(Outcome::fail(reason)) + } +} + +#[async_trait] +impl Handler for ToolHandler { + async fn execute( + &self, + node: &Node, + _context: &Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + let command = node + .attrs + .get("tool_command") + .and_then(|v| v.as_str()) + .unwrap_or(""); + + if command.is_empty() { + return Ok(Outcome::fail("No tool_command specified")); + } + + let cmd_future = tokio::process::Command::new("sh") + .arg("-c") + .arg(command) + .output(); + + let output = if let Some(timeout_dur) = node.timeout() { + match tokio::time::timeout(timeout_dur, cmd_future).await { + Ok(result) => result, + Err(_elapsed) => { + return Ok(Outcome::fail(format!( + "Tool timed out after {}ms: {command}", + timeout_dur.as_millis() + ))); + } + } + } else { + cmd_future.await + }; + + match output { + Ok(output) => process_output(output, command), + Err(e) => Ok(Outcome::fail(e.to_string())), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::AttrValue; + use crate::outcome::StageStatus; + use std::time::Duration; + + #[tokio::test] + async fn tool_handler_no_command() { + let handler = ToolHandler; + let node = Node::new("tool_node"); + let context = Context::new(); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Fail); + assert_eq!( + outcome.failure_reason.as_deref(), + Some("No tool_command specified") + ); + } + + #[tokio::test] + async fn tool_handler_echo_command() { + let handler = ToolHandler; + let mut node = Node::new("tool_node"); + node.attrs.insert( + "tool_command".to_string(), + AttrValue::String("echo hello".to_string()), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Success); + assert!(outcome.notes.as_deref().unwrap().contains("echo hello")); + let tool_output = outcome.context_updates.get("tool.output").unwrap(); + assert!(tool_output.as_str().unwrap().contains("hello")); + } + + #[tokio::test] + async fn tool_handler_failing_command() { + let handler = ToolHandler; + let mut node = Node::new("tool_node"); + node.attrs.insert( + "tool_command".to_string(), + AttrValue::String("false".to_string()), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Fail); + } + + #[tokio::test] + async fn tool_handler_timeout() { + let handler = ToolHandler; + let mut node = Node::new("tool_node"); + node.attrs.insert( + "tool_command".to_string(), + AttrValue::String("sleep 60".to_string()), + ); + node.attrs.insert( + "timeout".to_string(), + AttrValue::Duration(Duration::from_millis(50)), + ); + let context = Context::new(); + let graph = Graph::new("test"); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(&node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, StageStatus::Fail); + assert!( + outcome + .failure_reason + .as_deref() + .unwrap() + .contains("timed out"), + "expected timeout message, got: {:?}", + outcome.failure_reason + ); + } +} diff --git a/crates/attractor/src/handler/wait_human.rs b/crates/attractor/src/handler/wait_human.rs new file mode 100644 index 000000000..4dbd60ea0 --- /dev/null +++ b/crates/attractor/src/handler/wait_human.rs @@ -0,0 +1,413 @@ +use std::path::Path; +use std::sync::Arc; +use std::time::Instant; + +use async_trait::async_trait; + +use crate::context::Context; +use crate::error::AttractorError; +use crate::event::{EventEmitter, PipelineEvent}; +use crate::graph::{Graph, Node}; +use crate::interviewer::{ + Answer, AnswerValue, Interviewer, Question, QuestionOption, QuestionType, +}; +use crate::outcome::Outcome; + +use super::Handler; + +/// Convert a Duration's milliseconds to u64, saturating on overflow. +fn millis_u64(d: std::time::Duration) -> u64 { + u64::try_from(d.as_millis()).unwrap_or(u64::MAX) +} + +/// A choice derived from an outgoing edge. +struct Choice { + key: String, + label: String, + to: String, +} + +/// Parse an accelerator key from a label. +/// Patterns: `[K] Label`, `K) Label`, `K - Label`, or first character. +fn parse_accelerator_key(label: &str) -> String { + let trimmed = label.trim(); + + // Pattern: [K] Label + if trimmed.starts_with('[') { + if let Some(end) = trimmed.find(']') { + let key = &trimmed[1..end]; + if !key.is_empty() { + return key.to_string(); + } + } + } + + // Pattern: K) Label + if let Some(paren_pos) = trimmed.find(')') { + if paren_pos > 0 && paren_pos <= 3 { + let key = &trimmed[..paren_pos]; + if key.chars().all(char::is_alphanumeric) { + return key.to_string(); + } + } + } + + // Pattern: K - Label + if let Some(dash_pos) = trimmed.find(" - ") { + if dash_pos > 0 && dash_pos <= 3 { + let key = &trimmed[..dash_pos]; + if key.chars().all(char::is_alphanumeric) { + return key.to_string(); + } + } + } + + // Fallback: first character + trimmed + .chars() + .next() + .map(|c| c.to_string()) + .unwrap_or_default() +} + +/// Blocks until a human selects an option derived from outgoing edges. +pub struct WaitHumanHandler { + interviewer: Arc, + emitter: Option>, +} + +impl WaitHumanHandler { + pub fn new(interviewer: Arc) -> Self { + Self { + interviewer, + emitter: None, + } + } + + #[must_use] + pub fn with_emitter(mut self, emitter: Arc) -> Self { + self.emitter = Some(emitter); + self + } + + fn emit(&self, event: &PipelineEvent) { + if let Some(emitter) = &self.emitter { + emitter.emit(event); + } + } +} + +#[async_trait] +impl Handler for WaitHumanHandler { + async fn execute( + &self, + node: &Node, + _context: &Context, + graph: &Graph, + _logs_root: &Path, + ) -> Result { + // 1. Derive choices from outgoing edges + let edges = graph.outgoing_edges(&node.id); + let mut freeform_target: Option = None; + let mut choices: Vec = Vec::new(); + + for edge in &edges { + if edge.freeform() { + freeform_target = Some(edge.to.clone()); + continue; + } + let label = edge + .label() + .filter(|l| !l.is_empty()) + .unwrap_or(&edge.to); + let key = parse_accelerator_key(label); + choices.push(Choice { + key, + label: label.to_string(), + to: edge.to.clone(), + }); + } + + if choices.is_empty() && freeform_target.is_none() { + return Ok(Outcome::fail("No outgoing edges for human gate")); + } + + // 2. Build question + let options: Vec = choices + .iter() + .map(|c| QuestionOption { + key: c.key.clone(), + label: c.label.clone(), + }) + .collect(); + + let mut question = Question::new( + node.label(), + QuestionType::MultipleChoice, + ); + question.options = options; + question.allow_freeform = freeform_target.is_some(); + question.stage.clone_from(&node.id); + + // 3. Present to interviewer + let question_text = node.label().to_string(); + self.emit(&PipelineEvent::InterviewStarted { + question: question_text.clone(), + stage: node.id.clone(), + }); + let interview_start = Instant::now(); + let answer = self.interviewer.ask(question).await; + + // 4. Handle timeout + if answer.value == AnswerValue::Timeout { + self.emit(&PipelineEvent::InterviewTimeout { + question: question_text, + stage: node.id.clone(), + duration_ms: millis_u64(interview_start.elapsed()), + }); + let default_choice = node + .attrs + .get("human.default_choice") + .and_then(|v| v.as_str()); + if let Some(default_target) = default_choice { + return Ok(make_choice_outcome( + default_target, + default_target, + default_target, + )); + } + return Ok(Outcome::retry("human gate timeout, no default")); + } + + // 5. Handle skipped + if answer.value == AnswerValue::Skipped { + return Ok(Outcome::fail("human skipped interaction")); + } + + // Emit interview completed for successful interactions + self.emit(&PipelineEvent::InterviewCompleted { + question: question_text, + answer: answer_text(&answer), + duration_ms: millis_u64(interview_start.elapsed()), + }); + + // 6. Try fixed-choice match + if let Some(selected) = find_choice_match(&answer, &choices) { + return Ok(make_choice_outcome( + &selected.key, + &selected.label, + &selected.to, + )); + } + + // 7. Freeform fallback + if let Some(freeform_to) = &freeform_target { + let text = answer_text(&answer); + let mut outcome = Outcome::success(); + outcome.suggested_next_ids = vec![freeform_to.clone()]; + outcome.context_updates.insert( + "human.gate.selected".to_string(), + serde_json::json!("freeform"), + ); + outcome.context_updates.insert( + "human.gate.label".to_string(), + serde_json::json!(text), + ); + outcome.context_updates.insert( + "human.gate.text".to_string(), + serde_json::json!(text), + ); + return Ok(outcome); + } + + // 8. Fallback to first choice + if let Some(first) = choices.first() { + return Ok(make_choice_outcome(&first.key, &first.label, &first.to)); + } + + Ok(Outcome::fail("No matching choice")) + } +} + +fn make_choice_outcome(key: &str, label: &str, to: &str) -> Outcome { + let mut outcome = Outcome::success(); + outcome.suggested_next_ids = vec![to.to_string()]; + outcome.context_updates.insert( + "human.gate.selected".to_string(), + serde_json::json!(key), + ); + outcome.context_updates.insert( + "human.gate.label".to_string(), + serde_json::json!(label), + ); + outcome +} + +fn find_choice_match<'a>(answer: &Answer, choices: &'a [Choice]) -> Option<&'a Choice> { + match &answer.value { + AnswerValue::Selected(key) => { + choices.iter().find(|c| c.key == *key) + } + AnswerValue::Text(text) => { + // Try matching by key or label + choices + .iter() + .find(|c| c.key.eq_ignore_ascii_case(text) || c.label.eq_ignore_ascii_case(text)) + } + _ => None, + } +} + +fn answer_text(answer: &Answer) -> String { + if let Some(text) = &answer.text { + return text.clone(); + } + match &answer.value { + AnswerValue::Text(t) => t.clone(), + AnswerValue::Selected(s) => s.clone(), + AnswerValue::Yes => "yes".to_string(), + AnswerValue::No => "no".to_string(), + AnswerValue::Skipped => "skipped".to_string(), + AnswerValue::Timeout => "timeout".to_string(), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::{AttrValue, Edge}; + use crate::interviewer::auto_approve::AutoApproveInterviewer; + + fn build_graph_with_human_gate() -> Graph { + let mut graph = Graph::new("test"); + let mut gate = Node::new("gate"); + gate.attrs.insert( + "shape".to_string(), + AttrValue::String("hexagon".to_string()), + ); + gate.attrs.insert( + "label".to_string(), + AttrValue::String("Review Changes".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")); + + let mut e1 = Edge::new("gate", "approve"); + e1.attrs.insert( + "label".to_string(), + AttrValue::String("[A] Approve".to_string()), + ); + let mut e2 = Edge::new("gate", "reject"); + e2.attrs.insert( + "label".to_string(), + AttrValue::String("[R] Reject".to_string()), + ); + graph.edges.push(e1); + graph.edges.push(e2); + graph + } + + #[test] + fn parse_accelerator_key_bracket() { + assert_eq!(parse_accelerator_key("[A] Approve"), "A"); + assert_eq!(parse_accelerator_key("[Y] Yes, deploy"), "Y"); + } + + #[test] + fn parse_accelerator_key_paren() { + assert_eq!(parse_accelerator_key("Y) Yes, deploy"), "Y"); + } + + #[test] + fn parse_accelerator_key_dash() { + assert_eq!(parse_accelerator_key("Y - Yes, deploy"), "Y"); + } + + #[test] + fn parse_accelerator_key_first_char() { + assert_eq!(parse_accelerator_key("Yes, deploy"), "Y"); + } + + #[test] + fn parse_accelerator_key_empty() { + assert_eq!(parse_accelerator_key(""), ""); + } + + #[tokio::test] + async fn wait_human_auto_approve_selects_first() { + let interviewer = Arc::new(AutoApproveInterviewer); + let handler = WaitHumanHandler::new(interviewer); + let graph = build_graph_with_human_gate(); + let node = graph.nodes.get("gate").unwrap(); + let context = Context::new(); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, crate::outcome::StageStatus::Success); + // Auto-approve picks first option key "A" + assert_eq!( + outcome.context_updates.get("human.gate.selected"), + Some(&serde_json::json!("A")) + ); + assert_eq!(outcome.suggested_next_ids, vec!["approve"]); + } + + #[tokio::test] + async fn wait_human_no_edges_returns_fail() { + let interviewer = Arc::new(AutoApproveInterviewer); + let handler = WaitHumanHandler::new(interviewer); + let mut graph = Graph::new("test"); + let gate = Node::new("gate"); + graph.nodes.insert("gate".to_string(), gate); + let node = graph.nodes.get("gate").unwrap(); + let context = Context::new(); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, crate::outcome::StageStatus::Fail); + } + + #[tokio::test] + async fn wait_human_with_freeform_edge() { + let interviewer = Arc::new(crate::interviewer::callback::CallbackInterviewer::new( + |_| Answer::text("custom input"), + )); + let handler = WaitHumanHandler::new(interviewer); + + let mut graph = Graph::new("test"); + let mut gate = Node::new("gate"); + gate.attrs.insert( + "label".to_string(), + AttrValue::String("Choose".to_string()), + ); + graph.nodes.insert("gate".to_string(), gate); + graph.nodes.insert("freeform_target".to_string(), Node::new("freeform_target")); + + let mut edge = Edge::new("gate", "freeform_target"); + edge.attrs + .insert("freeform".to_string(), AttrValue::Boolean(true)); + graph.edges.push(edge); + + let node = graph.nodes.get("gate").unwrap(); + let context = Context::new(); + let logs_root = Path::new("/tmp/test"); + + let outcome = handler + .execute(node, &context, &graph, logs_root) + .await + .unwrap(); + assert_eq!(outcome.status, crate::outcome::StageStatus::Success); + assert_eq!(outcome.suggested_next_ids, vec!["freeform_target"]); + assert_eq!( + outcome.context_updates.get("human.gate.text"), + Some(&serde_json::json!("custom input")) + ); + } +} diff --git a/crates/attractor/src/interviewer/auto_approve.rs b/crates/attractor/src/interviewer/auto_approve.rs new file mode 100644 index 000000000..f26d06c26 --- /dev/null +++ b/crates/attractor/src/interviewer/auto_approve.rs @@ -0,0 +1,88 @@ +use async_trait::async_trait; + +use super::{Answer, AnswerValue, Interviewer, Question, QuestionType}; + +/// Always approves: YES for yes/no, first option for multiple choice, "auto-approved" for freeform. +pub struct AutoApproveInterviewer; + +#[async_trait] +impl Interviewer for AutoApproveInterviewer { + async fn ask(&self, question: Question) -> Answer { + match question.question_type { + QuestionType::YesNo | QuestionType::Confirmation => Answer::yes(), + QuestionType::MultipleChoice => question.options.first().map_or_else( + || Answer::text("auto-approved"), + |first| Answer { + value: AnswerValue::Selected(first.key.clone()), + selected_option: Some(first.clone()), + text: None, + }, + ), + QuestionType::Freeform => Answer::text("auto-approved"), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::interviewer::QuestionOption; + + #[tokio::test] + async fn yes_no_returns_yes() { + let interviewer = AutoApproveInterviewer; + let q = Question::new("Approve?", QuestionType::YesNo); + let answer = interviewer.ask(q).await; + assert_eq!(answer.value, AnswerValue::Yes); + } + + #[tokio::test] + async fn confirmation_returns_yes() { + let interviewer = AutoApproveInterviewer; + let q = Question::new("Confirm?", QuestionType::Confirmation); + let answer = interviewer.ask(q).await; + assert_eq!(answer.value, AnswerValue::Yes); + } + + #[tokio::test] + async fn multiple_choice_returns_first_option() { + let interviewer = AutoApproveInterviewer; + let mut q = Question::new("Choose:", QuestionType::MultipleChoice); + q.options = vec![ + QuestionOption { + key: "A".to_string(), + label: "Alpha".to_string(), + }, + QuestionOption { + key: "B".to_string(), + label: "Beta".to_string(), + }, + ]; + let answer = interviewer.ask(q).await; + assert_eq!(answer.value, AnswerValue::Selected("A".to_string())); + assert_eq!( + answer.selected_option, + Some(QuestionOption { + key: "A".to_string(), + label: "Alpha".to_string(), + }) + ); + } + + #[tokio::test] + async fn multiple_choice_no_options_returns_auto_approved() { + let interviewer = AutoApproveInterviewer; + let q = Question::new("Choose:", QuestionType::MultipleChoice); + let answer = interviewer.ask(q).await; + assert_eq!(answer.value, AnswerValue::Text("auto-approved".to_string())); + } + + #[tokio::test] + async fn freeform_returns_auto_approved() { + let interviewer = AutoApproveInterviewer; + let q = Question::new("Enter text:", QuestionType::Freeform); + let answer = interviewer.ask(q).await; + assert_eq!(answer.value, AnswerValue::Text("auto-approved".to_string())); + assert_eq!(answer.text, Some("auto-approved".to_string())); + } +} diff --git a/crates/attractor/src/interviewer/callback.rs b/crates/attractor/src/interviewer/callback.rs new file mode 100644 index 000000000..a0b8a99d2 --- /dev/null +++ b/crates/attractor/src/interviewer/callback.rs @@ -0,0 +1,56 @@ +use async_trait::async_trait; + +use super::{Answer, Interviewer, Question}; + +/// Delegates question answering to a provided callback function. +pub struct CallbackInterviewer { + callback: Box Answer + Send + Sync>, +} + +impl CallbackInterviewer { + pub fn new(callback: impl Fn(Question) -> Answer + Send + Sync + 'static) -> Self { + Self { + callback: Box::new(callback), + } + } +} + +#[async_trait] +impl Interviewer for CallbackInterviewer { + async fn ask(&self, question: Question) -> Answer { + (self.callback)(question) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::interviewer::{AnswerValue, QuestionType}; + + #[tokio::test] + async fn calls_callback_with_question() { + let interviewer = CallbackInterviewer::new(|q| { + if q.question_type == QuestionType::YesNo { + Answer::yes() + } else { + Answer::no() + } + }); + + let yes_q = Question::new("approve?", QuestionType::YesNo); + let answer = interviewer.ask(yes_q).await; + assert_eq!(answer.value, AnswerValue::Yes); + + let no_q = Question::new("choose:", QuestionType::MultipleChoice); + let answer = interviewer.ask(no_q).await; + assert_eq!(answer.value, AnswerValue::No); + } + + #[tokio::test] + async fn callback_receives_question_text() { + let interviewer = CallbackInterviewer::new(|q| Answer::text(q.text.clone())); + let q = Question::new("hello world", QuestionType::Freeform); + let answer = interviewer.ask(q).await; + assert_eq!(answer.text, Some("hello world".to_string())); + } +} diff --git a/crates/attractor/src/interviewer/console.rs b/crates/attractor/src/interviewer/console.rs new file mode 100644 index 000000000..255b947d5 --- /dev/null +++ b/crates/attractor/src/interviewer/console.rs @@ -0,0 +1,162 @@ +use async_trait::async_trait; +use tokio::io::{AsyncBufReadExt, BufReader}; + +use super::{Answer, AnswerValue, Interviewer, Question, QuestionType}; + +/// Reads from stdin to collect answers. Displays formatted prompts per spec 6.4. +pub struct ConsoleInterviewer; + +fn find_matching_option( + response: &str, + options: &[super::QuestionOption], +) -> Option { + let trimmed = response.trim(); + // Try matching by key (case-insensitive) + for opt in options { + if opt.key.eq_ignore_ascii_case(trimmed) { + return Some(Answer { + value: AnswerValue::Selected(opt.key.clone()), + selected_option: Some(opt.clone()), + text: None, + }); + } + } + // Try matching by 1-based index + if let Ok(idx) = trimmed.parse::() { + if idx >= 1 && idx <= options.len() { + let opt = &options[idx - 1]; + return Some(Answer { + value: AnswerValue::Selected(opt.key.clone()), + selected_option: Some(opt.clone()), + text: None, + }); + } + } + None +} + +async fn read_line(prompt: &str) -> std::io::Result { + // Print the prompt to stderr so it doesn't interfere with piped stdout + eprint!("{prompt}"); + let stdin = tokio::io::stdin(); + let mut reader = BufReader::new(stdin); + let mut line = String::new(); + reader.read_line(&mut line).await?; + Ok(line.trim_end().to_string()) +} + +#[async_trait] +impl Interviewer for ConsoleInterviewer { + async fn ask(&self, question: Question) -> Answer { + eprintln!("[?] {}", question.text); + + match question.question_type { + QuestionType::MultipleChoice => { + for (i, opt) in question.options.iter().enumerate() { + eprintln!(" [{}] {} - {}", i + 1, opt.key, opt.label); + } + if question.allow_freeform { + eprintln!(" Or type a free-text response"); + } + let response = read_line("Select: ").await.unwrap_or_default(); + if let Some(answer) = find_matching_option(&response, &question.options) { + return answer; + } + if question.allow_freeform { + return Answer::text(response); + } + // Fallback: try match again (spec says to do this) + find_matching_option(&response, &question.options) + .unwrap_or_else(Answer::skipped) + } + QuestionType::YesNo | QuestionType::Confirmation => { + let response = read_line("[Y/N]: ").await.unwrap_or_default(); + let trimmed = response.trim().to_lowercase(); + if trimmed == "y" || trimmed == "yes" { + Answer::yes() + } else { + Answer::no() + } + } + QuestionType::Freeform => { + let response = read_line("> ").await.unwrap_or_default(); + Answer::text(response) + } + } + } + + async fn inform(&self, message: &str, stage: &str) { + eprintln!("[{stage}] {message}"); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn find_matching_option_by_key() { + let options = vec![ + super::super::QuestionOption { + key: "A".to_string(), + label: "Approve".to_string(), + }, + super::super::QuestionOption { + key: "R".to_string(), + label: "Reject".to_string(), + }, + ]; + let result = find_matching_option("A", &options); + assert!(result.is_some()); + let answer = result.unwrap(); + assert_eq!(answer.value, AnswerValue::Selected("A".to_string())); + } + + #[test] + fn find_matching_option_by_key_case_insensitive() { + let options = vec![super::super::QuestionOption { + key: "Y".to_string(), + label: "Yes".to_string(), + }]; + let result = find_matching_option("y", &options); + assert!(result.is_some()); + } + + #[test] + fn find_matching_option_by_index() { + let options = vec![ + super::super::QuestionOption { + key: "A".to_string(), + label: "Alpha".to_string(), + }, + super::super::QuestionOption { + key: "B".to_string(), + label: "Beta".to_string(), + }, + ]; + let result = find_matching_option("2", &options); + assert!(result.is_some()); + let answer = result.unwrap(); + assert_eq!(answer.value, AnswerValue::Selected("B".to_string())); + } + + #[test] + fn find_matching_option_no_match() { + let options = vec![super::super::QuestionOption { + key: "A".to_string(), + label: "Alpha".to_string(), + }]; + let result = find_matching_option("zzz", &options); + assert!(result.is_none()); + } + + #[test] + fn find_matching_option_index_out_of_range() { + let options = vec![super::super::QuestionOption { + key: "A".to_string(), + label: "Alpha".to_string(), + }]; + let result = find_matching_option("5", &options); + assert!(result.is_none()); + } +} diff --git a/crates/attractor/src/interviewer/mod.rs b/crates/attractor/src/interviewer/mod.rs new file mode 100644 index 000000000..b6fcab581 --- /dev/null +++ b/crates/attractor/src/interviewer/mod.rs @@ -0,0 +1,297 @@ +pub mod auto_approve; +pub mod callback; +pub mod console; +pub mod queue; +pub mod recording; + +use std::collections::HashMap; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; + +/// The type of question being asked. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum QuestionType { + YesNo, + MultipleChoice, + Freeform, + Confirmation, +} + +/// An option presented to the user for multiple-choice questions. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct QuestionOption { + pub key: String, + pub label: String, +} + +/// A question presented to the user. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Question { + pub text: String, + pub question_type: QuestionType, + pub options: Vec, + pub allow_freeform: bool, + pub default: Option, + pub timeout_seconds: Option, + pub stage: String, + pub metadata: HashMap, +} + +impl Question { + pub fn new(text: impl Into, question_type: QuestionType) -> Self { + Self { + text: text.into(), + question_type, + options: Vec::new(), + allow_freeform: false, + default: None, + timeout_seconds: None, + stage: String::new(), + metadata: HashMap::new(), + } + } +} + +/// The value of an answer. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum AnswerValue { + Yes, + No, + Skipped, + Timeout, + Selected(String), + Text(String), +} + +/// An answer from the user. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Answer { + pub value: AnswerValue, + pub selected_option: Option, + pub text: Option, +} + +impl Answer { + #[must_use] + pub const fn yes() -> Self { + Self { + value: AnswerValue::Yes, + selected_option: None, + text: None, + } + } + + #[must_use] + pub const fn no() -> Self { + Self { + value: AnswerValue::No, + selected_option: None, + text: None, + } + } + + #[must_use] + pub const fn skipped() -> Self { + Self { + value: AnswerValue::Skipped, + selected_option: None, + text: None, + } + } + + #[must_use] + pub const fn timeout() -> Self { + Self { + value: AnswerValue::Timeout, + selected_option: None, + text: None, + } + } + + pub fn selected(key: impl Into, option: QuestionOption) -> Self { + let key = key.into(); + Self { + value: AnswerValue::Selected(key), + selected_option: Some(option), + text: None, + } + } + + pub fn text(text: impl Into) -> Self { + let t = text.into(); + Self { + value: AnswerValue::Text(t.clone()), + selected_option: None, + text: Some(t), + } + } +} + +/// Apply timeout enforcement to an interviewer ask call. +/// Per spec 6.5: if `timeout_seconds` is set, returns default answer or `Answer::timeout()`. +pub async fn ask_with_timeout( + interviewer: &dyn Interviewer, + question: Question, +) -> Answer { + let timeout_secs = question.timeout_seconds; + let default_answer = question.default.clone(); + + if let Some(secs) = timeout_secs { + let duration = std::time::Duration::from_secs_f64(secs); + match tokio::time::timeout(duration, interviewer.ask(question)).await { + Ok(answer) => answer, + Err(_elapsed) => default_answer.unwrap_or_else(Answer::timeout), + } + } else { + interviewer.ask(question).await + } +} + +/// The interviewer trait for human-in-the-loop interactions. +#[async_trait] +pub trait Interviewer: Send + Sync { + async fn ask(&self, question: Question) -> Answer; + + async fn ask_multiple(&self, questions: Vec) -> Vec { + let mut answers = Vec::with_capacity(questions.len()); + for q in questions { + answers.push(self.ask(q).await); + } + answers + } + + async fn inform(&self, _message: &str, _stage: &str) { + // Default no-op + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn question_new() { + let q = Question::new("Do you approve?", QuestionType::YesNo); + assert_eq!(q.text, "Do you approve?"); + assert_eq!(q.question_type, QuestionType::YesNo); + assert!(q.options.is_empty()); + assert!(!q.allow_freeform); + assert!(q.default.is_none()); + assert!(q.timeout_seconds.is_none()); + assert!(q.stage.is_empty()); + assert!(q.metadata.is_empty()); + } + + #[test] + fn answer_yes() { + let a = Answer::yes(); + assert_eq!(a.value, AnswerValue::Yes); + assert!(a.selected_option.is_none()); + assert!(a.text.is_none()); + } + + #[test] + fn answer_no() { + let a = Answer::no(); + assert_eq!(a.value, AnswerValue::No); + } + + #[test] + fn answer_skipped() { + let a = Answer::skipped(); + assert_eq!(a.value, AnswerValue::Skipped); + } + + #[test] + fn answer_timeout() { + let a = Answer::timeout(); + assert_eq!(a.value, AnswerValue::Timeout); + } + + #[test] + fn answer_selected() { + let opt = QuestionOption { + key: "A".to_string(), + label: "Approve".to_string(), + }; + let a = Answer::selected("A", opt.clone()); + assert_eq!(a.value, AnswerValue::Selected("A".to_string())); + assert_eq!(a.selected_option, Some(opt)); + } + + #[test] + fn answer_text() { + let a = Answer::text("free input"); + assert_eq!(a.value, AnswerValue::Text("free input".to_string())); + assert_eq!(a.text, Some("free input".to_string())); + } + + #[test] + fn question_option_eq() { + let a = QuestionOption { + key: "Y".to_string(), + label: "Yes".to_string(), + }; + let b = QuestionOption { + key: "Y".to_string(), + label: "Yes".to_string(), + }; + assert_eq!(a, b); + } + + #[test] + fn answer_value_variants() { + assert_ne!(AnswerValue::Yes, AnswerValue::No); + assert_ne!(AnswerValue::Skipped, AnswerValue::Timeout); + assert_eq!( + AnswerValue::Selected("x".to_string()), + AnswerValue::Selected("x".to_string()) + ); + assert_eq!( + AnswerValue::Text("hello".to_string()), + AnswerValue::Text("hello".to_string()) + ); + } + + /// A slow interviewer that waits before answering -- for testing timeouts. + struct SlowInterviewer; + + #[async_trait] + impl Interviewer for SlowInterviewer { + async fn ask(&self, _question: Question) -> Answer { + tokio::time::sleep(std::time::Duration::from_secs(60)).await; + Answer::yes() + } + } + + #[tokio::test] + async fn ask_with_timeout_returns_timeout_when_expired() { + let interviewer = SlowInterviewer; + let mut q = Question::new("approve?", QuestionType::YesNo); + q.timeout_seconds = Some(0.01); + + let answer = ask_with_timeout(&interviewer, q).await; + assert_eq!(answer.value, AnswerValue::Timeout); + } + + #[tokio::test] + async fn ask_with_timeout_returns_default_when_set() { + let interviewer = SlowInterviewer; + let mut q = Question::new("approve?", QuestionType::YesNo); + q.timeout_seconds = Some(0.01); + q.default = Some(Answer::no()); + + let answer = ask_with_timeout(&interviewer, q).await; + assert_eq!(answer.value, AnswerValue::No); + } + + #[tokio::test] + async fn ask_with_timeout_no_timeout_returns_normally() { + let interviewer = crate::interviewer::auto_approve::AutoApproveInterviewer; + let q = Question::new("approve?", QuestionType::YesNo); + + let answer = ask_with_timeout(&interviewer, q).await; + assert_eq!(answer.value, AnswerValue::Yes); + } +} diff --git a/crates/attractor/src/interviewer/queue.rs b/crates/attractor/src/interviewer/queue.rs new file mode 100644 index 000000000..d33617ec6 --- /dev/null +++ b/crates/attractor/src/interviewer/queue.rs @@ -0,0 +1,66 @@ +use std::collections::VecDeque; +use std::sync::Mutex; + +use async_trait::async_trait; + +use super::{Answer, Interviewer, Question}; + +/// Reads answers from a pre-filled queue. Returns Skipped when empty. +pub struct QueueInterviewer { + answers: Mutex>, +} + +impl QueueInterviewer { + #[must_use] + pub const fn new(answers: VecDeque) -> Self { + Self { + answers: Mutex::new(answers), + } + } +} + +#[async_trait] +impl Interviewer for QueueInterviewer { + async fn ask(&self, _question: Question) -> Answer { + let mut queue = self.answers.lock().expect("queue lock poisoned"); + queue.pop_front().unwrap_or_else(Answer::skipped) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::interviewer::{AnswerValue, QuestionType}; + + #[tokio::test] + async fn returns_queued_answers_in_order() { + let answers = VecDeque::from([Answer::yes(), Answer::no()]); + let interviewer = QueueInterviewer::new(answers); + let q = Question::new("q1", QuestionType::YesNo); + + let a1 = interviewer.ask(q.clone()).await; + assert_eq!(a1.value, AnswerValue::Yes); + + let a2 = interviewer.ask(q).await; + assert_eq!(a2.value, AnswerValue::No); + } + + #[tokio::test] + async fn returns_skipped_when_empty() { + let interviewer = QueueInterviewer::new(VecDeque::new()); + let q = Question::new("q", QuestionType::YesNo); + let answer = interviewer.ask(q).await; + assert_eq!(answer.value, AnswerValue::Skipped); + } + + #[tokio::test] + async fn returns_skipped_after_exhausted() { + let answers = VecDeque::from([Answer::yes()]); + let interviewer = QueueInterviewer::new(answers); + let q = Question::new("q", QuestionType::YesNo); + + let _ = interviewer.ask(q.clone()).await; + let answer = interviewer.ask(q).await; + assert_eq!(answer.value, AnswerValue::Skipped); + } +} diff --git a/crates/attractor/src/interviewer/recording.rs b/crates/attractor/src/interviewer/recording.rs new file mode 100644 index 000000000..9af342ec5 --- /dev/null +++ b/crates/attractor/src/interviewer/recording.rs @@ -0,0 +1,84 @@ +use std::sync::Mutex; + +use async_trait::async_trait; + +use super::{Answer, Interviewer, Question}; + +/// Wraps another interviewer and records all question-answer pairs. +pub struct RecordingInterviewer { + inner: Box, + recordings: Mutex>, +} + +impl RecordingInterviewer { + #[must_use] + pub fn new(inner: Box) -> Self { + Self { + inner, + recordings: Mutex::new(Vec::new()), + } + } + + /// # Panics + /// Panics if the internal mutex is poisoned. + #[must_use] + pub fn recordings(&self) -> Vec<(Question, Answer)> { + self.recordings.lock().expect("recordings lock poisoned").clone() + } +} + +#[async_trait] +impl Interviewer for RecordingInterviewer { + async fn ask(&self, question: Question) -> Answer { + let answer = self.inner.ask(question.clone()).await; + self.recordings + .lock() + .expect("recordings lock poisoned") + .push((question, answer.clone())); + answer + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::interviewer::auto_approve::AutoApproveInterviewer; + use crate::interviewer::{AnswerValue, QuestionType}; + + #[tokio::test] + async fn records_question_answer_pairs() { + let inner = Box::new(AutoApproveInterviewer); + let recorder = RecordingInterviewer::new(inner); + + let q1 = Question::new("approve?", QuestionType::YesNo); + let q2 = Question::new("confirm?", QuestionType::Confirmation); + + let a1 = recorder.ask(q1).await; + assert_eq!(a1.value, AnswerValue::Yes); + + let a2 = recorder.ask(q2).await; + assert_eq!(a2.value, AnswerValue::Yes); + + let recs = recorder.recordings(); + assert_eq!(recs.len(), 2); + assert_eq!(recs[0].0.text, "approve?"); + assert_eq!(recs[1].0.text, "confirm?"); + } + + #[tokio::test] + async fn delegates_to_inner() { + let inner = Box::new(AutoApproveInterviewer); + let recorder = RecordingInterviewer::new(inner); + + let q = Question::new("text input", QuestionType::Freeform); + let answer = recorder.ask(q).await; + assert_eq!(answer.value, AnswerValue::Text("auto-approved".to_string())); + } + + #[tokio::test] + async fn recordings_empty_initially() { + let inner = Box::new(AutoApproveInterviewer); + let recorder = RecordingInterviewer::new(inner); + assert!(recorder.recordings().is_empty()); + } +} diff --git a/crates/attractor/src/lib.rs b/crates/attractor/src/lib.rs new file mode 100644 index 000000000..8ed07702a --- /dev/null +++ b/crates/attractor/src/lib.rs @@ -0,0 +1,16 @@ +pub mod artifact; +pub mod checkpoint; +pub mod condition; +pub mod context; +pub mod engine; +pub mod error; +pub mod event; +pub mod graph; +pub mod handler; +pub mod interviewer; +pub mod outcome; +pub mod parser; +pub mod pipeline; +pub mod stylesheet; +pub mod transform; +pub mod validation; diff --git a/crates/attractor/src/outcome.rs b/crates/attractor/src/outcome.rs new file mode 100644 index 000000000..ddf201b8c --- /dev/null +++ b/crates/attractor/src/outcome.rs @@ -0,0 +1,197 @@ +use std::collections::HashMap; +use std::fmt; +use std::str::FromStr; + +use serde::{Deserialize, Serialize}; + +/// Status of a pipeline stage execution. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum StageStatus { + Success, + Fail, + PartialSuccess, + Retry, + Skipped, +} + +impl fmt::Display for StageStatus { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let s = match self { + Self::Success => "success", + Self::Fail => "fail", + Self::PartialSuccess => "partial_success", + Self::Retry => "retry", + Self::Skipped => "skipped", + }; + write!(f, "{s}") + } +} + +impl FromStr for StageStatus { + type Err = String; + + fn from_str(s: &str) -> std::result::Result { + match s { + "success" => Ok(Self::Success), + "fail" => Ok(Self::Fail), + "partial_success" => Ok(Self::PartialSuccess), + "retry" => Ok(Self::Retry), + "skipped" => Ok(Self::Skipped), + other => Err(format!("unknown stage status: {other}")), + } + } +} + +/// The result of executing a node handler. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Outcome { + pub status: StageStatus, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub preferred_label: Option, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub suggested_next_ids: Vec, + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub context_updates: HashMap, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub notes: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub failure_reason: Option, +} + +impl Outcome { + #[must_use] + pub fn success() -> Self { + Self { + status: StageStatus::Success, + preferred_label: None, + suggested_next_ids: Vec::new(), + context_updates: HashMap::new(), + notes: None, + failure_reason: None, + } + } + + pub fn fail(reason: impl Into) -> Self { + Self { + status: StageStatus::Fail, + preferred_label: None, + suggested_next_ids: Vec::new(), + context_updates: HashMap::new(), + notes: None, + failure_reason: Some(reason.into()), + } + } + + pub fn retry(reason: impl Into) -> Self { + Self { + status: StageStatus::Retry, + preferred_label: None, + suggested_next_ids: Vec::new(), + context_updates: HashMap::new(), + notes: None, + failure_reason: Some(reason.into()), + } + } + + #[must_use] + pub fn skipped() -> Self { + Self { + status: StageStatus::Skipped, + preferred_label: None, + suggested_next_ids: Vec::new(), + context_updates: HashMap::new(), + notes: None, + failure_reason: None, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn stage_status_display() { + assert_eq!(StageStatus::Success.to_string(), "success"); + assert_eq!(StageStatus::Fail.to_string(), "fail"); + assert_eq!(StageStatus::PartialSuccess.to_string(), "partial_success"); + assert_eq!(StageStatus::Retry.to_string(), "retry"); + assert_eq!(StageStatus::Skipped.to_string(), "skipped"); + } + + #[test] + fn stage_status_from_str() { + assert_eq!("success".parse::().unwrap(), StageStatus::Success); + assert_eq!("fail".parse::().unwrap(), StageStatus::Fail); + assert_eq!( + "partial_success".parse::().unwrap(), + StageStatus::PartialSuccess + ); + assert_eq!("retry".parse::().unwrap(), StageStatus::Retry); + assert_eq!("skipped".parse::().unwrap(), StageStatus::Skipped); + } + + #[test] + fn stage_status_from_str_invalid() { + assert!("unknown".parse::().is_err()); + } + + #[test] + fn outcome_success_factory() { + let o = Outcome::success(); + assert_eq!(o.status, StageStatus::Success); + assert!(o.preferred_label.is_none()); + assert!(o.suggested_next_ids.is_empty()); + assert!(o.context_updates.is_empty()); + assert!(o.notes.is_none()); + assert!(o.failure_reason.is_none()); + } + + #[test] + fn outcome_fail_factory() { + let o = Outcome::fail("something broke"); + assert_eq!(o.status, StageStatus::Fail); + assert_eq!(o.failure_reason.as_deref(), Some("something broke")); + } + + #[test] + fn outcome_retry_factory() { + let o = Outcome::retry("try again"); + assert_eq!(o.status, StageStatus::Retry); + assert_eq!(o.failure_reason.as_deref(), Some("try again")); + } + + #[test] + fn outcome_skipped_factory() { + let o = Outcome::skipped(); + assert_eq!(o.status, StageStatus::Skipped); + assert!(o.failure_reason.is_none()); + } + + #[test] + fn outcome_serialization_roundtrip() { + let mut o = Outcome::success(); + o.notes = Some("done".to_string()); + o.context_updates + .insert("key".to_string(), serde_json::json!("val")); + + let json = serde_json::to_string(&o).unwrap(); + let deserialized: Outcome = serde_json::from_str(&json).unwrap(); + + assert_eq!(deserialized.status, StageStatus::Success); + assert_eq!(deserialized.notes.as_deref(), Some("done")); + assert_eq!( + deserialized.context_updates.get("key"), + Some(&serde_json::json!("val")) + ); + } + + #[test] + fn stage_status_serde_roundtrip() { + let json = serde_json::to_string(&StageStatus::PartialSuccess).unwrap(); + assert_eq!(json, "\"partial_success\""); + let parsed: StageStatus = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed, StageStatus::PartialSuccess); + } +} diff --git a/crates/attractor/src/parser/ast.rs b/crates/attractor/src/parser/ast.rs new file mode 100644 index 000000000..0ba63bcff --- /dev/null +++ b/crates/attractor/src/parser/ast.rs @@ -0,0 +1,119 @@ +/// A parsed DOT value before semantic interpretation. +#[derive(Debug, Clone, PartialEq)] +pub enum AstValue { + Str(String), + Int(i64), + Float(f64), + Bool(bool), + /// A bare identifier used as a value (e.g., shape names, direction keywords). + Ident(String), +} + +/// A list of key-value attribute pairs from an attribute block `[k=v, ...]`. +pub type AttrBlock = Vec<(String, AstValue)>; + +/// A node statement: `id [attrs]?`. +#[derive(Debug, Clone, PartialEq)] +pub struct NodeStmt { + pub id: String, + pub attrs: Option, +} + +/// An edge statement: `A -> B -> C [attrs]?`. +#[derive(Debug, Clone, PartialEq)] +pub struct EdgeStmt { + /// Chain of node IDs (at least 2). + pub nodes: Vec, + pub attrs: Option, +} + +/// A subgraph statement: `subgraph name? { stmts }`. +#[derive(Debug, Clone, PartialEq)] +pub struct SubgraphStmt { + pub name: Option, + pub statements: Vec, +} + +/// A single statement in a DOT graph body. +#[derive(Debug, Clone, PartialEq)] +pub enum Statement { + /// `graph [attrs]` + GraphAttr(AttrBlock), + /// `node [attrs]` + NodeDefaults(AttrBlock), + /// `edge [attrs]` + EdgeDefaults(AttrBlock), + /// `subgraph name? { ... }` + Subgraph(SubgraphStmt), + /// `id [attrs]?` + Node(NodeStmt), + /// `A -> B -> C [attrs]?` + Edge(EdgeStmt), + /// Top-level `key = value` + GraphAttrDecl(String, AstValue), +} + +/// The top-level parsed DOT graph. +#[derive(Debug, Clone, PartialEq)] +pub struct DotGraph { + pub name: String, + pub statements: Vec, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn ast_value_variants() { + let s = AstValue::Str("hello".into()); + let i = AstValue::Int(42); + let f = AstValue::Float(3.14); + let b = AstValue::Bool(true); + let id = AstValue::Ident("LR".into()); + + assert_eq!(s, AstValue::Str("hello".into())); + assert_eq!(i, AstValue::Int(42)); + assert_eq!(f, AstValue::Float(3.14)); + assert_eq!(b, AstValue::Bool(true)); + assert_eq!(id, AstValue::Ident("LR".into())); + } + + #[test] + fn dot_graph_construction() { + let graph = DotGraph { + name: "test".into(), + statements: vec![ + Statement::GraphAttrDecl("rankdir".into(), AstValue::Ident("LR".into())), + Statement::Node(NodeStmt { + id: "start".into(), + attrs: Some(vec![("shape".into(), AstValue::Ident("Mdiamond".into()))]), + }), + ], + }; + assert_eq!(graph.name, "test"); + assert_eq!(graph.statements.len(), 2); + } + + #[test] + fn edge_stmt_chained() { + let edge = EdgeStmt { + nodes: vec!["A".into(), "B".into(), "C".into()], + attrs: Some(vec![("label".into(), AstValue::Str("next".into()))]), + }; + assert_eq!(edge.nodes.len(), 3); + } + + #[test] + fn subgraph_stmt() { + let sub = SubgraphStmt { + name: Some("cluster_loop".into()), + statements: vec![Statement::NodeDefaults(vec![( + "timeout".into(), + AstValue::Str("900s".into()), + )])], + }; + assert_eq!(sub.name.as_deref(), Some("cluster_loop")); + assert_eq!(sub.statements.len(), 1); + } +} diff --git a/crates/attractor/src/parser/grammar.rs b/crates/attractor/src/parser/grammar.rs new file mode 100644 index 000000000..40acdf29f --- /dev/null +++ b/crates/attractor/src/parser/grammar.rs @@ -0,0 +1,404 @@ +use nom::branch::alt; +use nom::character::complete::char; +use nom::combinator::opt; +use nom::multi::{many0, separated_list1}; +use nom::sequence::{delimited, preceded, tuple}; +use nom::IResult; + +use crate::parser::ast::{ + AstValue, AttrBlock, DotGraph, EdgeStmt, NodeStmt, Statement, SubgraphStmt, +}; +use crate::parser::lexer::combinators::{identifier, key, value, ws, ws_tag}; + +/// Parse a single attribute: `key = value`. +fn attr(input: &str) -> IResult<&str, (String, AstValue)> { + let (rest, (k, _, _, v)) = tuple(( + preceded(ws, key), + ws, + char('='), + preceded(ws, value), + ))(input)?; + Ok((rest, (k, v))) +} + +/// Parse an attribute block: `[ attr (, attr)* ]`. +fn attr_block(input: &str) -> IResult<&str, AttrBlock> { + delimited( + preceded(ws, char('[')), + separated_list1(preceded(ws, char(',')), attr), + preceded(ws, char(']')), + )(input) +} + +/// Parse optional semicolon. +fn opt_semi(input: &str) -> IResult<&str, Option> { + preceded(ws, opt(char(';')))(input) +} + +/// Parse a graph attr statement: `graph [attrs] ;?` +fn graph_attr_stmt(input: &str) -> IResult<&str, Statement> { + let (rest, (_, attrs, _)) = tuple((ws_tag("graph"), attr_block, opt_semi))(input)?; + Ok((rest, Statement::GraphAttr(attrs))) +} + +/// Parse node defaults: `node [attrs] ;?` +fn node_defaults(input: &str) -> IResult<&str, Statement> { + let (rest, (_, attrs, _)) = tuple((ws_tag("node"), attr_block, opt_semi))(input)?; + Ok((rest, Statement::NodeDefaults(attrs))) +} + +/// Parse edge defaults: `edge [attrs] ;?` +fn edge_defaults(input: &str) -> IResult<&str, Statement> { + let (rest, (_, attrs, _)) = tuple((ws_tag("edge"), attr_block, opt_semi))(input)?; + Ok((rest, Statement::EdgeDefaults(attrs))) +} + +/// Parse a graph attr declaration: `identifier = value ;?` +fn graph_attr_decl(input: &str) -> IResult<&str, Statement> { + let (rest, (k, _, _, v, _)) = tuple(( + preceded(ws, identifier), + ws, + char('='), + preceded(ws, value), + opt_semi, + ))(input)?; + Ok((rest, Statement::GraphAttrDecl(k.to_string(), v))) +} + +/// Parse a subgraph: `subgraph name? { statement* }` +fn subgraph_stmt(input: &str) -> IResult<&str, Statement> { + let (rest, _) = ws_tag("subgraph")(input)?; + let (rest, name) = opt(preceded(ws, identifier))(rest)?; + let (rest, _) = preceded(ws, char('{'))(rest)?; + let (rest, stmts) = many0(statement)(rest)?; + let (rest, _) = preceded(ws, char('}'))(rest)?; + Ok(( + rest, + Statement::Subgraph(SubgraphStmt { + name: name.map(String::from), + statements: stmts, + }), + )) +} + +/// Parse an edge or node statement. +/// If an identifier is followed by `->`, parse as edge; otherwise as node. +fn node_or_edge_stmt(input: &str) -> IResult<&str, Statement> { + let (rest, first_id) = preceded(ws, identifier)(input)?; + + // Try to parse as edge: first_id (-> id)+ [attrs]? ;? + if let Ok((rest2, _)) = arrow::>(rest) { + let (rest2, second_id) = preceded(ws, identifier)(rest2)?; + let mut nodes = vec![first_id.to_string(), second_id.to_string()]; + let mut remaining = rest2; + while let Ok((r, _)) = arrow::>(remaining) { + let (r, next_id) = preceded(ws, identifier)(r)?; + nodes.push(next_id.to_string()); + remaining = r; + } + let (remaining, attrs) = opt(attr_block)(remaining)?; + let (remaining, _) = opt_semi(remaining)?; + return Ok((remaining, Statement::Edge(EdgeStmt { nodes, attrs }))); + } + + // Parse as node: first_id [attrs]? ;? + let (rest, attrs) = opt(attr_block)(rest)?; + let (rest, _) = opt_semi(rest)?; + Ok(( + rest, + Statement::Node(NodeStmt { + id: first_id.to_string(), + attrs, + }), + )) +} + +/// Parse a single statement. +fn statement(input: &str) -> IResult<&str, Statement> { + preceded( + ws, + alt(( + graph_attr_stmt, + node_defaults, + edge_defaults, + subgraph_stmt, + // graph_attr_decl must be tried before node_or_edge because both start with an identifier. + // graph_attr_decl is `id = value` while node is `id [attrs]?` + // We try graph_attr_decl first; if it fails (no `=` after id) we fall through to node_or_edge. + graph_attr_decl, + node_or_edge_stmt, + )), + )(input) +} + +/// Parse a complete DOT graph: `digraph name { statement* }`. +/// +/// # Errors +/// +/// Returns a nom error if the input does not match the DOT grammar. +pub fn parse_dot_graph(input: &str) -> IResult<&str, DotGraph> { + let (rest, _) = ws_tag("digraph")(input)?; + let (rest, name) = preceded(ws, identifier)(rest)?; + let (rest, _) = preceded(ws, char('{'))(rest)?; + let (rest, stmts) = many0(statement)(rest)?; + let (rest, _) = preceded(ws, char('}'))(rest)?; + Ok(( + rest, + DotGraph { + name: name.to_string(), + statements: stmts, + }, + )) +} + +// We need arrow to work with explicit error types +fn arrow<'a, E: nom::error::ParseError<&'a str>>(input: &'a str) -> IResult<&'a str, &'a str, E> { + preceded( + nom::character::complete::multispace0, + nom::bytes::complete::tag("->"), + )(input) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::parser::ast::AstValue; + + #[test] + fn parse_single_attr() { + let (rest, (k, v)) = attr(" label = \"Hello\"").unwrap(); + assert_eq!(k, "label"); + assert_eq!(v, AstValue::Str("Hello".into())); + assert_eq!(rest, ""); + } + + #[test] + fn parse_attr_block_single() { + let (rest, attrs) = attr_block("[label=\"Hello\"]").unwrap(); + assert_eq!(attrs.len(), 1); + assert_eq!(attrs[0].0, "label"); + assert_eq!(rest, ""); + } + + #[test] + fn parse_attr_block_multiple() { + let (rest, attrs) = attr_block("[shape=Mdiamond, label=\"Start\"]").unwrap(); + assert_eq!(attrs.len(), 2); + assert_eq!(attrs[0].0, "shape"); + assert_eq!(attrs[0].1, AstValue::Ident("Mdiamond".into())); + assert_eq!(attrs[1].0, "label"); + assert_eq!(attrs[1].1, AstValue::Str("Start".into())); + assert_eq!(rest, ""); + } + + #[test] + fn parse_graph_attr_stmt() { + let (_, stmt) = graph_attr_stmt("graph [goal=\"Run tests\"]").unwrap(); + match stmt { + Statement::GraphAttr(attrs) => { + assert_eq!(attrs.len(), 1); + assert_eq!(attrs[0].0, "goal"); + } + _ => panic!("expected GraphAttr"), + } + } + + #[test] + fn parse_node_defaults_stmt() { + let (_, stmt) = node_defaults("node [shape=box, timeout=\"900s\"]").unwrap(); + assert!(matches!(stmt, Statement::NodeDefaults(_))); + } + + #[test] + fn parse_edge_defaults_stmt() { + let (_, stmt) = edge_defaults("edge [weight=0]").unwrap(); + assert!(matches!(stmt, Statement::EdgeDefaults(_))); + } + + #[test] + fn parse_graph_attr_decl_stmt() { + let (_, stmt) = graph_attr_decl("rankdir=LR").unwrap(); + match stmt { + Statement::GraphAttrDecl(k, v) => { + assert_eq!(k, "rankdir"); + assert_eq!(v, AstValue::Ident("LR".into())); + } + _ => panic!("expected GraphAttrDecl"), + } + } + + #[test] + fn parse_node_stmt_simple() { + let (_, stmt) = node_or_edge_stmt("start [shape=Mdiamond, label=\"Start\"]").unwrap(); + match stmt { + Statement::Node(n) => { + assert_eq!(n.id, "start"); + assert!(n.attrs.is_some()); + } + _ => panic!("expected Node"), + } + } + + #[test] + fn parse_node_stmt_no_attrs() { + let (_, stmt) = node_or_edge_stmt("run_tests ;").unwrap(); + match stmt { + Statement::Node(n) => { + assert_eq!(n.id, "run_tests"); + assert!(n.attrs.is_none()); + } + _ => panic!("expected Node"), + } + } + + #[test] + fn parse_edge_stmt_simple() { + let (_, stmt) = node_or_edge_stmt("start -> run_tests").unwrap(); + match stmt { + Statement::Edge(e) => { + assert_eq!(e.nodes, vec!["start", "run_tests"]); + assert!(e.attrs.is_none()); + } + _ => panic!("expected Edge"), + } + } + + #[test] + fn parse_edge_stmt_chained() { + let (_, stmt) = node_or_edge_stmt("start -> run_tests -> report -> exit").unwrap(); + match stmt { + Statement::Edge(e) => { + assert_eq!(e.nodes, vec!["start", "run_tests", "report", "exit"]); + } + _ => panic!("expected Edge"), + } + } + + #[test] + fn parse_edge_stmt_with_attrs() { + let (_, stmt) = + node_or_edge_stmt("gate -> exit [label=\"Yes\", condition=\"outcome=success\"]") + .unwrap(); + match stmt { + Statement::Edge(e) => { + assert_eq!(e.nodes, vec!["gate", "exit"]); + let attrs = e.attrs.unwrap(); + assert_eq!(attrs.len(), 2); + } + _ => panic!("expected Edge"), + } + } + + #[test] + fn parse_subgraph() { + let input = r#"subgraph cluster_loop { + label = "Loop A" + node [thread_id="loop-a"] + Plan [label="Plan next step"] + }"#; + let (_, stmt) = subgraph_stmt(input).unwrap(); + match stmt { + Statement::Subgraph(s) => { + assert_eq!(s.name.as_deref(), Some("cluster_loop")); + assert_eq!(s.statements.len(), 3); + } + _ => panic!("expected Subgraph"), + } + } + + #[test] + fn parse_full_simple_graph() { + let input = r#"digraph Simple { + graph [goal="Run tests and report"] + rankdir=LR + + start [shape=Mdiamond, label="Start"] + exit [shape=Msquare, label="Exit"] + + run_tests [label="Run Tests", prompt="Run the test suite and report results"] + report [label="Report", prompt="Summarize the test results"] + + start -> run_tests -> report -> exit + }"#; + let (_, graph) = parse_dot_graph(input).unwrap(); + assert_eq!(graph.name, "Simple"); + assert_eq!(graph.statements.len(), 7); + } + + #[test] + fn parse_full_branching_graph() { + let input = r#"digraph Branch { + graph [goal="Implement and validate a feature"] + rankdir=LR + node [shape=box, timeout="900s"] + + start [shape=Mdiamond, label="Start"] + exit [shape=Msquare, label="Exit"] + plan [label="Plan", prompt="Plan the implementation"] + implement [label="Implement", prompt="Implement the plan"] + validate [label="Validate", prompt="Run tests"] + gate [shape=diamond, label="Tests passing?"] + + start -> plan -> implement -> validate -> gate + gate -> exit [label="Yes", condition="outcome=success"] + gate -> implement [label="No", condition="outcome!=success"] + }"#; + let (_, graph) = parse_dot_graph(input).unwrap(); + assert_eq!(graph.name, "Branch"); + // graph [goal=...], rankdir=LR, node [defaults], 6 nodes, 1 chain + 2 edges = 12 + assert!(graph.statements.len() >= 11); + } + + #[test] + fn parse_human_gate_graph() { + let input = r#"digraph Review { + rankdir=LR + + start [shape=Mdiamond, label="Start"] + exit [shape=Msquare, label="Exit"] + + review_gate [ + shape=hexagon, + label="Review Changes", + type="wait.human" + ] + + start -> review_gate + review_gate -> ship_it [label="[A] Approve"] + review_gate -> fixes [label="[F] Fix"] + ship_it -> exit + fixes -> review_gate + }"#; + let (_, graph) = parse_dot_graph(input).unwrap(); + assert_eq!(graph.name, "Review"); + } + + #[test] + fn parse_qualified_key_attr() { + let (rest, (k, v)) = attr(" tool_hooks.pre = \"echo hello\"").unwrap(); + assert_eq!(k, "tool_hooks.pre"); + assert_eq!(v, AstValue::Str("echo hello".into())); + assert_eq!(rest, ""); + } + + #[test] + fn parse_duration_attr() { + let (_, (k, v)) = attr(" timeout = 900s").unwrap(); + assert_eq!(k, "timeout"); + assert_eq!(v, AstValue::Str("900s".into())); + } + + #[test] + fn parse_boolean_attr() { + let (_, (k, v)) = attr(" goal_gate = true").unwrap(); + assert_eq!(k, "goal_gate"); + assert_eq!(v, AstValue::Bool(true)); + } + + #[test] + fn parse_integer_attr() { + let (_, (k, v)) = attr(" max_retries = 3").unwrap(); + assert_eq!(k, "max_retries"); + assert_eq!(v, AstValue::Int(3)); + } +} diff --git a/crates/attractor/src/parser/lexer.rs b/crates/attractor/src/parser/lexer.rs new file mode 100644 index 000000000..54cdb58ba --- /dev/null +++ b/crates/attractor/src/parser/lexer.rs @@ -0,0 +1,413 @@ +#![allow(clippy::module_name_repetitions)] + +/// Strip `//` line comments and `/* */` block comments from DOT source. +#[must_use] +pub fn strip_comments(input: &str) -> String { + let mut result = String::with_capacity(input.len()); + let chars: Vec = input.chars().collect(); + let len = chars.len(); + let mut i = 0; + + while i < len { + if i + 1 < len && chars[i] == '/' && chars[i + 1] == '/' { + // Line comment: skip to end of line + i += 2; + while i < len && chars[i] != '\n' { + i += 1; + } + } else if i + 1 < len && chars[i] == '/' && chars[i + 1] == '*' { + // Block comment: skip to closing */ + i += 2; + while i + 1 < len && !(chars[i] == '*' && chars[i + 1] == '/') { + if chars[i] == '\n' { + result.push('\n'); + } + i += 1; + } + if i + 1 < len { + i += 2; // skip */ + } + } else if chars[i] == '"' { + // Quoted string: pass through without stripping + result.push(chars[i]); + i += 1; + while i < len && chars[i] != '"' { + result.push(chars[i]); + if chars[i] == '\\' && i + 1 < len { + i += 1; + result.push(chars[i]); + } + i += 1; + } + if i < len { + result.push(chars[i]); // closing quote + i += 1; + } + } else { + result.push(chars[i]); + i += 1; + } + } + + result +} + +/// nom combinators for whitespace and common tokens. +#[allow( + clippy::missing_errors_doc, + clippy::must_use_candidate, + clippy::missing_const_for_fn +)] +pub mod combinators { + use nom::branch::alt; + use nom::bytes::complete::{tag, take_while, take_while1}; + use nom::character::complete::{char, multispace0}; + use nom::combinator::{map, opt, recognize}; + use nom::sequence::{delimited, pair, preceded}; + use nom::IResult; + + use crate::parser::ast::AstValue; + + /// Parse optional whitespace (including newlines). + pub fn ws(input: &str) -> IResult<&str, &str> { + multispace0(input) + } + + /// Parse a token surrounded by optional whitespace. + pub fn ws_tag<'a>(t: &'a str) -> impl Fn(&'a str) -> IResult<&'a str, &'a str> { + move |input| delimited(ws, tag(t), ws)(input) + } + + /// Parse an identifier: `[A-Za-z_][A-Za-z0-9_]*`. + pub fn identifier(input: &str) -> IResult<&str, &str> { + recognize(pair( + take_while1(|c: char| c.is_ascii_alphabetic() || c == '_'), + take_while(|c: char| c.is_ascii_alphanumeric() || c == '_'), + ))(input) + } + + /// Parse a qualified ID: `identifier(.identifier)+`. + pub fn qualified_id(input: &str) -> IResult<&str, String> { + let (rest, first) = identifier(input)?; + let mut result = first.to_string(); + let mut remaining = rest; + let mut found_dot = false; + while let Ok((r, _)) = char::<&str, nom::error::Error<&str>>('.')(remaining) { + if let Ok((r2, segment)) = identifier(r) { + result.push('.'); + result.push_str(segment); + remaining = r2; + found_dot = true; + } else { + break; + } + } + if found_dot { + Ok((remaining, result)) + } else { + Err(nom::Err::Error(nom::error::Error::new( + input, + nom::error::ErrorKind::Tag, + ))) + } + } + + /// Parse a key: either a qualified ID or a simple identifier. + pub fn key(input: &str) -> IResult<&str, String> { + alt((qualified_id, map(identifier, String::from)))(input) + } + + /// Parse a double-quoted string with escape handling. + pub fn quoted_string(input: &str) -> IResult<&str, String> { + let (input, _) = char('"')(input)?; + let mut result = String::new(); + let mut chars = input.chars(); + let mut consumed = 0; + + loop { + match chars.next() { + Some('"') => { + consumed += 1; + return Ok((&input[consumed..], result)); + } + Some('\\') => { + consumed += 1; + match chars.next() { + Some('"') => { + result.push('"'); + consumed += 1; + } + Some('n') => { + result.push('\n'); + consumed += 1; + } + Some('t') => { + result.push('\t'); + consumed += 1; + } + Some('\\') => { + result.push('\\'); + consumed += 1; + } + Some(c) => { + result.push('\\'); + result.push(c); + consumed += c.len_utf8(); + } + None => { + return Err(nom::Err::Error(nom::error::Error::new( + input, + nom::error::ErrorKind::Char, + ))); + } + } + } + Some(c) => { + result.push(c); + consumed += c.len_utf8(); + } + None => { + return Err(nom::Err::Error(nom::error::Error::new( + input, + nom::error::ErrorKind::Char, + ))); + } + } + } + } + + /// Parse a boolean: `true` or `false`. + pub fn boolean(input: &str) -> IResult<&str, bool> { + let (rest, word) = identifier(input)?; + match word { + "true" => Ok((rest, true)), + "false" => Ok((rest, false)), + _ => Err(nom::Err::Error(nom::error::Error::new( + input, + nom::error::ErrorKind::Tag, + ))), + } + } + + /// Parse a float: optional sign, optional integer part, `.`, fractional digits. + pub fn float_value(input: &str) -> IResult<&str, f64> { + let (rest, raw) = recognize(pair( + pair(opt(char('-')), take_while(|c: char| c.is_ascii_digit())), + pair(char('.'), take_while1(|c: char| c.is_ascii_digit())), + ))(input)?; + let val: f64 = raw.parse().map_err(|_| { + nom::Err::Error(nom::error::Error::new( + input, + nom::error::ErrorKind::Float, + )) + })?; + Ok((rest, val)) + } + + /// Parse an integer: optional sign, digits. Not followed by `.` (that's a float). + pub fn integer_value(input: &str) -> IResult<&str, i64> { + let (rest, raw) = recognize(pair( + opt(char('-')), + take_while1(|c: char| c.is_ascii_digit()), + ))(input)?; + if rest.starts_with('.') { + return Err(nom::Err::Error(nom::error::Error::new( + input, + nom::error::ErrorKind::Digit, + ))); + } + let val: i64 = raw.parse().map_err(|_| { + nom::Err::Error(nom::error::Error::new( + input, + nom::error::ErrorKind::Digit, + )) + })?; + Ok((rest, val)) + } + + /// Parse a duration: integer followed by unit suffix (ms, s, m, h, d). + pub fn duration_value(input: &str) -> IResult<&str, AstValue> { + let (rest, num) = recognize(pair( + opt(char('-')), + take_while1(|c: char| c.is_ascii_digit()), + ))(input)?; + let (rest, unit) = alt((tag("ms"), tag("s"), tag("m"), tag("h"), tag("d")))(rest)?; + if rest + .chars() + .next() + .is_some_and(|c| c.is_ascii_alphanumeric()) + { + return Err(nom::Err::Error(nom::error::Error::new( + input, + nom::error::ErrorKind::Tag, + ))); + } + Ok((rest, AstValue::Str(format!("{num}{unit}")))) + } + + /// Parse an AST value: duration, float, integer, boolean, quoted string, or bare identifier. + pub fn value(input: &str) -> IResult<&str, AstValue> { + let input = input.trim_start(); + alt(( + map(quoted_string, AstValue::Str), + duration_value, + map(float_value, AstValue::Float), + map(integer_value, AstValue::Int), + map(boolean, AstValue::Bool), + map(identifier, |s: &str| AstValue::Ident(s.to_string())), + ))(input) + } + + /// Parse the arrow operator `->` surrounded by optional whitespace. + pub fn arrow(input: &str) -> IResult<&str, &str> { + preceded(ws, tag("->"))(input) + } +} + +#[cfg(test)] +mod tests { + use super::combinators::*; + use super::*; + use crate::parser::ast::AstValue; + + #[test] + fn strip_line_comments() { + let input = "hello // this is a comment\nworld"; + assert_eq!(strip_comments(input), "hello \nworld"); + } + + #[test] + fn strip_block_comments() { + let input = "before /* inside */ after"; + assert_eq!(strip_comments(input), "before after"); + } + + #[test] + fn strip_block_comments_multiline() { + let input = "a /* line1\nline2 */ b"; + let result = strip_comments(input); + assert_eq!(result, "a \n b"); + } + + #[test] + fn strip_preserves_strings() { + let input = r#""hello // not a comment" rest"#; + assert_eq!(strip_comments(input), r#""hello // not a comment" rest"#); + } + + #[test] + fn strip_string_with_escapes() { + let input = r#""escaped \" quote" rest"#; + assert_eq!(strip_comments(input), r#""escaped \" quote" rest"#); + } + + #[test] + fn parse_identifier() { + assert_eq!(identifier("hello_world123 "), Ok((" ", "hello_world123"))); + assert_eq!(identifier("_private rest"), Ok((" rest", "_private"))); + assert!(identifier("123abc").is_err()); + } + + #[test] + fn parse_qualified_id() { + assert_eq!( + qualified_id("tool_hooks.pre rest"), + Ok((" rest", "tool_hooks.pre".into())) + ); + assert_eq!( + qualified_id("a.b.c rest"), + Ok((" rest", "a.b.c".into())) + ); + assert!(qualified_id("simple rest").is_err()); + } + + #[test] + fn parse_key_simple_and_qualified() { + assert_eq!(key("label rest"), Ok((" rest", "label".into()))); + assert_eq!( + key("tool_hooks.pre rest"), + Ok((" rest", "tool_hooks.pre".into())) + ); + } + + #[test] + fn parse_quoted_string() { + assert_eq!(quoted_string(r#""hello""#), Ok(("", "hello".into()))); + assert_eq!( + quoted_string(r#""line1\nline2""#), + Ok(("", "line1\nline2".into())) + ); + assert_eq!( + quoted_string(r#""tab\there""#), + Ok(("", "tab\there".into())) + ); + assert_eq!( + quoted_string(r#""escaped \" quote""#), + Ok(("", "escaped \" quote".into())) + ); + assert_eq!( + quoted_string(r#""back\\slash""#), + Ok(("", "back\\slash".into())) + ); + } + + #[test] + fn parse_boolean() { + assert_eq!(boolean("true rest"), Ok((" rest", true))); + assert_eq!(boolean("false rest"), Ok((" rest", false))); + assert!(boolean("yes").is_err()); + } + + #[test] + fn parse_integer() { + assert_eq!(integer_value("42 rest"), Ok((" rest", 42))); + assert_eq!(integer_value("-1 rest"), Ok((" rest", -1))); + assert_eq!(integer_value("0 rest"), Ok((" rest", 0))); + assert!(integer_value("42.5").is_err()); + } + + #[test] + fn parse_float() { + assert_eq!(float_value("3.14 rest"), Ok((" rest", 3.14))); + assert_eq!(float_value("0.5 rest"), Ok((" rest", 0.5))); + assert_eq!(float_value("-3.14 rest"), Ok((" rest", -3.14))); + assert_eq!(float_value(".5 rest"), Ok((" rest", 0.5))); + } + + #[test] + fn parse_duration() { + assert_eq!( + duration_value("250ms rest"), + Ok((" rest", AstValue::Str("250ms".into()))) + ); + assert_eq!( + duration_value("900s rest"), + Ok((" rest", AstValue::Str("900s".into()))) + ); + assert_eq!( + duration_value("15m rest"), + Ok((" rest", AstValue::Str("15m".into()))) + ); + assert_eq!( + duration_value("2h rest"), + Ok((" rest", AstValue::Str("2h".into()))) + ); + assert_eq!( + duration_value("1d rest"), + Ok((" rest", AstValue::Str("1d".into()))) + ); + } + + #[test] + fn parse_value_all_types() { + assert_eq!( + value(r#""hello""#), + Ok(("", AstValue::Str("hello".into()))) + ); + assert_eq!(value("250ms"), Ok(("", AstValue::Str("250ms".into())))); + assert_eq!(value("3.14"), Ok(("", AstValue::Float(3.14)))); + assert_eq!(value("42"), Ok(("", AstValue::Int(42)))); + assert_eq!(value("true"), Ok(("", AstValue::Bool(true)))); + assert_eq!(value("LR"), Ok(("", AstValue::Ident("LR".into())))); + } +} diff --git a/crates/attractor/src/parser/mod.rs b/crates/attractor/src/parser/mod.rs new file mode 100644 index 000000000..106bf341f --- /dev/null +++ b/crates/attractor/src/parser/mod.rs @@ -0,0 +1,147 @@ +pub mod ast; +pub mod grammar; +pub mod lexer; +pub mod semantic; + +use crate::error::AttractorError; +use crate::graph::types::Graph; + +/// Parse a DOT source string into a semantic `Graph`. +/// +/// Strips comments, parses the grammar, and performs semantic transformation +/// (expanding chained edges, applying defaults, flattening subgraphs). +/// +/// # Errors +/// +/// Returns an error if the input is not valid DOT syntax or contains +/// trailing content after the graph definition. +pub fn parse(input: &str) -> Result { + let stripped = lexer::strip_comments(input); + let (rest, dot_graph) = grammar::parse_dot_graph(&stripped) + .map_err(|e| AttractorError::Parse(format!("grammar error: {e}")))?; + + let remaining = rest.trim(); + if !remaining.is_empty() { + return Err(AttractorError::Parse(format!( + "unexpected trailing content: {:?}", + &remaining[..remaining.len().min(50)] + ))); + } + + semantic::ast_to_graph(&dot_graph) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parse_simple_linear() { + let input = r#"digraph Simple { + graph [goal="Run tests and report"] + rankdir=LR + + start [shape=Mdiamond, label="Start"] + exit [shape=Msquare, label="Exit"] + + run_tests [label="Run Tests", prompt="Run the test suite and report results"] + report [label="Report", prompt="Summarize the test results"] + + start -> run_tests -> report -> exit + }"#; + let graph = parse(input).unwrap(); + assert_eq!(graph.name, "Simple"); + assert_eq!(graph.goal(), "Run tests and report"); + assert_eq!(graph.nodes.len(), 4); + // start->run_tests, run_tests->report, report->exit + assert_eq!(graph.edges.len(), 3); + assert!(graph.find_start_node().is_some()); + assert!(graph.find_exit_node().is_some()); + } + + #[test] + fn parse_branching_with_conditions() { + let input = r#"digraph Branch { + graph [goal="Implement and validate a feature"] + rankdir=LR + node [shape=box, timeout="900s"] + + start [shape=Mdiamond, label="Start"] + exit [shape=Msquare, label="Exit"] + plan [label="Plan", prompt="Plan the implementation"] + implement [label="Implement", prompt="Implement the plan"] + validate [label="Validate", prompt="Run tests"] + gate [shape=diamond, label="Tests passing?"] + + start -> plan -> implement -> validate -> gate + gate -> exit [label="Yes", condition="outcome=success"] + gate -> implement [label="No", condition="outcome!=success"] + }"#; + let graph = parse(input).unwrap(); + assert_eq!(graph.name, "Branch"); + assert_eq!(graph.nodes.len(), 6); + // chain: 4 edges + 2 conditional = 6 + assert_eq!(graph.edges.len(), 6); + + // Check condition on gate -> exit edge + let gate_exit = graph + .edges + .iter() + .find(|e| e.from == "gate" && e.to == "exit") + .unwrap(); + assert_eq!(gate_exit.condition(), Some("outcome=success")); + } + + #[test] + fn parse_human_gate() { + let input = r#"digraph Review { + rankdir=LR + + start [shape=Mdiamond, label="Start"] + exit [shape=Msquare, label="Exit"] + + review_gate [ + shape=hexagon, + label="Review Changes", + type="wait.human" + ] + + start -> review_gate + review_gate -> ship_it [label="[A] Approve"] + review_gate -> fixes [label="[F] Fix"] + ship_it -> exit + fixes -> review_gate + }"#; + let graph = parse(input).unwrap(); + assert_eq!(graph.name, "Review"); + let gate = &graph.nodes["review_gate"]; + assert_eq!(gate.node_type(), Some("wait.human")); + assert_eq!(gate.shape(), "hexagon"); + } + + #[test] + fn parse_with_comments() { + let input = r#"// This is a comment + digraph Test { + /* block comment */ + start [shape=Mdiamond] // inline comment + exit [shape=Msquare] + start -> exit + }"#; + let graph = parse(input).unwrap(); + assert_eq!(graph.nodes.len(), 2); + } + + #[test] + fn parse_error_on_invalid_input() { + let result = parse("not a graph"); + assert!(result.is_err()); + } + + #[test] + fn parse_error_on_trailing_content() { + let input = "digraph A { } extra stuff"; + let result = parse(input); + assert!(result.is_err()); + } +} diff --git a/crates/attractor/src/parser/semantic.rs b/crates/attractor/src/parser/semantic.rs new file mode 100644 index 000000000..abf65d1e2 --- /dev/null +++ b/crates/attractor/src/parser/semantic.rs @@ -0,0 +1,516 @@ +use std::collections::HashMap; +use std::time::Duration; + +use crate::error::AttractorError; +use crate::graph::types::{AttrValue, Edge, Graph, Node}; +use crate::parser::ast::{AstValue, AttrBlock, DotGraph, Statement}; + +/// Convert an AST `AstValue` to a semantic `AttrValue`. +fn convert_value(ast_val: &AstValue) -> AttrValue { + match ast_val { + AstValue::Str(s) | AstValue::Ident(s) => { + if let Some(dur) = parse_duration_str(s) { + return AttrValue::Duration(dur); + } + AttrValue::String(s.clone()) + } + AstValue::Int(n) => AttrValue::Integer(*n), + AstValue::Float(f) => AttrValue::Float(*f), + AstValue::Bool(b) => AttrValue::Boolean(*b), + } +} + +fn parse_duration_str(s: &str) -> Option { + if s.ends_with("ms") { + let num = s.strip_suffix("ms")?.parse::().ok()?; + return Some(Duration::from_millis(num)); + } + let (num_str, multiplier) = if let Some(n) = s.strip_suffix('s') { + (n, 1_000u64) + } else if let Some(n) = s.strip_suffix('m') { + (n, 60_000u64) + } else if let Some(n) = s.strip_suffix('h') { + (n, 3_600_000u64) + } else if let Some(n) = s.strip_suffix('d') { + (n, 86_400_000u64) + } else { + return None; + }; + let num: u64 = num_str.parse().ok()?; + Some(Duration::from_millis(num * multiplier)) +} + +fn convert_attrs(block: &AttrBlock) -> HashMap { + block + .iter() + .map(|(k, v)| (k.clone(), convert_value(v))) + .collect() +} + +/// Derive a CSS class name from a subgraph label. +fn derive_class_from_label(label: &str) -> String { + label + .to_lowercase() + .chars() + .map(|c| if c == ' ' { '-' } else { c }) + .filter(|c| c.is_ascii_alphanumeric() || *c == '-') + .collect() +} + +struct SemanticState { + graph: Graph, + node_defaults: HashMap, + edge_defaults: HashMap, +} + +impl SemanticState { + fn new(name: String) -> Self { + Self { + graph: Graph::new(name), + node_defaults: HashMap::new(), + edge_defaults: HashMap::new(), + } + } + + fn ensure_node(&mut self, id: &str) { + if !self.graph.nodes.contains_key(id) { + let mut node = Node::new(id); + for (k, v) in &self.node_defaults { + node.attrs.insert(k.clone(), v.clone()); + } + self.graph.nodes.insert(id.to_string(), node); + } + } + + fn add_class_to_node(node: &mut Node, cls: &str) { + let cls_string = cls.to_string(); + if !node.classes.contains(&cls_string) { + node.classes.push(cls_string); + } + } + + fn process_node(&mut self, node_stmt: &crate::parser::ast::NodeStmt, subgraph_class: Option<&str>) { + self.ensure_node(&node_stmt.id); + let node = self.graph.nodes.get_mut(&node_stmt.id).expect("just ensured"); + if let Some(attrs) = &node_stmt.attrs { + for (k, v) in attrs { + node.attrs.insert(k.clone(), convert_value(v)); + } + } + if let Some(cls) = subgraph_class { + Self::add_class_to_node(node, cls); + } + // Parse explicit class attr into classes vec + let class_str = node + .attrs + .get("class") + .and_then(AttrValue::as_str) + .map(String::from); + if let Some(class_str) = class_str { + let node = self.graph.nodes.get_mut(&node_stmt.id).expect("just ensured"); + for cls in class_str.split(',') { + let cls = cls.trim().to_string(); + if !cls.is_empty() && !node.classes.contains(&cls) { + node.classes.push(cls); + } + } + } + } + + fn process_edge(&mut self, edge_stmt: &crate::parser::ast::EdgeStmt, subgraph_class: Option<&str>) { + for id in &edge_stmt.nodes { + self.ensure_node(id); + if let Some(cls) = subgraph_class { + let node = self.graph.nodes.get_mut(id).expect("just ensured"); + Self::add_class_to_node(node, cls); + } + } + let edge_attrs = edge_stmt + .attrs + .as_ref() + .map_or_else(HashMap::new, convert_attrs); + for pair in edge_stmt.nodes.windows(2) { + let mut edge = Edge::new(&pair[0], &pair[1]); + for (k, v) in &self.edge_defaults { + edge.attrs.insert(k.clone(), v.clone()); + } + for (k, v) in &edge_attrs { + edge.attrs.insert(k.clone(), v.clone()); + } + self.graph.edges.push(edge); + } + } + + #[allow(clippy::too_many_lines)] + fn process_statements( + &mut self, + statements: &[Statement], + subgraph_class: Option<&str>, + scoped_node_defaults: &HashMap, + scoped_edge_defaults: &HashMap, + ) { + let saved_node_defaults = self.node_defaults.clone(); + let saved_edge_defaults = self.edge_defaults.clone(); + for (k, v) in scoped_node_defaults { + self.node_defaults.insert(k.clone(), v.clone()); + } + for (k, v) in scoped_edge_defaults { + self.edge_defaults.insert(k.clone(), v.clone()); + } + + for stmt in statements { + match stmt { + Statement::GraphAttr(attrs) => { + for (k, v) in attrs { + self.graph.attrs.insert(k.clone(), convert_value(v)); + } + } + Statement::NodeDefaults(attrs) => { + for (k, v) in convert_attrs(attrs) { + self.node_defaults.insert(k, v); + } + } + Statement::EdgeDefaults(attrs) => { + for (k, v) in convert_attrs(attrs) { + self.edge_defaults.insert(k, v); + } + } + Statement::GraphAttrDecl(key, val) => { + self.graph.attrs.insert(key.clone(), convert_value(val)); + } + Statement::Node(node_stmt) => { + self.process_node(node_stmt, subgraph_class); + } + Statement::Edge(edge_stmt) => { + self.process_edge(edge_stmt, subgraph_class); + } + Statement::Subgraph(sub) => { + let sub_class = sub.statements.iter().find_map(|s| match s { + Statement::GraphAttrDecl(k, AstValue::Str(s) | AstValue::Ident(s)) + if k == "label" => + { + Some(derive_class_from_label(s)) + } + Statement::GraphAttr(attrs) => attrs.iter().find_map(|(k, v)| { + if k == "label" { + match v { + AstValue::Str(s) | AstValue::Ident(s) => { + Some(derive_class_from_label(s)) + } + _ => None, + } + } else { + None + } + }), + _ => None, + }); + + let mut sub_node_defaults = HashMap::new(); + let mut sub_edge_defaults = HashMap::new(); + for s in &sub.statements { + match s { + Statement::NodeDefaults(attrs) => { + sub_node_defaults.extend(convert_attrs(attrs)); + } + Statement::EdgeDefaults(attrs) => { + sub_edge_defaults.extend(convert_attrs(attrs)); + } + _ => {} + } + } + + self.process_statements( + &sub.statements, + sub_class.as_deref(), + &sub_node_defaults, + &sub_edge_defaults, + ); + } + } + } + + self.node_defaults = saved_node_defaults; + self.edge_defaults = saved_edge_defaults; + } +} + +/// Convert a parsed `DotGraph` AST into a semantic `Graph`. +/// +/// # Errors +/// +/// Returns an error if the AST cannot be converted to a valid graph. +pub fn ast_to_graph(dot: &DotGraph) -> Result { + let mut state = SemanticState::new(dot.name.clone()); + let empty = HashMap::new(); + state.process_statements(&dot.statements, None, &empty, &empty); + Ok(state.graph) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn convert_ast_str_to_string() { + assert_eq!( + convert_value(&AstValue::Str("hello".into())), + AttrValue::String("hello".into()) + ); + } + + #[test] + fn convert_ast_duration_str() { + assert_eq!( + convert_value(&AstValue::Str("900s".into())), + AttrValue::Duration(Duration::from_secs(900)) + ); + assert_eq!( + convert_value(&AstValue::Str("250ms".into())), + AttrValue::Duration(Duration::from_millis(250)) + ); + assert_eq!( + convert_value(&AstValue::Str("15m".into())), + AttrValue::Duration(Duration::from_secs(900)) + ); + assert_eq!( + convert_value(&AstValue::Str("2h".into())), + AttrValue::Duration(Duration::from_secs(7200)) + ); + assert_eq!( + convert_value(&AstValue::Str("1d".into())), + AttrValue::Duration(Duration::from_secs(86400)) + ); + } + + #[test] + fn convert_ast_int() { + assert_eq!(convert_value(&AstValue::Int(42)), AttrValue::Integer(42)); + } + + #[test] + fn convert_ast_bool() { + assert_eq!( + convert_value(&AstValue::Bool(true)), + AttrValue::Boolean(true) + ); + } + + #[test] + fn convert_ast_float() { + assert_eq!( + convert_value(&AstValue::Float(3.14)), + AttrValue::Float(3.14) + ); + } + + #[test] + fn convert_ast_ident() { + assert_eq!( + convert_value(&AstValue::Ident("LR".into())), + AttrValue::String("LR".into()) + ); + } + + #[test] + fn derive_class_simple() { + assert_eq!(derive_class_from_label("Loop A"), "loop-a"); + assert_eq!(derive_class_from_label("Code Review"), "code-review"); + assert_eq!(derive_class_from_label("Hello World!!!"), "hello-world"); + } + + #[test] + fn ast_to_graph_simple_linear() { + let dot = DotGraph { + name: "Simple".into(), + statements: vec![ + Statement::GraphAttr(vec![("goal".into(), AstValue::Str("Run tests".into()))]), + Statement::GraphAttrDecl("rankdir".into(), AstValue::Ident("LR".into())), + Statement::Node(crate::parser::ast::NodeStmt { + id: "start".into(), + attrs: Some(vec![ + ("shape".into(), AstValue::Ident("Mdiamond".into())), + ("label".into(), AstValue::Str("Start".into())), + ]), + }), + Statement::Node(crate::parser::ast::NodeStmt { + id: "exit".into(), + attrs: Some(vec![ + ("shape".into(), AstValue::Ident("Msquare".into())), + ("label".into(), AstValue::Str("Exit".into())), + ]), + }), + Statement::Node(crate::parser::ast::NodeStmt { + id: "run_tests".into(), + attrs: Some(vec![("label".into(), AstValue::Str("Run Tests".into()))]), + }), + Statement::Edge(crate::parser::ast::EdgeStmt { + nodes: vec!["start".into(), "run_tests".into(), "exit".into()], + attrs: None, + }), + ], + }; + + let graph = ast_to_graph(&dot).unwrap(); + assert_eq!(graph.name, "Simple"); + assert_eq!(graph.goal(), "Run tests"); + assert_eq!(graph.nodes.len(), 3); + assert_eq!(graph.edges.len(), 2); + assert_eq!(graph.edges[0].from, "start"); + assert_eq!(graph.edges[0].to, "run_tests"); + assert_eq!(graph.edges[1].from, "run_tests"); + assert_eq!(graph.edges[1].to, "exit"); + } + + #[test] + fn ast_to_graph_node_defaults_applied() { + let dot = DotGraph { + name: "Defaults".into(), + statements: vec![ + Statement::NodeDefaults(vec![ + ("shape".into(), AstValue::Ident("box".into())), + ("timeout".into(), AstValue::Str("900s".into())), + ]), + Statement::Node(crate::parser::ast::NodeStmt { + id: "plan".into(), + attrs: Some(vec![("label".into(), AstValue::Str("Plan".into()))]), + }), + Statement::Node(crate::parser::ast::NodeStmt { + id: "implement".into(), + attrs: Some(vec![ + ("label".into(), AstValue::Str("Implement".into())), + ("timeout".into(), AstValue::Str("1800s".into())), + ]), + }), + ], + }; + + let graph = ast_to_graph(&dot).unwrap(); + let plan = &graph.nodes["plan"]; + assert_eq!( + plan.attrs.get("shape").and_then(AttrValue::as_str), + Some("box") + ); + assert_eq!( + plan.attrs.get("timeout").and_then(AttrValue::as_duration), + Some(Duration::from_secs(900)) + ); + + let implement = &graph.nodes["implement"]; + assert_eq!( + implement + .attrs + .get("timeout") + .and_then(AttrValue::as_duration), + Some(Duration::from_secs(1800)) + ); + } + + #[test] + fn ast_to_graph_subgraph_class_derivation() { + let dot = DotGraph { + name: "SubgraphTest".into(), + statements: vec![Statement::Subgraph(crate::parser::ast::SubgraphStmt { + name: Some("cluster_loop".into()), + statements: vec![ + Statement::GraphAttrDecl("label".into(), AstValue::Str("Loop A".into())), + Statement::Node(crate::parser::ast::NodeStmt { + id: "plan".into(), + attrs: None, + }), + ], + })], + }; + + let graph = ast_to_graph(&dot).unwrap(); + let plan = &graph.nodes["plan"]; + assert!(plan.classes.contains(&"loop-a".to_string())); + } + + #[test] + fn ast_to_graph_subgraph_class_from_graph_attr_block() { + let dot = DotGraph { + name: "SubgraphAttrBlock".into(), + statements: vec![Statement::Subgraph(crate::parser::ast::SubgraphStmt { + name: Some("cluster_review".into()), + statements: vec![ + Statement::GraphAttr(vec![ + ("label".into(), AstValue::Str("Code Review".into())), + ]), + Statement::Node(crate::parser::ast::NodeStmt { + id: "reviewer".into(), + attrs: None, + }), + ], + })], + }; + + let graph = ast_to_graph(&dot).unwrap(); + let reviewer = &graph.nodes["reviewer"]; + assert!(reviewer.classes.contains(&"code-review".to_string())); + } + + #[test] + fn ast_to_graph_edge_defaults_applied() { + let dot = DotGraph { + name: "EdgeDefaults".into(), + statements: vec![ + Statement::EdgeDefaults(vec![("weight".into(), AstValue::Int(5))]), + Statement::Edge(crate::parser::ast::EdgeStmt { + nodes: vec!["a".into(), "b".into()], + attrs: None, + }), + ], + }; + + let graph = ast_to_graph(&dot).unwrap(); + assert_eq!(graph.edges[0].weight(), 5); + } + + #[test] + fn ast_to_graph_chained_edges_with_attrs() { + let dot = DotGraph { + name: "Chained".into(), + statements: vec![Statement::Edge(crate::parser::ast::EdgeStmt { + nodes: vec!["a".into(), "b".into(), "c".into()], + attrs: Some(vec![("label".into(), AstValue::Str("next".into()))]), + })], + }; + + let graph = ast_to_graph(&dot).unwrap(); + assert_eq!(graph.edges.len(), 2); + assert_eq!(graph.edges[0].label(), Some("next")); + assert_eq!(graph.edges[1].label(), Some("next")); + } + + #[test] + fn ast_to_graph_class_attr_parsed() { + let dot = DotGraph { + name: "ClassTest".into(), + statements: vec![Statement::Node(crate::parser::ast::NodeStmt { + id: "review".into(), + attrs: Some(vec![("class".into(), AstValue::Str("code,critical".into()))]), + })], + }; + + let graph = ast_to_graph(&dot).unwrap(); + let review = &graph.nodes["review"]; + assert!(review.classes.contains(&"code".to_string())); + assert!(review.classes.contains(&"critical".to_string())); + } + + #[test] + fn ast_to_graph_implicit_nodes_from_edges() { + let dot = DotGraph { + name: "Implicit".into(), + statements: vec![Statement::Edge(crate::parser::ast::EdgeStmt { + nodes: vec!["a".into(), "b".into()], + attrs: None, + })], + }; + + let graph = ast_to_graph(&dot).unwrap(); + assert!(graph.nodes.contains_key("a")); + assert!(graph.nodes.contains_key("b")); + } +} diff --git a/crates/attractor/src/pipeline.rs b/crates/attractor/src/pipeline.rs new file mode 100644 index 000000000..086ed04be --- /dev/null +++ b/crates/attractor/src/pipeline.rs @@ -0,0 +1,177 @@ +use crate::error::AttractorError; +use crate::graph::Graph; +use crate::transform::{ + PreambleTransform, StylesheetApplicationTransform, Transform, VariableExpansionTransform, +}; +use crate::validation::{self, Diagnostic}; + +/// Builder for configuring and executing a pipeline preparation. +/// Collects custom transforms that run after the built-in ones. +pub struct PipelineBuilder { + transforms: Vec>, +} + +impl PipelineBuilder { + #[must_use] + pub fn new() -> Self { + Self { + transforms: Vec::new(), + } + } + + /// Register a custom transform. Custom transforms run after built-in transforms, + /// in registration order. + pub fn register_transform(&mut self, transform: Box) { + self.transforms.push(transform); + } + + /// Prepare a pipeline: parse DOT, apply built-in and custom transforms, validate. + /// + /// # Errors + /// + /// Returns an error if parsing or validation fails. + pub fn prepare(&self, dot_source: &str) -> Result<(Graph, Vec), AttractorError> { + let mut graph = crate::parser::parse(dot_source)?; + + // Built-in transforms + VariableExpansionTransform.apply(&mut graph); + StylesheetApplicationTransform.apply(&mut graph); + PreambleTransform.apply(&mut graph); + + // Custom transforms + for transform in &self.transforms { + transform.apply(&mut graph); + } + + let diagnostics = validation::validate(&graph, &[]); + Ok((graph, diagnostics)) + } +} + +impl Default for PipelineBuilder { + fn default() -> Self { + Self::new() + } +} + +/// Convenience function: parse DOT, apply built-in transforms, validate, return graph. +/// +/// # Errors +/// +/// Returns an error if parsing fails or if validation produces Error-severity diagnostics. +pub fn prepare_pipeline(dot_source: &str) -> Result { + let builder = PipelineBuilder::new(); + let (graph, diagnostics) = builder.prepare(dot_source)?; + + let errors: Vec<&Diagnostic> = diagnostics + .iter() + .filter(|d| d.severity == validation::Severity::Error) + .collect(); + if !errors.is_empty() { + let messages: Vec = errors.iter().map(|d| d.message.clone()).collect(); + return Err(AttractorError::Validation(messages.join("; "))); + } + + Ok(graph) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::AttrValue; + + const MINIMAL_DOT: &str = r#"digraph Test { + graph [goal="Build feature"] + start [shape=Mdiamond] + exit [shape=Msquare] + start -> exit + }"#; + + #[test] + fn prepare_pipeline_minimal() { + let graph = prepare_pipeline(MINIMAL_DOT).unwrap(); + assert_eq!(graph.name, "Test"); + assert!(graph.find_start_node().is_some()); + assert!(graph.find_exit_node().is_some()); + } + + #[test] + fn prepare_pipeline_applies_variable_expansion() { + let dot = r#"digraph Test { + graph [goal="Fix bugs"] + start [shape=Mdiamond] + work [prompt="Goal: $goal"] + exit [shape=Msquare] + start -> work -> exit + }"#; + let graph = prepare_pipeline(dot).unwrap(); + let prompt = graph.nodes["work"] + .attrs + .get("prompt") + .and_then(AttrValue::as_str) + .unwrap(); + assert_eq!(prompt, "Goal: Fix bugs"); + } + + #[test] + fn prepare_pipeline_applies_stylesheet() { + let dot = r#"digraph Test { + graph [goal="Test", model_stylesheet="* { llm_model: sonnet; }"] + start [shape=Mdiamond] + work [label="Work"] + exit [shape=Msquare] + start -> work -> exit + }"#; + let graph = prepare_pipeline(dot).unwrap(); + assert_eq!( + graph.nodes["work"].attrs.get("llm_model"), + Some(&AttrValue::String("sonnet".into())) + ); + } + + #[test] + fn prepare_pipeline_returns_error_on_invalid_dot() { + let result = prepare_pipeline("not a graph"); + assert!(result.is_err()); + } + + #[test] + fn prepare_pipeline_returns_error_on_validation_failure() { + let dot = r#"digraph Test { + graph [goal="Test"] + work [label="Work"] + }"#; + let result = prepare_pipeline(dot); + assert!(result.is_err()); + } + + #[test] + fn pipeline_builder_custom_transform() { + struct TagTransform; + impl Transform for TagTransform { + fn apply(&self, graph: &mut crate::graph::Graph) { + for node in graph.nodes.values_mut() { + node.attrs.insert( + "tagged".to_string(), + AttrValue::Boolean(true), + ); + } + } + } + + let mut builder = PipelineBuilder::new(); + builder.register_transform(Box::new(TagTransform)); + let (graph, _) = builder.prepare(MINIMAL_DOT).unwrap(); + assert_eq!( + graph.nodes["start"].attrs.get("tagged"), + Some(&AttrValue::Boolean(true)) + ); + } + + #[test] + fn pipeline_builder_default() { + let builder = PipelineBuilder::default(); + let (graph, _) = builder.prepare(MINIMAL_DOT).unwrap(); + assert_eq!(graph.name, "Test"); + } +} diff --git a/crates/attractor/src/stylesheet.rs b/crates/attractor/src/stylesheet.rs new file mode 100644 index 000000000..f3f547f74 --- /dev/null +++ b/crates/attractor/src/stylesheet.rs @@ -0,0 +1,518 @@ +use crate::error::AttractorError; +use crate::graph::types::{AttrValue, Graph}; + +/// A parsed stylesheet selector. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum Selector { + /// `*` -- matches all nodes, specificity 0. + Universal, + /// bare word like `box` -- matches nodes whose shape equals that word, specificity 1. + Shape(String), + /// `.classname` -- matches nodes with that class, specificity 2. + Class(String), + /// `#nodeid` -- matches a specific node, specificity 3. + Id(String), +} + +impl Selector { + #[must_use] + pub const fn specificity(&self) -> u8 { + match self { + Self::Universal => 0, + Self::Shape(_) => 1, + Self::Class(_) => 2, + Self::Id(_) => 3, + } + } +} + +/// A single CSS-like declaration: `property: value`. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Declaration { + pub property: String, + pub value: String, +} + +/// A stylesheet rule: selector + declarations. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Rule { + pub selector: Selector, + pub declarations: Vec, +} + +/// A parsed stylesheet containing multiple rules. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Stylesheet { + pub rules: Vec, +} + +/// Parse a stylesheet string into a `Stylesheet`. +/// +/// # Errors +/// +/// Returns an error if the input contains invalid stylesheet syntax. +pub fn parse_stylesheet(input: &str) -> Result { + let input = input.trim(); + if input.is_empty() { + return Ok(Stylesheet { rules: Vec::new() }); + } + + let mut rules = Vec::new(); + let mut remaining = input; + + while !remaining.trim().is_empty() { + remaining = remaining.trim(); + + let selector = parse_selector(&mut remaining)?; + if !remaining.starts_with('{') { + return Err(AttractorError::Stylesheet(format!( + "expected '{{' after selector, got: {:?}", + &remaining[..remaining.len().min(20)] + ))); + } + remaining = remaining[1..].trim(); + + let declarations = parse_declarations(&mut remaining)?; + remaining = remaining[1..].trim(); // skip '}' + + rules.push(Rule { + selector, + declarations, + }); + } + + Ok(Stylesheet { rules }) +} + +fn parse_selector(remaining: &mut &str) -> Result { + if remaining.starts_with('*') { + *remaining = remaining[1..].trim(); + Ok(Selector::Universal) + } else if remaining.starts_with('#') { + *remaining = remaining[1..].trim(); + let end = remaining + .find(|c: char| !c.is_ascii_alphanumeric() && c != '_' && c != '-') + .unwrap_or(remaining.len()); + if end == 0 { + return Err(AttractorError::Stylesheet( + "expected identifier after '#'".into(), + )); + } + let id = remaining[..end].to_string(); + *remaining = remaining[end..].trim(); + Ok(Selector::Id(id)) + } else if remaining.starts_with('.') { + *remaining = remaining[1..].trim(); + let end = remaining + .find(|c: char| !c.is_ascii_lowercase() && !c.is_ascii_digit() && c != '-') + .unwrap_or(remaining.len()); + if end == 0 { + return Err(AttractorError::Stylesheet( + "expected class name after '.'".into(), + )); + } + let class = remaining[..end].to_string(); + *remaining = remaining[end..].trim(); + Ok(Selector::Class(class)) + } else { + // Bare word: shape selector + let end = remaining + .find(|c: char| !c.is_ascii_alphanumeric() && c != '_' && c != '-') + .unwrap_or(remaining.len()); + if end == 0 { + return Err(AttractorError::Stylesheet(format!( + "expected selector ('*', '#id', '.class', or shape name), got: {:?}", + &remaining[..remaining.len().min(20)] + ))); + } + let shape = remaining[..end].to_string(); + *remaining = remaining[end..].trim(); + Ok(Selector::Shape(shape)) + } +} + +fn parse_declarations(remaining: &mut &str) -> Result, AttractorError> { + let mut declarations = Vec::new(); + while !remaining.starts_with('}') { + if remaining.is_empty() { + return Err(AttractorError::Stylesheet( + "unexpected end of stylesheet, expected '}'".into(), + )); + } + if remaining.starts_with(';') { + *remaining = remaining[1..].trim(); + continue; + } + + let prop_end = remaining + .find(|c: char| c == ':' || c.is_whitespace()) + .unwrap_or(remaining.len()); + let property = remaining[..prop_end].to_string(); + *remaining = remaining[prop_end..].trim(); + + if !remaining.starts_with(':') { + return Err(AttractorError::Stylesheet(format!( + "expected ':' after property name '{property}'" + ))); + } + *remaining = remaining[1..].trim(); + + let val_end = remaining + .find([';', '}']) + .unwrap_or(remaining.len()); + let value = remaining[..val_end].trim().to_string(); + *remaining = remaining[val_end..].trim(); + + if value.is_empty() { + return Err(AttractorError::Stylesheet(format!( + "empty value for property '{property}'" + ))); + } + + declarations.push(Declaration { property, value }); + + if remaining.starts_with(';') { + *remaining = remaining[1..].trim(); + } + } + Ok(declarations) +} + +/// Recognized stylesheet properties. +const STYLESHEET_PROPERTIES: &[&str] = &["llm_model", "llm_provider", "reasoning_effort"]; + +/// Apply a stylesheet to a graph. Rules are applied by specificity order; +/// higher specificity wins. Explicit node attributes are never overridden. +/// +/// # Panics +/// +/// Panics if the internal node map is inconsistent (should not happen). +pub fn apply_stylesheet(stylesheet: &Stylesheet, graph: &mut Graph) { + let mut sorted_rules: Vec<&Rule> = stylesheet.rules.iter().collect(); + sorted_rules.sort_by_key(|r| r.selector.specificity()); + + let node_ids: Vec = graph.nodes.keys().cloned().collect(); + + for node_id in &node_ids { + let mut applied: std::collections::HashMap = + std::collections::HashMap::new(); + + for rule in &sorted_rules { + let node = &graph.nodes[node_id.as_str()]; + let matches = match &rule.selector { + Selector::Universal => true, + Selector::Shape(shape) => node.shape() == shape.as_str(), + Selector::Class(cls) => node.classes.contains(cls), + Selector::Id(id) => node_id == id, + }; + + if matches { + for decl in &rule.declarations { + if STYLESHEET_PROPERTIES.contains(&decl.property.as_str()) { + let spec = rule.selector.specificity(); + match applied.get(&decl.property) { + Some((_, existing_spec)) if spec < *existing_spec => {} + _ => { + applied.insert( + decl.property.clone(), + (decl.value.clone(), spec), + ); + } + } + } + } + } + } + + let node = graph + .nodes + .get_mut(node_id.as_str()) + .expect("node must exist"); + for (prop, (val, _)) in &applied { + if !node.attrs.contains_key(prop) { + node.attrs + .insert(prop.clone(), AttrValue::String(val.clone())); + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::types::Node; + + #[test] + fn parse_empty_stylesheet() { + let ss = parse_stylesheet("").unwrap(); + assert!(ss.rules.is_empty()); + } + + #[test] + fn parse_universal_rule() { + let ss = + parse_stylesheet("* { llm_model: claude-sonnet-4-5; llm_provider: anthropic; }") + .unwrap(); + assert_eq!(ss.rules.len(), 1); + assert_eq!(ss.rules[0].selector, Selector::Universal); + assert_eq!(ss.rules[0].declarations.len(), 2); + assert_eq!(ss.rules[0].declarations[0].property, "llm_model"); + assert_eq!(ss.rules[0].declarations[0].value, "claude-sonnet-4-5"); + } + + #[test] + fn parse_class_rule() { + let ss = parse_stylesheet(".code { llm_model: claude-opus-4-6; }").unwrap(); + assert_eq!(ss.rules[0].selector, Selector::Class("code".into())); + } + + #[test] + fn parse_id_rule() { + let ss = parse_stylesheet( + "#critical_review { llm_model: gpt-5.2; reasoning_effort: high; }", + ) + .unwrap(); + assert_eq!( + ss.rules[0].selector, + Selector::Id("critical_review".into()) + ); + assert_eq!(ss.rules[0].declarations.len(), 2); + } + + #[test] + fn parse_multiple_rules() { + let input = r#" + * { llm_model: claude-sonnet-4-5; llm_provider: anthropic; } + .code { llm_model: claude-opus-4-6; llm_provider: anthropic; } + #critical_review { llm_model: gpt-5.2; llm_provider: openai; reasoning_effort: high; } + "#; + let ss = parse_stylesheet(input).unwrap(); + assert_eq!(ss.rules.len(), 3); + } + + #[test] + fn parse_error_missing_brace() { + let result = parse_stylesheet("* llm_model: test; }"); + assert!(result.is_err()); + } + + #[test] + fn parse_error_missing_selector() { + let result = parse_stylesheet("{ llm_model: test; }"); + assert!(result.is_err()); + } + + #[test] + fn apply_universal_to_all_nodes() { + let ss = parse_stylesheet("* { llm_model: sonnet; }").unwrap(); + let mut graph = Graph::new("test"); + graph.nodes.insert("a".into(), Node::new("a")); + graph.nodes.insert("b".into(), Node::new("b")); + apply_stylesheet(&ss, &mut graph); + + assert_eq!( + graph.nodes["a"].attrs.get("llm_model"), + Some(&AttrValue::String("sonnet".into())) + ); + assert_eq!( + graph.nodes["b"].attrs.get("llm_model"), + Some(&AttrValue::String("sonnet".into())) + ); + } + + #[test] + fn apply_class_overrides_universal() { + let ss = + parse_stylesheet("* { llm_model: sonnet; } .code { llm_model: opus; }").unwrap(); + let mut graph = Graph::new("test"); + + let mut code_node = Node::new("impl"); + code_node.classes.push("code".into()); + graph.nodes.insert("impl".into(), code_node); + + let plain_node = Node::new("plan"); + graph.nodes.insert("plan".into(), plain_node); + + apply_stylesheet(&ss, &mut graph); + + assert_eq!( + graph.nodes["impl"].attrs.get("llm_model"), + Some(&AttrValue::String("opus".into())) + ); + assert_eq!( + graph.nodes["plan"].attrs.get("llm_model"), + Some(&AttrValue::String("sonnet".into())) + ); + } + + #[test] + fn apply_id_overrides_class() { + let ss = parse_stylesheet( + ".code { llm_model: opus; } #special { llm_model: gpt; }", + ) + .unwrap(); + let mut graph = Graph::new("test"); + + let mut node = Node::new("special"); + node.classes.push("code".into()); + graph.nodes.insert("special".into(), node); + + apply_stylesheet(&ss, &mut graph); + + assert_eq!( + graph.nodes["special"].attrs.get("llm_model"), + Some(&AttrValue::String("gpt".into())) + ); + } + + #[test] + fn explicit_attrs_not_overridden() { + let ss = parse_stylesheet("* { llm_model: sonnet; }").unwrap(); + let mut graph = Graph::new("test"); + + let mut node = Node::new("a"); + node.attrs + .insert("llm_model".into(), AttrValue::String("explicit".into())); + graph.nodes.insert("a".into(), node); + + apply_stylesheet(&ss, &mut graph); + + assert_eq!( + graph.nodes["a"].attrs.get("llm_model"), + Some(&AttrValue::String("explicit".into())) + ); + } + + #[test] + fn selector_specificity_values() { + assert_eq!(Selector::Universal.specificity(), 0); + assert_eq!(Selector::Shape("box".into()).specificity(), 1); + assert_eq!(Selector::Class("x".into()).specificity(), 2); + assert_eq!(Selector::Id("x".into()).specificity(), 3); + } + + #[test] + fn spec_section_86_example() { + let input = r#" + * { llm_model: claude-sonnet-4-5; llm_provider: anthropic; } + .code { llm_model: claude-opus-4-6; llm_provider: anthropic; } + #critical_review { llm_model: gpt-5.2; llm_provider: openai; reasoning_effort: high; } + "#; + let ss = parse_stylesheet(input).unwrap(); + let mut graph = Graph::new("test"); + + let mut plan = Node::new("plan"); + plan.classes.push("planning".into()); + graph.nodes.insert("plan".into(), plan); + + let mut implement = Node::new("implement"); + implement.classes.push("code".into()); + graph.nodes.insert("implement".into(), implement); + + let mut review = Node::new("critical_review"); + review.classes.push("code".into()); + graph.nodes.insert("critical_review".into(), review); + + apply_stylesheet(&ss, &mut graph); + + assert_eq!( + graph.nodes["plan"].attrs.get("llm_model"), + Some(&AttrValue::String("claude-sonnet-4-5".into())) + ); + + assert_eq!( + graph.nodes["implement"].attrs.get("llm_model"), + Some(&AttrValue::String("claude-opus-4-6".into())) + ); + + assert_eq!( + graph.nodes["critical_review"].attrs.get("llm_model"), + Some(&AttrValue::String("gpt-5.2".into())) + ); + assert_eq!( + graph.nodes["critical_review"].attrs.get("llm_provider"), + Some(&AttrValue::String("openai".into())) + ); + assert_eq!( + graph.nodes["critical_review"] + .attrs + .get("reasoning_effort"), + Some(&AttrValue::String("high".into())) + ); + } + + #[test] + fn parse_shape_selector() { + let ss = parse_stylesheet("box { llm_model: opus; }").unwrap(); + assert_eq!(ss.rules.len(), 1); + assert_eq!(ss.rules[0].selector, Selector::Shape("box".into())); + } + + #[test] + fn apply_shape_selector_matches_by_shape() { + let ss = parse_stylesheet("box { llm_model: opus; }").unwrap(); + let mut graph = Graph::new("test"); + + // Default shape is "box" + let node_a = Node::new("a"); + graph.nodes.insert("a".into(), node_a); + + let mut node_b = Node::new("b"); + node_b + .attrs + .insert("shape".into(), AttrValue::String("diamond".into())); + graph.nodes.insert("b".into(), node_b); + + apply_stylesheet(&ss, &mut graph); + + assert_eq!( + graph.nodes["a"].attrs.get("llm_model"), + Some(&AttrValue::String("opus".into())) + ); + // diamond shape should not match "box" selector + assert_eq!(graph.nodes["b"].attrs.get("llm_model"), None); + } + + #[test] + fn shape_selector_specificity_between_universal_and_class() { + let ss = parse_stylesheet( + "* { llm_model: sonnet; } box { llm_model: opus; } .special { llm_model: gpt; }", + ) + .unwrap(); + let mut graph = Graph::new("test"); + + // Node with default shape "box" and class "special" + let mut node_a = Node::new("a"); + node_a.classes.push("special".into()); + graph.nodes.insert("a".into(), node_a); + + // Node with default shape "box" and no class + let node_b = Node::new("b"); + graph.nodes.insert("b".into(), node_b); + + // Node with shape "diamond" and no class + let mut node_c = Node::new("c"); + node_c + .attrs + .insert("shape".into(), AttrValue::String("diamond".into())); + graph.nodes.insert("c".into(), node_c); + + apply_stylesheet(&ss, &mut graph); + + // .special (specificity 2) overrides box (specificity 1) + assert_eq!( + graph.nodes["a"].attrs.get("llm_model"), + Some(&AttrValue::String("gpt".into())) + ); + // box (specificity 1) overrides * (specificity 0) + assert_eq!( + graph.nodes["b"].attrs.get("llm_model"), + Some(&AttrValue::String("opus".into())) + ); + // diamond doesn't match "box", so gets universal + assert_eq!( + graph.nodes["c"].attrs.get("llm_model"), + Some(&AttrValue::String("sonnet".into())) + ); + } +} diff --git a/crates/attractor/src/transform.rs b/crates/attractor/src/transform.rs new file mode 100644 index 000000000..1e9bcf564 --- /dev/null +++ b/crates/attractor/src/transform.rs @@ -0,0 +1,254 @@ +use crate::graph::{AttrValue, Graph}; +use crate::stylesheet::{apply_stylesheet, parse_stylesheet}; + +/// A transform that modifies the pipeline graph after parsing and before validation. +pub trait Transform { + fn apply(&self, graph: &mut Graph); +} + +/// Expands `$goal` in node `prompt` attributes to the graph-level `goal` value. +pub struct VariableExpansionTransform; + +impl Transform for VariableExpansionTransform { + fn apply(&self, graph: &mut Graph) { + let goal = graph.goal().to_string(); + for node in graph.nodes.values_mut() { + if let Some(AttrValue::String(prompt)) = node.attrs.get("prompt") { + if prompt.contains("$goal") { + let expanded = prompt.replace("$goal", &goal); + node.attrs + .insert("prompt".to_string(), AttrValue::String(expanded)); + } + } + } + } +} + +/// For nodes whose fidelity is not "full", prepend a context mode preamble to the prompt. +pub struct PreambleTransform; + +impl Transform for PreambleTransform { + fn apply(&self, graph: &mut Graph) { + let default_fidelity = graph.default_fidelity().unwrap_or("full").to_string(); + for node in graph.nodes.values_mut() { + let fidelity = node + .fidelity() + .unwrap_or(&default_fidelity) + .to_string(); + if fidelity == "full" { + continue; + } + let preamble = format!("[Context mode: {fidelity}]\n"); + if let Some(AttrValue::String(prompt)) = node.attrs.get("prompt") { + let new_prompt = format!("{preamble}{prompt}"); + node.attrs + .insert("prompt".to_string(), AttrValue::String(new_prompt)); + } + } + } +} + +/// Applies the `model_stylesheet` graph attribute to resolve LLM properties for each node. +pub struct StylesheetApplicationTransform; + +impl Transform for StylesheetApplicationTransform { + fn apply(&self, graph: &mut Graph) { + let stylesheet_text = graph.model_stylesheet().to_string(); + if stylesheet_text.is_empty() { + return; + } + let Ok(stylesheet) = parse_stylesheet(&stylesheet_text) else { + return; + }; + apply_stylesheet(&stylesheet, graph); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::Node; + + #[test] + fn variable_expansion_replaces_goal() { + let mut graph = Graph::new("test"); + graph + .attrs + .insert("goal".to_string(), AttrValue::String("Fix bugs".to_string())); + + let mut node = Node::new("plan"); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Achieve: $goal now".to_string()), + ); + graph.nodes.insert("plan".to_string(), node); + + let transform = VariableExpansionTransform; + transform.apply(&mut graph); + + let prompt = graph.nodes["plan"] + .attrs + .get("prompt") + .and_then(AttrValue::as_str) + .unwrap(); + assert_eq!(prompt, "Achieve: Fix bugs now"); + } + + #[test] + fn variable_expansion_no_goal_variable() { + let mut graph = Graph::new("test"); + graph + .attrs + .insert("goal".to_string(), AttrValue::String("Fix bugs".to_string())); + + let mut node = Node::new("plan"); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Do something".to_string()), + ); + graph.nodes.insert("plan".to_string(), node); + + let transform = VariableExpansionTransform; + transform.apply(&mut graph); + + let prompt = graph.nodes["plan"] + .attrs + .get("prompt") + .and_then(AttrValue::as_str) + .unwrap(); + assert_eq!(prompt, "Do something"); + } + + #[test] + fn variable_expansion_empty_goal() { + let mut graph = Graph::new("test"); + let mut node = Node::new("plan"); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Goal: $goal".to_string()), + ); + graph.nodes.insert("plan".to_string(), node); + + let transform = VariableExpansionTransform; + transform.apply(&mut graph); + + let prompt = graph.nodes["plan"] + .attrs + .get("prompt") + .and_then(AttrValue::as_str) + .unwrap(); + assert_eq!(prompt, "Goal: "); + } + + #[test] + fn variable_expansion_no_prompt() { + let mut graph = Graph::new("test"); + graph + .attrs + .insert("goal".to_string(), AttrValue::String("Fix bugs".to_string())); + let node = Node::new("plan"); + graph.nodes.insert("plan".to_string(), node); + + let transform = VariableExpansionTransform; + // Should not panic + transform.apply(&mut graph); + assert!(graph.nodes["plan"].attrs.get("prompt").is_none()); + } + + #[test] + fn stylesheet_transform_empty_stylesheet() { + let mut graph = Graph::new("test"); + graph.nodes.insert("a".to_string(), Node::new("a")); + + let transform = StylesheetApplicationTransform; + // Should not panic with empty stylesheet + transform.apply(&mut graph); + } + + #[test] + fn preamble_transform_prepends_for_non_full_fidelity() { + let mut graph = Graph::new("test"); + let mut node = Node::new("work"); + node.attrs.insert( + "fidelity".to_string(), + AttrValue::String("truncate".to_string()), + ); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Do the thing".to_string()), + ); + graph.nodes.insert("work".to_string(), node); + + PreambleTransform.apply(&mut graph); + + let prompt = graph.nodes["work"] + .attrs + .get("prompt") + .and_then(AttrValue::as_str) + .unwrap(); + assert_eq!(prompt, "[Context mode: truncate]\nDo the thing"); + } + + #[test] + fn preamble_transform_skips_full_fidelity() { + let mut graph = Graph::new("test"); + let mut node = Node::new("work"); + node.attrs.insert( + "fidelity".to_string(), + AttrValue::String("full".to_string()), + ); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Do the thing".to_string()), + ); + graph.nodes.insert("work".to_string(), node); + + PreambleTransform.apply(&mut graph); + + let prompt = graph.nodes["work"] + .attrs + .get("prompt") + .and_then(AttrValue::as_str) + .unwrap(); + assert_eq!(prompt, "Do the thing"); + } + + #[test] + fn preamble_transform_uses_graph_default_fidelity() { + let mut graph = Graph::new("test"); + graph.attrs.insert( + "default_fidelity".to_string(), + AttrValue::String("compact".to_string()), + ); + let mut node = Node::new("work"); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Do the thing".to_string()), + ); + graph.nodes.insert("work".to_string(), node); + + PreambleTransform.apply(&mut graph); + + let prompt = graph.nodes["work"] + .attrs + .get("prompt") + .and_then(AttrValue::as_str) + .unwrap(); + assert_eq!(prompt, "[Context mode: compact]\nDo the thing"); + } + + #[test] + fn preamble_transform_no_prompt_skips() { + let mut graph = Graph::new("test"); + let mut node = Node::new("work"); + node.attrs.insert( + "fidelity".to_string(), + AttrValue::String("truncate".to_string()), + ); + graph.nodes.insert("work".to_string(), node); + + PreambleTransform.apply(&mut graph); + + assert!(graph.nodes["work"].attrs.get("prompt").is_none()); + } +} diff --git a/crates/attractor/src/validation/mod.rs b/crates/attractor/src/validation/mod.rs new file mode 100644 index 000000000..e41e105bc --- /dev/null +++ b/crates/attractor/src/validation/mod.rs @@ -0,0 +1,160 @@ +pub mod rules; + +use serde::{Deserialize, Serialize}; + +use crate::error::AttractorError; +use crate::graph::Graph; + +/// Severity level for validation diagnostics. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum Severity { + Error, + Warning, + Info, +} + +/// A validation diagnostic produced by a lint rule. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Diagnostic { + pub rule: String, + pub severity: Severity, + pub message: String, + pub node_id: Option, + pub edge: Option<(String, String)>, + pub fix: Option, +} + +/// A lint rule that validates a graph. +pub trait LintRule { + fn name(&self) -> &'static str; + fn apply(&self, graph: &Graph) -> Vec; +} + +/// Run all built-in lint rules (and any extra rules) against the graph. +#[must_use] +pub fn validate(graph: &Graph, extra_rules: &[&dyn LintRule]) -> Vec { + let built_in = rules::built_in_rules(); + let mut diagnostics = Vec::new(); + for rule in &built_in { + diagnostics.extend(rule.apply(graph)); + } + for rule in extra_rules { + diagnostics.extend(rule.apply(graph)); + } + diagnostics +} + +/// Run all built-in lint rules (and any extra rules). Returns Err if any Error-severity +/// diagnostics are found. +/// +/// # Errors +/// Returns `AttractorError::Validation` if any Error-severity diagnostics are found. +pub fn validate_or_raise( + graph: &Graph, + extra_rules: &[&dyn LintRule], +) -> Result, AttractorError> { + let diagnostics = validate(graph, extra_rules); + let errors: Vec<&Diagnostic> = diagnostics + .iter() + .filter(|d| d.severity == Severity::Error) + .collect(); + if !errors.is_empty() { + let messages: Vec = errors.iter().map(|d| d.message.clone()).collect(); + return Err(AttractorError::Validation(messages.join("; "))); + } + Ok(diagnostics) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::{AttrValue, Edge, Graph, Node}; + + fn minimal_valid_graph() -> Graph { + let mut g = Graph::new("test"); + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + g.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs + .insert("shape".to_string(), AttrValue::String("Msquare".to_string())); + g.nodes.insert("exit".to_string(), exit); + + g.edges.push(Edge::new("start", "exit")); + g + } + + #[test] + fn validate_minimal_valid_graph_has_no_errors() { + let g = minimal_valid_graph(); + let diagnostics = validate(&g, &[]); + let errors: Vec<_> = diagnostics + .iter() + .filter(|d| d.severity == Severity::Error) + .collect(); + assert!(errors.is_empty(), "Expected no errors, got: {errors:?}"); + } + + #[test] + fn validate_or_raise_passes_for_valid_graph() { + let g = minimal_valid_graph(); + let result = validate_or_raise(&g, &[]); + assert!(result.is_ok()); + } + + #[test] + fn validate_or_raise_fails_for_missing_start() { + let mut g = Graph::new("test"); + let mut exit = Node::new("exit"); + exit.attrs + .insert("shape".to_string(), AttrValue::String("Msquare".to_string())); + g.nodes.insert("exit".to_string(), exit); + let result = validate_or_raise(&g, &[]); + assert!(result.is_err()); + } + + #[test] + fn validate_or_raise_fails_for_missing_exit() { + let mut g = Graph::new("test"); + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + g.nodes.insert("start".to_string(), start); + let result = validate_or_raise(&g, &[]); + assert!(result.is_err()); + } + + #[test] + fn validate_runs_extra_rules() { + struct AlwaysWarnRule; + impl LintRule for AlwaysWarnRule { + fn name(&self) -> &'static str { "always_warn" } + fn apply(&self, _graph: &Graph) -> Vec { + vec![Diagnostic { + rule: "always_warn".to_string(), + severity: Severity::Warning, + message: "custom warning".to_string(), + node_id: None, + edge: None, + fix: None, + }] + } + } + let g = minimal_valid_graph(); + let extra = AlwaysWarnRule; + let diagnostics = validate(&g, &[&extra]); + let custom: Vec<_> = diagnostics.iter().filter(|d| d.rule == "always_warn").collect(); + assert_eq!(custom.len(), 1); + } + + #[test] + fn diagnostic_severity_eq() { + assert_eq!(Severity::Error, Severity::Error); + assert_ne!(Severity::Error, Severity::Warning); + assert_ne!(Severity::Warning, Severity::Info); + } +} diff --git a/crates/attractor/src/validation/rules.rs b/crates/attractor/src/validation/rules.rs new file mode 100644 index 000000000..cd13ea44d --- /dev/null +++ b/crates/attractor/src/validation/rules.rs @@ -0,0 +1,1130 @@ +use std::collections::{HashSet, VecDeque}; + +use crate::graph::{AttrValue, Graph}; + +use super::{Diagnostic, LintRule, Severity}; + +/// Returns all 14 built-in lint rules. +#[must_use] +pub fn built_in_rules() -> Vec> { + vec![ + Box::new(StartNodeRule), + Box::new(TerminalNodeRule), + Box::new(ReachabilityRule), + Box::new(EdgeTargetExistsRule), + Box::new(StartNoIncomingRule), + Box::new(ExitNoOutgoingRule), + Box::new(ConditionSyntaxRule), + Box::new(StylesheetSyntaxRule), + Box::new(TypeKnownRule), + Box::new(FidelityValidRule), + Box::new(RetryTargetExistsRule), + Box::new(GoalGateHasRetryRule), + Box::new(PromptOnLlmNodesRule), + Box::new(FreeformEdgeCountRule), + ] +} + +// --- Rule 1: start_node (ERROR) --- + +struct StartNodeRule; + +impl LintRule for StartNodeRule { + fn name(&self) -> &'static str { + "start_node" + } + + fn apply(&self, graph: &Graph) -> Vec { + let start_count = graph + .nodes + .iter() + .filter(|(id, n)| { + n.shape() == "Mdiamond" || *id == "start" || *id == "Start" + }) + .count(); + if start_count == 0 { + return vec![Diagnostic { + rule: self.name().to_string(), + severity: Severity::Error, + message: "Pipeline must have exactly one start node (shape=Mdiamond or id start/Start)".to_string(), + node_id: None, + edge: None, + fix: Some("Add a node with shape=Mdiamond or id 'start'".to_string()), + }]; + } + if start_count > 1 { + return vec![Diagnostic { + rule: self.name().to_string(), + severity: Severity::Error, + message: format!( + "Pipeline has {start_count} start nodes but must have exactly one" + ), + node_id: None, + edge: None, + fix: Some("Remove extra start nodes".to_string()), + }]; + } + Vec::new() + } +} + +// --- Rule 2: terminal_node (ERROR) --- + +struct TerminalNodeRule; + +impl LintRule for TerminalNodeRule { + fn name(&self) -> &'static str { + "terminal_node" + } + + fn apply(&self, graph: &Graph) -> Vec { + let terminal_count = graph + .nodes + .iter() + .filter(|(id, n)| { + n.shape() == "Msquare" + || *id == "exit" + || *id == "Exit" + || *id == "end" + || *id == "End" + }) + .count(); + if terminal_count == 0 { + return vec![Diagnostic { + rule: self.name().to_string(), + severity: Severity::Error, + message: "Pipeline must have at least one terminal node (shape=Msquare or id exit/end)".to_string(), + node_id: None, + edge: None, + fix: Some("Add a node with shape=Msquare or id 'exit'/'end'".to_string()), + }]; + } + Vec::new() + } +} + +// --- Rule 3: reachability (ERROR) --- + +struct ReachabilityRule; + +impl LintRule for ReachabilityRule { + fn name(&self) -> &'static str { + "reachability" + } + + fn apply(&self, graph: &Graph) -> Vec { + let Some(start) = graph.find_start_node() else { + return Vec::new(); + }; + + let mut visited = HashSet::new(); + let mut queue = VecDeque::new(); + queue.push_back(start.id.clone()); + visited.insert(start.id.clone()); + + while let Some(node_id) = queue.pop_front() { + for edge in graph.outgoing_edges(&node_id) { + if visited.insert(edge.to.clone()) { + queue.push_back(edge.to.clone()); + } + } + } + + let mut unreachable: Vec<&str> = graph + .nodes + .keys() + .filter(|id| !visited.contains(id.as_str())) + .map(std::string::String::as_str) + .collect(); + unreachable.sort_unstable(); + + unreachable + .into_iter() + .map(|node_id| Diagnostic { + rule: self.name().to_string(), + severity: Severity::Error, + message: format!("Node '{node_id}' is not reachable from the start node"), + node_id: Some(node_id.to_string()), + edge: None, + fix: Some(format!( + "Add an edge path from the start node to '{node_id}'" + )), + }) + .collect() + } +} + +// --- Rule 4: edge_target_exists (ERROR) --- + +struct EdgeTargetExistsRule; + +impl LintRule for EdgeTargetExistsRule { + fn name(&self) -> &'static str { + "edge_target_exists" + } + + fn apply(&self, graph: &Graph) -> Vec { + let mut diagnostics = Vec::new(); + for edge in &graph.edges { + if !graph.nodes.contains_key(&edge.to) { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Error, + message: format!( + "Edge from '{}' targets non-existent node '{}'", + edge.from, edge.to + ), + node_id: None, + edge: Some((edge.from.clone(), edge.to.clone())), + fix: Some(format!("Define node '{}' or fix the edge target", edge.to)), + }); + } + if !graph.nodes.contains_key(&edge.from) { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Error, + message: format!( + "Edge source '{}' references non-existent node", + edge.from + ), + node_id: None, + edge: Some((edge.from.clone(), edge.to.clone())), + fix: Some(format!( + "Define node '{}' or fix the edge source", + edge.from + )), + }); + } + } + diagnostics + } +} + +// --- Rule 5: start_no_incoming (ERROR) --- + +struct StartNoIncomingRule; + +impl LintRule for StartNoIncomingRule { + fn name(&self) -> &'static str { + "start_no_incoming" + } + + fn apply(&self, graph: &Graph) -> Vec { + let Some(start) = graph.find_start_node() else { + return Vec::new(); + }; + let incoming = graph.incoming_edges(&start.id); + if !incoming.is_empty() { + return vec![Diagnostic { + rule: self.name().to_string(), + severity: Severity::Error, + message: format!( + "Start node '{}' has {} incoming edge(s) but must have none", + start.id, + incoming.len() + ), + node_id: Some(start.id.clone()), + edge: None, + fix: Some("Remove incoming edges to the start node".to_string()), + }]; + } + Vec::new() + } +} + +// --- Rule 6: exit_no_outgoing (ERROR) --- + +struct ExitNoOutgoingRule; + +impl LintRule for ExitNoOutgoingRule { + fn name(&self) -> &'static str { + "exit_no_outgoing" + } + + fn apply(&self, graph: &Graph) -> Vec { + let mut diagnostics = Vec::new(); + for (id, node) in &graph.nodes { + let is_terminal = node.shape() == "Msquare" + || *id == "exit" + || *id == "Exit" + || *id == "end" + || *id == "End"; + if is_terminal { + let outgoing = graph.outgoing_edges(&node.id); + if !outgoing.is_empty() { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Error, + message: format!( + "Exit node '{}' has {} outgoing edge(s) but must have none", + node.id, + outgoing.len() + ), + node_id: Some(node.id.clone()), + edge: None, + fix: Some("Remove outgoing edges from the exit node".to_string()), + }); + } + } + } + diagnostics + } +} + +// --- Rule 7: condition_syntax (ERROR) --- + +struct ConditionSyntaxRule; + +impl LintRule for ConditionSyntaxRule { + fn name(&self) -> &'static str { + "condition_syntax" + } + + fn apply(&self, graph: &Graph) -> Vec { + let mut diagnostics = Vec::new(); + for edge in &graph.edges { + let Some(condition) = edge.condition() else { + continue; + }; + if condition.is_empty() { + continue; + } + for clause in condition.split("&&") { + let clause = clause.trim(); + if clause.is_empty() { + continue; + } + // A clause must contain = or != operator, or be a bare key (truthy check) + let has_operator = clause.contains("!=") || clause.contains('='); + if !has_operator && clause.contains(' ') && !clause.starts_with("context.") { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Error, + message: format!( + "Invalid condition clause '{clause}' on edge {} -> {}", + edge.from, edge.to + ), + node_id: None, + edge: Some((edge.from.clone(), edge.to.clone())), + fix: Some("Use key=value or key!=value syntax".to_string()), + }); + } + } + } + diagnostics + } +} + +// --- Rule 8: stylesheet_syntax (ERROR) --- + +struct StylesheetSyntaxRule; + +impl LintRule for StylesheetSyntaxRule { + fn name(&self) -> &'static str { + "stylesheet_syntax" + } + + fn apply(&self, graph: &Graph) -> Vec { + let stylesheet = graph.model_stylesheet(); + if stylesheet.is_empty() { + return Vec::new(); + } + let open_count = stylesheet.chars().filter(|c| *c == '{').count(); + let close_count = stylesheet.chars().filter(|c| *c == '}').count(); + if open_count != close_count { + return vec![Diagnostic { + rule: self.name().to_string(), + severity: Severity::Error, + message: format!( + "Model stylesheet has unbalanced braces ({open_count} open, {close_count} close)" + ), + node_id: None, + edge: None, + fix: Some("Balance the curly braces in model_stylesheet".to_string()), + }]; + } + Vec::new() + } +} + +// --- Rule 9: type_known (WARNING) --- + +struct TypeKnownRule; + +const KNOWN_HANDLER_TYPES: &[&str] = &[ + "start", + "exit", + "codergen", + "wait.human", + "conditional", + "parallel", + "parallel.fan_in", + "tool", + "stack.manager_loop", +]; + +impl LintRule for TypeKnownRule { + fn name(&self) -> &'static str { + "type_known" + } + + fn apply(&self, graph: &Graph) -> Vec { + let mut diagnostics = Vec::new(); + for node in graph.nodes.values() { + if let Some(node_type) = node.node_type() { + if !KNOWN_HANDLER_TYPES.contains(&node_type) { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!( + "Node '{}' has unrecognized type '{node_type}'", + node.id + ), + node_id: Some(node.id.clone()), + edge: None, + fix: Some(format!("Use one of: {}", KNOWN_HANDLER_TYPES.join(", "))), + }); + } + } + } + diagnostics + } +} + +// --- Rule 10: fidelity_valid (WARNING) --- + +struct FidelityValidRule; + +const VALID_FIDELITY_MODES: &[&str] = &[ + "full", + "truncate", + "compact", + "summary:low", + "summary:medium", + "summary:high", +]; + +impl LintRule for FidelityValidRule { + fn name(&self) -> &'static str { + "fidelity_valid" + } + + fn apply(&self, graph: &Graph) -> Vec { + let mut diagnostics = Vec::new(); + for node in graph.nodes.values() { + if let Some(fidelity) = node.fidelity() { + if !VALID_FIDELITY_MODES.contains(&fidelity) { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!( + "Node '{}' has invalid fidelity mode '{fidelity}'", + node.id + ), + node_id: Some(node.id.clone()), + edge: None, + fix: Some(format!( + "Use one of: {}", + VALID_FIDELITY_MODES.join(", ") + )), + }); + } + } + } + for edge in &graph.edges { + if let Some(fidelity) = edge.fidelity() { + if !VALID_FIDELITY_MODES.contains(&fidelity) { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!( + "Edge {} -> {} has invalid fidelity mode '{fidelity}'", + edge.from, edge.to + ), + node_id: None, + edge: Some((edge.from.clone(), edge.to.clone())), + fix: Some(format!( + "Use one of: {}", + VALID_FIDELITY_MODES.join(", ") + )), + }); + } + } + } + if let Some(fidelity) = graph.default_fidelity() { + if !VALID_FIDELITY_MODES.contains(&fidelity) { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!("Graph has invalid default_fidelity '{fidelity}'"), + node_id: None, + edge: None, + fix: Some(format!( + "Use one of: {}", + VALID_FIDELITY_MODES.join(", ") + )), + }); + } + } + diagnostics + } +} + +// --- Rule 11: retry_target_exists (WARNING) --- + +struct RetryTargetExistsRule; + +impl LintRule for RetryTargetExistsRule { + fn name(&self) -> &'static str { + "retry_target_exists" + } + + fn apply(&self, graph: &Graph) -> Vec { + let mut diagnostics = Vec::new(); + for node in graph.nodes.values() { + if let Some(target) = node.retry_target() { + if !graph.nodes.contains_key(target) { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!( + "Node '{}' has retry_target '{}' that does not exist", + node.id, target + ), + node_id: Some(node.id.clone()), + edge: None, + fix: Some(format!("Define node '{target}' or fix retry_target")), + }); + } + } + if let Some(target) = node.fallback_retry_target() { + if !graph.nodes.contains_key(target) { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!( + "Node '{}' has fallback_retry_target '{}' that does not exist", + node.id, target + ), + node_id: Some(node.id.clone()), + edge: None, + fix: Some(format!( + "Define node '{target}' or fix fallback_retry_target" + )), + }); + } + } + } + if let Some(target) = graph.retry_target() { + if !graph.nodes.contains_key(target) { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!("Graph has retry_target '{target}' that does not exist"), + node_id: None, + edge: None, + fix: Some(format!("Define node '{target}' or fix graph retry_target")), + }); + } + } + if let Some(target) = graph.fallback_retry_target() { + if !graph.nodes.contains_key(target) { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!( + "Graph has fallback_retry_target '{target}' that does not exist" + ), + node_id: None, + edge: None, + fix: Some(format!( + "Define node '{target}' or fix graph fallback_retry_target" + )), + }); + } + } + diagnostics + } +} + +// --- Rule 12: goal_gate_has_retry (WARNING) --- + +struct GoalGateHasRetryRule; + +impl LintRule for GoalGateHasRetryRule { + fn name(&self) -> &'static str { + "goal_gate_has_retry" + } + + fn apply(&self, graph: &Graph) -> Vec { + let mut diagnostics = Vec::new(); + for node in graph.nodes.values() { + if node.goal_gate() { + let has_node_retry = + node.retry_target().is_some() || node.fallback_retry_target().is_some(); + let has_graph_retry = + graph.retry_target().is_some() || graph.fallback_retry_target().is_some(); + if !has_node_retry && !has_graph_retry { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!( + "Node '{}' has goal_gate=true but no retry_target or fallback_retry_target", + node.id + ), + node_id: Some(node.id.clone()), + edge: None, + fix: Some( + "Add retry_target or fallback_retry_target attribute".to_string(), + ), + }); + } + } + } + diagnostics + } +} + +// --- Rule 13: prompt_on_llm_nodes (WARNING) --- + +struct PromptOnLlmNodesRule; + +impl LintRule for PromptOnLlmNodesRule { + fn name(&self) -> &'static str { + "prompt_on_llm_nodes" + } + + fn apply(&self, graph: &Graph) -> Vec { + let mut diagnostics = Vec::new(); + for node in graph.nodes.values() { + if node.handler_type() == Some("codergen") { + let has_prompt = node.prompt().is_some_and(|p| !p.is_empty()); + let has_label = node + .attrs + .get("label") + .and_then(AttrValue::as_str) + .is_some_and(|l| !l.is_empty()); + if !has_prompt && !has_label { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!( + "Codergen node '{}' has no prompt or label attribute", + node.id + ), + node_id: Some(node.id.clone()), + edge: None, + fix: Some("Add a prompt or label attribute".to_string()), + }); + } + } + } + diagnostics + } +} + +// --- Rule 14: freeform_edge_count (ERROR) --- + +struct FreeformEdgeCountRule; + +impl LintRule for FreeformEdgeCountRule { + fn name(&self) -> &'static str { + "freeform_edge_count" + } + + fn apply(&self, graph: &Graph) -> Vec { + let mut diagnostics = Vec::new(); + for node in graph.nodes.values() { + if node.handler_type() == Some("wait.human") { + let freeform_count = graph + .outgoing_edges(&node.id) + .iter() + .filter(|e| e.freeform()) + .count(); + if freeform_count > 1 { + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Error, + message: format!( + "wait.human node '{}' has {freeform_count} freeform edges but at most one is allowed", + node.id + ), + node_id: Some(node.id.clone()), + edge: None, + fix: Some( + "Remove extra freeform=true edges so at most one remains".to_string(), + ), + }); + } + } + } + diagnostics + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::graph::{AttrValue, Edge, Node}; + + fn minimal_graph() -> Graph { + let mut g = Graph::new("test"); + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + g.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs + .insert("shape".to_string(), AttrValue::String("Msquare".to_string())); + g.nodes.insert("exit".to_string(), exit); + + g.edges.push(Edge::new("start", "exit")); + g + } + + // start_node rule tests + + #[test] + fn start_node_rule_no_start() { + let g = Graph::new("test"); + let rule = StartNodeRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Error); + } + + #[test] + fn start_node_rule_two_starts() { + let mut g = Graph::new("test"); + let mut s1 = Node::new("s1"); + s1.attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + let mut s2 = Node::new("s2"); + s2.attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + g.nodes.insert("s1".to_string(), s1); + g.nodes.insert("s2".to_string(), s2); + let rule = StartNodeRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Error); + } + + #[test] + fn start_node_rule_one_start() { + let g = minimal_graph(); + let rule = StartNodeRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + #[test] + fn start_node_rule_by_id() { + let mut g = Graph::new("test"); + // Node with id "start" but no Mdiamond shape + let node = Node::new("start"); + g.nodes.insert("start".to_string(), node); + let rule = StartNodeRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + #[test] + fn start_node_rule_by_capitalized_id() { + let mut g = Graph::new("test"); + let node = Node::new("Start"); + g.nodes.insert("Start".to_string(), node); + let rule = StartNodeRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // terminal_node rule tests + + #[test] + fn terminal_node_rule_no_terminal() { + let mut g = Graph::new("test"); + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + g.nodes.insert("start".to_string(), start); + let rule = TerminalNodeRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Error); + } + + #[test] + fn terminal_node_rule_with_terminal() { + let g = minimal_graph(); + let rule = TerminalNodeRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + #[test] + fn terminal_node_rule_by_exit_id() { + let mut g = Graph::new("test"); + // Node with id "exit" but no Msquare shape + let node = Node::new("exit"); + g.nodes.insert("exit".to_string(), node); + let rule = TerminalNodeRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + #[test] + fn terminal_node_rule_by_end_id() { + let mut g = Graph::new("test"); + let node = Node::new("end"); + g.nodes.insert("end".to_string(), node); + let rule = TerminalNodeRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + #[test] + fn terminal_node_rule_by_capitalized_end_id() { + let mut g = Graph::new("test"); + let node = Node::new("End"); + g.nodes.insert("End".to_string(), node); + let rule = TerminalNodeRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // reachability rule tests + + #[test] + fn reachability_rule_unreachable_node() { + let mut g = minimal_graph(); + g.nodes.insert("orphan".to_string(), Node::new("orphan")); + let rule = ReachabilityRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].node_id, Some("orphan".to_string())); + } + + #[test] + fn reachability_rule_all_reachable() { + let g = minimal_graph(); + let rule = ReachabilityRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // edge_target_exists rule tests + + #[test] + fn edge_target_exists_rule_missing_target() { + let mut g = minimal_graph(); + g.edges.push(Edge::new("start", "nonexistent")); + let rule = EdgeTargetExistsRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Error); + } + + #[test] + fn edge_target_exists_rule_valid() { + let g = minimal_graph(); + let rule = EdgeTargetExistsRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // start_no_incoming rule tests + + #[test] + fn start_no_incoming_rule_with_incoming() { + let mut g = minimal_graph(); + g.edges.push(Edge::new("exit", "start")); + let rule = StartNoIncomingRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Error); + } + + #[test] + fn start_no_incoming_rule_clean() { + let g = minimal_graph(); + let rule = StartNoIncomingRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // exit_no_outgoing rule tests + + #[test] + fn exit_no_outgoing_rule_with_outgoing() { + let mut g = minimal_graph(); + g.edges.push(Edge::new("exit", "start")); + let rule = ExitNoOutgoingRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Error); + } + + #[test] + fn exit_no_outgoing_rule_clean() { + let g = minimal_graph(); + let rule = ExitNoOutgoingRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // type_known rule tests + + #[test] + fn type_known_rule_unknown_type() { + let mut g = minimal_graph(); + let mut node = Node::new("custom"); + node.attrs.insert( + "type".to_string(), + AttrValue::String("unknown_type".to_string()), + ); + g.nodes.insert("custom".to_string(), node); + let rule = TypeKnownRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Warning); + } + + #[test] + fn type_known_rule_known_type() { + let mut g = minimal_graph(); + let mut node = Node::new("gate"); + node.attrs.insert( + "type".to_string(), + AttrValue::String("wait.human".to_string()), + ); + g.nodes.insert("gate".to_string(), node); + let rule = TypeKnownRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // fidelity_valid rule tests + + #[test] + fn fidelity_valid_rule_invalid_mode() { + let mut g = minimal_graph(); + let mut node = Node::new("work"); + node.attrs.insert( + "fidelity".to_string(), + AttrValue::String("invalid_mode".to_string()), + ); + g.nodes.insert("work".to_string(), node); + let rule = FidelityValidRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Warning); + } + + #[test] + fn fidelity_valid_rule_valid_mode() { + let mut g = minimal_graph(); + let mut node = Node::new("work"); + node.attrs.insert( + "fidelity".to_string(), + AttrValue::String("full".to_string()), + ); + g.nodes.insert("work".to_string(), node); + let rule = FidelityValidRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // freeform_edge_count rule tests + + #[test] + fn freeform_edge_count_rule_two_freeform() { + let mut g = minimal_graph(); + let mut gate = Node::new("gate"); + gate.attrs.insert( + "shape".to_string(), + AttrValue::String("hexagon".to_string()), + ); + g.nodes.insert("gate".to_string(), gate); + g.nodes.insert("a".to_string(), Node::new("a")); + g.nodes.insert("b".to_string(), Node::new("b")); + + let mut e1 = Edge::new("gate", "a"); + e1.attrs + .insert("freeform".to_string(), AttrValue::Boolean(true)); + let mut e2 = Edge::new("gate", "b"); + e2.attrs + .insert("freeform".to_string(), AttrValue::Boolean(true)); + g.edges.push(e1); + g.edges.push(e2); + + let rule = FreeformEdgeCountRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Error); + } + + #[test] + fn freeform_edge_count_rule_one_freeform() { + let mut g = minimal_graph(); + let mut gate = Node::new("gate"); + gate.attrs.insert( + "shape".to_string(), + AttrValue::String("hexagon".to_string()), + ); + g.nodes.insert("gate".to_string(), gate); + g.nodes.insert("a".to_string(), Node::new("a")); + + let mut e1 = Edge::new("gate", "a"); + e1.attrs + .insert("freeform".to_string(), AttrValue::Boolean(true)); + g.edges.push(e1); + + let rule = FreeformEdgeCountRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // goal_gate_has_retry rule tests + + #[test] + fn goal_gate_has_retry_rule_no_retry() { + let mut g = minimal_graph(); + let mut node = Node::new("work"); + node.attrs + .insert("goal_gate".to_string(), AttrValue::Boolean(true)); + g.nodes.insert("work".to_string(), node); + let rule = GoalGateHasRetryRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Warning); + } + + #[test] + fn goal_gate_has_retry_rule_with_retry() { + let mut g = minimal_graph(); + let mut node = Node::new("work"); + node.attrs + .insert("goal_gate".to_string(), AttrValue::Boolean(true)); + node.attrs.insert( + "retry_target".to_string(), + AttrValue::String("start".to_string()), + ); + g.nodes.insert("work".to_string(), node); + let rule = GoalGateHasRetryRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // prompt_on_llm_nodes rule tests + + #[test] + fn prompt_on_llm_nodes_rule_no_prompt_no_label() { + let mut g = minimal_graph(); + let node = Node::new("work"); + g.nodes.insert("work".to_string(), node); + let rule = PromptOnLlmNodesRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Warning); + } + + #[test] + fn prompt_on_llm_nodes_rule_with_prompt() { + let mut g = minimal_graph(); + let mut node = Node::new("work"); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Do the thing".to_string()), + ); + g.nodes.insert("work".to_string(), node); + let rule = PromptOnLlmNodesRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // condition_syntax rule tests + + #[test] + fn condition_syntax_rule_valid_condition() { + let mut g = minimal_graph(); + let mut edge = Edge::new("start", "exit"); + edge.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + g.edges = vec![edge]; + let rule = ConditionSyntaxRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // stylesheet_syntax rule tests + + #[test] + fn stylesheet_syntax_rule_unbalanced() { + let mut g = minimal_graph(); + g.attrs.insert( + "model_stylesheet".to_string(), + AttrValue::String("* { llm_model: foo;".to_string()), + ); + let rule = StylesheetSyntaxRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Error); + } + + #[test] + fn stylesheet_syntax_rule_balanced() { + let mut g = minimal_graph(); + g.attrs.insert( + "model_stylesheet".to_string(), + AttrValue::String("* { llm_model: foo; }".to_string()), + ); + let rule = StylesheetSyntaxRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // retry_target_exists rule tests + + #[test] + fn retry_target_exists_rule_missing() { + let mut g = minimal_graph(); + let mut node = Node::new("work"); + node.attrs.insert( + "retry_target".to_string(), + AttrValue::String("nonexistent".to_string()), + ); + g.nodes.insert("work".to_string(), node); + let rule = RetryTargetExistsRule; + let d = rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Warning); + } + + #[test] + fn retry_target_exists_rule_valid() { + let mut g = minimal_graph(); + let mut node = Node::new("work"); + node.attrs.insert( + "retry_target".to_string(), + AttrValue::String("start".to_string()), + ); + g.nodes.insert("work".to_string(), node); + let rule = RetryTargetExistsRule; + let d = rule.apply(&g); + assert!(d.is_empty()); + } + + // built_in_rules tests + + #[test] + fn built_in_rules_returns_14_rules() { + let rules = built_in_rules(); + assert_eq!(rules.len(), 14); + } +} diff --git a/crates/attractor/tests/integration.rs b/crates/attractor/tests/integration.rs new file mode 100644 index 000000000..962473751 --- /dev/null +++ b/crates/attractor/tests/integration.rs @@ -0,0 +1,1180 @@ +use std::collections::VecDeque; +use std::path::Path; +use std::sync::Arc; + +use attractor::checkpoint::Checkpoint; +use attractor::context::Context; +use attractor::engine::{PipelineEngine, RunConfig}; +use attractor::error::AttractorError; +use attractor::event::EventEmitter; +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::start::StartHandler; +use attractor::handler::wait_human::WaitHumanHandler; +use attractor::handler::{Handler, HandlerRegistry}; +use attractor::interviewer::queue::QueueInterviewer; +use attractor::interviewer::{Answer, AnswerValue}; +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; + +// --------------------------------------------------------------------------- +// 1. Parse and validate all 3 spec examples (Section 2.13) +// --------------------------------------------------------------------------- + +#[test] +fn parse_and_validate_simple_linear() { + let input = r#"digraph Simple { + graph [goal="Run tests and report"] + rankdir=LR + + start [shape=Mdiamond, label="Start"] + exit [shape=Msquare, label="Exit"] + + run_tests [label="Run Tests", prompt="Run the test suite and report results"] + report [label="Report", prompt="Summarize the test results"] + + start -> run_tests -> report -> exit + }"#; + + let graph = parse(input).expect("parsing should succeed"); + assert_eq!(graph.name, "Simple"); + assert_eq!(graph.goal(), "Run tests and report"); + assert_eq!(graph.nodes.len(), 4); + assert_eq!(graph.edges.len(), 3); + assert!(graph.find_start_node().is_some()); + assert!(graph.find_exit_node().is_some()); + + let diagnostics = validate_or_raise(&graph, &[]).expect("validation should pass"); + let errors: Vec<_> = diagnostics + .iter() + .filter(|d| d.severity == attractor::validation::Severity::Error) + .collect(); + assert!(errors.is_empty(), "expected no validation errors"); +} + +#[test] +fn parse_and_validate_branching_with_conditions() { + let input = r#"digraph Branch { + graph [goal="Implement and validate a feature"] + rankdir=LR + node [shape=box, timeout="900s"] + + start [shape=Mdiamond, label="Start"] + exit [shape=Msquare, label="Exit"] + plan [label="Plan", prompt="Plan the implementation"] + implement [label="Implement", prompt="Implement the plan"] + validate [label="Validate", prompt="Run tests"] + gate [shape=diamond, label="Tests passing?"] + + start -> plan -> implement -> validate -> gate + gate -> exit [label="Yes", condition="outcome=success"] + gate -> implement [label="No", condition="outcome!=success"] + }"#; + + let graph = parse(input).expect("parsing should succeed"); + assert_eq!(graph.name, "Branch"); + assert_eq!(graph.nodes.len(), 6); + assert_eq!(graph.edges.len(), 6); + + let gate_exit = graph + .edges + .iter() + .find(|e| e.from == "gate" && e.to == "exit") + .expect("gate -> exit edge should exist"); + assert_eq!(gate_exit.condition(), Some("outcome=success")); + + let gate_impl = graph + .edges + .iter() + .find(|e| e.from == "gate" && e.to == "implement") + .expect("gate -> implement edge should exist"); + assert_eq!(gate_impl.condition(), Some("outcome!=success")); + + let diagnostics = validate_or_raise(&graph, &[]).expect("validation should pass"); + let errors: Vec<_> = diagnostics + .iter() + .filter(|d| d.severity == attractor::validation::Severity::Error) + .collect(); + assert!(errors.is_empty(), "expected no validation errors"); +} + +#[test] +fn parse_and_validate_human_gate() { + let input = r#"digraph Review { + rankdir=LR + + start [shape=Mdiamond, label="Start"] + exit [shape=Msquare, label="Exit"] + + review_gate [ + shape=hexagon, + label="Review Changes", + type="wait.human" + ] + + start -> review_gate + review_gate -> ship_it [label="[A] Approve"] + review_gate -> fixes [label="[F] Fix"] + ship_it -> exit + fixes -> review_gate + }"#; + + let graph = parse(input).expect("parsing should succeed"); + assert_eq!(graph.name, "Review"); + assert_eq!(graph.nodes.len(), 5); + assert_eq!(graph.edges.len(), 5); + + let gate = &graph.nodes["review_gate"]; + assert_eq!(gate.node_type(), Some("wait.human")); + assert_eq!(gate.shape(), "hexagon"); + assert_eq!(gate.label(), "Review Changes"); + + let diagnostics = validate_or_raise(&graph, &[]).expect("validation should pass"); + let errors: Vec<_> = diagnostics + .iter() + .filter(|d| d.severity == attractor::validation::Severity::Error) + .collect(); + assert!(errors.is_empty(), "expected no validation errors"); +} + +// --------------------------------------------------------------------------- +// 2. End-to-end linear pipeline +// --------------------------------------------------------------------------- + +fn make_linear_registry() -> 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 +} + +#[tokio::test] +async fn end_to_end_linear_pipeline() { + let input = r#"digraph Linear { + graph [goal="Build the feature"] + start [shape=Mdiamond] + exit [shape=Msquare] + codergen_step [shape=box, label="Code", prompt="Implement the feature"] + start -> codergen_step -> exit + }"#; + + let graph = parse(input).expect("parse should succeed"); + validate_or_raise(&graph, &[]).expect("validation should pass"); + + 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 should succeed"); + assert_eq!(outcome.status, StageStatus::Success); + + // Checkpoint should exist + let checkpoint_path = dir.path().join("checkpoint.json"); + assert!(checkpoint_path.exists(), "checkpoint.json should exist"); + + let checkpoint = Checkpoint::load(&checkpoint_path).expect("checkpoint should load"); + assert!(checkpoint.completed_nodes.contains(&"start".to_string())); + assert!(checkpoint + .completed_nodes + .contains(&"codergen_step".to_string())); + + // Codergen handler writes prompt.md, response.md, status.json + let stage_dir = dir.path().join("codergen_step"); + assert!(stage_dir.join("prompt.md").exists(), "prompt.md should exist"); + assert!( + stage_dir.join("response.md").exists(), + "response.md should exist" + ); + assert!( + stage_dir.join("status.json").exists(), + "status.json should exist" + ); + + let prompt_content = std::fs::read_to_string(stage_dir.join("prompt.md")).unwrap(); + assert_eq!(prompt_content, "Implement the feature"); +} + +// --------------------------------------------------------------------------- +// 3. End-to-end branching pipeline +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn end_to_end_branching_pipeline() { + // Build a graph: + // start -> work -> gate (diamond) + // gate -> success_path [condition="outcome=success"] + // gate -> fail_path [condition="outcome=fail"] + // success_path -> exit + // fail_path -> exit + // + // Since work defaults to codergen (shape=box) which returns SUCCESS, + // the engine should route gate -> success_path via condition match. + + let mut graph = Graph::new("BranchTest"); + graph + .attrs + .insert("goal".to_string(), AttrValue::String("Test branching".to_string())); + + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + graph.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs + .insert("shape".to_string(), AttrValue::String("Msquare".to_string())); + graph.nodes.insert("exit".to_string(), exit); + + let mut work = Node::new("work"); + work.attrs + .insert("shape".to_string(), AttrValue::String("box".to_string())); + work.attrs.insert( + "prompt".to_string(), + AttrValue::String("Do work".to_string()), + ); + graph.nodes.insert("work".to_string(), work); + + 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("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")); + graph.edges.push(Edge::new("work", "gate")); + + let mut gate_success = Edge::new("gate", "success_path"); + gate_success.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + graph.edges.push(gate_success); + + let mut gate_fail = Edge::new("gate", "fail_path"); + gate_fail.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=fail".to_string()), + ); + graph.edges.push(gate_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(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)); + + 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 should succeed"); + assert_eq!(outcome.status, StageStatus::Success); + + let checkpoint = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!( + checkpoint + .completed_nodes + .contains(&"success_path".to_string()), + "should have traversed success_path" + ); + assert!( + !checkpoint + .completed_nodes + .contains(&"fail_path".to_string()), + "should NOT have traversed fail_path" + ); +} + +// --------------------------------------------------------------------------- +// 4. End-to-end human gate pipeline with QueueInterviewer +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn end_to_end_human_gate_pipeline() { + // Build a graph: + // start -> gate (hexagon, type=wait.human) + // gate -> approve [label="[A] Approve"] + // gate -> reject [label="[R] Reject"] + // approve -> exit + // reject -> exit + // + // QueueInterviewer pre-filled to select "R" -> should route to reject + + let mut graph = Graph::new("HumanGateTest"); + + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + graph.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs + .insert("shape".to_string(), AttrValue::String("Msquare".to_string())); + graph.nodes.insert("exit".to_string(), exit); + + let mut 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 Changes".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")); + + // Pre-fill the queue with an answer selecting "R" + let answers = VecDeque::from([Answer { + value: AnswerValue::Selected("R".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 should succeed"); + assert_eq!(outcome.status, StageStatus::Success); + + let checkpoint = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!( + checkpoint + .completed_nodes + .contains(&"reject".to_string()), + "should have traversed reject path" + ); + assert!( + !checkpoint + .completed_nodes + .contains(&"approve".to_string()), + "should NOT have traversed approve path" + ); +} + +// --------------------------------------------------------------------------- +// 5. Goal gate enforcement +// --------------------------------------------------------------------------- + +/// A custom handler that always returns FAIL for testing goal gate enforcement. +struct AlwaysFailHandler; + +#[async_trait::async_trait] +impl Handler for AlwaysFailHandler { + async fn execute( + &self, + node: &Node, + _context: &attractor::context::Context, + _graph: &Graph, + _logs_root: &Path, + ) -> Result { + Ok(Outcome::fail(format!("forced failure for {}", node.id))) + } +} + +#[tokio::test] +async fn goal_gate_routes_to_retry_target_on_failure() { + // Pipeline: + // start -> gated_work -> exit + // gated_work has goal_gate=true, retry_target=start + // gated_work always returns FAIL + // + // When engine reaches exit, it checks goal gates and finds gated_work failed. + // It should route back to retry_target (start). + // + // To avoid infinite loops, we set max_retries=0 on gated_work so it fails + // immediately each time. After looping once (start -> gated_work -> exit -> start + // -> gated_work -> exit), if goal gate is still unsatisfied and no retry_target + // changes, we need to limit iterations. The engine itself doesn't limit loops, + // so we test a simpler scenario: verify the error when retry_target is missing. + + // Test: goal_gate with NO retry_target returns an error + let mut graph = Graph::new("GoalGateNoRetry"); + + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + graph.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs + .insert("shape".to_string(), AttrValue::String("Msquare".to_string())); + graph.nodes.insert("exit".to_string(), exit); + + let mut gated_work = Node::new("gated_work"); + gated_work + .attrs + .insert("goal_gate".to_string(), AttrValue::Boolean(true)); + gated_work + .attrs + .insert("max_retries".to_string(), AttrValue::Integer(0)); + gated_work.attrs.insert( + "type".to_string(), + AttrValue::String("always_fail".to_string()), + ); + graph + .nodes + .insert("gated_work".to_string(), gated_work); + + graph.edges.push(Edge::new("start", "gated_work")); + graph.edges.push(Edge::new("gated_work", "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 result = engine.run(&graph, &config).await; + assert!(result.is_err(), "should fail when goal gate unsatisfied and no retry_target"); + let err_msg = result.unwrap_err().to_string(); + assert!( + err_msg.contains("goal gate unsatisfied"), + "error should mention goal gate, got: {err_msg}" + ); +} + +#[tokio::test] +async fn goal_gate_routes_to_retry_target_when_present() { + // Pipeline: + // start -> gated_work -> exit + // gated_work has goal_gate=true, retry_target=start + // gated_work always fails via AlwaysFailHandler. + // + // When engine reaches exit and finds goal gate unsatisfied, it should route + // to the retry_target. Since AlwaysFailHandler always fails, this creates a + // loop. However, the gated_work node will emit a FAIL outcome, and the + // edge gated_work -> exit is unconditional, so it still reaches exit. After + // the first retry (start -> gated_work -> exit), goal gate is still failed + // and retry_target is still start, so it loops. To prevent an infinite loop + // in tests, we use a custom handler that fails the first time and succeeds + // the second time. + + struct FailThenSucceedHandler { + call_count: std::sync::atomic::AtomicU32, + } + + #[async_trait::async_trait] + impl Handler for FailThenSucceedHandler { + async fn execute( + &self, + _node: &Node, + _context: &attractor::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 = Graph::new("GoalGateRetry"); + + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + graph.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs + .insert("shape".to_string(), AttrValue::String("Msquare".to_string())); + graph.nodes.insert("exit".to_string(), exit); + + let mut gated_work = Node::new("gated_work"); + gated_work + .attrs + .insert("goal_gate".to_string(), AttrValue::Boolean(true)); + gated_work + .attrs + .insert("max_retries".to_string(), AttrValue::Integer(0)); + gated_work.attrs.insert( + "retry_target".to_string(), + AttrValue::String("start".to_string()), + ); + gated_work.attrs.insert( + "type".to_string(), + AttrValue::String("fail_then_succeed".to_string()), + ); + graph + .nodes + .insert("gated_work".to_string(), gated_work); + + graph.edges.push(Edge::new("start", "gated_work")); + graph.edges.push(Edge::new("gated_work", "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( + "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 should eventually succeed after retry"); + assert_eq!(outcome.status, StageStatus::Success); + + let checkpoint = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + // gated_work should appear in completed nodes (at least twice -- first fail, then succeed) + let gated_work_count = checkpoint + .completed_nodes + .iter() + .filter(|n| *n == "gated_work") + .count(); + assert!( + gated_work_count >= 2, + "gated_work should have been executed at least twice, got {gated_work_count}" + ); +} + +// --------------------------------------------------------------------------- +// 6. Variable expansion transform +// --------------------------------------------------------------------------- + +#[test] +fn variable_expansion_replaces_goal_in_prompts() { + let mut graph = Graph::new("test"); + graph.attrs.insert( + "goal".to_string(), + AttrValue::String("Fix all bugs".to_string()), + ); + + let mut plan_node = Node::new("plan"); + plan_node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Plan to achieve: $goal".to_string()), + ); + graph.nodes.insert("plan".to_string(), plan_node); + + let mut impl_node = Node::new("implement"); + impl_node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Implement $goal now".to_string()), + ); + graph + .nodes + .insert("implement".to_string(), impl_node); + + let mut no_var_node = Node::new("report"); + no_var_node.attrs.insert( + "prompt".to_string(), + AttrValue::String("Generate a report".to_string()), + ); + graph + .nodes + .insert("report".to_string(), no_var_node); + + let transform = VariableExpansionTransform; + transform.apply(&mut graph); + + let plan_prompt = graph.nodes["plan"] + .attrs + .get("prompt") + .and_then(AttrValue::as_str) + .expect("plan prompt should exist"); + assert_eq!(plan_prompt, "Plan to achieve: Fix all bugs"); + + let impl_prompt = graph.nodes["implement"] + .attrs + .get("prompt") + .and_then(AttrValue::as_str) + .expect("implement prompt should exist"); + assert_eq!(impl_prompt, "Implement Fix all bugs now"); + + let report_prompt = graph.nodes["report"] + .attrs + .get("prompt") + .and_then(AttrValue::as_str) + .expect("report prompt should exist"); + assert_eq!(report_prompt, "Generate a report"); +} + +// --------------------------------------------------------------------------- +// 7. Stylesheet application +// --------------------------------------------------------------------------- + +#[test] +fn stylesheet_application_by_specificity() { + let stylesheet_text = r#" + * { llm_model: claude-sonnet-4-5; llm_provider: anthropic; } + .code { llm_model: claude-opus-4-6; llm_provider: anthropic; } + #critical_review { llm_model: gpt-5.2; llm_provider: openai; reasoning_effort: high; } + "#; + + let mut graph = Graph::new("test"); + graph.attrs.insert( + "model_stylesheet".to_string(), + AttrValue::String(stylesheet_text.to_string()), + ); + + // plan node: no class, should get universal defaults + let plan = Node::new("plan"); + graph.nodes.insert("plan".to_string(), plan); + + // implement node: class="code", should get .code overrides + let mut implement = Node::new("implement"); + implement.classes.push("code".to_string()); + graph + .nodes + .insert("implement".to_string(), implement); + + // critical_review node: class="code" AND id="critical_review", id wins + let mut critical = Node::new("critical_review"); + critical.classes.push("code".to_string()); + graph + .nodes + .insert("critical_review".to_string(), critical); + + // explicit node: has explicit llm_model, should NOT be overridden + let mut explicit = Node::new("explicit_node"); + explicit.attrs.insert( + "llm_model".to_string(), + AttrValue::String("my-custom-model".to_string()), + ); + graph + .nodes + .insert("explicit_node".to_string(), explicit); + + let transform = StylesheetApplicationTransform; + transform.apply(&mut graph); + + // plan: universal -> claude-sonnet-4-5 + assert_eq!( + graph.nodes["plan"].attrs.get("llm_model"), + Some(&AttrValue::String("claude-sonnet-4-5".to_string())) + ); + assert_eq!( + graph.nodes["plan"].attrs.get("llm_provider"), + Some(&AttrValue::String("anthropic".to_string())) + ); + + // implement: .code -> claude-opus-4-6 + assert_eq!( + graph.nodes["implement"].attrs.get("llm_model"), + Some(&AttrValue::String("claude-opus-4-6".to_string())) + ); + assert_eq!( + graph.nodes["implement"].attrs.get("llm_provider"), + Some(&AttrValue::String("anthropic".to_string())) + ); + + // critical_review: #critical_review -> gpt-5.2 (id overrides class) + assert_eq!( + graph.nodes["critical_review"].attrs.get("llm_model"), + Some(&AttrValue::String("gpt-5.2".to_string())) + ); + assert_eq!( + graph.nodes["critical_review"].attrs.get("llm_provider"), + Some(&AttrValue::String("openai".to_string())) + ); + assert_eq!( + graph.nodes["critical_review"] + .attrs + .get("reasoning_effort"), + Some(&AttrValue::String("high".to_string())) + ); + + // explicit_node: explicit attr NOT overridden by universal + assert_eq!( + graph.nodes["explicit_node"].attrs.get("llm_model"), + Some(&AttrValue::String("my-custom-model".to_string())) + ); +} + +#[test] +fn stylesheet_application_via_parsed_graph() { + let input = r#"digraph StyleTest { + graph [ + goal="Test stylesheet", + model_stylesheet="* { llm_model: sonnet; }" + ] + start [shape=Mdiamond] + exit [shape=Msquare] + work [shape=box, prompt="Do work"] + start -> work -> exit + }"#; + + let mut graph = parse(input).expect("parse should succeed"); + validate_or_raise(&graph, &[]).expect("validation should pass"); + + let transform = StylesheetApplicationTransform; + transform.apply(&mut graph); + + // All nodes without explicit llm_model should get "sonnet" + assert_eq!( + graph.nodes["work"].attrs.get("llm_model"), + Some(&AttrValue::String("sonnet".to_string())) + ); + assert_eq!( + graph.nodes["start"].attrs.get("llm_model"), + Some(&AttrValue::String("sonnet".to_string())) + ); + assert_eq!( + graph.nodes["exit"].attrs.get("llm_model"), + Some(&AttrValue::String("sonnet".to_string())) + ); +} + +#[test] +fn stylesheet_parse_and_apply_directly() { + let stylesheet_text = "* { llm_model: base; } .fast { llm_model: turbo; }"; + let stylesheet = parse_stylesheet(stylesheet_text).expect("stylesheet parse should succeed"); + assert_eq!(stylesheet.rules.len(), 2); + + let mut graph = Graph::new("test"); + let plain = Node::new("a"); + graph.nodes.insert("a".to_string(), plain); + + let mut fast_node = Node::new("b"); + fast_node.classes.push("fast".to_string()); + graph.nodes.insert("b".to_string(), fast_node); + + apply_stylesheet(&stylesheet, &mut graph); + + assert_eq!( + graph.nodes["a"].attrs.get("llm_model"), + Some(&AttrValue::String("base".to_string())) + ); + assert_eq!( + graph.nodes["b"].attrs.get("llm_model"), + Some(&AttrValue::String("turbo".to_string())) + ); +} + +// --------------------------------------------------------------------------- +// 8. Retry on failure (Gap #35.1) +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn retry_on_failure_then_succeed() { + // A handler that fails the first call and succeeds on the second. + 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 = Graph::new("RetryTest"); + + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + graph.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs + .insert("shape".to_string(), AttrValue::String("Msquare".to_string())); + graph.nodes.insert("exit".to_string(), exit); + + let mut retry_node = Node::new("work"); + retry_node.attrs.insert( + "type".to_string(), + AttrValue::String("retry_handler".to_string()), + ); + retry_node + .attrs + .insert("max_retries".to_string(), AttrValue::Integer(3)); + graph.nodes.insert("work".to_string(), retry_node); + + graph.edges.push(Edge::new("start", "work")); + graph.edges.push(Edge::new("work", "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("should succeed after retry"); + assert_eq!(outcome.status, StageStatus::Success); +} + +// --------------------------------------------------------------------------- +// 9. Pipeline with 10+ nodes (Gap #35.2) +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn pipeline_with_many_nodes() { + // Build a linear pipeline: start -> n1 -> n2 -> ... -> n10 -> exit (12 nodes) + let mut graph = Graph::new("ManyNodes"); + graph.attrs.insert( + "goal".to_string(), + AttrValue::String("Test large pipeline".to_string()), + ); + + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + graph.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs + .insert("shape".to_string(), AttrValue::String("Msquare".to_string())); + graph.nodes.insert("exit".to_string(), exit); + + let node_names: Vec = (1..=10).map(|i| format!("step_{i}")).collect(); + + for name in &node_names { + let mut node = Node::new(name.clone()); + node.attrs.insert( + "shape".to_string(), + AttrValue::String("box".to_string()), + ); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String(format!("Execute {name}")), + ); + graph.nodes.insert(name.clone(), node); + } + + graph.edges.push(Edge::new("start", &node_names[0])); + for pair in node_names.windows(2) { + graph.edges.push(Edge::new(&pair[0], &pair[1])); + } + graph.edges.push(Edge::new( + node_names.last().unwrap(), + "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(), + }; + + let outcome = engine + .run(&graph, &config) + .await + .expect("large pipeline should succeed"); + assert_eq!(outcome.status, StageStatus::Success); + + let checkpoint = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + // All 10 step nodes should be in completed_nodes + for name in &node_names { + assert!( + checkpoint.completed_nodes.contains(name), + "{name} should be in completed_nodes" + ); + } +} + +// --------------------------------------------------------------------------- +// 10. Checkpoint save and load round-trip (Gap #35.3) +// --------------------------------------------------------------------------- + +#[test] +fn checkpoint_save_and_resume_roundtrip() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("checkpoint.json"); + + let ctx = Context::new(); + ctx.set("goal", serde_json::json!("Test checkpoint")); + ctx.set("progress", serde_json::json!(42)); + ctx.append_log("started"); + ctx.append_log("step_1 completed"); + + let mut checkpoint = Checkpoint::from_context( + &ctx, + "step_2", + vec![ + "start".to_string(), + "step_1".to_string(), + ], + ); + checkpoint.node_retries.insert("step_1".to_string(), 1); + + checkpoint.save(&path).expect("save should succeed"); + + let loaded = Checkpoint::load(&path).expect("load should succeed"); + assert_eq!(loaded.current_node, "step_2"); + assert_eq!(loaded.completed_nodes.len(), 2); + assert!(loaded.completed_nodes.contains(&"start".to_string())); + assert!(loaded.completed_nodes.contains(&"step_1".to_string())); + assert_eq!(loaded.node_retries.get("step_1"), Some(&1)); + assert_eq!( + loaded.context_values.get("goal"), + Some(&serde_json::json!("Test checkpoint")) + ); + assert_eq!( + loaded.context_values.get("progress"), + Some(&serde_json::json!(42)) + ); + assert_eq!(loaded.logs.len(), 2); +} + +// --------------------------------------------------------------------------- +// 11. Smoke test with mock CodergenBackend (Gap #36) +// --------------------------------------------------------------------------- + +struct MockCodergenBackend; + +#[async_trait::async_trait] +impl CodergenBackend for MockCodergenBackend { + async fn run( + &self, + node: &Node, + prompt: &str, + _context: &Context, + ) -> Result { + Ok(CodergenResult::Text(format!( + "Response for {}: processed prompt '{}'", + node.id, + &prompt[..prompt.len().min(50)] + ))) + } +} + +#[tokio::test] +async fn smoke_test_with_mock_codergen_backend() { + // Pipeline: + // start -> plan -> gate (diamond) + // gate -> implement [condition="outcome=success"] + // gate -> fix [condition="outcome!=success"] + // implement -> exit + // fix -> exit + // + // codergen nodes use MockCodergenBackend which returns real Text responses. + // The gate is a conditional node. Since the mock backend returns success, + // we should route through implement. + + let mut graph = Graph::new("SmokeTest"); + graph.attrs.insert( + "goal".to_string(), + AttrValue::String("Build and validate".to_string()), + ); + + let mut start = Node::new("start"); + start + .attrs + .insert("shape".to_string(), AttrValue::String("Mdiamond".to_string())); + graph.nodes.insert("start".to_string(), start); + + let mut exit = Node::new("exit"); + exit.attrs + .insert("shape".to_string(), AttrValue::String("Msquare".to_string())); + graph.nodes.insert("exit".to_string(), exit); + + let mut plan = Node::new("plan"); + plan.attrs + .insert("shape".to_string(), AttrValue::String("box".to_string())); + plan.attrs.insert( + "prompt".to_string(), + AttrValue::String("Plan to achieve: $goal".to_string()), + ); + graph.nodes.insert("plan".to_string(), plan); + + let mut gate = Node::new("gate"); + gate.attrs + .insert("shape".to_string(), AttrValue::String("diamond".to_string())); + graph.nodes.insert("gate".to_string(), gate); + + 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 the plan".to_string()), + ); + graph + .nodes + .insert("implement".to_string(), implement); + + let mut fix = Node::new("fix"); + fix.attrs + .insert("shape".to_string(), AttrValue::String("box".to_string())); + fix.attrs.insert( + "prompt".to_string(), + AttrValue::String("Fix the issues".to_string()), + ); + graph.nodes.insert("fix".to_string(), fix); + + graph.edges.push(Edge::new("start", "plan")); + graph.edges.push(Edge::new("plan", "gate")); + + let mut gate_impl = Edge::new("gate", "implement"); + gate_impl.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome=success".to_string()), + ); + graph.edges.push(gate_impl); + + let mut gate_fix = Edge::new("gate", "fix"); + gate_fix.attrs.insert( + "condition".to_string(), + AttrValue::String("outcome!=success".to_string()), + ); + graph.edges.push(gate_fix); + + graph.edges.push(Edge::new("implement", "exit")); + graph.edges.push(Edge::new("fix", "exit")); + + let dir = tempfile::tempdir().unwrap(); + let backend = Box::new(MockCodergenBackend); + let mut registry = + HandlerRegistry::new(Box::new(CodergenHandler::new(Some(backend)))); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + registry.register( + "codergen", + Box::new(CodergenHandler::new(Some(Box::new(MockCodergenBackend)))), + ); + 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("smoke test should succeed"); + assert_eq!(outcome.status, StageStatus::Success); + + let checkpoint = Checkpoint::load(&dir.path().join("checkpoint.json")).unwrap(); + assert!( + checkpoint + .completed_nodes + .contains(&"plan".to_string()), + "plan should have executed" + ); + assert!( + checkpoint + .completed_nodes + .contains(&"implement".to_string()), + "should route through implement (success path)" + ); + assert!( + !checkpoint + .completed_nodes + .contains(&"fix".to_string()), + "should NOT have traversed fix path" + ); + + // Verify response.md was written by the mock backend + let plan_response = std::fs::read_to_string(dir.path().join("plan").join("response.md")) + .expect("plan response should exist"); + assert!( + plan_response.contains("Response for plan"), + "mock backend should have written response, got: {plan_response}" + ); + + // Verify prompt.md had $goal expanded by the CodergenHandler + let plan_prompt = std::fs::read_to_string(dir.path().join("plan").join("prompt.md")) + .expect("plan prompt should exist"); + assert_eq!(plan_prompt, "Plan to achieve: Build and validate"); +} diff --git a/docs/agent/reviews/attractor-spec-full-review.md b/docs/agent/reviews/attractor-spec-full-review.md new file mode 100644 index 000000000..f13404339 --- /dev/null +++ b/docs/agent/reviews/attractor-spec-full-review.md @@ -0,0 +1,243 @@ +# Attractor Spec Compliance Review + +Full review of `crates/attractor/` against `docs/specs/attractor-spec.md`. +Reviewed 2026-02-20 by 5 parallel agents. Second-pass false-positive analysis applied. + +--- + +## Section 1: Overview and Goals + +| # | Sub-section | Verdict | Notes | +|---|------------|---------|-------| +| 1 | 1.1 Problem Statement | ALIGNED | Narrative; no code requirements | +| 2 | 1.2 Why DOT Syntax | ALIGNED | Parser accepts DOT as specified | +| 3 | 1.3 Design Principles | ALIGNED | Graph layer supports pluggable handlers, checkpoint, HITL, edge routing | +| 4 | 1.4 Layering and LLM Backends | ALIGNED | Backend-agnostic; `CodergenBackend` trait exists | + +--- + +## Section 2: DOT DSL Schema + +| # | Sub-section | Verdict | Notes | +|---|------------|---------|-------| +| 5 | 2.1 Supported Subset | ALIGNED | Strict DOT subset enforced; one digraph per file (`parser/mod.rs:23-29`) | +| 6 | 2.2 BNF Grammar | ALIGNED | All productions implemented in `grammar.rs` and `lexer.rs` | +| 7 | 2.3 Key Constraints | ALIGNED | Directed-only, commas required, optional semicolons, comments stripped | +| 8 | 2.4 Value Types | ALIGNED | String, Integer, Float, Boolean, Duration all supported | +| 9 | 2.5 Graph Attributes | ALIGNED | `goal`, `model_stylesheet`, `default_max_retry`, `retry_target`, `fallback_retry_target`, `default_fidelity` all have accessors | +| 10 | 2.6 Node Attributes | ALIGNED | All 17 node attributes have typed accessors in `graph/types.rs` | +| 11 | 2.7 Edge Attributes | ALIGNED | All 7 edge attributes implemented | +| 12 | 2.8 Shape-to-Handler Mapping | ALIGNED | Complete 9-shape mapping at `types.rs:72-85` with tests | +| 13 | 2.9 Chained Edges | ALIGNED | `A -> B -> C` expanded via `windows(2)` in semantic analysis | +| 14 | 2.10 Subgraphs | ALIGNED | Scoped defaults, class derivation from label | +| 15 | 2.11 Default Blocks | ALIGNED | `node [...]` and `edge [...]` defaults applied correctly | +| 16 | 2.12 Class Attribute | ALIGNED | Comma-separated, trimmed, deduplicated | +| 17 | 2.13 Minimal Examples | ALIGNED | All 3 spec examples parse; tested | + +**Minor gaps in Section 2:** +- `Direction` values (`TB`/`LR`/`BT`/`RL`) not validated to allowed set +- No explicit error for `strict` modifier or undirected `graph` keyword +- `Graph` missing a `label()` convenience accessor +- `Node::max_retries()` returns `Option` instead of defaulting to `0` +- Subgraph class derivation only checks `GraphAttrDecl`, not `graph [label=...]` block form + +--- + +## Section 3: Pipeline Execution Engine + +| # | Sub-section | Verdict | Notes | +|---|------------|---------|-------| +| 18 | 3.1 Run Lifecycle | ALIGNED | 5 phases present; FINALIZE does not clean up resources (minor) | +| 19 | 3.2 Core Execution Loop | GAP | `loop_restart` (step 7) not implemented -- just jumps to target | +| 20 | 3.3 Edge Selection | ALIGNED | All 5 steps match spec; `normalize_label` handles prefixes | +| 21 | 3.4 Goal Gate Enforcement | ALIGNED | 4-level retry_target fallback implemented | +| 22 | 3.5 Retry Logic | GAP | `should_retry` predicate missing; retry counter not tracked; `reset_retry_counter` absent | +| 23 | 3.6 Retry Policy | GAP | Presets defined but never selectable from node attributes; `default_max_retry=50` vs spec default `0` | +| 24 | 3.7 Failure Routing | GAP | retry_target/fallback_retry_target not consulted on node FAIL -- only used for goal gates | +| 25 | 3.8 Concurrency Model | ALIGNED | Single-threaded traversal; parallel handler manages branches | + +--- + +## Section 4: Node Handlers + +| # | Sub-section | Verdict | Notes | +|---|------------|---------|-------| +| 26 | 4.1 Handler Interface | ALIGNED | `Handler` trait matches spec signature | +| 27 | 4.2 Handler Registry | ALIGNED | 3-step priority: explicit type > shape > default | +| 28 | 4.3 Start Handler | ALIGNED | Returns SUCCESS immediately | +| 29 | 4.4 Exit Handler | ALIGNED | Returns SUCCESS immediately | +| 30 | 4.5 Codergen Handler | ALIGNED | Prompt expansion, backend call, artifact writes all present | +| 31 | 4.6 Wait For Human | ALIGNED | Choices, freeform, accelerator keys, timeout/skip handling | +| 32 | 4.7 Conditional Handler | ALIGNED | Pass-through SUCCESS; routing via edge selection | +| 33 | 4.8 Parallel Handler | GAP | **Stub** -- no actual concurrent execution, no context cloning, no join/error policies | +| 34 | 4.9 Fan-In Handler | GAP | Heuristic select works; no LLM-based evaluation path | +| 35 | 4.10 Tool Handler | ALIGNED | Shell execution via `sh -c`; no command timeout (minor) | +| 36 | 4.11 Manager Loop Handler | GAP | **Stub** -- always returns FAIL | +| 37 | 4.12 Custom Handlers | ALIGNED | Trait + registry supports registration; panics not caught (minor) | + +--- + +## Section 5: State and Context + +| # | Sub-section | Verdict | Notes | +|---|------------|---------|-------| +| 38 | 5.1 PipelineContext | ALIGNED | Key-value store with get/set/merge; `internal.retry_count.` never written (minor) | +| 39 | 5.2 Outcome Model | ALIGNED | All fields and status values present | +| 40 | 5.3 Checkpoint/Resume | GAP | `Checkpoint::save` works; `load` exists but **no resume logic in engine**; `node_retries` never populated | +| 41 | 5.4 Context Fidelity | GAP | Attributes parsed/validated but **fidelity resolution precedence and session/thread management not implemented** in engine | +| 42 | 5.5 Artifact Store | ALIGNED | Full implementation including file-backing | +| 43 | 5.6 Run Directory | GAP | Missing `manifest.json`; per-node dirs only created by CodergenHandler, not other handlers | + +--- + +## Section 6: Human-in-the-Loop (Interviewer Pattern) + +| # | Sub-section | Verdict | Notes | +|---|------------|---------|-------| +| 44 | 6.1 Interviewer Trait | ALIGNED | `ask(Question) -> Answer` interface | +| 45 | 6.2 Question Model | ALIGNED | Options, allow_freeform, timeout, stage | +| 46 | 6.3 Answer Model | ALIGNED | AnswerValue variants cover spec cases | +| 47 | 6.4 Built-in Interviewers | GAP | AutoApprove, Callback, Queue, Recording present; **ConsoleInterviewer missing** | +| 48 | 6.5 Timeout Handling | GAP | Data model present but **no runtime timeout enforcement** in any interviewer | + +--- + +## Section 7: Validation and Linting + +| # | Sub-section | Verdict | Notes | +|---|------------|---------|-------| +| 49 | 7.1 Diagnostic Model | ALIGNED | Struct matches spec exactly (rule, severity, message, node_id, edge, fix) | +| 50 | 7.2 Built-In Rules | GAP | 14 rules present; `start_node` and `terminal_node` only check shape, not ID fallback; `stylesheet_syntax` only checks brace balance | +| 51 | 7.3 Validation API | GAP | No `extra_rules` parameter on `validate()`/`validate_or_raise()` | +| 52 | 7.4 Custom Lint Rules | GAP | `LintRule` trait exists but **no registration mechanism** | + +--- + +## Section 8: Model Stylesheet + +| # | Sub-section | Verdict | Notes | +|---|------------|---------|-------| +| 53 | 8.1 Overview | ALIGNED | Stylesheet applied as transform after parsing | +| 54 | 8.2 Grammar | ALIGNED | `*`, `.class`, `#id` selectors; ClassName accepts uppercase (minor) | +| 55 | 8.3 Specificity | ALIGNED | Universal(0) < Class(1) < ID(2); explicit attrs never overridden | +| 56 | 8.4 Recognized Properties | ALIGNED | `llm_model`, `llm_provider`, `reasoning_effort` | +| 57 | 8.5 Application Order | ALIGNED | Explicit > stylesheet > default | +| 58 | 8.6 Example | ALIGNED | Spec example validated by dedicated test | + +**Minor gap:** No shape selector support (e.g., `box { ... }`) -- spec section 11.10 mentions it. + +--- + +## Section 9: Transforms and Extensibility + +| # | Sub-section | Verdict | Notes | +|---|------------|---------|-------| +| 59 | 9.1 AST Transforms | GAP | Trait takes `&mut Graph` (in-place) vs spec's "returns new Graph"; no `prepare_pipeline` function | +| 60 | 9.2 Built-In Transforms | GAP | Variable expansion and stylesheet transforms present; **Preamble Transform missing** | +| 61 | 9.3 Custom Transforms | GAP | Trait is public but **no `register_transform` API** on engine | +| 62 | 9.4 Pipeline Composition | GAP | Manager loop stub; no graph merging transform | +| 63 | 9.5 HTTP Server Mode | ALIGNED | Spec says "may expose" -- not required | +| 64 | 9.6 Events | ALIGNED | All 16 event types defined and serializable | +| 65 | 9.7 Tool Call Hooks | GAP | `tool_hooks.pre`/`tool_hooks.post` not read or executed | + +**Minor gaps in 9.6:** Parallel events and interview events defined but never emitted by their handlers. + +--- + +## Section 10: Condition Expression Language + +| # | Sub-section | Verdict | Notes | +|---|------------|---------|-------| +| 66 | 10.1 Overview | ALIGNED | Dedicated condition module | +| 67 | 10.2 Grammar | ALIGNED | `&&`-separated clauses, `=`/`!=` operators | +| 68 | 10.3 Semantics | ALIGNED | AND-combined, exact case-sensitive string comparison | +| 69 | 10.4 Variable Resolution | ALIGNED | `outcome`, `preferred_label`, `context.*`, missing-as-empty | +| 70 | 10.5 Evaluation | GAP | Bare key truthiness check returns error instead of evaluating as truthy | +| 71 | 10.6 Examples | ALIGNED | All spec examples verified by tests | +| 72 | 10.7 Extended Operators | ALIGNED | Not implemented (spec says future) | + +--- + +## Section 11: Definition of Done + +| # | Sub-section | Verdict | Notes | +|---|------------|---------|-------| +| 73 | 11.1 DOT Parsing | ALIGNED | Integration tests verify all spec examples | +| 74 | 11.2 Validation | ALIGNED | 14 rules, `validate_or_raise` blocks on errors | +| 75 | 11.3 Execution Engine | ALIGNED | Full loop, edge selection, handler dispatch | +| 76 | 11.4 Goal Gates | ALIGNED | Checked at terminal; retry_target fallback chain | +| 77 | 11.5 Retry Logic | ALIGNED | Exponential backoff with jitter; `allow_partial` | +| 78 | 11.6 Node Handlers | GAP | Manager loop stub; parallel stub | +| 79 | 11.7 State and Context | GAP | Checkpoint resume not implemented | +| 80 | 11.8 Human-in-the-Loop | GAP | No ConsoleInterviewer; no SINGLE_SELECT/MULTI_SELECT distinction | +| 81 | 11.9 Conditions | ALIGNED | With bare-key gap noted above | +| 82 | 11.10 Stylesheet | GAP | No shape selector | +| 83 | 11.11 Transforms | GAP | No `register_transform` API; no preamble transform | +| 84 | 11.12 Cross-Feature Matrix | GAP | Missing integration tests: retry-on-failure, checkpoint resume, 10+ node pipeline | +| 85 | 11.13 Integration Smoke Test | GAP | No end-to-end test with real LLM callback | + +--- + +## Gap Summary by Severity (after false-positive analysis) + +9 of 36 original gaps were false positives. **27 legitimate gaps remain.** + +### False Positives Removed + +| Original # | Reason | +|------------|--------| +| 8 (Retry Presets) | Spec does not require presets be selectable from node attributes | +| 9 (default_max_retry) | Spec attribute tables say default is 50; code matches | +| 24 (Direction) | Direction is a Graphviz layout hint, not an Attractor semantic attribute | +| 25 (strict/graph) | Implicit rejection via grammar is sufficient | +| 26 (Graph::label()) | Attribute is stored and accessible via `attrs` map | +| 27 (max_retries default) | Spec tables and code agree on graph-level default of 50 | +| 29 (Variable Expansion) | Spec says only `$goal` from graph; code matches exactly | +| 32 (Transform Signature) | `&mut Graph` is Rust-idiomatic; spec allows "modified graph" | +| 34 (QuestionType) | Spec 6.2 defines the actual types; 11.8 uses inconsistent names | + +### HIGH (5 gaps -- blocking or core functionality missing) + +| # | Location | Gap | +|---|----------|-----| +| 1 | 3.7 Failure Routing | retry_target/fallback_retry_target not consulted on node FAIL (only on goal gate) | +| 2 | 4.8 Parallel Handler | Stub -- no concurrent execution, no join/error policies | +| 3 | 4.11 Manager Loop | Stub -- always returns FAIL | +| 4 | 5.3 Checkpoint Resume | `Checkpoint::load` exists but engine has no resume-from-checkpoint logic | +| 5 | 5.4 Context Fidelity | Fidelity resolution and session/thread management not implemented | + +### MEDIUM (15 gaps -- functional but incomplete) + +| # | Location | Gap | +|---|----------|-----| +| 6 | 3.2 loop_restart | Not implemented -- jumps to target instead of restarting run | +| 7 | 3.5 should_retry | No retryable vs non-retryable error classification | +| 10 | 4.9 Fan-In | No LLM-based evaluation path | +| 11 | 4.12 Panic Safety | Handler panics not caught by engine | +| 12 | 5.6 Run Directory | Missing `manifest.json`; per-node dirs only from CodergenHandler | +| 13 | 6.4 ConsoleInterviewer | Not implemented | +| 14 | 6.5 Timeout Handling | No runtime timeout enforcement | +| 15 | 7.2 start_node rule | Only checks shape=Mdiamond, not ID-based fallback | +| 16 | 7.2 terminal_node rule | Only checks shape=Msquare, not ID-based fallback | +| 17 | 7.3/7.4 Custom Rules | LintRule trait exists but no registration or extra_rules param | +| 18 | 8.2/11.10 Shape Selector | Stylesheet only supports `*`, `.class`, `#id` -- no shape selector | +| 19 | 9.1 prepare_pipeline | No function chaining parse -> transforms -> validate | +| 20 | 9.2 Preamble Transform | Not implemented | +| 21 | 9.3 register_transform | No API on engine for registering transforms | +| 22 | 9.6 Event Emission | Parallel and interview events defined but never emitted | +| 23 | 9.7 Tool Call Hooks | pre/post hooks not read or executed | + +### LOW (7 gaps -- minor, cosmetic, or optional) + +| # | Location | Gap | +|---|----------|-----| +| 28 | 2.10 Subgraph label | Class derivation misses `graph [label=...]` block form | +| 30 | 4.10 Tool Timeout | No command timeout on shell execution | +| 31 | 8.2 ClassName | Accepts uppercase; spec says `[a-z0-9-]+` | +| 33 | 10.5 Bare Key | Returns error instead of truthiness check | +| 35 | 11.12 Test Coverage | Missing integration tests for retry, resume, 10+ nodes | +| 36 | 11.13 Smoke Test | No real LLM end-to-end test | + +--- + +**Overall: 58 of 85 items ALIGNED. 27 legitimate gaps (5 HIGH, 15 MEDIUM, 7 LOW).** diff --git a/docs/agent/reviews/spec-sections-5-6-review.md b/docs/agent/reviews/spec-sections-5-6-review.md new file mode 100644 index 000000000..36c886c5b --- /dev/null +++ b/docs/agent/reviews/spec-sections-5-6-review.md @@ -0,0 +1,250 @@ +# Spec Compliance Review: Sections 5-6 + +## Section 5: State and Context + +### 5.1 Context + +**ALIGNED** (with minor gaps) + +The `Context` struct in `/Users/bhelmkamp/p/brynary/attractor-rust/crates/attractor/src/context.rs` correctly implements: + +- Thread-safe key-value store using `Arc>>` (context.rs:9) +- Append-only logs using `Arc>>` (context.rs:10) +- `set(key, value)` with write lock (context.rs:33-38) +- `get(key)` with read lock, returns `Option` (context.rs:46-52) -- spec says `default=NONE` which maps to Rust's `Option::None` +- `get_string(key, default)` with string coercion (context.rs:56-60) +- `append_log(entry)` with write lock (context.rs:67-72) +- `snapshot()` returning a cloned map (context.rs:80-85) +- `clone_context()` for deep copy / parallel isolation (context.rs:99-106) -- called `clone_context` instead of `clone` to avoid conflict with Rust's `Clone` trait +- `apply_updates(updates)` merging a map into context (context.rs:113-118) + +**Minor gap**: `get()` does not accept a `default` parameter like the spec's `get(key, default=NONE)`. The Rust version returns `Option` instead, which is idiomatic but means callers must handle the default themselves. This is an acceptable Rust adaptation. + +**Built-in context keys** set by the engine: + +| Key | Status | Evidence | +|-----|--------|----------| +| `outcome` | ALIGNED | engine.rs:533 sets `context.set("outcome", ...)` | +| `preferred_label` | ALIGNED | engine.rs:535 sets `context.set("preferred_label", ...)` | +| `graph.goal` | ALIGNED | engine.rs:349-351 mirrors graph goal | +| `current_node` | ALIGNED | engine.rs:492 sets `context.set("current_node", ...)` | +| `last_stage` | ALIGNED | codergen.rs:100-103 sets `last_stage` via context_updates | +| `last_response` | ALIGNED | codergen.rs:104-107 sets `last_response` (truncated to 200 chars) | +| `internal.retry_count.` | **GAP** | Not implemented anywhere. The engine tracks retry attempts locally in `execute_with_retry` but never writes `internal.retry_count.` to the context. | + +**Context key namespace conventions**: The code uses `graph.*` namespace (engine.rs:354) and `context.*` would be user-driven. No enforcement of namespaces exists (which is expected -- they're conventions). + +### 5.2 Outcome + +**ALIGNED** + +The `Outcome` struct in `/Users/bhelmkamp/p/brynary/attractor-rust/crates/attractor/src/outcome.rs` matches the spec exactly: + +- `status: StageStatus` (outcome.rs:49) -- all five values present: `Success`, `Fail`, `PartialSuccess`, `Retry`, `Skipped` (outcome.rs:10-16) +- `preferred_label: Option` (outcome.rs:51) +- `suggested_next_ids: Vec` (outcome.rs:53) +- `context_updates: HashMap` (outcome.rs:55) +- `notes: Option` (outcome.rs:57) +- `failure_reason: Option` (outcome.rs:59) + +Factory methods: `success()`, `fail(reason)`, `retry(reason)`, `skipped()` all present (outcome.rs:63-108). + +Serialization with `serde` roundtrips correctly (outcome.rs:173-188). + +### 5.3 Checkpoint + +**ALIGNED** (with minor gaps) + +The `Checkpoint` struct in `/Users/bhelmkamp/p/brynary/attractor-rust/crates/attractor/src/checkpoint.rs` implements: + +- `timestamp: DateTime` (checkpoint.rs:14) +- `current_node: String` (checkpoint.rs:15) +- `completed_nodes: Vec` (checkpoint.rs:16) +- `node_retries: HashMap` (checkpoint.rs:17) +- `context_values: HashMap` (checkpoint.rs:18) +- `logs: Vec` (checkpoint.rs:19) +- `save(path)` serializes to JSON (checkpoint.rs:44-49) +- `load(path)` deserializes from JSON (checkpoint.rs:56-61) + +**GAP: node_retries not populated by engine**: The `Checkpoint::from_context` (checkpoint.rs:24-37) always initializes `node_retries` to an empty map. The engine (engine.rs:539-543) never populates retry counts into the checkpoint. Tests manually set retries (checkpoint.rs:103) but the engine never does. + +**GAP: Resume behavior not implemented**: Spec 5.3 describes a 6-step resume process (load checkpoint, restore context, restore completed_nodes, restore retry counters, determine next node, degrade fidelity). The engine has no `resume_from_checkpoint` method. The `Checkpoint::load` exists but nothing consumes it for resumption. + +### 5.4 Context Fidelity + +**GAP** (data model present, runtime not implemented) + +The spec defines `FidelityMode` with values: `full`, `truncate`, `compact`, `summary:low`, `summary:medium`, `summary:high`. + +- Graph types support `fidelity` attribute on nodes (graph/types.rs:159-160) and edges (graph/types.rs:257-258) +- Graph supports `default_fidelity` (graph/types.rs:368-372) +- Nodes and edges support `thread_id` attribute (graph/types.rs:164-165, 262-263) +- Validation rule `fidelity_valid` validates fidelity modes (validation/rules.rs:381-446) + +**However**: +- No `FidelityMode` enum exists as a first-class type -- fidelity is only a string attribute +- The fidelity resolution precedence (edge -> node -> graph default -> `compact`) is not implemented in the engine +- Thread resolution for `full` fidelity is not implemented +- The engine does not use fidelity to control context passing between nodes +- No session reuse / thread management exists + +### 5.5 Artifact Store + +**ALIGNED** + +The `ArtifactStore` in `/Users/bhelmkamp/p/brynary/attractor-rust/crates/attractor/src/artifact.rs` fully implements the spec: + +- `store(id, name, data) -> ArtifactInfo` (artifact.rs:62-96) with file-backing for large artifacts +- `retrieve(id) -> Value` (artifact.rs:107-130) reading from memory or disk +- `has(id) -> bool` (artifact.rs:137-142) +- `list() -> Vec` (artifact.rs:150-157) +- `remove(id)` (artifact.rs:164-169) including disk cleanup +- `clear()` (artifact.rs:176-184) including disk cleanup +- `FILE_BACKING_THRESHOLD = 100 * 1024` (artifact.rs:12) matching spec's 100KB + +`ArtifactInfo` fields match spec (artifact.rs:16-22): +- `id`, `name`, `size_bytes`, `stored_at`, `is_file_backed` + +Thread safety via `RwLock` (artifact.rs:33). + +### 5.6 Run Directory Structure + +**PARTIALLY ALIGNED** + +Spec directory structure: +``` +{logs_root}/ + checkpoint.json -- present (engine.rs:544) + manifest.json -- MISSING + {node_id}/ + status.json -- present (codergen.rs:82, 111) + prompt.md -- present (codergen.rs:74) + response.md -- present (codergen.rs:95) + artifacts/ + {artifact_id}.json -- present (artifact.rs:75) +``` + +**GAP: manifest.json**: The spec requires a `manifest.json` with pipeline metadata (name, goal, start time). This file is never written by the engine. + +**GAP: Per-node directories only for codergen**: Only the `CodergenHandler` creates `{node_id}/` subdirectories with `status.json`, `prompt.md`, and `response.md`. Other handlers (start, exit, tool, parallel, etc.) do not write any per-node log files. + +--- + +## Section 6: Human-in-the-Loop (Interviewer Pattern) + +### 6.1 Interviewer Interface + +**ALIGNED** + +The `Interviewer` trait in `/Users/bhelmkamp/p/brynary/attractor-rust/crates/attractor/src/interviewer/mod.rs` matches the spec: + +- `ask(question: Question) -> Answer` (mod.rs:133) +- `ask_multiple(questions: Vec) -> Vec` with default sequential implementation (mod.rs:135-140) +- `inform(message, stage)` with default no-op (mod.rs:143-145) + +The trait is `async` (using `#[async_trait]`) and requires `Send + Sync` (mod.rs:132), which is appropriate for Rust. + +### 6.2 Question Model + +**ALIGNED** + +`Question` struct (mod.rs:29-38): +- `text: String` +- `question_type: QuestionType` (named `question_type` instead of `type` since `type` is a Rust keyword) +- `options: Vec` +- `allow_freeform: bool` +- `default: Option` +- `timeout_seconds: Option` +- `stage: String` +- `metadata: HashMap` + +`QuestionType` enum (mod.rs:13-18): +- `YesNo`, `MultipleChoice`, `Freeform`, `Confirmation` -- all four spec variants present + +`QuestionOption` (mod.rs:21-24): +- `key: String`, `label: String` -- matches spec's `Option` (renamed to avoid Rust keyword collision) + +### 6.3 Answer Model + +**ALIGNED** + +`Answer` struct (mod.rs:68-72): +- `value: AnswerValue` +- `selected_option: Option` +- `text: Option` + +`AnswerValue` enum (mod.rs:57-64): +- `Yes`, `No`, `Skipped`, `Timeout` -- matches spec +- `Selected(String)` -- represents a multiple-choice selection (spec used `value: String`) +- `Text(String)` -- represents freeform text + +The spec uses a single `value` field that can be either an `AnswerValue` enum or a string. The Rust implementation cleanly separates these via enum variants, which is a good adaptation. + +### 6.4 Built-In Interviewer Implementations + +**AutoApproveInterviewer: ALIGNED** + +`/Users/bhelmkamp/p/brynary/attractor-rust/crates/attractor/src/interviewer/auto_approve.rs`: +- YesNo/Confirmation -> `Answer::yes()` (auto_approve.rs:12) +- MultipleChoice -> first option or "auto-approved" text (auto_approve.rs:13-20) +- Freeform -> `Answer::text("auto-approved")` (auto_approve.rs:21) +- Matches spec pseudocode exactly + +**ConsoleInterviewer: GAP (not implemented)** + +No `ConsoleInterviewer` exists in the codebase. Grep for `ConsoleInterviewer` returns no matches. The spec describes a CLI-based interviewer that reads from stdin. + +**CallbackInterviewer: ALIGNED** + +`/Users/bhelmkamp/p/brynary/attractor-rust/crates/attractor/src/interviewer/callback.rs`: +- Accepts a `Fn(Question) -> Answer` callback (callback.rs:7) +- `ask()` delegates to callback (callback.rs:21) +- Matches spec exactly + +**QueueInterviewer: ALIGNED** + +`/Users/bhelmkamp/p/brynary/attractor-rust/crates/attractor/src/interviewer/queue.rs`: +- Pre-filled `VecDeque` (queue.rs:10) +- `ask()` dequeues or returns `Answer::skipped()` (queue.rs:24-26) +- Thread-safe via `Mutex` (queue.rs:10) +- Matches spec exactly + +**RecordingInterviewer: ALIGNED** + +`/Users/bhelmkamp/p/brynary/attractor-rust/crates/attractor/src/interviewer/recording.rs`: +- Wraps an inner `Box` (recording.rs:9) +- Records `(Question, Answer)` pairs in `Mutex>` (recording.rs:10) +- `ask()` delegates to inner, then records (recording.rs:32-38) +- `recordings()` accessor (recording.rs:25-27) +- Matches spec exactly + +### 6.5 Timeout Handling + +**GAP** (partially modeled, not implemented at runtime) + +- The `Question` struct has a `timeout_seconds: Option` field (mod.rs:35) +- The `Answer` has a `timeout()` factory (mod.rs:103-108) and `AnswerValue::Timeout` variant (mod.rs:61) +- The `Question` struct has a `default: Option` field (mod.rs:34) + +**However**: +- No interviewer implementation actually enforces timeouts (no tokio timeout wrapper) +- The spec's timeout behavior (use default if available, else return Timeout) is not implemented in any interviewer +- `wait.human` node's `human.default_choice` for timeout behavior is not checked at runtime + +--- + +## Summary of Gaps + +| Section | Status | Gap Description | +|---------|--------|-----------------| +| 5.1 | ALIGNED (minor) | `internal.retry_count.` context key never written by engine | +| 5.2 | ALIGNED | Fully matches spec | +| 5.3 | GAP | Checkpoint save works but resume from checkpoint not implemented; node_retries never populated by engine | +| 5.4 | GAP | Fidelity attributes parsed and validated, but fidelity resolution/application not in engine; no session/thread management | +| 5.5 | ALIGNED | Fully matches spec | +| 5.6 | GAP | Missing `manifest.json`; per-node directories only created by CodergenHandler, not other handlers | +| 6.1 | ALIGNED | Trait matches spec interface | +| 6.2 | ALIGNED | Question model complete | +| 6.3 | ALIGNED | Answer model complete | +| 6.4 | GAP | Missing `ConsoleInterviewer`; other four implementations aligned | +| 6.5 | GAP | Timeout data model present but no runtime enforcement in any interviewer | diff --git a/docs/agent/reviews/spec-sections-7-9-review.md b/docs/agent/reviews/spec-sections-7-9-review.md new file mode 100644 index 000000000..7986ee841 --- /dev/null +++ b/docs/agent/reviews/spec-sections-7-9-review.md @@ -0,0 +1,197 @@ +# Spec Compliance Review: Sections 7-9 (Subagents + Definition of Done) + +## Section 7: Subagents + +### 7.1 Concept -- ALIGNED +- `SubAgent` in `subagent.rs` spawns a child session via `SubAgentManager::spawn()` which takes a `Session` and runs `session.process_input()` in a tokio task. +- The child session has its own conversation history (its own `History` instance). +- The child session shares the parent's execution environment (passed through the `SessionFactory` / session construction). + +### 7.2 Spawn Interface -- ALIGNED +All four tools are implemented in `subagent.rs`: +- `spawn_agent`: Correct params (`task` required, `working_dir`/`model`/`max_turns` optional). Returns agent ID. +- `send_input`: Correct params (`agent_id`, `message` required). Returns acknowledgement. +- `wait`: Correct params (`agent_id` required). Returns `SubAgentResult` (output, success, turns_used). +- `close_agent`: Correct params (`agent_id` required). Returns final status. + +**GAP**: `spawn_agent` tool executor does not use the `working_dir`, `model`, or `max_turns` optional parameters. They are defined in the schema but ignored in the executor at line 185-197. The session factory creates a default session regardless of these overrides. + +### 7.3 SubAgent Lifecycle -- ALIGNED (with one minor gap) +- `SubAgentResult` record: matches spec (`output: String`, `success: bool`, `turns_used: usize`). +- `SubAgent` struct has `id` and `depth` fields but no explicit `status` enum (`"running" | "completed" | "failed"`). Status is implicit via whether the tokio task is running. +- `SubAgentHandle` is not a separate record; `SubAgent` serves this role. +- Depth limiting: Implemented in `SubAgentManager::spawn()` at line 56 (`depth >= self.max_depth`). Default `max_subagent_depth: 1` in `SessionConfig`. +- Independent history: Each subagent gets its own `Session` with its own `History`. + +**GAP**: No explicit `SubAgentHandle` record with a `status` field as spec defines. Status is implicit. + +### 7.4 Use Cases -- ALIGNED +The architecture supports all listed use cases (parallel exploration, focused refactoring, test execution, alternative approaches) through the spawn/wait/close interface. Subagents run as independent tokio tasks sharing the execution environment. + +--- + +## Section 9: Definition of Done + +### 9.1 Core Loop + +| Item | Status | Evidence | +|------|--------|----------| +| Session created with ProviderProfile + ExecutionEnvironment | DONE | `Session::new(client, profile, env, config)` in `session.rs:37` | +| `process_input()` runs agentic loop | DONE | `session.rs:118-157` -- LLM call -> tool exec -> loop | +| Natural completion (text only, no tool calls) | DONE | `session.rs:259-261` -- breaks when `tool_calls.is_empty()` | +| Round limits (`max_tool_rounds_per_input`) | DONE | `session.rs:184-191` -- checked each iteration | +| Session turn limits (`max_turns`) | DONE | `session.rs:194-201` -- checked each iteration | +| Abort signal -> CLOSED | DONE | `session.rs:204-207` -- checks `abort_flag`, transitions to Closed | +| Loop detection -> warning SteeringTurn | DONE | `session.rs:278-290` -- calls `detect_loop`, injects Steering turn | +| Multiple sequential inputs | DONE | Test `sequential_inputs` in `session.rs:1437-1457` confirms this works | + +**Result: 8/8 DONE** + +### 9.2 Provider Profiles + +| Item | Status | Evidence | +|------|--------|----------| +| OpenAI profile with `apply_patch` (v4a) | DONE | `profiles/openai.rs` -- registers `apply_patch`, full v4a parser + applier | +| Anthropic profile with `edit_file` (old_string/new_string) | DONE | `profiles/anthropic.rs` -- registers `edit_file` tool | +| Gemini profile with gemini-cli-aligned tools | DONE | `profiles/gemini.rs` -- registers read/write/edit/shell/grep/glob | +| Each profile has provider-specific system prompt | DONE | Each profile's `build_system_prompt()` includes identity + env context + tool guidance | +| Custom tools can be registered | DONE | `tool_registry_mut()` exposed on `ProviderProfile` trait, `ToolRegistry::register()` available | +| Tool name collisions resolved (override) | DONE | `ToolRegistry::register()` uses `HashMap::insert` which overwrites. Test `name_collision_overrides` in `tool_registry.rs:111` | + +**Result: 6/6 DONE** + +### 9.3 Tool Execution + +| Item | Status | Evidence | +|------|--------|----------| +| Tool calls dispatched through ToolRegistry | DONE | `session.rs:352-353` -- `registry.get(tool_name)` | +| Unknown tool -> error result to LLM | DONE | `session.rs:386-393` -- returns `is_error: true` with "Unknown tool" | +| Tool argument JSON validated against schema | DONE | `session.rs:356-366` -- `validate_tool_args()` using `jsonschema` crate | +| Tool execution errors caught and returned as error results | DONE | `session.rs:377-384` -- `Err(err)` mapped to `is_error: true` | +| Parallel tool execution when `supports_parallel_tool_calls` | DONE | `session.rs:400-404` -- routes to `execute_tool_calls_parallel` when supported | + +**Result: 5/5 DONE** + +### 9.4 Execution Environment + +| Item | Status | Evidence | +|------|--------|----------| +| `LocalExecutionEnvironment` implements all file/command ops | DONE | `local_env.rs` -- read/write/exists/list/exec/grep/glob all implemented | +| Command timeout default is 10 seconds | DONE | `config.rs:23` -- `default_command_timeout_ms: 10_000` | +| Command timeout overridable per-call via `timeout_ms` param | DONE | `tools.rs:180-184` -- shell tool reads `timeout_ms` from args | +| Timed-out: SIGTERM then SIGKILL after 2 seconds | DONE | `local_env.rs:152-173` -- sends SIGTERM, waits 2s, then SIGKILL | +| Env var filtering excludes sensitive variables | DONE | `local_env.rs:31-38` -- filters `*_API_KEY`, `*_SECRET`, `*_TOKEN`, `*_PASSWORD`, `*_CREDENTIAL` | +| `ExecutionEnvironment` interface implementable by consumers | DONE | `execution_env.rs:27` -- `trait ExecutionEnvironment: Send + Sync` with all methods | + +**Result: 6/6 DONE** + +### 9.5 Tool Output Truncation + +| Item | Status | Evidence | +|------|--------|----------| +| Character-based truncation runs FIRST | DONE | `truncation.rs:99-109` -- char truncation applied first | +| Line-based truncation runs SECOND (shell:256, grep:200, glob:500) | DONE | `truncation.rs:111-121` -- line truncation after chars; `default_line_limits()` has correct values | +| Truncation inserts visible marker | DONE | `truncation.rs:54-55, 62-63` -- `[WARNING: Output truncated...]` markers | +| Full untruncated output in `TOOL_CALL_END` event | DONE | `session.rs:554-571` -- emits event with full output BEFORE truncation | +| Default char limits match spec Section 5.2 | DONE | `truncation.rs:10-21` -- read_file:50k, shell:30k, grep:20k, glob:20k, edit_file:10k, write_file:1k | +| Both char and line limits overridable via `SessionConfig` | DONE | `config.rs:10-11` -- `tool_output_limits` and `tool_line_limits` HashMaps; `truncation.rs:100-103, 112-116` checks config first | + +**Result: 6/6 DONE** + +### 9.6 Steering + +| Item | Status | Evidence | +|------|--------|----------| +| `steer()` queues message injected after current tool round | DONE | `session.rs:80-85` -- pushes to `steering_queue`; `session.rs:275` -- `drain_steering()` called after tool execution | +| `follow_up()` queues message processed after current input completes | DONE | `session.rs:87-92` -- pushes to `followup_queue`; `session.rs:136-146` -- processed after `run_single_input` | +| Steering messages appear as SteeringTurn in history | DONE | `session.rs:304` -- `Turn::Steering` pushed to history | +| SteeringTurns converted to user-role messages for LLM | DONE | `history.rs:75-81` -- `Turn::Steering` maps to `Role::User` message | + +**Result: 4/4 DONE** + +### 9.7 Reasoning Effort + +| Item | Status | Evidence | +|------|--------|----------| +| `reasoning_effort` passed through to LLM SDK Request | DONE | `session.rs:340` -- `reasoning_effort: self.config.reasoning_effort.clone()` | +| Changing mid-session takes effect on next LLM call | DONE | `session.rs:110-112` -- `set_reasoning_effort()` mutates config; test at line 1722 confirms | +| Valid values: "low", "medium", "high", null | DONE | Stored as `Option` and passed through to SDK; no validation in this layer (SDK handles it) | + +**Result: 3/3 DONE** + +### 9.8 System Prompts + +| Item | Status | Evidence | +|------|--------|----------| +| Provider-specific base instructions | DONE | Each profile has distinct identity text ("You are Claude...", "You are a coding assistant", "powered by Gemini") | +| Environment context (platform, git, working dir, date, model) | DONE | `profiles/mod.rs:21-48` -- `build_env_context_block` includes platform, working_directory, OS version, git branch, date, model | +| Tool descriptions from active profile | DONE | `session.rs:323` -- `self.provider_profile.tools()` included in request | +| Project docs (AGENTS.md + provider-specific) discovered and included | DONE | `project_docs.rs` -- discovers AGENTS.md plus provider-specific files | +| User instruction overrides appended last | NOT DONE | No mechanism for user instruction overrides in the system prompt pipeline. The `build_system_prompt` method appends project docs but has no separate "user overrides" parameter. | +| Only relevant project files loaded per provider | DONE | `project_docs.rs:13-18` -- filters by provider_id: anthropic gets CLAUDE.md, openai gets .codex/instructions.md, gemini gets GEMINI.md | + +**Result: 5/6 DONE** + +**GAP**: No explicit user instruction override mechanism in the system prompt. The spec says "User instruction overrides are appended last (highest priority)". The current `build_system_prompt` takes `project_docs` but has no separate parameter or config field for user-supplied instruction overrides. + +### 9.9 Subagents + +| Item | Status | Evidence | +|------|--------|----------| +| Subagents spawned with scoped task via `spawn_agent` tool | DONE | `subagent.rs:153-200` | +| Subagents share parent's execution environment | DONE | Session factory creates session with shared env | +| Subagents maintain independent conversation history | DONE | Each Session has its own History | +| Depth limiting prevents recursive spawning (default max: 1) | DONE | `subagent.rs:56` and `config.rs:30` -- `max_subagent_depth: 1` | +| Subagent results returned to parent as tool results | DONE | `wait` tool returns formatted result string | +| `send_input`, `wait`, `close_agent` tools work correctly | DONE | All three tools implemented with correct params and tested | + +**GAP (minor)**: The `spawn_agent` tool ignores `working_dir`, `model`, and `max_turns` optional parameters. They are in the schema but not wired to session creation. + +**Result: 6/6 DONE** (core behavior works; optional param wiring is a gap but not a DoD blocker) + +### 9.10 Event System + +| Item | Status | Evidence | +|------|--------|----------| +| All event kinds from Section 2.9 emitted at correct times | PARTIAL | Most events are emitted. `ASSISTANT_TEXT_START` is defined in `EventKind` enum but never emitted in `session.rs`. `CONTEXT_WINDOW_WARNING` is emitted but not in the spec's enum (it's an extension). | +| Events delivered via async iterator / equivalent | DONE | `event.rs` -- `tokio::sync::broadcast` channel with `subscribe()` returning `Receiver` | +| `TOOL_CALL_END` events carry full untruncated output | DONE | `session.rs:554-571` -- emits before truncation | +| Session lifecycle events (SESSION_START, SESSION_END) bracket session | DONE | `session.rs:123-127` emits SessionStart; `session.rs:150-154` emits SessionEnd | + +**Result: 3/4 DONE, 1 PARTIAL** + +**GAP**: `AssistantTextStart` event kind is defined in the enum but never emitted anywhere in the session code. The spec lists `ASSISTANT_TEXT_START` as a required event. + +### 9.11 Error Handling + +| Item | Status | Evidence | +|------|--------|----------| +| Tool execution errors -> error result sent to LLM | DONE | `session.rs:377-384` -- tool errors returned as `is_error: true` ToolResult | +| LLM API transient errors -> retry with backoff (via SDK) | DONE | Spec explicitly says "handled by Unified LLM SDK layer" | +| Authentication errors -> surface immediately, session CLOSED | DONE | `session.rs:226-229` -- `is_auth_error()` check, transitions to Closed | +| Context window overflow -> emit warning event | DONE | `session.rs:629-654` -- `check_context_usage()` emits `ContextWindowWarning` | +| Graceful shutdown: abort -> cancel -> kill -> flush -> SESSION_END | PARTIAL | Abort flag checked, returns `AgentError::Aborted`, transitions to Closed. But `SESSION_END` is NOT emitted on abort (the abort short-circuits before the `emit(SessionEnd)` call). Also no explicit process killing on abort -- the session just stops looping. | + +**Result: 4/5 DONE, 1 PARTIAL** + +**GAP**: On abort, `SESSION_END` event is not emitted. The abort at `session.rs:204-207` returns an `Err(AgentError::Aborted)` which skips the `SessionEnd` emit at line 150-154. Running processes are not explicitly killed on abort either (only the loop stops). + +--- + +## Summary of All Gaps + +### Functional Gaps (should fix): + +1. **Spawn agent ignores optional params** (`subagent.rs:185-197`): `working_dir`, `model`, `max_turns` params are in the tool schema but the executor does not use them when creating the session. The session factory ignores these overrides. + +2. **`AssistantTextStart` event never emitted** (`session.rs`): The event kind exists in the enum but is never emitted. Should be emitted before/when the LLM starts generating text. + +3. **No `SESSION_END` event on abort** (`session.rs:204-207`): When abort triggers, the method returns early with `Err(AgentError::Aborted)` without emitting `SESSION_END`. The spec says graceful shutdown should "flush events -> emit SESSION_END". + +4. **No user instruction overrides in system prompt** (`session.rs` / `ProviderProfile`): Spec 9.8 item 5 says "User instruction overrides are appended last (highest priority)". There is no mechanism to pass user instruction overrides into the system prompt pipeline. `SessionConfig` lacks an `instructions` or `user_overrides` field. + +### Minor / Non-blocking Gaps: + +5. **No explicit `SubAgentHandle` with `status` field**: Status is implicit based on tokio task state rather than an explicit enum field. Functionally equivalent but structurally different from spec. + +6. **`SESSION_END` not emitted on abort path for running processes**: No explicit kill of running child processes on abort. The session just stops the loop, but any child process from `exec_command` may continue running. The `close()` method for subagents does handle this properly.