1use std::collections::HashMap;
36use std::convert::Infallible;
37use std::future::Future;
38use std::pin::Pin;
39use std::sync::Arc;
40use std::sync::atomic::{AtomicU64, Ordering};
41use std::task::{Context, Poll};
42
43use tower::{Layer, Service};
44use tower_mcp::router::{Extensions, RouterRequest, RouterResponse};
45use tower_mcp_types::protocol::{CallToolParams, GetPromptParams, McpRequest, ReadResourceParams};
46
47#[derive(Clone)]
49pub struct CanaryLayer {
50 canaries: HashMap<String, (String, u32, u32)>,
51 separator: String,
52}
53
54impl CanaryLayer {
55 pub fn new(
59 canaries: HashMap<String, (String, u32, u32)>,
60 separator: impl Into<String>,
61 ) -> Self {
62 Self {
63 canaries,
64 separator: separator.into(),
65 }
66 }
67}
68
69impl<S> Layer<S> for CanaryLayer {
70 type Service = CanaryService<S>;
71
72 fn layer(&self, inner: S) -> Self::Service {
73 CanaryService::new(inner, self.canaries.clone(), &self.separator)
74 }
75}
76
77#[derive(Debug, Clone)]
79struct CanaryMapping {
80 primary_prefix: String,
82 canary_prefix: String,
84 primary_weight: u32,
86 total_weight: u32,
88 counter: Arc<AtomicU64>,
90}
91
92#[derive(Clone)]
97pub struct CanaryService<S> {
98 inner: S,
99 mappings: Arc<Vec<CanaryMapping>>,
100}
101
102impl<S> CanaryService<S> {
103 pub fn new(inner: S, canaries: HashMap<String, (String, u32, u32)>, separator: &str) -> Self {
108 let mappings = canaries
109 .into_iter()
110 .map(
111 |(primary, (canary, primary_weight, canary_weight))| CanaryMapping {
112 primary_prefix: format!("{primary}{separator}"),
113 canary_prefix: format!("{canary}{separator}"),
114 primary_weight,
115 total_weight: primary_weight + canary_weight,
116 counter: Arc::new(AtomicU64::new(0)),
117 },
118 )
119 .collect();
120
121 Self {
122 inner,
123 mappings: Arc::new(mappings),
124 }
125 }
126}
127
128fn find_canary<'a>(name: &str, mappings: &'a [CanaryMapping]) -> Option<&'a CanaryMapping> {
130 mappings
131 .iter()
132 .find(|m| name.starts_with(&m.primary_prefix))
133}
134
135fn should_route_to_canary(mapping: &CanaryMapping) -> bool {
137 let count = mapping.counter.fetch_add(1, Ordering::Relaxed);
138 let position = count % mapping.total_weight as u64;
139 position >= mapping.primary_weight as u64
141}
142
143fn rewrite_to_canary(req: RouterRequest, mapping: &CanaryMapping) -> RouterRequest {
145 let new_inner = match req.inner {
146 McpRequest::CallTool(params) if params.name.starts_with(&mapping.primary_prefix) => {
147 let suffix = ¶ms.name[mapping.primary_prefix.len()..];
148 McpRequest::CallTool(CallToolParams {
149 name: format!("{}{suffix}", mapping.canary_prefix),
150 arguments: params.arguments,
151 input_responses: params.input_responses,
152 request_state: params.request_state,
153 meta: params.meta,
154 task: params.task,
155 })
156 }
157 McpRequest::ReadResource(params) if params.uri.starts_with(&mapping.primary_prefix) => {
158 let suffix = ¶ms.uri[mapping.primary_prefix.len()..];
159 McpRequest::ReadResource(ReadResourceParams {
160 uri: format!("{}{suffix}", mapping.canary_prefix),
161 input_responses: params.input_responses,
162 request_state: params.request_state,
163 meta: params.meta,
164 })
165 }
166 McpRequest::GetPrompt(params) if params.name.starts_with(&mapping.primary_prefix) => {
167 let suffix = ¶ms.name[mapping.primary_prefix.len()..];
168 McpRequest::GetPrompt(GetPromptParams {
169 name: format!("{}{suffix}", mapping.canary_prefix),
170 arguments: params.arguments,
171 input_responses: params.input_responses,
172 request_state: params.request_state,
173 meta: params.meta,
174 })
175 }
176 other => other,
177 };
178
179 RouterRequest {
180 id: req.id,
181 inner: new_inner,
182 extensions: Extensions::new(),
183 }
184}
185
186fn request_name(req: &McpRequest) -> Option<&str> {
188 match req {
189 McpRequest::CallTool(params) => Some(¶ms.name),
190 McpRequest::ReadResource(params) => Some(¶ms.uri),
191 McpRequest::GetPrompt(params) => Some(¶ms.name),
192 _ => None,
193 }
194}
195
196impl<S> Service<RouterRequest> for CanaryService<S>
197where
198 S: Service<RouterRequest, Response = RouterResponse, Error = Infallible>
199 + Clone
200 + Send
201 + 'static,
202 S::Future: Send,
203{
204 type Response = RouterResponse;
205 type Error = Infallible;
206 type Future = Pin<Box<dyn Future<Output = Result<RouterResponse, Infallible>> + Send>>;
207
208 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
209 self.inner.poll_ready(cx)
210 }
211
212 fn call(&mut self, req: RouterRequest) -> Self::Future {
213 let should_canary = request_name(&req.inner)
215 .and_then(|name| find_canary(name, &self.mappings))
216 .filter(|mapping| should_route_to_canary(mapping))
217 .cloned();
218
219 let req = if let Some(ref mapping) = should_canary {
220 tracing::debug!(
221 primary = %mapping.primary_prefix,
222 canary = %mapping.canary_prefix,
223 "Routing request to canary backend"
224 );
225 rewrite_to_canary(req, mapping)
226 } else {
227 req
228 };
229
230 let fut = self.inner.call(req);
231 Box::pin(fut)
232 }
233}
234
235#[cfg(test)]
236mod tests {
237 use super::*;
238 use crate::test_util::{MockService, call_service};
239 use tower_mcp::protocol::RequestId;
240
241 fn make_canaries(
242 primary: &str,
243 canary: &str,
244 primary_weight: u32,
245 canary_weight: u32,
246 ) -> HashMap<String, (String, u32, u32)> {
247 let mut m = HashMap::new();
248 m.insert(
249 primary.to_string(),
250 (canary.to_string(), primary_weight, canary_weight),
251 );
252 m
253 }
254
255 #[test]
256 fn test_find_canary_match() {
257 let mappings = vec![CanaryMapping {
258 primary_prefix: "api/".to_string(),
259 canary_prefix: "api-canary/".to_string(),
260 primary_weight: 90,
261 total_weight: 100,
262 counter: Arc::new(AtomicU64::new(0)),
263 }];
264 assert!(find_canary("api/search", &mappings).is_some());
265 assert!(find_canary("other/search", &mappings).is_none());
266 }
267
268 #[test]
269 fn test_should_route_to_canary_weights() {
270 let mapping = CanaryMapping {
271 primary_prefix: "api/".to_string(),
272 canary_prefix: "api-canary/".to_string(),
273 primary_weight: 90,
274 total_weight: 100,
275 counter: Arc::new(AtomicU64::new(0)),
276 };
277
278 let canary_count: u32 = (0..100)
280 .filter(|_| should_route_to_canary(&mapping))
281 .count() as u32;
282 assert_eq!(canary_count, 10);
283 }
284
285 #[test]
286 fn test_should_route_to_canary_50_50() {
287 let mapping = CanaryMapping {
288 primary_prefix: "api/".to_string(),
289 canary_prefix: "api-canary/".to_string(),
290 primary_weight: 50,
291 total_weight: 100,
292 counter: Arc::new(AtomicU64::new(0)),
293 };
294
295 let canary_count: u32 = (0..100)
296 .filter(|_| should_route_to_canary(&mapping))
297 .count() as u32;
298 assert_eq!(canary_count, 50);
299 }
300
301 #[test]
302 fn test_rewrite_to_canary_call_tool() {
303 let mapping = CanaryMapping {
304 primary_prefix: "api/".to_string(),
305 canary_prefix: "api-canary/".to_string(),
306 primary_weight: 90,
307 total_weight: 100,
308 counter: Arc::new(AtomicU64::new(0)),
309 };
310
311 let req = RouterRequest {
312 id: RequestId::Number(1),
313 inner: McpRequest::CallTool(CallToolParams {
314 name: "api/search".to_string(),
315 arguments: serde_json::json!({"q": "test"}),
316 input_responses: Some(Default::default()),
317 request_state: Some("continuation-1".to_string()),
318 meta: None,
319 task: None,
320 }),
321 extensions: Extensions::new(),
322 };
323
324 let rewritten = rewrite_to_canary(req, &mapping);
325 match &rewritten.inner {
326 McpRequest::CallTool(params) => {
327 assert_eq!(params.name, "api-canary/search");
328 assert_eq!(params.arguments, serde_json::json!({"q": "test"}));
329 assert!(params.input_responses.is_some());
330 assert_eq!(params.request_state.as_deref(), Some("continuation-1"));
331 }
332 _ => panic!("expected CallTool"),
333 }
334 }
335
336 #[test]
337 fn test_rewrite_to_canary_read_resource() {
338 let mapping = CanaryMapping {
339 primary_prefix: "api/".to_string(),
340 canary_prefix: "api-canary/".to_string(),
341 primary_weight: 90,
342 total_weight: 100,
343 counter: Arc::new(AtomicU64::new(0)),
344 };
345
346 let req = RouterRequest {
347 id: RequestId::Number(1),
348 inner: McpRequest::ReadResource(ReadResourceParams {
349 uri: "api/docs/readme".to_string(),
350 input_responses: None,
351 request_state: None,
352 meta: None,
353 }),
354 extensions: Extensions::new(),
355 };
356
357 let rewritten = rewrite_to_canary(req, &mapping);
358 match &rewritten.inner {
359 McpRequest::ReadResource(params) => {
360 assert_eq!(params.uri, "api-canary/docs/readme");
361 }
362 _ => panic!("expected ReadResource"),
363 }
364 }
365
366 #[test]
367 fn test_rewrite_leaves_non_matching_unchanged() {
368 let mapping = CanaryMapping {
369 primary_prefix: "api/".to_string(),
370 canary_prefix: "api-canary/".to_string(),
371 primary_weight: 90,
372 total_weight: 100,
373 counter: Arc::new(AtomicU64::new(0)),
374 };
375
376 let req = RouterRequest {
377 id: RequestId::Number(1),
378 inner: McpRequest::ListTools(Default::default()),
379 extensions: Extensions::new(),
380 };
381
382 let rewritten = rewrite_to_canary(req, &mapping);
383 assert!(matches!(rewritten.inner, McpRequest::ListTools(_)));
384 }
385
386 #[tokio::test]
387 async fn test_canary_service_routes_to_canary() {
388 let mock = MockService::with_tools(&["api/search", "api-canary/search"]);
390 let canaries = make_canaries("api", "api-canary", 0, 100);
391 let mut svc = CanaryService::new(mock, canaries, "/");
392
393 let resp = call_service(
394 &mut svc,
395 McpRequest::CallTool(CallToolParams {
396 name: "api/search".to_string(),
397 arguments: serde_json::json!({}),
398 input_responses: None,
399 request_state: None,
400 meta: None,
401 task: None,
402 }),
403 )
404 .await;
405
406 assert!(resp.inner.is_ok());
408 }
409
410 #[tokio::test]
411 async fn test_canary_service_passes_through_primary() {
412 let mock = MockService::with_tools(&["api/search"]);
414 let canaries = make_canaries("api", "api-canary", 100, 1);
415 let mut svc = CanaryService::new(mock, canaries, "/");
416
417 let resp = call_service(
419 &mut svc,
420 McpRequest::CallTool(CallToolParams {
421 name: "api/search".to_string(),
422 arguments: serde_json::json!({}),
423 input_responses: None,
424 request_state: None,
425 meta: None,
426 task: None,
427 }),
428 )
429 .await;
430
431 assert!(resp.inner.is_ok());
432 }
433
434 #[tokio::test]
435 async fn test_canary_service_non_matching_passes_through() {
436 let mock = MockService::with_tools(&["other/tool"]);
437 let canaries = make_canaries("api", "api-canary", 0, 100);
438 let mut svc = CanaryService::new(mock, canaries, "/");
439
440 let resp = call_service(
441 &mut svc,
442 McpRequest::CallTool(CallToolParams {
443 name: "other/tool".to_string(),
444 arguments: serde_json::json!({}),
445 input_responses: None,
446 request_state: None,
447 meta: None,
448 task: None,
449 }),
450 )
451 .await;
452
453 assert!(resp.inner.is_ok());
454 }
455
456 #[tokio::test]
457 async fn test_canary_service_list_tools_not_affected() {
458 let mock = MockService::with_tools(&["api/search"]);
459 let canaries = make_canaries("api", "api-canary", 0, 100);
460 let mut svc = CanaryService::new(mock, canaries, "/");
461
462 let resp = call_service(&mut svc, McpRequest::ListTools(Default::default())).await;
463 assert!(resp.inner.is_ok());
464 }
465}