Skip to main content

mcp_proxy/
failover.rs

1//! Backend failover middleware.
2//!
3//! Routes requests to a primary backend, automatically falling over to
4//! secondary backends when the primary returns an error. Multiple failover
5//! backends can be configured per primary, ordered by [`priority`](crate::config::BackendConfig::priority)
6//! (lower values tried first).
7//!
8//! # Configuration
9//!
10//! ```toml
11//! [[backends]]
12//! name = "api"
13//! transport = "http"
14//! url = "http://primary:8080"
15//!
16//! [[backends]]
17//! name = "api-backup"
18//! transport = "http"
19//! url = "http://secondary:8080"
20//! failover_for = "api"
21//! priority = 0            # tried first (default)
22//!
23//! [[backends]]
24//! name = "api-backup-2"
25//! transport = "http"
26//! url = "http://tertiary:8080"
27//! failover_for = "api"
28//! priority = 10           # tried second
29//! ```
30//!
31//! # How it works
32//!
33//! 1. Request arrives targeting `api/search`
34//! 2. Request is forwarded to the `api` backend
35//! 3. If `api` returns an error, the request is retried against `api-backup/search`
36//! 4. If `api-backup` also fails, the request is retried against `api-backup-2/search`
37//! 5. Failover backend tools are hidden from `ListTools` (like canary backends)
38
39use 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/// Resolved failover mapping for a single primary backend.
51#[derive(Debug, Clone)]
52struct FailoverMapping {
53    /// Primary namespace prefix (e.g. "api/").
54    primary_prefix: String,
55    /// Ordered list of failover namespace prefixes (e.g. ["api-backup/", "api-backup-2/"]).
56    /// Tried in order until one succeeds.
57    failover_prefixes: Vec<String>,
58}
59
60/// Tower layer that produces a [`FailoverService`].
61#[derive(Clone)]
62pub struct FailoverLayer {
63    failovers: HashMap<String, Vec<String>>,
64    separator: String,
65}
66
67impl FailoverLayer {
68    /// Create a new failover layer.
69    ///
70    /// `failovers` maps primary backend names to an ordered list of failover
71    /// backend names (sorted by priority, lowest first).
72    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/// Tower service that fails over to secondary backends on primary error.
89///
90/// When a primary backend returns an error, failover backends are tried
91/// in priority order until one succeeds or all have been exhausted.
92#[derive(Clone)]
93pub struct FailoverService<S> {
94    inner: S,
95    mappings: Arc<Vec<FailoverMapping>>,
96}
97
98impl<S> FailoverService<S> {
99    /// Create a new failover service.
100    ///
101    /// `failovers` maps primary backend names to an ordered list of failover
102    /// backend names (sorted by priority, lowest first).
103    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
122/// Rewrite a request's namespace from primary to failover.
123fn 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            // Find if this request targets a primary that has failovers
190            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                    // No failover configured for this request, pass through
201                    return inner.call(req).await;
202                }
203            };
204
205            // Try primary
206            let primary_resp = inner.call(req.clone()).await?;
207
208            // If primary succeeded, return it
209            if primary_resp.inner.is_ok() {
210                return Ok(primary_resp);
211            }
212
213            // Primary failed -- attempt failovers in priority order
214            // TODO: When outlier detection is integrated, check if the primary
215            // backend is ejected and skip directly to failover without waiting
216            // for an error response. This requires sharing ejection state
217            // between the OutlierDetectionService and FailoverService layers.
218            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            // All failovers exhausted, return the last error
247            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        // Create a mock that returns errors for "primary/" calls
319        // but succeeds for "backup/" calls
320        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        // Mock that fails for primary and backup-1, succeeds for backup-2
390        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        // Should get the last failover's error
526        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}