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
22pub 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 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}