#![cfg(feature = "scheduler")]
use std::sync::Mutex;
use zeph_llm::any::AnyProvider;
use zeph_llm::mock::MockProvider;
use zeph_llm::provider::{ChatResponse, ToolUseRequest};
use zeph_tools::executor::{ToolCall, ToolError, ToolExecutor, ToolOutput};
use crate::agent::Agent;
use crate::agent::agent_tests::{MockChannel, create_test_registry};
struct CallableToolExecutor {
outputs: Mutex<Vec<Result<Option<ToolOutput>, ToolError>>>,
}
impl CallableToolExecutor {
fn new(outputs: Vec<Result<Option<ToolOutput>, ToolError>>) -> Self {
Self {
outputs: Mutex::new(outputs),
}
}
fn fixed_output(summary: &str) -> Self {
Self::new(vec![Ok(Some(ToolOutput {
tool_name: "test_tool".into(),
summary: summary.to_owned(),
blocks_executed: 1,
filter_stats: None,
diff: None,
streamed: false,
terminal_id: None,
locations: None,
raw_response: None,
claim_source: None,
..Default::default()
}))])
}
fn failing() -> Self {
Self::new(vec![Err(ToolError::InvalidParams {
message: "tool failed".into(),
})])
}
}
impl ToolExecutor for CallableToolExecutor {
async fn execute(&self, _response: &str) -> Result<Option<ToolOutput>, ToolError> {
Ok(None)
}
async fn execute_tool_call(&self, _call: &ToolCall) -> Result<Option<ToolOutput>, ToolError> {
let mut outputs = self.outputs.lock().unwrap();
if outputs.is_empty() {
Ok(None)
} else {
outputs.remove(0)
}
}
zeph_tools::tool_executor_no_inner_defaults!();
}
fn tool_use_response(tool_id: &str, tool_name: &str) -> ChatResponse {
ChatResponse::ToolUse {
text: None,
tool_calls: vec![ToolUseRequest {
id: tool_id.to_owned(),
name: tool_name.into(),
input: serde_json::json!({"arg": "val"}),
}],
thinking_blocks: vec![],
}
}
#[tokio::test]
async fn text_only_response_returns_immediately() {
let (mock, _counter) =
MockProvider::default().with_tool_use(vec![ChatResponse::Text("the answer".into())]);
let provider = AnyProvider::Mock(mock);
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let executor = CallableToolExecutor::new(vec![]);
let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
let result = agent.run_inline_tool_loop("what is 2+2?", 10).await;
assert_eq!(result.unwrap().text, "the answer");
}
#[tokio::test]
async fn single_tool_iteration_returns_final_text() {
let (mock, counter) = MockProvider::default().with_tool_use(vec![
tool_use_response("call-1", "test_tool"),
ChatResponse::Text("done".into()),
]);
let provider = AnyProvider::Mock(mock);
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let executor = CallableToolExecutor::fixed_output("tool result");
let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
let result = agent.run_inline_tool_loop("run a tool", 10).await;
assert_eq!(result.unwrap().text, "done");
assert_eq!(*counter.lock().unwrap(), 2);
}
#[tokio::test]
async fn loop_terminates_at_max_iterations() {
let responses: Vec<ChatResponse> = (0..25)
.map(|i| tool_use_response(&format!("call-{i}"), "test_tool"))
.collect();
let (mock, counter) = MockProvider::default().with_tool_use(responses);
let provider = AnyProvider::Mock(mock);
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let executor = CallableToolExecutor::fixed_output("ok");
let max_iter = 5usize;
let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
let result = agent.run_inline_tool_loop("loop forever", max_iter).await;
assert!(result.is_ok());
assert_eq!(*counter.lock().unwrap(), u32::try_from(max_iter).unwrap());
}
#[tokio::test]
async fn tool_error_produces_is_error_result_and_loop_continues() {
let (mock, _counter) = MockProvider::default().with_tool_use(vec![
tool_use_response("call-err", "test_tool"),
ChatResponse::Text("recovered".into()),
]);
let provider = AnyProvider::Mock(mock);
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let executor = CallableToolExecutor::failing();
let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
let result = agent.run_inline_tool_loop("trigger error", 10).await;
assert_eq!(result.unwrap().text, "recovered");
}
#[tokio::test]
async fn multiple_tool_iterations_before_text() {
let (mock, counter) = MockProvider::default().with_tool_use(vec![
tool_use_response("call-1", "test_tool"),
tool_use_response("call-2", "test_tool"),
ChatResponse::Text("all done".into()),
]);
let provider = AnyProvider::Mock(mock);
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let executor = CallableToolExecutor::new(vec![
Ok(Some(ToolOutput {
tool_name: "test_tool".into(),
summary: "result-1".into(),
blocks_executed: 1,
filter_stats: None,
diff: None,
streamed: false,
terminal_id: None,
locations: None,
raw_response: None,
claim_source: None,
..Default::default()
})),
Ok(Some(ToolOutput {
tool_name: "test_tool".into(),
summary: "result-2".into(),
blocks_executed: 1,
filter_stats: None,
diff: None,
streamed: false,
terminal_id: None,
locations: None,
raw_response: None,
claim_source: None,
..Default::default()
})),
]);
let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
let result = agent
.run_inline_tool_loop("two tools then answer", 10)
.await
.unwrap();
assert_eq!(result.text, "all done");
assert_eq!(*counter.lock().unwrap(), 3);
assert_eq!(result.tool_trace.len(), 2);
assert!(result.tool_trace.iter().all(|t| t.tool == "test_tool"));
assert!(result.tool_trace.iter().all(|t| t.ok));
assert!(
result
.tool_trace
.iter()
.all(|t| t.args_summary.as_deref() == Some("val"))
);
}
#[tokio::test]
async fn provider_error_is_propagated() {
let provider = AnyProvider::Mock(zeph_llm::mock::MockProvider::failing());
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let executor = CallableToolExecutor::new(vec![]);
let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
let result = agent.run_inline_tool_loop("this will fail", 10).await;
assert!(result.is_err());
}
#[tokio::test]
async fn network_deny_wrapped_executor_blocks_fetch_before_reaching_inner() {
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
struct FlaggingExecutor {
called: Arc<AtomicBool>,
}
impl ToolExecutor for FlaggingExecutor {
async fn execute(&self, _response: &str) -> Result<Option<ToolOutput>, ToolError> {
Ok(None)
}
async fn execute_tool_call(
&self,
_call: &ToolCall,
) -> Result<Option<ToolOutput>, ToolError> {
self.called.store(true, Ordering::SeqCst);
Ok(Some(ToolOutput {
tool_name: "fetch".into(),
summary: "should not be reached".into(),
blocks_executed: 1,
filter_stats: None,
diff: None,
streamed: false,
terminal_id: None,
locations: None,
raw_response: None,
claim_source: None,
..Default::default()
}))
}
zeph_tools::tool_executor_no_inner_defaults!();
}
let (mock, _counter) = MockProvider::default().with_tool_use(vec![
tool_use_response("call-1", "fetch"),
ChatResponse::Text("done".into()),
]);
let provider = AnyProvider::Mock(mock);
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let called = Arc::new(AtomicBool::new(false));
let executor = FlaggingExecutor {
called: called.clone(),
};
let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
agent.tool_executor = Arc::new(zeph_subagent::NetworkDenyToolExecutor::new(
agent.tool_executor.clone(),
));
let result = agent.run_inline_tool_loop("fetch a url", 10).await;
assert_eq!(result.unwrap().text, "done");
assert!(
!called.load(Ordering::SeqCst),
"fetch tool call must be blocked before reaching the inner executor"
);
}
#[tokio::test]
async fn elicitation_event_during_tool_execution_is_handled() {
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{mpsc, oneshot};
use zeph_mcp::ElicitationEvent;
struct BlockingElicitingExecutor {
elic_tx: mpsc::Sender<ElicitationEvent>,
unblock_rx: Arc<std::sync::Mutex<Option<oneshot::Receiver<()>>>>,
sent: Arc<std::sync::atomic::AtomicBool>,
}
impl ToolExecutor for BlockingElicitingExecutor {
async fn execute(&self, _response: &str) -> Result<Option<ToolOutput>, ToolError> {
Ok(None)
}
async fn execute_tool_call(
&self,
_call: &ToolCall,
) -> Result<Option<ToolOutput>, ToolError> {
if !self.sent.swap(true, std::sync::atomic::Ordering::SeqCst) {
let (response_tx, _response_rx) = oneshot::channel();
let event = ElicitationEvent {
server_id: "test-server".to_owned(),
request: rmcp::model::ElicitRequestParams::FormElicitationParams {
meta: None,
message: "please fill in".to_owned(),
requested_schema: rmcp::model::ElicitationSchema::new(
std::collections::BTreeMap::new(),
),
},
response_tx,
};
let _ = self.elic_tx.send(event).await;
let rx = self.unblock_rx.lock().unwrap().take();
if let Some(rx) = rx {
let _ = rx.await;
}
}
Ok(Some(ToolOutput {
tool_name: "elicit_tool".into(),
summary: "result".into(),
blocks_executed: 1,
filter_stats: None,
diff: None,
streamed: false,
terminal_id: None,
locations: None,
raw_response: None,
claim_source: None,
..Default::default()
}))
}
zeph_tools::tool_executor_no_inner_defaults!();
}
let (elic_tx, elic_rx) = mpsc::channel::<ElicitationEvent>(4);
let (_unblock_tx, unblock_rx) = oneshot::channel::<()>();
let (mock, _counter) = MockProvider::default().with_tool_use(vec![
tool_use_response("call-elic", "elicit_tool"),
ChatResponse::Text("done".into()),
]);
let provider = AnyProvider::Mock(mock);
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let executor = BlockingElicitingExecutor {
elic_tx,
unblock_rx: Arc::new(std::sync::Mutex::new(Some(unblock_rx))),
sent: Arc::new(std::sync::atomic::AtomicBool::new(false)),
};
let mut agent =
Agent::new(provider, channel, registry, None, 5, executor).with_mcp_elicitation_rx(elic_rx);
let result = tokio::time::timeout(
Duration::from_secs(5),
agent.run_inline_tool_loop("trigger elicitation", 10),
)
.await
.expect("run_inline_tool_loop timed out — elicitation deadlock not fixed")
.unwrap();
assert_eq!(result.text, "done");
}
mod run_inline_timeout {
use std::time::Duration;
use zeph_orchestration::{
DagScheduler, GraphStatus, RuleBasedRouter, TaskGraph, TaskNode, TaskStatus, TimeoutPolicy,
};
use super::*;
struct SlowToolExecutor {
delay: Duration,
}
impl ToolExecutor for SlowToolExecutor {
async fn execute(&self, _response: &str) -> Result<Option<ToolOutput>, ToolError> {
Ok(None)
}
async fn execute_tool_call(
&self,
_call: &ToolCall,
) -> Result<Option<ToolOutput>, ToolError> {
tokio::time::sleep(self.delay).await;
Ok(Some(ToolOutput {
tool_name: "test_tool".into(),
summary: "slow result".into(),
blocks_executed: 1,
filter_stats: None,
diff: None,
streamed: false,
terminal_id: None,
locations: None,
raw_response: None,
claim_source: None,
..Default::default()
}))
}
zeph_tools::tool_executor_no_inner_defaults!();
}
#[tokio::test]
async fn short_override_fires_before_slow_tool_loop_completes() {
let (mock, _counter) = MockProvider::default().with_tool_use(vec![
tool_use_response("call-1", "test_tool"),
ChatResponse::Text("done".into()),
]);
let provider = AnyProvider::Mock(mock);
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let executor = SlowToolExecutor {
delay: Duration::from_secs(3),
};
let mut graph = TaskGraph::new("slow run-inline task");
let mut node = TaskNode::new(0, "slow task", "run something slow");
node.timeout = Some(TimeoutPolicy {
run_timeout_secs: Some(1),
idle_timeout_secs: None,
});
graph.tasks.push(node);
let config = zeph_config::OrchestrationConfig {
task_timeout_secs: 300, ..zeph_config::OrchestrationConfig::default()
};
let mut scheduler =
DagScheduler::new(graph, &config, Box::new(RuleBasedRouter), vec![], None).unwrap();
let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
agent.services.orchestration.orchestration_config = config;
let token = tokio_util::sync::CancellationToken::new();
let status = tokio::time::timeout(
Duration::from_secs(10),
agent.run_scheduler_loop(&mut scheduler, 1, token),
)
.await
.expect("run_scheduler_loop must not hang past the 1s override")
.unwrap();
assert_eq!(
status,
GraphStatus::Failed,
"timed-out RunInline task with default Abort strategy fails the graph"
);
assert_eq!(scheduler.graph().tasks[0].status, TaskStatus::Failed);
}
#[tokio::test]
async fn no_override_fast_completion_is_unaffected() {
let (mock, _counter) = MockProvider::default().with_tool_use(vec![
tool_use_response("call-1", "test_tool"),
ChatResponse::Text("done".into()),
]);
let provider = AnyProvider::Mock(mock);
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let executor = CallableToolExecutor::fixed_output("fast result");
let mut graph = TaskGraph::new("fast run-inline task");
let node = TaskNode::new(0, "fast task", "run something fast");
graph.tasks.push(node);
let config = zeph_config::OrchestrationConfig::default();
let mut scheduler =
DagScheduler::new(graph, &config, Box::new(RuleBasedRouter), vec![], None).unwrap();
let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
agent.services.orchestration.orchestration_config = config;
let token = tokio_util::sync::CancellationToken::new();
let status = agent
.run_scheduler_loop(&mut scheduler, 1, token)
.await
.unwrap();
assert_eq!(status, GraphStatus::Completed);
assert_eq!(scheduler.graph().tasks[0].status, TaskStatus::Completed);
}
#[tokio::test]
async fn no_override_task_is_capped_by_global_default_previously_unbounded() {
let (mock, _counter) = MockProvider::default().with_tool_use(vec![
tool_use_response("call-1", "test_tool"),
ChatResponse::Text("unused".into()),
]);
let provider = AnyProvider::Mock(mock);
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let executor = SlowToolExecutor {
delay: Duration::from_secs(3),
};
let mut graph = TaskGraph::new("no-override slow run-inline task");
let node = TaskNode::new(0, "slow task, no override", "run something slow");
graph.tasks.push(node);
let config = zeph_config::OrchestrationConfig {
task_timeout_secs: 1, ..zeph_config::OrchestrationConfig::default()
};
let mut scheduler =
DagScheduler::new(graph, &config, Box::new(RuleBasedRouter), vec![], None).unwrap();
let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
agent.services.orchestration.orchestration_config = config;
let token = tokio_util::sync::CancellationToken::new();
let status = tokio::time::timeout(
Duration::from_secs(10),
agent.run_scheduler_loop(&mut scheduler, 1, token),
)
.await
.expect("run_scheduler_loop must not hang past the 1s global default")
.unwrap();
assert_eq!(
status,
GraphStatus::Failed,
"a RunInline task with no override must now be capped by the global default \
(previously this dispatch path was entirely unbounded)"
);
assert_eq!(scheduler.graph().tasks[0].status, TaskStatus::Failed);
}
#[tokio::test]
async fn timeout_and_recovery_together_unblocks_dependent() {
let (mock, _counter) = MockProvider::default().with_tool_use(vec![
tool_use_response("call-1", "test_tool"),
ChatResponse::Text("dependent done".into()),
]);
let provider = AnyProvider::Mock(mock);
let channel = MockChannel::new(vec![]);
let registry = create_test_registry();
let executor = SlowToolExecutor {
delay: Duration::from_secs(3),
};
let mut graph = TaskGraph::new("timeout + recovery run-inline test");
let mut node0 = TaskNode::new(0, "slow recoverable task", "run something slow");
node0.timeout = Some(TimeoutPolicy {
run_timeout_secs: Some(1),
idle_timeout_secs: None,
});
node0.recovery = Some(zeph_orchestration::RecoveryAction {
state_injection: Some("recovered output".to_string()),
});
let mut node1 = TaskNode::new(1, "dependent task", "consume the recovered output");
node1.depends_on = vec![zeph_orchestration::TaskId(0)];
graph.tasks.push(node0);
graph.tasks.push(node1);
let config = zeph_config::OrchestrationConfig {
task_timeout_secs: 300,
..zeph_config::OrchestrationConfig::default()
};
let mut scheduler =
DagScheduler::new(graph, &config, Box::new(RuleBasedRouter), vec![], None).unwrap();
let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
agent.services.orchestration.orchestration_config = config;
let token = tokio_util::sync::CancellationToken::new();
let status = tokio::time::timeout(
Duration::from_secs(10),
agent.run_scheduler_loop(&mut scheduler, 2, token),
)
.await
.expect("run_scheduler_loop must not hang past the 1s override")
.unwrap();
assert_eq!(
status,
GraphStatus::Completed,
"recovery absorbs task 0's timeout; graph continues and completes via task 1"
);
assert_eq!(scheduler.graph().tasks[0].status, TaskStatus::Completed);
assert_eq!(
scheduler.graph().tasks[0]
.result
.as_ref()
.unwrap()
.agent_def
.as_deref(),
Some("__recovery__")
);
assert_eq!(scheduler.graph().tasks[1].status, TaskStatus::Completed);
}
}