Skip to main content

mcp_utils/testing/
fake_mcp.rs

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/// A fake MCP server preloaded with the classic math tools (`add_numbers`,
28/// `divide_numbers`, `slow_tool`); add scripted tools with [`Self::with_tool`].
29#[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    /// Compute the response from the request, for tools whose output depends
239    /// on their arguments. Scripted responses take precedence.
240    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}