Skip to main content

mcp_proxy/
retry.rs

1//! Retry middleware for per-backend request retries with exponential backoff.
2//!
3//! Uses tower-resilience's [`RetryLayer`] with a response-based predicate to
4//! retry MCP error responses with transient error codes (internal errors,
5//! timeouts). Tool-not-found and other client errors are not retried.
6//!
7//! # Retry Budget
8//!
9//! When `budget_percent` is configured, a token bucket budget limits retries.
10//! The budget is sized as a percentage of expected request volume, preventing
11//! retry storms during widespread backend failures. A `min_retries_per_sec`
12//! floor ensures low-traffic backends can still retry.
13//!
14//! # Configuration
15//!
16//! ```toml
17//! [[backends]]
18//! name = "flaky-api"
19//! transport = "http"
20//! url = "http://localhost:8080"
21//!
22//! [backends.retry]
23//! max_retries = 3
24//! initial_backoff_ms = 100
25//! max_backoff_ms = 5000
26//! budget_percent = 20.0    # max 20% of requests can be retries
27//! min_retries_per_sec = 10 # floor for low-traffic backends
28//! ```
29
30use std::time::Duration;
31
32use tower_mcp::router::RouterResponse;
33use tower_resilience::retry::{RetryBudgetBuilder, RetryLayer};
34
35use crate::config::RetryConfig;
36
37/// Returns true if the MCP error code indicates a transient/retriable error.
38fn is_retriable_error(code: i32) -> bool {
39    // JSON-RPC internal error (-32603) and server errors (-32000 to -32099)
40    // are potentially transient. Method not found, invalid params, etc. are not.
41    code == -32603 || (-32099..=-32000).contains(&code)
42}
43
44/// Returns true if a `RouterResponse` contains a retriable MCP error.
45fn is_retriable_response(resp: &RouterResponse) -> bool {
46    match &resp.inner {
47        Err(err) => is_retriable_error(err.code),
48        Ok(_) => false,
49    }
50}
51
52/// Build a tower-resilience [`RetryLayer`] from our [`RetryConfig`].
53///
54/// The layer uses a response-based predicate (since MCP services use
55/// `Error = Infallible` and encode errors inside `RouterResponse`).
56pub fn build_retry_layer(
57    config: &RetryConfig,
58    backend_name: &str,
59) -> RetryLayer<tower_mcp::router::RouterRequest, RouterResponse, std::convert::Infallible> {
60    let mut builder = RetryLayer::builder()
61        // max_attempts includes the initial attempt, so max_retries + 1
62        .max_attempts((config.max_retries + 1) as usize)
63        .exponential_backoff(Duration::from_millis(config.initial_backoff_ms))
64        .retry_on_response(is_retriable_response)
65        .name(format!("retry-{backend_name}"));
66
67    // Configure budget if percent-based limiting is enabled
68    if let Some(percent) = config.budget_percent {
69        // Map budget_percent to a token bucket: scale tokens to approximate
70        // the percentage model. We use min_retries_per_sec as the refill rate
71        // and size the bucket relative to expected request volume.
72        //
73        // For a 20% budget at 100 req/s, we'd want ~20 retries/s capacity.
74        // The token bucket's initial_tokens acts as burst capacity.
75        let min_per_sec = config.min_retries_per_sec as f64;
76        let max_tokens = ((percent / 100.0) * 1000.0).max(min_per_sec * 10.0) as usize;
77
78        let budget = RetryBudgetBuilder::new()
79            .token_bucket()
80            .tokens_per_second(min_per_sec)
81            .max_tokens(max_tokens.max(1))
82            .initial_tokens(max_tokens.max(1))
83            .build();
84
85        builder = builder.budget(budget);
86    }
87
88    builder.build()
89}
90
91#[cfg(test)]
92mod tests {
93    use super::*;
94    use crate::config::RetryConfig;
95    use crate::test_util::{ErrorMockService, MockService, call_service};
96    use tower::Layer;
97    use tower_mcp_types::protocol::McpRequest;
98
99    fn make_config(max_retries: u32) -> RetryConfig {
100        RetryConfig {
101            max_retries,
102            initial_backoff_ms: 1, // fast for tests
103            max_backoff_ms: 10,
104            budget_percent: None,
105            min_retries_per_sec: 10,
106        }
107    }
108
109    fn make_config_with_budget(max_retries: u32, budget_percent: f64) -> RetryConfig {
110        RetryConfig {
111            max_retries,
112            initial_backoff_ms: 1,
113            max_backoff_ms: 10,
114            budget_percent: Some(budget_percent),
115            min_retries_per_sec: 0,
116        }
117    }
118
119    #[tokio::test]
120    async fn test_retries_internal_error() {
121        // ErrorMockService always returns a -32603 error
122        let svc = ErrorMockService;
123        let layer = build_retry_layer(&make_config(3), "test");
124        let mut svc = layer.layer(svc);
125
126        let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
127        // Should still get an error (all attempts fail), but it should have retried
128        assert!(resp.inner.is_err());
129    }
130
131    #[tokio::test]
132    async fn test_does_not_retry_success() {
133        let svc = MockService::with_tools(&["tool1"]);
134        let layer = build_retry_layer(&make_config(3), "test");
135        let mut svc = layer.layer(svc);
136
137        let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
138        assert!(resp.inner.is_ok());
139    }
140
141    #[tokio::test]
142    async fn test_response_predicate_matches_transient_errors() {
143        // Verify the predicate function directly
144        use tower_mcp::protocol::RequestId;
145        use tower_mcp_types::JsonRpcError;
146
147        let transient = RouterResponse {
148            id: RequestId::Number(1),
149            inner: Err(JsonRpcError {
150                code: -32603,
151                message: "internal error".to_string(),
152                data: None,
153            }),
154        };
155        assert!(is_retriable_response(&transient));
156
157        let server_err = RouterResponse {
158            id: RequestId::Number(1),
159            inner: Err(JsonRpcError {
160                code: -32000,
161                message: "server error".to_string(),
162                data: None,
163            }),
164        };
165        assert!(is_retriable_response(&server_err));
166
167        // Client error -- should NOT retry
168        let client_err = RouterResponse {
169            id: RequestId::Number(1),
170            inner: Err(JsonRpcError {
171                code: -32601,
172                message: "method not found".to_string(),
173                data: None,
174            }),
175        };
176        assert!(!is_retriable_response(&client_err));
177
178        // Success -- should NOT retry
179        let success = RouterResponse {
180            id: RequestId::Number(1),
181            inner: Ok(tower_mcp_types::protocol::McpResponse::ListTools(
182                tower_mcp_types::protocol::ListToolsResult {
183                    tools: vec![],
184                    next_cursor: None,
185                    ttl_ms: None,
186                    cache_scope: None,
187                    meta: None,
188                },
189            )),
190        };
191        assert!(!is_retriable_response(&success));
192    }
193
194    #[tokio::test]
195    async fn test_budget_limits_retries() {
196        // With a very small budget, retries should be limited
197        let config = make_config_with_budget(10, 1.0); // tiny budget
198        let layer = build_retry_layer(&config, "test");
199
200        let svc = ErrorMockService;
201        let mut svc = layer.layer(svc);
202
203        // Should still eventually return (budget exhaustion returns the response)
204        let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
205        assert!(resp.inner.is_err());
206    }
207
208    #[tokio::test]
209    async fn test_no_budget_allows_all_retries() {
210        let config = make_config(2); // No budget, 2 retries
211        let _layer = build_retry_layer(&config, "test");
212        // Just verify it builds without a budget
213        assert!(config.budget_percent.is_none());
214    }
215}