Skip to main content

aether_core/testing/
mcp_test.rs

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