|
5 | 5 | #![cfg(not(feature = "local"))] |
6 | 6 | #![cfg(feature = "client")] |
7 | 7 |
|
8 | | -use std::borrow::Cow; |
| 8 | +use std::{ |
| 9 | + borrow::Cow, |
| 10 | + sync::{ |
| 11 | + Arc, |
| 12 | + atomic::{AtomicUsize, Ordering}, |
| 13 | + }, |
| 14 | +}; |
9 | 15 |
|
10 | 16 | use rmcp::{ |
11 | 17 | ClientHandler, ErrorData, RoleServer, ServerHandler, ServiceExt, |
12 | 18 | model::{ |
13 | | - ClientInfo, ErrorCode, InitializeRequestParams, InitializeResult, ProtocolVersion, |
14 | | - ServerInfo, |
| 19 | + ClientCapabilities, ClientInfo, ErrorCode, Implementation, InitializeRequestParams, |
| 20 | + InitializeResult, ProtocolVersion, ServerInfo, |
15 | 21 | }, |
16 | 22 | service::{ClientInitializeError, RequestContext}, |
17 | 23 | }; |
@@ -232,3 +238,91 @@ async fn narrowed_server_caps_even_when_it_overrides_initialize() { |
232 | 238 | "the handshake layer should not raise the version above what the server supports" |
233 | 239 | ); |
234 | 240 | } |
| 241 | + |
| 242 | +/// Overrides `initialize` to run a side effect, then delegates the version |
| 243 | +/// answer back to the SDK with [`ServerHandler::negotiate_initialize`]. |
| 244 | +#[derive(Debug, Clone, Default)] |
| 245 | +struct DelegatingServer { |
| 246 | + initializations: Arc<AtomicUsize>, |
| 247 | +} |
| 248 | + |
| 249 | +impl ServerHandler for DelegatingServer { |
| 250 | + fn get_info(&self) -> ServerInfo { |
| 251 | + ServerInfo::default() |
| 252 | + } |
| 253 | + |
| 254 | + fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { |
| 255 | + Cow::Borrowed(HANDSHAKE_VERSIONS) |
| 256 | + } |
| 257 | + |
| 258 | + async fn initialize( |
| 259 | + &self, |
| 260 | + request: InitializeRequestParams, |
| 261 | + context: RequestContext<RoleServer>, |
| 262 | + ) -> Result<InitializeResult, ErrorData> { |
| 263 | + self.initializations.fetch_add(1, Ordering::Relaxed); |
| 264 | + context.peer.set_peer_info(request.clone()); |
| 265 | + self.negotiate_initialize(&request) |
| 266 | + } |
| 267 | +} |
| 268 | + |
| 269 | +fn initialize_params(protocol_version: ProtocolVersion) -> InitializeRequestParams { |
| 270 | + let mut params = InitializeRequestParams::new( |
| 271 | + ClientCapabilities::default(), |
| 272 | + Implementation::new("test-client", "0.0.0"), |
| 273 | + ); |
| 274 | + params.protocol_version = protocol_version; |
| 275 | + params |
| 276 | +} |
| 277 | + |
| 278 | +#[test] |
| 279 | +fn negotiate_initialize_echoes_a_supported_version() { |
| 280 | + let result = NarrowedServer |
| 281 | + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2025_06_18)) |
| 282 | + .expect("a supported handshake version should negotiate"); |
| 283 | + assert_eq!(result.protocol_version, ProtocolVersion::V_2025_06_18); |
| 284 | +} |
| 285 | + |
| 286 | +#[test] |
| 287 | +fn negotiate_initialize_caps_at_supported_versions() { |
| 288 | + let result = NarrowedServer |
| 289 | + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2026_07_28)) |
| 290 | + .expect("an unsupported version should fall back rather than fail"); |
| 291 | + assert_eq!(result.protocol_version, ProtocolVersion::V_2025_11_25); |
| 292 | +} |
| 293 | + |
| 294 | +#[test] |
| 295 | +fn negotiate_initialize_keeps_the_rest_of_get_info() { |
| 296 | + let server = NarrowedServer; |
| 297 | + let result = server |
| 298 | + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2026_07_28)) |
| 299 | + .expect("an unsupported version should fall back rather than fail"); |
| 300 | + assert_eq!(result.capabilities, server.get_info().capabilities); |
| 301 | +} |
| 302 | + |
| 303 | +#[test] |
| 304 | +fn negotiate_initialize_rejects_when_no_handshake_version_is_supported() { |
| 305 | + let error = ModernOnlyServer |
| 306 | + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2026_07_28)) |
| 307 | + .expect_err("a server with no handshake version cannot answer initialize"); |
| 308 | + assert_eq!(error.code, ErrorCode::UNSUPPORTED_PROTOCOL_VERSION); |
| 309 | +} |
| 310 | + |
| 311 | +#[tokio::test] |
| 312 | +async fn delegating_server_negotiates_like_the_default_initialize() { |
| 313 | + let negotiated = |
| 314 | + negotiated_version_with(DelegatingServer::default(), ProtocolVersion::V_2026_07_28).await; |
| 315 | + assert_eq!( |
| 316 | + negotiated, |
| 317 | + ProtocolVersion::V_2025_11_25, |
| 318 | + "an override that delegates should answer what the default initialize would" |
| 319 | + ); |
| 320 | +} |
| 321 | + |
| 322 | +#[tokio::test] |
| 323 | +async fn delegating_server_still_runs_its_own_side_effect() { |
| 324 | + let server = DelegatingServer::default(); |
| 325 | + let initializations = Arc::clone(&server.initializations); |
| 326 | + negotiated_version_with(server, ProtocolVersion::V_2025_06_18).await; |
| 327 | + assert_eq!(initializations.load(Ordering::Relaxed), 1); |
| 328 | +} |
0 commit comments