mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-06 02:48:25 +00:00
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:
parent
a2f9e87b7c
commit
c15b532bcf
10 changed files with 3085 additions and 5 deletions
82
Cargo.lock
generated
82
Cargo.lock
generated
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
371
crates/attractor/src/handler/sub_pipeline.rs
Normal file
371
crates/attractor/src/handler/sub_pipeline.rs
Normal 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(¤t_node_id) {
|
||||
Some(n) => n,
|
||||
None => {
|
||||
return Ok(Outcome::fail(format!(
|
||||
"Sub-pipeline node not found: {current_node_id}"
|
||||
)));
|
||||
}
|
||||
};
|
||||
|
||||
// Check for terminal node
|
||||
if is_terminal(sub_node) {
|
||||
break;
|
||||
}
|
||||
|
||||
// Execute the node handler
|
||||
let handler = self.registry.resolve(sub_node);
|
||||
last_outcome = handler
|
||||
.execute(sub_node, &sub_context, &sub_graph, &sub_logs_root)
|
||||
.await?;
|
||||
|
||||
// Apply context updates from the outcome
|
||||
sub_context.apply_updates(&last_outcome.context_updates);
|
||||
sub_context.set("outcome", serde_json::json!(last_outcome.status.to_string()));
|
||||
|
||||
// Select next edge
|
||||
match select_edge(¤t_node_id, &last_outcome, &sub_context, &sub_graph) {
|
||||
Some(edge) => {
|
||||
current_node_id.clone_from(&edge.to);
|
||||
}
|
||||
None => break,
|
||||
}
|
||||
}
|
||||
|
||||
// 6. Compute context diff (sub_context changes vs parent's original snapshot)
|
||||
let after_snapshot = sub_context.snapshot();
|
||||
let mut context_updates = std::collections::HashMap::new();
|
||||
for (key, value) in &after_snapshot {
|
||||
match before_snapshot.get(key) {
|
||||
Some(old_value) if old_value == value => {}
|
||||
_ => {
|
||||
context_updates.insert(key.clone(), value.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 7. Return the last outcome with context updates propagated
|
||||
let mut result = last_outcome;
|
||||
result.context_updates.extend(context_updates);
|
||||
Ok(result)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::graph::AttrValue;
|
||||
use crate::handler::exit::ExitHandler;
|
||||
use crate::handler::start::StartHandler;
|
||||
use crate::outcome::StageStatus;
|
||||
|
||||
fn make_registry() -> Arc<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);
|
||||
}
|
||||
}
|
||||
|
|
@ -3,6 +3,7 @@ pub mod callback;
|
|||
pub mod console;
|
||||
pub mod queue;
|
||||
pub mod recording;
|
||||
pub mod web;
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
|
|
|
|||
256
crates/attractor/src/interviewer/web.rs
Normal file
256
crates/attractor/src/interviewer/web.rs
Normal 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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
727
crates/attractor/src/server.rs
Normal file
727
crates/attractor/src/server.rs
Normal 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");
|
||||
}
|
||||
}
|
||||
|
|
@ -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
Loading…
Add table
Reference in a new issue