use std::convert::Infallible;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::time::Duration;
use tower::Service;
use tower::ServiceBuilder;
use tower::timeout::TimeoutLayer;
use mcp_proxy::config::OutlierDetectionConfig;
use mcp_proxy::outlier::{OutlierDetectionLayer, OutlierDetector};
use tower_mcp::client::ChannelTransport;
use tower_mcp::protocol::{CallToolParams, McpRequest, McpResponse, RequestId};
use tower_mcp::proxy::McpProxy;
use tower_mcp::router::{Extensions, RouterRequest, RouterResponse};
use tower_mcp::{CallToolResult, McpRouter, ToolBuilder};
async fn call<S>(svc: &mut S, request: McpRequest) -> RouterResponse
where
S: Service<RouterRequest, Response = RouterResponse, Error = Infallible>,
{
let req = RouterRequest {
id: RequestId::Number(1),
inner: request,
extensions: Extensions::new(),
};
svc.call(req).await.expect("infallible")
}
fn tool_call(name: &str, args: serde_json::Value) -> McpRequest {
McpRequest::CallTool(CallToolParams {
name: name.to_string(),
arguments: args,
input_responses: None,
request_state: None,
meta: None,
task: None,
})
}
fn ping_router(hits: Arc<AtomicUsize>, slow: Arc<AtomicBool>) -> McpRouter {
let ping = ToolBuilder::new("ping")
.description("Ping; slow when the fault flag is set")
.handler(move |_: tower_mcp::NoParams| {
let hits = hits.clone();
let slow = slow.clone();
async move {
hits.fetch_add(1, Ordering::SeqCst);
if slow.load(Ordering::SeqCst) {
tokio::time::sleep(Duration::from_millis(100)).await;
}
Ok(CallToolResult::text("pong"))
}
})
.build();
McpRouter::new()
.server_info("ping-server", "1.0.0")
.tool(ping)
}
fn mcp_error_router(hits: Arc<AtomicUsize>) -> McpRouter {
let fail = ToolBuilder::new("fail")
.description("Always returns an MCP error result")
.handler(move |_: tower_mcp::NoParams| {
let hits = hits.clone();
async move {
hits.fetch_add(1, Ordering::SeqCst);
Ok(CallToolResult::error("tool-level failure"))
}
})
.build();
McpRouter::new()
.server_info("mcp-error-server", "1.0.0")
.tool(fail)
}
fn breaker_layer(
name: &str,
minimum_calls: usize,
wait_in_open: Duration,
permitted_in_half_open: usize,
) -> tower_resilience::circuitbreaker::CircuitBreakerLayer {
let (layer, _handle) = tower_resilience::circuitbreaker::CircuitBreakerLayer::builder()
.failure_rate_threshold(0.5)
.minimum_number_of_calls(minimum_calls)
.sliding_window_size(minimum_calls)
.wait_duration_in_open(wait_in_open)
.permitted_calls_in_half_open(permitted_in_half_open)
.name(format!("{name}-cb"))
.build_with_handle();
layer
}
fn ok_text(resp: &RouterResponse) -> String {
match resp.inner.as_ref().expect("expected success") {
McpResponse::CallTool(result) => result.all_text(),
other => panic!("expected CallTool, got: {other:?}"),
}
}
#[tokio::test]
async fn chaos_latency_timeouts_open_the_breaker() {
let injected = Arc::new(AtomicUsize::new(0));
let injected_cb = injected.clone();
let chaos = tower_resilience_chaos::ChaosLayer::builder()
.name("flaky-chaos")
.latency_rate(1.0)
.min_latency(Duration::from_millis(50))
.max_latency(Duration::from_millis(60))
.on_latency_injected(move |_| {
injected_cb.fetch_add(1, Ordering::SeqCst);
})
.seed(42)
.build();
let hits = Arc::new(AtomicUsize::new(0));
let slow = Arc::new(AtomicBool::new(false));
let stack = ServiceBuilder::new()
.layer(breaker_layer("flaky", 5, Duration::from_secs(60), 1))
.layer(TimeoutLayer::new(Duration::from_millis(10)))
.layer(chaos);
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"flaky",
ChannelTransport::new(ping_router(hits.clone(), slow)),
)
.await
.backend_layer(stack)
.build_strict()
.await
.expect("proxy should build");
for i in 1..=5 {
let resp = call(&mut proxy, tool_call("flaky/ping", serde_json::json!({}))).await;
assert!(resp.inner.is_err(), "call {i} should time out");
}
assert_eq!(
injected.load(Ordering::SeqCst),
5,
"all 5 calls reach chaos"
);
let resp = call(&mut proxy, tool_call("flaky/ping", serde_json::json!({}))).await;
assert!(resp.inner.is_err(), "open breaker rejects");
assert_eq!(
injected.load(Ordering::SeqCst),
5,
"rejected call must not reach the chaos layer"
);
}
#[tokio::test]
async fn timeout_breaker_opens_then_recovers_through_half_open() {
let hits = Arc::new(AtomicUsize::new(0));
let slow = Arc::new(AtomicBool::new(true));
let stack = ServiceBuilder::new()
.layer(breaker_layer("slowpoke", 4, Duration::from_millis(200), 2))
.layer(TimeoutLayer::new(Duration::from_millis(10)));
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"slowpoke",
ChannelTransport::new(ping_router(hits.clone(), slow.clone())),
)
.await
.backend_layer(stack)
.build_strict()
.await
.expect("proxy should build");
for i in 1..=4 {
let resp = call(
&mut proxy,
tool_call("slowpoke/ping", serde_json::json!({})),
)
.await;
assert!(resp.inner.is_err(), "call {i} should time out");
}
let hits_after_trip = hits.load(Ordering::SeqCst);
assert_eq!(hits_after_trip, 4, "all four calls reach the backend");
let resp = call(
&mut proxy,
tool_call("slowpoke/ping", serde_json::json!({})),
)
.await;
assert!(resp.inner.is_err(), "open breaker rejects");
assert_eq!(
hits.load(Ordering::SeqCst),
hits_after_trip,
"rejected call must not reach the backend"
);
tokio::time::sleep(Duration::from_millis(300)).await;
slow.store(false, Ordering::SeqCst);
for i in 1..=2 {
let resp = call(
&mut proxy,
tool_call("slowpoke/ping", serde_json::json!({})),
)
.await;
assert_eq!(ok_text(&resp), "pong", "half-open probe {i} succeeds");
}
let resp = call(
&mut proxy,
tool_call("slowpoke/ping", serde_json::json!({})),
)
.await;
assert_eq!(ok_text(&resp), "pong");
assert_eq!(hits.load(Ordering::SeqCst), hits_after_trip + 3);
}
#[tokio::test]
async fn mcp_error_results_do_not_trip_the_breaker() {
let hits = Arc::new(AtomicUsize::new(0));
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"erroring",
ChannelTransport::new(mcp_error_router(hits.clone())),
)
.await
.backend_layer(breaker_layer("erroring", 3, Duration::from_secs(60), 1))
.build_strict()
.await
.expect("proxy should build");
for _ in 0..6 {
let resp = call(
&mut proxy,
tool_call("erroring/fail", serde_json::json!({})),
)
.await;
match resp
.inner
.as_ref()
.expect("MCP error result is an Ok response")
{
McpResponse::CallTool(result) => {
assert!(result.is_error, "tool reports an MCP-level error");
}
other => panic!("expected CallTool, got: {other:?}"),
}
}
let before = hits.load(Ordering::SeqCst);
let resp = call(
&mut proxy,
tool_call("erroring/fail", serde_json::json!({})),
)
.await;
assert!(resp.inner.is_ok(), "breaker must not have opened");
assert_eq!(hits.load(Ordering::SeqCst), before + 1);
}
fn transient_failure_layer(
failures_remaining: Arc<AtomicUsize>,
) -> impl tower::Layer<
tower_mcp::proxy::BackendService,
Service = tower::util::BoxCloneService<RouterRequest, RouterResponse, Infallible>,
> {
tower::layer::layer_fn(move |inner: tower_mcp::proxy::BackendService| {
let failures = failures_remaining.clone();
tower::util::BoxCloneService::new(tower::ServiceExt::map_response(
inner,
move |resp: RouterResponse| {
if failures
.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |n| n.checked_sub(1))
.is_ok()
{
RouterResponse {
id: resp.id,
inner: Err(tower_mcp_types::JsonRpcError::internal_error(
"transient transport failure",
)),
}
} else {
resp
}
},
))
})
}
fn retry_config(
max_retries: u32,
budget_percent: Option<f64>,
min_retries_per_sec: u32,
) -> mcp_proxy::config::RetryConfig {
mcp_proxy::config::RetryConfig {
max_retries,
initial_backoff_ms: 1,
max_backoff_ms: 5,
budget_percent,
min_retries_per_sec,
}
}
async fn build_transient_proxy(
cfg: &mcp_proxy::config::RetryConfig,
hits: Arc<AtomicUsize>,
failures: Arc<AtomicUsize>,
) -> McpProxy {
let stack = ServiceBuilder::new()
.layer(mcp_proxy::retry::build_retry_layer(cfg, "flaky"))
.layer(transient_failure_layer(failures));
McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"flaky",
ChannelTransport::new(ping_router(hits, Arc::new(AtomicBool::new(false)))),
)
.await
.backend_layer(stack)
.build_strict()
.await
.expect("proxy should build")
}
#[tokio::test]
async fn retry_absorbs_transient_failures() {
let hits = Arc::new(AtomicUsize::new(0));
let failures = Arc::new(AtomicUsize::new(2));
let cfg = retry_config(3, None, 10);
let mut proxy = build_transient_proxy(&cfg, hits.clone(), failures).await;
let resp = call(&mut proxy, tool_call("flaky/ping", serde_json::json!({}))).await;
assert_eq!(ok_text(&resp), "pong");
assert_eq!(
hits.load(Ordering::SeqCst),
3,
"two failed attempts plus the success"
);
}
#[tokio::test]
async fn retry_gives_up_after_max_retries() {
let hits = Arc::new(AtomicUsize::new(0));
let failures = Arc::new(AtomicUsize::new(usize::MAX));
let cfg = retry_config(2, None, 10);
let mut proxy = build_transient_proxy(&cfg, hits.clone(), failures).await;
let resp = call(&mut proxy, tool_call("flaky/ping", serde_json::json!({}))).await;
let err = resp.inner.as_ref().expect_err("failure surfaces");
assert_eq!(err.code, -32603, "the backend error is returned as-is");
assert_eq!(
hits.load(Ordering::SeqCst),
3,
"initial attempt plus exactly max_retries"
);
}
#[tokio::test]
async fn retry_budget_caps_attempts_under_sustained_failure() {
let hits = Arc::new(AtomicUsize::new(0));
let failures = Arc::new(AtomicUsize::new(usize::MAX));
let cfg = retry_config(3, Some(1.0), 0);
let mut proxy = build_transient_proxy(&cfg, hits.clone(), failures).await;
let requests = 30;
for _ in 0..requests {
let resp = call(&mut proxy, tool_call("flaky/ping", serde_json::json!({}))).await;
assert!(resp.inner.is_err(), "sustained failure surfaces every time");
}
let total_attempts = hits.load(Ordering::SeqCst);
assert!(
total_attempts >= requests,
"every request attempts at least once (got {total_attempts})"
);
assert!(
total_attempts > requests,
"the budget allows some retries before it is spent (got {total_attempts})"
);
assert!(
total_attempts <= requests + 15,
"attempts stay within the ~10-token budget plus slack, no retry storm \
(got {total_attempts})"
);
}
fn outlier_config(consecutive_errors: u32, base_ejection_seconds: u64) -> OutlierDetectionConfig {
OutlierDetectionConfig {
consecutive_errors,
interval_seconds: 10,
base_ejection_seconds,
max_ejection_percent: 50,
}
}
#[tokio::test]
async fn outlier_ejects_after_consecutive_errors_and_fails_fast() {
let detector = OutlierDetector::new(50);
let hits = Arc::new(AtomicUsize::new(0));
let failures = Arc::new(AtomicUsize::new(usize::MAX));
let stack = ServiceBuilder::new()
.layer(OutlierDetectionLayer::new(
"sick".to_string(),
outlier_config(3, 60),
detector.clone(),
))
.layer(transient_failure_layer(failures));
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"sick",
ChannelTransport::new(ping_router(hits.clone(), Arc::new(AtomicBool::new(false)))),
)
.await
.backend_layer(stack)
.build_strict()
.await
.expect("proxy should build");
for i in 1..=3 {
let resp = call(&mut proxy, tool_call("sick/ping", serde_json::json!({}))).await;
assert!(resp.inner.is_err(), "call {i} fails at the response plane");
}
assert_eq!(
hits.load(Ordering::SeqCst),
3,
"all three reach the backend"
);
assert_eq!(detector.ejected_count(), 1, "backend is ejected");
let resp = call(&mut proxy, tool_call("sick/ping", serde_json::json!({}))).await;
assert!(resp.inner.is_err(), "ejected backend fails fast");
assert_eq!(
hits.load(Ordering::SeqCst),
3,
"ejected call must not reach the backend"
);
}
#[tokio::test]
async fn outlier_success_resets_the_error_streak() {
let detector = OutlierDetector::new(50);
let hits = Arc::new(AtomicUsize::new(0));
let failures = Arc::new(AtomicUsize::new(2));
let stack = ServiceBuilder::new()
.layer(OutlierDetectionLayer::new(
"wobbly".to_string(),
outlier_config(3, 60),
detector.clone(),
))
.layer(transient_failure_layer(failures.clone()));
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"wobbly",
ChannelTransport::new(ping_router(hits.clone(), Arc::new(AtomicBool::new(false)))),
)
.await
.backend_layer(stack)
.build_strict()
.await
.expect("proxy should build");
for _ in 0..2 {
let resp = call(&mut proxy, tool_call("wobbly/ping", serde_json::json!({}))).await;
assert!(resp.inner.is_err());
}
let resp = call(&mut proxy, tool_call("wobbly/ping", serde_json::json!({}))).await;
assert_eq!(ok_text(&resp), "pong", "streak broken by a success");
failures.store(2, Ordering::SeqCst);
for _ in 0..2 {
let resp = call(&mut proxy, tool_call("wobbly/ping", serde_json::json!({}))).await;
assert!(resp.inner.is_err());
}
assert_eq!(detector.ejected_count(), 0, "never ejected");
assert_eq!(
hits.load(Ordering::SeqCst),
5,
"every call reached the backend"
);
}
#[tokio::test]
async fn outlier_max_ejection_percent_caps_ejections() {
let detector = OutlierDetector::new(50);
let hits_a = Arc::new(AtomicUsize::new(0));
let hits_b = Arc::new(AtomicUsize::new(0));
let failures_a = Arc::new(AtomicUsize::new(usize::MAX));
let failures_b = Arc::new(AtomicUsize::new(usize::MAX));
let stack_a = ServiceBuilder::new()
.layer(OutlierDetectionLayer::new(
"a".to_string(),
outlier_config(3, 60),
detector.clone(),
))
.layer(transient_failure_layer(failures_a));
let stack_b = ServiceBuilder::new()
.layer(OutlierDetectionLayer::new(
"b".to_string(),
outlier_config(3, 60),
detector.clone(),
))
.layer(transient_failure_layer(failures_b));
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"a",
ChannelTransport::new(ping_router(
hits_a.clone(),
Arc::new(AtomicBool::new(false)),
)),
)
.await
.backend_layer(stack_a)
.backend(
"b",
ChannelTransport::new(ping_router(
hits_b.clone(),
Arc::new(AtomicBool::new(false)),
)),
)
.await
.backend_layer(stack_b)
.build_strict()
.await
.expect("proxy should build");
for _ in 0..3 {
let resp = call(&mut proxy, tool_call("a/ping", serde_json::json!({}))).await;
assert!(resp.inner.is_err());
}
assert_eq!(detector.ejected_count(), 1, "A is ejected");
for _ in 0..5 {
let resp = call(&mut proxy, tool_call("b/ping", serde_json::json!({}))).await;
assert!(resp.inner.is_err());
}
assert_eq!(detector.ejected_count(), 1, "B is not ejected (cap)");
assert_eq!(
hits_b.load(Ordering::SeqCst),
5,
"B keeps reaching the backend"
);
}
#[tokio::test]
#[ignore = "needs 1s+ wall clock (base_ejection_seconds granularity)"]
async fn outlier_ejection_expires_and_traffic_resumes() {
let detector = OutlierDetector::new(50);
let hits = Arc::new(AtomicUsize::new(0));
let failures = Arc::new(AtomicUsize::new(3));
let stack = ServiceBuilder::new()
.layer(OutlierDetectionLayer::new(
"healing".to_string(),
outlier_config(3, 1),
detector.clone(),
))
.layer(transient_failure_layer(failures));
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"healing",
ChannelTransport::new(ping_router(hits.clone(), Arc::new(AtomicBool::new(false)))),
)
.await
.backend_layer(stack)
.build_strict()
.await
.expect("proxy should build");
for _ in 0..3 {
let resp = call(&mut proxy, tool_call("healing/ping", serde_json::json!({}))).await;
assert!(resp.inner.is_err());
}
assert_eq!(detector.ejected_count(), 1);
tokio::time::sleep(Duration::from_millis(1200)).await;
let resp = call(&mut proxy, tool_call("healing/ping", serde_json::json!({}))).await;
assert_eq!(
ok_text(&resp),
"pong",
"traffic resumes after ejection expires"
);
assert_eq!(detector.ejected_count(), 0, "uneject recorded");
assert_eq!(hits.load(Ordering::SeqCst), 4);
}
fn hedge_layer(
name: &str,
delay: Duration,
max_hedges: u32,
) -> tower_resilience::hedge::HedgeLayer {
tower_resilience::hedge::HedgeLayer::builder()
.delay(delay)
.max_hedged_attempts((max_hedges + 1) as usize)
.name(format!("{name}-hedge"))
.build()
}
#[tokio::test]
async fn hedge_fires_on_slow_request_and_first_response_wins() {
let hits = Arc::new(AtomicUsize::new(0));
let slow = Arc::new(AtomicBool::new(true));
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"laggy",
ChannelTransport::new(ping_router(hits.clone(), slow)),
)
.await
.backend_layer(hedge_layer("laggy", Duration::from_millis(20), 1))
.build_strict()
.await
.expect("proxy should build");
let resp = call(&mut proxy, tool_call("laggy/ping", serde_json::json!({}))).await;
assert_eq!(ok_text(&resp), "pong", "client sees one winning response");
assert_eq!(
hits.load(Ordering::SeqCst),
2,
"the slow original plus exactly one hedge reach the backend"
);
}
#[tokio::test]
async fn hedge_does_not_fire_on_fast_requests() {
let hits = Arc::new(AtomicUsize::new(0));
let slow = Arc::new(AtomicBool::new(false));
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"snappy",
ChannelTransport::new(ping_router(hits.clone(), slow)),
)
.await
.backend_layer(hedge_layer("snappy", Duration::from_millis(50), 2))
.build_strict()
.await
.expect("proxy should build");
for _ in 0..3 {
let resp = call(&mut proxy, tool_call("snappy/ping", serde_json::json!({}))).await;
assert_eq!(ok_text(&resp), "pong");
}
assert_eq!(
hits.load(Ordering::SeqCst),
3,
"no hedges fire for fast responses"
);
}
#[tokio::test]
async fn hedge_respects_max_hedges() {
let hits = Arc::new(AtomicUsize::new(0));
let slow = Arc::new(AtomicBool::new(true));
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"laggy",
ChannelTransport::new(ping_router(hits.clone(), slow)),
)
.await
.backend_layer(hedge_layer("laggy", Duration::from_millis(10), 2))
.build_strict()
.await
.expect("proxy should build");
let resp = call(&mut proxy, tool_call("laggy/ping", serde_json::json!({}))).await;
assert_eq!(ok_text(&resp), "pong");
let attempts = hits.load(Ordering::SeqCst);
assert!(
(2..=3).contains(&attempts),
"at least one and at most max_hedges duplicates fire (got {attempts} attempts)"
);
}
#[tokio::test]
async fn hedge_composes_with_retry_without_amplification() {
let hits = Arc::new(AtomicUsize::new(0));
let failures = Arc::new(AtomicUsize::new(usize::MAX));
let cfg = retry_config(1, None, 10);
let stack = ServiceBuilder::new()
.layer(hedge_layer("flaky", Duration::from_millis(20), 1))
.layer(mcp_proxy::retry::build_retry_layer(&cfg, "flaky"))
.layer(transient_failure_layer(failures));
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"flaky",
ChannelTransport::new(ping_router(hits.clone(), Arc::new(AtomicBool::new(false)))),
)
.await
.backend_layer(stack)
.build_strict()
.await
.expect("proxy should build");
let resp = call(&mut proxy, tool_call("flaky/ping", serde_json::json!({}))).await;
assert!(resp.inner.is_err(), "sustained failure surfaces");
assert_eq!(
hits.load(Ordering::SeqCst),
2,
"fast failures never hedge: initial attempt plus one retry"
);
let hits = Arc::new(AtomicUsize::new(0));
let failures = Arc::new(AtomicUsize::new(usize::MAX));
let stack = ServiceBuilder::new()
.layer(hedge_layer("slowflaky", Duration::from_millis(20), 1))
.layer(mcp_proxy::retry::build_retry_layer(&cfg, "slowflaky"))
.layer(transient_failure_layer(failures));
let mut proxy = McpProxy::builder("chaos-proxy", "1.0.0")
.separator("/")
.backend(
"slowflaky",
ChannelTransport::new(ping_router(hits.clone(), Arc::new(AtomicBool::new(true)))),
)
.await
.backend_layer(stack)
.build_strict()
.await
.expect("proxy should build");
let resp = call(
&mut proxy,
tool_call("slowflaky/ping", serde_json::json!({})),
)
.await;
assert!(resp.inner.is_err(), "sustained failure surfaces");
tokio::time::sleep(Duration::from_millis(250)).await;
let attempts = hits.load(Ordering::SeqCst);
assert!(
(2..=4).contains(&attempts),
"attempts bounded by (1 + max_hedges) * (1 + max_retries) (got {attempts})"
);
}