1use std::collections::HashMap;
8use std::convert::Infallible;
9use std::future::Future;
10use std::pin::Pin;
11use std::sync::Arc;
12use std::task::{Context, Poll};
13
14use tokio::sync::{Mutex, broadcast};
15use tower::{Layer, Service};
16use tower_mcp::router::{RouterRequest, RouterResponse};
17use tower_mcp_types::protocol::{InputResponses, McpRequest};
18
19#[derive(Clone)]
21pub struct CoalesceLayer;
22
23impl CoalesceLayer {
24 pub fn new() -> Self {
26 Self
27 }
28}
29
30impl Default for CoalesceLayer {
31 fn default() -> Self {
32 Self::new()
33 }
34}
35
36impl<S> Layer<S> for CoalesceLayer {
37 type Service = CoalesceService<S>;
38
39 fn layer(&self, inner: S) -> Self::Service {
40 CoalesceService::new(inner)
41 }
42}
43
44#[derive(Clone)]
46pub struct CoalesceService<S> {
47 inner: S,
48 in_flight: Arc<Mutex<HashMap<String, broadcast::Sender<RouterResponse>>>>,
49}
50
51impl<S> CoalesceService<S> {
52 pub fn new(inner: S) -> Self {
54 Self {
55 inner,
56 in_flight: Arc::new(Mutex::new(HashMap::new())),
57 }
58 }
59}
60
61fn continuation_identity(
62 input_responses: &Option<InputResponses>,
63 request_state: &Option<String>,
64) -> String {
65 serde_json::to_string(&(input_responses, request_state)).unwrap_or_default()
66}
67
68fn coalesce_key(req: &McpRequest) -> Option<String> {
69 match req {
70 McpRequest::CallTool(params) => {
71 let args = serde_json::to_string(¶ms.arguments).unwrap_or_default();
72 let continuation =
73 continuation_identity(¶ms.input_responses, ¶ms.request_state);
74 Some(format!("tool:{}:{args}:{continuation}", params.name))
75 }
76 McpRequest::ReadResource(params) => {
77 let continuation =
78 continuation_identity(¶ms.input_responses, ¶ms.request_state);
79 Some(format!("res:{}:{continuation}", params.uri))
80 }
81 _ => None,
82 }
83}
84
85impl<S> Service<RouterRequest> for CoalesceService<S>
86where
87 S: Service<RouterRequest, Response = RouterResponse, Error = Infallible>
88 + Clone
89 + Send
90 + 'static,
91 S::Future: Send,
92{
93 type Response = RouterResponse;
94 type Error = Infallible;
95 type Future = Pin<Box<dyn Future<Output = Result<RouterResponse, Infallible>> + Send>>;
96
97 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
98 self.inner.poll_ready(cx)
99 }
100
101 fn call(&mut self, req: RouterRequest) -> Self::Future {
102 let Some(key) = coalesce_key(&req.inner) else {
103 let fut = self.inner.call(req);
105 return Box::pin(fut);
106 };
107
108 let in_flight = Arc::clone(&self.in_flight);
109 let mut inner = self.inner.clone();
110 let request_id = req.id.clone();
111
112 Box::pin(async move {
113 {
115 let map = in_flight.lock().await;
116 if let Some(tx) = map.get(&key) {
117 let mut rx = tx.subscribe();
118 drop(map);
119 if let Ok(resp) = rx.recv().await {
121 return Ok(RouterResponse {
122 id: request_id,
123 inner: resp.inner,
124 });
125 }
126 }
128 }
129
130 let (tx, _) = broadcast::channel(1);
132 {
133 let mut map = in_flight.lock().await;
134 map.insert(key.clone(), tx.clone());
135 }
136
137 let result = inner.call(req).await;
138
139 let Ok(ref resp) = result;
141 let _ = tx.send(resp.clone());
142 {
143 let mut map = in_flight.lock().await;
144 map.remove(&key);
145 }
146
147 result
148 })
149 }
150}
151
152#[cfg(test)]
153mod tests {
154 use tower_mcp::protocol::{McpRequest, McpResponse};
155
156 use super::CoalesceService;
157 use crate::test_util::{MockService, call_service};
158
159 #[tokio::test]
160 async fn test_coalesce_passes_through_single_request() {
161 let mock = MockService::with_tools(&["fs/read"]);
162 let mut svc = CoalesceService::new(mock);
163
164 let resp = call_service(
165 &mut svc,
166 McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
167 name: "fs/read".to_string(),
168 arguments: serde_json::json!({}),
169 input_responses: None,
170 request_state: None,
171 meta: None,
172 task: None,
173 }),
174 )
175 .await;
176
177 match resp.inner.unwrap() {
178 McpResponse::CallTool(r) => assert_eq!(r.all_text(), "called: fs/read"),
179 other => panic!("expected CallTool, got: {:?}", other),
180 }
181 }
182
183 #[tokio::test]
184 async fn test_coalesce_non_coalesceable_passes_through() {
185 let mock = MockService::with_tools(&["tool"]);
186 let mut svc = CoalesceService::new(mock);
187
188 let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
189 assert!(resp.inner.is_ok(), "list_tools should pass through");
190 }
191
192 #[tokio::test]
193 async fn test_coalesce_key_includes_arguments() {
194 let key1 =
196 super::coalesce_key(&McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
197 name: "tool".to_string(),
198 arguments: serde_json::json!({"a": 1}),
199 input_responses: None,
200 request_state: None,
201 meta: None,
202 task: None,
203 }));
204 let key2 =
205 super::coalesce_key(&McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
206 name: "tool".to_string(),
207 arguments: serde_json::json!({"a": 2}),
208 input_responses: None,
209 request_state: None,
210 meta: None,
211 task: None,
212 }));
213 assert_ne!(key1, key2, "different args should have different keys");
214 }
215
216 #[tokio::test]
217 async fn test_coalesce_key_same_arguments_produce_same_key() {
218 let key1 =
219 super::coalesce_key(&McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
220 name: "tool".to_string(),
221 arguments: serde_json::json!({"a": 1}),
222 input_responses: None,
223 request_state: None,
224 meta: None,
225 task: None,
226 }));
227 let key2 =
228 super::coalesce_key(&McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
229 name: "tool".to_string(),
230 arguments: serde_json::json!({"a": 1}),
231 input_responses: None,
232 request_state: None,
233 meta: None,
234 task: None,
235 }));
236 assert_eq!(key1, key2, "same tool+args should have the same key");
237 }
238
239 #[tokio::test]
240 async fn test_coalesce_key_read_resource() {
241 let key = super::coalesce_key(&McpRequest::ReadResource(
242 tower_mcp::protocol::ReadResourceParams {
243 uri: "file:///tmp/test.txt".to_string(),
244 input_responses: None,
245 request_state: None,
246 meta: None,
247 },
248 ));
249 assert_eq!(
250 key,
251 Some("res:file:///tmp/test.txt:[null,null]".to_string())
252 );
253 }
254
255 #[tokio::test]
256 async fn test_coalesce_key_includes_continuation_state() {
257 let request = |state: &str| {
258 McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
259 name: "tool".to_string(),
260 arguments: serde_json::json!({"a": 1}),
261 input_responses: None,
262 request_state: Some(state.to_string()),
263 meta: None,
264 task: None,
265 })
266 };
267
268 let key1 = super::coalesce_key(&request("first"));
269 let key2 = super::coalesce_key(&request("second"));
270 assert_ne!(key1, key2, "continuations must not be coalesced together");
271 }
272
273 #[tokio::test]
274 async fn test_coalesce_key_non_coalesceable_returns_none() {
275 let key = super::coalesce_key(&McpRequest::ListTools(Default::default()));
276 assert!(key.is_none(), "ListTools should not be coalesceable");
277
278 let key = super::coalesce_key(&McpRequest::ListResources(Default::default()));
279 assert!(key.is_none(), "ListResources should not be coalesceable");
280 }
281
282 #[tokio::test]
283 async fn test_concurrent_identical_requests_coalesced() {
284 use std::sync::Arc;
285 use std::sync::atomic::{AtomicUsize, Ordering};
286 use tower::Service;
287
288 #[derive(Clone)]
290 struct CountingService {
291 call_count: Arc<AtomicUsize>,
292 }
293
294 impl Service<tower_mcp::router::RouterRequest> for CountingService {
295 type Response = tower_mcp::router::RouterResponse;
296 type Error = std::convert::Infallible;
297 type Future = std::pin::Pin<
298 Box<
299 dyn std::future::Future<
300 Output = Result<
301 tower_mcp::router::RouterResponse,
302 std::convert::Infallible,
303 >,
304 > + Send,
305 >,
306 >;
307
308 fn poll_ready(
309 &mut self,
310 _cx: &mut std::task::Context<'_>,
311 ) -> std::task::Poll<Result<(), Self::Error>> {
312 std::task::Poll::Ready(Ok(()))
313 }
314
315 fn call(&mut self, req: tower_mcp::router::RouterRequest) -> Self::Future {
316 let count = self.call_count.clone();
317 let id = req.id.clone();
318 Box::pin(async move {
319 count.fetch_add(1, Ordering::SeqCst);
320 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
322 Ok(tower_mcp::router::RouterResponse {
323 id,
324 inner: Ok(McpResponse::CallTool(
325 tower_mcp::protocol::CallToolResult::text("result"),
326 )),
327 })
328 })
329 }
330 }
331
332 let call_count = Arc::new(AtomicUsize::new(0));
333 let svc = CountingService {
334 call_count: call_count.clone(),
335 };
336 let coalesce = CoalesceService::new(svc);
337
338 let make_request = || {
339 let mut c = coalesce.clone();
340 let req = tower_mcp::router::RouterRequest {
341 id: tower_mcp::protocol::RequestId::Number(1),
342 inner: McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
343 name: "tool".to_string(),
344 arguments: serde_json::json!({"x": 42}),
345 input_responses: None,
346 request_state: None,
347 meta: None,
348 task: None,
349 }),
350 extensions: tower_mcp::router::Extensions::new(),
351 };
352 async move { c.call(req).await }
353 };
354
355 let (r1, r2, r3) = tokio::join!(make_request(), make_request(), make_request());
357
358 assert!(r1.is_ok());
360 assert!(r2.is_ok());
361 assert!(r3.is_ok());
362
363 let count = call_count.load(Ordering::SeqCst);
367 assert!(
368 count < 3,
369 "expected fewer than 3 backend calls due to coalescing, got {count}"
370 );
371 }
372
373 #[tokio::test]
374 async fn test_different_requests_not_coalesced() {
375 let mock = MockService::with_tools(&["tool"]);
376 let coalesce = CoalesceService::new(mock);
377
378 let mut c1 = coalesce.clone();
380 let req1 = tower_mcp::router::RouterRequest {
381 id: tower_mcp::protocol::RequestId::Number(1),
382 inner: McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
383 name: "tool".to_string(),
384 arguments: serde_json::json!({"x": 1}),
385 input_responses: None,
386 request_state: None,
387 meta: None,
388 task: None,
389 }),
390 extensions: tower_mcp::router::Extensions::new(),
391 };
392
393 let mut c2 = coalesce.clone();
394 let req2 = tower_mcp::router::RouterRequest {
395 id: tower_mcp::protocol::RequestId::Number(2),
396 inner: McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
397 name: "tool".to_string(),
398 arguments: serde_json::json!({"x": 2}),
399 input_responses: None,
400 request_state: None,
401 meta: None,
402 task: None,
403 }),
404 extensions: tower_mcp::router::Extensions::new(),
405 };
406
407 let (r1, r2) = tokio::join!(
408 tower::Service::call(&mut c1, req1),
409 tower::Service::call(&mut c2, req2)
410 );
411
412 assert!(r1.is_ok());
414 assert!(r2.is_ok());
415 }
416
417 #[tokio::test]
418 async fn test_coalesce_with_error_response() {
419 use crate::test_util::ErrorMockService;
420
421 let mock = ErrorMockService;
422 let mut svc = CoalesceService::new(mock);
423
424 let resp = call_service(
425 &mut svc,
426 McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
427 name: "failing_tool".to_string(),
428 arguments: serde_json::json!({}),
429 input_responses: None,
430 request_state: None,
431 meta: None,
432 task: None,
433 }),
434 )
435 .await;
436
437 assert!(
439 resp.inner.is_err(),
440 "error response should propagate through coalesce"
441 );
442 let err = resp.inner.unwrap_err();
443 assert_eq!(err.code, -32603);
444 assert_eq!(err.message, "internal error");
445 }
446}