prism_mcp_rs/transport/
endpoint_pool.rs1use 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#[derive(Debug, Clone)]
13pub struct EndpointPoolConfig {
14 pub failure_threshold: u32,
16 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
36pub 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
114pub 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 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}