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
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,10 @@ where
{
type Error = C::Error;

fn preserves_raw_responses() -> bool {
C::preserves_raw_responses()
}

async fn delete_session(
&self,
uri: std::sync::Arc<str>,
Expand Down
103 changes: 100 additions & 3 deletions crates/rmcp/src/transport/streamable_http_client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,36 @@ fn cache_tools_from_response(
}
}

fn cache_tools_from_raw_response(
cache: &mut HashMap<String, Arc<JsonObject>>,
message: &mut RawRxJsonRpcMessage<RoleClient>,
protocol_version: &ProtocolVersion,
) {
if protocol_version < &ProtocolVersion::STANDARD_HEADERS {
return;
}
if let crate::model::JsonRpcMessage::Response(response) = message
&& let Some(tools) = response
.result
.get_mut("tools")
.and_then(serde_json::Value::as_array_mut)
{
tools.retain(|value| {
// Preserve the original JSON, including extension fields. Malformed
// tools remain for the normal response decoder to reject.
let Ok(tool) = serde_json::from_value::<crate::model::Tool>(value.clone()) else {
return true;
};
if let Err(reason) = mcp_headers::validate_param_header_annotations(&tool.input_schema) {
tracing::warn!(tool = %tool.name, "rejecting invalid x-mcp-header annotations: {reason}");
return false;
}
cache.insert(tool.name.to_string(), tool.input_schema);
true
});
}
}

fn negotiate_version_headers(
init_response: &ServerJsonRpcMessage,
base: HashMap<HeaderName, HeaderValue>,
Expand Down Expand Up @@ -245,7 +275,7 @@ pub enum StreamableHttpProtocolError {
MissingSessionIdInResponse,
}

#[expect(
#[allow(
clippy::large_enum_variant,
reason = "boxing the streaming response would add an allocation to the common response path"
)]
Expand Down Expand Up @@ -1681,8 +1711,17 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
context.send_to_handler(message).await?;
Ok(())
}
Ok(StreamableHttpPostResponse::RawJson(message, ..)) => {
context.send_to_handler(message).await?;
Ok(StreamableHttpPostResponse::RawJson(mut raw_message, ..)) => {
if matches!(&message, ClientJsonRpcMessage::Request(request)
if matches!(&request.request, ClientRequest::ListToolsRequest(_)))
{
cache_tools_from_raw_response(
&mut tool_header_cache,
&mut raw_message,
&version,
);
}
context.send_to_handler(raw_message).await?;
Ok(())
}
Ok(StreamableHttpPostResponse::Sse(stream, ..)) => {
Expand Down Expand Up @@ -2531,6 +2570,36 @@ mod tests {
);
}

#[test]
fn raw_tool_cache_preserves_extensions_and_filters_invalid_annotations() {
let mut valid = serde_json::to_value(tool(
"valid",
json!({"type": "string", "x-mcp-header": "Value"}),
))
.unwrap();
valid["vendorExtension"] = json!({"retained": true});
let invalid = serde_json::to_value(tool(
"invalid",
json!({"type": "string", "x-mcp-header": ""}),
))
.unwrap();
let mut message = RawRxJsonRpcMessage::<RoleClient>::response(
json!({"tools": [valid.clone(), invalid], "vendorResult": 42}),
NumberOrString::Number(1),
);
let mut cache = HashMap::new();
cache_tools_from_raw_response(&mut cache, &mut message, &ProtocolVersion::V_2026_07_28);
assert!(cache.contains_key("valid"));
assert!(!cache.contains_key("invalid"));
let crate::model::JsonRpcMessage::Response(response) = message else {
panic!("response")
};
assert_eq!(
response.result,
json!({"tools": [valid], "vendorResult": 42})
);
}

#[test]
fn cache_tools_preserves_pre_standard_header_results() {
let invalid = tool("legacy", json!({ "type": "string", "x-mcp-header": "" }));
Expand Down Expand Up @@ -2558,6 +2627,34 @@ mod tests {
);
}

#[test]
fn raw_tool_cache_preserves_legacy_results_and_malformed_tools() {
let malformed = json!({"name": "broken", "inputSchema": "not-an-object"});
let invalid = serde_json::to_value(tool("legacy", json!({"x-mcp-header": ""}))).unwrap();
let original = json!({"tools": [invalid, malformed.clone()], "vendorResult": 42});
let mut message = RawRxJsonRpcMessage::<RoleClient>::response(
original.clone(),
NumberOrString::Number(1),
);
let mut cache = HashMap::new();
cache_tools_from_raw_response(&mut cache, &mut message, &ProtocolVersion::V_2025_11_25);
assert!(cache.is_empty());
let crate::model::JsonRpcMessage::Response(response) = &message else {
panic!("response")
};
assert_eq!(response.result, original);

cache_tools_from_raw_response(&mut cache, &mut message, &ProtocolVersion::V_2026_07_28);
assert!(cache.is_empty());
let crate::model::JsonRpcMessage::Response(response) = message else {
panic!("response")
};
assert_eq!(
response.result,
json!({"tools": [malformed], "vendorResult": 42})
);
}

#[cfg(feature = "transport-streamable-http-client-reqwest")]
#[test]
fn clear_stream_response_pending_accepts_stringified_numeric_id() {
Expand Down
211 changes: 211 additions & 0 deletions crates/rmcp/tests/test_custom_request.rs
Original file line number Diff line number Diff line change
Expand Up @@ -170,6 +170,217 @@ async fn typed_test_client() -> anyhow::Result<rmcp::service::RunningService<rmc
Ok(().serve(client_transport).await?)
}

#[cfg(all(
feature = "auth",
feature = "transport-streamable-http-server",
feature = "transport-streamable-http-client-reqwest"
))]
#[tokio::test]
async fn typed_skills_request_survives_oauth_http_wrapper() -> anyhow::Result<()> {
use rmcp::transport::{
auth::{AuthClient, AuthorizationManager},
streamable_http_client::{StreamableHttpClientTransportConfig, StreamableHttpClientWorker},
streamable_http_server::{
StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager,
},
};
let stop = tokio_util::sync::CancellationToken::new();
let server: StreamableHttpService<TypedCustomRequestServer, LocalSessionManager> =
StreamableHttpService::new(
|| Ok(TypedCustomRequestServer),
Default::default(),
StreamableHttpServerConfig::default().with_cancellation_token(stop.child_token()),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
let endpoint = format!("http://{}/mcp", listener.local_addr()?);
let router = axum::Router::new().nest_service("/mcp", server);
let shutdown = stop.clone();
let task = tokio::spawn(async move {
axum::serve(listener, router)
.with_graceful_shutdown(shutdown.cancelled_owned())
.await
});
let result = tokio::time::timeout(Duration::from_secs(10), async {
let manager = AuthorizationManager::new(&endpoint).await?;
let auth = AuthClient::new(reqwest::Client::new(), manager);
let transport = StreamableHttpClientWorker::new(
auth,
StreamableHttpClientTransportConfig::with_uri(endpoint),
);
let client = ().serve(transport).await?;
let response = client
.send_request_as::<SkillsListResult>(ClientRequest::CustomRequest(CustomRequest::new(
"skills/list",
Some(json!({})),
)))
.await;
client.cancel().await?;
let response = response?;
assert_eq!(response.skills, ["example"]);
assert_eq!(
response.meta["io.modelcontextprotocol/serverInfo"]["name"],
"skills"
);
anyhow::Ok(())
})
.await;
stop.cancel();
tokio::time::timeout(Duration::from_secs(5), task).await???;
result??;
Ok(())
}

#[cfg(all(
feature = "auth",
feature = "transport-streamable-http-server",
feature = "transport-streamable-http-client-reqwest"
))]
#[tokio::test]
async fn json_tool_listing_populates_parameter_headers_through_auth_client() -> anyhow::Result<()> {
use rmcp::{
ClientServiceExt,
model::{
CallToolRequestParams, CallToolResponse, CallToolResult, ListToolsResult,
PaginatedRequestParams, ServerCapabilities, ServerInfo, Tool,
},
service::RequestContext,
transport::{
auth::{AuthClient, AuthorizationManager},
streamable_http_client::{
StreamableHttpClientTransportConfig, StreamableHttpClientWorker,
},
streamable_http_server::{
StreamableHttpServerConfig, StreamableHttpService,
session::local::LocalSessionManager,
},
},
};

struct AnnotatedToolServer;
impl ServerHandler for AnnotatedToolServer {
fn supported_protocol_versions(
&self,
) -> std::borrow::Cow<'static, [rmcp::model::ProtocolVersion]> {
std::borrow::Cow::Borrowed(&[rmcp::model::ProtocolVersion::V_2026_07_28])
}

fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}

fn get_tool(&self, name: &str) -> Option<Tool> {
(name == "deploy").then(|| {
Tool::new(
"deploy",
"Deploy in a region",
Arc::new(
json!({"type": "object", "properties": {
"region": {"type": "string", "x-mcp-header": "Region"}
}})
.as_object()
.unwrap()
.clone(),
),
)
})
}

async fn list_tools(
&self,
_: Option<PaginatedRequestParams>,
_: RequestContext<rmcp::RoleServer>,
) -> Result<ListToolsResult, ErrorData> {
Ok(ListToolsResult::with_all_items(vec![
self.get_tool("deploy").unwrap(),
]))
}

async fn call_tool(
&self,
_: CallToolRequestParams,
_: RequestContext<rmcp::RoleServer>,
) -> Result<CallToolResponse, ErrorData> {
Ok(CallToolResult::success(vec![]).into())
}
}

let observed_headers = Arc::new(Mutex::new(Vec::new()));
let capture = observed_headers.clone();
let stop = tokio_util::sync::CancellationToken::new();
let server: StreamableHttpService<AnnotatedToolServer, LocalSessionManager> =
StreamableHttpService::new(
|| Ok(AnnotatedToolServer),
Default::default(),
StreamableHttpServerConfig::default()
.with_legacy_session_mode(false)
.with_json_response(true)
.with_cancellation_token(stop.child_token()),
);
let router = axum::Router::new()
.nest_service("/mcp", server)
.layer(axum::middleware::from_fn(
move |request: axum::extract::Request, next: axum::middleware::Next| {
let capture = capture.clone();
async move {
if request
.headers()
.get("Mcp-Method")
.is_some_and(|method| method == "tools/call")
{
capture
.lock()
.await
.push(request.headers().get("Mcp-Param-Region").cloned());
}
next.run(request).await
}
},
));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
let endpoint = format!("http://{}/mcp", listener.local_addr()?);
let shutdown = stop.clone();
let task = tokio::spawn(async move {
axum::serve(listener, router)
.with_graceful_shutdown(shutdown.cancelled_owned())
.await
});
let result = tokio::time::timeout(Duration::from_secs(10), async {
let manager = AuthorizationManager::new(&endpoint).await?;
let transport = StreamableHttpClientWorker::new(
AuthClient::new(reqwest::Client::new(), manager),
StreamableHttpClientTransportConfig::with_uri(endpoint),
);
let client = ()
.serve_with_lifecycle(
transport,
rmcp::service::ClientLifecycleMode::Discover {
preferred_versions: vec![rmcp::model::ProtocolVersion::V_2026_07_28],
},
)
.await?;
let listed = client.list_tools(None).await?;
assert_eq!(listed.tools.len(), 1);
let response = client
.call_tool(
CallToolRequestParams::new("deploy")
.with_arguments(json!({"region": "us-east-1"}).as_object().unwrap().clone()),
)
.await;
client.cancel().await?;
response?;
assert_eq!(
*observed_headers.lock().await,
vec![Some(http::HeaderValue::from_static("us-east-1"))]
);
anyhow::Ok(())
})
.await;
stop.cancel();
tokio::time::timeout(Duration::from_secs(5), task).await???;
result??;
Ok(())
}

#[tokio::test]
async fn typed_custom_request_bypasses_the_core_response_union() -> anyhow::Result<()> {
let client = typed_test_client().await?;
Expand Down