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 = 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 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}