use std::collections::BTreeMap;
use std::sync::Mutex;
use async_trait::async_trait;
use serde_json::Value;
use crate::{
McpAudit, McpCallClaim, McpCallContext, McpEndpoint, McpError, McpGuard, McpTool,
McpToolResult, McpTransport,
};
pub struct FakeMcpTransport {
pub tools: Vec<McpTool>,
pub result: Result<McpToolResult, McpError>,
}
#[async_trait]
impl McpTransport for FakeMcpTransport {
async fn list_tools(&self, _: &McpEndpoint) -> Result<Vec<McpTool>, McpError> {
Ok(self.tools.clone())
}
async fn call_tool(
&self,
_: &McpEndpoint,
_: &McpCallContext,
_: &str,
_: Value,
) -> Result<McpToolResult, McpError> {
self.result.clone()
}
}
pub struct AllowAllMcpGuard;
#[async_trait]
impl McpGuard for AllowAllMcpGuard {
async fn authorize(
&self,
_: &McpCallContext,
_: &str,
_: &str,
_: &Value,
) -> Result<(), McpError> {
Ok(())
}
}
#[derive(Default)]
pub struct MemoryMcpAudit(Mutex<BTreeMap<String, McpToolResult>>);
#[async_trait]
impl McpAudit for MemoryMcpAudit {
async fn claim(
&self,
context: &McpCallContext,
endpoint: &str,
tool: &str,
_: &Value,
) -> Result<McpCallClaim, McpError> {
let key = format!(
"{}:{endpoint}:{tool}:{}",
context.tenant_id, context.call_id
);
Ok(self
.0
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get(&key)
.cloned()
.map_or(McpCallClaim::Execute, McpCallClaim::Completed))
}
async fn complete(
&self,
context: &McpCallContext,
endpoint: &str,
tool: &str,
result: &McpToolResult,
) -> Result<(), McpError> {
let key = format!(
"{}:{endpoint}:{tool}:{}",
context.tenant_id, context.call_id
);
self.0
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(key, result.clone());
Ok(())
}
async fn outcome_unknown(
&self,
_: &McpCallContext,
_: &str,
_: &str,
_: &str,
) -> Result<(), McpError> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn fake_injects_transport_failure() {
let transport = FakeMcpTransport {
tools: Vec::new(),
result: Err(McpError::Unavailable("offline".into())),
};
assert!(matches!(
transport
.call_tool(
&McpEndpoint {
id: "id".into(),
url: "https://example.invalid".into(),
namespace: "test".into(),
allowed_hosts: Default::default(),
allowed_tools: Default::default(),
credential_ref: None,
timeout_ms: 1,
failure_threshold: 1,
recovery_ms: 1
},
&McpCallContext {
tenant_id: "tenant".parse().unwrap(),
subject_id: "subject".parse().unwrap(),
session_id: "session".parse().unwrap(),
run_id: "run".parse().unwrap(),
call_id: "call".parse().unwrap(),
source_event_seq: 1,
request_id: "request".parse().unwrap()
},
"tool",
Value::Null
)
.await,
Err(McpError::Unavailable(_))
));
}
}