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