use crate::runtime::error::Result;
use crate::runtime::mcp_server::{McpServerActsAs, is_mcp_tool, parse_mcp_tool_name};
use crate::runtime::tool_context::ToolContext;
use crate::runtime::tool_types::{BuiltinTool, ToolCall, ToolDefinition, ToolHints};
use crate::runtime::tools::{Tool, ToolExecutionResult};
use async_trait::async_trait;
use serde_json::Value;
use std::collections::HashSet;
use std::sync::Arc;
#[async_trait]
pub trait McpToolInvoker: Send + Sync {
fn for_execution(&self, _input_message_id: uuid::Uuid) -> Option<Arc<dyn McpToolInvoker>> {
None
}
async fn invoke(&self, tool_call: &ToolCall) -> Result<crate::runtime::tool_types::ToolResult>;
async fn invoke_recorded(
&self,
tool_call: &ToolCall,
) -> Result<(
crate::runtime::tool_types::ToolResult,
Option<McpServerActsAs>,
)> {
self.invoke(tool_call).await.map(|result| (result, None))
}
async fn list_server_tools(
&self,
_server_prefix: &str,
_session_id: uuid::Uuid,
) -> Result<Option<McpServerTools>> {
Ok(None)
}
}
#[derive(Debug, Clone)]
pub enum McpServerTools {
Listed(Vec<ToolDefinition>),
ConnectionRequired(crate::runtime::tool_types::ToolResult),
}
#[derive(Debug, Default)]
pub struct McpCallIdentity(std::sync::Mutex<Option<McpServerActsAs>>);
impl McpCallIdentity {
pub fn record(&self, acted_as: Option<McpServerActsAs>) {
*self.0.lock().unwrap_or_else(|error| error.into_inner()) = acted_as;
}
pub fn get(&self) -> Option<McpServerActsAs> {
*self.0.lock().unwrap_or_else(|error| error.into_inner())
}
}
pub struct ScopedMcpToolInvoker {
inner: Arc<dyn McpToolInvoker>,
allowed_tool_names: HashSet<String>,
deferred_prefixes: HashSet<String>,
}
impl ScopedMcpToolInvoker {
pub fn new(definitions: &[ToolDefinition], inner: Arc<dyn McpToolInvoker>) -> Self {
let allowed_tool_names: HashSet<String> = definitions
.iter()
.filter_map(|def| match def {
ToolDefinition::Builtin(builtin) if is_mcp_tool(&builtin.name) => {
Some(builtin.name.clone())
}
ToolDefinition::ClientSide(_) | ToolDefinition::Builtin(_) => None,
})
.collect();
let deferred_prefixes = allowed_tool_names
.iter()
.filter_map(|name| crate::runtime::mcp_deferred::deferred_mcp_server_prefix(name))
.map(str::to_string)
.collect();
Self {
inner,
allowed_tool_names,
deferred_prefixes,
}
}
fn check(&self, tool_name: &str) -> Result<()> {
let on_deferred_server = parse_mcp_tool_name(tool_name)
.is_some_and(|(prefix, _)| self.deferred_prefixes.contains(&prefix));
if self.allowed_tool_names.contains(tool_name) || on_deferred_server {
return Ok(());
}
Err(crate::runtime::AgentLoopError::tool(format!(
"MCP tool '{tool_name}' is not allowed in the current tool scope"
)))
}
}
#[async_trait]
impl McpToolInvoker for ScopedMcpToolInvoker {
fn for_execution(&self, input_message_id: uuid::Uuid) -> Option<Arc<dyn McpToolInvoker>> {
self.inner.for_execution(input_message_id).map(|inner| {
Arc::new(Self {
inner,
allowed_tool_names: self.allowed_tool_names.clone(),
deferred_prefixes: self.deferred_prefixes.clone(),
}) as Arc<dyn McpToolInvoker>
})
}
async fn invoke(&self, tool_call: &ToolCall) -> Result<crate::runtime::tool_types::ToolResult> {
self.check(&tool_call.name)?;
self.inner.invoke(tool_call).await
}
async fn invoke_recorded(
&self,
tool_call: &ToolCall,
) -> Result<(
crate::runtime::tool_types::ToolResult,
Option<McpServerActsAs>,
)> {
self.check(&tool_call.name)?;
self.inner.invoke_recorded(tool_call).await
}
async fn list_server_tools(
&self,
server_prefix: &str,
session_id: uuid::Uuid,
) -> Result<Option<McpServerTools>> {
if !self.deferred_prefixes.contains(server_prefix) {
return Err(crate::runtime::AgentLoopError::tool(format!(
"MCP server '{server_prefix}' is not a deferred server in the current tool scope"
)));
}
self.inner
.list_server_tools(server_prefix, session_id)
.await
}
}
pub struct McpProxyTool {
definition: BuiltinTool,
invoker: Arc<dyn McpToolInvoker>,
}
impl McpProxyTool {
pub fn new(definition: BuiltinTool, invoker: Arc<dyn McpToolInvoker>) -> Self {
Self {
definition,
invoker,
}
}
async fn invoke(
&self,
tool_call_id: String,
arguments: Value,
identity: Option<&McpCallIdentity>,
) -> ToolExecutionResult {
let call = ToolCall {
id: tool_call_id,
name: self.definition.name.clone(),
arguments,
};
match self.invoker.invoke_recorded(&call).await {
Ok((result, acted_as)) => {
if let Some(identity) = identity {
identity.record(acted_as);
}
tool_result_to_execution(result)
}
Err(error) => ToolExecutionResult::tool_error(error.to_string()),
}
}
}
#[async_trait]
impl Tool for McpProxyTool {
fn name(&self) -> &str {
&self.definition.name
}
fn display_name(&self) -> Option<&str> {
self.definition.display_name.as_deref()
}
fn description(&self) -> &str {
&self.definition.description
}
fn parameters_schema(&self) -> Value {
self.definition.parameters.clone()
}
fn hints(&self) -> ToolHints {
self.definition.hints.clone()
}
fn requires_context(&self) -> bool {
true
}
fn to_definition(&self) -> ToolDefinition {
ToolDefinition::Builtin(self.definition.clone())
}
async fn execute(&self, arguments: Value) -> ToolExecutionResult {
self.invoke(String::new(), arguments, None).await
}
async fn execute_with_context(
&self,
arguments: Value,
context: &ToolContext,
) -> ToolExecutionResult {
let tool_call_id = context.tool_call_id.clone().unwrap_or_default();
let identity = context.extensions.get::<McpCallIdentity>();
if let Some(id) = context
.event_context
.as_ref()
.and_then(|c| c.input_message_id)
&& let Some(invoker) = self.invoker.for_execution(id.uuid())
{
let scoped = Self {
definition: self.definition.clone(),
invoker,
};
return scoped
.invoke(tool_call_id, arguments, identity.as_deref())
.await;
}
self.invoke(tool_call_id, arguments, identity.as_deref())
.await
}
}
pub fn build_mcp_proxy_tools(
definitions: &[ToolDefinition],
invoker: Arc<dyn McpToolInvoker>,
) -> Vec<Box<dyn Tool>> {
definitions
.iter()
.filter(|def| is_mcp_tool(def.name()))
.filter_map(|def| match def {
ToolDefinition::Builtin(builtin)
if crate::runtime::mcp_deferred::deferred_mcp_server_prefix(&builtin.name)
.is_some() =>
{
Some(
Box::new(crate::runtime::mcp_deferred::DeferredMcpServerTool::new(
builtin.clone(),
)) as Box<dyn Tool>,
)
}
ToolDefinition::Builtin(builtin) => {
Some(Box::new(McpProxyTool::new(builtin.clone(), invoker.clone())) as Box<dyn Tool>)
}
ToolDefinition::ClientSide(_) => None,
})
.collect()
}
fn tool_result_to_execution(result: crate::runtime::tool_types::ToolResult) -> ToolExecutionResult {
if let Some(required) = result.connection_required {
return ToolExecutionResult::ConnectionRequired {
provider: required.provider,
subject: required.subject,
setup_url: required.setup_url,
};
}
if let Some(error) = result.error {
return ToolExecutionResult::ToolError(error);
}
let value = result.result.unwrap_or(Value::Null);
match result.images {
Some(images) if !images.is_empty() => ToolExecutionResult::SuccessWithImages {
result: value,
images,
},
_ => ToolExecutionResult::Success(value),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::runtime::tool_types::{
ClientSideTool, ConnectionRequired, DeferrablePolicy, ToolPolicy, ToolResult,
};
use std::sync::Mutex;
fn builtin_def(name: &str) -> BuiltinTool {
BuiltinTool {
name: name.to_string(),
display_name: None,
description: "an mcp tool".to_string(),
parameters: serde_json::json!({
"type": "object",
"properties": { "q": { "type": "string" } }
}),
policy: ToolPolicy::Auto,
category: Some("MCP Servers".to_string()),
deferrable: DeferrablePolicy::Automatic,
hints: ToolHints::default().with_open_world(true),
full_parameters: None,
}
}
fn mcp_def(name: &str) -> ToolDefinition {
ToolDefinition::Builtin(builtin_def(name))
}
fn client_side_mcp_def(name: &str) -> ToolDefinition {
ToolDefinition::ClientSide(
ClientSideTool::new(
name,
"an mcp tool",
serde_json::json!({
"type": "object",
"properties": { "q": { "type": "string" } }
}),
)
.with_category("MCP Servers")
.with_deferrable(DeferrablePolicy::Automatic)
.with_hints(ToolHints::default().with_open_world(true)),
)
}
struct RecordingInvoker {
calls: Mutex<Vec<ToolCall>>,
result: ToolResult,
}
#[async_trait]
impl McpToolInvoker for RecordingInvoker {
async fn invoke(&self, tool_call: &ToolCall) -> Result<ToolResult> {
self.calls.lock().unwrap().push(tool_call.clone());
Ok(self.result.clone())
}
}
fn ok_result(value: Value) -> ToolResult {
ToolResult {
tool_call_id: String::new(),
result: Some(value),
images: None,
error: None,
connection_required: None,
raw_output: None,
}
}
#[tokio::test]
async fn registry_preserves_complete_mcp_definition_and_executes_context() {
let mut definition = builtin_def("mcp_docs__search");
definition.display_name = Some("Search docs".into());
definition.full_parameters = Some(
serde_json::json!({"type":"object","properties":{"q":{"type":"string","description":"query"}}}),
);
definition.deferrable = DeferrablePolicy::Always;
let expected = ToolDefinition::Builtin(definition.clone());
let definitions = [
expected.clone(),
mcp_def("read_file"),
client_side_mcp_def("mcp_secret__capture"),
];
let invoker = Arc::new(RecordingInvoker {
calls: Mutex::new(vec![]),
result: ok_result(serde_json::json!({"answer":42,"nested":[true,null]})),
});
let mut registry = crate::runtime::tools::ToolRegistry::new();
for tool in build_mcp_proxy_tools(&definitions, invoker.clone()) {
registry.register_boxed(tool);
}
assert_eq!(registry.len(), 1);
assert_eq!(
serde_json::to_value(registry.tool_definitions()).unwrap(),
serde_json::json!([expected])
);
let tool = registry.get("mcp_docs__search").unwrap();
assert_eq!(tool.display_name(), Some("Search docs"));
assert_eq!(tool.description(), "an mcp tool");
assert_eq!(tool.parameters_schema(), definition.parameters);
assert_eq!(tool.hints(), definition.hints);
assert!(tool.requires_context());
let mut context = ToolContext::new(crate::runtime::typed_id::SessionId::from_seed(1));
context.tool_call_id = Some("call-1".into());
let arguments = serde_json::json!({"q":"query","nested":{"exact":true}});
for result in [
tool.execute_with_context(arguments.clone(), &context).await,
tool.execute(arguments.clone()).await,
] {
match result {
ToolExecutionResult::Success(value) => {
assert_eq!(value, serde_json::json!({"answer":42,"nested":[true,null]}))
}
other => panic!("{other:?}"),
}
}
assert_eq!(
serde_json::to_value(&*invoker.calls.lock().unwrap()).unwrap(),
serde_json::json!([
{"id":"call-1","name":"mcp_docs__search","arguments":arguments},
{"id":"","name":"mcp_docs__search","arguments":arguments}
])
);
}
#[tokio::test]
async fn proxy_result_mapping_preserves_images_and_connection_error_precedence() {
let image = crate::runtime::tool_types::ToolResultImage {
media_type: "image/png".into(),
base64: "YWJj+/==".into(),
};
for (connection, error, images, expected) in [
(
Some("github"),
Some("boom"),
Some(vec![image.clone()]),
"connection",
),
(None, Some("boom"), Some(vec![image.clone()]), "error"),
(None, None, Some(vec![image.clone()]), "images"),
(None, None, Some(vec![]), "success"),
(None, None, None, "success"),
] {
let mut raw = ok_result(serde_json::Value::Null);
raw.result = None;
raw.connection_required = connection.map(ConnectionRequired::provider_only);
raw.error = error.map(str::to_string);
raw.images = images;
let tool = McpProxyTool::new(
builtin_def("mcp_docs__search"),
Arc::new(RecordingInvoker {
calls: Mutex::new(vec![]),
result: raw,
}),
);
match (expected, tool.execute(serde_json::json!({})).await) {
("connection", ToolExecutionResult::ConnectionRequired { provider, .. }) => {
assert_eq!(provider, "github")
}
("error", ToolExecutionResult::ToolError(message)) => assert_eq!(message, "boom"),
("images", ToolExecutionResult::SuccessWithImages { result, images }) => {
assert_eq!(result, serde_json::Value::Null);
assert_eq!(
serde_json::to_value(images).unwrap(),
serde_json::json!([{"media_type":"image/png","base64":"YWJj+/=="}])
);
}
("success", ToolExecutionResult::Success(value)) => {
assert_eq!(value, serde_json::Value::Null)
}
other => panic!("unexpected mapping: {other:?}"),
}
}
}
#[tokio::test]
async fn proxy_records_the_account_the_call_ran_as() {
struct ServiceInvoker;
#[async_trait]
impl McpToolInvoker for ServiceInvoker {
async fn invoke(&self, _call: &ToolCall) -> Result<ToolResult> {
unreachable!("the proxy asks for the recorded variant")
}
async fn invoke_recorded(
&self,
_call: &ToolCall,
) -> Result<(ToolResult, Option<McpServerActsAs>)> {
Ok((ok_result(Value::Null), Some(McpServerActsAs::Service)))
}
}
let definitions = [mcp_def("mcp_docs__search")];
let scoped = Arc::new(ScopedMcpToolInvoker::new(
&definitions,
Arc::new(ServiceInvoker),
));
let tool = McpProxyTool::new(builtin_def("mcp_docs__search"), scoped);
let identity = Arc::new(McpCallIdentity::default());
let context = ToolContext::new(crate::runtime::typed_id::SessionId::from_seed(1))
.with_extension(identity.clone());
assert!(
tool.execute_with_context(serde_json::json!({}), &context)
.await
.is_success()
);
assert_eq!(identity.get(), Some(McpServerActsAs::Service));
identity.record(Some(McpServerActsAs::User));
let plain = McpProxyTool::new(
builtin_def("mcp_docs__search"),
Arc::new(RecordingInvoker {
calls: Mutex::new(vec![]),
result: ok_result(Value::Null),
}),
);
plain
.execute_with_context(serde_json::json!({}), &context)
.await;
assert_eq!(identity.get(), None);
}
#[tokio::test]
async fn proxy_maps_invoker_error_to_tool_error() {
struct FailingInvoker;
#[async_trait]
impl McpToolInvoker for FailingInvoker {
async fn invoke(&self, _call: &ToolCall) -> Result<ToolResult> {
Err(crate::runtime::AgentLoopError::tool(
"MCP server not found for prefix: docs",
))
}
}
let tool = McpProxyTool::new(builtin_def("mcp_docs__search"), Arc::new(FailingInvoker));
match tool.execute(serde_json::json!({})).await {
ToolExecutionResult::ToolError(message) => assert_eq!(
message,
"Tool execution error: MCP server not found for prefix: docs"
),
other => panic!("{other:?}"),
}
}
#[tokio::test]
async fn scoped_invoker_rejects_unlisted_non_mcp_and_client_side_tools_before_backend() {
let expected = ok_result(serde_json::json!({"ok":true,"data":[1,2]}));
let inner = Arc::new(RecordingInvoker {
calls: Mutex::new(vec![]),
result: expected.clone(),
});
let scoped = ScopedMcpToolInvoker::new(
&[
mcp_def("mcp_docs__search"),
mcp_def("read_file"),
client_side_mcp_def("mcp_secret__capture"),
],
inner.clone(),
);
for name in [
"mcp_other__search",
"read_file",
"mcp_secret__capture",
"mcp_docs__search_extra",
] {
let call = ToolCall {
id: "denied".into(),
name: name.into(),
arguments: serde_json::json!({"q":"private"}),
};
assert_eq!(
scoped.invoke(&call).await.unwrap_err().to_string(),
format!(
"Tool execution error: MCP tool '{name}' is not allowed in the current tool scope"
)
);
}
assert!(inner.calls.lock().unwrap().is_empty());
let allowed = ToolCall {
id: "allowed".into(),
name: "mcp_docs__search".into(),
arguments: serde_json::json!({"q":"exact","nested":[true]}),
};
let result = scoped.invoke(&allowed).await.unwrap();
assert_eq!(
serde_json::to_value(result).unwrap(),
serde_json::to_value(expected).unwrap()
);
assert_eq!(
serde_json::to_value(&*inner.calls.lock().unwrap()).unwrap(),
serde_json::json!([allowed])
);
}
#[tokio::test]
async fn a_deferred_placeholder_scopes_its_whole_server() {
let placeholder =
crate::runtime::mcp_deferred::deferred_mcp_server_definition("docs", None);
let inner = Arc::new(RecordingInvoker {
calls: Mutex::new(vec![]),
result: ok_result(Value::Null),
});
let scoped = ScopedMcpToolInvoker::new(&[placeholder, mcp_def("mcp_wiki__read")], inner);
let call = |name: &str| ToolCall {
id: String::new(),
name: name.to_string(),
arguments: serde_json::json!({}),
};
assert!(scoped.invoke(&call("mcp_docs__search")).await.is_ok());
assert!(scoped.invoke(&call("mcp_wiki__read")).await.is_ok());
assert!(scoped.invoke(&call("mcp_wiki__write")).await.is_err());
assert!(scoped.invoke(&call("mcp_other__search")).await.is_err());
let session = uuid::Uuid::nil();
assert!(
scoped
.list_server_tools("docs", session)
.await
.unwrap()
.is_none()
);
assert!(scoped.list_server_tools("wiki", session).await.is_err());
}
}