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