#![warn(missing_docs)]
mod error;
pub mod fork;
mod walk;
use std::collections::{HashMap, HashSet};
use salvor_core::{Effect, RunId};
use salvor_graph::expr::Expr;
use salvor_graph::{BranchCondition, BranchNode, Edge, GateNode, Graph, MapBody, MapNode, Node};
use salvor_runtime::{
Agent, LoopOutcome, ParkReason, Resumption, RunCtx, ToolCallResult, drive_loop, hash_value,
};
use salvor_tools::DynTool;
use serde_json::{Value, json};
pub use error::EngineError;
pub use fork::{ForkError, ForkPlan, WriteHazard, plan_fork};
pub trait AgentResolver {
fn resolve_agent(&self, agent_hash: &str) -> Option<&Agent>;
}
pub trait ToolResolver {
fn resolve_tool(&self, name: &str) -> Option<&dyn DynTool>;
}
impl AgentResolver for HashMap<String, Agent> {
fn resolve_agent(&self, agent_hash: &str) -> Option<&Agent> {
self.get(agent_hash)
}
}
impl ToolResolver for HashMap<String, Box<dyn DynTool>> {
fn resolve_tool(&self, name: &str) -> Option<&dyn DynTool> {
self.get(name).map(AsRef::as_ref)
}
}
#[derive(Debug)]
pub enum GraphOutcome {
Completed {
output: Value,
},
Parked {
node: String,
reason: ParkReason,
},
}
pub fn graph_hash(graph: &Graph) -> Result<String, EngineError> {
let value = serde_json::to_value(graph).map_err(EngineError::GraphEncode)?;
Ok(hash_value(&value))
}
pub async fn run_graph(
ctx: &mut RunCtx,
graph: &Graph,
input: &Value,
agents: &impl AgentResolver,
tools: &impl ToolResolver,
) -> Result<GraphOutcome, EngineError> {
let hash = graph_hash(graph)?;
let graph_input = ctx.begin_graph(&hash, input).await?;
let by_id: HashMap<&str, &Node> = graph.nodes.iter().map(|n| (n.id(), n)).collect();
let mut inbound: HashMap<&str, Vec<&Edge>> = HashMap::new();
for edge in &graph.edges {
inbound.entry(edge.to.as_str()).or_default().push(edge);
}
let branches = parse_branches(graph)?;
let map_body_targets: HashSet<&str> = graph
.nodes
.iter()
.filter_map(|node| match node {
Node::Map(map) => match &map.body {
MapBody::Node(target) => Some(target.as_str()),
MapBody::Subgraph(_) => None,
},
_ => None,
})
.collect();
let mut outputs: HashMap<&str, Value> = HashMap::new();
let mut skipped: HashSet<&str> = HashSet::new();
let mut branch_case: HashMap<&str, String> = HashMap::new();
let mut last_output = graph_input.clone();
for node in walk::walk_order(graph)? {
let id = node.id();
if map_body_targets.contains(id) {
continue;
}
let Some(node_input) = select_input(
id,
&inbound,
&by_id,
&branch_case,
&skipped,
&outputs,
&graph_input,
) else {
ctx.node_skipped(id, SKIP_REASON).await?;
skipped.insert(id);
continue;
};
match node {
Node::Agent(agent_node) => {
let agent = agents
.resolve_agent(&agent_node.agent_hash)
.ok_or_else(|| EngineError::UnknownAgent {
node: agent_node.id.clone(),
agent_hash: agent_node.agent_hash.clone(),
})?;
ctx.node_entered(id).await?;
match drive_loop(ctx, agent, &node_input).await? {
LoopOutcome::Completed(output) => {
ctx.node_exited(id).await?;
last_output = output.clone();
outputs.insert(id, output);
}
LoopOutcome::Parked(reason) => {
return Ok(GraphOutcome::Parked {
node: agent_node.id.clone(),
reason,
});
}
}
}
Node::Tool(tool_node) => {
let tool = tools.resolve_tool(&tool_node.tool).ok_or_else(|| {
EngineError::UnknownTool {
node: tool_node.id.clone(),
tool: tool_node.tool.clone(),
}
})?;
ctx.node_entered(id).await?;
let idempotency_key = match tool.effect() {
Effect::Idempotent => Some(fork_safe_idempotency_key(&hash, id, 0)),
Effect::Read | Effect::Write => None,
};
match ctx
.tool_call(tool, &node_input, idempotency_key.as_deref())
.await?
{
ToolCallResult::Output(output) => {
ctx.node_exited(id).await?;
last_output = output.clone();
outputs.insert(id, output);
}
ToolCallResult::Failed(failure) => {
return Err(EngineError::ToolFailed {
node: tool_node.id.clone(),
message: failure.message,
});
}
ToolCallResult::Suspended(suspension) => {
ctx.suspend(&suspension.reason, &suspension.input_schema)
.await?;
match ctx.await_resume().await? {
Resumption::Parked => {
return Ok(GraphOutcome::Parked {
node: tool_node.id.clone(),
reason: ParkReason::Suspended {
reason: suspension.reason,
input_schema: suspension.input_schema,
},
});
}
Resumption::Resumed(resume_input) => {
ctx.node_exited(id).await?;
last_output = resume_input.clone();
outputs.insert(id, resume_input);
}
}
}
}
}
Node::Gate(gate) => {
ctx.node_entered(id).await?;
let reason = gate_reason(gate);
ctx.suspend(&reason, &gate.approval_schema).await?;
match ctx.await_resume().await? {
Resumption::Parked => {
return Ok(GraphOutcome::Parked {
node: gate.id.clone(),
reason: ParkReason::Suspended {
reason,
input_schema: gate.approval_schema.clone(),
},
});
}
Resumption::Resumed(resume_input) => {
ctx.node_exited(id).await?;
last_output = resume_input.clone();
outputs.insert(id, resume_input);
}
}
}
Node::Branch(branch) => {
let cases = branches.get(id).expect("every branch node is parsed");
let chosen: String = match &branch.agent_hash {
None => {
let case = choose_expression_case(id, cases, &node_input)?;
ctx.node_entered(id).await?;
case.to_owned()
}
Some(agent_hash) => {
let agent = agents.resolve_agent(agent_hash).ok_or_else(|| {
EngineError::UnknownAgent {
node: branch.id.clone(),
agent_hash: agent_hash.clone(),
}
})?;
ctx.node_entered(id).await?;
let reply = match drive_loop(ctx, agent, &node_input).await? {
LoopOutcome::Completed(output) => output,
LoopOutcome::Parked(reason) => {
return Ok(GraphOutcome::Parked {
node: branch.id.clone(),
reason,
});
}
};
match_decision(branch, &reply)?.to_owned()
}
};
ctx.branch_taken(id, &chosen).await?;
ctx.node_exited(id).await?;
branch_case.insert(id, chosen);
last_output = node_input.clone();
outputs.insert(id, node_input);
}
Node::Map(map_node) => {
match drive_map(ctx, map_node, &node_input, &by_id, agents, tools, &hash).await? {
MapOutcome::Joined(output) => {
last_output = output.clone();
outputs.insert(id, output);
}
MapOutcome::Parked { node, reason } => {
return Ok(GraphOutcome::Parked { node, reason });
}
}
}
Node::Fold(fold) => {
return Err(EngineError::UnsupportedNode {
node: fold.id.clone(),
kind: "fold",
});
}
}
}
ctx.complete_run(&last_output).await?;
Ok(GraphOutcome::Completed {
output: last_output,
})
}
enum MapOutcome {
Joined(Value),
Parked {
node: String,
reason: ParkReason,
},
}
async fn drive_map(
ctx: &mut RunCtx,
map_node: &MapNode,
routed: &Value,
by_id: &HashMap<&str, &Node>,
agents: &impl AgentResolver,
tools: &impl ToolResolver,
graph_hash: &str,
) -> Result<MapOutcome, EngineError> {
let node_id = map_node.id.as_str();
let body: &Node = match &map_node.body {
MapBody::Node(target) => {
let body_node =
by_id
.get(target.as_str())
.copied()
.ok_or_else(|| EngineError::MalformedGraph {
detail: format!("map node `{node_id}`: body names unknown node `{target}`"),
})?;
match body_node {
Node::Agent(_) | Node::Tool(_) => body_node,
other => {
return Err(EngineError::UnsupportedMapBody {
node: node_id.to_owned(),
detail: format!(
"a `{}` body node cannot be a per-item worker; only `agent` and `tool` bodies run",
other.kind_name()
),
});
}
}
}
MapBody::Subgraph(_) => {
return Err(EngineError::UnsupportedMapBody {
node: node_id.to_owned(),
detail: "an embedded `subgraph` body is not executed yet".to_owned(),
});
}
};
let items = resolve_over(node_id, &map_node.over, routed)?;
ctx.node_entered(node_id).await?;
ctx.map_fanned_out(node_id, &Value::Array(items.clone()))
.await?;
let mut joined: Vec<Value> = Vec::with_capacity(items.len());
for (position, item) in items.iter().enumerate() {
let index = position as u64;
let child_run = map_child_run_id(ctx.run_id(), node_id, index);
ctx.map_iteration_started(node_id, index, &child_run)
.await?;
let call = MapCall {
graph_hash,
node_id,
index,
};
match run_map_body(ctx, body, item, agents, tools, call).await? {
IterationOutcome::Output(output) => joined.push(output),
IterationOutcome::Parked(reason) => {
return Ok(MapOutcome::Parked {
node: node_id.to_owned(),
reason,
});
}
}
ctx.map_iteration_joined(node_id, index).await?;
}
ctx.node_exited(node_id).await?;
Ok(MapOutcome::Joined(Value::Array(joined)))
}
enum IterationOutcome {
Output(Value),
Parked(ParkReason),
}
struct MapCall<'a> {
graph_hash: &'a str,
node_id: &'a str,
index: u64,
}
async fn run_map_body(
ctx: &mut RunCtx,
body: &Node,
item: &Value,
agents: &impl AgentResolver,
tools: &impl ToolResolver,
call: MapCall<'_>,
) -> Result<IterationOutcome, EngineError> {
match body {
Node::Agent(agent_node) => {
let agent = agents
.resolve_agent(&agent_node.agent_hash)
.ok_or_else(|| EngineError::UnknownAgent {
node: agent_node.id.clone(),
agent_hash: agent_node.agent_hash.clone(),
})?;
match drive_loop(ctx, agent, item).await? {
LoopOutcome::Completed(output) => Ok(IterationOutcome::Output(output)),
LoopOutcome::Parked(reason) => Ok(IterationOutcome::Parked(reason)),
}
}
Node::Tool(tool_node) => {
let tool =
tools
.resolve_tool(&tool_node.tool)
.ok_or_else(|| EngineError::UnknownTool {
node: tool_node.id.clone(),
tool: tool_node.tool.clone(),
})?;
let idempotency_key = match tool.effect() {
Effect::Idempotent => Some(fork_safe_idempotency_key(
call.graph_hash,
call.node_id,
call.index,
)),
Effect::Read | Effect::Write => None,
};
match ctx
.tool_call(tool, item, idempotency_key.as_deref())
.await?
{
ToolCallResult::Output(output) => Ok(IterationOutcome::Output(output)),
ToolCallResult::Failed(failure) => Err(EngineError::ToolFailed {
node: tool_node.id.clone(),
message: failure.message,
}),
ToolCallResult::Suspended(suspension) => {
ctx.suspend(&suspension.reason, &suspension.input_schema)
.await?;
match ctx.await_resume().await? {
Resumption::Parked => Ok(IterationOutcome::Parked(ParkReason::Suspended {
reason: suspension.reason,
input_schema: suspension.input_schema,
})),
Resumption::Resumed(resume_input) => {
Ok(IterationOutcome::Output(resume_input))
}
}
}
}
}
other => Err(EngineError::UnsupportedMapBody {
node: call.node_id.to_owned(),
detail: format!(
"a `{}` body node cannot be a per-item worker",
other.kind_name()
),
}),
}
}
fn map_child_run_id(parent_run: RunId, node_id: &str, index: u64) -> String {
hash_value(&json!({
"parent_run": parent_run,
"node": node_id,
"index": index,
}))
}
fn resolve_over(node_id: &str, over: &str, routed: &Value) -> Result<Vec<Value>, EngineError> {
let reference =
salvor_graph::expr::parse_reference(over).map_err(|error| EngineError::MalformedGraph {
detail: format!(
"map node `{node_id}`: `over` reference `{over}` is unparseable: {error}"
),
})?;
match reference.resolve(routed) {
Some(Value::Array(items)) => Ok(items.clone()),
_ => Err(EngineError::MapOverNotAList {
node: node_id.to_owned(),
over: over.to_owned(),
}),
}
}
fn fork_safe_idempotency_key(graph_hash: &str, node_id: &str, call_index: u64) -> String {
hash_value(&serde_json::json!({
"graph_hash": graph_hash,
"node": node_id,
"call": call_index,
}))
}
const SKIP_REASON: &str = "no live inbound edge: an upstream branch routed to another case";
type ParsedBranches<'a> = HashMap<&'a str, Vec<(&'a str, Option<Expr>)>>;
fn parse_branches(graph: &Graph) -> Result<ParsedBranches<'_>, EngineError> {
let mut parsed = HashMap::new();
for node in &graph.nodes {
let Node::Branch(branch) = node else {
continue;
};
let mut cases = Vec::with_capacity(branch.cases.len());
for case in &branch.cases {
let expr = match &case.when {
BranchCondition::Expression(source) => {
Some(salvor_graph::expr::parse(source).map_err(|error| {
EngineError::MalformedGraph {
detail: format!(
"branch node `{}`: case `{}` has an unparseable condition: {error}",
branch.id, case.name
),
}
})?)
}
BranchCondition::ModelDecision => None,
};
cases.push((case.name.as_str(), expr));
}
parsed.insert(branch.id.as_str(), cases);
}
Ok(parsed)
}
fn select_input(
id: &str,
inbound: &HashMap<&str, Vec<&Edge>>,
by_id: &HashMap<&str, &Node>,
branch_case: &HashMap<&str, String>,
skipped: &HashSet<&str>,
outputs: &HashMap<&str, Value>,
graph_input: &Value,
) -> Option<Value> {
let edges = inbound.get(id).map(Vec::as_slice).unwrap_or_default();
if edges.is_empty() {
return Some(graph_input.clone());
}
let mut chosen: Option<&Edge> = None;
for edge in edges {
if !is_live_inbound(edge, by_id, branch_case, skipped) {
continue;
}
chosen = match chosen {
Some(best) if best.from <= edge.from => Some(best),
_ => Some(edge),
};
}
chosen.map(|edge| {
outputs
.get(edge.from.as_str())
.cloned()
.unwrap_or(Value::Null)
})
}
fn is_live_inbound(
edge: &Edge,
by_id: &HashMap<&str, &Node>,
branch_case: &HashMap<&str, String>,
skipped: &HashSet<&str>,
) -> bool {
if skipped.contains(edge.from.as_str()) {
return false;
}
match by_id.get(edge.from.as_str()) {
Some(Node::Branch(_)) => {
branch_case.get(edge.from.as_str()).map(String::as_str) == edge.label.as_deref()
}
_ => true,
}
}
fn gate_reason(gate: &GateNode) -> String {
gate.prompt
.clone()
.unwrap_or_else(|| format!("approval required at gate `{}`", gate.id))
}
fn choose_expression_case<'a>(
node_id: &str,
cases: &'a [(&'a str, Option<Expr>)],
value: &Value,
) -> Result<&'a str, EngineError> {
for (name, expr) in cases {
match expr {
Some(expr) if expr.eval(value) => return Ok(name),
Some(_) => {}
None => {
return Err(EngineError::MalformedGraph {
detail: format!(
"branch node `{node_id}`: an expression branch must not carry a model-decision case"
),
});
}
}
}
Err(EngineError::NoBranchCaseMatched {
node: node_id.to_owned(),
})
}
fn match_decision<'a>(branch: &'a BranchNode, reply: &Value) -> Result<&'a str, EngineError> {
let reply_text = reply
.as_str()
.map_or_else(|| reply.to_string(), |text| text.trim().to_owned());
for case in &branch.cases {
if case.name == reply_text {
return Ok(case.name.as_str());
}
}
Err(EngineError::BranchDecisionUnmatched {
node: branch.id.clone(),
reply: reply_text,
cases: branch.cases.iter().map(|case| case.name.clone()).collect(),
})
}