1use std::collections::{BTreeMap, HashMap};
19use std::sync::{Arc, Mutex};
20use std::time::{Duration, Instant};
21
22use serde::{Deserialize, Serialize};
23use serde_json::Value as JsonValue;
24use sha2::{Digest, Sha256};
25
26use crate::mcp::{call_mcp_tool_with_hint, VmMcpClientHandle};
27use crate::mcp_protocol::McpCacheHint;
28use crate::mcp_registry::{self, RegisteredMcpServer};
29use crate::value::VmError;
30
31pub const DEFAULT_MAX_RESTARTS: u32 = 5;
36pub const DEFAULT_RESTART_WINDOW: Duration = Duration::from_mins(5);
37
38pub const DEFAULT_CIRCUIT_THRESHOLD: u32 = 5;
44pub const DEFAULT_CIRCUIT_RESET: Duration = Duration::from_secs(30);
45
46pub const INITIAL_RESTART_BACKOFF: Duration = Duration::from_millis(100);
49pub const MAX_RESTART_BACKOFF: Duration = Duration::from_secs(5);
50
51pub const RESPONSE_CACHE_MAX_ENTRIES_PER_TOOL: usize = 64;
56
57#[derive(Debug)]
62struct SupervisionState {
63 restart_attempts: Vec<Instant>,
66 consecutive_failures: u32,
69 breaker_opens_until: Option<Instant>,
74 ejected: bool,
77 circuit_threshold: u32,
79 circuit_reset: Duration,
81 max_restarts: u32,
83 restart_window: Duration,
85}
86
87impl SupervisionState {
88 fn new(policy: SupervisionPolicy) -> Self {
89 Self {
90 restart_attempts: Vec::new(),
91 consecutive_failures: 0,
92 breaker_opens_until: None,
93 ejected: false,
94 circuit_threshold: policy.circuit_threshold,
95 circuit_reset: policy.circuit_reset,
96 max_restarts: policy.max_restarts,
97 restart_window: policy.restart_window,
98 }
99 }
100
101 fn breaker_state(&mut self, now: Instant) -> BreakerState {
105 match self.breaker_opens_until {
106 Some(deadline) if now < deadline => BreakerState::Open,
107 Some(_) => BreakerState::HalfOpen,
108 None => BreakerState::Closed,
109 }
110 }
111
112 fn record_success(&mut self) {
115 self.consecutive_failures = 0;
116 self.breaker_opens_until = None;
117 }
118
119 fn record_failure(&mut self, now: Instant) {
123 self.consecutive_failures = self.consecutive_failures.saturating_add(1);
124 if self.consecutive_failures >= self.circuit_threshold {
125 self.breaker_opens_until = Some(now + self.circuit_reset);
126 }
127 }
128
129 fn record_restart(&mut self, now: Instant) -> bool {
134 self.prune_restart_window(now);
135 self.restart_attempts.push(now);
136 if self.restart_attempts.len() as u32 > self.max_restarts {
137 self.ejected = true;
138 return false;
139 }
140 true
141 }
142
143 fn backoff_delay(&self) -> Duration {
146 let attempt = self.restart_attempts.len() as u32;
147 let exp = attempt.saturating_sub(1).min(6);
148 let mul = 1u64 << exp;
149 let nanos = INITIAL_RESTART_BACKOFF.as_nanos() as u64 * mul;
150 Duration::from_nanos(nanos).min(MAX_RESTART_BACKOFF)
151 }
152
153 fn prune_restart_window(&mut self, now: Instant) {
154 let window = self.restart_window;
155 self.restart_attempts
156 .retain(|t| now.duration_since(*t) <= window);
157 }
158
159 fn clear(&mut self) {
160 self.restart_attempts.clear();
161 self.consecutive_failures = 0;
162 self.breaker_opens_until = None;
163 self.ejected = false;
164 }
165}
166
167#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
169#[serde(rename_all = "snake_case")]
170pub enum BreakerState {
171 Closed,
172 Open,
173 HalfOpen,
174}
175
176impl BreakerState {
177 pub fn as_str(self) -> &'static str {
178 match self {
179 BreakerState::Closed => "closed",
180 BreakerState::Open => "open",
181 BreakerState::HalfOpen => "half_open",
182 }
183 }
184}
185
186#[derive(Clone, Copy, Debug)]
190pub struct SupervisionPolicy {
191 pub circuit_threshold: u32,
192 pub circuit_reset: Duration,
193 pub max_restarts: u32,
194 pub restart_window: Duration,
195}
196
197impl Default for SupervisionPolicy {
198 fn default() -> Self {
199 Self {
200 circuit_threshold: DEFAULT_CIRCUIT_THRESHOLD,
201 circuit_reset: DEFAULT_CIRCUIT_RESET,
202 max_restarts: DEFAULT_MAX_RESTARTS,
203 restart_window: DEFAULT_RESTART_WINDOW,
204 }
205 }
206}
207
208#[derive(Clone, Debug)]
212struct CachedResponse {
213 payload: JsonValue,
214 inserted_at: Instant,
215 expires_at: Instant,
216 #[allow(dead_code)]
221 scope: Option<&'static str>,
222}
223
224#[derive(Clone, Debug, PartialEq, Eq)]
228pub enum AllowlistDecision {
229 Allow,
230 Deny { reason: String },
231}
232
233pub type AllowlistGuard = Arc<dyn Fn(&str, Option<&str>) -> AllowlistDecision + Send + Sync>;
238
239#[derive(Clone, Debug, Serialize)]
242pub struct McpHostStatus {
243 pub name: String,
244 pub transport: String,
245 pub url: Option<String>,
246 pub active: bool,
247 pub lazy: bool,
248 pub ref_count: usize,
249 pub restart_count: u32,
250 pub consecutive_failures: u32,
251 pub circuit: BreakerState,
252 pub ejected: bool,
253 pub cache_entries: usize,
256 pub display_identity: Option<String>,
260}
261
262#[derive(Clone, Debug, Default, Deserialize)]
265pub struct SpawnOptions {
266 #[serde(default)]
270 pub lazy: bool,
271 #[serde(default)]
274 pub keep_alive_ms: Option<u64>,
275 #[serde(default)]
278 pub card: Option<String>,
279 #[serde(default)]
282 pub circuit_threshold: Option<u32>,
283 #[serde(default)]
284 pub circuit_reset_ms: Option<u64>,
285 #[serde(default)]
286 pub max_restarts: Option<u32>,
287 #[serde(default)]
288 pub restart_window_ms: Option<u64>,
289}
290
291impl SpawnOptions {
292 fn into_policy(self) -> (SupervisionPolicy, RegisteredMcpServerMeta) {
293 let default = SupervisionPolicy::default();
294 let policy = SupervisionPolicy {
295 circuit_threshold: self.circuit_threshold.unwrap_or(default.circuit_threshold),
296 circuit_reset: self
297 .circuit_reset_ms
298 .map(Duration::from_millis)
299 .unwrap_or(default.circuit_reset),
300 max_restarts: self.max_restarts.unwrap_or(default.max_restarts),
301 restart_window: self
302 .restart_window_ms
303 .map(Duration::from_millis)
304 .unwrap_or(default.restart_window),
305 };
306 let meta = RegisteredMcpServerMeta {
307 lazy: self.lazy,
308 keep_alive: self.keep_alive_ms.map(Duration::from_millis),
309 card: self.card,
310 };
311 (policy, meta)
312 }
313}
314
315struct RegisteredMcpServerMeta {
316 lazy: bool,
317 keep_alive: Option<Duration>,
318 card: Option<String>,
319}
320
321struct HostInner {
324 supervision: HashMap<String, SupervisionState>,
328 response_cache: HashMap<(String, String), HashMap<String, CachedResponse>>,
332 allowlist: Option<AllowlistGuard>,
334 cache_hits: u64,
337 cache_misses: u64,
338}
339
340impl HostInner {
341 fn new() -> Self {
342 Self {
343 supervision: HashMap::new(),
344 response_cache: HashMap::new(),
345 allowlist: None,
346 cache_hits: 0,
347 cache_misses: 0,
348 }
349 }
350}
351
352static HOST: Mutex<Option<HostInner>> = Mutex::new(None);
353
354fn with_inner<F, R>(f: F) -> R
355where
356 F: FnOnce(&mut HostInner) -> R,
357{
358 let mut guard = HOST.lock().expect("mcp host mutex poisoned");
359 if guard.is_none() {
360 *guard = Some(HostInner::new());
361 }
362 f(guard.as_mut().expect("host inner just initialized"))
363}
364
365pub fn set_allowlist(guard: Option<AllowlistGuard>) {
367 with_inner(|inner| inner.allowlist = guard);
368}
369
370pub fn reset_for_tests() {
373 with_inner(|inner| {
374 inner.supervision.clear();
375 inner.response_cache.clear();
376 inner.allowlist = None;
377 inner.cache_hits = 0;
378 inner.cache_misses = 0;
379 });
380 mcp_registry::reset();
381}
382
383#[derive(Clone, Copy, Debug)]
387pub struct CacheStats {
388 pub hits: u64,
389 pub misses: u64,
390}
391
392pub fn cache_stats() -> CacheStats {
393 with_inner(|inner| CacheStats {
394 hits: inner.cache_hits,
395 misses: inner.cache_misses,
396 })
397}
398
399pub async fn spawn(spec: JsonValue, options: SpawnOptions) -> Result<String, VmError> {
403 let name = spec
404 .get("name")
405 .and_then(|v| v.as_str())
406 .ok_or_else(|| VmError::Runtime("mcp.spawn: spec must include a `name` field".into()))?
407 .to_string();
408 if name.is_empty() {
409 return Err(VmError::Runtime(
410 "mcp.spawn: spec.name must be a non-empty string".into(),
411 ));
412 }
413
414 if let Some(guard) = current_allowlist() {
415 if let AllowlistDecision::Deny { reason } = guard(&name, None) {
416 return Err(VmError::Runtime(format!(
417 "mcp.spawn({name}): denied by allowlist: {reason}"
418 )));
419 }
420 }
421
422 let (policy, meta) = options.into_policy();
423 mcp_registry::register_servers(vec![RegisteredMcpServer {
424 name: name.clone(),
425 spec: spec.clone(),
426 preparation: None,
427 lazy: meta.lazy,
428 card: meta.card,
429 keep_alive: meta.keep_alive,
430 }]);
431
432 with_inner(|inner| {
433 inner
434 .supervision
435 .insert(name.clone(), SupervisionState::new(policy));
436 });
437
438 if !meta.lazy {
439 let _ = mcp_registry::ensure_active(&name).await.inspect_err(|_| {
442 with_inner(|inner| {
443 inner.supervision.remove(&name);
444 });
445 })?;
446 }
447
448 Ok(name)
449}
450
451pub fn stop(name: &str) -> Result<(), VmError> {
457 if !mcp_registry::is_registered(name) {
458 return Err(VmError::Runtime(format!(
459 "mcp.stop: no server named '{name}' is hosted"
460 )));
461 }
462 mcp_registry::release(name);
463 with_inner(|inner| {
464 inner.supervision.remove(name);
465 inner.response_cache.retain(|(s, _), _| s != name);
466 });
467 Ok(())
468}
469
470pub fn reload(name: &str) -> Result<(), VmError> {
475 if !mcp_registry::is_registered(name) {
476 return Err(VmError::Runtime(format!(
477 "mcp.reload: no server named '{name}' is hosted"
478 )));
479 }
480 mcp_registry::release(name);
481 with_inner(|inner| {
482 if let Some(state) = inner.supervision.get_mut(name) {
483 state.clear();
484 }
485 inner.response_cache.retain(|(s, _), _| s != name);
486 });
487 Ok(())
488}
489
490pub async fn tools(name: &str) -> Result<Vec<JsonValue>, VmError> {
495 let handle = ensure_or_restart(name).await?;
496 let result = supervised_call(name, || async {
497 handle.call("tools/list", serde_json::json!({})).await
498 })
499 .await?;
500
501 let mut tools = result
502 .get("tools")
503 .and_then(|t| t.as_array())
504 .cloned()
505 .unwrap_or_default();
506 for tool in tools.iter_mut() {
507 if let Some(obj) = tool.as_object_mut() {
508 obj.entry("_mcp_server")
509 .or_insert_with(|| JsonValue::String(name.to_string()));
510 }
511 }
512 let security_policy = crate::security::current_policy();
517 if security_policy.pin_mcp_schemas && !security_policy.server_is_trusted(name) {
518 for tool in tools.iter_mut() {
519 let hash = crate::security::tool_schema_hash(tool);
520 let tool_name = tool
521 .get("name")
522 .and_then(|v| v.as_str())
523 .unwrap_or_default()
524 .to_string();
525 if tool_name.is_empty() {
526 continue;
527 }
528 if crate::security::pin_and_detect_change(name, &tool_name, &hash) {
529 if let Some(obj) = tool.as_object_mut() {
530 obj.insert("_schema_changed".to_string(), JsonValue::Bool(true));
531 }
532 }
533 }
534 }
535 Ok(tools)
536}
537
538pub async fn call(name: &str, tool: &str, args: JsonValue) -> Result<JsonValue, VmError> {
545 if let Some(guard) = current_allowlist() {
546 if let AllowlistDecision::Deny { reason } = guard(name, Some(tool)) {
547 return Err(VmError::Runtime(format!(
548 "mcp.call({name}/{tool}): denied by allowlist: {reason}"
549 )));
550 }
551 }
552
553 crate::call_budget::charge_mcp_call()?;
558
559 let now = Instant::now();
560 let args_hash = hash_args(&args);
561 if let Some(payload) = take_cache_hit(name, tool, &args_hash, now) {
562 return Ok(payload);
563 }
564 with_inner(|inner| inner.cache_misses = inner.cache_misses.saturating_add(1));
565
566 breaker_gate(name, now)?;
567
568 let handle = ensure_or_restart(name).await?;
569 let envelope_hint: Arc<Mutex<Option<McpCacheHint>>> = Arc::new(Mutex::new(None));
574 let hint_slot = Arc::clone(&envelope_hint);
575 let result = supervised_call(name, move || {
576 let handle = handle.clone();
577 let tool = tool.to_string();
578 let args = args.clone();
579 let hint_slot = Arc::clone(&hint_slot);
580 async move {
581 let (content, hint) = call_mcp_tool_with_hint(&handle, &tool, args).await?;
582 if let Ok(mut slot) = hint_slot.lock() {
583 *slot = hint;
584 }
585 Ok(content)
586 }
587 })
588 .await?;
589
590 let hint = envelope_hint.lock().ok().and_then(|slot| *slot);
591 if let Some(hint) = hint {
592 insert_cache(name, tool, &args_hash, &result, hint, now);
593 }
594
595 Ok(result)
596}
597
598pub async fn discover() -> Result<Vec<JsonValue>, VmError> {
602 let names: Vec<String> = mcp_registry::snapshot_status()
603 .into_iter()
604 .map(|s| s.name)
605 .collect();
606 let mut out: Vec<JsonValue> = Vec::new();
607 for name in names {
608 if let Some(guard) = current_allowlist() {
611 if matches!(guard(&name, None), AllowlistDecision::Deny { .. }) {
612 continue;
613 }
614 }
615 match tools(&name).await {
619 Ok(tools) => {
620 for tool in tools {
621 let tool_name = tool
622 .get("name")
623 .and_then(|v| v.as_str())
624 .unwrap_or("")
625 .to_string();
626 out.push(serde_json::json!({
627 "server": name,
628 "tool": tool_name,
629 "schema": tool,
630 }));
631 }
632 }
633 Err(err) => {
634 out.push(serde_json::json!({
635 "server": name,
636 "error": err.to_string(),
637 }));
638 }
639 }
640 }
641 Ok(out)
642}
643
644pub async fn status() -> Vec<McpHostStatus> {
646 let registry: BTreeMap<String, mcp_registry::RegistryStatus> = mcp_registry::snapshot_status()
647 .into_iter()
648 .map(|s| (s.name.clone(), s))
649 .collect();
650 let mut statuses = with_inner(|inner| {
651 let mut out = Vec::new();
652 let now = Instant::now();
653 for (name, reg) in ®istry {
654 let (restart_count, consecutive_failures, circuit, ejected) =
655 if let Some(state) = inner.supervision.get_mut(name) {
656 let st = state.breaker_state(now);
657 (
658 state.restart_attempts.len() as u32,
659 state.consecutive_failures,
660 st,
661 state.ejected,
662 )
663 } else {
664 (0, 0, BreakerState::Closed, false)
665 };
666 let cache_entries = inner
667 .response_cache
668 .iter()
669 .filter(|((s, _), _)| s == name)
670 .map(|(_, v)| v.len())
671 .sum();
672 out.push(McpHostStatus {
673 name: name.clone(),
674 transport: reg.transport.clone(),
675 url: reg.url.clone(),
676 active: reg.active,
677 lazy: reg.lazy,
678 ref_count: reg.ref_count,
679 restart_count,
680 consecutive_failures,
681 circuit,
682 ejected,
683 cache_entries,
684 display_identity: None,
685 });
686 }
687 out
688 });
689 for status in &mut statuses {
690 if !status.active || status.transport != "http" {
691 continue;
692 }
693 let Some(url) = status.url.as_deref() else {
694 continue;
695 };
696 status.display_identity = crate::mcp_identity::display_identity_from_store(url, None).await;
697 }
698 statuses
699}
700
701fn current_allowlist() -> Option<AllowlistGuard> {
702 with_inner(|inner| inner.allowlist.clone())
703}
704
705fn breaker_gate(name: &str, now: Instant) -> Result<(), VmError> {
706 with_inner(|inner| {
707 let Some(state) = inner.supervision.get_mut(name) else {
708 return Ok(());
709 };
710 if state.ejected {
711 return Err(VmError::Runtime(format!(
712 "mcp.call({name}): server is ejected after exhausting its restart budget; call `harn.mcp.reload({name:?})` to clear"
713 )));
714 }
715 match state.breaker_state(now) {
716 BreakerState::Open => Err(VmError::Runtime(format!(
717 "mcp.call({name}): circuit breaker is open (last {n} consecutive failures); retry after the breaker resets",
718 n = state.consecutive_failures
719 ))),
720 BreakerState::Closed | BreakerState::HalfOpen => Ok(()),
723 }
724 })
725}
726
727async fn ensure_or_restart(name: &str) -> Result<VmMcpClientHandle, VmError> {
728 if let Some(handle) = mcp_registry::active_handle(name) {
730 return Ok(handle);
731 }
732
733 mcp_registry::ensure_active(name).await
739}
740
741async fn supervised_call<F, Fut>(name: &str, op: F) -> Result<JsonValue, VmError>
745where
746 F: Fn() -> Fut,
747 Fut: std::future::Future<Output = Result<JsonValue, VmError>>,
748{
749 let span = tracing::info_span!(
750 "harn.mcp.call",
751 otel.name = "harn.mcp.call",
752 harn.mcp.server = name,
753 );
754 let _enter = span.enter();
755
756 let first = op().await;
757 match first {
758 Ok(v) => {
759 with_inner(|inner| {
760 if let Some(state) = inner.supervision.get_mut(name) {
761 state.record_success();
762 }
763 });
764 Ok(v)
765 }
766 Err(err) => {
767 let now = Instant::now();
768 let (should_retry, backoff) = with_inner(|inner| {
769 let Some(state) = inner.supervision.get_mut(name) else {
770 return (false, Duration::ZERO);
771 };
772 state.record_failure(now);
773 if !looks_like_transport_failure(&err) {
778 return (false, Duration::ZERO);
779 }
780 let ok = state.record_restart(now);
781 if !ok {
782 return (false, Duration::ZERO);
783 }
784 (true, state.backoff_delay())
785 });
786 if !should_retry {
787 tracing::warn!(
788 server = name,
789 error = %err,
790 "harn.mcp.call: failure (no retry)"
791 );
792 return Err(err);
793 }
794
795 tracing::info!(
796 server = name,
797 error = %err,
798 backoff_ms = backoff.as_millis() as u64,
799 "harn.mcp.call: retrying after transport failure"
800 );
801
802 mcp_registry::release(name);
805 tokio::time::sleep(backoff).await;
806 let _handle = ensure_or_restart(name).await?;
807 let second = op().await;
808 match &second {
809 Ok(_) => with_inner(|inner| {
810 if let Some(state) = inner.supervision.get_mut(name) {
811 state.record_success();
812 }
813 }),
814 Err(err) => with_inner(|inner| {
815 if let Some(state) = inner.supervision.get_mut(name) {
816 state.record_failure(Instant::now());
817 }
818 tracing::warn!(
819 server = name,
820 error = %err,
821 "harn.mcp.call: second attempt failed"
822 );
823 }),
824 }
825 second
826 }
827 }
828}
829
830fn looks_like_transport_failure(err: &VmError) -> bool {
831 let text = err.to_string();
832 let needles = [
833 "server closed connection",
834 "disconnected",
835 "MCP read error",
836 "MCP write error",
837 "did not respond to",
838 "MCP flush error",
839 "connect",
840 ];
841 needles.iter().any(|n| text.contains(n))
842}
843
844fn hash_args(args: &JsonValue) -> String {
845 let mut hasher = Sha256::new();
846 let canonical = crate::canonical_json::to_string(args);
847 hasher.update(canonical.as_bytes());
848 let digest = hasher.finalize();
849 let mut hex = String::with_capacity(digest.len() * 2);
850 for byte in digest {
851 use std::fmt::Write;
852 let _ = write!(&mut hex, "{byte:02x}");
853 }
854 hex
855}
856
857fn take_cache_hit(server: &str, tool: &str, args_hash: &str, now: Instant) -> Option<JsonValue> {
858 with_inner(|inner| {
859 let key = (server.to_string(), tool.to_string());
860 let entry = inner.response_cache.get_mut(&key)?;
861 let cached = entry.get(args_hash)?;
862 if now >= cached.expires_at {
863 entry.remove(args_hash);
864 return None;
865 }
866 let payload = cached.payload.clone();
867 inner.cache_hits = inner.cache_hits.saturating_add(1);
868 Some(payload)
869 })
870}
871
872fn insert_cache(
877 server: &str,
878 tool: &str,
879 args_hash: &str,
880 payload: &JsonValue,
881 hint: McpCacheHint,
882 now: Instant,
883) {
884 let Some(ttl_ms) = hint.ttl_ms else {
885 return;
886 };
887 if ttl_ms == 0 {
888 return;
889 }
890 let expires_at = now + Duration::from_millis(ttl_ms);
891 let cached = CachedResponse {
892 payload: payload.clone(),
893 inserted_at: now,
894 expires_at,
895 scope: hint.scope,
896 };
897 with_inner(|inner| {
898 let key = (server.to_string(), tool.to_string());
899 let bucket = inner.response_cache.entry(key).or_default();
900 if bucket.len() >= RESPONSE_CACHE_MAX_ENTRIES_PER_TOOL {
901 if let Some(oldest_key) = bucket
903 .iter()
904 .min_by_key(|(_, v)| v.inserted_at)
905 .map(|(k, _)| k.clone())
906 {
907 bucket.remove(&oldest_key);
908 }
909 }
910 bucket.insert(args_hash.to_string(), cached);
911 });
912}
913
914#[cfg(test)]
915mod tests {
916 use super::*;
917
918 static TEST_LOCK: Mutex<()> = Mutex::new(());
919
920 fn lock() -> std::sync::MutexGuard<'static, ()> {
921 TEST_LOCK.lock().unwrap_or_else(|p| p.into_inner())
922 }
923
924 #[test]
925 fn supervision_breaker_opens_after_threshold() {
926 let _g = lock();
927 let mut state = SupervisionState::new(SupervisionPolicy {
928 circuit_threshold: 3,
929 circuit_reset: Duration::from_millis(100),
930 ..SupervisionPolicy::default()
931 });
932 let t0 = Instant::now();
933 assert_eq!(state.breaker_state(t0), BreakerState::Closed);
934 state.record_failure(t0);
935 state.record_failure(t0);
936 assert_eq!(state.breaker_state(t0), BreakerState::Closed);
937 state.record_failure(t0);
938 assert_eq!(state.breaker_state(t0), BreakerState::Open);
939 assert_eq!(
941 state.breaker_state(t0 + Duration::from_millis(200)),
942 BreakerState::HalfOpen
943 );
944 }
945
946 #[test]
947 fn supervision_restart_budget_ejects_after_n_attempts() {
948 let _g = lock();
949 let mut state = SupervisionState::new(SupervisionPolicy {
950 max_restarts: 2,
951 restart_window: Duration::from_mins(1),
952 ..SupervisionPolicy::default()
953 });
954 let t = Instant::now();
955 assert!(state.record_restart(t));
956 assert!(state.record_restart(t));
957 assert!(!state.record_restart(t));
958 assert!(state.ejected);
959 }
960
961 #[test]
962 fn supervision_backoff_grows_exponentially_then_caps() {
963 let _g = lock();
964 let mut state = SupervisionState::new(SupervisionPolicy::default());
965 let t = Instant::now();
966 state.record_restart(t);
967 let d1 = state.backoff_delay();
968 state.record_restart(t);
969 let d2 = state.backoff_delay();
970 state.record_restart(t);
971 let d3 = state.backoff_delay();
972 assert!(
973 d2 > d1,
974 "second backoff ({d2:?}) should exceed first ({d1:?})"
975 );
976 assert!(d3 > d2);
977 for _ in 0..16 {
978 state.record_restart(t);
979 }
980 assert!(state.backoff_delay() <= MAX_RESTART_BACKOFF);
981 }
982
983 #[test]
984 fn hash_args_is_stable_across_key_order() {
985 let h1 = hash_args(&serde_json::json!({"x": 1, "y": [1, 2]}));
986 let h2 = hash_args(&serde_json::json!({"y": [1, 2], "x": 1}));
987 assert_eq!(h1, h2);
988 }
989
990 #[test]
991 fn cache_insert_and_take_respects_ttl() {
992 let _g = lock();
993 reset_for_tests();
994 let payload = serde_json::json!({
995 "ttlMs": 100,
996 "cacheScope": "private",
997 "value": 1
998 });
999 let now = Instant::now();
1000 insert_cache(
1001 "srv",
1002 "ping",
1003 "deadbeef",
1004 &payload,
1005 McpCacheHint::from_result(&payload).unwrap(),
1006 now,
1007 );
1008 let hit = take_cache_hit("srv", "ping", "deadbeef", now);
1009 assert!(hit.is_some(), "fresh entry should hit");
1010 let stale = take_cache_hit("srv", "ping", "deadbeef", now + Duration::from_millis(200));
1011 assert!(stale.is_none(), "expired entry should miss");
1012 }
1013
1014 #[test]
1015 fn allowlist_denies_disallowed_tool() {
1016 let _g = lock();
1017 reset_for_tests();
1018 set_allowlist(Some(Arc::new(|server, tool| {
1019 if server == "github" && tool == Some("delete_repo") {
1020 AllowlistDecision::Deny {
1021 reason: "destructive tool blocked".into(),
1022 }
1023 } else {
1024 AllowlistDecision::Allow
1025 }
1026 })));
1027 let runtime = tokio::runtime::Builder::new_current_thread()
1028 .enable_all()
1029 .build()
1030 .unwrap();
1031 let err = runtime
1032 .block_on(call("github", "delete_repo", serde_json::json!({})))
1033 .unwrap_err();
1034 assert!(err.to_string().contains("denied by allowlist"));
1035 set_allowlist(None);
1036 }
1037
1038 #[test]
1039 fn stop_unregistered_server_errors() {
1040 let _g = lock();
1041 reset_for_tests();
1042 let err = stop("nope").unwrap_err();
1043 assert!(err.to_string().contains("no server named 'nope'"));
1044 }
1045
1046 #[test]
1047 fn supervision_record_success_resets_counters() {
1048 let _g = lock();
1049 let mut state = SupervisionState::new(SupervisionPolicy::default());
1050 let t = Instant::now();
1051 state.record_failure(t);
1052 state.record_failure(t);
1053 state.record_success();
1054 assert_eq!(state.consecutive_failures, 0);
1055 assert!(state.breaker_opens_until.is_none());
1056 }
1057
1058 #[test]
1059 fn looks_like_transport_failure_matches_common_errors() {
1060 let cases = [
1061 "MCP: server closed connection",
1062 "MCP: server did not respond to 'tools/call' within 60s",
1063 "MCP write error: broken pipe",
1064 "MCP client is disconnected",
1065 ];
1066 for msg in cases {
1067 assert!(
1068 looks_like_transport_failure(&VmError::Runtime(msg.into())),
1069 "expected {msg:?} to be classified as transport failure"
1070 );
1071 }
1072 assert!(
1073 !looks_like_transport_failure(&VmError::Runtime(
1074 "tool 'foo' rejected arguments".into()
1075 )),
1076 "tool-level errors must not trigger an auto-restart"
1077 );
1078 }
1079}