Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
100 changes: 40 additions & 60 deletions crates/rmcp/src/model.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1120,9 +1120,9 @@ impl schemars::JsonSchema for DiscoverRequestParams {
pub type DiscoverRequest = Request<DiscoverRequestMethod, DiscoverRequestParams>;

/// The server's response to a [`DiscoverRequest`].
#[derive(Debug, Serialize, Clone, PartialEq)]
#[serde(rename_all = "camelCase")]
#[derive(Debug, Serialize, Deserialize, Clone, PartialEq)]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
#[serde(rename_all = "camelCase")]
#[non_exhaustive]
pub struct DiscoverResult {
/// Identifies how the result should be parsed.
Expand All @@ -1131,8 +1131,6 @@ pub struct DiscoverResult {
pub supported_versions: Vec<ProtocolVersion>,
/// Capabilities provided by this server.
pub capabilities: ServerCapabilities,
/// Information about the server implementation.
pub server_info: Implementation,
/// Optional guidance for using the server.
#[serde(skip_serializing_if = "Option::is_none")]
pub instructions: Option<String>,
Expand All @@ -1145,72 +1143,47 @@ pub struct DiscoverResult {
pub meta: Option<MetaObject>,
}

impl<'de> Deserialize<'de> for DiscoverResult {
fn deserialize<__D>(deserializer: __D) -> Result<Self, __D::Error>
where
__D: serde::Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct Helper {
result_type: ResultType,
supported_versions: Vec<ProtocolVersion>,
capabilities: ServerCapabilities,
server_info: Option<Implementation>,
instructions: Option<String>,
ttl_ms: u64,
cache_scope: CacheScope,
#[serde(rename = "_meta")]
meta: Option<MetaObject>,
}

let helper = Helper::deserialize(deserializer)?;
let server_info = match helper.server_info {
Some(server_info) => server_info,
None => {
let metadata_server_info = helper
.meta
.as_ref()
.and_then(|metadata| metadata.0.get("io.modelcontextprotocol/serverInfo"))
.ok_or_else(|| serde::de::Error::missing_field("serverInfo"))?;

serde_json::from_value(metadata_server_info.clone())
.map_err(serde::de::Error::custom)?
}
};

Ok(Self {
result_type: helper.result_type,
supported_versions: helper.supported_versions,
capabilities: helper.capabilities,
server_info,
instructions: helper.instructions,
ttl_ms: helper.ttl_ms,
cache_scope: helper.cache_scope,
meta: helper.meta,
})
}
}

impl DiscoverResult {
const SERVER_INFO_META_KEY: &str = "io.modelcontextprotocol/serverInfo";

/// Create a non-cacheable private discovery result.
pub fn new(
supported_versions: Vec<ProtocolVersion>,
capabilities: ServerCapabilities,
server_info: Implementation,
) -> Self {
pub fn new(supported_versions: Vec<ProtocolVersion>, capabilities: ServerCapabilities) -> Self {
Self {
result_type: ResultType::COMPLETE,
supported_versions,
capabilities,
server_info,
instructions: None,
ttl_ms: 0,
cache_scope: CacheScope::Private,
meta: None,
}
}

/// Return the server implementation information stored in result metadata.
pub fn server_info(&self) -> Option<Implementation> {
self.meta
.as_ref()?
.0
.get(Self::SERVER_INFO_META_KEY)
.and_then(|value| serde_json::from_value(value.clone()).ok())
}

/// Store server implementation information in result metadata.
pub fn set_server_info(&mut self, server_info: Implementation) {
let server_info =
serde_json::to_value(server_info).expect("Implementation serialization cannot fail");
self.meta
.get_or_insert_default()
.0
.insert(Self::SERVER_INFO_META_KEY.to_owned(), server_info);
}

/// Store server implementation information in result metadata.
pub fn with_server_info(mut self, server_info: Implementation) -> Self {
self.set_server_info(server_info);
self
}

/// Create a discovery result from the server's initialization information.
pub fn from_server_info(
supported_versions: Vec<ProtocolVersion>,
Expand All @@ -1223,9 +1196,16 @@ impl DiscoverResult {
meta,
..
} = server_info;
let mut result = Self::new(supported_versions, capabilities, server_info);
result.instructions = instructions;
result.meta = meta;
let mut result = Self {
result_type: ResultType::COMPLETE,
supported_versions,
capabilities,
instructions,
ttl_ms: 0,
cache_scope: CacheScope::Private,
meta,
};
result.set_server_info(server_info);
result
}

Expand Down
27 changes: 13 additions & 14 deletions crates/rmcp/src/service/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -846,13 +846,15 @@ where
server_supported: result.supported_versions,
});
};
peer.set_peer_info(ServerInfo {
protocol_version: selected.clone(),
capabilities: result.capabilities,
server_info: result.server_info,
instructions: result.instructions,
meta: result.meta,
});
if let Some(server_info) = result.server_info() {
peer.set_peer_info(ServerInfo {
protocol_version: selected.clone(),
capabilities: result.capabilities,
server_info,
instructions: result.instructions,
meta: result.meta,
});
}
peer.set_client_request_metadata(ClientRequestMetadata {
protocol_version: selected,
client_info: client_info.client_info.clone(),
Expand Down Expand Up @@ -2197,13 +2199,10 @@ mod tests {
let peer = disconnected_peer();
let meta = RequestMetaObject::default();
let key = discover_cache_key();
let expected = DiscoverResult::new(
vec![ProtocolVersion::default()],
Default::default(),
crate::model::Implementation::from_build_env(),
)
.with_ttl_ms(5_000)
.with_cache_scope(CacheScope::Public);
let expected = DiscoverResult::new(vec![ProtocolVersion::default()], Default::default())
.with_server_info(crate::model::Implementation::from_build_env())
.with_ttl_ms(5_000)
.with_cache_scope(CacheScope::Public);
peer.cache_response(
key,
ServerResult::DiscoverResult(expected.clone()),
Expand Down
103 changes: 83 additions & 20 deletions crates/rmcp/tests/test_client_lifecycle_modes.rs
Original file line number Diff line number Diff line change
Expand Up @@ -36,11 +36,13 @@ async fn discover_startup_accepts_stringified_numeric_response_id() {
};
server
.send(ServerJsonRpcMessage::response(
ServerResult::DiscoverResult(DiscoverResult::new(
vec![ProtocolVersion::V_2026_07_28],
ServerCapabilities::default(),
Implementation::new("discover-server", "1.0.0"),
)),
ServerResult::DiscoverResult(
DiscoverResult::new(
vec![ProtocolVersion::V_2026_07_28],
ServerCapabilities::default(),
)
.with_server_info(Implementation::new("discover-server", "1.0.0")),
),
RequestId::String(response_id.to_string().into()),
))
.await
Expand All @@ -60,6 +62,61 @@ async fn discover_startup_accepts_stringified_numeric_response_id() {
server_task.await.expect("server task");
}

#[tokio::test]
async fn discover_startup_accepts_missing_optional_server_info() {
let (server_transport, client_transport) = tokio::io::duplex(4096);
let mut server = IntoTransport::<rmcp::RoleServer, _, _>::into_transport(server_transport);
let server_task = tokio::spawn(async move {
let ClientJsonRpcMessage::Request(discover_request) =
server.receive().await.expect("expected discover request")
else {
panic!("expected discover request");
};
let result = DiscoverResult::new(
vec![ProtocolVersion::V_2026_07_28],
ServerCapabilities::default(),
);
server
.send(ServerJsonRpcMessage::response(
ServerResult::DiscoverResult(result),
discover_request.id,
))
.await
.expect("send discover response");

let ClientJsonRpcMessage::Request(request) =
server.receive().await.expect("expected normal request")
else {
panic!("expected normal request");
};
assert_eq!(
request.request.get_meta().protocol_version(),
Some(ProtocolVersion::V_2026_07_28)
);
server
.send(ServerJsonRpcMessage::response(
ServerResult::ListToolsResult(Default::default()),
request.id,
))
.await
.expect("send list tools response");
});

let client = DiscoverClient
.serve_with_lifecycle(
client_transport,
ClientLifecycleMode::Discover {
preferred_versions: vec![ProtocolVersion::V_2026_07_28],
},
)
.await
.expect("missing optional server info should not fail discovery");
assert!(client.peer_info().is_none());
client.list_tools(None).await.expect("list tools");
client.cancel().await.expect("cancel client");
server_task.await.expect("server task");
}

#[tokio::test]
async fn high_level_server_accepts_discover_startup_without_initialize() {
let (server_transport, client_transport) = tokio::io::duplex(4096);
Expand Down Expand Up @@ -103,11 +160,13 @@ async fn discover_startup_omits_initialize() {

server
.send(ServerJsonRpcMessage::response(
ServerResult::DiscoverResult(DiscoverResult::new(
vec![ProtocolVersion::V_2026_07_28],
ServerCapabilities::default(),
Implementation::new("discover-server", "1.0.0"),
)),
ServerResult::DiscoverResult(
DiscoverResult::new(
vec![ProtocolVersion::V_2026_07_28],
ServerCapabilities::default(),
)
.with_server_info(Implementation::new("discover-server", "1.0.0")),
),
request.id,
))
.await
Expand Down Expand Up @@ -267,11 +326,13 @@ async fn discover_startup_retries_a_mutually_supported_version() {
);
server
.send(ServerJsonRpcMessage::response(
ServerResult::DiscoverResult(DiscoverResult::new(
vec![ProtocolVersion::V_2026_07_28],
ServerCapabilities::default(),
Implementation::new("discover-server", "1.0.0"),
)),
ServerResult::DiscoverResult(
DiscoverResult::new(
vec![ProtocolVersion::V_2026_07_28],
ServerCapabilities::default(),
)
.with_server_info(Implementation::new("discover-server", "1.0.0")),
),
second.id,
))
.await
Expand Down Expand Up @@ -326,11 +387,13 @@ async fn discover_startup_retries_current_version_once_when_server_reports_it_su
);
server
.send(ServerJsonRpcMessage::response(
ServerResult::DiscoverResult(DiscoverResult::new(
vec![ProtocolVersion::V_2026_07_28],
ServerCapabilities::default(),
Implementation::new("discover-server", "1.0.0"),
)),
ServerResult::DiscoverResult(
DiscoverResult::new(
vec![ProtocolVersion::V_2026_07_28],
ServerCapabilities::default(),
)
.with_server_info(Implementation::new("discover-server", "1.0.0")),
),
second.id,
))
.await
Expand Down
7 changes: 7 additions & 0 deletions crates/rmcp/tests/test_message_schema.rs
Original file line number Diff line number Diff line change
Expand Up @@ -121,6 +121,13 @@ mod tests {
let schema = settings
.into_generator()
.into_root_schema_for::<ServerJsonRpcMessage>();
let schema_value =
serde_json::to_value(&schema).expect("Failed to serialize server schema");
let discover_result = &schema_value["definitions"]["DiscoverResult"];
assert!(
discover_result["properties"].get("serverInfo").is_none(),
"DiscoverResult serverInfo belongs in namespaced _meta"
);
let schema_str = serde_json::to_string_pretty(&schema).expect("Failed to serialize schema");

compare_schemas(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -759,14 +759,6 @@
}
]
},
"serverInfo": {
"description": "Information about the server implementation.",
"allOf": [
{
"$ref": "#/definitions/Implementation"
}
]
},
"supportedVersions": {
"description": "Protocol versions implemented by this server.",
"type": "array",
Expand All @@ -785,7 +777,6 @@
"resultType",
"supportedVersions",
"capabilities",
"serverInfo",
"ttlMs",
"cacheScope"
]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -759,14 +759,6 @@
}
]
},
"serverInfo": {
"description": "Information about the server implementation.",
"allOf": [
{
"$ref": "#/definitions/Implementation"
}
]
},
"supportedVersions": {
"description": "Protocol versions implemented by this server.",
"type": "array",
Expand All @@ -785,7 +777,6 @@
"resultType",
"supportedVersions",
"capabilities",
"serverInfo",
"ttlMs",
"cacheScope"
]
Expand Down
Loading
Loading