use std::collections::BTreeMap;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use serde_json::Value;
pub const SCHEMA_VERSION: u32 = 1;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct Graph {
pub schema_version: u32,
pub nodes: Vec<Node>,
#[serde(default)]
pub edges: Vec<Edge>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(
tag = "kind",
content = "payload",
rename_all = "snake_case",
deny_unknown_fields
)]
pub enum Node {
Agent(AgentNode),
Tool(ToolNode),
Gate(GateNode),
Branch(BranchNode),
Map(MapNode),
Fold(FoldNode),
}
impl Node {
#[must_use]
pub fn id(&self) -> &str {
match self {
Node::Agent(n) => &n.id,
Node::Tool(n) => &n.id,
Node::Gate(n) => &n.id,
Node::Branch(n) => &n.id,
Node::Map(n) => &n.id,
Node::Fold(n) => &n.id,
}
}
#[must_use]
pub fn kind_name(&self) -> &'static str {
match self {
Node::Agent(_) => "agent",
Node::Tool(_) => "tool",
Node::Gate(_) => "gate",
Node::Branch(_) => "branch",
Node::Map(_) => "map",
Node::Fold(_) => "fold",
}
}
#[must_use]
pub fn name(&self) -> Option<&str> {
match self {
Node::Agent(n) => n.name.as_deref(),
Node::Tool(n) => n.name.as_deref(),
Node::Gate(n) => n.name.as_deref(),
Node::Branch(n) => n.name.as_deref(),
Node::Map(n) => n.name.as_deref(),
Node::Fold(n) => n.name.as_deref(),
}
}
#[must_use]
pub fn input_schema(&self) -> Option<&Value> {
match self {
Node::Agent(n) => n.input_schema.as_ref(),
Node::Tool(n) => n.input_schema.as_ref(),
Node::Gate(_) | Node::Branch(_) | Node::Map(_) | Node::Fold(_) => None,
}
}
#[must_use]
pub fn output_schema(&self) -> Option<&Value> {
match self {
Node::Agent(n) => n.output_schema.as_ref(),
Node::Tool(n) => n.output_schema.as_ref(),
Node::Map(n) => n.output_schema.as_ref(),
Node::Gate(_) | Node::Branch(_) | Node::Fold(_) => None,
}
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct AgentNode {
pub id: String,
pub agent_hash: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_schema: Option<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_schema: Option<Value>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct ToolNode {
pub id: String,
pub tool: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub input: BTreeMap<String, String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_schema: Option<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_schema: Option<Value>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct GateNode {
pub id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prompt: Option<String>,
pub approval_schema: Value,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct BranchNode {
pub id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub on: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub agent_hash: Option<String>,
pub cases: Vec<BranchCase>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct BranchCase {
pub name: String,
pub when: BranchCondition,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(
tag = "kind",
content = "value",
rename_all = "snake_case",
deny_unknown_fields
)]
pub enum BranchCondition {
Expression(String),
ModelDecision,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct MapNode {
pub id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
pub over: String,
pub concurrency: u32,
pub body: MapBody,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_schema: Option<Value>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(
tag = "kind",
content = "value",
rename_all = "snake_case",
deny_unknown_fields
)]
pub enum MapBody {
Node(String),
Subgraph(Box<Graph>),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct FoldNode {
pub id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
pub body: FoldBody,
pub max_iterations: u32,
pub stop_when: String,
pub join: FoldJoin,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub accumulator_schema: Option<Value>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(
tag = "kind",
content = "value",
rename_all = "snake_case",
deny_unknown_fields
)]
pub enum FoldBody {
Node(String),
Subgraph(Box<Graph>),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(
tag = "kind",
content = "value",
rename_all = "snake_case",
deny_unknown_fields
)]
pub enum FoldJoin {
BestBy(String),
Last,
All,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct Edge {
pub from: String,
pub to: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub label: Option<String>,
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn sample() -> Graph {
Graph {
schema_version: SCHEMA_VERSION,
nodes: vec![
Node::Agent(AgentNode {
id: "research".into(),
agent_hash: format!("sha256:{}", "a".repeat(64)),
name: None,
input_schema: None,
output_schema: Some(json!({"type": "object"})),
}),
Node::Gate(GateNode {
id: "approve".into(),
name: None,
prompt: Some("Approve publication?".into()),
approval_schema: json!({"type": "object"}),
}),
],
edges: vec![Edge {
from: "research".into(),
to: "approve".into(),
label: None,
}],
}
}
#[test]
fn round_trips_through_json() {
let original = sample();
let json = serde_json::to_string(&original).expect("serialize");
let restored: Graph = serde_json::from_str(&json).expect("deserialize");
assert_eq!(original, restored, "round trip changed the value: {json}");
}
#[test]
fn node_uses_adjacent_kind_payload_shape() {
let node = Node::Tool(ToolNode {
id: "publish".into(),
tool: "http_post".into(),
name: None,
input: BTreeMap::new(),
input_schema: None,
output_schema: None,
});
let json = serde_json::to_string(&node).expect("serialize");
assert_eq!(
json,
r#"{"kind":"tool","payload":{"id":"publish","tool":"http_post"}}"#
);
}
#[test]
fn node_name_is_present_only_when_set() {
let named = Node::Tool(ToolNode {
id: "publish".into(),
tool: "http_post".into(),
name: Some("Publish the draft".into()),
input: BTreeMap::new(),
input_schema: None,
output_schema: None,
});
let json = serde_json::to_string(&named).expect("serialize");
assert_eq!(
json,
r#"{"kind":"tool","payload":{"id":"publish","tool":"http_post","name":"Publish the draft"}}"#
);
let unnamed = Node::Tool(ToolNode {
id: "publish".into(),
tool: "http_post".into(),
name: None,
input: BTreeMap::new(),
input_schema: None,
output_schema: None,
});
assert_eq!(
serde_json::to_string(&unnamed).expect("serialize"),
r#"{"kind":"tool","payload":{"id":"publish","tool":"http_post"}}"#,
"an unset name must not appear on the wire"
);
}
#[test]
fn fold_node_serializes_with_join_and_body_shapes() {
let node = Node::Fold(FoldNode {
id: "refine".into(),
name: None,
body: FoldBody::Node("tailor".into()),
max_iterations: 3,
stop_when: "score >= 0.85".into(),
join: FoldJoin::BestBy("score".into()),
accumulator_schema: None,
});
let json = serde_json::to_string(&node).expect("serialize");
assert_eq!(
json,
r#"{"kind":"fold","payload":{"id":"refine","body":{"kind":"node","value":"tailor"},"max_iterations":3,"stop_when":"score >= 0.85","join":{"kind":"best_by","value":"score"}}}"#
);
let restored: Node = serde_json::from_str(&json).expect("deserialize");
assert_eq!(node, restored, "fold round trip changed the value: {json}");
}
#[test]
fn fold_join_unit_variants_carry_only_the_kind_tag() {
assert_eq!(
serde_json::to_string(&FoldJoin::Last).expect("serialize"),
r#"{"kind":"last"}"#
);
assert_eq!(
serde_json::to_string(&FoldJoin::All).expect("serialize"),
r#"{"kind":"all"}"#
);
}
#[test]
fn unknown_document_field_is_rejected() {
let text = r#"{"schema_version":1,"nodes":[],"edges":[],"surprise":true}"#;
let error = serde_json::from_str::<Graph>(text).expect_err("must reject");
assert!(
error.to_string().contains("surprise"),
"error should name the stray field: {error}"
);
}
#[test]
fn unknown_payload_field_is_rejected() {
let text = r#"{"kind":"gate","payload":{"id":"g","approval_schema":{},"oops":1}}"#;
let error = serde_json::from_str::<Node>(text).expect_err("must reject");
assert!(
error.to_string().contains("oops"),
"error should name the stray field: {error}"
);
}
#[test]
fn unknown_node_envelope_field_is_rejected() {
let text = r#"{"kind":"gate","payload":{"id":"g","approval_schema":{}},"extra":1}"#;
let error = serde_json::from_str::<Node>(text).expect_err("must reject");
assert!(
error.to_string().contains("extra"),
"error should name the stray key: {error}"
);
}
}