1use crate::client::{RuntimeMcpServer, RuntimeMcpTransport, ToolExposure};
2use rmcp::{
3 ErrorData as McpError, Peer, RoleServer, ServerHandler,
4 model::{
5 CacheScope, CallToolRequestParams, CallToolResponse, CallToolResult, CancelTaskParams, ClientCapabilities,
6 ContentBlock, CreateTaskResult, DetailedTask, DiscoverResult, GetTaskParams, GetTaskResult, Implementation,
7 ListToolsResult, PaginatedRequestParams, ProgressNotificationParam, ProtocolVersion, ResultType,
8 ServerCapabilities, ServerConfig, TaskPayload, Tool, UpdateTaskParams,
9 },
10 service::{DynService, RequestContext},
11};
12use serde_json::json;
13use std::collections::{BTreeMap, HashMap, VecDeque};
14use std::future::Future;
15use std::sync::{Arc, Mutex};
16use std::time::Duration;
17
18pub fn fake_mcp(name: &str, server: FakeMcpServer) -> RuntimeMcpServer {
19 RuntimeMcpServer::new(name, RuntimeMcpTransport::InMemory { server: server.into_dyn() }, ToolExposure::ModelVisible)
20}
21
22pub fn completed_task_payload(result: CallToolResult) -> TaskPayload {
23 let result = serde_json::to_value(result).and_then(serde_json::from_value).expect("a tool result is a JSON object");
24 TaskPayload::Completed { result }
25}
26
27#[derive(Clone)]
30pub struct FakeMcpServer {
31 state: FakeMcpState,
32}
33
34#[derive(Clone, Default)]
35pub struct FakeMcpState {
36 inner: Arc<Mutex<FakeMcpStateInner>>,
37}
38
39#[derive(Clone)]
40pub struct CapturedToolCall {
41 pub request: CallToolRequestParams,
42 pub context_meta: serde_json::Map<String, serde_json::Value>,
43}
44
45#[derive(Clone)]
46pub struct CapturedTaskUpdate {
47 pub task_id: String,
48 pub input_responses: rmcp::model::InputResponses,
49}
50
51#[derive(Clone)]
52pub struct FakeTool {
53 definition: Tool,
54 responses: HashMap<Option<String>, FakeToolResponse>,
55 handler: Option<ToolHandler>,
56}
57
58#[derive(Clone)]
59pub struct FakeToolResponse {
60 response: CallToolResponse,
61 delay: Duration,
62 progress: Vec<FakeProgress>,
63 task_progress: Vec<(f64, Option<f64>)>,
64}
65
66#[derive(Clone)]
67struct FakeProgress {
68 progress: f64,
69 total: Option<f64>,
70 message: Option<String>,
71}
72
73impl FakeMcpServer {
74 pub fn new() -> Self {
75 Self::default()
76 }
77
78 pub fn with_tool(self, tool: FakeTool) -> Self {
79 self.state.add_tool(tool);
80 self
81 }
82
83 pub fn with_task(self, task_id: impl Into<String>, states: impl IntoIterator<Item = DetailedTask>) -> Self {
84 self.state.script_task(task_id, states);
85 self
86 }
87
88 pub fn with_task_get_failures(self, failures: usize) -> Self {
89 self.state.lock().task_get_failures = failures;
90 self
91 }
92
93 pub fn with_task_update_failures(self, failures: usize) -> Self {
94 self.state.lock().task_update_failures = failures;
95 self
96 }
97
98 pub fn state(&self) -> FakeMcpState {
99 self.state.clone()
100 }
101
102 pub fn into_dyn(self) -> Box<dyn DynService<RoleServer>> {
103 Box::new(self)
104 }
105}
106
107impl FakeMcpState {
108 pub fn calls_for(&self, tool: &str) -> Vec<CapturedToolCall> {
109 self.lock().calls.iter().filter(|call| call.request.name.as_ref() == tool).cloned().collect()
110 }
111
112 pub fn task_get_ids(&self) -> Vec<String> {
113 self.lock().task_get_ids.clone()
114 }
115
116 pub fn task_updates(&self) -> Vec<CapturedTaskUpdate> {
117 self.lock().task_updates.clone()
118 }
119
120 pub fn task_cancel_ids(&self) -> Vec<String> {
121 self.lock().task_cancel_ids.clone()
122 }
123
124 pub fn client_capabilities(&self) -> Option<ClientCapabilities> {
125 self.lock().client_capabilities.clone()
126 }
127
128 pub fn script_task(&self, task_id: impl Into<String>, states: impl IntoIterator<Item = DetailedTask>) {
129 self.lock().tasks.insert(task_id.into(), states.into_iter().collect());
130 }
131
132 fn task_for(&self, task_id: &str) -> Result<Option<DetailedTask>, ()> {
133 let mut inner = self.lock();
134 inner.task_get_ids.push(task_id.to_string());
135 if inner.task_get_failures > 0 {
136 inner.task_get_failures -= 1;
137 return Err(());
138 }
139 let Some(states) = inner.tasks.get_mut(task_id) else {
140 return Ok(None);
141 };
142 Ok(if states.len() > 1 { states.pop_front() } else { states.front().cloned() })
143 }
144
145 fn record_task_update(&self, request: UpdateTaskParams) -> bool {
146 let mut inner = self.lock();
147 inner
148 .task_updates
149 .push(CapturedTaskUpdate { task_id: request.task_id, input_responses: request.input_responses });
150 if inner.task_update_failures > 0 {
151 inner.task_update_failures -= 1;
152 false
153 } else {
154 true
155 }
156 }
157
158 fn record_task_cancel(&self, request: CancelTaskParams) {
159 self.lock().task_cancel_ids.push(request.task_id);
160 }
161
162 pub fn add_tool(&self, tool: FakeTool) {
163 self.lock().tools.insert(tool.definition.name.to_string(), tool);
164 }
165
166 pub async fn add_tool_and_notify(&self, tool: FakeTool) {
167 let peers = {
168 let mut inner = self.lock();
169 inner.tools.insert(tool.definition.name.to_string(), tool);
170 inner.peers.clone()
171 };
172 for peer in peers {
173 let _ = peer.notify_tool_list_changed().await;
174 }
175 }
176
177 pub async fn clear_tools_and_notify(&self) {
178 let peers = {
179 let mut inner = self.lock();
180 inner.tools.clear();
181 inner.peers.clone()
182 };
183 for peer in peers {
184 let _ = peer.notify_tool_list_changed().await;
185 }
186 }
187
188 pub fn fail_next_tool_list(&self) {
189 self.lock().tool_list_failures += 1;
190 }
191
192 fn definitions(&self) -> Vec<Tool> {
193 self.lock().tools.values().map(|tool| tool.definition.clone()).collect()
194 }
195
196 fn response_for(
197 &self,
198 request: &CallToolRequestParams,
199 context_meta: serde_json::Map<String, serde_json::Value>,
200 ) -> Option<FakeToolResponse> {
201 let mut inner = self.lock();
202 inner.calls.push(CapturedToolCall { request: request.clone(), context_meta });
203 inner.tools.get(request.name.as_ref()).and_then(|tool| tool.response_for(request))
204 }
205
206 fn lock(&self) -> std::sync::MutexGuard<'_, FakeMcpStateInner> {
207 self.inner.lock().unwrap_or_else(std::sync::PoisonError::into_inner)
208 }
209}
210
211impl FakeTool {
212 pub fn new(name: impl Into<String>) -> Self {
213 let name = name.into();
214 let schema = serde_json::from_value(json!({ "type": "object", "properties": {} }))
215 .expect("empty object schema is valid");
216 Self {
217 definition: Tool::new(name, "Fake MCP tool", Arc::new(schema)),
218 responses: HashMap::new(),
219 handler: None,
220 }
221 }
222
223 pub fn description(mut self, description: impl Into<String>) -> Self {
224 self.definition.description = Some(description.into().into());
225 self
226 }
227
228 pub fn responds(mut self, response: impl Into<FakeToolResponse>) -> Self {
229 self.responses.insert(None, response.into());
230 self
231 }
232
233 pub fn when_state(mut self, state: impl Into<String>, response: impl Into<FakeToolResponse>) -> Self {
234 self.responses.insert(Some(state.into()), response.into());
235 self
236 }
237
238 pub fn responds_with(
241 mut self,
242 handler: impl Fn(&CallToolRequestParams) -> FakeToolResponse + Send + Sync + 'static,
243 ) -> Self {
244 self.handler = Some(Arc::new(handler));
245 self
246 }
247
248 fn response_for(&self, request: &CallToolRequestParams) -> Option<FakeToolResponse> {
249 self.responses
250 .get(&request.request_state.as_deref().map(str::to_string))
251 .cloned()
252 .or_else(|| self.handler.as_ref().map(|handler| handler(request)))
253 }
254}
255
256impl FakeToolResponse {
257 pub fn new(response: impl Into<CallToolResponse>) -> Self {
258 Self { response: response.into(), delay: Duration::ZERO, progress: Vec::new(), task_progress: Vec::new() }
259 }
260
261 pub fn text(text: impl Into<String>) -> Self {
262 Self::new(CallToolResult::success(vec![ContentBlock::text(text.into())]))
263 }
264
265 pub fn task(task: CreateTaskResult) -> Self {
266 Self::new(CallToolResponse::Task(task))
267 }
268
269 pub fn delay(mut self, delay: Duration) -> Self {
270 self.delay = delay;
271 self
272 }
273
274 pub fn progress(mut self, progress: f64, total: Option<f64>) -> Self {
275 self.progress.push(FakeProgress { progress, total, message: None });
276 self
277 }
278
279 pub fn progress_message(mut self, progress: f64, message: impl Into<String>) -> Self {
280 self.progress.push(FakeProgress { progress, total: None, message: Some(message.into()) });
281 self
282 }
283
284 pub fn task_progress(mut self, progress: f64, total: Option<f64>) -> Self {
285 self.task_progress.push((progress, total));
286 self
287 }
288}
289
290impl<T> From<T> for FakeToolResponse
291where
292 T: Into<CallToolResponse>,
293{
294 fn from(response: T) -> Self {
295 Self::new(response)
296 }
297}
298
299impl Default for FakeMcpServer {
300 fn default() -> Self {
301 Self { state: FakeMcpState::default() }
302 .with_tool(add_numbers())
303 .with_tool(divide_numbers())
304 .with_tool(slow_tool())
305 }
306}
307
308impl ServerHandler for FakeMcpServer {
309 fn discover(
310 &self,
311 context: RequestContext<RoleServer>,
312 ) -> impl Future<Output = Result<DiscoverResult, McpError>> + Send + '_ {
313 self.state.lock().client_capabilities = context.meta.client_capabilities();
314 std::future::ready(Ok(DiscoverResult::from_server_info(
315 ServerHandler::supported_protocol_versions(self).into_owned(),
316 ServerHandler::get_info(self),
317 )))
318 }
319
320 fn get_info(&self) -> ServerConfig {
321 ServerConfig::new(ServerCapabilities::builder().enable_tools().enable_tasks().build())
322 .with_server_info(
323 Implementation::new("fake-mcp-server", "0.1.0").with_description("A fake MCP server for testing"),
324 )
325 .with_instructions("A fake MCP server for testing")
326 }
327
328 fn get_task(
329 &self,
330 request: GetTaskParams,
331 _context: RequestContext<RoleServer>,
332 ) -> impl Future<Output = Result<GetTaskResult, McpError>> + Send + '_ {
333 let result = match self.state.task_for(&request.task_id) {
334 Ok(Some(task)) => Ok(GetTaskResult::new(task)),
335 Ok(None) => Err(McpError::invalid_params(format!("unknown task: {}", request.task_id), None)),
336 Err(()) => Err(McpError::internal_error("scripted tasks/get failure", None)),
337 };
338 std::future::ready(result)
339 }
340
341 fn update_task(
342 &self,
343 request: UpdateTaskParams,
344 _context: RequestContext<RoleServer>,
345 ) -> impl Future<Output = Result<(), McpError>> + Send + '_ {
346 std::future::ready(
347 self.state
348 .record_task_update(request)
349 .then_some(())
350 .ok_or_else(|| McpError::internal_error("scripted tasks/update failure", None)),
351 )
352 }
353
354 fn cancel_task(
355 &self,
356 request: CancelTaskParams,
357 _context: RequestContext<RoleServer>,
358 ) -> impl Future<Output = Result<(), McpError>> + Send + '_ {
359 self.state.record_task_cancel(request);
360 std::future::ready(Ok(()))
361 }
362
363 fn list_tools(
364 &self,
365 _request: Option<PaginatedRequestParams>,
366 context: RequestContext<RoleServer>,
367 ) -> impl Future<Output = Result<ListToolsResult, McpError>> + Send + '_ {
368 let supports_cache_hints =
369 context.protocol_version().is_some_and(|version| version >= ProtocolVersion::V_2026_07_28);
370 let tools = {
371 let mut inner = self.state.lock();
372 if inner.tool_list_failures > 0 {
373 inner.tool_list_failures -= 1;
374 return std::future::ready(Err(McpError::internal_error("scripted tools/list failure", None)));
375 }
376 if inner.peers.is_empty() {
377 inner.peers.push(context.peer);
378 }
379 inner.tools.values().map(|tool| tool.definition.clone()).collect()
380 };
381 std::future::ready(Ok(ListToolsResult {
382 result_type: Some(ResultType::COMPLETE),
383 tools,
384 meta: None,
385 next_cursor: None,
386 ttl_ms: supports_cache_hints.then_some(0),
387 cache_scope: supports_cache_hints.then_some(CacheScope::Public),
388 }))
389 }
390
391 fn get_tool(&self, name: &str) -> Option<Tool> {
392 self.state.definitions().into_iter().find(|tool| tool.name == name)
393 }
394
395 async fn call_tool(
396 &self,
397 request: CallToolRequestParams,
398 context: RequestContext<RoleServer>,
399 ) -> Result<CallToolResponse, McpError> {
400 let response = self.state.response_for(&request, context.meta.0.0.clone());
401 let Some(response) = response else {
402 return Err(McpError::invalid_params(format!("unknown tool: {}", request.name), None));
403 };
404
405 if !response.delay.is_zero() {
406 tokio::time::sleep(response.delay).await;
407 }
408 if let Some(token) = context.meta.get_progress_token() {
409 for FakeProgress { progress, total, message } in response.progress {
410 let mut notification = ProgressNotificationParam::new(token.clone(), progress);
411 if let Some(total) = total {
412 notification = notification.with_total(total);
413 }
414 if let Some(message) = message {
415 notification = notification.with_message(message);
416 }
417 let _ = context.peer.notify_progress(notification).await;
418 }
419 if !response.task_progress.is_empty() {
420 let peer = context.peer.clone();
421 let token = token.clone();
422 tokio::spawn(async move {
423 tokio::task::yield_now().await;
424 for (progress, total) in response.task_progress {
425 let mut notification = ProgressNotificationParam::new(token.clone(), progress);
426 if let Some(total) = total {
427 notification = notification.with_total(total);
428 }
429 let _ = peer.notify_progress(notification).await;
430 }
431 });
432 }
433 }
434 Ok(response.response)
435 }
436}
437
438type ToolHandler = Arc<dyn Fn(&CallToolRequestParams) -> FakeToolResponse + Send + Sync>;
439
440#[derive(Default)]
441struct FakeMcpStateInner {
442 tools: BTreeMap<String, FakeTool>,
443 calls: Vec<CapturedToolCall>,
444 client_capabilities: Option<ClientCapabilities>,
445 tasks: HashMap<String, VecDeque<DetailedTask>>,
446 task_get_ids: Vec<String>,
447 task_updates: Vec<CapturedTaskUpdate>,
448 task_cancel_ids: Vec<String>,
449 task_get_failures: usize,
450 task_update_failures: usize,
451 tool_list_failures: usize,
452 peers: Vec<Peer<RoleServer>>,
453}
454
455fn add_numbers() -> FakeTool {
456 FakeTool::new("add_numbers").description("Adds two numbers together").responds_with(|request| {
457 let sum = int_arg(request, "a") + int_arg(request, "b");
458 FakeToolResponse::new(CallToolResult::structured(json!({ "sum": sum })))
459 })
460}
461
462fn divide_numbers() -> FakeTool {
463 FakeTool::new("divide_numbers").description("Divides two numbers").responds_with(|request| {
464 let (a, b) = (int_arg(request, "a"), int_arg(request, "b"));
465 if b == 0 {
466 return FakeToolResponse::new(CallToolResult::error(vec![ContentBlock::text("Division by zero")]));
467 }
468 FakeToolResponse::new(CallToolResult::structured(json!({ "quotient": a / b })))
469 })
470}
471
472fn slow_tool() -> FakeTool {
473 FakeTool::new("slow_tool")
474 .description("A tool that sleeps for a specified duration (for testing timeouts)")
475 .responds_with(|request| {
476 let sleep_ms = int_arg(request, "sleep_ms").unsigned_abs();
477 FakeToolResponse::new(CallToolResult::structured(json!({ "message": format!("Slept for {sleep_ms}ms") })))
478 .delay(Duration::from_millis(sleep_ms))
479 })
480}
481
482fn int_arg(request: &CallToolRequestParams, name: &str) -> i64 {
483 request.arguments.as_ref().and_then(|args| args.get(name)).and_then(serde_json::Value::as_i64).unwrap_or_default()
484}