use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, RwLock};
use std::task::{Context, Poll};
use tower::{Layer, ServiceExt};
use tower_service::Service;
use crate::error::{Error, JsonRpcError, Result};
use crate::inspection::{
McpDirection, McpInspection, McpInspectionError, McpInspectionErrorKind, McpInspector,
McpProtocolRevision,
};
use crate::protocol::{
JsonRpcMessage, JsonRpcRequest, JsonRpcResponse, JsonRpcResponseMessage, McpRequest, ResultType,
};
use crate::router::{Extensions, RouterRequest, RouterResponse};
use crate::{ProtocolSupport, ProtocolSupportError};
#[derive(Debug, Clone, Copy, Default)]
pub struct JsonRpcLayer {
_priv: (),
}
impl JsonRpcLayer {
pub fn new() -> Self {
Self { _priv: () }
}
}
impl<S> Layer<S> for JsonRpcLayer {
type Service = JsonRpcService<S>;
fn layer(&self, inner: S) -> Self::Service {
JsonRpcService::new(inner)
}
}
pub struct JsonRpcService<S> {
inner: S,
extensions: Extensions,
protocol_support: ProtocolSupport,
negotiated_revision: Arc<RwLock<Option<McpProtocolRevision>>>,
}
impl<S> JsonRpcService<S> {
pub fn new(inner: S) -> Self {
Self {
inner,
extensions: Extensions::new(),
protocol_support: ProtocolSupport::default(),
negotiated_revision: Arc::new(RwLock::new(None)),
}
}
pub fn with_extensions(mut self, ext: Extensions) -> Self {
self.extensions = ext;
self
}
pub fn protocol_support(mut self, support: ProtocolSupport) -> Self {
self.protocol_support = support;
self
}
pub fn protocol_versions<I, V>(
self,
versions: I,
) -> std::result::Result<Self, ProtocolSupportError>
where
I: IntoIterator<Item = V>,
V: Into<String>,
{
Ok(self.protocol_support(ProtocolSupport::try_new(versions)?))
}
pub(crate) fn configured_protocol_support(&self) -> &ProtocolSupport {
&self.protocol_support
}
pub(crate) fn inspect_incoming_value(
&self,
value: &serde_json::Value,
direction: McpDirection,
) -> std::result::Result<McpInspection, JsonRpcError> {
let protocol_support = self
.extensions
.get::<ProtocolSupport>()
.unwrap_or(&self.protocol_support);
let revision = self.resolve_revision(value, protocol_support)?;
inspect_runtime_value(value, revision, protocol_support, direction)
}
#[cfg_attr(not(feature = "http"), allow(dead_code))]
pub(crate) fn with_negotiated_protocol_version(self, version: &str) -> Self {
if let Ok(revision) = version.parse()
&& let Ok(mut selected) = self.negotiated_revision.write()
{
*selected = Some(revision);
}
self
}
fn resolve_revision(
&self,
value: &serde_json::Value,
protocol_support: &ProtocolSupport,
) -> std::result::Result<McpProtocolRevision, JsonRpcError> {
if let Some(items) = value.as_array() {
let mut declared = items.iter().filter_map(|item| {
item.pointer("/params/_meta/io.modelcontextprotocol~1protocolVersion")
.and_then(serde_json::Value::as_str)
});
if let Some(version) = declared.next() {
if declared.any(|candidate| candidate != version) {
return Err(JsonRpcError::invalid_request(
"A JSON-RPC batch cannot mix MCP protocol revisions",
));
}
return allowed_revision(version, protocol_support);
}
}
if let Some(version) = value
.pointer("/params/_meta/io.modelcontextprotocol~1protocolVersion")
.and_then(serde_json::Value::as_str)
{
return allowed_revision(version, protocol_support);
}
if value.get("method").and_then(serde_json::Value::as_str) == Some("initialize") {
let requested = value
.pointer("/params/protocolVersion")
.and_then(serde_json::Value::as_str);
let selected = requested
.filter(|version| {
crate::protocol::SUPPORTED_PROTOCOL_VERSIONS.contains(version)
&& protocol_support.contains(version)
})
.or_else(|| {
protocol_support.versions().iter().find_map(|version| {
crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
.contains(&version.as_str())
.then_some(version.as_str())
})
})
.ok_or_else(|| {
JsonRpcError::unsupported_protocol_version(
requested.unwrap_or("unknown"),
protocol_support.versions().iter().map(String::as_str),
)
})?;
return allowed_revision(selected, protocol_support);
}
if let Some(revision) = self
.extensions
.get::<McpProtocolRevision>()
.copied()
.or_else(|| {
self.negotiated_revision
.read()
.ok()
.and_then(|revision| *revision)
})
{
if protocol_support.contains(revision.as_str()) {
return Ok(revision);
}
return Err(JsonRpcError::unsupported_protocol_version(
revision.as_str(),
protocol_support.versions().iter().map(String::as_str),
));
}
if protocol_support.versions().len() == 1 {
return allowed_revision(protocol_support.preferred(), protocol_support);
}
if value.is_array() {
return Err(JsonRpcError::invalid_request(
"Cannot determine the exact MCP revision for a batch before protocol negotiation",
));
}
let provisional = protocol_support
.versions()
.iter()
.find(|version| {
crate::protocol::SUPPORTED_PROTOCOL_VERSIONS.contains(&version.as_str())
})
.ok_or_else(|| {
JsonRpcError::invalid_request(
"Request does not declare an MCP revision and no legacy revision is enabled",
)
})?;
allowed_revision(provisional, protocol_support)
}
#[cfg(feature = "stateless")]
pub(crate) fn validate_request_protocol(
&self,
req: &JsonRpcRequest,
) -> std::result::Result<Option<String>, JsonRpcError> {
req.validate()?;
let value = serde_json::to_value(req)
.map_err(|error| JsonRpcError::invalid_request(error.to_string()))?;
self.inspect_incoming_value(&value, McpDirection::ClientToServer)?;
let mut extensions = self.extensions.clone();
let protocol_support = extensions
.get::<ProtocolSupport>()
.cloned()
.unwrap_or_else(|| self.protocol_support.clone());
prepare_modern_request(req, &mut extensions, &protocol_support)
}
pub async fn call_single(&mut self, req: JsonRpcRequest) -> Result<JsonRpcResponse>
where
S: Service<RouterRequest, Response = RouterResponse, Error = std::convert::Infallible>
+ Clone
+ Send
+ 'static,
S::Future: Send,
{
process_single_request(
self.inner.clone(),
req,
self.extensions.clone(),
self.protocol_support.clone(),
self.negotiated_revision.clone(),
)
.await
}
pub async fn call_batch(
&mut self,
requests: Vec<JsonRpcRequest>,
) -> Result<Vec<JsonRpcResponse>>
where
S: Service<RouterRequest, Response = RouterResponse, Error = std::convert::Infallible>
+ Clone
+ Send
+ 'static,
S::Future: Send,
{
if requests.is_empty() {
return Err(Error::JsonRpc(JsonRpcError::invalid_request(
"Empty batch request",
)));
}
let value = serde_json::to_value(JsonRpcMessage::Batch(requests.clone()))
.map_err(Error::Serialization)?;
self.inspect_incoming_value(&value, McpDirection::ClientToServer)
.map_err(Error::JsonRpc)?;
let futures: Vec<_> = requests
.into_iter()
.map(|req| {
let inner = self.inner.clone();
let extensions = self.extensions.clone();
let protocol_support = self.protocol_support.clone();
let negotiated_revision = self.negotiated_revision.clone();
let req_id = req.id.clone();
async move {
match process_single_request(
inner,
req,
extensions,
protocol_support,
negotiated_revision,
)
.await
{
Ok(resp) => resp,
Err(e) => {
JsonRpcResponse::error(
Some(req_id),
JsonRpcError::internal_error(e.to_string()),
)
}
}
}
})
.collect();
let results: Vec<JsonRpcResponse> = futures::future::join_all(futures).await;
Ok(results)
}
pub async fn call_message(&mut self, msg: JsonRpcMessage) -> Result<JsonRpcResponseMessage>
where
S: Service<RouterRequest, Response = RouterResponse, Error = std::convert::Infallible>
+ Clone
+ Send
+ 'static,
S::Future: Send,
{
match msg {
JsonRpcMessage::Single(req) => {
let response = self.call_single(req).await?;
Ok(JsonRpcResponseMessage::Single(response))
}
JsonRpcMessage::Batch(requests) => match self.call_batch(requests).await {
Ok(responses) => Ok(JsonRpcResponseMessage::Batch(responses)),
Err(Error::JsonRpc(error)) => Ok(JsonRpcResponseMessage::Single(
JsonRpcResponse::error(None, error),
)),
Err(error) => Err(error),
},
_ => Ok(JsonRpcResponseMessage::Single(JsonRpcResponse::error(
None,
JsonRpcError::invalid_request("Unsupported message type"),
))),
}
}
}
impl<S> Clone for JsonRpcService<S>
where
S: Clone,
{
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
extensions: self.extensions.clone(),
protocol_support: self.protocol_support.clone(),
negotiated_revision: self.negotiated_revision.clone(),
}
}
}
impl<S> Service<JsonRpcRequest> for JsonRpcService<S>
where
S: Service<RouterRequest, Response = RouterResponse, Error = std::convert::Infallible>
+ Clone
+ Send
+ 'static,
S::Future: Send,
{
type Response = JsonRpcResponse;
type Error = Error;
type Future =
Pin<Box<dyn Future<Output = std::result::Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<std::result::Result<(), Self::Error>> {
self.inner.poll_ready(cx).map_err(|_| unreachable!())
}
fn call(&mut self, req: JsonRpcRequest) -> Self::Future {
let mut service = self.clone();
Box::pin(async move { service.call_single(req).await })
}
}
impl<S> Service<JsonRpcMessage> for JsonRpcService<S>
where
S: Service<RouterRequest, Response = RouterResponse, Error = std::convert::Infallible>
+ Clone
+ Send
+ 'static,
S::Future: Send,
{
type Response = JsonRpcResponseMessage;
type Error = Error;
type Future =
Pin<Box<dyn Future<Output = std::result::Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<std::result::Result<(), Self::Error>> {
self.inner.poll_ready(cx).map_err(|_| unreachable!())
}
fn call(&mut self, msg: JsonRpcMessage) -> Self::Future {
let mut service = self.clone();
Box::pin(async move { service.call_message(msg).await })
}
}
async fn process_single_request<S>(
inner: S,
req: JsonRpcRequest,
mut extensions: Extensions,
configured_protocol_support: ProtocolSupport,
negotiated_revision: Arc<RwLock<Option<McpProtocolRevision>>>,
) -> std::result::Result<JsonRpcResponse, Error>
where
S: Service<RouterRequest, Response = RouterResponse, Error = std::convert::Infallible>
+ Clone
+ Send
+ 'static,
S::Future: Send,
{
if let Err(e) = req.validate() {
return Ok(JsonRpcResponse::error(Some(req.id), e));
}
let method = req.method.clone();
#[cfg(feature = "stateless")]
let request_id = req.id.clone();
let protocol_support = extensions
.get::<ProtocolSupport>()
.cloned()
.unwrap_or(configured_protocol_support);
extensions.insert(protocol_support.clone());
let inspection_service = JsonRpcService {
inner: inner.clone(),
extensions: extensions.clone(),
protocol_support: protocol_support.clone(),
negotiated_revision: negotiated_revision.clone(),
};
let request_value = serde_json::to_value(&req).map_err(Error::Serialization)?;
if let Err(error) =
inspection_service.inspect_incoming_value(&request_value, McpDirection::ClientToServer)
{
return Ok(JsonRpcResponse::error(Some(req.id), error));
}
#[cfg(feature = "stateless")]
let protocol_version = match prepare_modern_request(&req, &mut extensions, &protocol_support) {
Ok(version) => version,
Err(error) => return Ok(JsonRpcResponse::error(Some(request_id), error)),
};
#[cfg(not(feature = "stateless"))]
let protocol_version: Option<String> = None;
let mcp_request = match McpRequest::from_jsonrpc(&req) {
Ok(r) => r,
Err(e) => {
return Ok(JsonRpcResponse::error(
Some(req.id),
JsonRpcError::invalid_params(e.to_string()),
));
}
};
let router_req = RouterRequest {
id: req.id,
inner: mcp_request,
extensions,
};
let response = inner.oneshot(router_req).await.unwrap();
let mut response = response.into_jsonrpc();
if method == "initialize"
&& let JsonRpcResponse::Result(result) = &response
&& let Some(version) = result
.result
.get("protocolVersion")
.and_then(serde_json::Value::as_str)
&& protocol_support.contains(version)
&& let Ok(revision) = version.parse::<McpProtocolRevision>()
&& let Ok(mut selected) = negotiated_revision.write()
{
*selected = Some(revision);
}
if let Some(version) = protocol_version.as_deref() {
apply_protocol_result_fields(&mut response, &method, version);
}
Ok(response)
}
fn allowed_revision(
version: &str,
protocol_support: &ProtocolSupport,
) -> std::result::Result<McpProtocolRevision, JsonRpcError> {
if !protocol_support.contains(version) {
return Err(JsonRpcError::unsupported_protocol_version(
version,
protocol_support.versions().iter().map(String::as_str),
));
}
version.parse().map_err(|_| {
JsonRpcError::unsupported_protocol_version(
version,
protocol_support.versions().iter().map(String::as_str),
)
})
}
pub(crate) fn inspect_runtime_value(
value: &serde_json::Value,
revision: McpProtocolRevision,
protocol_support: &ProtocolSupport,
direction: McpDirection,
) -> std::result::Result<McpInspection, JsonRpcError> {
if !protocol_support.contains(revision.as_str()) {
return Err(JsonRpcError::unsupported_protocol_version(
revision.as_str(),
protocol_support.versions().iter().map(String::as_str),
));
}
McpInspector::for_revision(revision)
.inspect(value, Some(direction))
.map_err(inspection_error_to_json_rpc)
}
fn inspection_error_to_json_rpc(error: McpInspectionError) -> JsonRpcError {
match error.kind() {
McpInspectionErrorKind::MissingParams | McpInspectionErrorKind::InvalidParams => {
JsonRpcError::invalid_params(error.to_string())
}
McpInspectionErrorKind::UnsupportedProfile => JsonRpcError::unsupported_protocol_version(
error.revision().unwrap_or("unknown"),
crate::inspection::MCP_INSPECTION_PROFILES.iter().copied(),
),
McpInspectionErrorKind::JsonRpc
| McpInspectionErrorKind::BatchUnavailable
| McpInspectionErrorKind::InitializeInBatch
| McpInspectionErrorKind::MessageKindMismatch
| McpInspectionErrorKind::DirectionMismatch => {
JsonRpcError::invalid_request(error.to_string())
}
_ => JsonRpcError::invalid_request(error.to_string()),
}
}
#[cfg(feature = "stateless")]
fn prepare_modern_request(
req: &JsonRpcRequest,
extensions: &mut Extensions,
protocol_support: &ProtocolSupport,
) -> std::result::Result<Option<String>, JsonRpcError> {
let Some(params) = req.params.as_ref() else {
return Ok(None);
};
let Some(meta_value) = params.as_object().and_then(|params| params.get("_meta")) else {
return Ok(None);
};
let claims_modern = meta_value
.as_object()
.is_some_and(|meta| meta.contains_key("io.modelcontextprotocol/protocolVersion"));
if !claims_modern {
return Ok(None);
}
crate::protocol::validate_meta_object(meta_value)
.map_err(|error| JsonRpcError::invalid_params(error.to_string()))?;
let meta_object = meta_value
.as_object()
.expect("validate_meta_object accepted a JSON object");
let protocol_version = meta_object
.get("io.modelcontextprotocol/protocolVersion")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| {
JsonRpcError::invalid_params(
"Missing or invalid _meta.io.modelcontextprotocol/protocolVersion",
)
})?;
let client_capabilities = meta_object
.get("io.modelcontextprotocol/clientCapabilities")
.ok_or_else(|| {
JsonRpcError::invalid_params("Missing _meta.io.modelcontextprotocol/clientCapabilities")
})?;
if !client_capabilities.is_object()
|| serde_json::from_value::<crate::protocol::ClientCapabilities>(
client_capabilities.clone(),
)
.is_err()
{
return Err(JsonRpcError::invalid_params(
"Invalid _meta.io.modelcontextprotocol/clientCapabilities",
));
}
if !protocol_support.contains(protocol_version) {
return Err(JsonRpcError::unsupported_protocol_version(
protocol_version,
protocol_support.versions().iter().map(String::as_str),
));
}
if protocol_version == crate::protocol::PROTOCOL_VERSION_2026_07_28
&& is_removed_modern_method(&req.method)
{
return Err(JsonRpcError::method_not_found(&req.method));
}
let meta: crate::stateless::StatelessRequestMeta =
serde_json::from_value(meta_value.clone())
.map_err(|error| JsonRpcError::invalid_params(error.to_string()))?;
extensions.insert(meta);
Ok(Some(protocol_version.to_string()))
}
#[cfg(feature = "stateless")]
fn is_removed_modern_method(method: &str) -> bool {
matches!(
method,
"initialize"
| "notifications/initialized"
| "ping"
| "logging/setLevel"
| "resources/subscribe"
| "resources/unsubscribe"
| "notifications/roots/list_changed"
)
}
pub(crate) fn apply_protocol_result_fields(
response: &mut JsonRpcResponse,
method: &str,
protocol_version: &str,
) {
if protocol_version != crate::protocol::PROTOCOL_VERSION_2026_07_28 {
return;
}
let JsonRpcResponse::Result(result) = response else {
return;
};
ResultType::Complete.stamp_into(&mut result.result, protocol_version);
if !is_cacheable_result_method(method) {
return;
}
let Some(object) = result.result.as_object_mut() else {
return;
};
object
.entry("ttlMs")
.or_insert_with(|| serde_json::Value::Number(0.into()));
object
.entry("cacheScope")
.or_insert_with(|| serde_json::Value::String("private".to_string()));
}
fn is_cacheable_result_method(method: &str) -> bool {
matches!(
method,
"server/discover"
| "tools/list"
| "prompts/list"
| "resources/list"
| "resources/read"
| "resources/templates/list"
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::McpRouter;
use crate::tool::ToolBuilder;
use schemars::JsonSchema;
use serde::Deserialize;
#[derive(Debug, Deserialize, JsonSchema)]
struct AddInput {
a: i32,
b: i32,
}
fn create_test_router() -> McpRouter {
let add_tool = ToolBuilder::new("add")
.description("Add two numbers")
.handler(|input: AddInput| async move {
Ok(crate::CallToolResult::text(format!(
"{}",
input.a + input.b
)))
})
.build();
McpRouter::new()
.server_info("test-server", "1.0.0")
.tool(add_tool)
}
#[tokio::test]
async fn test_jsonrpc_service() {
let router = create_test_router();
let mut service = JsonRpcService::new(router.clone());
let init_req = JsonRpcRequest::new(1, "initialize").with_params(serde_json::json!({
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}));
let resp = service.call_single(init_req).await.unwrap();
assert!(matches!(resp, JsonRpcResponse::Result(_)));
router.handle_notification(crate::protocol::McpNotification::Initialized);
let req = JsonRpcRequest::new(2, "tools/list").with_params(serde_json::json!({}));
let resp = service.call_single(req).await.unwrap();
match resp {
JsonRpcResponse::Result(r) => {
let tools = r.result.get("tools").unwrap().as_array().unwrap();
assert_eq!(tools.len(), 1);
}
JsonRpcResponse::Error(e) => panic!("Expected result, got error: {:?}", e),
_ => panic!("unexpected response variant"),
}
}
#[tokio::test]
async fn test_batch_request() {
let router = create_test_router();
let mut service = JsonRpcService::new(router.clone());
let init_req = JsonRpcRequest::new(1, "initialize").with_params(serde_json::json!({
"protocolVersion": "2025-03-26",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}));
service.call_single(init_req).await.unwrap();
router.handle_notification(crate::protocol::McpNotification::Initialized);
let requests = vec![
JsonRpcRequest::new(2, "tools/list").with_params(serde_json::json!({})),
JsonRpcRequest::new(3, "tools/call").with_params(serde_json::json!({
"name": "add",
"arguments": { "a": 1, "b": 2 }
})),
];
let responses = service.call_batch(requests).await.unwrap();
assert_eq!(responses.len(), 2);
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn modern_batch_is_rejected_by_exact_profile() {
let router = create_test_router();
let mut service = JsonRpcService::new(router);
let final_request = JsonRpcRequest::new(1, "tools/list").with_params(serde_json::json!({
"_meta": {
"io.modelcontextprotocol/protocolVersion":
crate::protocol::PROTOCOL_VERSION_2026_07_28,
"io.modelcontextprotocol/clientCapabilities": {}
}
}));
let legacy_request =
JsonRpcRequest::new(2, "tools/list").with_params(serde_json::json!({}));
let error = service
.call_batch(vec![final_request, legacy_request])
.await
.unwrap_err();
let Error::JsonRpc(error) = error else {
panic!("final batch should fail with a JSON-RPC error");
};
assert_eq!(error.code, -32600);
assert!(
error
.message
.contains("does not permit top-level JSON-RPC batches")
);
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn modern_request_requires_client_capabilities() {
let router = create_test_router();
let mut service = JsonRpcService::new(router);
let request = JsonRpcRequest::new(1, "server/discover").with_params(serde_json::json!({
"_meta": {
"io.modelcontextprotocol/protocolVersion":
crate::protocol::PROTOCOL_VERSION_2026_07_28
}
}));
let response = service.call_single(request).await.unwrap();
let JsonRpcResponse::Error(response) = response else {
panic!("missing clientCapabilities must be rejected");
};
assert_eq!(response.error.code, -32602);
assert!(response.error.message.contains("clientCapabilities"));
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn final_only_policy_never_negotiates_final_via_initialize() {
let router = create_test_router();
let mut service = JsonRpcService::new(router).protocol_support(
ProtocolSupport::try_new([crate::protocol::PROTOCOL_VERSION_2026_07_28]).unwrap(),
);
let request = JsonRpcRequest::new(1, "initialize").with_params(serde_json::json!({
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {"name": "legacy-client", "version": "1.0.0"}
}));
let response = service.call_single(request).await.unwrap();
let JsonRpcResponse::Error(response) = response else {
panic!("final-only policy must reject the removed initialize lifecycle");
};
assert_eq!(response.error.code, -32022);
assert_eq!(
response.error.data.unwrap()["supported"],
serde_json::json!([crate::protocol::PROTOCOL_VERSION_2026_07_28])
);
}
#[tokio::test]
async fn test_empty_batch_error() {
let router = create_test_router();
let mut service = JsonRpcService::new(router);
let result = service.call_batch(vec![]).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_jsonrpc_layer() {
use tower::ServiceBuilder;
let router = create_test_router();
let router_clone = router.clone();
let mut service = ServiceBuilder::new()
.layer(JsonRpcLayer::new())
.service(router);
let init_req = JsonRpcRequest::new(1, "initialize").with_params(serde_json::json!({
"protocolVersion": "2025-03-26",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}));
let resp = Service::<JsonRpcRequest>::call(&mut service, init_req)
.await
.unwrap();
assert!(matches!(resp, JsonRpcResponse::Result(_)));
router_clone.handle_notification(crate::protocol::McpNotification::Initialized);
let req = JsonRpcRequest::new(2, "tools/list").with_params(serde_json::json!({}));
let resp = Service::<JsonRpcRequest>::call(&mut service, req)
.await
.unwrap();
match resp {
JsonRpcResponse::Result(r) => {
let tools = r.result.get("tools").unwrap().as_array().unwrap();
assert_eq!(tools.len(), 1);
}
JsonRpcResponse::Error(e) => panic!("Expected result, got error: {:?}", e),
_ => panic!("unexpected response variant"),
}
}
#[test]
fn test_jsonrpc_layer_default() {
let _layer = JsonRpcLayer::default();
}
#[test]
fn test_jsonrpc_layer_clone() {
let layer = JsonRpcLayer::new();
let _cloned = layer;
let _copied = layer;
}
#[tokio::test]
async fn test_invalid_jsonrpc_version() {
let router = create_test_router();
let mut service = JsonRpcService::new(router);
let req = JsonRpcRequest {
jsonrpc: "1.0".to_string(),
id: crate::protocol::RequestId::Number(1),
method: "ping".to_string(),
params: None,
};
let resp = service.call_single(req).await.unwrap();
match resp {
JsonRpcResponse::Error(e) => {
assert_eq!(e.error.code, -32600); }
_ => panic!("Expected error for invalid jsonrpc version"),
}
}
#[tokio::test]
async fn test_unknown_method() {
let router = create_test_router();
let mut service = JsonRpcService::new(router.clone());
let init_req = JsonRpcRequest::new(1, "initialize").with_params(serde_json::json!({
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}));
service.call_single(init_req).await.unwrap();
router.handle_notification(crate::protocol::McpNotification::Initialized);
let req = JsonRpcRequest::new(2, "nonexistent/method");
let resp = service.call_single(req).await.unwrap();
match resp {
JsonRpcResponse::Error(e) => {
assert_eq!(e.error.code, -32601); }
_ => panic!("Expected error for unknown method"),
}
}
#[tokio::test]
async fn test_invalid_params() {
let router = create_test_router();
let mut service = JsonRpcService::new(router.clone());
let init_req = JsonRpcRequest::new(1, "initialize").with_params(serde_json::json!({
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}));
service.call_single(init_req).await.unwrap();
router.handle_notification(crate::protocol::McpNotification::Initialized);
let req = JsonRpcRequest::new(2, "tools/call").with_params(serde_json::json!({
"wrong_field": "value"
}));
let resp = service.call_single(req).await.unwrap();
match resp {
JsonRpcResponse::Error(e) => {
assert_eq!(e.error.code, -32602); }
_ => panic!("Expected error for invalid params"),
}
}
#[tokio::test]
async fn exact_profile_rejects_malformed_present_params() {
let router = create_test_router();
let mut service = JsonRpcService::new(router.clone());
let initialize = JsonRpcRequest::new(0, "initialize").with_params(serde_json::json!({
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {"name": "test", "version": "1.0"}
}));
service.call_single(initialize).await.unwrap();
router.handle_notification(crate::protocol::McpNotification::Initialized);
let request = JsonRpcRequest::new(1, "tools/list")
.with_params(serde_json::json!(["not", "an", "object"]));
let response = service.call_single(request).await.unwrap();
let JsonRpcResponse::Error(error) = response else {
panic!("malformed present params should be rejected");
};
assert_eq!(error.error.code, -32602);
assert!(error.error.message.contains("`tools/list` params"));
}
#[tokio::test]
async fn test_request_before_initialize() {
let router = create_test_router();
let mut service = JsonRpcService::new(router);
let req = JsonRpcRequest::new(1, "tools/list").with_params(serde_json::json!({}));
let resp = service.call_single(req).await.unwrap();
match resp {
JsonRpcResponse::Error(e) => {
assert_eq!(e.error.code, -32600); }
_ => panic!("Expected error for request before initialize"),
}
}
#[tokio::test]
async fn test_ping_before_initialize() {
let router = create_test_router();
let mut service = JsonRpcService::new(router);
let req = JsonRpcRequest::new(1, "ping");
let resp = service.call_single(req).await.unwrap();
assert!(matches!(resp, JsonRpcResponse::Result(_)));
}
#[tokio::test]
async fn test_call_message_single() {
let router = create_test_router();
let mut service = JsonRpcService::new(router);
let msg = JsonRpcMessage::Single(JsonRpcRequest::new(1, "ping"));
let resp = service.call_message(msg).await.unwrap();
assert!(matches!(resp, JsonRpcResponseMessage::Single(_)));
}
#[tokio::test]
async fn test_call_message_batch() {
let router = create_test_router();
let mut service = JsonRpcService::new(router.clone());
let init_req = JsonRpcRequest::new(1, "initialize").with_params(serde_json::json!({
"protocolVersion": "2025-03-26",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}));
service.call_single(init_req).await.unwrap();
router.handle_notification(crate::protocol::McpNotification::Initialized);
let msg = JsonRpcMessage::Batch(vec![
JsonRpcRequest::new(2, "ping"),
JsonRpcRequest::new(3, "tools/list").with_params(serde_json::json!({})),
]);
let resp = service.call_message(msg).await.unwrap();
match resp {
JsonRpcResponseMessage::Batch(responses) => {
assert_eq!(responses.len(), 2);
}
_ => panic!("Expected batch response"),
}
}
#[tokio::test]
async fn negotiated_2025_11_rejects_batch() {
let router = create_test_router();
let mut service = JsonRpcService::new(router.clone());
let init_req = JsonRpcRequest::new(1, "initialize").with_params(serde_json::json!({
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}));
service.call_single(init_req).await.unwrap();
router.handle_notification(crate::protocol::McpNotification::Initialized);
let message = JsonRpcMessage::Batch(vec![JsonRpcRequest::new(2, "ping")]);
let response = service.call_message(message).await.unwrap();
let JsonRpcResponseMessage::Single(JsonRpcResponse::Error(error)) = response else {
panic!("2025-11-25 batch should produce one JSON-RPC error");
};
assert_eq!(error.error.code, -32600);
assert!(
error
.error
.message
.contains("does not permit top-level JSON-RPC batches")
);
}
#[tokio::test]
async fn test_call_message_empty_batch() {
let router = create_test_router();
let mut service = JsonRpcService::new(router);
let msg = JsonRpcMessage::Batch(vec![]);
let result = service.call_message(msg).await.unwrap();
let JsonRpcResponseMessage::Single(JsonRpcResponse::Error(error)) = result else {
panic!("empty batch should produce one JSON-RPC error response");
};
assert_eq!(error.error.code, -32600);
}
#[tokio::test]
async fn test_extensions_bridging() {
let router = create_test_router();
#[derive(Debug, Clone)]
#[allow(dead_code)]
struct TestClaim(String);
let mut ext = Extensions::new();
ext.insert(TestClaim("admin".to_string()));
let mut service = JsonRpcService::new(router).with_extensions(ext);
let req = JsonRpcRequest::new(1, "ping");
let resp = service.call_single(req).await.unwrap();
assert!(matches!(resp, JsonRpcResponse::Result(_)));
}
#[tokio::test]
async fn test_batch_with_mixed_valid_invalid() {
let router = create_test_router();
let mut service = JsonRpcService::new(router.clone());
let init_req = JsonRpcRequest::new(1, "initialize").with_params(serde_json::json!({
"protocolVersion": "2025-03-26",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}));
service.call_single(init_req).await.unwrap();
router.handle_notification(crate::protocol::McpNotification::Initialized);
let requests = vec![
JsonRpcRequest::new(2, "ping"),
JsonRpcRequest::new(3, "nonexistent/method"),
];
let responses = service.call_batch(requests).await.unwrap();
assert_eq!(responses.len(), 2);
assert!(matches!(&responses[0], JsonRpcResponse::Result(_)));
assert!(matches!(&responses[1], JsonRpcResponse::Error(_)));
}
}