Skip to content

Commit ca4564a

Browse files
committed
feat: add ServerHandler::negotiate_initialize
1 parent 3023198 commit ca4564a

5 files changed

Lines changed: 246 additions & 13 deletions

File tree

‎crates/rmcp/src/handler/server.rs‎

Lines changed: 60 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -321,16 +321,59 @@ macro_rules! server_handler_methods {
321321
context: RequestContext<RoleServer>,
322322
) -> impl Future<Output = Result<InitializeResult, McpError>> + MaybeSendFuture + '_ {
323323
context.peer.set_peer_info(request.clone());
324+
std::future::ready(self.negotiate_initialize(&request))
325+
}
326+
/// Build the `initialize` response for `request`, negotiating the
327+
/// protocol version against [`Self::supported_protocol_versions`].
328+
///
329+
/// This is the whole body of the default [`Self::initialize`] minus its
330+
/// `set_peer_info` side effect, so a server that overrides `initialize`
331+
/// to add its own can call this instead of restating the negotiation
332+
/// rule:
333+
///
334+
/// ```
335+
/// use rmcp::{
336+
/// ErrorData as McpError, RoleServer, ServerHandler,
337+
/// model::{InitializeRequestParams, InitializeResult, ServerInfo},
338+
/// service::RequestContext,
339+
/// };
340+
///
341+
/// struct MyServer;
342+
///
343+
/// impl ServerHandler for MyServer {
344+
/// fn get_info(&self) -> ServerInfo {
345+
/// ServerInfo::default()
346+
/// }
347+
///
348+
/// async fn initialize(
349+
/// &self,
350+
/// request: InitializeRequestParams,
351+
/// context: RequestContext<RoleServer>,
352+
/// ) -> Result<InitializeResult, McpError> {
353+
/// // ... record telemetry, register the peer, etc.
354+
/// context.peer.set_peer_info(request.clone());
355+
/// self.negotiate_initialize(&request)
356+
/// }
357+
/// }
358+
/// ```
359+
///
360+
/// # Errors
361+
///
362+
/// Returns [`ErrorCode::UNSUPPORTED_PROTOCOL_VERSION`] when this server
363+
/// supports no version that still has an `initialize` handshake.
364+
///
365+
/// [`ErrorCode::UNSUPPORTED_PROTOCOL_VERSION`]: crate::model::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION
366+
fn negotiate_initialize(
367+
&self,
368+
request: &InitializeRequestParams,
369+
) -> Result<InitializeResult, McpError> {
324370
let mut info = self.get_info();
325-
let negotiated = negotiate_protocol_version(
371+
info.protocol_version = negotiate_protocol_version(
326372
&request.protocol_version,
327373
std::mem::take(&mut info.protocol_version),
328374
&self.supported_protocol_versions(),
329-
);
330-
std::future::ready(negotiated.map(|version| {
331-
info.protocol_version = version;
332-
info
333-
}))
375+
)?;
376+
Ok(info)
334377
}
335378
/// Return the protocol versions supported by this server.
336379
///
@@ -339,6 +382,10 @@ macro_rules! server_handler_methods {
339382
/// list is advertised by [`Self::discover`], bounds what `initialize`
340383
/// negotiation may agree to, and is what per-request versions are
341384
/// validated against.
385+
///
386+
/// To support everything up to some ceiling, use
387+
/// [`ProtocolVersion::known_up_to`] rather than filtering
388+
/// [`ProtocolVersion::KNOWN_VERSIONS`] by hand.
342389
fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
343390
Cow::Borrowed(ProtocolVersion::KNOWN_VERSIONS)
344391
}
@@ -621,6 +668,13 @@ macro_rules! impl_server_handler_for_wrapper {
621668
(**self).initialize(request, context)
622669
}
623670

671+
fn negotiate_initialize(
672+
&self,
673+
request: &InitializeRequestParams,
674+
) -> Result<InitializeResult, McpError> {
675+
(**self).negotiate_initialize(request)
676+
}
677+
624678
fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
625679
(**self).supported_protocol_versions()
626680
}

‎crates/rmcp/src/model.rs‎

Lines changed: 83 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -177,7 +177,7 @@ impl ProtocolVersion {
177177
/// First protocol version that requires SEP-2243 standard HTTP headers.
178178
pub const STANDARD_HEADERS: Self = Self::V_2026_07_28;
179179

180-
/// All protocol versions known to this SDK.
180+
/// All protocol versions known to this SDK, oldest first.
181181
pub const KNOWN_VERSIONS: &[Self] = &[
182182
Self::V_2024_11_05,
183183
Self::V_2025_03_26,
@@ -190,6 +190,43 @@ impl ProtocolVersion {
190190
pub fn as_str(&self) -> &str {
191191
&self.0
192192
}
193+
194+
/// The known versions up to and including `max`, oldest first.
195+
///
196+
/// Servers that implement every revision up to some ceiling can return
197+
/// this from `supported_protocol_versions` instead of filtering
198+
/// [`Self::KNOWN_VERSIONS`] by hand. `max` itself need not be a known
199+
/// version; the result is empty when it predates all of them.
200+
///
201+
/// The result borrows from [`Self::KNOWN_VERSIONS`], so call it directly
202+
/// in the method body — it needs no `static` and no `LazyLock`:
203+
///
204+
/// ```rust,ignore
205+
/// const MAX_SUPPORTED: ProtocolVersion = ProtocolVersion::V_2025_11_25;
206+
///
207+
/// fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> {
208+
/// Cow::Borrowed(ProtocolVersion::known_up_to(&MAX_SUPPORTED))
209+
/// }
210+
/// ```
211+
///
212+
/// ```
213+
/// # use rmcp::model::ProtocolVersion;
214+
/// assert_eq!(
215+
/// ProtocolVersion::known_up_to(&ProtocolVersion::V_2025_06_18),
216+
/// &[
217+
/// ProtocolVersion::V_2024_11_05,
218+
/// ProtocolVersion::V_2025_03_26,
219+
/// ProtocolVersion::V_2025_06_18,
220+
/// ],
221+
/// );
222+
/// ```
223+
pub fn known_up_to(max: &Self) -> &'static [Self] {
224+
let count = Self::KNOWN_VERSIONS
225+
.iter()
226+
.take_while(|version| version.as_str() <= max.as_str())
227+
.count();
228+
&Self::KNOWN_VERSIONS[..count]
229+
}
193230
}
194231

195232
impl Serialize for ProtocolVersion {
@@ -4643,6 +4680,51 @@ mod tests {
46434680

46444681
use super::*;
46454682

4683+
#[test]
4684+
fn known_versions_are_ordered_oldest_first() {
4685+
// `known_up_to` walks the list as a sorted prefix.
4686+
assert!(
4687+
ProtocolVersion::KNOWN_VERSIONS
4688+
.windows(2)
4689+
.all(|pair| pair[0].as_str() < pair[1].as_str())
4690+
);
4691+
}
4692+
4693+
#[test]
4694+
fn known_up_to_includes_the_ceiling_itself() {
4695+
assert_eq!(
4696+
ProtocolVersion::known_up_to(&ProtocolVersion::V_2024_11_05),
4697+
&[ProtocolVersion::V_2024_11_05]
4698+
);
4699+
}
4700+
4701+
#[test]
4702+
fn known_up_to_the_newest_version_yields_every_known_version() {
4703+
assert_eq!(
4704+
ProtocolVersion::known_up_to(&ProtocolVersion::V_2026_07_28),
4705+
ProtocolVersion::KNOWN_VERSIONS
4706+
);
4707+
}
4708+
4709+
#[test]
4710+
fn known_up_to_an_unknown_ceiling_stops_at_the_versions_below_it() {
4711+
let unknown = ProtocolVersion(Cow::Borrowed("2025-07-01"));
4712+
assert_eq!(
4713+
ProtocolVersion::known_up_to(&unknown),
4714+
&[
4715+
ProtocolVersion::V_2024_11_05,
4716+
ProtocolVersion::V_2025_03_26,
4717+
ProtocolVersion::V_2025_06_18,
4718+
]
4719+
);
4720+
}
4721+
4722+
#[test]
4723+
fn known_up_to_a_ceiling_below_every_known_version_is_empty() {
4724+
let ancient = ProtocolVersion(Cow::Borrowed("1999-01-01"));
4725+
assert!(ProtocolVersion::known_up_to(&ancient).is_empty());
4726+
}
4727+
46464728
#[cfg(feature = "transport-streamable-http-client")]
46474729
#[test]
46484730
fn transport_closed_marker_accepts_only_the_process_local_token() {

‎crates/rmcp/src/service/server.rs‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -499,7 +499,9 @@ pub(crate) fn negotiate_protocol_version(
499499
server_supported,
500500
));
501501
};
502-
tracing::warn!(
502+
// Falling back is the designed answer for a pinned client, and stateless
503+
// HTTP re-runs it on every request, so this is not a warning.
504+
tracing::debug!(
503505
client_requested = %client_requested,
504506
server_fallback = %legacy_fallback,
505507
"client requested a protocol version unavailable over initialize; falling back to server default"

‎crates/rmcp/tests/test_protocol_version_negotiation.rs‎

Lines changed: 97 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,13 +5,19 @@
55
#![cfg(not(feature = "local"))]
66
#![cfg(feature = "client")]
77

8-
use std::borrow::Cow;
8+
use std::{
9+
borrow::Cow,
10+
sync::{
11+
Arc,
12+
atomic::{AtomicUsize, Ordering},
13+
},
14+
};
915

1016
use rmcp::{
1117
ClientHandler, ErrorData, RoleServer, ServerHandler, ServiceExt,
1218
model::{
13-
ClientInfo, ErrorCode, InitializeRequestParams, InitializeResult, ProtocolVersion,
14-
ServerInfo,
19+
ClientCapabilities, ClientInfo, ErrorCode, Implementation, InitializeRequestParams,
20+
InitializeResult, ProtocolVersion, ServerInfo,
1521
},
1622
service::{ClientInitializeError, RequestContext},
1723
};
@@ -232,3 +238,91 @@ async fn narrowed_server_caps_even_when_it_overrides_initialize() {
232238
"the handshake layer should not raise the version above what the server supports"
233239
);
234240
}
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+
}

‎examples/servers/src/common/counter.rs‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -270,15 +270,16 @@ impl ServerHandler for Counter {
270270

271271
async fn initialize(
272272
&self,
273-
_request: InitializeRequestParams,
273+
request: InitializeRequestParams,
274274
context: RequestContext<RoleServer>,
275275
) -> Result<InitializeResult, McpError> {
276276
if let Some(http_request_part) = context.extensions.get::<axum::http::request::Parts>() {
277277
let initialize_headers = &http_request_part.headers;
278278
let initialize_uri = &http_request_part.uri;
279279
tracing::info!(?initialize_headers, %initialize_uri, "initialize from http server");
280280
}
281-
Ok(self.get_info())
281+
context.peer.set_peer_info(request.clone());
282+
self.negotiate_initialize(&request)
282283
}
283284
}
284285

0 commit comments

Comments
 (0)