Skip to main content

mcp_proxy/
coalesce.rs

1//! Request coalescing middleware for the proxy.
2//!
3//! Deduplicates identical in-flight `CallTool` and `ReadResource` requests.
4//! When multiple identical requests arrive concurrently, only one is forwarded
5//! to the backend; all callers receive the same response.
6
7use 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/// Tower layer that produces a [`CoalesceService`].
20#[derive(Clone)]
21pub struct CoalesceLayer;
22
23impl CoalesceLayer {
24    /// Create a new request coalescing layer.
25    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/// Tower service that coalesces identical in-flight requests.
45#[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    /// Create a new request coalescing service wrapping `inner`.
53    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(&params.arguments).unwrap_or_default();
72            let continuation =
73                continuation_identity(&params.input_responses, &params.request_state);
74            Some(format!("tool:{}:{args}:{continuation}", params.name))
75        }
76        McpRequest::ReadResource(params) => {
77            let continuation =
78                continuation_identity(&params.input_responses, &params.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            // Non-coalesceable request, pass through
104            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            // Check if there's already an in-flight request for this key
114            {
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                    // Wait for the in-flight request to complete
120                    if let Ok(resp) = rx.recv().await {
121                        return Ok(RouterResponse {
122                            id: request_id,
123                            inner: resp.inner,
124                        });
125                    }
126                    // Sender dropped (shouldn't happen), fall through to make our own request
127                }
128            }
129
130            // We're the first — register ourselves
131            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            // Broadcast result to any waiters and clean up
140            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        // Different arguments should produce different keys
195        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        // A mock that counts how many times it's actually called
289        #[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                    // Small delay to ensure concurrent requests overlap
321                    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        // Fire 3 identical requests concurrently
356        let (r1, r2, r3) = tokio::join!(make_request(), make_request(), make_request());
357
358        // All should succeed
359        assert!(r1.is_ok());
360        assert!(r2.is_ok());
361        assert!(r3.is_ok());
362
363        // The backend should be called at most twice (the first caller registers,
364        // some others may arrive before the lock is acquired). The key invariant
365        // is that it's called fewer times than the number of requests.
366        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        // Two requests with different arguments
379        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        // Both should succeed independently
413        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        // The error response should pass through correctly
438        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}