1use std::time::Duration;
31
32use tower_mcp::router::RouterResponse;
33use tower_resilience::retry::{RetryBudgetBuilder, RetryLayer};
34
35use crate::config::RetryConfig;
36
37fn is_retriable_error(code: i32) -> bool {
39 code == -32603 || (-32099..=-32000).contains(&code)
42}
43
44fn is_retriable_response(resp: &RouterResponse) -> bool {
46 match &resp.inner {
47 Err(err) => is_retriable_error(err.code),
48 Ok(_) => false,
49 }
50}
51
52pub 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((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 if let Some(percent) = config.budget_percent {
69 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, 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 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 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 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 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 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 let config = make_config_with_budget(10, 1.0); let layer = build_retry_layer(&config, "test");
199
200 let svc = ErrorMockService;
201 let mut svc = layer.layer(svc);
202
203 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); let _layer = build_retry_layer(&config, "test");
212 assert!(config.budget_percent.is_none());
214 }
215}