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
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) = tokio::io::duplex(64 * 1024);
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 futures::future::BoxFuture;
56 use rmcp::model::{ElicitRequestParams, ElicitResult, ElicitationAction};
57 use std::collections::VecDeque;
58 use std::future::Future;
59 use std::sync::{Arc, Mutex, PoisonError};
60 use tokio::sync::mpsc;
61 use tokio::task::JoinHandle;
62
63 pub type UrlElicitationHandler = Arc<dyn Fn(String, String) -> BoxFuture<'static, ()> + Send + Sync>;
64
65 pub struct ElicitationScript {
69 captured: Arc<Mutex<Vec<CapturedElicitation>>>,
70 task: JoinHandle<()>,
71 }
72
73 pub fn url_elicitation_handler<F, Fut>(handler: F) -> UrlElicitationHandler
74 where
75 F: Fn(String, String) -> Fut + Send + Sync + 'static,
76 Fut: Future<Output = ()> + Send + 'static,
77 {
78 Arc::new(move |url, id| Box::pin(handler(url, id)))
79 }
80
81 #[derive(Clone)]
82 pub struct CapturedElicitation {
83 pub server_name: String,
84 pub request: ElicitRequestParams,
85 }
86
87 impl ElicitationScript {
88 pub fn spawn(
89 event_rx: mpsc::Receiver<McpClientEvent>,
90 responses: impl IntoIterator<Item = ElicitResult>,
91 ) -> Self {
92 Self::spawn_with_url_handler(event_rx, responses, None)
93 }
94
95 pub fn spawn_with_url_handler(
96 mut event_rx: mpsc::Receiver<McpClientEvent>,
97 responses: impl IntoIterator<Item = ElicitResult>,
98 on_url: Option<UrlElicitationHandler>,
99 ) -> Self {
100 let mut responses = responses.into_iter().collect::<VecDeque<_>>();
101 let captured = Arc::new(Mutex::new(Vec::new()));
102 let recorder = Arc::clone(&captured);
103 let task = tokio::spawn(async move {
104 while let Some(event) = event_rx.recv().await {
105 let McpClientEvent::Elicitation(event) = event else { continue };
106 recorder
107 .lock()
108 .unwrap_or_else(PoisonError::into_inner)
109 .push(CapturedElicitation { server_name: event.server_name, request: event.request.clone() });
110 if let (Some(on_url), ElicitRequestParams::UrlElicitationParams { url, elicitation_id, .. }) =
111 (&on_url, event.request)
112 {
113 on_url(url, elicitation_id).await;
114 }
115 let response =
116 responses.pop_front().unwrap_or_else(|| ElicitResult::new(ElicitationAction::Cancel));
117 let _ = event.response_sender.send(response);
118 }
119 });
120 Self { captured, task }
121 }
122
123 pub fn captured(&self) -> Vec<CapturedElicitation> {
124 self.captured.lock().unwrap_or_else(PoisonError::into_inner).clone()
125 }
126 }
127
128 impl Drop for ElicitationScript {
129 fn drop(&mut self) {
130 self.task.abort();
131 }
132 }
133}
134
135#[cfg(test)]
136mod tests {
137 use super::connect;
138 use rmcp::{
139 ClientHandler, ServerHandler,
140 model::{
141 ErrorData, Implementation, InitializeRequestParams, ProtocolVersion, ServerCapabilities, ServerConfig,
142 },
143 service::RequestContext,
144 };
145 use std::borrow::Cow;
146
147 #[tokio::test]
148 async fn connect_prefers_stateless_discovery_for_modern_servers() {
149 let (_server, client) = connect(McpServer728, TestClient).await.expect("connect");
150 assert_eq!(client.peer_info().expect("peer info").protocol_version, ProtocolVersion::V_2026_07_28);
151 client.list_tools(None).await.expect("list tools");
152 client.cancel().await.expect("cancel client");
153 }
154
155 #[tokio::test]
156 async fn connect_selects_an_older_mutually_supported_revision() {
157 let (_server, client) = connect(McpServer618, TestClient).await.expect("connect");
158
159 assert_eq!(client.peer_info().expect("peer info").protocol_version, ProtocolVersion::V_2025_06_18);
160 client.list_tools(None).await.expect("list tools");
161 client.cancel().await.expect("cancel client");
162 }
163
164 #[tokio::test]
165 async fn connect_falls_back_to_legacy_initialization() {
166 let (_server, client) = connect(McpServer1125, TestClient).await.expect("connect");
167
168 assert_eq!(client.peer_info().expect("peer info").protocol_version, ProtocolVersion::V_2025_11_25);
169 client.cancel().await.expect("cancel client");
170 }
171
172 #[derive(Clone, Default)]
173 struct TestClient;
174
175 impl ClientHandler for TestClient {}
176
177 #[derive(Clone, Default)]
178 struct McpServer728;
179
180 impl ServerHandler for McpServer728 {
181 fn get_info(&self) -> ServerConfig {
182 ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
183 .with_server_info(Implementation::new("modern-only", "1.0.0"))
184 .with_protocol_version(ProtocolVersion::V_2026_07_28)
185 }
186
187 fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
188 Cow::Owned(vec![ProtocolVersion::V_2026_07_28])
189 }
190
191 fn initialize(
192 &self,
193 _request: InitializeRequestParams,
194 _context: RequestContext<rmcp::RoleServer>,
195 ) -> impl std::future::Future<Output = Result<rmcp::model::InitializeResult, ErrorData>> + Send + '_ {
196 std::future::ready(Err(ErrorData::new(
197 rmcp::model::ErrorCode::METHOD_NOT_FOUND,
198 "initialize is not supported",
199 None,
200 )))
201 }
202 }
203
204 #[derive(Clone, Default)]
205 struct McpServer618;
206
207 impl ServerHandler for McpServer618 {
208 fn get_info(&self) -> ServerConfig {
209 ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
210 .with_server_info(Implementation::new("older-revision", "1.0.0"))
211 .with_protocol_version(ProtocolVersion::V_2025_06_18)
212 }
213
214 fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
215 Cow::Owned(vec![ProtocolVersion::V_2025_06_18])
216 }
217 }
218
219 #[derive(Clone, Default)]
220 struct McpServer1125;
221
222 impl ServerHandler for McpServer1125 {
223 fn get_info(&self) -> ServerConfig {
224 ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
225 .with_server_info(Implementation::new("legacy", "1.0.0"))
226 .with_protocol_version(ProtocolVersion::V_2025_11_25)
227 }
228
229 fn discover(
230 &self,
231 _context: RequestContext<rmcp::RoleServer>,
232 ) -> impl Future<Output = Result<rmcp::model::DiscoverResult, ErrorData>> + Send + '_ {
233 std::future::ready(Err(ErrorData::new(
234 rmcp::model::ErrorCode::METHOD_NOT_FOUND,
235 "server/discover is not supported",
236 None,
237 )))
238 }
239 }
240}