1use crate::events::{TaskOutcomeState, TraceContext, task_created_result};
2use crate::mcp::tool_bridge::{convert_tool_result, map_task_result_to_outcome};
3use crate::mcp::{McpHandle, McpRuntime, ServerFactory, ToolCallStream, mcp};
4use futures::{FutureExt, StreamExt};
5use mcp_utils::client::{
6 CallToolOptions, CancellationToken, InMemoryServerSpec, McpConnectionDetails, McpServer, McpTransport,
7 ToolCallEvent, ToolExposure,
8};
9use mcp_utils::testing::ElicitationScript;
10use rmcp::model::{CreateTaskResult, ElicitResult, ProgressNotificationParam};
11use rmcp::{RoleServer, ServerHandler, service::DynService};
12use serde_json::Value;
13use std::collections::{HashMap, VecDeque};
14use std::sync::Mutex;
15use std::sync::atomic::{AtomicU64, Ordering};
16use std::time::Duration;
17use tokio::sync::watch;
18
19pub use mcp_utils::testing::CapturedElicitation;
20
21const DEFAULT_TOOL_TIMEOUT: Duration = Duration::from_secs(10);
22
23#[derive(Default)]
24pub struct McpTestBuilder {
25 servers: Vec<McpServer>,
26 factories: Vec<(String, ServerFactory)>,
27 elicitation_responses: Vec<ElicitResult>,
28 trace_context: Option<TraceContext>,
29 tool_timeout: Duration,
30}
31
32fn task_outcome(outcome: crate::events::TaskOutcome) -> TaskOutcome {
33 let (status, body) = match outcome.state {
34 TaskOutcomeState::Completed { result, .. } => ("completed", result.result),
35 TaskOutcomeState::Failed { error } => ("failed", error.error),
36 TaskOutcomeState::Cancelled => {
37 ("cancelled", "The background task was cancelled and will not produce a result.".into())
38 }
39 };
40 TaskOutcome { task_id: outcome.task_id, status: status.into(), body }
41}
42
43pub struct McpTest {
44 mcp: McpHandle,
45 _runtime: McpRuntime,
46 snapshot: McpConnectionDetails,
47 elicitations: ElicitationScript,
48 deferred_tools: tokio::sync::Mutex<VecDeque<DeferredTool>>,
49 cancel_tokens: Mutex<HashMap<String, CancellationToken>>,
50 trace_context: Option<TraceContext>,
51 tool_timeout: Duration,
52 next_call_id: AtomicU64,
53}
54
55pub struct TaskOutcome {
56 pub task_id: String,
57 pub status: String,
58 pub body: String,
59}
60
61pub struct ToolCallOutcome {
62 pub result: Result<llm::ToolCallResult, llm::ToolCallError>,
63 pub progress: Vec<ProgressNotificationParam>,
64 pub deferred_task: Option<CreateTaskResult>,
65}
66
67struct DeferredTool {
68 request: llm::ToolCallRequest,
69 events: ToolCallStream,
70}
71
72impl McpTestBuilder {
73 pub fn new() -> Self {
74 Self::default()
75 }
76
77 pub fn server<S>(self, name: impl Into<String>, server: S) -> Self
78 where
79 S: ServerHandler + Clone + Send + Sync + 'static,
80 {
81 self.server_with_exposure(name, server, ToolExposure::ModelVisible)
82 }
83
84 pub fn deferred_server<T>(self, name: impl Into<String>, server: T) -> Self
85 where
86 T: ServerHandler + Clone + Send + Sync + 'static,
87 {
88 self.server_with_exposure(name, server, ToolExposure::deferred_all())
89 }
90
91 pub fn server_with_exposure<T>(mut self, name: impl Into<String>, server: T, exposure: ToolExposure) -> Self
92 where
93 T: ServerHandler + Clone + Send + Sync + 'static,
94 {
95 let name = name.into();
96 let factory_name = format!("test-{}", self.factories.len());
97 let factory_server = server;
98 let factory: ServerFactory = Box::new(move |_spec, _services| {
99 let server = factory_server.clone();
100 async move { Box::new(server) as Box<dyn DynService<RoleServer>> }.boxed()
101 });
102 self.factories.push((factory_name.clone(), factory));
103 self.servers.push(McpServer::new(
104 name,
105 McpTransport::InMemory {
106 spec: InMemoryServerSpec { factory: factory_name, args: Vec::new(), input: None },
107 },
108 exposure,
109 ));
110 self
111 }
112
113 pub fn elicitation_response(mut self, response: ElicitResult) -> Self {
114 self.elicitation_responses.push(response);
115 self
116 }
117
118 pub fn trace_context(mut self, trace_context: TraceContext) -> Self {
119 self.trace_context = Some(trace_context);
120 self
121 }
122
123 pub fn tool_timeout(mut self, timeout: Duration) -> Self {
124 self.tool_timeout = timeout;
125 self
126 }
127
128 pub async fn build(self) -> McpTest {
129 let builder = mcp("/workspace").with_servers(self.servers);
130 let builder = self
131 .factories
132 .into_iter()
133 .fold(builder, |builder, (name, factory)| builder.register_in_memory_server(name, factory));
134 let mut spawn = builder.spawn().await.expect("MCP test manager spawns");
135 let snapshot = spawn.block_until_ready().await.expect("MCP test manager becomes ready");
136 let (runtime, event_rx) = spawn.split();
137
138 McpTest {
139 mcp: runtime.handle().clone(),
140 _runtime: runtime,
141 snapshot,
142 elicitations: ElicitationScript::spawn(event_rx, self.elicitation_responses),
143 deferred_tools: tokio::sync::Mutex::new(VecDeque::new()),
144 cancel_tokens: Mutex::new(HashMap::new()),
145 trace_context: self.trace_context,
146 tool_timeout: if self.tool_timeout.is_zero() { DEFAULT_TOOL_TIMEOUT } else { self.tool_timeout },
147 next_call_id: AtomicU64::new(1),
148 }
149 }
150}
151
152impl McpTest {
153 pub async fn call(&self, server: &str, tool: &str, arguments: Value) -> ToolCallOutcome {
154 let id = self.next_call_id.fetch_add(1, Ordering::Relaxed);
155 let request = llm::ToolCallRequest {
156 id: format!("mcp-test-{id}"),
157 name: format!("{server}__{tool}"),
158 arguments: arguments.to_string(),
159 };
160 let request_for_outcome = request.clone();
161 let cancel = CancellationToken::new();
162 self.cancel_tokens.lock().expect("cancel token lock").insert(request.id.clone(), cancel.clone());
163 let options = CallToolOptions {
164 timeout: self.tool_timeout,
165 meta: self.trace_context.as_ref().map(TraceContext::to_meta),
166 cancel,
167 };
168 let mut events = self.mcp.call_model_visible(request.name, &request.arguments, options);
169
170 let mut progress = Vec::new();
171 while let Some(event) = events.next().await {
172 match event {
173 ToolCallEvent::Progress(event) => progress.push(event),
174 ToolCallEvent::TaskCreated(task) => {
175 self.deferred_tools
176 .lock()
177 .await
178 .push_back(DeferredTool { request: request_for_outcome.clone(), events });
179 return ToolCallOutcome {
180 result: Ok(task_created_result(&request_for_outcome, &task.task.task_id)),
181 progress,
182 deferred_task: Some(task),
183 };
184 }
185 ToolCallEvent::Complete(outcome) => {
186 let result = convert_tool_result(&request_for_outcome, outcome).map(|(result, _)| result);
187 return ToolCallOutcome { result, progress, deferred_task: None };
188 }
189 ToolCallEvent::TaskStatus(_) | ToolCallEvent::TaskComplete { .. } | ToolCallEvent::Cancelled { .. } => {
190 panic!("MCP task lifecycle event arrived before deferral")
191 }
192 }
193 }
194 panic!("MCP test tool event stream ended before completion");
195 }
196
197 pub fn cancel_tool(&self, tool_id: &str) {
198 let tokens = self.cancel_tokens.lock().expect("cancel token lock");
199 tokens.get(tool_id).expect("cancel_tool targets a tool started with call()").cancel();
200 }
201
202 pub async fn next_tool_event(&self) -> Option<ToolCallEvent> {
203 self.next_deferred_event().await.map(|(_, event)| event)
204 }
205
206 pub async fn next_task_outcome(&self) -> Option<TaskOutcome> {
207 while let Some((request, event)) = self.next_deferred_event().await {
208 match event {
209 ToolCallEvent::TaskComplete { task, result } => {
210 return Some(task_outcome(map_task_result_to_outcome(request, task, result)));
211 }
212 ToolCallEvent::Cancelled { task_id } => {
213 return Some(task_outcome(crate::events::TaskOutcome {
214 request,
215 task_id: task_id.unwrap_or_else(|| "pending".to_string()),
216 state: TaskOutcomeState::Cancelled,
217 }));
218 }
219 ToolCallEvent::Progress(_)
220 | ToolCallEvent::TaskCreated(_)
221 | ToolCallEvent::TaskStatus(_)
222 | ToolCallEvent::Complete(_) => {}
223 }
224 }
225 None
226 }
227
228 async fn next_deferred_event(&self) -> Option<(llm::ToolCallRequest, ToolCallEvent)> {
229 let mut deferred = self.deferred_tools.lock().await;
230 loop {
231 let front = deferred.front_mut()?;
232 match front.events.next().await {
233 Some(event) => return Some((front.request.clone(), event)),
234 None => {
235 deferred.pop_front();
236 }
237 }
238 }
239 }
240
241 pub fn snapshot(&self) -> &McpConnectionDetails {
242 &self.snapshot
243 }
244
245 pub fn subscribe(&self) -> watch::Receiver<McpConnectionDetails> {
246 self.mcp.subscribe()
247 }
248
249 pub fn elicitations(&self) -> Vec<CapturedElicitation> {
250 self.elicitations.captured()
251 }
252}