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