af-mcp-client 0.17.2

Fail-closed Streamable HTTP MCP client policy for cloud Agent hosts.
Documentation
//! In-memory MCP transport, guard and audit doubles.

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,
};

/// Deterministic transport with optional injected failure.
pub struct FakeMcpTransport {
    /// Catalog returned by `list_tools`.
    pub tools: Vec<McpTool>,
    /// Result returned by `call_tool`.
    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()
    }
}

/// Guard that authorizes every call.
pub struct AllowAllMcpGuard;

#[async_trait]
impl McpGuard for AllowAllMcpGuard {
    async fn authorize(
        &self,
        _: &McpCallContext,
        _: &str,
        _: &str,
        _: &Value,
    ) -> Result<(), McpError> {
        Ok(())
    }
}

/// Tenant-scoped in-memory idempotency audit.
#[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(_))
        ));
    }
}