Skip to content

Commit 9234e2b

Browse files
committed
feat: Implement SEP-2260 require server requests to associate with client requests
1 parent 07abfdf commit 9234e2b

3 files changed

Lines changed: 139 additions & 2 deletions

File tree

crates/rmcp/src/service.rs

Lines changed: 11 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -146,6 +146,14 @@ pub(crate) fn uses_legacy_lifecycle(
146146
&& protocol_version.is_none_or(|version| version < &ProtocolVersion::V_2026_07_28)
147147
}
148148

149+
tokio::task_local! {
150+
pub(crate) static ORIGINATING_REQUEST: RequestId;
151+
}
152+
153+
pub(crate) fn in_request_handler_scope() -> bool {
154+
ORIGINATING_REQUEST.try_with(|_| ()).is_ok()
155+
}
156+
149157
pub type TxJsonRpcMessage<R> =
150158
JsonRpcMessage<<R as ServiceRole>::Req, <R as ServiceRole>::Resp, <R as ServiceRole>::Not>;
151159
pub type RxJsonRpcMessage<R> = JsonRpcMessage<
@@ -1379,9 +1387,10 @@ where
13791387
extensions,
13801388
};
13811389
let current_span = tracing::Span::current();
1390+
let handler_id = id.clone();
13821391
spawn_service_task(async move {
1383-
let result = service
1384-
.handle_request(request, context)
1392+
let result = ORIGINATING_REQUEST
1393+
.scope(handler_id, service.handle_request(request, context))
13851394
.await;
13861395
let response = match result {
13871396
Ok(result) => {

crates/rmcp/src/service/server.rs

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -596,6 +596,7 @@ macro_rules! method {
596596
($(#[$meta:meta])* peer_req $method:ident $Req:ident() => $Resp: ident ) => {
597597
$(#[$meta])*
598598
pub async fn $method(&self) -> Result<$Resp, ServiceError> {
599+
self.ensure_request_association()?;
599600
let result = self
600601
.send_request(ServerRequest::$Req($Req {
601602
method: Default::default(),
@@ -611,6 +612,7 @@ macro_rules! method {
611612
($(#[$meta:meta])* peer_req $method:ident $Req:ident($Param: ident) => $Resp: ident ) => {
612613
$(#[$meta])*
613614
pub async fn $method(&self, params: $Param) -> Result<$Resp, ServiceError> {
615+
self.ensure_request_association()?;
614616
let result = self
615617
.send_request(ServerRequest::$Req($Req {
616618
method: Default::default(),
@@ -680,6 +682,7 @@ macro_rules! method {
680682
method: Default::default(),
681683
extensions: Default::default(),
682684
});
685+
self.ensure_request_association()?;
683686
let options = crate::service::PeerRequestOptions {
684687
timeout,
685688
meta: None,
@@ -705,6 +708,7 @@ macro_rules! method {
705708
params: $Param,
706709
timeout: Option<std::time::Duration>,
707710
) -> Result<$Resp, ServiceError> {
711+
self.ensure_request_association()?;
708712
let request = ServerRequest::$Req($Req {
709713
method: Default::default(),
710714
params,
@@ -730,6 +734,19 @@ macro_rules! method {
730734
}
731735

732736
impl Peer<RoleServer> {
737+
fn ensure_request_association(&self) -> Result<(), ServiceError> {
738+
let strict = self
739+
.peer_info()
740+
.is_some_and(|info| info.protocol_version >= ProtocolVersion::V_2026_07_28);
741+
if strict && !crate::service::in_request_handler_scope() {
742+
return Err(ServiceError::McpError(ErrorData::invalid_request(
743+
"SEP-2260: server-to-client requests must be associated with an originating client request",
744+
None,
745+
)));
746+
}
747+
Ok(())
748+
}
749+
733750
/// Check if the client supports sampling tools capability.
734751
pub fn supports_sampling_tools(&self) -> bool {
735752
if let Some(client_info) = self.peer_info() {
@@ -752,6 +769,7 @@ impl Peer<RoleServer> {
752769
&self,
753770
params: CreateMessageRequestParams,
754771
) -> Result<CreateMessageResult, ServiceError> {
772+
self.ensure_request_association()?;
755773
// MUST throw error when tools/toolChoice provided without capability
756774
if (params.tools.is_some() || params.tool_choice.is_some())
757775
&& !self.supports_sampling_tools()
Lines changed: 110 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,110 @@
1+
#![cfg(all(feature = "server", feature = "client", not(feature = "local")))]
2+
#![allow(deprecated)]
3+
4+
use std::sync::{Arc, Mutex};
5+
6+
use rmcp::{
7+
ClientHandler, RoleClient, RoleServer, ServerHandler, ServiceError, ServiceExt,
8+
model::{
9+
CallToolRequestParams, CallToolResponse, CallToolResult, ClientInfo, ContentBlock,
10+
CreateMessageRequestParams, CreateMessageResult, ProtocolVersion, SamplingMessage,
11+
ServerCapabilities, ServerInfo,
12+
},
13+
service::RequestContext,
14+
};
15+
use tokio::sync::oneshot;
16+
17+
#[derive(Clone)]
18+
struct SamplingServer {
19+
outside: Arc<Mutex<Option<oneshot::Sender<Result<(), ServiceError>>>>>,
20+
}
21+
22+
impl ServerHandler for SamplingServer {
23+
fn get_info(&self) -> ServerInfo {
24+
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
25+
}
26+
27+
async fn call_tool(
28+
&self,
29+
_request: CallToolRequestParams,
30+
context: RequestContext<RoleServer>,
31+
) -> Result<CallToolResponse, rmcp::ErrorData> {
32+
let peer = context.peer.clone();
33+
let nested = peer
34+
.create_message(CreateMessageRequestParams::new(
35+
vec![SamplingMessage::user_text("nested")],
36+
16,
37+
))
38+
.await;
39+
let slot = self.outside.clone();
40+
tokio::spawn(async move {
41+
let outside = peer
42+
.create_message(CreateMessageRequestParams::new(
43+
vec![SamplingMessage::user_text("standalone")],
44+
16,
45+
))
46+
.await
47+
.map(|_| ());
48+
if let Some(tx) = slot.lock().unwrap().take() {
49+
let _ = tx.send(outside);
50+
}
51+
});
52+
nested.map_err(|e| rmcp::ErrorData::internal_error(e.to_string(), None))?;
53+
Ok(CallToolResult::success(vec![ContentBlock::text("ok")]).into())
54+
}
55+
}
56+
57+
#[derive(Clone)]
58+
struct SamplingClient;
59+
60+
impl ClientHandler for SamplingClient {
61+
async fn create_message(
62+
&self,
63+
_params: CreateMessageRequestParams,
64+
_context: RequestContext<RoleClient>,
65+
) -> Result<CreateMessageResult, rmcp::ErrorData> {
66+
Ok(CreateMessageResult::new(
67+
SamplingMessage::assistant_text("pong"),
68+
"test-model".to_string(),
69+
)
70+
.with_stop_reason(CreateMessageResult::STOP_REASON_END_TURN))
71+
}
72+
73+
fn get_info(&self) -> ClientInfo {
74+
let mut info = ClientInfo::default();
75+
info.protocol_version = ProtocolVersion::V_2026_07_28;
76+
info
77+
}
78+
}
79+
80+
#[tokio::test]
81+
async fn nested_sampling_allowed_standalone_rejected() -> anyhow::Result<()> {
82+
let (server_transport, client_transport) = tokio::io::duplex(4096);
83+
let (tx, rx) = oneshot::channel();
84+
let server = SamplingServer {
85+
outside: Arc::new(Mutex::new(Some(tx))),
86+
};
87+
let server_handle = tokio::spawn(async move {
88+
let running = server.serve(server_transport).await?;
89+
running.waiting().await?;
90+
anyhow::Ok(())
91+
});
92+
93+
let client = SamplingClient.serve(client_transport).await?;
94+
95+
let result = client
96+
.peer()
97+
.call_tool(CallToolRequestParams::new("sample"))
98+
.await?;
99+
assert_eq!(
100+
result.content.first().unwrap().as_text().unwrap().text,
101+
"ok"
102+
);
103+
104+
let outside = rx.await?;
105+
assert!(matches!(outside, Err(ServiceError::McpError(_))));
106+
107+
client.cancel().await?;
108+
let _ = server_handle.await?;
109+
Ok(())
110+
}

0 commit comments

Comments
 (0)