Implement three missing spec features: SubPipelineHandler, GraphMergeTransform, HTTP Server

Add SubPipelineHandler (handler/sub_pipeline.rs) that inline-executes a parsed
sub-graph within the same engine, reading DOT source from node attributes and
propagating context diffs back to the parent pipeline.

Add GraphMergeTransform (transform.rs) that merges nodes and edges from secondary
graphs into a primary graph with namespace-prefixed IDs to avoid collisions.

Add WebInterviewer (interviewer/web.rs) backed by oneshot channels for async
question/answer flow, and HTTP server (server.rs) with 8 axum endpoints behind
a "server" feature flag for pipeline management and human-in-the-loop via web.

Fix tempfile dev-dependency usage in server production code by using std::env::temp_dir.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Bryan Helmkamp 2026-02-22 12:07:44 -04:00
parent a2f9e87b7c
commit c15b532bcf
10 changed files with 3085 additions and 5 deletions

82
Cargo.lock generated
View file

@ -133,6 +133,7 @@ name = "attractor"
version = "0.1.0"
dependencies = [
"async-trait",
"axum",
"chrono",
"coding-agent-loop",
"futures",
@ -143,6 +144,8 @@ dependencies = [
"tempfile",
"thiserror 2.0.18",
"tokio",
"tokio-stream",
"tower",
"unified-llm",
"uuid",
]
@ -175,6 +178,58 @@ dependencies = [
"fs_extra",
]
[[package]]
name = "axum"
version = "0.8.8"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8b52af3cb4058c895d37317bb27508dccc8e5f2d39454016b297bf4a400597b8"
dependencies = [
"axum-core",
"bytes",
"form_urlencoded",
"futures-util",
"http",
"http-body",
"http-body-util",
"hyper",
"hyper-util",
"itoa",
"matchit",
"memchr",
"mime",
"percent-encoding",
"pin-project-lite",
"serde_core",
"serde_json",
"serde_path_to_error",
"serde_urlencoded",
"sync_wrapper",
"tokio",
"tower",
"tower-layer",
"tower-service",
"tracing",
]
[[package]]
name = "axum-core"
version = "0.5.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "08c78f31d7b1291f7ee735c1c6780ccde7785daae9a9206026862dab7d8792d1"
dependencies = [
"bytes",
"futures-core",
"http",
"http-body",
"http-body-util",
"mime",
"pin-project-lite",
"sync_wrapper",
"tower-layer",
"tower-service",
"tracing",
]
[[package]]
name = "base64"
version = "0.22.1"
@ -770,6 +825,12 @@ version = "1.10.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87"
[[package]]
name = "httpdate"
version = "1.0.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9"
[[package]]
name = "hyper"
version = "1.8.1"
@ -784,6 +845,7 @@ dependencies = [
"http",
"http-body",
"httparse",
"httpdate",
"itoa",
"pin-project-lite",
"pin-utils",
@ -1137,6 +1199,12 @@ version = "0.4.29"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897"
[[package]]
name = "matchit"
version = "0.8.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3"
[[package]]
name = "memchr"
version = "2.8.0"
@ -1863,6 +1931,17 @@ dependencies = [
"zmij",
]
[[package]]
name = "serde_path_to_error"
version = "0.1.20"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "10a9ff822e371bb5403e391ecd83e182e0e77ba7f6fe0160b795797109d1b457"
dependencies = [
"itoa",
"serde",
"serde_core",
]
[[package]]
name = "serde_urlencoded"
version = "0.7.1"
@ -2109,6 +2188,7 @@ dependencies = [
"futures-core",
"pin-project-lite",
"tokio",
"tokio-util",
]
[[package]]
@ -2137,6 +2217,7 @@ dependencies = [
"tokio",
"tower-layer",
"tower-service",
"tracing",
]
[[package]]
@ -2175,6 +2256,7 @@ version = "0.1.44"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100"
dependencies = [
"log",
"pin-project-lite",
"tracing-core",
]

View file

@ -9,6 +9,10 @@ keywords = ["llm", "ai", "pipeline", "workflow", "dot"]
categories = ["development-tools"]
readme = "README.md"
[features]
default = []
server = ["axum", "tower", "tokio-stream"]
[dependencies]
coding-agent-loop = { path = "../coding-agent-loop" }
unified-llm = { path = "../unified-llm" }
@ -22,10 +26,16 @@ async-trait.workspace = true
futures.workspace = true
chrono = { workspace = true, features = ["serde"] }
nom = "7"
axum = { version = "0.8", optional = true }
tower = { version = "0.5", optional = true }
tokio-stream = { workspace = true, optional = true, features = ["sync"] }
[dev-dependencies]
tokio = { workspace = true, features = ["test-util", "macros"] }
tempfile = "3"
axum = "0.8"
tower = "0.5"
tokio-stream = { workspace = true, features = ["sync"] }
[lints]
workspace = true

View file

@ -5,6 +5,7 @@ pub mod fan_in;
pub mod manager_loop;
pub mod parallel;
pub mod start;
pub mod sub_pipeline;
pub mod tool;
pub mod wait_human;

View file

@ -0,0 +1,371 @@
use std::path::Path;
use std::sync::Arc;
use async_trait::async_trait;
use crate::context::Context;
use crate::engine::select_edge;
use crate::error::AttractorError;
use crate::event::EventEmitter;
use crate::graph::{Graph, Node};
use crate::outcome::Outcome;
use crate::pipeline::prepare_pipeline;
use super::{Handler, HandlerRegistry};
/// Executes a sub-pipeline defined by inline DOT source in a node attribute.
/// The sub-pipeline runs with a cloned context; context updates propagate back.
pub struct SubPipelineHandler {
registry: Arc<HandlerRegistry>,
_emitter: Arc<EventEmitter>,
}
impl SubPipelineHandler {
#[must_use]
pub fn new(registry: Arc<HandlerRegistry>, emitter: Arc<EventEmitter>) -> Self {
Self {
registry,
_emitter: emitter,
}
}
}
/// Check whether a node is a terminal (exit) node.
fn is_terminal(node: &Node) -> bool {
node.shape() == "Msquare" || node.handler_type() == Some("exit")
}
#[async_trait]
impl Handler for SubPipelineHandler {
async fn execute(
&self,
node: &Node,
context: &Context,
_graph: &Graph,
logs_root: &Path,
) -> Result<Outcome, AttractorError> {
// 1. Get DOT source from node attribute
let dot_source = match node.attrs.get("sub_pipeline.dot_source").and_then(|v| v.as_str()) {
Some(s) if !s.is_empty() => s,
_ => return Ok(Outcome::fail("No sub_pipeline.dot_source attribute specified")),
};
// 2. Parse the sub-pipeline DOT
let sub_graph = match prepare_pipeline(dot_source) {
Ok(g) => g,
Err(e) => return Ok(Outcome::fail(format!("Failed to parse sub-pipeline: {e}"))),
};
// 3. Find start node
let start_node = match sub_graph.find_start_node() {
Some(n) => n.id.clone(),
None => return Ok(Outcome::fail("Sub-pipeline has no start node")),
};
// 4. Clone parent context for isolation
let sub_context = context.clone_context();
let before_snapshot = context.snapshot();
// 5. Walk the sub-graph
let sub_logs_root = logs_root.join(&node.id);
let mut current_node_id = start_node;
let mut last_outcome = Outcome::success();
let max_steps: usize = 1000;
let mut steps: usize = 0;
while steps < max_steps {
steps += 1;
let sub_node = match sub_graph.nodes.get(&current_node_id) {
Some(n) => n,
None => {
return Ok(Outcome::fail(format!(
"Sub-pipeline node not found: {current_node_id}"
)));
}
};
// Check for terminal node
if is_terminal(sub_node) {
break;
}
// Execute the node handler
let handler = self.registry.resolve(sub_node);
last_outcome = handler
.execute(sub_node, &sub_context, &sub_graph, &sub_logs_root)
.await?;
// Apply context updates from the outcome
sub_context.apply_updates(&last_outcome.context_updates);
sub_context.set("outcome", serde_json::json!(last_outcome.status.to_string()));
// Select next edge
match select_edge(&current_node_id, &last_outcome, &sub_context, &sub_graph) {
Some(edge) => {
current_node_id.clone_from(&edge.to);
}
None => break,
}
}
// 6. Compute context diff (sub_context changes vs parent's original snapshot)
let after_snapshot = sub_context.snapshot();
let mut context_updates = std::collections::HashMap::new();
for (key, value) in &after_snapshot {
match before_snapshot.get(key) {
Some(old_value) if old_value == value => {}
_ => {
context_updates.insert(key.clone(), value.clone());
}
}
}
// 7. Return the last outcome with context updates propagated
let mut result = last_outcome;
result.context_updates.extend(context_updates);
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::graph::AttrValue;
use crate::handler::exit::ExitHandler;
use crate::handler::start::StartHandler;
use crate::outcome::StageStatus;
fn make_registry() -> Arc<HandlerRegistry> {
let mut registry = HandlerRegistry::new(Box::new(StartHandler));
registry.register("start", Box::new(StartHandler));
registry.register("exit", Box::new(ExitHandler));
Arc::new(registry)
}
fn make_emitter() -> Arc<EventEmitter> {
Arc::new(EventEmitter::new())
}
#[tokio::test]
async fn executes_simple_sub_pipeline() {
let registry = make_registry();
let handler = SubPipelineHandler::new(registry, make_emitter());
let mut node = Node::new("sub");
node.attrs.insert(
"sub_pipeline.dot_source".to_string(),
AttrValue::String(
r#"digraph Sub {
start [shape=Mdiamond]
exit [shape=Msquare]
start -> exit
}"#
.to_string(),
),
);
let context = Context::new();
let graph = Graph::new("parent");
let tmp = tempfile::tempdir().unwrap();
let outcome = handler
.execute(&node, &context, &graph, tmp.path())
.await
.unwrap();
assert_eq!(outcome.status, StageStatus::Success);
}
#[tokio::test]
async fn parent_context_available_in_sub_pipeline() {
let registry = make_registry();
let handler = SubPipelineHandler::new(registry, make_emitter());
let mut node = Node::new("sub");
node.attrs.insert(
"sub_pipeline.dot_source".to_string(),
AttrValue::String(
r#"digraph Sub {
start [shape=Mdiamond]
exit [shape=Msquare]
start -> exit
}"#
.to_string(),
),
);
let context = Context::new();
context.set("parent.value", serde_json::json!("hello"));
let graph = Graph::new("parent");
let tmp = tempfile::tempdir().unwrap();
let outcome = handler
.execute(&node, &context, &graph, tmp.path())
.await
.unwrap();
assert_eq!(outcome.status, StageStatus::Success);
// The sub-pipeline clones the context, so the parent value should be
// available during sub-execution. After execution, any sub-pipeline
// context updates should be in the outcome's context_updates.
}
#[tokio::test]
async fn context_updates_propagate_back() {
// Use a handler that sets a context value, register it in the sub-pipeline registry
struct ContextSettingHandler;
#[async_trait]
impl Handler for ContextSettingHandler {
async fn execute(
&self,
_node: &Node,
context: &Context,
_graph: &Graph,
_logs_root: &Path,
) -> Result<Outcome, AttractorError> {
context.set("sub.result", serde_json::json!("from_sub"));
Ok(Outcome::success())
}
}
let mut registry = HandlerRegistry::new(Box::new(ContextSettingHandler));
registry.register("start", Box::new(StartHandler));
registry.register("exit", Box::new(ExitHandler));
let registry = Arc::new(registry);
let handler = SubPipelineHandler::new(registry, make_emitter());
let mut node = Node::new("sub");
node.attrs.insert(
"sub_pipeline.dot_source".to_string(),
AttrValue::String(
r#"digraph Sub {
start [shape=Mdiamond]
work [shape=box]
exit [shape=Msquare]
start -> work -> exit
}"#
.to_string(),
),
);
let context = Context::new();
let graph = Graph::new("parent");
let tmp = tempfile::tempdir().unwrap();
let outcome = handler
.execute(&node, &context, &graph, tmp.path())
.await
.unwrap();
assert_eq!(outcome.status, StageStatus::Success);
// Context updates from the sub-pipeline should be in the outcome
assert!(
outcome.context_updates.contains_key("sub.result"),
"sub-pipeline context updates should propagate back"
);
assert_eq!(
outcome.context_updates.get("sub.result"),
Some(&serde_json::json!("from_sub"))
);
}
#[tokio::test]
async fn failing_sub_pipeline_returns_fail() {
struct AlwaysFailHandler;
#[async_trait]
impl Handler for AlwaysFailHandler {
async fn execute(
&self,
_node: &Node,
_context: &Context,
_graph: &Graph,
_logs_root: &Path,
) -> Result<Outcome, AttractorError> {
Ok(Outcome::fail("sub-pipeline failure"))
}
}
let mut registry = HandlerRegistry::new(Box::new(AlwaysFailHandler));
registry.register("start", Box::new(StartHandler));
registry.register("exit", Box::new(ExitHandler));
let registry = Arc::new(registry);
let handler = SubPipelineHandler::new(registry, make_emitter());
let mut node = Node::new("sub");
// Sub-pipeline where the work node fails and there's a fail edge to exit
node.attrs.insert(
"sub_pipeline.dot_source".to_string(),
AttrValue::String(
r#"digraph Sub {
start [shape=Mdiamond]
work [shape=box, max_retries="0"]
exit [shape=Msquare]
start -> work
work -> exit [condition="outcome=fail"]
}"#
.to_string(),
),
);
let context = Context::new();
let graph = Graph::new("parent");
let tmp = tempfile::tempdir().unwrap();
let outcome = handler
.execute(&node, &context, &graph, tmp.path())
.await
.unwrap();
assert_eq!(outcome.status, StageStatus::Fail);
}
#[tokio::test]
async fn missing_dot_source_returns_fail() {
let registry = make_registry();
let handler = SubPipelineHandler::new(registry, make_emitter());
let node = Node::new("sub");
let context = Context::new();
let graph = Graph::new("parent");
let tmp = tempfile::tempdir().unwrap();
let outcome = handler
.execute(&node, &context, &graph, tmp.path())
.await
.unwrap();
assert_eq!(outcome.status, StageStatus::Fail);
assert!(
outcome
.failure_reason
.as_deref()
.unwrap()
.contains("sub_pipeline.dot_source"),
"should mention the missing attribute"
);
}
#[tokio::test]
async fn invalid_dot_source_returns_fail() {
let registry = make_registry();
let handler = SubPipelineHandler::new(registry, make_emitter());
let mut node = Node::new("sub");
node.attrs.insert(
"sub_pipeline.dot_source".to_string(),
AttrValue::String("not valid dot".to_string()),
);
let context = Context::new();
let graph = Graph::new("parent");
let tmp = tempfile::tempdir().unwrap();
let outcome = handler
.execute(&node, &context, &graph, tmp.path())
.await
.unwrap();
assert_eq!(outcome.status, StageStatus::Fail);
}
}

View file

@ -3,6 +3,7 @@ pub mod callback;
pub mod console;
pub mod queue;
pub mod recording;
pub mod web;
use std::collections::HashMap;

View file

@ -0,0 +1,256 @@
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use tokio::sync::oneshot;
use super::{Answer, Interviewer, Question};
/// A pending question waiting for an answer from an external source (e.g., HTTP endpoint).
#[derive(Debug)]
pub struct PendingQuestion {
pub id: String,
pub question: Question,
}
/// Internal state: maps question ID to its oneshot sender.
struct WebInterviewerInner {
pending: HashMap<String, oneshot::Sender<Answer>>,
questions: Vec<PendingQuestion>,
next_id: u64,
}
/// An interviewer that holds questions until answers are submitted externally.
///
/// When `ask()` is called, the question is enqueued with a unique ID and the call
/// blocks until `submit_answer()` is called with the matching ID.
pub struct WebInterviewer {
inner: Arc<Mutex<WebInterviewerInner>>,
}
impl WebInterviewer {
#[must_use]
pub fn new() -> Self {
Self {
inner: Arc::new(Mutex::new(WebInterviewerInner {
pending: HashMap::new(),
questions: Vec::new(),
next_id: 1,
})),
}
}
/// Returns a snapshot of currently pending questions.
///
/// # Panics
///
/// Panics if the internal lock is poisoned.
#[must_use]
pub fn pending_questions(&self) -> Vec<PendingQuestion> {
let inner = self.inner.lock().expect("web interviewer lock poisoned");
inner
.questions
.iter()
.map(|pq| PendingQuestion {
id: pq.id.clone(),
question: pq.question.clone(),
})
.collect()
}
/// Submit an answer for a pending question by ID.
/// Returns `true` if the question was found and the answer was delivered,
/// `false` if no such question was pending.
///
/// # Panics
///
/// Panics if the internal lock is poisoned.
#[must_use]
pub fn submit_answer(&self, question_id: &str, answer: Answer) -> bool {
let sender = {
let mut inner = self.inner.lock().expect("web interviewer lock poisoned");
let sender = inner.pending.remove(question_id);
if sender.is_some() {
inner.questions.retain(|pq| pq.id != question_id);
}
sender
};
sender.is_some_and(|tx| tx.send(answer).is_ok())
}
}
impl Default for WebInterviewer {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Interviewer for WebInterviewer {
async fn ask(&self, question: Question) -> Answer {
let (tx, rx) = oneshot::channel();
{
let mut inner = self.inner.lock().expect("web interviewer lock poisoned");
let id = format!("q-{}", inner.next_id);
inner.next_id += 1;
inner.pending.insert(id.clone(), tx);
inner.questions.push(PendingQuestion {
id,
question: question.clone(),
});
}
// Block until answer arrives or sender is dropped
rx.await.unwrap_or_else(|_| Answer::skipped())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::interviewer::{AnswerValue, QuestionType};
use std::sync::Arc;
#[tokio::test]
async fn ask_blocks_until_answer_submitted() {
let interviewer = Arc::new(WebInterviewer::new());
let interviewer_clone = Arc::clone(&interviewer);
let ask_handle = tokio::spawn(async move {
let q = Question::new("approve?", QuestionType::YesNo);
interviewer_clone.ask(q).await
});
// Give the ask task a moment to register the question
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
// Question should be pending
let pending = interviewer.pending_questions();
assert_eq!(pending.len(), 1);
assert_eq!(pending[0].question.text, "approve?");
// Submit answer
let submitted = interviewer.submit_answer(&pending[0].id, Answer::yes());
assert!(submitted);
// ask() should now return
let answer = ask_handle.await.expect("task should complete");
assert_eq!(answer.value, AnswerValue::Yes);
}
#[tokio::test]
async fn submit_answer_unblocks_ask() {
let interviewer = Arc::new(WebInterviewer::new());
let interviewer_clone = Arc::clone(&interviewer);
let ask_handle = tokio::spawn(async move {
let q = Question::new("name?", QuestionType::Freeform);
interviewer_clone.ask(q).await
});
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let pending = interviewer.pending_questions();
assert_eq!(pending.len(), 1);
let _ = interviewer.submit_answer(&pending[0].id, Answer::text("Alice"));
let answer = ask_handle.await.expect("task should complete");
assert_eq!(answer.value, AnswerValue::Text("Alice".to_string()));
assert_eq!(answer.text, Some("Alice".to_string()));
}
#[tokio::test]
async fn timeout_returns_default_or_timeout_answer() {
let interviewer = Arc::new(WebInterviewer::new());
let mut q = Question::new("approve?", QuestionType::YesNo);
q.timeout_seconds = Some(0.05);
// Use ask_with_timeout from the parent module
let answer = crate::interviewer::ask_with_timeout(interviewer.as_ref(), q).await;
assert_eq!(answer.value, AnswerValue::Timeout);
}
#[tokio::test]
async fn question_id_correlation() {
let interviewer = Arc::new(WebInterviewer::new());
let i1 = Arc::clone(&interviewer);
let i2 = Arc::clone(&interviewer);
// Spawn two concurrent asks
let handle1 = tokio::spawn(async move {
let q = Question::new("first?", QuestionType::YesNo);
i1.ask(q).await
});
let handle2 = tokio::spawn(async move {
let q = Question::new("second?", QuestionType::YesNo);
i2.ask(q).await
});
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let pending = interviewer.pending_questions();
assert_eq!(pending.len(), 2);
// Find which ID corresponds to which question
let first_id = pending
.iter()
.find(|pq| pq.question.text == "first?")
.expect("first question should be pending")
.id
.clone();
let second_id = pending
.iter()
.find(|pq| pq.question.text == "second?")
.expect("second question should be pending")
.id
.clone();
// Answer them in reverse order
let _ = interviewer.submit_answer(&second_id, Answer::no());
let _ = interviewer.submit_answer(&first_id, Answer::yes());
let answer1 = handle1.await.expect("task should complete");
let answer2 = handle2.await.expect("task should complete");
assert_eq!(answer1.value, AnswerValue::Yes);
assert_eq!(answer2.value, AnswerValue::No);
}
#[test]
fn submit_answer_for_unknown_id_returns_false() {
let interviewer = WebInterviewer::new();
let result = interviewer.submit_answer("nonexistent", Answer::yes());
assert!(!result);
}
#[tokio::test]
async fn pending_questions_empty_initially() {
let interviewer = WebInterviewer::new();
assert!(interviewer.pending_questions().is_empty());
}
#[tokio::test]
async fn pending_questions_cleared_after_answer() {
let interviewer = Arc::new(WebInterviewer::new());
let i_clone = Arc::clone(&interviewer);
let handle = tokio::spawn(async move {
let q = Question::new("q?", QuestionType::YesNo);
i_clone.ask(q).await
});
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let pending = interviewer.pending_questions();
assert_eq!(pending.len(), 1);
let _ = interviewer.submit_answer(&pending[0].id, Answer::yes());
handle.await.expect("task should complete");
assert!(interviewer.pending_questions().is_empty());
}
}

View file

@ -11,6 +11,8 @@ pub mod interviewer;
pub mod outcome;
pub mod parser;
pub mod pipeline;
#[cfg(any(feature = "server", test))]
pub mod server;
pub mod stylesheet;
pub mod transform;
pub mod validation;

View file

@ -0,0 +1,727 @@
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use axum::extract::{Path, State};
use axum::http::StatusCode;
use axum::response::sse::{Event, Sse};
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post};
use axum::{Json, Router};
use serde::{Deserialize, Serialize};
use tokio::sync::broadcast;
use tokio_stream::wrappers::BroadcastStream;
use tokio_stream::StreamExt;
use crate::checkpoint::Checkpoint;
use crate::context::Context;
use crate::engine::{PipelineEngine, RunConfig};
use crate::event::{EventEmitter, PipelineEvent};
use crate::handler::HandlerRegistry;
use crate::interviewer::web::WebInterviewer;
use crate::interviewer::{Answer, AnswerValue};
/// Status of a managed pipeline.
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum PipelineStatus {
Running,
Completed,
Failed,
Cancelled,
}
/// A pending question exposed via the API.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ApiQuestion {
pub id: String,
pub text: String,
pub question_type: String,
}
/// Snapshot of a managed pipeline.
struct ManagedPipeline {
status: PipelineStatus,
error: Option<String>,
interviewer: Arc<WebInterviewer>,
event_tx: broadcast::Sender<PipelineEvent>,
context: Option<Context>,
checkpoint: Option<Checkpoint>,
cancel_tx: Option<tokio::sync::oneshot::Sender<()>>,
}
/// Shared application state for the server.
pub struct AppState {
pipelines: Mutex<HashMap<String, ManagedPipeline>>,
registry_factory: Box<dyn Fn() -> HandlerRegistry + Send + Sync>,
}
/// Request body for POST /pipelines.
#[derive(Debug, Deserialize)]
pub struct StartPipelineRequest {
pub dot_source: String,
}
/// Response body for POST /pipelines.
#[derive(Debug, Serialize)]
pub struct StartPipelineResponse {
pub id: String,
}
/// Response body for GET /pipelines/{id}.
#[derive(Debug, Serialize)]
pub struct PipelineStatusResponse {
pub id: String,
pub status: PipelineStatus,
#[serde(skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
}
/// Request body for POST /pipelines/{id}/questions/{qid}/answer.
#[derive(Debug, Deserialize)]
pub struct SubmitAnswerRequest {
pub value: String,
}
/// Response for answer submission.
#[derive(Debug, Serialize)]
pub struct SubmitAnswerResponse {
pub accepted: bool,
}
/// Build the axum Router with all pipeline endpoints.
pub fn build_router(state: Arc<AppState>) -> Router {
Router::new()
.route("/pipelines", post(start_pipeline))
.route("/pipelines/{id}", get(get_pipeline_status))
.route("/pipelines/{id}/questions", get(get_questions))
.route(
"/pipelines/{id}/questions/{qid}/answer",
post(submit_answer),
)
.route("/pipelines/{id}/events", get(get_events))
.route("/pipelines/{id}/checkpoint", get(get_checkpoint))
.route("/pipelines/{id}/context", get(get_context))
.route("/pipelines/{id}/cancel", post(cancel_pipeline))
.with_state(state)
}
/// Create an `AppState` with the given registry factory.
pub fn create_app_state(
registry_factory: impl Fn() -> HandlerRegistry + Send + Sync + 'static,
) -> Arc<AppState> {
Arc::new(AppState {
pipelines: Mutex::new(HashMap::new()),
registry_factory: Box::new(registry_factory),
})
}
async fn start_pipeline(
State(state): State<Arc<AppState>>,
Json(req): Json<StartPipelineRequest>,
) -> Response {
// Parse the DOT source
let graph = match crate::pipeline::prepare_pipeline(&req.dot_source) {
Ok(g) => g,
Err(e) => {
return (StatusCode::BAD_REQUEST, Json(serde_json::json!({"error": e.to_string()})))
.into_response();
}
};
let pipeline_id = uuid::Uuid::new_v4().to_string();
let interviewer = Arc::new(WebInterviewer::new());
let (event_tx, _) = broadcast::channel(256);
let (cancel_tx, cancel_rx) = tokio::sync::oneshot::channel::<()>();
let context = Context::new();
// Set up event emitter that broadcasts to the channel
let mut emitter = EventEmitter::new();
let tx_clone = event_tx.clone();
emitter.on_event(move |event| {
let _ = tx_clone.send(event.clone());
});
let registry = (state.registry_factory)();
let engine = PipelineEngine::new(registry, emitter);
{
let mut pipelines = state.pipelines.lock().expect("pipelines lock poisoned");
pipelines.insert(
pipeline_id.clone(),
ManagedPipeline {
status: PipelineStatus::Running,
error: None,
interviewer: Arc::clone(&interviewer),
event_tx: event_tx.clone(),
context: Some(context.clone()),
checkpoint: None,
cancel_tx: Some(cancel_tx),
},
);
}
// Spawn pipeline execution
let state_clone = Arc::clone(&state);
let id_clone = pipeline_id.clone();
tokio::spawn(async move {
let logs_root = std::env::temp_dir().join(format!("attractor-{}", uuid::Uuid::new_v4()));
std::fs::create_dir_all(&logs_root).expect("failed to create logs directory");
let config = RunConfig { logs_root };
let result = tokio::select! {
result = engine.run(&graph, &config) => result,
_ = cancel_rx => {
let mut pipelines = state_clone.pipelines.lock().expect("pipelines lock poisoned");
if let Some(pipeline) = pipelines.get_mut(&id_clone) {
pipeline.status = PipelineStatus::Cancelled;
}
return;
}
};
// Save final checkpoint
let checkpoint = Checkpoint::load(&config.logs_root.join("checkpoint.json")).ok();
let mut pipelines = state_clone.pipelines.lock().expect("pipelines lock poisoned");
if let Some(pipeline) = pipelines.get_mut(&id_clone) {
match result {
Ok(_) => {
pipeline.status = PipelineStatus::Completed;
}
Err(e) => {
pipeline.status = PipelineStatus::Failed;
pipeline.error = Some(e.to_string());
}
}
pipeline.checkpoint = checkpoint;
}
});
(
StatusCode::CREATED,
Json(StartPipelineResponse { id: pipeline_id }),
)
.into_response()
}
async fn get_pipeline_status(
State(state): State<Arc<AppState>>,
Path(id): Path<String>,
) -> Response {
let pipelines = state.pipelines.lock().expect("pipelines lock poisoned");
match pipelines.get(&id) {
Some(pipeline) => (
StatusCode::OK,
Json(PipelineStatusResponse {
id: id.clone(),
status: pipeline.status.clone(),
error: pipeline.error.clone(),
}),
)
.into_response(),
None => StatusCode::NOT_FOUND.into_response(),
}
}
async fn get_questions(
State(state): State<Arc<AppState>>,
Path(id): Path<String>,
) -> Response {
let pipelines = state.pipelines.lock().expect("pipelines lock poisoned");
match pipelines.get(&id) {
Some(pipeline) => {
let pending = pipeline.interviewer.pending_questions();
let questions: Vec<ApiQuestion> = pending
.into_iter()
.map(|pq| ApiQuestion {
id: pq.id,
text: pq.question.text,
question_type: format!("{:?}", pq.question.question_type),
})
.collect();
(StatusCode::OK, Json(questions)).into_response()
}
None => StatusCode::NOT_FOUND.into_response(),
}
}
async fn submit_answer(
State(state): State<Arc<AppState>>,
Path((id, qid)): Path<(String, String)>,
Json(req): Json<SubmitAnswerRequest>,
) -> Response {
let pipelines = state.pipelines.lock().expect("pipelines lock poisoned");
match pipelines.get(&id) {
Some(pipeline) => {
let answer = Answer {
value: AnswerValue::Text(req.value.clone()),
selected_option: None,
text: Some(req.value),
};
let accepted = pipeline.interviewer.submit_answer(&qid, answer);
(StatusCode::OK, Json(SubmitAnswerResponse { accepted })).into_response()
}
None => StatusCode::NOT_FOUND.into_response(),
}
}
async fn get_events(
State(state): State<Arc<AppState>>,
Path(id): Path<String>,
) -> Response {
let rx = {
let pipelines = state.pipelines.lock().expect("pipelines lock poisoned");
match pipelines.get(&id) {
Some(pipeline) => pipeline.event_tx.subscribe(),
None => return StatusCode::NOT_FOUND.into_response(),
}
};
let stream = BroadcastStream::new(rx).filter_map(|result| match result {
Ok(event) => {
let data = serde_json::to_string(&event).unwrap_or_default();
Some(Ok::<Event, std::convert::Infallible>(
Event::default().data(data),
))
}
Err(_) => None,
});
Sse::new(stream).into_response()
}
async fn get_checkpoint(
State(state): State<Arc<AppState>>,
Path(id): Path<String>,
) -> Response {
let pipelines = state.pipelines.lock().expect("pipelines lock poisoned");
match pipelines.get(&id) {
Some(pipeline) => match &pipeline.checkpoint {
Some(cp) => (StatusCode::OK, Json(cp.clone())).into_response(),
None => (StatusCode::OK, Json(serde_json::json!(null))).into_response(),
},
None => StatusCode::NOT_FOUND.into_response(),
}
}
async fn get_context(
State(state): State<Arc<AppState>>,
Path(id): Path<String>,
) -> Response {
let pipelines = state.pipelines.lock().expect("pipelines lock poisoned");
match pipelines.get(&id) {
Some(pipeline) => match &pipeline.context {
Some(ctx) => (StatusCode::OK, Json(ctx.snapshot())).into_response(),
None => (StatusCode::OK, Json(serde_json::json!({}))).into_response(),
},
None => StatusCode::NOT_FOUND.into_response(),
}
}
async fn cancel_pipeline(
State(state): State<Arc<AppState>>,
Path(id): Path<String>,
) -> Response {
let mut pipelines = state.pipelines.lock().expect("pipelines lock poisoned");
match pipelines.get_mut(&id) {
Some(pipeline) => {
if pipeline.status != PipelineStatus::Running {
return (
StatusCode::CONFLICT,
Json(serde_json::json!({"error": "pipeline is not running"})),
)
.into_response();
}
if let Some(cancel_tx) = pipeline.cancel_tx.take() {
let _ = cancel_tx.send(());
}
pipeline.status = PipelineStatus::Cancelled;
(StatusCode::OK, Json(serde_json::json!({"cancelled": true}))).into_response()
}
None => StatusCode::NOT_FOUND.into_response(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use axum::http::Request;
use tower::ServiceExt;
use crate::handler::exit::ExitHandler;
use crate::handler::start::StartHandler;
const MINIMAL_DOT: &str = r#"digraph Test {
graph [goal="Test"]
start [shape=Mdiamond]
exit [shape=Msquare]
start -> exit
}"#;
fn test_registry() -> HandlerRegistry {
let mut registry = HandlerRegistry::new(Box::new(StartHandler));
registry.register("start", Box::new(StartHandler));
registry.register("exit", Box::new(ExitHandler));
registry
}
fn test_app() -> Router {
let state = create_app_state(test_registry);
build_router(state)
}
async fn body_json(body: Body) -> serde_json::Value {
let bytes = axum::body::to_bytes(body, usize::MAX).await.unwrap();
serde_json::from_slice(&bytes).unwrap()
}
#[tokio::test]
async fn post_pipelines_starts_pipeline_and_returns_id() {
let app = test_app();
let req = Request::builder()
.method("POST")
.uri("/pipelines")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::CREATED);
let body = body_json(response.into_body()).await;
assert!(body["id"].is_string());
assert!(!body["id"].as_str().unwrap().is_empty());
}
#[tokio::test]
async fn post_pipelines_invalid_dot_returns_bad_request() {
let app = test_app();
let req = Request::builder()
.method("POST")
.uri("/pipelines")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"dot_source": "not a graph"})).unwrap(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
}
#[tokio::test]
async fn get_pipeline_status_returns_status() {
let state = create_app_state(test_registry);
let app = build_router(Arc::clone(&state));
// Start a pipeline
let req = Request::builder()
.method("POST")
.uri("/pipelines")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(),
))
.unwrap();
let response = app.clone().oneshot(req).await.unwrap();
let body = body_json(response.into_body()).await;
let pipeline_id = body["id"].as_str().unwrap().to_string();
// Give pipeline a moment to run
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
// Check status
let req = Request::builder()
.method("GET")
.uri(format!("/pipelines/{pipeline_id}"))
.body(Body::empty())
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = body_json(response.into_body()).await;
assert_eq!(body["id"].as_str().unwrap(), pipeline_id);
// Status should be either "running" or "completed"
let status = body["status"].as_str().unwrap();
assert!(
status == "running" || status == "completed",
"unexpected status: {status}"
);
}
#[tokio::test]
async fn get_pipeline_status_not_found() {
let app = test_app();
let req = Request::builder()
.method("GET")
.uri("/pipelines/nonexistent")
.body(Body::empty())
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::NOT_FOUND);
}
#[tokio::test]
async fn get_questions_returns_empty_list() {
let state = create_app_state(test_registry);
let app = build_router(Arc::clone(&state));
// Start a pipeline
let req = Request::builder()
.method("POST")
.uri("/pipelines")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(),
))
.unwrap();
let response = app.clone().oneshot(req).await.unwrap();
let body = body_json(response.into_body()).await;
let pipeline_id = body["id"].as_str().unwrap().to_string();
// Get questions (should be empty for a pipeline without wait.human nodes)
let req = Request::builder()
.method("GET")
.uri(format!("/pipelines/{pipeline_id}/questions"))
.body(Body::empty())
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = body_json(response.into_body()).await;
assert!(body.is_array());
}
#[tokio::test]
async fn submit_answer_not_found_pipeline() {
let app = test_app();
let req = Request::builder()
.method("POST")
.uri("/pipelines/nonexistent/questions/q1/answer")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"value": "yes"})).unwrap(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::NOT_FOUND);
}
#[tokio::test]
async fn get_events_not_found() {
let app = test_app();
let req = Request::builder()
.method("GET")
.uri("/pipelines/nonexistent/events")
.body(Body::empty())
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::NOT_FOUND);
}
#[tokio::test]
async fn get_checkpoint_returns_null_initially() {
let state = create_app_state(test_registry);
let app = build_router(Arc::clone(&state));
// Start a pipeline
let req = Request::builder()
.method("POST")
.uri("/pipelines")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(),
))
.unwrap();
let response = app.clone().oneshot(req).await.unwrap();
let body = body_json(response.into_body()).await;
let pipeline_id = body["id"].as_str().unwrap().to_string();
// Get checkpoint immediately (before pipeline completes, may be null)
let req = Request::builder()
.method("GET")
.uri(format!("/pipelines/{pipeline_id}/checkpoint"))
.body(Body::empty())
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn get_context_returns_map() {
let state = create_app_state(test_registry);
let app = build_router(Arc::clone(&state));
// Start a pipeline
let req = Request::builder()
.method("POST")
.uri("/pipelines")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(),
))
.unwrap();
let response = app.clone().oneshot(req).await.unwrap();
let body = body_json(response.into_body()).await;
let pipeline_id = body["id"].as_str().unwrap().to_string();
// Get context
let req = Request::builder()
.method("GET")
.uri(format!("/pipelines/{pipeline_id}/context"))
.body(Body::empty())
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = body_json(response.into_body()).await;
assert!(body.is_object());
}
#[tokio::test]
async fn cancel_pipeline_succeeds() {
let state = create_app_state(test_registry);
let app = build_router(Arc::clone(&state));
// Start a pipeline
let req = Request::builder()
.method("POST")
.uri("/pipelines")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(),
))
.unwrap();
let response = app.clone().oneshot(req).await.unwrap();
let body = body_json(response.into_body()).await;
let pipeline_id = body["id"].as_str().unwrap().to_string();
// Cancel it
let req = Request::builder()
.method("POST")
.uri(format!("/pipelines/{pipeline_id}/cancel"))
.body(Body::empty())
.unwrap();
let response = app.clone().oneshot(req).await.unwrap();
// Could be OK (cancelled) or CONFLICT (already completed)
let status = response.status();
assert!(
status == StatusCode::OK || status == StatusCode::CONFLICT,
"unexpected status: {status}"
);
}
#[tokio::test]
async fn cancel_nonexistent_pipeline_returns_not_found() {
let app = test_app();
let req = Request::builder()
.method("POST")
.uri("/pipelines/nonexistent/cancel")
.body(Body::empty())
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::NOT_FOUND);
}
#[tokio::test]
async fn get_events_returns_sse_stream() {
let state = create_app_state(test_registry);
let app = build_router(Arc::clone(&state));
// Start a pipeline
let req = Request::builder()
.method("POST")
.uri("/pipelines")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(),
))
.unwrap();
let response = app.clone().oneshot(req).await.unwrap();
let body = body_json(response.into_body()).await;
let pipeline_id = body["id"].as_str().unwrap().to_string();
// Request the SSE stream
let req = Request::builder()
.method("GET")
.uri(format!("/pipelines/{pipeline_id}/events"))
.body(Body::empty())
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
// Check content-type is text/event-stream
let content_type = response
.headers()
.get("content-type")
.expect("content-type header should be present")
.to_str()
.unwrap();
assert!(
content_type.contains("text/event-stream"),
"expected text/event-stream, got: {content_type}"
);
}
#[tokio::test]
async fn pipeline_completes_and_status_is_completed() {
let state = create_app_state(test_registry);
let app = build_router(Arc::clone(&state));
// Start a pipeline
let req = Request::builder()
.method("POST")
.uri("/pipelines")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(),
))
.unwrap();
let response = app.clone().oneshot(req).await.unwrap();
let body = body_json(response.into_body()).await;
let pipeline_id = body["id"].as_str().unwrap().to_string();
// Wait for pipeline to complete
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
// Check status
let req = Request::builder()
.method("GET")
.uri(format!("/pipelines/{pipeline_id}"))
.body(Body::empty())
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = body_json(response.into_body()).await;
assert_eq!(body["status"].as_str().unwrap(), "completed");
}
}

View file

@ -1,4 +1,4 @@
use crate::graph::{AttrValue, Graph};
use crate::graph::{AttrValue, Edge, Graph, Node};
use crate::stylesheet::{apply_stylesheet, parse_stylesheet};
/// A transform that modifies the pipeline graph after parsing and before validation.
@ -48,6 +48,43 @@ impl Transform for PreambleTransform {
}
}
/// Merges nodes and edges from secondary graphs into the primary graph.
/// Node IDs from secondary graphs are prefixed with a namespace to avoid collisions.
pub struct GraphMergeTransform {
secondary_graphs: Vec<Graph>,
}
impl GraphMergeTransform {
pub fn new(secondary_graphs: Vec<Graph>) -> Self {
Self { secondary_graphs }
}
}
impl Transform for GraphMergeTransform {
fn apply(&self, graph: &mut Graph) {
for secondary in &self.secondary_graphs {
let prefix = &secondary.name;
for (id, node) in &secondary.nodes {
let prefixed_id = format!("{prefix}.{id}");
let mut merged_node = Node::new(&prefixed_id);
merged_node.attrs = node.attrs.clone();
merged_node.classes = node.classes.clone();
graph.nodes.insert(prefixed_id, merged_node);
}
for edge in &secondary.edges {
let mut merged_edge = Edge::new(
format!("{prefix}.{}", edge.from),
format!("{prefix}.{}", edge.to),
);
merged_edge.attrs = edge.attrs.clone();
graph.edges.push(merged_edge);
}
}
}
}
/// Applies the `model_stylesheet` graph attribute to resolve LLM properties for each node.
pub struct StylesheetApplicationTransform;
@ -67,7 +104,6 @@ impl Transform for StylesheetApplicationTransform {
#[cfg(test)]
mod tests {
use super::*;
use crate::graph::Node;
#[test]
fn variable_expansion_replaces_goal() {
@ -251,4 +287,191 @@ mod tests {
assert!(graph.nodes["work"].attrs.get("prompt").is_none());
}
// -----------------------------------------------------------------------
// GraphMergeTransform tests
// -----------------------------------------------------------------------
#[test]
fn graph_merge_combines_nodes_and_edges() {
let mut primary = Graph::new("primary");
primary.nodes.insert("a".to_string(), Node::new("a"));
primary.nodes.insert("b".to_string(), Node::new("b"));
primary.edges.push(Edge::new("a", "b"));
let mut secondary = Graph::new("secondary");
secondary.nodes.insert("x".to_string(), Node::new("x"));
secondary.nodes.insert("y".to_string(), Node::new("y"));
secondary.edges.push(Edge::new("x", "y"));
let transform = GraphMergeTransform::new(vec![secondary]);
transform.apply(&mut primary);
// Primary should now have 4 nodes: a, b, secondary.x, secondary.y
assert_eq!(primary.nodes.len(), 4);
assert!(primary.nodes.contains_key("secondary.x"));
assert!(primary.nodes.contains_key("secondary.y"));
// Should have 2 edges: a->b and secondary.x->secondary.y
assert_eq!(primary.edges.len(), 2);
}
#[test]
fn graph_merge_prefixes_node_ids_to_avoid_collisions() {
let mut primary = Graph::new("primary");
primary.nodes.insert("work".to_string(), Node::new("work"));
let mut secondary = Graph::new("sub");
secondary.nodes.insert("work".to_string(), Node::new("work"));
let transform = GraphMergeTransform::new(vec![secondary]);
transform.apply(&mut primary);
// Primary "work" is preserved, secondary "work" becomes "sub.work"
assert!(primary.nodes.contains_key("work"));
assert!(primary.nodes.contains_key("sub.work"));
assert_eq!(primary.nodes.len(), 2);
}
#[test]
fn graph_merge_remaps_edges_to_prefixed_ids() {
let mut primary = Graph::new("primary");
primary.nodes.insert("a".to_string(), Node::new("a"));
let mut secondary = Graph::new("sub");
secondary.nodes.insert("x".to_string(), Node::new("x"));
secondary.nodes.insert("y".to_string(), Node::new("y"));
secondary.edges.push(Edge::new("x", "y"));
let transform = GraphMergeTransform::new(vec![secondary]);
transform.apply(&mut primary);
// The edge from secondary should be remapped to sub.x -> sub.y
let merged_edge = primary
.edges
.iter()
.find(|e| e.from == "sub.x")
.expect("should have edge from sub.x");
assert_eq!(merged_edge.to, "sub.y");
}
#[test]
fn graph_merge_preserves_primary_attributes() {
let mut primary = Graph::new("primary");
primary.attrs.insert(
"goal".to_string(),
AttrValue::String("Build feature".to_string()),
);
primary.attrs.insert(
"model_stylesheet".to_string(),
AttrValue::String("* { llm_model: sonnet; }".to_string()),
);
let mut secondary = Graph::new("sub");
secondary.attrs.insert(
"goal".to_string(),
AttrValue::String("Sub goal".to_string()),
);
secondary.nodes.insert("x".to_string(), Node::new("x"));
let transform = GraphMergeTransform::new(vec![secondary]);
transform.apply(&mut primary);
assert_eq!(primary.goal(), "Build feature");
assert_eq!(primary.model_stylesheet(), "* { llm_model: sonnet; }");
}
#[test]
fn graph_merge_empty_secondary_is_noop() {
let mut primary = Graph::new("primary");
primary.nodes.insert("a".to_string(), Node::new("a"));
primary.edges.push(Edge::new("a", "a"));
let secondary = Graph::new("empty");
let transform = GraphMergeTransform::new(vec![secondary]);
transform.apply(&mut primary);
assert_eq!(primary.nodes.len(), 1);
assert_eq!(primary.edges.len(), 1);
}
#[test]
fn graph_merge_multiple_secondary_graphs() {
let mut primary = Graph::new("primary");
primary.nodes.insert("a".to_string(), Node::new("a"));
let mut sub1 = Graph::new("sub1");
sub1.nodes.insert("n1".to_string(), Node::new("n1"));
let mut sub2 = Graph::new("sub2");
sub2.nodes.insert("n2".to_string(), Node::new("n2"));
let transform = GraphMergeTransform::new(vec![sub1, sub2]);
transform.apply(&mut primary);
assert_eq!(primary.nodes.len(), 3);
assert!(primary.nodes.contains_key("a"));
assert!(primary.nodes.contains_key("sub1.n1"));
assert!(primary.nodes.contains_key("sub2.n2"));
}
#[test]
fn graph_merge_preserves_node_attributes() {
let mut primary = Graph::new("primary");
let mut secondary = Graph::new("sub");
let mut node = Node::new("worker");
node.attrs.insert(
"prompt".to_string(),
AttrValue::String("Do the work".to_string()),
);
node.attrs.insert(
"shape".to_string(),
AttrValue::String("box".to_string()),
);
secondary.nodes.insert("worker".to_string(), node);
let transform = GraphMergeTransform::new(vec![secondary]);
transform.apply(&mut primary);
let merged = &primary.nodes["sub.worker"];
assert_eq!(merged.id, "sub.worker");
assert_eq!(
merged.attrs.get("prompt").and_then(AttrValue::as_str),
Some("Do the work")
);
assert_eq!(
merged.attrs.get("shape").and_then(AttrValue::as_str),
Some("box")
);
}
#[test]
fn graph_merge_preserves_edge_attributes() {
let mut primary = Graph::new("primary");
let mut secondary = Graph::new("sub");
secondary.nodes.insert("x".to_string(), Node::new("x"));
secondary.nodes.insert("y".to_string(), Node::new("y"));
let mut edge = Edge::new("x", "y");
edge.attrs.insert(
"condition".to_string(),
AttrValue::String("outcome=success".to_string()),
);
secondary.edges.push(edge);
let transform = GraphMergeTransform::new(vec![secondary]);
transform.apply(&mut primary);
let merged_edge = primary
.edges
.iter()
.find(|e| e.from == "sub.x")
.expect("should have merged edge");
assert_eq!(merged_edge.to, "sub.y");
assert_eq!(
merged_edge.attrs.get("condition").and_then(AttrValue::as_str),
Some("outcome=success")
);
}
}

File diff suppressed because it is too large Load diff