use crate::errors::GraphError;
use crate::graph::{GraphBuilder, END, START};
use crate::state::{AgentState, StateUpdate};
#[tokio::test]
async fn test_simple_linear_graph() {
let compiled = GraphBuilder::<AgentState>::new()
.add_node_fn("step1", |state| {
Ok(StateUpdate::full(AgentState::new(state.input.clone())))
})
.add_node_fn("step2", |state| {
let mut new_state = state.clone();
new_state.set_output("done".to_string());
Ok(StateUpdate::full(new_state))
})
.add_edge(START, "step1")
.add_edge("step1", "step2")
.add_edge("step2", END)
.compile()
.unwrap();
let input = AgentState::new("test input".to_string());
let result = compiled.invoke(input).await.unwrap();
assert!(result.final_state.output.is_some());
assert_eq!(result.recursion_count, 2);
}
#[tokio::test]
async fn test_stream_execution() {
let compiled = GraphBuilder::<AgentState>::new()
.add_node_fn("process", |state| Ok(StateUpdate::full(state.clone())))
.add_edge(START, "process")
.add_edge("process", END)
.compile()
.unwrap();
let input = AgentState::new("test".to_string());
let events = compiled.stream_collected(input).await.unwrap();
assert!(!events.is_empty());
}
fn chain_of(n: usize) -> crate::compiled::CompiledGraph<AgentState> {
let mut builder = GraphBuilder::<AgentState>::new();
for i in 1..=n {
let name = format!("n{}", i);
builder = builder.add_node_fn(name.clone(), |state| Ok(StateUpdate::full(state.clone())));
if i == 1 {
builder = builder.add_edge(START, name.clone());
}
if i == n {
builder = builder.add_edge(name, END);
}
}
for i in 1..n {
builder = builder.add_edge(format!("n{}", i), format!("n{}", i + 1));
}
builder.compile().unwrap()
}
#[tokio::test]
async fn test_recursion_limit_exact_fit_not_misreported() {
let compiled = chain_of(3).with_recursion_limit(3);
let result = compiled
.invoke(AgentState::new("x".to_string()))
.await
.unwrap();
assert_eq!(result.recursion_count, 3);
}
#[tokio::test]
async fn test_recursion_limit_exceeded_errors() {
let compiled = chain_of(3).with_recursion_limit(2);
let err = compiled
.invoke(AgentState::new("x".to_string()))
.await
.unwrap_err();
assert!(matches!(err, GraphError::RecursionLimitReached(2)));
}
#[tokio::test]
async fn test_invoke_from_node_enforces_recursion_limit() {
let compiled = chain_of(3).with_recursion_limit(1);
let err = compiled
.invoke_from_node("n1".to_string(), AgentState::new("x".to_string()))
.await
.unwrap_err();
assert!(matches!(err, GraphError::RecursionLimitReached(1)));
}
#[tokio::test]
async fn test_stream_reports_recursion_limit_hit() {
let compiled = chain_of(3).with_recursion_limit(1);
let events = compiled
.stream_collected(AgentState::new("x".to_string()))
.await;
let err = events.unwrap_err();
assert!(matches!(err, GraphError::RecursionLimitReached(1)));
}