Skip to content

Commit d772e4c

Browse files
committed
feat: implement SEP-2260 require server requests to associate with client requests
1 parent 1519707 commit d772e4c

4 files changed

Lines changed: 288 additions & 3 deletions

File tree

crates/rmcp/src/service.rs

Lines changed: 25 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -149,6 +149,15 @@ pub trait ServiceRole: std::fmt::Debug + Send + Sync + 'static + Copy + Clone {
149149
) -> impl Future<Output = ()> + MaybeSendFuture {
150150
async {}
151151
}
152+
153+
#[doc(hidden)]
154+
fn enforce_request_association(
155+
_request: &Self::Req,
156+
_peer_info: Option<&Self::PeerInfo>,
157+
_in_request_handler_scope: bool,
158+
) -> Result<(), ServiceError> {
159+
Ok(())
160+
}
152161
}
153162

154163
pub(crate) fn uses_legacy_lifecycle(
@@ -159,6 +168,14 @@ pub(crate) fn uses_legacy_lifecycle(
159168
&& protocol_version.is_none_or(|version| version < &ProtocolVersion::V_2026_07_28)
160169
}
161170

171+
tokio::task_local! {
172+
pub(crate) static ORIGINATING_REQUEST: RequestId;
173+
}
174+
175+
pub(crate) fn in_request_handler_scope() -> bool {
176+
ORIGINATING_REQUEST.try_with(|_| ()).is_ok()
177+
}
178+
162179
pub type TxJsonRpcMessage<R> =
163180
JsonRpcMessage<<R as ServiceRole>::Req, <R as ServiceRole>::Resp, <R as ServiceRole>::Not>;
164181
pub type RxJsonRpcMessage<R> = JsonRpcMessage<
@@ -725,6 +742,11 @@ impl<R: ServiceRole> Peer<R> {
725742
options: PeerRequestOptions,
726743
subscription_sender: Option<SubscriptionChannel<R::PeerNot>>,
727744
) -> Result<RequestHandle<R>, ServiceError> {
745+
R::enforce_request_association(
746+
&request,
747+
self.peer_info().as_deref(),
748+
in_request_handler_scope(),
749+
)?;
728750
let id = self.request_id_provider.next_request_id();
729751
let progress_token = self.progress_token_provider.next_progress_token();
730752
if let Some(metadata) = self.client_request_metadata.get() {
@@ -1398,9 +1420,10 @@ where
13981420
extensions,
13991421
};
14001422
let current_span = tracing::Span::current();
1423+
let handler_id = id.clone();
14011424
spawn_service_task(async move {
1402-
let result = service
1403-
.handle_request(request, context)
1425+
let result = ORIGINATING_REQUEST
1426+
.scope(handler_id, service.handle_request(request, context))
14041427
.await;
14051428
let response = match result {
14061429
Ok(result) => {

crates/rmcp/src/service/server.rs

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,31 @@ impl ServiceRole for RoleServer {
4949
_ => None,
5050
}
5151
}
52+
53+
fn enforce_request_association(
54+
request: &Self::Req,
55+
peer_info: Option<&Self::PeerInfo>,
56+
in_request_handler_scope: bool,
57+
) -> Result<(), ServiceError> {
58+
let restricted = matches!(
59+
request,
60+
ServerRequest::CreateMessageRequest(_)
61+
| ServerRequest::ListRootsRequest(_)
62+
| ServerRequest::ElicitRequest(_)
63+
);
64+
if !restricted {
65+
return Ok(());
66+
}
67+
let strict =
68+
peer_info.is_some_and(|info| info.protocol_version >= ProtocolVersion::V_2026_07_28);
69+
if strict && !in_request_handler_scope {
70+
return Err(ServiceError::McpError(ErrorData::invalid_request(
71+
"SEP-2260: server-to-client requests must be associated with an originating client request",
72+
None,
73+
)));
74+
}
75+
Ok(())
76+
}
5277
}
5378

5479
/// It represents the error that may occur when serving the server.

crates/rmcp/src/task_manager.rs

Lines changed: 79 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -349,8 +349,11 @@ impl TaskManager {
349349
let future = make_future(context);
350350
let inner = self.inner.clone();
351351
let id_for_task = task_id.clone();
352+
let originating_request = crate::service::ORIGINATING_REQUEST
353+
.try_with(|id| id.clone())
354+
.ok();
352355
let handle = tokio::spawn(async move {
353-
let result = future.await;
356+
let result = run_task_operation(originating_request, future).await;
354357
let mut inner = inner.lock().expect("task manager lock poisoned");
355358
if let Some(entry) = inner.tasks.get_mut(&id_for_task) {
356359
if entry.terminal.is_none() {
@@ -538,6 +541,16 @@ fn unknown_task(task_id: &str) -> McpError {
538541
McpError::invalid_params(format!("unknown task: {task_id}"), None)
539542
}
540543

544+
async fn run_task_operation(
545+
originating_request: Option<crate::model::RequestId>,
546+
future: TaskFuture,
547+
) -> Result<CallToolResult, TaskExit> {
548+
match originating_request {
549+
Some(id) => crate::service::ORIGINATING_REQUEST.scope(id, future).await,
550+
None => future.await,
551+
}
552+
}
553+
541554
fn result_to_object(result: &CallToolResult) -> JsonObject {
542555
match serde_json::to_value(result) {
543556
Ok(serde_json::Value::Object(map)) => map,
@@ -917,4 +930,69 @@ mod tests {
917930
}
918931
panic!("task did not complete after input response");
919932
}
933+
934+
#[tokio::test]
935+
async fn task_operation_reestablishes_request_association_scope() {
936+
use crate::{
937+
model::RequestId,
938+
service::{ORIGINATING_REQUEST, in_request_handler_scope},
939+
};
940+
941+
let manager = TaskManager::new();
942+
let observed = Arc::new(Mutex::new(None::<bool>));
943+
let observed_in_task = observed.clone();
944+
945+
ORIGINATING_REQUEST
946+
.scope(RequestId::Number(7), async {
947+
manager.spawn(TaskOptions::default(), move |_ctx| {
948+
let observed_in_task = observed_in_task.clone();
949+
Box::pin(async move {
950+
*observed_in_task.lock().unwrap() = Some(in_request_handler_scope());
951+
Ok(ok_result("done"))
952+
})
953+
})
954+
})
955+
.await;
956+
957+
for _ in 0..100 {
958+
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
959+
if let Some(scoped) = *observed.lock().unwrap() {
960+
assert!(
961+
scoped,
962+
"task operation must run inside the originating request's association scope"
963+
);
964+
return;
965+
}
966+
}
967+
panic!("task operation did not run");
968+
}
969+
970+
#[tokio::test]
971+
async fn task_operation_without_originating_request_is_unscoped() {
972+
use crate::service::in_request_handler_scope;
973+
974+
let manager = TaskManager::new();
975+
let observed = Arc::new(Mutex::new(None::<bool>));
976+
let observed_in_task = observed.clone();
977+
978+
manager.spawn(TaskOptions::default(), move |_ctx| {
979+
let observed_in_task = observed_in_task.clone();
980+
Box::pin(async move {
981+
*observed_in_task.lock().unwrap() = Some(in_request_handler_scope());
982+
Ok(ok_result("done"))
983+
})
984+
});
985+
986+
for _ in 0..100 {
987+
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
988+
if let Some(scoped) = *observed.lock().unwrap() {
989+
assert!(
990+
!scoped,
991+
"task operation started without an originating request must remain unscoped"
992+
);
993+
return;
994+
}
995+
}
996+
panic!("task operation did not run");
997+
}
920998
}
Lines changed: 159 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,159 @@
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+
CreateMessageRequest, CreateMessageRequestParams, CreateMessageResult, ProtocolVersion,
11+
SamplingMessage, ServerCapabilities, ServerInfo, ServerRequest,
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 slot = self.outside.clone();
34+
35+
let use_generic = request.name == "sample_generic";
36+
tokio::spawn(async move {
37+
let outside = if use_generic {
38+
peer.send_request(ServerRequest::CreateMessageRequest(
39+
CreateMessageRequest::new(CreateMessageRequestParams::new(
40+
vec![SamplingMessage::user_text("standalone-generic")],
41+
16,
42+
)),
43+
))
44+
.await
45+
.map(|_| ())
46+
} else {
47+
peer.create_message(CreateMessageRequestParams::new(
48+
vec![SamplingMessage::user_text("standalone")],
49+
16,
50+
))
51+
.await
52+
.map(|_| ())
53+
};
54+
if let Some(tx) = slot.lock().unwrap().take() {
55+
let _ = tx.send(outside);
56+
}
57+
});
58+
59+
let nested = context
60+
.peer
61+
.create_message(CreateMessageRequestParams::new(
62+
vec![SamplingMessage::user_text("nested")],
63+
16,
64+
))
65+
.await;
66+
nested.map_err(|e| rmcp::ErrorData::internal_error(e.to_string(), None))?;
67+
Ok(CallToolResult::success(vec![ContentBlock::text("ok")]).into())
68+
}
69+
}
70+
71+
#[derive(Clone)]
72+
struct SamplingClient;
73+
74+
impl ClientHandler for SamplingClient {
75+
async fn create_message(
76+
&self,
77+
_params: CreateMessageRequestParams,
78+
_context: RequestContext<RoleClient>,
79+
) -> Result<CreateMessageResult, rmcp::ErrorData> {
80+
Ok(CreateMessageResult::new(
81+
SamplingMessage::assistant_text("pong"),
82+
"test-model".to_string(),
83+
)
84+
.with_stop_reason(CreateMessageResult::STOP_REASON_END_TURN))
85+
}
86+
87+
fn get_info(&self) -> ClientInfo {
88+
let mut info = ClientInfo::default();
89+
info.protocol_version = ProtocolVersion::V_2026_07_28;
90+
info
91+
}
92+
}
93+
94+
#[tokio::test]
95+
async fn nested_sampling_allowed_standalone_rejected() -> anyhow::Result<()> {
96+
let (server_transport, client_transport) = tokio::io::duplex(4096);
97+
let (tx, rx) = oneshot::channel();
98+
let server = SamplingServer {
99+
outside: Arc::new(Mutex::new(Some(tx))),
100+
};
101+
let server_handle = tokio::spawn(async move {
102+
let running = server.serve(server_transport).await?;
103+
running.waiting().await?;
104+
anyhow::Ok(())
105+
});
106+
107+
let client = SamplingClient.serve(client_transport).await?;
108+
109+
let result = client
110+
.peer()
111+
.call_tool(CallToolRequestParams::new("sample"))
112+
.await?;
113+
assert_eq!(
114+
result.content.first().unwrap().as_text().unwrap().text,
115+
"ok"
116+
);
117+
118+
let outside = rx.await?;
119+
assert!(matches!(outside, Err(ServiceError::McpError(_))));
120+
121+
client.cancel().await?;
122+
let _ = server_handle.await?;
123+
Ok(())
124+
}
125+
126+
#[tokio::test]
127+
async fn generic_send_request_bypass_rejected() -> anyhow::Result<()> {
128+
let (server_transport, client_transport) = tokio::io::duplex(4096);
129+
let (tx, rx) = oneshot::channel();
130+
let server = SamplingServer {
131+
outside: Arc::new(Mutex::new(Some(tx))),
132+
};
133+
let server_handle = tokio::spawn(async move {
134+
let running = server.serve(server_transport).await?;
135+
running.waiting().await?;
136+
anyhow::Ok(())
137+
});
138+
139+
let client = SamplingClient.serve(client_transport).await?;
140+
141+
let result = client
142+
.peer()
143+
.call_tool(CallToolRequestParams::new("sample_generic"))
144+
.await?;
145+
assert_eq!(
146+
result.content.first().unwrap().as_text().unwrap().text,
147+
"ok"
148+
);
149+
150+
let outside = rx.await?;
151+
assert!(
152+
matches!(outside, Err(ServiceError::McpError(_))),
153+
"generic send_request must not bypass SEP-2260 enforcement"
154+
);
155+
156+
client.cancel().await?;
157+
let _ = server_handle.await?;
158+
Ok(())
159+
}

0 commit comments

Comments
 (0)