#![deny(missing_docs)]
#![deny(rustdoc::broken_intra_doc_links)]
mod transport;
pub use transport::{McpCredential, McpCredentialProvider, StreamableHttpTransport};
use af_context::{RequestId, RunId, SessionId, SubjectId, TenantId, ToolCallId};
use std::collections::{BTreeMap, BTreeSet};
use std::net::IpAddr;
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use af_agent::{
validate_json_schema, AgentPlugin, AgentRegistrar, PluginError, PluginLease, PluginManifest,
PluginMountContext, PluginPermission, Tool,
};
use async_trait::async_trait;
use reqwest::Url;
use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct McpEndpoint {
pub id: String,
pub url: String,
pub namespace: String,
pub allowed_hosts: BTreeSet<String>,
pub allowed_tools: BTreeSet<String>,
pub credential_ref: Option<String>,
pub timeout_ms: u64,
pub failure_threshold: u32,
pub recovery_ms: u64,
}
impl McpEndpoint {
pub fn validate(&self) -> Result<Url, McpError> {
let url = Url::parse(&self.url).map_err(|_| McpError::Rejected("invalid URL".into()))?;
if url.scheme() != "https"
|| !url.username().is_empty()
|| url.password().is_some()
|| url.fragment().is_some()
{
return Err(McpError::Rejected(
"MCP requires an HTTPS URL without credentials or fragments".into(),
));
}
let host = url
.host_str()
.ok_or_else(|| McpError::Rejected("missing host".into()))?;
if host.parse::<IpAddr>().is_ok()
|| host.eq_ignore_ascii_case("localhost")
|| !self.allowed_hosts.contains(host)
{
return Err(McpError::Rejected(
"host is not an allowlisted DNS name".into(),
));
}
if self.id.trim().is_empty()
|| self.namespace.trim().is_empty()
|| self.allowed_tools.is_empty()
|| self.timeout_ms == 0
|| self.failure_threshold == 0
{
return Err(McpError::Rejected(
"id, namespace, tool allowlist, timeout and failure threshold are required".into(),
));
}
Ok(url)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct McpCallContext {
pub tenant_id: TenantId,
pub subject_id: SubjectId,
pub session_id: SessionId,
pub run_id: RunId,
pub call_id: ToolCallId,
pub source_event_seq: u64,
pub request_id: RequestId,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct McpTool {
pub name: String,
pub description: String,
pub input_schema: Value,
#[serde(default)]
pub output_schema: Value,
}
#[derive(Debug, Clone, PartialEq)]
pub struct McpToolResult {
pub content: Value,
pub is_error: bool,
}
#[derive(Debug, Clone, PartialEq)]
pub enum McpCallClaim {
Execute,
Completed(McpToolResult),
OutcomeUnknown,
}
#[async_trait]
pub trait McpTransport: Send + Sync {
async fn list_tools(&self, endpoint: &McpEndpoint) -> Result<Vec<McpTool>, McpError>;
async fn call_tool(
&self,
endpoint: &McpEndpoint,
context: &McpCallContext,
name: &str,
arguments: Value,
) -> Result<McpToolResult, McpError>;
}
#[async_trait]
pub trait McpGuard: Send + Sync {
async fn authorize(
&self,
context: &McpCallContext,
endpoint: &str,
tool: &str,
arguments: &Value,
) -> Result<(), McpError>;
}
#[async_trait]
pub trait McpAudit: Send + Sync {
async fn claim(
&self,
context: &McpCallContext,
endpoint: &str,
tool: &str,
arguments: &Value,
) -> Result<McpCallClaim, McpError>;
async fn complete(
&self,
context: &McpCallContext,
endpoint: &str,
tool: &str,
result: &McpToolResult,
) -> Result<(), McpError>;
async fn outcome_unknown(
&self,
context: &McpCallContext,
endpoint: &str,
tool: &str,
error: &str,
) -> Result<(), McpError>;
}
#[derive(Default)]
struct Circuit {
failures: u32,
opened_at: Option<Instant>,
}
pub struct McpClient<T, G, A> {
endpoint: McpEndpoint,
transport: T,
guard: G,
audit: A,
circuit: Mutex<Circuit>,
}
pub struct McpClientPlugin<T, G, A> {
manifest: PluginManifest,
client: Arc<McpClient<T, G, A>>,
}
impl<T, G, A> McpClientPlugin<T, G, A> {
pub fn new(client: McpClient<T, G, A>) -> Self {
let id = format!("agentfactory.mcp.{}", client.endpoint.id);
Self {
manifest: PluginManifest {
id,
version: "0.3.0".into(),
dependencies: BTreeMap::new(),
config_schema: serde_json::json!({"type":"object"}),
permissions: BTreeSet::from([PluginPermission::Tool]),
},
client: Arc::new(client),
}
}
}
#[async_trait]
impl<T, G, A> AgentPlugin for McpClientPlugin<T, G, A>
where
T: McpTransport + 'static,
G: McpGuard + 'static,
A: McpAudit + 'static,
{
fn manifest(&self) -> &PluginManifest {
&self.manifest
}
fn lease(&self, _: &PluginMountContext, _: &Value) -> Box<dyn PluginLease> {
Box::new(Lease)
}
async fn activate(
&self,
_scope: &PluginMountContext,
_: &Value,
registrar: &mut AgentRegistrar<'_>,
) -> Result<(), PluginError> {
let tools = self
.client
.list_tools()
.await
.map_err(|error| PluginError::Mount(format!("{}: {error}", self.manifest.id)))?;
for tool in tools {
registrar.tool(Arc::new(RemoteTool {
client: Arc::clone(&self.client),
tool,
}))?;
}
Ok(())
}
}
struct RemoteTool<T, G, A> {
client: Arc<McpClient<T, G, A>>,
tool: McpTool,
}
#[async_trait]
impl<T, G, A> Tool for RemoteTool<T, G, A>
where
T: McpTransport + 'static,
G: McpGuard + 'static,
A: McpAudit + 'static,
{
fn name(&self) -> &str {
&self.tool.name
}
fn description(&self) -> &str {
&self.tool.description
}
fn parameters(&self) -> Value {
self.tool.input_schema.clone()
}
fn output_schema(&self) -> Value {
if self.tool.output_schema.is_null() {
serde_json::json!({})
} else {
self.tool.output_schema.clone()
}
}
async fn call(&self, arguments: Value) -> Result<Value, String> {
let _ = arguments;
Err("durable tool execution context is required".into())
}
async fn call_with_context(
&self,
execution: &af_agent::ToolExecutionContext,
arguments: Value,
) -> Result<Value, String> {
let context = McpCallContext {
tenant_id: execution.request.tenant_id.clone(),
subject_id: execution.request.subject_id.clone(),
session_id: execution.session_id.clone(),
run_id: execution.run_id.clone(),
call_id: execution.call_id.clone(),
source_event_seq: execution.source_event_seq,
request_id: execution.request.request_id.clone(),
};
let result = self
.client
.call(&context, &self.tool, arguments)
.await
.map_err(|error| error.to_string())?;
if result.is_error {
Err(format!("remote MCP tool failed: {}", result.content))
} else {
Ok(result.content)
}
}
}
struct Lease;
#[async_trait]
impl PluginLease for Lease {
async fn unmount(&mut self) -> Result<(), PluginError> {
Ok(())
}
}
impl<T: McpTransport, G: McpGuard, A: McpAudit> McpClient<T, G, A> {
pub fn new(endpoint: McpEndpoint, transport: T, guard: G, audit: A) -> Result<Self, McpError> {
endpoint.validate()?;
Ok(Self {
endpoint,
transport,
guard,
audit,
circuit: Mutex::new(Circuit::default()),
})
}
pub async fn list_tools(&self) -> Result<Vec<McpTool>, McpError> {
self.ensure_closed()?;
let prefix = format!("{}.", self.endpoint.namespace);
let mut tools = self
.timed(self.transport.list_tools(&self.endpoint))
.await?;
let names = tools
.iter()
.map(|tool| tool.name.as_str())
.collect::<BTreeSet<_>>();
if tools.iter().any(|tool| {
tool.name.contains('.')
|| tool.name.trim().is_empty()
|| !tool.input_schema.is_object()
|| (!tool.output_schema.is_null() && !tool.output_schema.is_object())
}) || names.len() != tools.len()
|| !self
.endpoint
.allowed_tools
.iter()
.all(|name| names.contains(name.as_str()))
{
return Err(McpError::Rejected("remote tool catalog is invalid".into()));
}
tools.retain(|tool| self.endpoint.allowed_tools.contains(&tool.name));
for tool in &mut tools {
tool.name = format!("{prefix}{}", tool.name);
}
self.success();
Ok(tools)
}
pub async fn call(
&self,
context: &McpCallContext,
tool: &McpTool,
arguments: Value,
) -> Result<McpToolResult, McpError> {
if context.tenant_id.trim().is_empty()
|| context.subject_id.trim().is_empty()
|| context.session_id.trim().is_empty()
|| context.run_id.trim().is_empty()
|| context.call_id.trim().is_empty()
|| context.source_event_seq == 0
|| context.request_id.trim().is_empty()
{
return Err(McpError::Rejected(
"authenticated tenant context is required".into(),
));
}
self.ensure_closed()?;
let name = tool
.name
.strip_prefix(&format!("{}.", self.endpoint.namespace))
.ok_or_else(|| McpError::Rejected("tool is outside endpoint namespace".into()))?;
validate_json_schema(&tool.input_schema, &arguments).map_err(McpError::Rejected)?;
self.guard
.authorize(context, &self.endpoint.id, name, &arguments)
.await?;
match self
.audit
.claim(context, &self.endpoint.id, name, &arguments)
.await?
{
McpCallClaim::Completed(result) => {
validate_result(tool, &result)?;
return Ok(result);
}
McpCallClaim::OutcomeUnknown => return Err(McpError::OutcomeUnknown),
McpCallClaim::Execute => {}
}
let result = match self
.timed(
self.transport
.call_tool(&self.endpoint, context, name, arguments),
)
.await
{
Ok(result) => result,
Err(error) => {
let _ = self
.audit
.outcome_unknown(context, &self.endpoint.id, name, &error.to_string())
.await;
return Err(McpError::OutcomeUnknown);
}
};
self.audit
.complete(context, &self.endpoint.id, name, &result)
.await
.map_err(|_| McpError::OutcomeUnknown)?;
self.success();
validate_result(tool, &result)?;
Ok(result)
}
async fn timed<R>(
&self,
future: impl std::future::Future<Output = Result<R, McpError>>,
) -> Result<R, McpError> {
match tokio::time::timeout(Duration::from_millis(self.endpoint.timeout_ms), future).await {
Ok(Ok(value)) => Ok(value),
Ok(Err(error)) => {
self.failure();
Err(error)
}
Err(_) => {
self.failure();
Err(McpError::Timeout)
}
}
}
fn ensure_closed(&self) -> Result<(), McpError> {
let circuit = self
.circuit
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if circuit.failures < self.endpoint.failure_threshold {
return Ok(());
}
if circuit.opened_at.is_some_and(|opened| {
opened.elapsed() >= Duration::from_millis(self.endpoint.recovery_ms)
}) {
return Ok(());
}
Err(McpError::CircuitOpen)
}
fn failure(&self) {
let mut circuit = self
.circuit
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
circuit.failures += 1;
if circuit.failures >= self.endpoint.failure_threshold {
circuit.opened_at = Some(Instant::now());
}
}
fn success(&self) {
*self
.circuit
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner) = Circuit::default();
}
}
fn validate_result(tool: &McpTool, result: &McpToolResult) -> Result<(), McpError> {
if result.is_error || tool.output_schema.is_null() {
return Ok(());
}
validate_json_schema(&tool.output_schema, &result.content)
.map_err(|error| McpError::Rejected(format!("remote tool output is invalid: {error}")))
}
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
pub enum McpError {
#[error("MCP request rejected: {0}")]
Rejected(String),
#[error("MCP dependency unavailable: {0}")]
Unavailable(String),
#[error("MCP request timed out")]
Timeout,
#[error("MCP circuit is open")]
CircuitOpen,
#[error("MCP tool outcome is unknown; verify remote state before retrying")]
OutcomeUnknown,
}
#[cfg(test)]
mod tests {
use super::*;
struct Transport;
#[async_trait]
impl McpTransport for Transport {
async fn list_tools(&self, _: &McpEndpoint) -> Result<Vec<McpTool>, McpError> {
Ok(vec![McpTool {
name: "echo".into(),
description: String::new(),
input_schema: serde_json::json!({"type":"object","required":["text"],"properties":{"text":{"type":"string"}}}),
output_schema: serde_json::json!({"type":"object","required":["text"],"properties":{"text":{"type":"string"}}}),
}])
}
async fn call_tool(
&self,
_: &McpEndpoint,
_: &McpCallContext,
_: &str,
arguments: Value,
) -> Result<McpToolResult, McpError> {
Ok(McpToolResult {
content: arguments,
is_error: false,
})
}
}
struct Allow;
#[async_trait]
impl McpGuard for Allow {
async fn authorize(
&self,
_: &McpCallContext,
_: &str,
_: &str,
_: &Value,
) -> Result<(), McpError> {
Ok(())
}
}
#[async_trait]
impl McpAudit for Allow {
async fn claim(
&self,
_: &McpCallContext,
_: &str,
_: &str,
_: &Value,
) -> Result<McpCallClaim, McpError> {
Ok(McpCallClaim::Execute)
}
async fn complete(
&self,
_: &McpCallContext,
_: &str,
_: &str,
_: &McpToolResult,
) -> Result<(), McpError> {
Ok(())
}
async fn outcome_unknown(
&self,
_: &McpCallContext,
_: &str,
_: &str,
_: &str,
) -> Result<(), McpError> {
Ok(())
}
}
fn endpoint(url: &str) -> McpEndpoint {
McpEndpoint {
id: "docs".into(),
url: url.into(),
namespace: "docs".into(),
allowed_hosts: BTreeSet::from(["mcp.example.com".into()]),
allowed_tools: BTreeSet::from(["echo".into()]),
credential_ref: Some("secret-ref".into()),
timeout_ms: 100,
failure_threshold: 2,
recovery_ms: 1000,
}
}
#[tokio::test]
async fn rejects_ssrf_and_validates_namespaced_calls() {
assert!(
McpClient::new(endpoint("https://127.0.0.1/mcp"), Transport, Allow, Allow).is_err()
);
let client = McpClient::new(
endpoint("https://mcp.example.com/mcp"),
Transport,
Allow,
Allow,
)
.unwrap();
let tool = client.list_tools().await.unwrap().remove(0);
assert_eq!(tool.name, "docs.echo");
let context = McpCallContext {
tenant_id: "t".parse().unwrap(),
subject_id: "s".parse().unwrap(),
session_id: "session".parse().unwrap(),
run_id: "run".parse().unwrap(),
call_id: "call".parse().unwrap(),
source_event_seq: 1,
request_id: "r".parse().unwrap(),
};
assert!(client
.call(&context, &tool, serde_json::json!({}))
.await
.is_err());
assert_eq!(
client
.call(&context, &tool, serde_json::json!({"text":"ok"}))
.await
.unwrap()
.content["text"],
"ok"
);
}
struct ToolErrorTransport;
#[async_trait]
impl McpTransport for ToolErrorTransport {
async fn list_tools(&self, _: &McpEndpoint) -> Result<Vec<McpTool>, McpError> {
Transport
.list_tools(&endpoint("https://mcp.example.com/mcp"))
.await
}
async fn call_tool(
&self,
_: &McpEndpoint,
_: &McpCallContext,
_: &str,
_: Value,
) -> Result<McpToolResult, McpError> {
Ok(McpToolResult {
content: serde_json::json!({"content":[{"type":"text","text":"invalid input"}]}),
is_error: true,
})
}
}
#[tokio::test]
async fn application_tool_errors_do_not_open_transport_circuit() {
let client = McpClient::new(
endpoint("https://mcp.example.com/mcp"),
ToolErrorTransport,
Allow,
Allow,
)
.unwrap();
let tool = client.list_tools().await.unwrap().remove(0);
let context = McpCallContext {
tenant_id: "t".parse().unwrap(),
subject_id: "s".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(),
};
for _ in 0..3 {
assert!(
client
.call(&context, &tool, serde_json::json!({"text":"bad"}))
.await
.unwrap()
.is_error
);
}
assert_eq!(client.list_tools().await.unwrap().len(), 1);
}
struct Catalog(Value);
#[async_trait]
impl McpTransport for Catalog {
async fn list_tools(&self, _: &McpEndpoint) -> Result<Vec<McpTool>, McpError> {
Ok(vec![McpTool {
name: "echo".into(),
description: String::new(),
input_schema: serde_json::json!({"type":"object"}),
output_schema: self.0.clone(),
}])
}
async fn call_tool(
&self,
_: &McpEndpoint,
_: &McpCallContext,
_: &str,
_: Value,
) -> Result<McpToolResult, McpError> {
unreachable!()
}
}
#[tokio::test]
async fn rejects_missing_allowlisted_tools_and_invalid_output_schemas() {
let client = McpClient::new(
endpoint("https://mcp.example.com/mcp"),
Catalog(Value::String("invalid".into())),
Allow,
Allow,
)
.unwrap();
assert!(matches!(
client.list_tools().await,
Err(McpError::Rejected(_))
));
let client = McpClient::new(
endpoint("https://mcp.example.com/mcp"),
Catalog(Value::Null),
Allow,
Allow,
)
.unwrap();
assert_eq!(
client.list_tools().await.unwrap()[0].output_schema,
Value::Null
);
let mut missing = endpoint("https://mcp.example.com/mcp");
missing.allowed_tools = BTreeSet::from(["missing".into()]);
let client = McpClient::new(missing, Transport, Allow, Allow).unwrap();
assert!(matches!(
client.list_tools().await,
Err(McpError::Rejected(_))
));
}
struct InvalidOutput;
#[async_trait]
impl McpTransport for InvalidOutput {
async fn list_tools(&self, endpoint: &McpEndpoint) -> Result<Vec<McpTool>, McpError> {
Transport.list_tools(endpoint).await
}
async fn call_tool(
&self,
_: &McpEndpoint,
_: &McpCallContext,
_: &str,
_: Value,
) -> Result<McpToolResult, McpError> {
Ok(McpToolResult {
content: serde_json::json!({"wrong":true}),
is_error: false,
})
}
}
#[tokio::test]
async fn validates_normalized_success_output() {
let client = McpClient::new(
endpoint("https://mcp.example.com/mcp"),
InvalidOutput,
Allow,
Allow,
)
.unwrap();
let tool = client.list_tools().await.unwrap().remove(0);
let context = McpCallContext {
tenant_id: "t".parse().unwrap(),
subject_id: "s".parse().unwrap(),
session_id: "session".parse().unwrap(),
run_id: "run".parse().unwrap(),
call_id: "invalid-output".parse().unwrap(),
source_event_seq: 1,
request_id: "request".parse().unwrap(),
};
assert!(matches!(
client
.call(&context, &tool, serde_json::json!({"text":"ok"}))
.await,
Err(McpError::Rejected(_))
));
}
}