use std::collections::BTreeMap;
use schemars::JsonSchema;
use serde::de::Error as _;
use serde::{Deserialize, Deserializer, 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, JsonSchema)]
#[serde(
tag = "kind",
content = "value",
rename_all = "snake_case",
deny_unknown_fields
)]
pub enum BranchCondition {
Expression(String),
ModelDecision,
}
#[derive(Deserialize)]
#[serde(
tag = "kind",
content = "value",
rename_all = "snake_case",
deny_unknown_fields
)]
enum BranchConditionShape {
Expression(String),
ModelDecision,
}
impl From<BranchConditionShape> for BranchCondition {
fn from(shape: BranchConditionShape) -> Self {
match shape {
BranchConditionShape::Expression(expr) => BranchCondition::Expression(expr),
BranchConditionShape::ModelDecision => BranchCondition::ModelDecision,
}
}
}
impl<'de> Deserialize<'de> for BranchCondition {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let value = Value::deserialize(deserializer)?;
serde_json::from_value::<BranchConditionShape>(value.clone())
.map(Into::into)
.map_err(|_| D::Error::custom(describe_branch_condition_error(&value)))
}
}
fn describe_branch_condition_error(value: &Value) -> String {
format!(
"a branch condition must be an object shaped \
`{{\"kind\": \"expression\", \"value\": \"<expr>\"}}` or \
`{{\"kind\": \"model_decision\"}}`; got {}",
describe_json_value(value)
)
}
fn describe_json_value(value: &Value) -> String {
let text = serde_json::to_string(value).unwrap_or_else(|_| "<unrepresentable>".to_string());
match value {
Value::Null => "null".to_string(),
Value::Bool(_) => format!("a bare boolean {text}"),
Value::Number(_) => format!("a bare number {text}"),
Value::String(_) => format!("a bare string {text}"),
Value::Array(_) => format!("a bare array {text}"),
Value::Object(_) => format!("an object {text}"),
}
}
#[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}"
);
}
#[test]
fn bare_string_branch_condition_names_accepted_shapes() {
let text = r#"{"name":"big","when":"value > 10000"}"#;
let error = serde_json::from_str::<BranchCase>(text).expect_err("must reject");
let message = error.to_string();
assert!(
!message.contains("adjacently tagged enum"),
"error should not leak serde internals: {message}"
);
assert!(
message.contains(r#"{"kind": "expression", "value": "<expr>"}"#)
&& message.contains(r#"{"kind": "model_decision"}"#),
"error should name both accepted shapes: {message}"
);
assert!(
message.contains(r#"a bare string "value > 10000""#),
"error should echo the offending value: {message}"
);
}
#[test]
fn non_string_branch_condition_also_names_accepted_shapes() {
let text = r#"{"name":"big","when":10000}"#;
let error = serde_json::from_str::<BranchCase>(text).expect_err("must reject");
let message = error.to_string();
assert!(
!message.contains("adjacently tagged enum"),
"error should not leak serde internals: {message}"
);
assert!(
message.contains("a bare number 10000"),
"error should describe the actual value: {message}"
);
}
#[test]
fn wrong_kind_branch_condition_names_accepted_shapes() {
let text = r#"{"name":"big","when":{"kind":"regex","value":"x"}}"#;
let error = serde_json::from_str::<BranchCase>(text).expect_err("must reject");
let message = error.to_string();
assert!(
!message.contains("adjacently tagged enum"),
"error should not leak serde internals: {message}"
);
assert!(
message.contains("an object"),
"error should describe the value as an object: {message}"
);
}
#[test]
fn valid_branch_conditions_still_round_trip() {
let expression = BranchCase {
name: "big".into(),
when: BranchCondition::Expression("value > 10000".into()),
};
let json = serde_json::to_string(&expression).expect("serialize");
assert_eq!(
json,
r#"{"name":"big","when":{"kind":"expression","value":"value > 10000"}}"#
);
let restored: BranchCase = serde_json::from_str(&json).expect("deserialize");
assert_eq!(expression, restored, "round trip changed the value: {json}");
let model_decision = BranchCase {
name: "review".into(),
when: BranchCondition::ModelDecision,
};
let json = serde_json::to_string(&model_decision).expect("serialize");
assert_eq!(
json,
r#"{"name":"review","when":{"kind":"model_decision"}}"#
);
let restored: BranchCase = serde_json::from_str(&json).expect("deserialize");
assert_eq!(
model_decision, restored,
"round trip changed the value: {json}"
);
}
#[test]
fn branch_node_document_serializes_byte_identical() {
let mut graph = sample();
graph.nodes.push(Node::Branch(BranchNode {
id: "route".into(),
name: None,
on: None,
agent_hash: None,
cases: vec![
BranchCase {
name: "big".into(),
when: BranchCondition::Expression("value > 10000".into()),
},
BranchCase {
name: "small".into(),
when: BranchCondition::ModelDecision,
},
],
}));
let json = serde_json::to_string(&graph).expect("serialize");
let restored: Graph = serde_json::from_str(&json).expect("deserialize");
assert_eq!(graph, restored, "round trip changed the value: {json}");
let json_again = serde_json::to_string(&restored).expect("re-serialize");
assert_eq!(
json, json_again,
"serialization is not byte-identical across a round trip"
);
}
}