Skip to main content

prism_mcp_rs/transport/
endpoint_pool.rs

1//! Multi-endpoint transport with conservative, idempotency-aware failover.
2
3use async_trait::async_trait;
4use std::time::{Duration, Instant};
5
6use crate::core::error::{McpError, McpResult};
7use crate::protocol::methods;
8use crate::protocol::types::{JsonRpcNotification, JsonRpcRequest, JsonRpcResponse};
9use crate::transport::traits::Transport;
10
11/// Failover and circuit configuration for an endpoint pool.
12#[derive(Debug, Clone)]
13pub struct EndpointPoolConfig {
14    /// Consecutive recoverable failures before an endpoint is temporarily open.
15    pub failure_threshold: u32,
16    /// How long an open endpoint is excluded from selection.
17    pub cooldown: Duration,
18}
19
20impl Default for EndpointPoolConfig {
21    fn default() -> Self {
22        Self {
23            failure_threshold: 3,
24            cooldown: Duration::from_secs(30),
25        }
26    }
27}
28
29struct Endpoint {
30    name: String,
31    transport: Box<dyn Transport>,
32    consecutive_failures: u32,
33    open_until: Option<Instant>,
34}
35
36/// A round-robin endpoint pool with per-endpoint circuit state.
37///
38/// Read-only MCP methods may fail over to another endpoint. Mutating methods,
39/// including `tools/call`, are attempted once unless the request contains
40/// `params._meta.idempotencyKey`.
41pub struct EndpointPoolTransport {
42    endpoints: Vec<Endpoint>,
43    config: EndpointPoolConfig,
44    cursor: usize,
45}
46
47impl EndpointPoolTransport {
48    pub fn new(config: EndpointPoolConfig) -> Self {
49        Self {
50            endpoints: Vec::new(),
51            config,
52            cursor: 0,
53        }
54    }
55
56    pub fn add_endpoint(
57        mut self,
58        name: impl Into<String>,
59        transport: impl Transport + 'static,
60    ) -> Self {
61        self.endpoints.push(Endpoint {
62            name: name.into(),
63            transport: Box::new(transport),
64            consecutive_failures: 0,
65            open_until: None,
66        });
67        self
68    }
69
70    pub fn endpoint_count(&self) -> usize {
71        self.endpoints.len()
72    }
73
74    fn selectable_indices(&mut self) -> Vec<usize> {
75        let now = Instant::now();
76        for endpoint in &mut self.endpoints {
77            if endpoint.open_until.is_some_and(|until| until <= now) {
78                endpoint.open_until = None;
79                endpoint.consecutive_failures = 0;
80            }
81        }
82
83        let len = self.endpoints.len();
84        if len == 0 {
85            return Vec::new();
86        }
87        let start = self.cursor % len;
88        self.cursor = (self.cursor + 1) % len;
89        (0..len)
90            .map(|offset| (start + offset) % len)
91            .filter(|index| self.endpoints[*index].open_until.is_none())
92            .collect()
93    }
94
95    fn record_success(&mut self, index: usize) {
96        self.endpoints[index].consecutive_failures = 0;
97        self.endpoints[index].open_until = None;
98    }
99
100    fn record_failure(&mut self, index: usize) {
101        let endpoint = &mut self.endpoints[index];
102        endpoint.consecutive_failures = endpoint.consecutive_failures.saturating_add(1);
103        if endpoint.consecutive_failures >= self.config.failure_threshold.max(1) {
104            endpoint.open_until = Some(Instant::now() + self.config.cooldown);
105            tracing::warn!(
106                endpoint = %endpoint.name,
107                cooldown_ms = self.config.cooldown.as_millis() as u64,
108                "endpoint circuit opened"
109            );
110        }
111    }
112}
113
114/// Whether a request is safe to replay on another endpoint.
115pub fn is_request_idempotent(request: &JsonRpcRequest) -> bool {
116    let naturally_idempotent = matches!(
117        request.method.as_str(),
118        methods::PING
119            | methods::TOOLS_LIST
120            | methods::RESOURCES_LIST
121            | methods::RESOURCES_READ
122            | methods::RESOURCES_TEMPLATES_LIST
123            | methods::PROMPTS_LIST
124            | methods::PROMPTS_GET
125            | methods::RPC_DISCOVER
126    );
127    naturally_idempotent
128        || request
129            .params
130            .as_ref()
131            .and_then(|params| params.get("_meta"))
132            .and_then(|meta| meta.get("idempotencyKey"))
133            .and_then(serde_json::Value::as_str)
134            .is_some_and(|key| !key.is_empty())
135}
136
137#[async_trait]
138impl Transport for EndpointPoolTransport {
139    async fn send_request(&mut self, request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
140        let indices = self.selectable_indices();
141        if indices.is_empty() {
142            return Err(McpError::Connection(
143                "endpoint pool has no available endpoints".to_string(),
144            ));
145        }
146
147        let max_attempts = if is_request_idempotent(&request) {
148            indices.len()
149        } else {
150            1
151        };
152        let mut last_error = None;
153
154        for index in indices.into_iter().take(max_attempts) {
155            match self.endpoints[index]
156                .transport
157                .send_request(request.clone())
158                .await
159            {
160                Ok(response) => {
161                    self.record_success(index);
162                    return Ok(response);
163                }
164                Err(error) => {
165                    let recoverable = error.is_recoverable();
166                    if recoverable {
167                        self.record_failure(index);
168                    } else {
169                        // The endpoint responded; a protocol, validation, or
170                        // authorization rejection is not endpoint unhealthiness.
171                        self.record_success(index);
172                    }
173                    last_error = Some(error);
174                    if !recoverable {
175                        break;
176                    }
177                }
178            }
179        }
180
181        Err(last_error
182            .unwrap_or_else(|| McpError::Connection("all selected endpoints failed".to_string())))
183    }
184
185    async fn send_notification(&mut self, notification: JsonRpcNotification) -> McpResult<()> {
186        let index = self
187            .selectable_indices()
188            .into_iter()
189            .next()
190            .ok_or_else(|| McpError::Connection("no available endpoint".to_string()))?;
191        let result = self.endpoints[index]
192            .transport
193            .send_notification(notification)
194            .await;
195        match &result {
196            Ok(()) => self.record_success(index),
197            Err(error) if error.is_recoverable() => self.record_failure(index),
198            Err(_) => self.record_success(index),
199        }
200        result
201    }
202
203    async fn receive_notification(&mut self) -> McpResult<Option<JsonRpcNotification>> {
204        for index in self.selectable_indices() {
205            match self.endpoints[index].transport.receive_notification().await {
206                Ok(Some(notification)) => return Ok(Some(notification)),
207                Ok(None) => {}
208                Err(error) if error.is_recoverable() => self.record_failure(index),
209                Err(error) => return Err(error),
210            }
211        }
212        Ok(None)
213    }
214
215    async fn close(&mut self) -> McpResult<()> {
216        let mut first_error = None;
217        for endpoint in &mut self.endpoints {
218            if let Err(error) = endpoint.transport.close().await {
219                first_error.get_or_insert(error);
220            }
221        }
222        first_error.map_or(Ok(()), Err)
223    }
224
225    fn is_connected(&self) -> bool {
226        self.endpoints
227            .iter()
228            .any(|endpoint| endpoint.open_until.is_none() && endpoint.transport.is_connected())
229    }
230
231    fn connection_info(&self) -> String {
232        format!("endpoint pool ({} endpoints)", self.endpoints.len())
233    }
234}
235
236#[cfg(test)]
237mod tests {
238    use super::*;
239    use serde_json::{json, Value};
240    use std::collections::VecDeque;
241    use std::sync::{Arc, Mutex};
242
243    struct MockTransport {
244        calls: Arc<Mutex<usize>>,
245        results: VecDeque<McpResult<JsonRpcResponse>>,
246    }
247
248    #[async_trait]
249    impl Transport for MockTransport {
250        async fn send_request(&mut self, _request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
251            *self.calls.lock().unwrap() += 1;
252            self.results.pop_front().unwrap()
253        }
254
255        async fn send_notification(&mut self, _notification: JsonRpcNotification) -> McpResult<()> {
256            Ok(())
257        }
258
259        async fn receive_notification(&mut self) -> McpResult<Option<JsonRpcNotification>> {
260            Ok(None)
261        }
262
263        async fn close(&mut self) -> McpResult<()> {
264            Ok(())
265        }
266    }
267
268    fn request(method: &str, params: Option<Value>) -> JsonRpcRequest {
269        JsonRpcRequest::new(json!(1), method.to_string(), params).unwrap()
270    }
271
272    #[tokio::test]
273    async fn reads_fail_over_to_the_next_endpoint() {
274        let first_calls = Arc::new(Mutex::new(0));
275        let second_calls = Arc::new(Mutex::new(0));
276        let first = MockTransport {
277            calls: first_calls.clone(),
278            results: VecDeque::from([Err(McpError::connection("down"))]),
279        };
280        let second = MockTransport {
281            calls: second_calls.clone(),
282            results: VecDeque::from([Ok(JsonRpcResponse::success(json!(1), json!([])).unwrap())]),
283        };
284        let mut pool = EndpointPoolTransport::new(EndpointPoolConfig::default())
285            .add_endpoint("first", first)
286            .add_endpoint("second", second);
287
288        assert!(pool
289            .send_request(request(methods::TOOLS_LIST, None))
290            .await
291            .is_ok());
292        assert_eq!(*first_calls.lock().unwrap(), 1);
293        assert_eq!(*second_calls.lock().unwrap(), 1);
294    }
295
296    #[tokio::test]
297    async fn unkeyed_tool_calls_are_never_replayed() {
298        let first_calls = Arc::new(Mutex::new(0));
299        let second_calls = Arc::new(Mutex::new(0));
300        let first = MockTransport {
301            calls: first_calls.clone(),
302            results: VecDeque::from([Err(McpError::connection("response lost"))]),
303        };
304        let second = MockTransport {
305            calls: second_calls.clone(),
306            results: VecDeque::from([Ok(JsonRpcResponse::success(json!(1), json!({})).unwrap())]),
307        };
308        let mut pool = EndpointPoolTransport::new(EndpointPoolConfig::default())
309            .add_endpoint("first", first)
310            .add_endpoint("second", second);
311
312        assert!(pool
313            .send_request(request(
314                methods::TOOLS_CALL,
315                Some(json!({"name": "charge", "arguments": {}}))
316            ))
317            .await
318            .is_err());
319        assert_eq!(*first_calls.lock().unwrap(), 1);
320        assert_eq!(*second_calls.lock().unwrap(), 0);
321    }
322
323    #[test]
324    fn tool_call_requires_an_idempotency_key_for_replay() {
325        assert!(!is_request_idempotent(&request(
326            methods::TOOLS_CALL,
327            Some(json!({"name": "charge"}))
328        )));
329        assert!(is_request_idempotent(&request(
330            methods::TOOLS_CALL,
331            Some(json!({
332                "name": "charge",
333                "_meta": {"idempotencyKey": "operation-42"}
334            }))
335        )));
336    }
337}