1use std::collections::HashMap;
40use std::convert::Infallible;
41use std::future::Future;
42use std::pin::Pin;
43use std::sync::Arc;
44use std::task::{Context, Poll};
45
46use tower::{Layer, Service};
47use tower_mcp::router::{Extensions, RouterRequest, RouterResponse};
48use tower_mcp_types::protocol::{CallToolParams, GetPromptParams, McpRequest, ReadResourceParams};
49
50#[derive(Debug, Clone)]
52struct FailoverMapping {
53 primary_prefix: String,
55 failover_prefixes: Vec<String>,
58}
59
60#[derive(Clone)]
62pub struct FailoverLayer {
63 failovers: HashMap<String, Vec<String>>,
64 separator: String,
65}
66
67impl FailoverLayer {
68 pub fn new(failovers: HashMap<String, Vec<String>>, separator: impl Into<String>) -> Self {
73 Self {
74 failovers,
75 separator: separator.into(),
76 }
77 }
78}
79
80impl<S> Layer<S> for FailoverLayer {
81 type Service = FailoverService<S>;
82
83 fn layer(&self, inner: S) -> Self::Service {
84 FailoverService::new(inner, self.failovers.clone(), &self.separator)
85 }
86}
87
88#[derive(Clone)]
93pub struct FailoverService<S> {
94 inner: S,
95 mappings: Arc<Vec<FailoverMapping>>,
96}
97
98impl<S> FailoverService<S> {
99 pub fn new(inner: S, failovers: HashMap<String, Vec<String>>, separator: &str) -> Self {
104 let mappings = failovers
105 .into_iter()
106 .map(|(primary, failover_names)| FailoverMapping {
107 primary_prefix: format!("{primary}{separator}"),
108 failover_prefixes: failover_names
109 .into_iter()
110 .map(|name| format!("{name}{separator}"))
111 .collect(),
112 })
113 .collect();
114
115 Self {
116 inner,
117 mappings: Arc::new(mappings),
118 }
119 }
120}
121
122fn rewrite_request(req: &McpRequest, primary_prefix: &str, failover_prefix: &str) -> McpRequest {
124 match req {
125 McpRequest::CallTool(params) => {
126 if let Some(local) = params.name.strip_prefix(primary_prefix) {
127 McpRequest::CallTool(CallToolParams {
128 name: format!("{failover_prefix}{local}"),
129 arguments: params.arguments.clone(),
130 input_responses: params.input_responses.clone(),
131 request_state: params.request_state.clone(),
132 meta: params.meta.clone(),
133 task: params.task.clone(),
134 })
135 } else {
136 req.clone()
137 }
138 }
139 McpRequest::ReadResource(params) => {
140 if let Some(local) = params.uri.strip_prefix(primary_prefix) {
141 McpRequest::ReadResource(ReadResourceParams {
142 uri: format!("{failover_prefix}{local}"),
143 input_responses: params.input_responses.clone(),
144 request_state: params.request_state.clone(),
145 meta: params.meta.clone(),
146 })
147 } else {
148 req.clone()
149 }
150 }
151 McpRequest::GetPrompt(params) => {
152 if let Some(local) = params.name.strip_prefix(primary_prefix) {
153 McpRequest::GetPrompt(GetPromptParams {
154 name: format!("{failover_prefix}{local}"),
155 arguments: params.arguments.clone(),
156 input_responses: params.input_responses.clone(),
157 request_state: params.request_state.clone(),
158 meta: params.meta.clone(),
159 })
160 } else {
161 req.clone()
162 }
163 }
164 other => other.clone(),
165 }
166}
167
168impl<S> Service<RouterRequest> for FailoverService<S>
169where
170 S: Service<RouterRequest, Response = RouterResponse, Error = Infallible>
171 + Clone
172 + Send
173 + 'static,
174 S::Future: Send,
175{
176 type Response = RouterResponse;
177 type Error = Infallible;
178 type Future = Pin<Box<dyn Future<Output = Result<RouterResponse, Infallible>> + Send>>;
179
180 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
181 self.inner.poll_ready(cx)
182 }
183
184 fn call(&mut self, req: RouterRequest) -> Self::Future {
185 let mappings = Arc::clone(&self.mappings);
186 let mut inner = self.inner.clone();
187
188 Box::pin(async move {
189 let mapping = mappings.iter().find(|m| match &req.inner {
191 McpRequest::CallTool(p) => p.name.starts_with(&m.primary_prefix),
192 McpRequest::ReadResource(p) => p.uri.starts_with(&m.primary_prefix),
193 McpRequest::GetPrompt(p) => p.name.starts_with(&m.primary_prefix),
194 _ => false,
195 });
196
197 let mapping = match mapping {
198 Some(m) => m.clone(),
199 None => {
200 return inner.call(req).await;
202 }
203 };
204
205 let primary_resp = inner.call(req.clone()).await?;
207
208 if primary_resp.inner.is_ok() {
210 return Ok(primary_resp);
211 }
212
213 let mut last_resp = primary_resp;
219
220 for failover_prefix in &mapping.failover_prefixes {
221 let failover_name = failover_prefix.trim_end_matches('/');
222 tracing::warn!(
223 primary = %mapping.primary_prefix.trim_end_matches('/'),
224 failover = %failover_name,
225 "Backend failed, attempting failover"
226 );
227
228 let failover_request =
229 rewrite_request(&req.inner, &mapping.primary_prefix, failover_prefix);
230
231 let failover_req = RouterRequest {
232 id: req.id.clone(),
233 inner: failover_request,
234 extensions: Extensions::new(),
235 };
236
237 let resp = inner.call(failover_req).await?;
238
239 if resp.inner.is_ok() {
240 return Ok(resp);
241 }
242
243 last_resp = resp;
244 }
245
246 Ok(last_resp)
248 })
249 }
250}
251
252#[cfg(test)]
253mod tests {
254 use tower_mcp::protocol::{McpRequest, McpResponse};
255
256 use super::{FailoverService, rewrite_request};
257 use crate::test_util::{MockService, call_service};
258
259 fn make_failover_svc(mock: MockService) -> FailoverService<MockService> {
260 let failovers = [("primary".to_string(), vec!["backup".to_string()])]
261 .into_iter()
262 .collect();
263 FailoverService::new(mock, failovers, "/")
264 }
265
266 #[test]
267 fn test_rewrite_preserves_continuation_state() {
268 let request = McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
269 name: "primary/tool".to_string(),
270 arguments: serde_json::json!({"q": "test"}),
271 input_responses: Some(Default::default()),
272 request_state: Some("continuation-1".to_string()),
273 meta: None,
274 task: None,
275 });
276
277 let rewritten = rewrite_request(&request, "primary/", "backup/");
278 let McpRequest::CallTool(params) = rewritten else {
279 panic!("expected CallTool");
280 };
281 assert_eq!(params.name, "backup/tool");
282 assert!(params.input_responses.is_some());
283 assert_eq!(params.request_state.as_deref(), Some("continuation-1"));
284 }
285
286 #[tokio::test]
287 async fn test_failover_passes_through_when_no_mapping() {
288 let mock = MockService::with_tools(&["other/tool"]);
289 let mut svc = make_failover_svc(mock);
290
291 let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
292 assert!(resp.inner.is_ok());
293 }
294
295 #[tokio::test]
296 async fn test_failover_passes_through_on_success() {
297 let mock = MockService::with_tools(&["primary/tool", "backup/tool"]);
298 let mut svc = make_failover_svc(mock);
299
300 let resp = call_service(
301 &mut svc,
302 McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
303 name: "primary/tool".to_string(),
304 arguments: serde_json::json!({}),
305 input_responses: None,
306 request_state: None,
307 meta: None,
308 task: None,
309 }),
310 )
311 .await;
312
313 assert!(resp.inner.is_ok(), "successful primary should pass through");
314 }
315
316 #[tokio::test]
317 async fn test_failover_retries_on_primary_error() {
318 use std::convert::Infallible;
321 use std::future::Future;
322 use std::pin::Pin;
323 use std::task::{Context, Poll};
324 use tower::Service;
325 use tower_mcp::protocol::CallToolResult;
326 use tower_mcp::router::{RouterRequest, RouterResponse};
327
328 #[derive(Clone)]
329 struct FailPrimaryMock;
330
331 impl Service<RouterRequest> for FailPrimaryMock {
332 type Response = RouterResponse;
333 type Error = Infallible;
334 type Future = Pin<Box<dyn Future<Output = Result<RouterResponse, Infallible>> + Send>>;
335
336 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
337 Poll::Ready(Ok(()))
338 }
339
340 fn call(&mut self, req: RouterRequest) -> Self::Future {
341 let id = req.id.clone();
342 Box::pin(async move {
343 let inner = match &req.inner {
344 McpRequest::CallTool(params) if params.name.starts_with("primary/") => {
345 Err(tower_mcp_types::JsonRpcError {
346 code: -32603,
347 message: "primary down".to_string(),
348 data: None,
349 })
350 }
351 McpRequest::CallTool(params) if params.name.starts_with("backup/") => {
352 Ok(McpResponse::CallTool(CallToolResult::text("from backup")))
353 }
354 _ => Ok(McpResponse::Pong(Default::default())),
355 };
356 Ok(RouterResponse { id, inner })
357 })
358 }
359 }
360
361 let failovers = [("primary".to_string(), vec!["backup".to_string()])]
362 .into_iter()
363 .collect();
364 let mut svc = FailoverService::new(FailPrimaryMock, failovers, "/");
365
366 let resp = call_service(
367 &mut svc,
368 McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
369 name: "primary/tool".to_string(),
370 arguments: serde_json::json!({}),
371 input_responses: None,
372 request_state: None,
373 meta: None,
374 task: None,
375 }),
376 )
377 .await;
378
379 match resp.inner.unwrap() {
380 McpResponse::CallTool(result) => {
381 assert_eq!(result.all_text(), "from backup");
382 }
383 other => panic!("expected CallTool, got: {:?}", other),
384 }
385 }
386
387 #[tokio::test]
388 async fn test_failover_chain_tries_in_order() {
389 use std::convert::Infallible;
391 use std::future::Future;
392 use std::pin::Pin;
393 use std::task::{Context, Poll};
394 use tower::Service;
395 use tower_mcp::protocol::CallToolResult;
396 use tower_mcp::router::{RouterRequest, RouterResponse};
397
398 #[derive(Clone)]
399 struct ChainMock;
400
401 impl Service<RouterRequest> for ChainMock {
402 type Response = RouterResponse;
403 type Error = Infallible;
404 type Future = Pin<Box<dyn Future<Output = Result<RouterResponse, Infallible>> + Send>>;
405
406 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
407 Poll::Ready(Ok(()))
408 }
409
410 fn call(&mut self, req: RouterRequest) -> Self::Future {
411 let id = req.id.clone();
412 Box::pin(async move {
413 let inner = match &req.inner {
414 McpRequest::CallTool(params) if params.name.starts_with("primary/") => {
415 Err(tower_mcp_types::JsonRpcError {
416 code: -32603,
417 message: "primary down".to_string(),
418 data: None,
419 })
420 }
421 McpRequest::CallTool(params) if params.name.starts_with("backup-1/") => {
422 Err(tower_mcp_types::JsonRpcError {
423 code: -32603,
424 message: "backup-1 down".to_string(),
425 data: None,
426 })
427 }
428 McpRequest::CallTool(params) if params.name.starts_with("backup-2/") => {
429 Ok(McpResponse::CallTool(CallToolResult::text("from backup-2")))
430 }
431 _ => Ok(McpResponse::Pong(Default::default())),
432 };
433 Ok(RouterResponse { id, inner })
434 })
435 }
436 }
437
438 let failovers = [(
439 "primary".to_string(),
440 vec!["backup-1".to_string(), "backup-2".to_string()],
441 )]
442 .into_iter()
443 .collect();
444 let mut svc = FailoverService::new(ChainMock, failovers, "/");
445
446 let resp = call_service(
447 &mut svc,
448 McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
449 name: "primary/tool".to_string(),
450 arguments: serde_json::json!({}),
451 input_responses: None,
452 request_state: None,
453 meta: None,
454 task: None,
455 }),
456 )
457 .await;
458
459 match resp.inner.unwrap() {
460 McpResponse::CallTool(result) => {
461 assert_eq!(result.all_text(), "from backup-2");
462 }
463 other => panic!("expected CallTool, got: {:?}", other),
464 }
465 }
466
467 #[tokio::test]
468 async fn test_failover_chain_all_fail_returns_last_error() {
469 use std::convert::Infallible;
470 use std::future::Future;
471 use std::pin::Pin;
472 use std::task::{Context, Poll};
473 use tower::Service;
474 use tower_mcp::router::{RouterRequest, RouterResponse};
475
476 #[derive(Clone)]
477 struct AllFailMock;
478
479 impl Service<RouterRequest> for AllFailMock {
480 type Response = RouterResponse;
481 type Error = Infallible;
482 type Future = Pin<Box<dyn Future<Output = Result<RouterResponse, Infallible>> + Send>>;
483
484 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
485 Poll::Ready(Ok(()))
486 }
487
488 fn call(&mut self, req: RouterRequest) -> Self::Future {
489 let id = req.id.clone();
490 Box::pin(async move {
491 let inner = match &req.inner {
492 McpRequest::CallTool(params) => Err(tower_mcp_types::JsonRpcError {
493 code: -32603,
494 message: format!("{} down", params.name),
495 data: None,
496 }),
497 _ => Ok(McpResponse::Pong(Default::default())),
498 };
499 Ok(RouterResponse { id, inner })
500 })
501 }
502 }
503
504 let failovers = [(
505 "primary".to_string(),
506 vec!["backup-1".to_string(), "backup-2".to_string()],
507 )]
508 .into_iter()
509 .collect();
510 let mut svc = FailoverService::new(AllFailMock, failovers, "/");
511
512 let resp = call_service(
513 &mut svc,
514 McpRequest::CallTool(tower_mcp::protocol::CallToolParams {
515 name: "primary/tool".to_string(),
516 arguments: serde_json::json!({}),
517 input_responses: None,
518 request_state: None,
519 meta: None,
520 task: None,
521 }),
522 )
523 .await;
524
525 let err = resp.inner.unwrap_err();
527 assert!(
528 err.message.contains("backup-2"),
529 "expected last failover error, got: {}",
530 err.message
531 );
532 }
533}