Skip to main content

mcp_proxy/
canary.rs

1//! Canary / weighted routing middleware.
2//!
3//! Routes a percentage of requests to a canary backend instead of the primary.
4//! The canary backend is registered as a separate backend with its own namespace,
5//! but its tools are hidden from `ListTools` (via capability filtering). When a
6//! request targets the primary namespace, this middleware probabilistically
7//! rewrites it to target the canary namespace instead.
8//!
9//! # Configuration
10//!
11//! ```toml
12//! [[backends]]
13//! name = "api"
14//! transport = "http"
15//! url = "http://api-v1.internal:8080"
16//! weight = 90
17//!
18//! [[backends]]
19//! name = "api-canary"
20//! transport = "http"
21//! url = "http://api-v2.internal:8080"
22//! weight = 10
23//! canary_of = "api"  # share namespace with api
24//! ```
25//!
26//! # How it works
27//!
28//! 1. Both `api` and `api-canary` are registered as separate backends
29//! 2. `api-canary`'s tools are auto-hidden via capability filtering
30//! 3. When `CallTool("api/search")` arrives, this middleware rolls a weighted
31//!    random selection: 90% chance it passes through to `api`, 10% chance it
32//!    rewrites to `CallTool("api-canary/search")`
33//! 4. `ListTools` always returns only the primary's tools
34
35use 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/// Tower layer that produces a [`CanaryService`].
48#[derive(Clone)]
49pub struct CanaryLayer {
50    canaries: HashMap<String, (String, u32, u32)>,
51    separator: String,
52}
53
54impl CanaryLayer {
55    /// Create a new canary routing layer.
56    ///
57    /// `canaries` maps primary backend names to `(canary_name, primary_weight, canary_weight)`.
58    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/// Mapping from a primary backend namespace to its canary configuration.
78#[derive(Debug, Clone)]
79struct CanaryMapping {
80    /// Primary namespace prefix (e.g. "api/").
81    primary_prefix: String,
82    /// Canary namespace prefix (e.g. "api-canary/").
83    canary_prefix: String,
84    /// Weight of the primary (e.g. 90).
85    primary_weight: u32,
86    /// Total weight (primary + canary, e.g. 100).
87    total_weight: u32,
88    /// Atomic counter for deterministic weight-based routing.
89    counter: Arc<AtomicU64>,
90}
91
92/// Canary routing middleware.
93///
94/// Wraps the proxy service and probabilistically rewrites requests from
95/// the primary namespace to the canary namespace based on configured weights.
96#[derive(Clone)]
97pub struct CanaryService<S> {
98    inner: S,
99    mappings: Arc<Vec<CanaryMapping>>,
100}
101
102impl<S> CanaryService<S> {
103    /// Create a new canary service.
104    ///
105    /// `canaries` maps primary backend names to `(canary_name, primary_weight, canary_weight)`.
106    /// The `separator` is used to construct namespace prefixes.
107    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
128/// Check if a request targets a primary namespace and return the mapping.
129fn 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
135/// Deterministic check: should this request go to the canary?
136fn 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    // Primary gets the first primary_weight slots, canary gets the rest
140    position >= mapping.primary_weight as u64
141}
142
143/// Rewrite a request to target the canary namespace.
144fn 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 = &params.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 = &params.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 = &params.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
186/// Extract the request name for namespace matching.
187fn request_name(req: &McpRequest) -> Option<&str> {
188    match req {
189        McpRequest::CallTool(params) => Some(&params.name),
190        McpRequest::ReadResource(params) => Some(&params.uri),
191        McpRequest::GetPrompt(params) => Some(&params.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        // Check if this request should be routed to a canary
214        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        // Over 100 requests, exactly 10 should go to canary
279        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        // Weight 0 primary / 100 canary = always canary
389        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        // Should succeed (rewritten to api-canary/search)
407        assert!(resp.inner.is_ok());
408    }
409
410    #[tokio::test]
411    async fn test_canary_service_passes_through_primary() {
412        // Weight 100 primary / 0 would panic, so use 100/1 (99% primary)
413        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        // First request goes to primary (position 0 < 100)
418        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}