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;
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    /// Filter which tools the servers expose and the gateway serves.
143    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    /// The Unix socket path of the deferred-tool gateway, if one was started.
272    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}