use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use serde::{Deserialize, Serialize};
use tokio::sync::Mutex as AsyncMutex;
use crate::tenant::TenantContext;
use super::checkpoint::{CheckpointError, CheckpointStore, GraphRunRecord, RunStatus};
use super::model::{GraphConfig, GraphError, NodeKind};
use super::router::route;
use super::state::{GraphRole, GraphRunState};
pub type BoxFut<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
#[derive(Debug, Clone)]
pub struct AgentTurnRequest {
pub node_id: String,
pub system_prompt: String,
pub model: String,
pub state: GraphRunState,
pub provider: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentTurnResult {
pub reply: String,
pub resolved: bool,
}
#[derive(Debug, Clone)]
pub struct ToolCallRequest {
pub node_id: String,
pub tool_name: String,
pub state: GraphRunState,
}
pub type AgentTurnFn = Arc<
dyn Fn(AgentTurnRequest) -> BoxFut<'static, Result<AgentTurnResult, GraphExecError>>
+ Send
+ Sync,
>;
pub type ToolFn = Arc<
dyn Fn(ToolCallRequest) -> BoxFut<'static, Result<serde_json::Value, GraphExecError>>
+ Send
+ Sync,
>;
#[derive(Debug, Clone)]
pub struct SupervisorRequest {
pub node_id: String,
pub system_prompt: String,
pub model: String,
pub routes: Vec<crate::graph::model::SupervisorRoute>,
pub state: GraphRunState,
pub provider: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SupervisorResult {
pub branch: String,
pub raw_reply: String,
}
pub type SupervisorFn = Arc<
dyn Fn(SupervisorRequest) -> BoxFut<'static, Result<SupervisorResult, GraphExecError>>
+ Send
+ Sync,
>;
#[derive(Debug, Clone)]
pub struct ApprovalRequest {
pub run_id: String,
pub node_id: String,
pub tenant: String,
pub title: String,
pub mode: String,
pub risk_threshold: Option<f64>,
pub confidence_threshold: Option<f64>,
pub deadline_ms: Option<u64>,
pub state: GraphRunState,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum ApprovalOutcome {
Awaiting,
Decided { branch: String },
}
pub type ApprovalFn = Arc<
dyn Fn(ApprovalRequest) -> BoxFut<'static, Result<ApprovalOutcome, GraphExecError>>
+ Send
+ Sync,
>;
#[derive(Debug, thiserror::Error)]
pub enum GraphExecError {
#[error("graph run {run_id} exceeded the node-visit cap")]
IterationCap { run_id: String },
#[error("unknown node `{0}` (cursor corrupt or graph changed)")]
UnknownNode(String),
#[error("unknown run `{0}`")]
UnknownRun(String),
#[error("run `{0}` already completed")]
AlreadyCompleted(String),
#[error(transparent)]
Graph(#[from] GraphError),
#[error(transparent)]
Checkpoint(#[from] CheckpointError),
#[error("agent turn failed: {0}")]
AgentTurn(String),
#[error("tool call failed: {0}")]
Tool(String),
#[error("supervisor routing failed: {0}")]
Supervisor(String),
}
#[derive(Debug, Clone)]
pub struct GraphRunOutcome {
pub status: RunStatus,
pub reply: String,
pub trail: Vec<serde_json::Value>,
}
pub const MAX_NODE_VISITS: u32 = 64;
#[derive(Debug, Clone, Serialize, Deserialize)]
struct BranchCursor {
branch: String,
cursor: String,
state_json: String,
parked: bool,
}
struct FrontierCheckpoint {
run_id: String,
graph_json: String,
parallel_node: String,
trunk_state_json: String,
trunk_visits: HashMap<String, u32>,
inner: AsyncMutex<Vec<BranchCursor>>,
branch_visits: AsyncMutex<HashMap<String, u32>>,
}
pub struct GraphExecutor {
store: Arc<dyn CheckpointStore>,
agent_turn: AgentTurnFn,
tool: ToolFn,
supervisor: SupervisorFn,
approval: ApprovalFn,
}
impl GraphExecutor {
pub fn new(
store: Arc<dyn CheckpointStore>,
agent_turn: AgentTurnFn,
tool: ToolFn,
supervisor: SupervisorFn,
approval: ApprovalFn,
) -> Self {
Self {
store,
agent_turn,
tool,
supervisor,
approval,
}
}
pub async fn start(
&self,
tenant: &TenantContext,
run_id: &str,
cfg: &GraphConfig,
user_text: &str,
) -> Result<GraphRunOutcome, GraphExecError> {
if let Some(existing) = self.store.load(tenant, run_id).await? {
return match existing.status {
RunStatus::Running | RunStatus::AwaitingInput => {
self.drive_from_record(tenant, run_id, existing).await
}
RunStatus::Succeeded | RunStatus::Failed => {
Err(GraphExecError::AlreadyCompleted(run_id.to_owned()))
}
};
}
let mut state = GraphRunState::default();
state.push_message(GraphRole::User, user_text);
let graph_json = serde_json::to_string(cfg)
.map_err(|e| GraphExecError::Checkpoint(CheckpointError::Serde(e)))?;
let cursor = cfg.graph.entry.clone();
let visits: HashMap<String, u32> = HashMap::new();
let rec = build_record(
run_id,
&graph_json,
&cursor,
&state,
&visits,
RunStatus::Running,
)?;
self.store.save(tenant, &rec).await?;
self.drive(tenant, run_id, cfg.clone(), cursor, state, visits)
.await
}
pub async fn resume(
&self,
tenant: &TenantContext,
run_id: &str,
) -> Result<GraphRunOutcome, GraphExecError> {
let rec = self
.store
.load(tenant, run_id)
.await?
.ok_or_else(|| GraphExecError::UnknownRun(run_id.to_owned()))?;
match rec.status {
RunStatus::Succeeded | RunStatus::Failed => {
let state: GraphRunState =
serde_json::from_str(&rec.state_json).map_err(CheckpointError::Serde)?;
let reply = last_assistant_message(&state);
let trail = rebuild_trail_from_state(&state);
Ok(GraphRunOutcome {
status: rec.status,
reply,
trail,
})
}
RunStatus::Running | RunStatus::AwaitingInput => {
self.drive_from_record(tenant, run_id, rec).await
}
}
}
async fn drive_from_record(
&self,
tenant: &TenantContext,
run_id: &str,
rec: GraphRunRecord,
) -> Result<GraphRunOutcome, GraphExecError> {
let cfg: GraphConfig = GraphConfig::from_json(&rec.graph_json)?;
let state: GraphRunState =
serde_json::from_str(&rec.state_json).map_err(CheckpointError::Serde)?;
let visits: HashMap<String, u32> =
serde_json::from_str(&rec.visits_json).map_err(CheckpointError::Serde)?;
if let Some(frontier_json) = &rec.frontier_json {
let frontier: Vec<BranchCursor> =
serde_json::from_str(frontier_json).map_err(CheckpointError::Serde)?;
return self
.resume_parallel(tenant, run_id, &cfg, &rec.cursor, state, visits, frontier)
.await;
}
let cursor = rec.cursor.clone();
self.drive(tenant, run_id, cfg, cursor, state, visits).await
}
async fn drive(
&self,
tenant: &TenantContext,
run_id: &str,
cfg: GraphConfig,
mut cursor: String,
mut state: GraphRunState,
mut visits: HashMap<String, u32>,
) -> Result<GraphRunOutcome, GraphExecError> {
let mut trail: Vec<serde_json::Value> = Vec::new();
for _ in 0..MAX_NODE_VISITS {
let node = cfg
.graph
.node(&cursor)
.ok_or_else(|| GraphExecError::UnknownNode(cursor.clone()))?
.clone();
match &node.kind {
NodeKind::Agent {
system_prompt,
model,
provider,
..
} => {
let attempt = *visits.get(&cursor).unwrap_or(&0) + 1;
let node_id_for_err = cursor.clone();
let provider_clone = provider.clone();
let (raw, replayed) = self
.visit_effect(tenant, run_id, &cursor, attempt, || {
let req = AgentTurnRequest {
node_id: node_id_for_err.clone(),
system_prompt: system_prompt.clone(),
model: model.clone(),
state: state.clone(),
provider: provider_clone,
};
let fut = (self.agent_turn)(req);
Box::pin(async move {
let r = fut.await.map_err(|e| {
GraphExecError::AgentTurn(format!(
"node '{}' attempt {}: {}",
node_id_for_err, attempt, e
))
})?;
serde_json::to_value(&r)
.map_err(CheckpointError::Serde)
.map_err(GraphExecError::Checkpoint)
})
})
.await?;
let result: AgentTurnResult =
serde_json::from_value(raw).map_err(CheckpointError::Serde)?;
trail.push(serde_json::json!({
"node": cursor,
"kind": "agent",
"attempt": attempt,
"replayed": replayed,
}));
visits.insert(cursor.clone(), attempt);
state.iterations += 1;
state.push_message(GraphRole::Assistant, &result.reply);
if result.resolved {
state.resolved = true;
}
cursor = next_linear(&cfg, &cursor)?;
let rec = build_record(
run_id,
&serde_json::to_string(&cfg).map_err(CheckpointError::Serde)?,
&cursor,
&state,
&visits,
RunStatus::Running,
)?;
self.store.save(tenant, &rec).await?;
}
NodeKind::Tool { tool_name } => {
let attempt = *visits.get(&cursor).unwrap_or(&0) + 1;
let node_id_for_err = cursor.clone();
let (result, replayed) = self
.visit_effect(tenant, run_id, &cursor, attempt, || {
let req = ToolCallRequest {
node_id: node_id_for_err.clone(),
tool_name: tool_name.clone(),
state: state.clone(),
};
let fut = (self.tool)(req);
Box::pin(async move {
fut.await.map_err(|e| {
GraphExecError::Tool(format!(
"node '{}' attempt {}: {}",
node_id_for_err, attempt, e
))
})
})
})
.await?;
trail.push(serde_json::json!({
"node": cursor,
"kind": "tool",
"attempt": attempt,
"replayed": replayed,
}));
visits.insert(cursor.clone(), attempt);
state.push_message(GraphRole::Tool, result.to_string());
cursor = next_linear(&cfg, &cursor)?;
let rec = build_record(
run_id,
&serde_json::to_string(&cfg).map_err(CheckpointError::Serde)?,
&cursor,
&state,
&visits,
RunStatus::Running,
)?;
self.store.save(tenant, &rec).await?;
}
NodeKind::Router { .. } => {
let attempt = *visits.get(&cursor).unwrap_or(&0) + 1;
let next = route(&cfg.graph, &cursor, &state)?;
trail.push(serde_json::json!({
"node": cursor,
"kind": "router",
"attempt": attempt,
"replayed": false,
}));
visits.insert(cursor.clone(), attempt);
cursor = next;
let rec = build_record(
run_id,
&serde_json::to_string(&cfg).map_err(CheckpointError::Serde)?,
&cursor,
&state,
&visits,
RunStatus::Running,
)?;
self.store.save(tenant, &rec).await?;
}
NodeKind::Respond => {
let attempt = *visits.get(&cursor).unwrap_or(&0) + 1;
trail.push(serde_json::json!({
"node": cursor,
"kind": "respond",
"attempt": attempt,
"replayed": false,
}));
visits.insert(cursor.clone(), attempt);
let rec = build_record(
run_id,
&serde_json::to_string(&cfg).map_err(CheckpointError::Serde)?,
&cursor,
&state,
&visits,
RunStatus::Succeeded,
)?;
self.store.save(tenant, &rec).await?;
let reply = last_assistant_message(&state);
return Ok(GraphRunOutcome {
status: RunStatus::Succeeded,
reply,
trail,
});
}
NodeKind::Supervisor {
system_prompt,
model,
routes,
provider,
} => {
let attempt = *visits.get(&cursor).unwrap_or(&0) + 1;
let node_id_for_err = cursor.clone();
let routes_clone = routes.clone();
let system_prompt_clone = system_prompt.clone();
let model_clone = model.clone();
let provider_clone = provider.clone();
let (raw, replayed) = self
.visit_effect(tenant, run_id, &cursor, attempt, || {
let req = SupervisorRequest {
node_id: node_id_for_err.clone(),
system_prompt: system_prompt_clone,
model: model_clone,
routes: routes_clone.clone(),
state: state.clone(),
provider: provider_clone,
};
let fut = (self.supervisor)(req);
Box::pin(async move {
let r = fut.await.map_err(|e| {
GraphExecError::Supervisor(format!(
"node '{}' attempt {}: {}",
node_id_for_err, attempt, e
))
})?;
serde_json::to_value(&r)
.map_err(CheckpointError::Serde)
.map_err(GraphExecError::Checkpoint)
})
})
.await?;
let result: SupervisorResult =
serde_json::from_value(raw).map_err(CheckpointError::Serde)?;
let branch = &result.branch;
let branch_is_valid_route = routes.iter().any(|r| &r.branch == branch);
let matching_edge = cfg
.graph
.edges_from(&cursor)
.find(|e| e.branch.as_deref() == Some(branch.as_str()));
let next_cursor = match (branch_is_valid_route, matching_edge) {
(true, Some(edge)) => edge.to.clone(),
_ => {
return Err(GraphExecError::Graph(super::model::GraphError::Invalid(
format!(
"supervisor node '{}': branch '{}' returned by supervisor \
does not match any declared route or outgoing edge",
cursor, branch
),
)));
}
};
trail.push(serde_json::json!({
"node": cursor,
"kind": "supervisor",
"attempt": attempt,
"replayed": replayed,
"branch": result.branch,
}));
visits.insert(cursor.clone(), attempt);
state.push_message(GraphRole::Assistant, &result.raw_reply);
cursor = next_cursor;
let rec = build_record(
run_id,
&serde_json::to_string(&cfg).map_err(CheckpointError::Serde)?,
&cursor,
&state,
&visits,
RunStatus::Running,
)?;
self.store.save(tenant, &rec).await?;
}
NodeKind::Parallel => {
let (next_cursor, merged_state, merged_visits, region_trail) = self
.run_parallel_region(tenant, run_id, &cfg, &cursor, &state, &visits)
.await?;
trail.extend(region_trail);
cursor = next_cursor;
state = merged_state;
visits = merged_visits;
}
NodeKind::Join => {
tracing::warn!(
node = %cursor,
"trunk drive loop reached a join node directly; \
passing through to its successor (expected to be \
consumed by the parallel arm)"
);
cursor = next_linear(&cfg, &cursor)?;
}
NodeKind::Approval {
title,
mode,
risk_threshold,
confidence_threshold,
deadline_ms,
} => {
let req = ApprovalRequest {
run_id: run_id.to_string(),
node_id: cursor.clone(),
tenant: tenant.tenant_id.clone(),
title: title.clone(),
mode: mode.clone(),
risk_threshold: *risk_threshold,
confidence_threshold: *confidence_threshold,
deadline_ms: *deadline_ms,
state: state.clone(),
};
match (self.approval)(req).await? {
ApprovalOutcome::Awaiting => {
let rec = build_record(
run_id,
&serde_json::to_string(&cfg).map_err(CheckpointError::Serde)?,
&cursor,
&state,
&visits,
RunStatus::AwaitingInput,
)?;
self.store.save(tenant, &rec).await?;
let reply = last_assistant_message(&state);
return Ok(GraphRunOutcome {
status: RunStatus::AwaitingInput,
reply,
trail,
});
}
ApprovalOutcome::Decided { branch } => {
let matching_edge = cfg
.graph
.edges_from(&cursor)
.find(|e| e.branch.as_deref() == Some(branch.as_str()));
let next_cursor = match matching_edge {
Some(edge) => edge.to.clone(),
None => {
return Err(GraphExecError::Graph(GraphError::Invalid(
format!(
"approval node '{}': decision branch '{}' does not \
match any outgoing edge",
cursor, branch
),
)));
}
};
let attempt = *visits.get(&cursor).unwrap_or(&0) + 1;
self.store
.record_node_visit(
tenant,
run_id,
&cursor,
attempt,
&serde_json::json!({"decision": branch}),
)
.await?;
trail.push(serde_json::json!({
"node": cursor,
"kind": "approval",
"attempt": attempt,
"replayed": false,
"branch": branch,
}));
visits.insert(cursor.clone(), attempt);
cursor = next_cursor;
let rec = build_record(
run_id,
&serde_json::to_string(&cfg).map_err(CheckpointError::Serde)?,
&cursor,
&state,
&visits,
RunStatus::Running,
)?;
self.store.save(tenant, &rec).await?;
}
}
}
}
}
let graph_json = serde_json::to_string(&cfg).map_err(CheckpointError::Serde)?;
let rec = build_record(
run_id,
&graph_json,
&cursor,
&state,
&visits,
RunStatus::Failed,
)?;
self.store.save(tenant, &rec).await?;
Err(GraphExecError::IterationCap {
run_id: run_id.to_owned(),
})
}
async fn visit_effect(
&self,
tenant: &TenantContext,
run_id: &str,
node_id: &str,
attempt: u32,
invoke: impl FnOnce() -> BoxFut<'static, Result<serde_json::Value, GraphExecError>>,
) -> Result<(serde_json::Value, bool), GraphExecError> {
if let Some(cached) = self
.store
.load_node_visit(tenant, run_id, node_id, attempt)
.await?
{
return Ok((cached, true));
}
let value = invoke().await?;
self.store
.record_node_visit(tenant, run_id, node_id, attempt, &value)
.await?;
Ok((value, false))
}
#[allow(clippy::type_complexity)]
async fn run_parallel_region(
&self,
tenant: &TenantContext,
run_id: &str,
cfg: &GraphConfig,
parallel_node: &str,
trunk_state: &GraphRunState,
trunk_visits: &HashMap<String, u32>,
) -> Result<
(
String,
GraphRunState,
HashMap<String, u32>,
Vec<serde_json::Value>,
),
GraphExecError,
> {
let mut branch_edges: Vec<(String, String)> = cfg
.graph
.edges_from(parallel_node)
.filter_map(|e| e.branch.clone().map(|b| (b, e.to.clone())))
.collect();
branch_edges.sort_by(|a, b| a.0.cmp(&b.0));
let join_id = find_join_for_parallel(cfg, parallel_node)?;
let trunk_state_json =
serde_json::to_string(trunk_state).map_err(CheckpointError::Serde)?;
let frontier: Vec<BranchCursor> = branch_edges
.iter()
.map(|(branch, target)| BranchCursor {
branch: branch.clone(),
cursor: target.clone(),
state_json: trunk_state_json.clone(),
parked: false,
})
.collect();
self.drive_frontier(
tenant,
run_id,
cfg,
parallel_node,
&join_id,
trunk_state,
trunk_visits,
frontier,
true,
)
.await
}
#[allow(clippy::too_many_arguments, clippy::type_complexity)]
async fn resume_parallel(
&self,
tenant: &TenantContext,
run_id: &str,
cfg: &GraphConfig,
parallel_node: &str,
trunk_state: GraphRunState,
trunk_visits: HashMap<String, u32>,
frontier: Vec<BranchCursor>,
) -> Result<GraphRunOutcome, GraphExecError> {
let join_id = find_join_for_parallel(cfg, parallel_node)?;
let (next_cursor, merged_state, merged_visits, _region_trail) = self
.drive_frontier(
tenant,
run_id,
cfg,
parallel_node,
&join_id,
&trunk_state,
&trunk_visits,
frontier,
false,
)
.await?;
self.drive(
tenant,
run_id,
cfg.clone(),
next_cursor,
merged_state,
merged_visits,
)
.await
}
#[allow(clippy::too_many_arguments, clippy::type_complexity)]
async fn drive_frontier(
&self,
tenant: &TenantContext,
run_id: &str,
cfg: &GraphConfig,
parallel_node: &str,
join_id: &str,
trunk_state: &GraphRunState,
trunk_visits: &HashMap<String, u32>,
frontier: Vec<BranchCursor>,
persist_before_driving: bool,
) -> Result<
(
String,
GraphRunState,
HashMap<String, u32>,
Vec<serde_json::Value>,
),
GraphExecError,
> {
let graph_json = serde_json::to_string(cfg).map_err(CheckpointError::Serde)?;
let trunk_state_json =
serde_json::to_string(trunk_state).map_err(CheckpointError::Serde)?;
let coord = Arc::new(FrontierCheckpoint {
run_id: run_id.to_owned(),
graph_json: graph_json.clone(),
parallel_node: parallel_node.to_owned(),
trunk_state_json,
trunk_visits: trunk_visits.clone(),
inner: AsyncMutex::new(frontier.clone()),
branch_visits: AsyncMutex::new(HashMap::new()),
});
if persist_before_driving {
coord.checkpoint(self.store.as_ref(), tenant).await?;
}
let trunk_visit_total: u32 = trunk_visits.values().copied().sum();
let global_visits = Arc::new(AtomicU32::new(trunk_visit_total));
let mut futs: Vec<BoxFut<'_, Result<BranchOutcome, GraphExecError>>> = Vec::new();
for (slot, bc) in frontier.iter().enumerate() {
if bc.parked {
continue;
}
let bc = bc.clone();
let coord = coord.clone();
let global_visits = global_visits.clone();
futs.push(Box::pin(self.drive_branch(
tenant,
run_id,
cfg,
join_id,
slot,
bc,
trunk_visits.clone(),
coord,
global_visits,
)));
}
let results = futures::future::join_all(futs).await;
let mut first_err: Option<GraphExecError> = None;
let mut trail: Vec<serde_json::Value> = Vec::new();
let mut merged_visits = trunk_visits.clone();
for r in results {
match r {
Ok(outcome) => {
trail.extend(outcome.trail);
for (k, v) in outcome.visits {
merged_visits.insert(k, v);
}
}
Err(e) if first_err.is_none() => first_err = Some(e),
Err(_) => { }
}
}
if let Some(e) = first_err {
if matches!(e, GraphExecError::IterationCap { .. }) {
let last_frontier = coord.inner.lock().await.clone();
let frontier_json =
Some(serde_json::to_string(&last_frontier).map_err(CheckpointError::Serde)?);
let rec = build_record_with_frontier(
run_id,
&graph_json,
parallel_node,
trunk_state,
trunk_visits,
RunStatus::Failed,
frontier_json,
)?;
self.store.save(tenant, &rec).await?;
}
return Err(e);
}
let mut ordered = coord.inner.lock().await.clone();
ordered.sort_by(|a, b| a.branch.cmp(&b.branch));
let snapshot_len = trunk_state.messages.len();
let mut merged_state = trunk_state.clone();
for bc in &ordered {
let branch_state: GraphRunState =
serde_json::from_str(&bc.state_json).map_err(CheckpointError::Serde)?;
for msg in branch_state.messages.iter().skip(snapshot_len) {
merged_state.messages.push(msg.clone());
}
if branch_state.resolved {
merged_state.resolved = true;
}
merged_state.iterations = merged_state.iterations.max(branch_state.iterations);
if !branch_state.scratchpad.is_null() {
if !merged_state.scratchpad.is_object() {
merged_state.scratchpad = serde_json::json!({});
}
if let Some(obj) = merged_state.scratchpad.as_object_mut() {
let branches = obj
.entry("branches")
.or_insert_with(|| serde_json::json!({}));
if let Some(branches_obj) = branches.as_object_mut() {
branches_obj.insert(bc.branch.clone(), branch_state.scratchpad.clone());
}
}
}
}
for (k, v) in coord.collect_branch_visits().await {
merged_visits.entry(k).or_insert(v);
}
let next_cursor = next_linear(cfg, join_id)?;
let rec = build_record(
run_id,
&graph_json,
&next_cursor,
&merged_state,
&merged_visits,
RunStatus::Running,
)?;
self.store.save(tenant, &rec).await?;
Ok((next_cursor, merged_state, merged_visits, trail))
}
#[allow(clippy::too_many_arguments)]
fn drive_branch<'a>(
&'a self,
tenant: &'a TenantContext,
run_id: &'a str,
cfg: &'a GraphConfig,
join_id: &'a str,
slot: usize,
mut bc: BranchCursor,
mut visits: HashMap<String, u32>,
coord: Arc<FrontierCheckpoint>,
global_visits: Arc<AtomicU32>,
) -> BoxFut<'a, Result<BranchOutcome, GraphExecError>> {
Box::pin(async move {
let mut state: GraphRunState =
serde_json::from_str(&bc.state_json).map_err(CheckpointError::Serde)?;
let mut trail: Vec<serde_json::Value> = Vec::new();
let branch_label = bc.branch.clone();
loop {
if bc.cursor == join_id {
bc.parked = true;
bc.state_json =
serde_json::to_string(&state).map_err(CheckpointError::Serde)?;
coord
.save_slot(self.store.as_ref(), tenant, slot, &bc)
.await?;
return Ok(BranchOutcome {
branch: branch_label,
visits,
trail,
});
}
let total = global_visits.fetch_add(1, Ordering::SeqCst) + 1;
if total > MAX_NODE_VISITS {
return Err(GraphExecError::IterationCap {
run_id: run_id.to_owned(),
});
}
let node = cfg
.graph
.node(&bc.cursor)
.ok_or_else(|| GraphExecError::UnknownNode(bc.cursor.clone()))?
.clone();
match &node.kind {
NodeKind::Agent {
system_prompt,
model,
provider,
..
} => {
let attempt = *visits.get(&bc.cursor).unwrap_or(&0) + 1;
let node_id = bc.cursor.clone();
let sp = system_prompt.clone();
let md = model.clone();
let pv = provider.clone();
let state_for_call = state.clone();
let (raw, replayed) = self
.visit_effect(tenant, run_id, &bc.cursor, attempt, || {
let req = AgentTurnRequest {
node_id: node_id.clone(),
system_prompt: sp,
model: md,
state: state_for_call,
provider: pv,
};
let fut = (self.agent_turn)(req);
Box::pin(async move {
let r = fut.await.map_err(|e| {
GraphExecError::AgentTurn(format!(
"node '{}' attempt {}: {}",
node_id, attempt, e
))
})?;
serde_json::to_value(&r)
.map_err(CheckpointError::Serde)
.map_err(GraphExecError::Checkpoint)
})
})
.await?;
let result: AgentTurnResult =
serde_json::from_value(raw).map_err(CheckpointError::Serde)?;
trail.push(serde_json::json!({
"node": bc.cursor, "kind": "agent",
"attempt": attempt, "replayed": replayed, "branch": branch_label,
}));
visits.insert(bc.cursor.clone(), attempt);
state.iterations += 1;
state.push_message(GraphRole::Assistant, &result.reply);
if result.resolved {
state.resolved = true;
}
bc.cursor = next_linear(cfg, &bc.cursor)?;
}
NodeKind::Tool { tool_name } => {
let attempt = *visits.get(&bc.cursor).unwrap_or(&0) + 1;
let node_id = bc.cursor.clone();
let tn = tool_name.clone();
let state_for_call = state.clone();
let (result, replayed) = self
.visit_effect(tenant, run_id, &bc.cursor, attempt, || {
let req = ToolCallRequest {
node_id: node_id.clone(),
tool_name: tn,
state: state_for_call,
};
let fut = (self.tool)(req);
Box::pin(async move {
fut.await.map_err(|e| {
GraphExecError::Tool(format!(
"node '{}' attempt {}: {}",
node_id, attempt, e
))
})
})
})
.await?;
trail.push(serde_json::json!({
"node": bc.cursor, "kind": "tool",
"attempt": attempt, "replayed": replayed, "branch": branch_label,
}));
visits.insert(bc.cursor.clone(), attempt);
state.push_message(GraphRole::Tool, result.to_string());
bc.cursor = next_linear(cfg, &bc.cursor)?;
}
NodeKind::Router { .. } => {
let attempt = *visits.get(&bc.cursor).unwrap_or(&0) + 1;
let next = route(&cfg.graph, &bc.cursor, &state)?;
trail.push(serde_json::json!({
"node": bc.cursor, "kind": "router",
"attempt": attempt, "replayed": false, "branch": branch_label,
}));
visits.insert(bc.cursor.clone(), attempt);
bc.cursor = next;
}
NodeKind::Supervisor {
system_prompt,
model,
routes,
provider,
} => {
let attempt = *visits.get(&bc.cursor).unwrap_or(&0) + 1;
let node_id = bc.cursor.clone();
let sp = system_prompt.clone();
let md = model.clone();
let pv = provider.clone();
let routes_clone = routes.clone();
let state_for_call = state.clone();
let (raw, replayed) = self
.visit_effect(tenant, run_id, &bc.cursor, attempt, || {
let req = SupervisorRequest {
node_id: node_id.clone(),
system_prompt: sp,
model: md,
routes: routes_clone.clone(),
state: state_for_call,
provider: pv,
};
let fut = (self.supervisor)(req);
Box::pin(async move {
let r = fut.await.map_err(|e| {
GraphExecError::Supervisor(format!(
"node '{}' attempt {}: {}",
node_id, attempt, e
))
})?;
serde_json::to_value(&r)
.map_err(CheckpointError::Serde)
.map_err(GraphExecError::Checkpoint)
})
})
.await?;
let result: SupervisorResult =
serde_json::from_value(raw).map_err(CheckpointError::Serde)?;
let branch = &result.branch;
let matching_edge = cfg
.graph
.edges_from(&bc.cursor)
.find(|e| e.branch.as_deref() == Some(branch.as_str()));
let next_cursor =
match (routes.iter().any(|r| &r.branch == branch), matching_edge) {
(true, Some(edge)) => edge.to.clone(),
_ => {
return Err(GraphExecError::Graph(GraphError::Invalid(
format!(
"supervisor node '{}': branch '{}' does not match any \
declared route or outgoing edge",
bc.cursor, branch
),
)));
}
};
trail.push(serde_json::json!({
"node": bc.cursor, "kind": "supervisor",
"attempt": attempt, "replayed": replayed,
"branch": result.branch, "branch_path": branch_label,
}));
visits.insert(bc.cursor.clone(), attempt);
state.push_message(GraphRole::Assistant, &result.raw_reply);
bc.cursor = next_cursor;
}
NodeKind::Respond => {
return Err(GraphExecError::Graph(GraphError::Invalid(format!(
"respond node '{}' inside parallel branch '{}' is not allowed",
bc.cursor, branch_label
))));
}
NodeKind::Parallel | NodeKind::Join => {
return Err(GraphExecError::Graph(GraphError::Invalid(format!(
"branch '{}' reached unexpected '{}' node '{}'",
branch_label,
node.kind.kind_name(),
bc.cursor
))));
}
NodeKind::Approval { .. } => {
return Err(GraphExecError::Graph(GraphError::Invalid(format!(
"approval node '{}' inside parallel branch '{}': not supported in v1 (parallel-branch approval parking is unimplemented)",
bc.cursor, branch_label
))));
}
}
bc.state_json = serde_json::to_string(&state).map_err(CheckpointError::Serde)?;
coord
.save_slot(self.store.as_ref(), tenant, slot, &bc)
.await?;
coord.record_branch_visits(&visits).await;
}
})
}
}
struct BranchOutcome {
#[allow(dead_code)]
branch: String,
visits: HashMap<String, u32>,
trail: Vec<serde_json::Value>,
}
impl FrontierCheckpoint {
async fn save_slot(
&self,
store: &dyn CheckpointStore,
tenant: &TenantContext,
slot: usize,
bc: &BranchCursor,
) -> Result<(), GraphExecError> {
let mut guard = self.inner.lock().await;
if let Some(existing) = guard.get_mut(slot) {
*existing = bc.clone();
}
let frontier_json = Some(serde_json::to_string(&*guard).map_err(CheckpointError::Serde)?);
let trunk_state: GraphRunState =
serde_json::from_str(&self.trunk_state_json).map_err(CheckpointError::Serde)?;
let rec = build_record_with_frontier(
&self.run_id,
&self.graph_json,
&self.parallel_node,
&trunk_state,
&self.trunk_visits,
RunStatus::Running,
frontier_json,
)?;
store.save(tenant, &rec).await?;
Ok(())
}
async fn checkpoint(
&self,
store: &dyn CheckpointStore,
tenant: &TenantContext,
) -> Result<(), GraphExecError> {
let guard = self.inner.lock().await;
let frontier_json = Some(serde_json::to_string(&*guard).map_err(CheckpointError::Serde)?);
let trunk_state: GraphRunState =
serde_json::from_str(&self.trunk_state_json).map_err(CheckpointError::Serde)?;
let rec = build_record_with_frontier(
&self.run_id,
&self.graph_json,
&self.parallel_node,
&trunk_state,
&self.trunk_visits,
RunStatus::Running,
frontier_json,
)?;
store.save(tenant, &rec).await?;
Ok(())
}
async fn record_branch_visits(&self, visits: &HashMap<String, u32>) {
let mut guard = self.branch_visits.lock().await;
for (k, v) in visits {
guard.insert(k.clone(), *v);
}
}
async fn collect_branch_visits(&self) -> HashMap<String, u32> {
self.branch_visits.lock().await.clone()
}
}
fn build_record(
run_id: &str,
graph_json: &str,
cursor: &str,
state: &GraphRunState,
visits: &HashMap<String, u32>,
status: RunStatus,
) -> Result<GraphRunRecord, GraphExecError> {
build_record_with_frontier(run_id, graph_json, cursor, state, visits, status, None)
}
fn build_record_with_frontier(
run_id: &str,
graph_json: &str,
cursor: &str,
state: &GraphRunState,
visits: &HashMap<String, u32>,
status: RunStatus,
frontier_json: Option<String>,
) -> Result<GraphRunRecord, GraphExecError> {
let state_json = serde_json::to_string(state).map_err(CheckpointError::Serde)?;
let visits_json = serde_json::to_string(visits).map_err(CheckpointError::Serde)?;
Ok(GraphRunRecord {
run_id: run_id.to_owned(),
graph_json: graph_json.to_owned(),
cursor: cursor.to_owned(),
state_json,
status,
visits_json,
frontier_json,
})
}
fn next_linear(cfg: &GraphConfig, id: &str) -> Result<String, GraphExecError> {
cfg.graph
.edges_from(id)
.next()
.map(|e| e.to.clone())
.ok_or_else(|| {
GraphExecError::Graph(super::model::GraphError::Invalid(format!(
"node '{id}' has no outgoing edge"
)))
})
}
fn find_join_for_parallel(
cfg: &GraphConfig,
parallel_node: &str,
) -> Result<String, GraphExecError> {
use std::collections::HashSet;
let mut visited: HashSet<String> = HashSet::new();
let mut queue: Vec<String> = cfg
.graph
.edges_from(parallel_node)
.map(|e| e.to.clone())
.collect();
while let Some(current) = queue.pop() {
if !visited.insert(current.clone()) {
continue;
}
match cfg.graph.node(¤t) {
Some(n) if matches!(n.kind, NodeKind::Join) => return Ok(current),
Some(_) => {
for e in cfg.graph.edges_from(¤t) {
queue.push(e.to.clone());
}
}
None => {
return Err(GraphExecError::UnknownNode(current));
}
}
}
Err(GraphExecError::Graph(GraphError::Invalid(format!(
"parallel node '{parallel_node}' has no reachable join node"
))))
}
fn last_assistant_message(state: &GraphRunState) -> String {
state
.messages
.iter()
.rev()
.find(|m| m.role == GraphRole::Assistant)
.map(|m| m.content.clone())
.unwrap_or_default()
}
fn rebuild_trail_from_state(state: &GraphRunState) -> Vec<serde_json::Value> {
state
.messages
.iter()
.map(|m| {
let kind = match m.role {
GraphRole::User => "user",
GraphRole::Assistant => "agent",
GraphRole::Tool => "tool",
};
serde_json::json!({"kind": kind, "content": m.content})
})
.collect()
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use super::*;
use crate::graph::test_fixtures::{parallel_json, supervisor_json, triage_json};
use crate::graph::{GraphConfig, InMemoryCheckpointStore};
use crate::tenant::TenantContext;
fn tenant() -> TenantContext {
TenantContext::new("test", "dev")
}
fn triage_cfg() -> GraphConfig {
GraphConfig::from_json(&triage_json()).expect("fixture is valid")
}
fn supervisor_cfg() -> GraphConfig {
GraphConfig::from_json(&supervisor_json()).expect("supervisor fixture is valid")
}
fn parallel_cfg() -> GraphConfig {
GraphConfig::from_json(¶llel_json()).expect("parallel fixture is valid")
}
fn agent_fn_resolves_on(counter: Arc<AtomicU32>, resolve_on_call: u32) -> AgentTurnFn {
Arc::new(move |req: AgentTurnRequest| {
let n = counter.fetch_add(1, Ordering::SeqCst) + 1; let resolved = n >= resolve_on_call;
let reply = format!("reply-{n} from {}", req.node_id);
Box::pin(async move { Ok(AgentTurnResult { reply, resolved }) })
})
}
fn tool_fn_counting(counter: Arc<AtomicU32>) -> ToolFn {
Arc::new(move |_req: ToolCallRequest| {
counter.fetch_add(1, Ordering::SeqCst);
Box::pin(async move { Ok(serde_json::json!({"found": true})) })
})
}
fn supervisor_fn_unreachable() -> SupervisorFn {
Arc::new(|_req: SupervisorRequest| {
Box::pin(async move {
Err(GraphExecError::Supervisor(
"supervisor fn should not be called in this test".into(),
))
})
})
}
fn supervisor_fn_always_routes_to(
counter: Arc<AtomicU32>,
branch: &'static str,
) -> SupervisorFn {
Arc::new(move |req: SupervisorRequest| {
counter.fetch_add(1, Ordering::SeqCst);
let node = req.node_id.clone();
Box::pin(async move {
Ok(SupervisorResult {
branch: branch.to_string(),
raw_reply: format!("[[ROUTE:{branch}]] from supervisor at {node}"),
})
})
})
}
fn approval_fn_awaiting() -> ApprovalFn {
Arc::new(|_req: ApprovalRequest| Box::pin(async move { Ok(ApprovalOutcome::Awaiting) }))
}
fn approval_fn_decides(counter: Arc<AtomicU32>, branch: &'static str) -> ApprovalFn {
Arc::new(move |_req: ApprovalRequest| {
counter.fetch_add(1, Ordering::SeqCst);
Box::pin(async move {
Ok(ApprovalOutcome::Decided {
branch: branch.to_string(),
})
})
})
}
#[tokio::test]
async fn happy_path_resolves_first_pass() {
let store = Arc::new(InMemoryCheckpointStore::default());
let agent_count = Arc::new(AtomicU32::new(0));
let tool_count = Arc::new(AtomicU32::new(0));
let exec = GraphExecutor::new(
store.clone(),
agent_fn_resolves_on(agent_count.clone(), 1),
tool_fn_counting(tool_count.clone()),
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
let outcome = exec
.start(&tenant(), "run-happy", &triage_cfg(), "help me")
.await
.expect("should succeed");
assert_eq!(outcome.status, RunStatus::Succeeded, "status");
assert!(
outcome.reply.contains("reply-1"),
"reply should contain agent output: {:?}",
outcome.reply
);
assert_eq!(agent_count.load(Ordering::SeqCst), 1, "agent invoked once");
assert_eq!(tool_count.load(Ordering::SeqCst), 1, "tool invoked once");
assert_eq!(outcome.trail.len(), 4, "trail: {:?}", outcome.trail);
}
#[tokio::test]
async fn loops_until_router_cap_then_resolves_via_cap() {
let store = Arc::new(InMemoryCheckpointStore::default());
let agent_count = Arc::new(AtomicU32::new(0));
let exec = GraphExecutor::new(
store.clone(),
agent_fn_resolves_on(agent_count.clone(), u32::MAX),
tool_fn_counting(Arc::new(AtomicU32::new(0))),
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
let outcome = exec
.start(&tenant(), "run-cap", &triage_cfg(), "loop me")
.await
.expect("should succeed via iteration cap");
assert_eq!(outcome.status, RunStatus::Succeeded);
assert_eq!(
agent_count.load(Ordering::SeqCst),
3,
"agent should be invoked exactly 3 times (maxIterations=3)"
);
}
#[tokio::test]
async fn resume_on_succeeded_run_returns_stored_outcome_without_reinvoking() {
let store = Arc::new(InMemoryCheckpointStore::default());
let agent_count = Arc::new(AtomicU32::new(0));
let tool_count = Arc::new(AtomicU32::new(0));
let exec = GraphExecutor::new(
store.clone(),
agent_fn_resolves_on(agent_count.clone(), 1),
tool_fn_counting(tool_count.clone()),
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
exec.start(&tenant(), "run-resume-done", &triage_cfg(), "hi")
.await
.expect("first run succeeds");
let after_start_agent = agent_count.load(Ordering::SeqCst);
let after_start_tool = tool_count.load(Ordering::SeqCst);
let outcome = exec
.resume(&tenant(), "run-resume-done")
.await
.expect("resume should succeed");
assert_eq!(outcome.status, RunStatus::Succeeded);
assert_eq!(
agent_count.load(Ordering::SeqCst),
after_start_agent,
"agent must NOT be called again on resume of terminal run"
);
assert_eq!(
tool_count.load(Ordering::SeqCst),
after_start_tool,
"tool must NOT be called again on resume of terminal run"
);
}
#[tokio::test]
async fn global_visit_cap_fails_run() {
let store = Arc::new(InMemoryCheckpointStore::default());
let mut v: serde_json::Value = serde_json::from_str(&triage_json()).expect("fixture JSON");
for node in v["nodes"].as_array_mut().expect("nodes array") {
if node["kind"] == "router" {
node["maxIterations"] = serde_json::json!(1000);
}
}
let cfg = GraphConfig::from_json(&v.to_string()).expect("patched graph valid");
let exec = GraphExecutor::new(
store.clone(),
agent_fn_resolves_on(Arc::new(AtomicU32::new(0)), u32::MAX),
tool_fn_counting(Arc::new(AtomicU32::new(0))),
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
let err = exec
.start(&tenant(), "run-global-cap", &cfg, "infinite loop")
.await
.expect_err("should fail with IterationCap");
assert!(
matches!(err, GraphExecError::IterationCap { .. }),
"expected IterationCap, got {err:?}"
);
let rec = store
.load(&tenant(), "run-global-cap")
.await
.expect("store accessible")
.expect("record must exist");
assert_eq!(
rec.status,
RunStatus::Failed,
"stored status must be Failed"
);
}
#[tokio::test]
async fn effect_error_leaves_run_resumable() {
let store = Arc::new(InMemoryCheckpointStore::default());
let t = tenant();
let phase1_count = Arc::new(AtomicU32::new(0));
{
let pc = phase1_count.clone();
let agent_phase1: AgentTurnFn = Arc::new(move |req: AgentTurnRequest| {
let n = pc.fetch_add(1, Ordering::SeqCst) + 1;
let node = req.node_id.clone();
Box::pin(async move {
if n == 1 {
Ok(AgentTurnResult {
reply: format!("pass-{n} from {node}"),
resolved: false,
})
} else {
Err(GraphExecError::AgentTurn(
"simulated failure on attempt 2".into(),
))
}
})
});
let tool_count = Arc::new(AtomicU32::new(0));
let store_ref = store.clone();
let exec = GraphExecutor::new(
store_ref,
agent_phase1,
tool_fn_counting(tool_count.clone()),
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
let err = exec
.start(&t, "run-resumable", &triage_cfg(), "retry me")
.await
.expect_err("should fail on agent attempt 2");
assert!(
matches!(err, GraphExecError::AgentTurn(_)),
"expected AgentTurn error, got {err:?}"
);
let rec = store
.load(&t, "run-resumable")
.await
.expect("store ok")
.expect("record exists");
assert_eq!(
rec.status,
RunStatus::Running,
"run should stay Running after effect error"
);
}
let phase2_count = Arc::new(AtomicU32::new(0));
let tool_phase2_count = Arc::new(AtomicU32::new(0));
{
let pc2 = phase2_count.clone();
let agent_phase2: AgentTurnFn = Arc::new(move |_req: AgentTurnRequest| {
pc2.fetch_add(1, Ordering::SeqCst);
Box::pin(async move {
Ok(AgentTurnResult {
reply: "resolved!".into(),
resolved: true,
})
})
});
let exec2 = GraphExecutor::new(
store.clone(),
agent_phase2,
tool_fn_counting(tool_phase2_count.clone()),
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
let outcome = exec2
.resume(&t, "run-resumable")
.await
.expect("resume should succeed");
assert_eq!(outcome.status, RunStatus::Succeeded, "outcome status");
}
assert_eq!(
phase2_count.load(Ordering::SeqCst),
1,
"phase2 agent must be called exactly once (attempt-1 was replayed)"
);
assert_eq!(
tool_phase2_count.load(Ordering::SeqCst),
1,
"phase2 tool must be called once (attempt-1 replayed, attempt-2 is fresh)"
);
}
#[tokio::test]
async fn start_twice_with_same_run_id_after_completion_errors() {
let store = Arc::new(InMemoryCheckpointStore::default());
let exec = GraphExecutor::new(
store.clone(),
agent_fn_resolves_on(Arc::new(AtomicU32::new(0)), 1),
tool_fn_counting(Arc::new(AtomicU32::new(0))),
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
exec.start(&tenant(), "run-dup", &triage_cfg(), "first")
.await
.expect("first start succeeds");
let err = exec
.start(&tenant(), "run-dup", &triage_cfg(), "second attempt")
.await
.expect_err("second start must fail");
assert!(
matches!(err, GraphExecError::AlreadyCompleted(_)),
"expected AlreadyCompleted, got {err:?}"
);
}
#[tokio::test]
async fn supervisor_routes_to_billing_branch() {
let store = Arc::new(InMemoryCheckpointStore::default());
let sup_count = Arc::new(AtomicU32::new(0));
let agent_count = Arc::new(AtomicU32::new(0));
let exec = GraphExecutor::new(
store.clone(),
agent_fn_resolves_on(agent_count.clone(), 1),
tool_fn_counting(Arc::new(AtomicU32::new(0))),
supervisor_fn_always_routes_to(sup_count.clone(), "billing"),
approval_fn_awaiting(),
);
let outcome = exec
.start(
&tenant(),
"run-sup-billing",
&supervisor_cfg(),
"I have a billing question",
)
.await
.expect("supervisor billing run should succeed");
assert_eq!(outcome.status, RunStatus::Succeeded, "status");
assert_eq!(
sup_count.load(Ordering::SeqCst),
1,
"supervisor invoked once"
);
assert_eq!(
agent_count.load(Ordering::SeqCst),
1,
"agent invoked once on billing branch"
);
let sup_entry = outcome.trail.iter().find(|e| e["kind"] == "supervisor");
assert!(
sup_entry.is_some(),
"trail must contain a supervisor entry: {:?}",
outcome.trail
);
let sup_entry = sup_entry.unwrap();
assert_eq!(
sup_entry["branch"], "billing",
"supervisor trail branch must be 'billing'"
);
assert_eq!(
sup_entry["replayed"], false,
"fresh run: supervisor not replayed"
);
}
#[tokio::test]
async fn supervisor_routes_to_tech_branch() {
let store = Arc::new(InMemoryCheckpointStore::default());
let sup_count = Arc::new(AtomicU32::new(0));
let agent_count = Arc::new(AtomicU32::new(0));
let exec = GraphExecutor::new(
store.clone(),
agent_fn_resolves_on(agent_count.clone(), 1),
tool_fn_counting(Arc::new(AtomicU32::new(0))),
supervisor_fn_always_routes_to(sup_count.clone(), "tech"),
approval_fn_awaiting(),
);
let outcome = exec
.start(
&tenant(),
"run-sup-tech",
&supervisor_cfg(),
"I have a tech issue",
)
.await
.expect("supervisor tech run should succeed");
assert_eq!(outcome.status, RunStatus::Succeeded, "status");
assert_eq!(
sup_count.load(Ordering::SeqCst),
1,
"supervisor invoked once"
);
assert_eq!(
agent_count.load(Ordering::SeqCst),
1,
"agent invoked once on tech branch"
);
let sup_entry = outcome.trail.iter().find(|e| e["kind"] == "supervisor");
assert!(sup_entry.is_some(), "trail must have supervisor entry");
assert_eq!(sup_entry.unwrap()["branch"], "tech");
}
#[tokio::test]
async fn supervisor_replay_determinism() {
let store = Arc::new(InMemoryCheckpointStore::default());
let sup_count = Arc::new(AtomicU32::new(0));
let agent_count = Arc::new(AtomicU32::new(0));
let exec = GraphExecutor::new(
store.clone(),
agent_fn_resolves_on(agent_count.clone(), 1),
tool_fn_counting(Arc::new(AtomicU32::new(0))),
supervisor_fn_always_routes_to(sup_count.clone(), "billing"),
approval_fn_awaiting(),
);
let first = exec
.start(
&tenant(),
"run-sup-replay",
&supervisor_cfg(),
"billing question",
)
.await
.expect("first run succeeds");
assert_eq!(first.status, RunStatus::Succeeded);
let after_first_sup = sup_count.load(Ordering::SeqCst);
let second = exec
.resume(&tenant(), "run-sup-replay")
.await
.expect("resume should succeed");
assert_eq!(second.status, RunStatus::Succeeded);
assert_eq!(
sup_count.load(Ordering::SeqCst),
after_first_sup,
"supervisor must NOT be called again on resume of a terminal run"
);
}
#[tokio::test]
async fn supervisor_trail_entry_has_branch_and_replayed_flag() {
let store = Arc::new(InMemoryCheckpointStore::default());
let sup_count = Arc::new(AtomicU32::new(0));
let exec = GraphExecutor::new(
store.clone(),
agent_fn_resolves_on(Arc::new(AtomicU32::new(0)), 1),
tool_fn_counting(Arc::new(AtomicU32::new(0))),
supervisor_fn_always_routes_to(sup_count.clone(), "tech"),
approval_fn_awaiting(),
);
let outcome = exec
.start(&tenant(), "run-sup-trail", &supervisor_cfg(), "need help")
.await
.expect("run should succeed");
let sup_entry = outcome
.trail
.iter()
.find(|e| e["kind"] == "supervisor")
.expect("trail must contain a supervisor entry");
assert_eq!(sup_entry["node"], "sup", "supervisor node id");
assert_eq!(sup_entry["kind"], "supervisor");
assert_eq!(sup_entry["attempt"], 1u32);
assert_eq!(sup_entry["replayed"], false);
assert_eq!(sup_entry["branch"], "tech");
}
#[tokio::test]
async fn agent_turn_request_carries_node_provider_when_set() {
use std::sync::Mutex;
let store = Arc::new(InMemoryCheckpointStore::default());
let captured: Arc<Mutex<Option<Option<String>>>> = Arc::new(Mutex::new(None));
let cap = captured.clone();
let agent: AgentTurnFn = Arc::new(move |req: AgentTurnRequest| {
*cap.lock().unwrap() = Some(req.provider.clone());
Box::pin(async move {
Ok(AgentTurnResult {
reply: "ok".into(),
resolved: true,
})
})
});
let cfg_json = serde_json::json!({
"schemaVersion": 1,
"entry": "agent",
"nodes": [
{
"id": "agent",
"kind": "agent",
"systemPrompt": "You help.",
"model": "claude-3-5-sonnet",
"provider": "anthropic"
},
{"id": "respond", "kind": "respond"}
],
"edges": [
{"from": "agent", "to": "respond"}
]
})
.to_string();
let cfg = GraphConfig::from_json(&cfg_json).expect("fixture valid");
let exec = GraphExecutor::new(
store.clone(),
agent,
Arc::new(|_| Box::pin(async { Ok(serde_json::json!({})) })),
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
exec.start(&tenant(), "run-provider-set", &cfg, "hi")
.await
.expect("run should succeed");
let got = captured.lock().unwrap().take().expect("agent was called");
assert_eq!(
got,
Some("anthropic".to_string()),
"AgentTurnRequest.provider must be Some(\"anthropic\") when set on the node"
);
}
#[tokio::test]
async fn agent_turn_request_provider_is_none_when_absent() {
use std::sync::Mutex;
let store = Arc::new(InMemoryCheckpointStore::default());
let captured: Arc<Mutex<Option<Option<String>>>> = Arc::new(Mutex::new(None));
let cap = captured.clone();
let agent: AgentTurnFn = Arc::new(move |req: AgentTurnRequest| {
*cap.lock().unwrap() = Some(req.provider.clone());
Box::pin(async move {
Ok(AgentTurnResult {
reply: "ok".into(),
resolved: true,
})
})
});
let exec = GraphExecutor::new(
store.clone(),
agent,
Arc::new(|_| Box::pin(async { Ok(serde_json::json!({})) })),
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
let cfg_json = serde_json::json!({
"schemaVersion": 1,
"entry": "agent",
"nodes": [
{
"id": "agent",
"kind": "agent",
"systemPrompt": "You help.",
"model": "gpt-4o-mini"
},
{"id": "respond", "kind": "respond"}
],
"edges": [{"from": "agent", "to": "respond"}]
})
.to_string();
let cfg = GraphConfig::from_json(&cfg_json).expect("fixture valid");
exec.start(&tenant(), "run-provider-absent", &cfg, "hi")
.await
.expect("run should succeed");
let got = captured.lock().unwrap().take().expect("agent was called");
assert_eq!(
got, None,
"AgentTurnRequest.provider must be None when the node has no provider field"
);
}
#[tokio::test]
async fn parallel_happy_path_merges_both_branches() {
let store = Arc::new(InMemoryCheckpointStore::default());
let agent: AgentTurnFn = Arc::new(|req: AgentTurnRequest| {
let node = req.node_id.clone();
Box::pin(async move {
Ok(AgentTurnResult {
reply: format!("agent-reply from {node}"),
resolved: true,
})
})
});
let tool: ToolFn = Arc::new(|_req: ToolCallRequest| {
Box::pin(async move { Ok(serde_json::json!({"branch_b": "done"})) })
});
let exec = GraphExecutor::new(
store.clone(),
agent,
tool,
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
let outcome = exec
.start(&tenant(), "run-par-happy", ¶llel_cfg(), "go")
.await
.expect("parallel run should succeed");
assert_eq!(outcome.status, RunStatus::Succeeded, "status");
let rec = store
.load(&tenant(), "run-par-happy")
.await
.unwrap()
.unwrap();
assert_eq!(rec.frontier_json, None, "frontier cleared after merge");
let state: GraphRunState = serde_json::from_str(&rec.state_json).unwrap();
let contents: Vec<&str> = state.messages.iter().map(|m| m.content.as_str()).collect();
let a_idx = contents
.iter()
.position(|c| c.contains("agent-reply from agent_a"))
.expect("branch a message present");
let b_idx = contents
.iter()
.position(|c| c.contains("branch_b"))
.expect("branch b message present");
assert!(
a_idx < b_idx,
"branch a must merge before branch b: {contents:?}"
);
}
#[tokio::test]
async fn parallel_branch_isolation_no_cross_bleed() {
let store = Arc::new(InMemoryCheckpointStore::default());
let agent: AgentTurnFn = Arc::new(|req: AgentTurnRequest| {
let node = req.node_id.clone();
Box::pin(async move {
let reply = if node == "agent_a" {
"SECRET-A-MESSAGE from agent_a".to_string()
} else {
format!("neutral reply from {node}")
};
Ok(AgentTurnResult {
reply,
resolved: true,
})
})
});
let saw_secret = Arc::new(std::sync::atomic::AtomicBool::new(false));
let saw_secret_probe = saw_secret.clone();
let tool: ToolFn = Arc::new(move |req: ToolCallRequest| {
let leaked = req
.state
.messages
.iter()
.any(|m| m.content.contains("SECRET-A-MESSAGE"));
if leaked {
saw_secret_probe.store(true, Ordering::SeqCst);
}
Box::pin(async move { Ok(serde_json::json!({"branch_b": "done"})) })
});
let exec = GraphExecutor::new(
store.clone(),
agent,
tool,
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
exec.start(&tenant(), "run-par-iso", ¶llel_cfg(), "go")
.await
.expect("run should succeed");
assert!(
!saw_secret.load(Ordering::SeqCst),
"branch B observed branch A's message — isolation violated"
);
}
#[tokio::test]
async fn parallel_merge_is_deterministic_under_delay() {
for _ in 0..3 {
let store = Arc::new(InMemoryCheckpointStore::default());
let agent: AgentTurnFn = Arc::new(|req: AgentTurnRequest| {
let node = req.node_id.clone();
Box::pin(async move {
if node == "agent_a" {
tokio::time::sleep(std::time::Duration::from_millis(40)).await;
}
Ok(AgentTurnResult {
reply: format!("reply from {node}"),
resolved: true,
})
})
});
let tool: ToolFn = Arc::new(|_req: ToolCallRequest| {
Box::pin(async move { Ok(serde_json::json!({"branch_b_fast": true})) })
});
let exec = GraphExecutor::new(
store.clone(),
agent,
tool,
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
exec.start(&tenant(), "run-par-det", ¶llel_cfg(), "go")
.await
.expect("run should succeed");
let rec = store.load(&tenant(), "run-par-det").await.unwrap().unwrap();
let state: GraphRunState = serde_json::from_str(&rec.state_json).unwrap();
let contents: Vec<&str> = state.messages.iter().map(|m| m.content.as_str()).collect();
let a_idx = contents
.iter()
.position(|c| c.contains("reply from agent_a"))
.expect("branch a present");
let b_idx = contents
.iter()
.position(|c| c.contains("branch_b_fast"))
.expect("branch b present");
assert!(
a_idx < b_idx,
"slow branch a must still merge before fast branch b: {contents:?}"
);
}
}
#[tokio::test]
async fn parallel_global_visit_cap_across_branches_fails() {
let v = serde_json::json!({
"schemaVersion": 2,
"entry": "fan",
"nodes": [
{"id": "fan", "kind": "parallel"},
{"id": "agent_a", "kind": "agent", "systemPrompt": "a", "model": "m", "tools": []},
{"id": "router_a", "kind": "router", "maxIterations": 1000},
{"id": "agent_b", "kind": "agent", "systemPrompt": "b", "model": "m", "tools": []},
{"id": "router_b", "kind": "router", "maxIterations": 1000},
{"id": "meet", "kind": "join"},
{"id": "respond", "kind": "respond"}
],
"edges": [
{"from": "fan", "to": "agent_a", "branch": "a"},
{"from": "fan", "to": "agent_b", "branch": "b"},
{"from": "agent_a", "to": "router_a"},
{"from": "router_a", "to": "agent_a", "branch": "loop"},
{"from": "router_a", "to": "meet", "branch": "resolved"},
{"from": "agent_b", "to": "router_b"},
{"from": "router_b", "to": "agent_b", "branch": "loop"},
{"from": "router_b", "to": "meet", "branch": "resolved"},
{"from": "meet", "to": "respond"}
]
});
let cfg = GraphConfig::from_json(&v.to_string()).expect("graph valid");
let store = Arc::new(InMemoryCheckpointStore::default());
let agent: AgentTurnFn = Arc::new(|req: AgentTurnRequest| {
let node = req.node_id.clone();
Box::pin(async move {
Ok(AgentTurnResult {
reply: format!("loop from {node}"),
resolved: false,
})
})
});
let exec = GraphExecutor::new(
store.clone(),
agent,
tool_fn_counting(Arc::new(AtomicU32::new(0))),
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
let err = exec
.start(&tenant(), "run-par-cap", &cfg, "loop forever")
.await
.expect_err("should hit global cap");
assert!(
matches!(err, GraphExecError::IterationCap { .. }),
"expected IterationCap, got {err:?}"
);
let rec = store.load(&tenant(), "run-par-cap").await.unwrap().unwrap();
assert_eq!(rec.status, RunStatus::Failed, "run must be Failed on cap");
}
#[tokio::test]
async fn parallel_mid_branch_error_keeps_run_running_with_frontier() {
let store = Arc::new(InMemoryCheckpointStore::default());
let agent: AgentTurnFn = Arc::new(|req: AgentTurnRequest| {
let node = req.node_id.clone();
Box::pin(async move {
if node == "agent_a" {
Err(GraphExecError::AgentTurn("branch a boom".into()))
} else {
Ok(AgentTurnResult {
reply: format!("ok from {node}"),
resolved: true,
})
}
})
});
let tool: ToolFn = Arc::new(|_req: ToolCallRequest| {
Box::pin(async move { Ok(serde_json::json!({"branch_b": "ok"})) })
});
let exec = GraphExecutor::new(
store.clone(),
agent,
tool,
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
let err = exec
.start(&tenant(), "run-par-err", ¶llel_cfg(), "go")
.await
.expect_err("branch a error should propagate");
assert!(
matches!(err, GraphExecError::AgentTurn(_)),
"expected AgentTurn error, got {err:?}"
);
let rec = store.load(&tenant(), "run-par-err").await.unwrap().unwrap();
assert_eq!(
rec.status,
RunStatus::Running,
"run must stay Running after a mid-branch error"
);
let frontier_json = rec
.frontier_json
.as_ref()
.expect("frontier must be persisted mid-parallel");
let frontier: Vec<BranchCursor> = serde_json::from_str(frontier_json).unwrap();
let a = frontier
.iter()
.find(|b| b.branch == "a")
.expect("branch a slot present");
assert_eq!(
a.cursor, "agent_a",
"failed branch cursor at last good node"
);
assert!(!a.parked, "failed branch must not be parked");
}
#[tokio::test]
async fn parallel_resume_after_branch_error_completes() {
let store = Arc::new(InMemoryCheckpointStore::default());
let t = tenant();
{
let agent: AgentTurnFn = Arc::new(|req: AgentTurnRequest| {
let node = req.node_id.clone();
Box::pin(async move {
if node == "agent_a" {
Err(GraphExecError::AgentTurn("boom".into()))
} else {
Ok(AgentTurnResult {
reply: format!("ok from {node}"),
resolved: true,
})
}
})
});
let tool: ToolFn =
Arc::new(|_r| Box::pin(async move { Ok(serde_json::json!({"branch_b": "ok"})) }));
let exec = GraphExecutor::new(
store.clone(),
agent,
tool,
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
exec.start(&t, "run-par-resume", ¶llel_cfg(), "go")
.await
.expect_err("phase 1 errors");
}
let branch_b_calls = Arc::new(AtomicU32::new(0));
{
let agent: AgentTurnFn = Arc::new(|req: AgentTurnRequest| {
let node = req.node_id.clone();
Box::pin(async move {
Ok(AgentTurnResult {
reply: format!("recovered from {node}"),
resolved: true,
})
})
});
let bcalls = branch_b_calls.clone();
let tool: ToolFn = Arc::new(move |_r| {
bcalls.fetch_add(1, Ordering::SeqCst);
Box::pin(async move { Ok(serde_json::json!({"branch_b": "ok"})) })
});
let exec = GraphExecutor::new(
store.clone(),
agent,
tool,
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
let outcome = exec
.resume(&t, "run-par-resume")
.await
.expect("resume should complete");
assert_eq!(outcome.status, RunStatus::Succeeded, "resumed run succeeds");
}
assert_eq!(
branch_b_calls.load(Ordering::SeqCst),
0,
"already-parked branch b must replay, not re-invoke its tool"
);
let rec = store.load(&t, "run-par-resume").await.unwrap().unwrap();
assert_eq!(rec.status, RunStatus::Succeeded);
assert_eq!(rec.frontier_json, None, "frontier cleared after merge");
let state: GraphRunState = serde_json::from_str(&rec.state_json).unwrap();
let contents: Vec<&str> = state.messages.iter().map(|m| m.content.as_str()).collect();
assert!(
contents
.iter()
.any(|c| c.contains("recovered from agent_a")),
"branch a recovered: {contents:?}"
);
assert!(
contents.iter().any(|c| c.contains("branch_b")),
"branch b present: {contents:?}"
);
}
fn approval_json() -> String {
serde_json::json!({
"schemaVersion": 2,
"entry": "start",
"nodes": [
{"id": "start", "kind": "agent", "systemPrompt": "greet the user", "model": "gpt-4o-mini"},
{"id": "approval", "kind": "approval", "title": "Approve refund?", "mode": "always"},
{"id": "respond", "kind": "respond"}
],
"edges": [
{"from": "start", "to": "approval"},
{"from": "approval", "to": "respond", "branch": "approved"}
]
})
.to_string()
}
fn approval_cfg() -> GraphConfig {
GraphConfig::from_json(&approval_json()).expect("approval fixture is valid")
}
#[tokio::test]
async fn approval_parks_then_resumes() {
let store = Arc::new(InMemoryCheckpointStore::default());
let t = tenant();
let agent_count = Arc::new(AtomicU32::new(0));
{
let exec = GraphExecutor::new(
store.clone(),
agent_fn_resolves_on(agent_count.clone(), 1),
tool_fn_counting(Arc::new(AtomicU32::new(0))),
supervisor_fn_unreachable(),
approval_fn_awaiting(),
);
let outcome = exec
.start(&t, "run-approval", &approval_cfg(), "please approve this")
.await
.expect("start should return cleanly when parked, not error");
assert_eq!(
outcome.status,
RunStatus::AwaitingInput,
"start() outcome must report AwaitingInput"
);
let rec = store
.load(&t, "run-approval")
.await
.expect("store accessible")
.expect("record must exist");
assert_eq!(
rec.status,
RunStatus::AwaitingInput,
"persisted record must be AwaitingInput"
);
assert_eq!(
rec.cursor, "approval",
"cursor must stay at the approval node while parked"
);
}
let decision_count = Arc::new(AtomicU32::new(0));
{
let exec = GraphExecutor::new(
store.clone(),
agent_fn_resolves_on(agent_count.clone(), 1),
tool_fn_counting(Arc::new(AtomicU32::new(0))),
supervisor_fn_unreachable(),
approval_fn_decides(decision_count.clone(), "approved"),
);
let outcome = exec
.resume(&t, "run-approval")
.await
.expect("resume with a decision should succeed");
assert_eq!(
outcome.status,
RunStatus::Succeeded,
"resumed run must reach Succeeded at respond"
);
}
assert_eq!(
decision_count.load(Ordering::SeqCst),
1,
"approval fn must be called exactly once on resume"
);
let rec = store
.load(&t, "run-approval")
.await
.unwrap()
.expect("record must exist after resume");
assert_eq!(rec.status, RunStatus::Succeeded);
assert_eq!(rec.cursor, "respond", "cursor must land on respond node");
let visit = store
.load_node_visit(&t, "run-approval", "approval", 1)
.await
.unwrap()
.expect("approval node visit must be recorded");
assert_eq!(visit["decision"], "approved");
}
}