Skip to main content

mcp_utils/
testing.rs

1use std::future::Future;
2
3use crate::protocol::client_lifecycle_mode;
4use crate::transport::create_in_memory_transport;
5use rmcp::{
6    RoleClient, RoleServer, Service, serve_client_with_lifecycle, serve_server,
7    service::{ClientInitializeError, RunningService, ServerInitializeError},
8};
9
10#[cfg(feature = "client")]
11pub use elicitation_script::{CapturedElicitation, ElicitationScript};
12#[cfg(all(feature = "client", any(test, feature = "testing")))]
13pub use fake_mcp::{
14    CapturedTaskUpdate, CapturedToolCall, FakeMcpServer, FakeMcpState, FakeTool, FakeToolResponse, 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) = create_in_memory_transport();
31
32        let (server_result, client_result) = tokio::join!(
33            serve_server(server, server_transport),
34            serve_client_with_lifecycle(client, client_transport, client_lifecycle_mode())
35        );
36
37        let server = server_result.map_err(|error| ConnectError::ServerInit(Box::new(error)))?;
38        let client = client_result.map_err(|error| ConnectError::ClientInit(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}
51
52#[cfg(feature = "client")]
53mod elicitation_script {
54    use crate::client::McpClientEvent;
55    use rmcp::model::{ElicitRequestParams, ElicitResult, ElicitationAction};
56    use std::collections::VecDeque;
57    use std::sync::{Arc, Mutex, PoisonError};
58    use tokio::sync::mpsc;
59    use tokio::task::JoinHandle;
60
61    /// Scripts the user's side of elicitation round trips: answers each
62    /// incoming request with the next queued response (Cancel once the queue
63    /// is empty) and records what arrived for assertions.
64    pub struct ElicitationScript {
65        captured: Arc<Mutex<Vec<CapturedElicitation>>>,
66        task: JoinHandle<()>,
67    }
68
69    #[derive(Clone)]
70    pub struct CapturedElicitation {
71        pub server_name: String,
72        pub request: ElicitRequestParams,
73    }
74
75    impl ElicitationScript {
76        pub fn spawn(
77            mut event_rx: mpsc::Receiver<McpClientEvent>,
78            responses: impl IntoIterator<Item = ElicitResult>,
79        ) -> Self {
80            let mut responses = responses.into_iter().collect::<VecDeque<_>>();
81            let captured = Arc::new(Mutex::new(Vec::new()));
82            let recorder = Arc::clone(&captured);
83            let task = tokio::spawn(async move {
84                while let Some(event) = event_rx.recv().await {
85                    if let McpClientEvent::Elicitation(event) = event {
86                        recorder
87                            .lock()
88                            .unwrap_or_else(PoisonError::into_inner)
89                            .push(CapturedElicitation { server_name: event.server_name, request: event.request });
90                        let response =
91                            responses.pop_front().unwrap_or_else(|| ElicitResult::new(ElicitationAction::Cancel));
92                        let _ = event.response_sender.send(response);
93                    }
94                }
95            });
96            Self { captured, task }
97        }
98
99        pub fn captured(&self) -> Vec<CapturedElicitation> {
100            self.captured.lock().unwrap_or_else(PoisonError::into_inner).clone()
101        }
102    }
103
104    impl Drop for ElicitationScript {
105        fn drop(&mut self) {
106            self.task.abort();
107        }
108    }
109}
110
111#[cfg(test)]
112mod tests {
113    use super::connect;
114    use rmcp::{
115        ClientHandler, ServerHandler,
116        model::{
117            ErrorData, Implementation, InitializeRequestParams, ProtocolVersion, ServerCapabilities, ServerConfig,
118        },
119        service::RequestContext,
120    };
121    use std::borrow::Cow;
122
123    #[tokio::test]
124    async fn connect_prefers_stateless_discovery_for_modern_servers() {
125        let (_server, client) = connect(McpServer728, TestClient).await.expect("connect");
126        assert_eq!(client.peer_info().expect("peer info").protocol_version, ProtocolVersion::V_2026_07_28);
127        client.list_tools(None).await.expect("list tools");
128        client.cancel().await.expect("cancel client");
129    }
130
131    #[tokio::test]
132    async fn connect_selects_an_older_mutually_supported_revision() {
133        let (_server, client) = connect(McpServer618, TestClient).await.expect("connect");
134
135        assert_eq!(client.peer_info().expect("peer info").protocol_version, ProtocolVersion::V_2025_06_18);
136        client.list_tools(None).await.expect("list tools");
137        client.cancel().await.expect("cancel client");
138    }
139
140    #[tokio::test]
141    async fn connect_falls_back_to_legacy_initialization() {
142        let (_server, client) = connect(McpServer1125, TestClient).await.expect("connect");
143
144        assert_eq!(client.peer_info().expect("peer info").protocol_version, ProtocolVersion::V_2025_11_25);
145        client.cancel().await.expect("cancel client");
146    }
147
148    #[derive(Clone, Default)]
149    struct TestClient;
150
151    impl ClientHandler for TestClient {}
152
153    #[derive(Clone, Default)]
154    struct McpServer728;
155
156    impl ServerHandler for McpServer728 {
157        fn get_info(&self) -> ServerConfig {
158            ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
159                .with_server_info(Implementation::new("modern-only", "1.0.0"))
160                .with_protocol_version(ProtocolVersion::V_2026_07_28)
161        }
162
163        fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
164            Cow::Owned(vec![ProtocolVersion::V_2026_07_28])
165        }
166
167        fn initialize(
168            &self,
169            _request: InitializeRequestParams,
170            _context: RequestContext<rmcp::RoleServer>,
171        ) -> impl std::future::Future<Output = Result<rmcp::model::InitializeResult, ErrorData>> + Send + '_ {
172            std::future::ready(Err(ErrorData::new(
173                rmcp::model::ErrorCode::METHOD_NOT_FOUND,
174                "initialize is not supported",
175                None,
176            )))
177        }
178    }
179
180    #[derive(Clone, Default)]
181    struct McpServer618;
182
183    impl ServerHandler for McpServer618 {
184        fn get_info(&self) -> ServerConfig {
185            ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
186                .with_server_info(Implementation::new("older-revision", "1.0.0"))
187                .with_protocol_version(ProtocolVersion::V_2025_06_18)
188        }
189
190        fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
191            Cow::Owned(vec![ProtocolVersion::V_2025_06_18])
192        }
193    }
194
195    #[derive(Clone, Default)]
196    struct McpServer1125;
197
198    impl ServerHandler for McpServer1125 {
199        fn get_info(&self) -> ServerConfig {
200            ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
201                .with_server_info(Implementation::new("legacy", "1.0.0"))
202                .with_protocol_version(ProtocolVersion::V_2025_11_25)
203        }
204
205        fn discover(
206            &self,
207            _context: RequestContext<rmcp::RoleServer>,
208        ) -> impl Future<Output = Result<rmcp::model::DiscoverResult, ErrorData>> + Send + '_ {
209            std::future::ready(Err(ErrorData::new(
210                rmcp::model::ErrorCode::METHOD_NOT_FOUND,
211                "server/discover is not supported",
212                None,
213            )))
214        }
215    }
216}