Skip to main content

mcp_utils/
testing.rs

1use std::future::Future;
2
3use crate::protocol::client_lifecycle_mode;
4use rmcp::{
5    RoleClient, RoleServer, Service, serve_client_with_lifecycle, serve_server,
6    service::{ClientInitializeError, RunningService, ServerInitializeError},
7};
8
9#[cfg(feature = "client")]
10pub use elicitation_script::{CapturedElicitation, ElicitationScript, UrlElicitationHandler, url_elicitation_handler};
11#[cfg(all(feature = "client", any(test, feature = "testing")))]
12pub use fake_mcp::{
13    CapturedTaskUpdate, CapturedToolCall, FakeMcpServer, FakeMcpState, FakeTool, FakeToolResponse,
14    completed_task_payload, fake_mcp,
15};
16
17#[cfg(all(feature = "client", any(test, feature = "testing")))]
18mod fake_mcp;
19
20pub type ConnectedServices<T, U> = (RunningService<RoleServer, T>, RunningService<RoleClient, U>);
21
22/// Helper function to connect an MCP server and client via in-memory transport
23/// This handles the dual-era discovery/initialization handshake by running both concurrently
24pub fn connect<T, U>(server: T, client: U) -> impl Future<Output = Result<ConnectedServices<T, U>, ConnectError>>
25where
26    T: Service<RoleServer>,
27    U: Service<RoleClient>,
28{
29    Box::pin(async move {
30        let (client_transport, server_transport) = tokio::io::duplex(64 * 1024);
31
32        let server = tokio::spawn(serve_server(server, server_transport));
33        let client = serve_client_with_lifecycle(client, client_transport, client_lifecycle_mode())
34            .await
35            .map_err(|error| ConnectError::ClientInit(Box::new(error)))?;
36
37        let _ = client.list_tools(None).await;
38        let server = server.await?.map_err(|error| ConnectError::ServerInit(Box::new(error)))?;
39
40        Ok((server, client))
41    })
42}
43
44#[derive(Debug, thiserror::Error)]
45pub enum ConnectError {
46    #[error("Server initialization failed: {0}")]
47    ServerInit(Box<ServerInitializeError>),
48    #[error("Client initialization failed: {0}")]
49    ClientInit(Box<ClientInitializeError>),
50    #[error("Server task failed: {0}")]
51    ServerTask(#[from] tokio::task::JoinError),
52}
53
54#[cfg(feature = "client")]
55mod elicitation_script {
56    use crate::client::McpClientEvent;
57    use futures::future::BoxFuture;
58    use rmcp::model::{ElicitRequestParams, ElicitResult, ElicitationAction};
59    use std::collections::VecDeque;
60    use std::future::Future;
61    use std::sync::{Arc, Mutex, PoisonError};
62    use tokio::sync::mpsc;
63    use tokio::task::JoinHandle;
64
65    pub type UrlElicitationHandler = Arc<dyn Fn(String, String) -> BoxFuture<'static, ()> + Send + Sync>;
66
67    /// Scripts the user's side of elicitation round trips: answers each
68    /// incoming request with the next queued response (Cancel once the queue
69    /// is empty) and records what arrived for assertions.
70    pub struct ElicitationScript {
71        captured: Arc<Mutex<Vec<CapturedElicitation>>>,
72        task: JoinHandle<()>,
73    }
74
75    pub fn url_elicitation_handler<F, Fut>(handler: F) -> UrlElicitationHandler
76    where
77        F: Fn(String, String) -> Fut + Send + Sync + 'static,
78        Fut: Future<Output = ()> + Send + 'static,
79    {
80        Arc::new(move |url, id| Box::pin(handler(url, id)))
81    }
82
83    #[derive(Clone)]
84    pub struct CapturedElicitation {
85        pub server_name: String,
86        pub request: ElicitRequestParams,
87    }
88
89    impl ElicitationScript {
90        pub fn spawn(
91            event_rx: mpsc::Receiver<McpClientEvent>,
92            responses: impl IntoIterator<Item = ElicitResult>,
93        ) -> Self {
94            Self::spawn_with_url_handler(event_rx, responses, None)
95        }
96
97        pub fn spawn_with_url_handler(
98            mut event_rx: mpsc::Receiver<McpClientEvent>,
99            responses: impl IntoIterator<Item = ElicitResult>,
100            on_url: Option<UrlElicitationHandler>,
101        ) -> Self {
102            let mut responses = responses.into_iter().collect::<VecDeque<_>>();
103            let captured = Arc::new(Mutex::new(Vec::new()));
104            let recorder = Arc::clone(&captured);
105            let task = tokio::spawn(async move {
106                while let Some(event) = event_rx.recv().await {
107                    let McpClientEvent::Elicitation(event) = event else { continue };
108                    recorder
109                        .lock()
110                        .unwrap_or_else(PoisonError::into_inner)
111                        .push(CapturedElicitation { server_name: event.server_name, request: event.request.clone() });
112                    if let (Some(on_url), ElicitRequestParams::UrlElicitationParams { url, elicitation_id, .. }) =
113                        (&on_url, event.request)
114                    {
115                        on_url(url, elicitation_id).await;
116                    }
117                    let response =
118                        responses.pop_front().unwrap_or_else(|| ElicitResult::new(ElicitationAction::Cancel));
119                    let _ = event.response_sender.send(response);
120                }
121            });
122            Self { captured, task }
123        }
124
125        pub fn captured(&self) -> Vec<CapturedElicitation> {
126            self.captured.lock().unwrap_or_else(PoisonError::into_inner).clone()
127        }
128    }
129
130    impl Drop for ElicitationScript {
131        fn drop(&mut self) {
132            self.task.abort();
133        }
134    }
135}
136
137#[cfg(test)]
138mod tests {
139    use super::connect;
140    use rmcp::{
141        ClientHandler, ServerHandler,
142        model::{
143            ErrorData, Implementation, InitializeRequestParams, ProtocolVersion, ServerCapabilities, ServerConfig,
144        },
145        service::RequestContext,
146    };
147    use std::borrow::Cow;
148
149    #[tokio::test]
150    async fn connect_prefers_stateless_discovery_for_modern_servers() {
151        let (_server, client) = connect(McpServer728, TestClient).await.expect("connect");
152        assert_eq!(client.peer_info().expect("peer info").protocol_version, ProtocolVersion::V_2026_07_28);
153        client.list_tools(None).await.expect("list tools");
154        client.cancel().await.expect("cancel client");
155    }
156
157    #[tokio::test]
158    async fn connect_selects_an_older_mutually_supported_revision() {
159        let (_server, client) = connect(McpServer618, TestClient).await.expect("connect");
160
161        assert_eq!(client.peer_info().expect("peer info").protocol_version, ProtocolVersion::V_2025_06_18);
162        client.list_tools(None).await.expect("list tools");
163        client.cancel().await.expect("cancel client");
164    }
165
166    #[tokio::test]
167    async fn connect_falls_back_to_legacy_initialization() {
168        let (_server, client) = connect(McpServer1125, TestClient).await.expect("connect");
169
170        assert_eq!(client.peer_info().expect("peer info").protocol_version, ProtocolVersion::V_2025_11_25);
171        client.cancel().await.expect("cancel client");
172    }
173
174    #[derive(Clone, Default)]
175    struct TestClient;
176
177    impl ClientHandler for TestClient {}
178
179    #[derive(Clone, Default)]
180    struct McpServer728;
181
182    impl ServerHandler for McpServer728 {
183        fn get_info(&self) -> ServerConfig {
184            ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
185                .with_server_info(Implementation::new("modern-only", "1.0.0"))
186                .with_protocol_version(ProtocolVersion::V_2026_07_28)
187        }
188
189        fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
190            Cow::Owned(vec![ProtocolVersion::V_2026_07_28])
191        }
192
193        fn initialize(
194            &self,
195            _request: InitializeRequestParams,
196            _context: RequestContext<rmcp::RoleServer>,
197        ) -> impl std::future::Future<Output = Result<rmcp::model::InitializeResult, ErrorData>> + Send + '_ {
198            std::future::ready(Err(ErrorData::new(
199                rmcp::model::ErrorCode::METHOD_NOT_FOUND,
200                "initialize is not supported",
201                None,
202            )))
203        }
204    }
205
206    #[derive(Clone, Default)]
207    struct McpServer618;
208
209    impl ServerHandler for McpServer618 {
210        fn get_info(&self) -> ServerConfig {
211            ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
212                .with_server_info(Implementation::new("older-revision", "1.0.0"))
213                .with_protocol_version(ProtocolVersion::V_2025_06_18)
214        }
215
216        fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
217            Cow::Owned(vec![ProtocolVersion::V_2025_06_18])
218        }
219    }
220
221    #[derive(Clone, Default)]
222    struct McpServer1125;
223
224    impl ServerHandler for McpServer1125 {
225        fn get_info(&self) -> ServerConfig {
226            ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
227                .with_server_info(Implementation::new("legacy", "1.0.0"))
228                .with_protocol_version(ProtocolVersion::V_2025_11_25)
229        }
230
231        fn discover(
232            &self,
233            _context: RequestContext<rmcp::RoleServer>,
234        ) -> impl Future<Output = Result<rmcp::model::DiscoverResult, ErrorData>> + Send + '_ {
235            std::future::ready(Err(ErrorData::new(
236                rmcp::model::ErrorCode::METHOD_NOT_FOUND,
237                "server/discover is not supported",
238                None,
239            )))
240        }
241    }
242}