mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-09 03:20:56 +00:00
Implement attractor crate: DOT-based pipeline runner with full spec compliance
Adds the attractor crate implementing all 11 sections of the attractor spec: - DOT parser (lexer, grammar, semantic analysis) for strict DOT subset - Pipeline execution engine with edge selection, goal gates, retry logic, failure routing, checkpoint save/resume, and loop_restart - 9 node handlers: start, exit, codergen, wait_human, conditional, parallel (concurrent with join/error policies), fan_in (with LLM eval), tool, manager_loop - State management: PipelineContext, Outcome, Artifact store, fidelity resolution - Human-in-the-loop: Interviewer trait with auto_approve, callback, queue, recording, and console implementations, plus timeout enforcement - Validation: 14 built-in lint rules with custom rule registration API - Model stylesheet with universal/shape/class/ID selectors and specificity - Transforms: variable expansion, stylesheet application, preamble; plus PipelineBuilder with register_transform and prepare_pipeline - Condition expression language with =, !=, bare-key truthiness, && combinator - Event system with all 16 event types emitted by engine and handlers - Tool call hooks (pre/post) for CodergenHandler - Run directory with manifest.json and per-node status.json 370 tests (354 unit + 16 integration) covering all spec sections. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
db32f00cd0
commit
34836200e5
44 changed files with 12563 additions and 0 deletions
1
Cargo.lock
generated
1
Cargo.lock
generated
|
|
@ -135,6 +135,7 @@ dependencies = [
|
|||
"async-trait",
|
||||
"chrono",
|
||||
"coding-agent-loop",
|
||||
"futures",
|
||||
"nom",
|
||||
"rand",
|
||||
"serde",
|
||||
|
|
|
|||
31
crates/attractor/Cargo.toml
Normal file
31
crates/attractor/Cargo.toml
Normal file
|
|
@ -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
|
||||
3
crates/attractor/README.md
Normal file
3
crates/attractor/README.md
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
# attractor
|
||||
|
||||
A DOT-based pipeline runner for multi-stage AI workflows.
|
||||
300
crates/attractor/src/artifact.rs
Normal file
300
crates/attractor/src/artifact.rs
Normal file
|
|
@ -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<Utc>,
|
||||
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<PathBuf>,
|
||||
artifacts: RwLock<HashMap<String, (ArtifactInfo, StoredData)>>,
|
||||
}
|
||||
|
||||
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<PathBuf>) -> 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<String>, name: impl Into<String>, data: Value) -> Result<ArtifactInfo> {
|
||||
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<Value> {
|
||||
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<ArtifactInfo> {
|
||||
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);
|
||||
}
|
||||
}
|
||||
143
crates/attractor/src/checkpoint.rs
Normal file
143
crates/attractor/src/checkpoint.rs
Normal file
|
|
@ -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<Utc>,
|
||||
pub current_node: String,
|
||||
pub completed_nodes: Vec<String>,
|
||||
pub node_retries: HashMap<String, u32>,
|
||||
pub context_values: HashMap<String, Value>,
|
||||
pub logs: Vec<String>,
|
||||
}
|
||||
|
||||
impl Checkpoint {
|
||||
/// Create a checkpoint from the current execution state.
|
||||
pub fn from_context(
|
||||
context: &Context,
|
||||
current_node: impl Into<String>,
|
||||
completed_nodes: Vec<String>,
|
||||
) -> 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<Self> {
|
||||
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");
|
||||
}
|
||||
}
|
||||
324
crates/attractor/src/condition.rs
Normal file
324
crates/attractor/src/condition.rs
Normal file
|
|
@ -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<Vec<Clause>, 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
|
||||
));
|
||||
}
|
||||
}
|
||||
225
crates/attractor/src/context.rs
Normal file
225
crates/attractor/src/context.rs
Normal file
|
|
@ -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<RwLock<HashMap<String, Value>>>,
|
||||
logs: Arc<RwLock<Vec<String>>>,
|
||||
}
|
||||
|
||||
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<String>, 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<Value> {
|
||||
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<String>) {
|
||||
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<String, Value> {
|
||||
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<String> {
|
||||
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<String, Value>) {
|
||||
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());
|
||||
}
|
||||
}
|
||||
1570
crates/attractor/src/engine.rs
Normal file
1570
crates/attractor/src/engine.rs
Normal file
File diff suppressed because it is too large
Load diff
91
crates/attractor/src/error.rs
Normal file
91
crates/attractor/src/error.rs
Normal file
|
|
@ -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<std::io::Error> for AttractorError {
|
||||
fn from(err: std::io::Error) -> Self {
|
||||
Self::Io(err.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
pub type Result<T> = std::result::Result<T, AttractorError>;
|
||||
|
||||
#[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<i32> = Ok(42);
|
||||
assert_eq!(ok.unwrap(), 42);
|
||||
|
||||
let err: Result<i32> = Err(AttractorError::Parse("bad".to_string()));
|
||||
assert!(err.is_err());
|
||||
}
|
||||
}
|
||||
165
crates/attractor/src/event.rs
Normal file
165
crates/attractor/src/event.rs
Normal file
|
|
@ -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<dyn Fn(&PipelineEvent) + Send + Sync>;
|
||||
|
||||
/// Callback-based event emitter for pipeline events.
|
||||
pub struct EventEmitter {
|
||||
listeners: Vec<EventListener>,
|
||||
}
|
||||
|
||||
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);
|
||||
}
|
||||
}
|
||||
3
crates/attractor/src/graph/mod.rs
Normal file
3
crates/attractor/src/graph/mod.rs
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
pub mod types;
|
||||
|
||||
pub use types::*;
|
||||
620
crates/attractor/src/graph/types.rs
Normal file
620
crates/attractor/src/graph/types.rs
Normal file
|
|
@ -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<i64> {
|
||||
match self {
|
||||
Self::Integer(n) => Some(*n),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub const fn as_f64(&self) -> Option<f64> {
|
||||
match self {
|
||||
Self::Float(n) => Some(*n),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub const fn as_bool(&self) -> Option<bool> {
|
||||
match self {
|
||||
Self::Boolean(b) => Some(*b),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub const fn as_duration(&self) -> Option<Duration> {
|
||||
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<String, AttrValue>,
|
||||
/// CSS-like classes for model stylesheet targeting (from `class` attr and subgraph derivation).
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub classes: Vec<String>,
|
||||
}
|
||||
|
||||
impl Node {
|
||||
pub fn new(id: impl Into<String>) -> 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<bool> {
|
||||
self.attrs.get(key).and_then(AttrValue::as_bool)
|
||||
}
|
||||
|
||||
fn int_attr(&self, key: &str) -> Option<i64> {
|
||||
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<i64> {
|
||||
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<Duration> {
|
||||
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<String, AttrValue>,
|
||||
}
|
||||
|
||||
impl Edge {
|
||||
pub fn new(from: impl Into<String>, to: impl Into<String>) -> 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<bool> {
|
||||
self.attrs.get(key).and_then(AttrValue::as_bool)
|
||||
}
|
||||
|
||||
fn int_attr(&self, key: &str) -> Option<i64> {
|
||||
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<String, Node>,
|
||||
pub edges: Vec<Edge>,
|
||||
pub attrs: HashMap<String, AttrValue>,
|
||||
}
|
||||
|
||||
impl Graph {
|
||||
pub fn new(name: impl Into<String>) -> 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());
|
||||
}
|
||||
}
|
||||
383
crates/attractor/src/handler/codergen.rs
Normal file
383
crates/attractor/src/handler/codergen.rs
Normal file
|
|
@ -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<CodergenResult, AttractorError>;
|
||||
}
|
||||
|
||||
/// The default handler for LLM task nodes.
|
||||
pub struct CodergenHandler {
|
||||
backend: Option<Box<dyn CodergenBackend>>,
|
||||
}
|
||||
|
||||
impl CodergenHandler {
|
||||
#[must_use]
|
||||
pub fn new(backend: Option<Box<dyn CodergenBackend>>) -> 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<String> {
|
||||
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<Outcome, AttractorError> {
|
||||
// 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);
|
||||
}
|
||||
}
|
||||
52
crates/attractor/src/handler/conditional.rs
Normal file
52
crates/attractor/src/handler/conditional.rs
Normal file
|
|
@ -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<Outcome, AttractorError> {
|
||||
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")
|
||||
);
|
||||
}
|
||||
}
|
||||
45
crates/attractor/src/handler/exit.rs
Normal file
45
crates/attractor/src/handler/exit.rs
Normal file
|
|
@ -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<Outcome, AttractorError> {
|
||||
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);
|
||||
}
|
||||
}
|
||||
341
crates/attractor/src/handler/fan_in.rs
Normal file
341
crates/attractor/src/handler/fan_in.rs
Normal file
|
|
@ -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<Box<dyn CodergenBackend>>,
|
||||
}
|
||||
|
||||
impl FanInHandler {
|
||||
#[must_use]
|
||||
pub fn new(backend: Option<Box<dyn CodergenBackend>>) -> Self {
|
||||
Self { backend }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Handler for FanInHandler {
|
||||
async fn execute(
|
||||
&self,
|
||||
node: &Node,
|
||||
context: &Context,
|
||||
_graph: &Graph,
|
||||
_logs_root: &Path,
|
||||
) -> Result<Outcome, AttractorError> {
|
||||
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<Candidate> = 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<Candidate, AttractorError> {
|
||||
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<CodergenResult, AttractorError> {
|
||||
// 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"))
|
||||
);
|
||||
}
|
||||
}
|
||||
309
crates/attractor/src/handler/manager_loop.rs
Normal file
309
crates/attractor/src/handler/manager_loop.rs
Normal file
|
|
@ -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<Box<dyn ChildObserver>>,
|
||||
}
|
||||
|
||||
impl ManagerLoopHandler {
|
||||
#[must_use]
|
||||
pub fn new(observer: Option<Box<dyn ChildObserver>>) -> 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::<u64>() {
|
||||
return Duration::from_millis(val);
|
||||
}
|
||||
} else if let Ok(val) = secs.parse::<u64>() {
|
||||
return Duration::from_secs(val);
|
||||
}
|
||||
}
|
||||
if let Some(mins) = s.strip_suffix('m') {
|
||||
if let Ok(val) = mins.parse::<u64>() {
|
||||
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<Outcome, AttractorError> {
|
||||
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));
|
||||
}
|
||||
}
|
||||
178
crates/attractor/src/handler/mod.rs
Normal file
178
crates/attractor/src/handler/mod.rs
Normal file
|
|
@ -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<Outcome, AttractorError>;
|
||||
}
|
||||
|
||||
/// Maps handler type strings to handler implementations.
|
||||
pub struct HandlerRegistry {
|
||||
handlers: HashMap<String, Box<dyn Handler>>,
|
||||
default_handler: Box<dyn Handler>,
|
||||
}
|
||||
|
||||
impl HandlerRegistry {
|
||||
#[must_use]
|
||||
pub fn new(default_handler: Box<dyn Handler>) -> Self {
|
||||
Self {
|
||||
handlers: HashMap::new(),
|
||||
default_handler,
|
||||
}
|
||||
}
|
||||
|
||||
/// Register a handler for a given type string.
|
||||
pub fn register(&mut self, type_string: impl Into<String>, handler: Box<dyn Handler>) {
|
||||
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<Outcome, AttractorError> {
|
||||
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;
|
||||
}
|
||||
}
|
||||
462
crates/attractor/src/handler/parallel.rs
Normal file
462
crates/attractor/src/handler/parallel.rs
Normal file
|
|
@ -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<HandlerRegistry>,
|
||||
emitter: Arc<EventEmitter>,
|
||||
}
|
||||
|
||||
impl ParallelHandler {
|
||||
#[must_use]
|
||||
pub fn new(registry: Arc<HandlerRegistry>, emitter: Arc<EventEmitter>) -> 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::<usize>() {
|
||||
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::<f64>() {
|
||||
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<Outcome, AttractorError> {
|
||||
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, AttractorError>(BranchResult {
|
||||
id: target_id,
|
||||
outcome,
|
||||
})
|
||||
});
|
||||
handles.push(handle);
|
||||
}
|
||||
|
||||
// Collect results
|
||||
let mut results: Vec<BranchResult> = 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<serde_json::Value> = 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<String> = 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<HandlerRegistry> {
|
||||
let registry = HandlerRegistry::new(Box::new(StartHandler));
|
||||
Arc::new(registry)
|
||||
}
|
||||
|
||||
fn make_emitter() -> Arc<EventEmitter> {
|
||||
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);
|
||||
}
|
||||
}
|
||||
45
crates/attractor/src/handler/start.rs
Normal file
45
crates/attractor/src/handler/start.rs
Normal file
|
|
@ -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<Outcome, AttractorError> {
|
||||
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);
|
||||
}
|
||||
}
|
||||
185
crates/attractor/src/handler/tool.rs
Normal file
185
crates/attractor/src/handler/tool.rs
Normal file
|
|
@ -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<Outcome, AttractorError> {
|
||||
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<Outcome, AttractorError> {
|
||||
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
|
||||
);
|
||||
}
|
||||
}
|
||||
413
crates/attractor/src/handler/wait_human.rs
Normal file
413
crates/attractor/src/handler/wait_human.rs
Normal file
|
|
@ -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<dyn Interviewer>,
|
||||
emitter: Option<Arc<EventEmitter>>,
|
||||
}
|
||||
|
||||
impl WaitHumanHandler {
|
||||
pub fn new(interviewer: Arc<dyn Interviewer>) -> Self {
|
||||
Self {
|
||||
interviewer,
|
||||
emitter: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn with_emitter(mut self, emitter: Arc<EventEmitter>) -> 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<Outcome, AttractorError> {
|
||||
// 1. Derive choices from outgoing edges
|
||||
let edges = graph.outgoing_edges(&node.id);
|
||||
let mut freeform_target: Option<String> = None;
|
||||
let mut choices: Vec<Choice> = 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<QuestionOption> = 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"))
|
||||
);
|
||||
}
|
||||
}
|
||||
88
crates/attractor/src/interviewer/auto_approve.rs
Normal file
88
crates/attractor/src/interviewer/auto_approve.rs
Normal file
|
|
@ -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()));
|
||||
}
|
||||
}
|
||||
56
crates/attractor/src/interviewer/callback.rs
Normal file
56
crates/attractor/src/interviewer/callback.rs
Normal file
|
|
@ -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<dyn Fn(Question) -> 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()));
|
||||
}
|
||||
}
|
||||
162
crates/attractor/src/interviewer/console.rs
Normal file
162
crates/attractor/src/interviewer/console.rs
Normal file
|
|
@ -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<Answer> {
|
||||
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::<usize>() {
|
||||
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<String> {
|
||||
// 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());
|
||||
}
|
||||
}
|
||||
297
crates/attractor/src/interviewer/mod.rs
Normal file
297
crates/attractor/src/interviewer/mod.rs
Normal file
|
|
@ -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<QuestionOption>,
|
||||
pub allow_freeform: bool,
|
||||
pub default: Option<Answer>,
|
||||
pub timeout_seconds: Option<f64>,
|
||||
pub stage: String,
|
||||
pub metadata: HashMap<String, serde_json::Value>,
|
||||
}
|
||||
|
||||
impl Question {
|
||||
pub fn new(text: impl Into<String>, 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<QuestionOption>,
|
||||
pub text: Option<String>,
|
||||
}
|
||||
|
||||
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<String>, option: QuestionOption) -> Self {
|
||||
let key = key.into();
|
||||
Self {
|
||||
value: AnswerValue::Selected(key),
|
||||
selected_option: Some(option),
|
||||
text: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn text(text: impl Into<String>) -> 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<Question>) -> Vec<Answer> {
|
||||
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);
|
||||
}
|
||||
}
|
||||
66
crates/attractor/src/interviewer/queue.rs
Normal file
66
crates/attractor/src/interviewer/queue.rs
Normal file
|
|
@ -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<VecDeque<Answer>>,
|
||||
}
|
||||
|
||||
impl QueueInterviewer {
|
||||
#[must_use]
|
||||
pub const fn new(answers: VecDeque<Answer>) -> 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);
|
||||
}
|
||||
}
|
||||
84
crates/attractor/src/interviewer/recording.rs
Normal file
84
crates/attractor/src/interviewer/recording.rs
Normal file
|
|
@ -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<dyn Interviewer>,
|
||||
recordings: Mutex<Vec<(Question, Answer)>>,
|
||||
}
|
||||
|
||||
impl RecordingInterviewer {
|
||||
#[must_use]
|
||||
pub fn new(inner: Box<dyn Interviewer>) -> 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());
|
||||
}
|
||||
}
|
||||
16
crates/attractor/src/lib.rs
Normal file
16
crates/attractor/src/lib.rs
Normal file
|
|
@ -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;
|
||||
197
crates/attractor/src/outcome.rs
Normal file
197
crates/attractor/src/outcome.rs
Normal file
|
|
@ -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<Self, Self::Err> {
|
||||
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<String>,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub suggested_next_ids: Vec<String>,
|
||||
#[serde(default, skip_serializing_if = "HashMap::is_empty")]
|
||||
pub context_updates: HashMap<String, serde_json::Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub notes: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub failure_reason: Option<String>,
|
||||
}
|
||||
|
||||
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<String>) -> 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<String>) -> 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::<StageStatus>().unwrap(), StageStatus::Success);
|
||||
assert_eq!("fail".parse::<StageStatus>().unwrap(), StageStatus::Fail);
|
||||
assert_eq!(
|
||||
"partial_success".parse::<StageStatus>().unwrap(),
|
||||
StageStatus::PartialSuccess
|
||||
);
|
||||
assert_eq!("retry".parse::<StageStatus>().unwrap(), StageStatus::Retry);
|
||||
assert_eq!("skipped".parse::<StageStatus>().unwrap(), StageStatus::Skipped);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stage_status_from_str_invalid() {
|
||||
assert!("unknown".parse::<StageStatus>().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);
|
||||
}
|
||||
}
|
||||
119
crates/attractor/src/parser/ast.rs
Normal file
119
crates/attractor/src/parser/ast.rs
Normal file
|
|
@ -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<AttrBlock>,
|
||||
}
|
||||
|
||||
/// An edge statement: `A -> B -> C [attrs]?`.
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct EdgeStmt {
|
||||
/// Chain of node IDs (at least 2).
|
||||
pub nodes: Vec<String>,
|
||||
pub attrs: Option<AttrBlock>,
|
||||
}
|
||||
|
||||
/// A subgraph statement: `subgraph name? { stmts }`.
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct SubgraphStmt {
|
||||
pub name: Option<String>,
|
||||
pub statements: Vec<Statement>,
|
||||
}
|
||||
|
||||
/// 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<Statement>,
|
||||
}
|
||||
|
||||
#[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);
|
||||
}
|
||||
}
|
||||
404
crates/attractor/src/parser/grammar.rs
Normal file
404
crates/attractor/src/parser/grammar.rs
Normal file
|
|
@ -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<char>> {
|
||||
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::<nom::error::Error<&str>>(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::<nom::error::Error<&str>>(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));
|
||||
}
|
||||
}
|
||||
413
crates/attractor/src/parser/lexer.rs
Normal file
413
crates/attractor/src/parser/lexer.rs
Normal file
|
|
@ -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<char> = 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()))));
|
||||
}
|
||||
}
|
||||
147
crates/attractor/src/parser/mod.rs
Normal file
147
crates/attractor/src/parser/mod.rs
Normal file
|
|
@ -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<Graph, AttractorError> {
|
||||
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());
|
||||
}
|
||||
}
|
||||
516
crates/attractor/src/parser/semantic.rs
Normal file
516
crates/attractor/src/parser/semantic.rs
Normal file
|
|
@ -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<Duration> {
|
||||
if s.ends_with("ms") {
|
||||
let num = s.strip_suffix("ms")?.parse::<u64>().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<String, AttrValue> {
|
||||
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<String, AttrValue>,
|
||||
edge_defaults: HashMap<String, AttrValue>,
|
||||
}
|
||||
|
||||
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<String, AttrValue>,
|
||||
scoped_edge_defaults: &HashMap<String, AttrValue>,
|
||||
) {
|
||||
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<Graph, AttractorError> {
|
||||
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"));
|
||||
}
|
||||
}
|
||||
177
crates/attractor/src/pipeline.rs
Normal file
177
crates/attractor/src/pipeline.rs
Normal file
|
|
@ -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<Box<dyn Transform>>,
|
||||
}
|
||||
|
||||
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<dyn Transform>) {
|
||||
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<Diagnostic>), 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<Graph, AttractorError> {
|
||||
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<String> = 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");
|
||||
}
|
||||
}
|
||||
518
crates/attractor/src/stylesheet.rs
Normal file
518
crates/attractor/src/stylesheet.rs
Normal file
|
|
@ -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<Declaration>,
|
||||
}
|
||||
|
||||
/// A parsed stylesheet containing multiple rules.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct Stylesheet {
|
||||
pub rules: Vec<Rule>,
|
||||
}
|
||||
|
||||
/// 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<Stylesheet, AttractorError> {
|
||||
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<Selector, AttractorError> {
|
||||
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<Vec<Declaration>, 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<String> = graph.nodes.keys().cloned().collect();
|
||||
|
||||
for node_id in &node_ids {
|
||||
let mut applied: std::collections::HashMap<String, (String, u8)> =
|
||||
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()))
|
||||
);
|
||||
}
|
||||
}
|
||||
254
crates/attractor/src/transform.rs
Normal file
254
crates/attractor/src/transform.rs
Normal file
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
160
crates/attractor/src/validation/mod.rs
Normal file
160
crates/attractor/src/validation/mod.rs
Normal file
|
|
@ -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<String>,
|
||||
pub edge: Option<(String, String)>,
|
||||
pub fix: Option<String>,
|
||||
}
|
||||
|
||||
/// A lint rule that validates a graph.
|
||||
pub trait LintRule {
|
||||
fn name(&self) -> &'static str;
|
||||
fn apply(&self, graph: &Graph) -> Vec<Diagnostic>;
|
||||
}
|
||||
|
||||
/// 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<Diagnostic> {
|
||||
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<Vec<Diagnostic>, 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<String> = 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<Diagnostic> {
|
||||
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);
|
||||
}
|
||||
}
|
||||
1130
crates/attractor/src/validation/rules.rs
Normal file
1130
crates/attractor/src/validation/rules.rs
Normal file
File diff suppressed because it is too large
Load diff
1180
crates/attractor/tests/integration.rs
Normal file
1180
crates/attractor/tests/integration.rs
Normal file
File diff suppressed because it is too large
Load diff
243
docs/agent/reviews/attractor-spec-full-review.md
Normal file
243
docs/agent/reviews/attractor-spec-full-review.md
Normal file
|
|
@ -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<i64>` 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.<node_id>` 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).**
|
||||
250
docs/agent/reviews/spec-sections-5-6-review.md
Normal file
250
docs/agent/reviews/spec-sections-5-6-review.md
Normal file
|
|
@ -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<RwLock<HashMap<String, Value>>>` (context.rs:9)
|
||||
- Append-only logs using `Arc<RwLock<Vec<String>>>` (context.rs:10)
|
||||
- `set(key, value)` with write lock (context.rs:33-38)
|
||||
- `get(key)` with read lock, returns `Option<Value>` (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<Value>` 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.<node_id>` | **GAP** | Not implemented anywhere. The engine tracks retry attempts locally in `execute_with_retry` but never writes `internal.retry_count.<node_id>` 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<String>` (outcome.rs:51)
|
||||
- `suggested_next_ids: Vec<String>` (outcome.rs:53)
|
||||
- `context_updates: HashMap<String, Value>` (outcome.rs:55)
|
||||
- `notes: Option<String>` (outcome.rs:57)
|
||||
- `failure_reason: Option<String>` (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<Utc>` (checkpoint.rs:14)
|
||||
- `current_node: String` (checkpoint.rs:15)
|
||||
- `completed_nodes: Vec<String>` (checkpoint.rs:16)
|
||||
- `node_retries: HashMap<String, u32>` (checkpoint.rs:17)
|
||||
- `context_values: HashMap<String, Value>` (checkpoint.rs:18)
|
||||
- `logs: Vec<String>` (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<ArtifactInfo>` (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<Question>) -> Vec<Answer>` 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<QuestionOption>`
|
||||
- `allow_freeform: bool`
|
||||
- `default: Option<Answer>`
|
||||
- `timeout_seconds: Option<f64>`
|
||||
- `stage: String`
|
||||
- `metadata: HashMap<String, Value>`
|
||||
|
||||
`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<QuestionOption>`
|
||||
- `text: Option<String>`
|
||||
|
||||
`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<Answer>` (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<dyn Interviewer>` (recording.rs:9)
|
||||
- Records `(Question, Answer)` pairs in `Mutex<Vec<...>>` (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<f64>` 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<Answer>` 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.<node_id>` 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 |
|
||||
197
docs/agent/reviews/spec-sections-7-9-review.md
Normal file
197
docs/agent/reviews/spec-sections-7-9-review.md
Normal file
|
|
@ -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<String>` 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<SessionEvent>` |
|
||||
| `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.
|
||||
Loading…
Add table
Reference in a new issue