Skip to main content

agentic_core/tool/
executors.rs

1use std::collections::HashMap;
2use std::sync::Arc;
3
4use tokio::sync::RwLock;
5
6use super::mcp::handler::McpServerToolSet;
7use super::mcp::{McpClientPool, McpDiscoveredHandler, McpHandler};
8use super::registry::ToolType;
9use super::web_search::WebSearchHandler;
10use super::{GatewayExecutor, ToolError};
11use crate::config::ToolRuntimeConfig;
12use crate::types::tools::McpToolParam;
13
14pub enum GatewayExecutorRegistration {
15    Shared(Arc<dyn GatewayExecutor>),
16    Mcp {
17        server_label: String,
18        handlers: Vec<McpDiscoveredHandler>,
19    },
20}
21
22impl<T> From<Arc<T>> for GatewayExecutorRegistration
23where
24    T: GatewayExecutor,
25{
26    fn from(executor: Arc<T>) -> Self {
27        Self::Shared(executor)
28    }
29}
30
31impl From<Arc<dyn GatewayExecutor>> for GatewayExecutorRegistration {
32    fn from(executor: Arc<dyn GatewayExecutor>) -> Self {
33        Self::Shared(executor)
34    }
35}
36
37/// Shared, per-server registry of gateway-owned tool executors.
38///
39/// Built once at startup ([`GatewayExecutors::from_env`]) and reused across
40/// every request. Configured and request-declared MCP servers are discovered
41/// lazily unless handlers were registered with [`Self::insert`].
42#[derive(Clone, Default)]
43pub struct GatewayExecutors {
44    mcp: HashMap<String, Vec<McpDiscoveredHandler>>,
45    mcp_configs: HashMap<String, super::mcp::McpServerEntry>,
46    mcp_clients: Arc<RwLock<HashMap<String, Arc<super::mcp::McpClient>>>>,
47    mcp_discovered: Arc<RwLock<HashMap<String, Vec<McpDiscoveredHandler>>>>,
48    mcp_allowed_hosts: Vec<String>,
49    web_search: Option<Arc<dyn GatewayExecutor>>,
50}
51
52impl GatewayExecutors {
53    #[must_use]
54    pub fn from_env(client: Arc<reqwest::Client>) -> Self {
55        Self {
56            mcp: HashMap::new(),
57            mcp_configs: HashMap::new(),
58            mcp_clients: Arc::new(RwLock::new(HashMap::new())),
59            mcp_discovered: Arc::new(RwLock::new(HashMap::new())),
60            mcp_allowed_hosts: super::mcp::pool::allowed_hosts_from_env(),
61            web_search: Some(Arc::new(WebSearchHandler::from_env(client))),
62        }
63    }
64
65    /// Builds the shared executors without contacting configured MCP servers.
66    ///
67    /// Configured `allowed_tools` are applied during discovery, so the stored
68    /// handler set is the maximum set that a request may use.
69    ///
70    /// # Errors
71    ///
72    /// Invalid policy configuration is returned as an error. Connection and
73    /// discovery happen when a configured server is requested.
74    pub fn from_config(client: Arc<reqwest::Client>, config: &ToolRuntimeConfig) -> Result<Self, ToolError> {
75        let executors = Self {
76            mcp: HashMap::new(),
77            mcp_configs: config.mcp_servers.clone(),
78            mcp_clients: Arc::new(RwLock::new(HashMap::new())),
79            mcp_discovered: Arc::new(RwLock::new(HashMap::new())),
80            mcp_allowed_hosts: if config.mcp_allowed_hosts.is_empty() {
81                super::mcp::pool::allowed_hosts_from_env()
82            } else {
83                config.mcp_allowed_hosts.clone()
84            },
85            web_search: Some(Arc::new(WebSearchHandler::from_values(
86                client,
87                config.web_search.api_key.clone(),
88                config.web_search.base_url.clone(),
89            ))),
90        };
91        if config.mcp_servers.is_empty() {
92            return Ok(executors);
93        }
94
95        for (server_label, entry) in &config.mcp_servers {
96            if entry.require_approval() != Some("never") {
97                return Err(ToolError::Config(format!(
98                    "configured MCP server '{server_label}' must set require_approval to 'never'"
99                )));
100            }
101        }
102
103        Ok(executors)
104    }
105
106    pub fn insert(&mut self, registration: impl Into<GatewayExecutorRegistration>) {
107        match registration.into() {
108            GatewayExecutorRegistration::Shared(executor) => match executor.tool_type() {
109                ToolType::WebSearch => self.web_search = Some(executor),
110                ToolType::Mcp => {
111                    tracing::debug!("MCP executors must be registered with a server_label and discovered handlers");
112                }
113                other => tracing::debug!(tool_type = ?other, "gateway executor type has no executor slot"),
114            },
115            GatewayExecutorRegistration::Mcp { server_label, handlers } => {
116                if handlers.is_empty() {
117                    tracing::debug!(server_label, "empty MCP discovered handler registration skipped");
118                    return;
119                }
120                if self.mcp.insert(server_label.clone(), handlers).is_some() {
121                    tracing::debug!(server_label, "replaced MCP discovered handler registration");
122                }
123            }
124        }
125    }
126
127    #[must_use]
128    pub fn web_search_handler(&self) -> Option<Arc<dyn GatewayExecutor>> {
129        self.web_search.clone()
130    }
131
132    #[must_use]
133    pub(crate) fn request_scoped(&self) -> Self {
134        self.clone()
135    }
136
137    /// Returns the discovered handlers for one request-declared MCP server.
138    ///
139    /// # Errors
140    ///
141    /// Returns a configuration error for an invalid declaration or an empty
142    /// allowed tool set, and an execution error when the server cannot connect.
143    pub async fn mcp_handler(&mut self, param: &McpToolParam) -> Result<Vec<McpDiscoveredHandler>, ToolError> {
144        Ok(self.mcp_server_tools(param).await?.discovered_handlers)
145    }
146
147    /// Returns the request-scoped tools and public discovery item for one MCP server.
148    ///
149    /// # Errors
150    ///
151    /// Returns a configuration error for an invalid declaration or an empty
152    /// allowed tool set, and an execution error when the server cannot connect.
153    pub(crate) async fn mcp_server_tools(&mut self, param: &McpToolParam) -> Result<McpServerToolSet, ToolError> {
154        let server_label = param.server_label.trim();
155        if server_label.is_empty() {
156            return Err(ToolError::Config(
157                "MCP declaration requires a non-empty server_label".to_owned(),
158            ));
159        }
160        let configured_handlers = self.mcp.get(server_label);
161        let configured_server = self.mcp_configs.contains_key(server_label);
162        validate_mcp_execution_options(param, configured_server || configured_handlers.is_some())?;
163        if (configured_server || configured_handlers.is_some()) && param.server_url.is_some() {
164            return Err(ToolError::Config(format!(
165                "MCP server '{server_label}' is configured by the gateway; omit server_url from the request"
166            )));
167        }
168        if let Some(configured_handlers) = configured_handlers {
169            let discovered_handlers = require_non_empty_mcp_handlers(
170                server_label,
171                filter_allowed_mcp_handlers(configured_handlers, param.allowed_tools.as_deref()),
172            )?;
173            return Ok(McpHandler::server_tool_set_from_handlers(
174                server_label,
175                discovered_handlers,
176            ));
177        }
178
179        if configured_server {
180            let Some(entry) = self.mcp_configs.get(server_label).cloned() else {
181                return Err(ToolError::Config(format!(
182                    "configured MCP server '{server_label}' is missing"
183                )));
184            };
185            let cached_client = self.mcp_clients.read().await.get(server_label).cloned();
186            let client = if let Some(client) = cached_client {
187                client
188            } else {
189                let mut servers = HashMap::new();
190                servers.insert(server_label.to_owned(), entry.clone());
191                let pool = McpClientPool::from_config(servers).await;
192                let Some(client) = pool.get(server_label).cloned() else {
193                    return Err(ToolError::Execution(format!(
194                        "configured MCP server '{server_label}' failed to connect: {}",
195                        pool.connection_error(server_label)
196                            .unwrap_or("unknown connection error")
197                    )));
198                };
199                self.mcp_clients
200                    .write()
201                    .await
202                    .insert(server_label.to_owned(), Arc::clone(&client));
203                client
204            };
205            let discovered_handlers = if let Some(discovered_handlers) =
206                self.mcp_discovered.read().await.get(server_label).cloned()
207            {
208                discovered_handlers
209            } else {
210                let tool_set = McpHandler::discover_tools(server_label, client, entry.allowed_tools()).await?;
211                let discovered_handlers = require_non_empty_mcp_handlers(server_label, tool_set.discovered_handlers)?;
212                self.mcp_discovered
213                    .write()
214                    .await
215                    .insert(server_label.to_owned(), discovered_handlers.clone());
216                discovered_handlers
217            };
218            let discovered_handlers = require_non_empty_mcp_handlers(
219                server_label,
220                filter_allowed_mcp_handlers(&discovered_handlers, param.allowed_tools.as_deref()),
221            )?;
222            return Ok(McpHandler::server_tool_set_from_handlers(
223                server_label,
224                discovered_handlers,
225            ));
226        }
227
228        let pool =
229            McpClientPool::from_params_with_allowed_hosts(std::slice::from_ref(param), &self.mcp_allowed_hosts).await;
230        let Some(client) = pool.get(server_label).cloned() else {
231            return Err(pool.connection_error(server_label).map_or_else(
232                || {
233                    ToolError::Config(format!(
234                        "MCP server '{server_label}' has no valid request-declared configuration"
235                    ))
236                },
237                |error| ToolError::Execution(format!("MCP server '{server_label}' failed to connect: {error}")),
238            ));
239        };
240        let tool_set = McpHandler::discover_tools(server_label, client, param.allowed_tools.as_deref()).await?;
241        let discovered_handlers = require_non_empty_mcp_handlers(server_label, tool_set.discovered_handlers)?;
242        self.mcp.insert(server_label.to_owned(), discovered_handlers.clone());
243        Ok(McpHandler::server_tool_set_from_handlers(
244            server_label,
245            discovered_handlers,
246        ))
247    }
248}
249
250fn filter_allowed_mcp_handlers(
251    handlers: &[McpDiscoveredHandler],
252    allowed_tools: Option<&[String]>,
253) -> Vec<McpDiscoveredHandler> {
254    handlers
255        .iter()
256        .filter(|handler| {
257            allowed_tools.is_none_or(|allowed| allowed.iter().any(|name| name == &handler.param.tool_name))
258        })
259        .cloned()
260        .collect()
261}
262
263fn require_non_empty_mcp_handlers(
264    server_label: &str,
265    handlers: Vec<McpDiscoveredHandler>,
266) -> Result<Vec<McpDiscoveredHandler>, ToolError> {
267    if handlers.is_empty() {
268        return Err(ToolError::Config(format!(
269            "MCP server '{server_label}' has an empty final allowed tool set"
270        )));
271    }
272    Ok(handlers)
273}
274
275fn validate_mcp_execution_options(param: &McpToolParam, configured_server: bool) -> Result<(), ToolError> {
276    if param.connector_id.is_some() {
277        return Err(ToolError::Config(
278            "MCP connector_id is not supported; configure server_url instead".to_owned(),
279        ));
280    }
281    if param
282        .require_approval
283        .as_deref()
284        .is_some_and(|policy| policy != "never")
285    {
286        return Err(ToolError::Config(
287            "MCP require_approval supports only 'never'; approval gating is not yet supported".to_owned(),
288        ));
289    }
290    if !configured_server && param.require_approval.is_none() {
291        return Err(ToolError::Config(
292            "MCP require_approval must be set to 'never' in gateway configuration or the request".to_owned(),
293        ));
294    }
295    Ok(())
296}
297
298impl std::fmt::Debug for GatewayExecutors {
299    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
300        f.debug_struct("GatewayExecutors")
301            .field("mcp_server_handlers", &self.mcp.len())
302            .field("mcp_server_configs", &self.mcp_configs.len())
303            .field("mcp_clients", &Arc::strong_count(&self.mcp_clients))
304            .field("mcp_discovered", &Arc::strong_count(&self.mcp_discovered))
305            .field("mcp_allowed_hosts", &self.mcp_allowed_hosts)
306            .field("web_search", &self.web_search.is_some())
307            .finish()
308    }
309}
310
311#[cfg(test)]
312mod tests {
313    use std::collections::HashMap;
314    use std::sync::Arc;
315
316    use super::{GatewayExecutorRegistration, GatewayExecutors, validate_mcp_execution_options};
317    use crate::config::ToolRuntimeConfig;
318    use crate::tool::mcp::McpServerEntry;
319    use crate::tool::mcp::{McpDiscoveredHandler, McpHandler};
320    use crate::types::tools::{McpDiscoveredToolParam, McpToolParam};
321
322    fn mcp_param(value: serde_json::Value) -> McpToolParam {
323        serde_json::from_value(value).unwrap()
324    }
325
326    fn discovered_handler(tool_name: &str) -> McpDiscoveredHandler {
327        McpDiscoveredHandler {
328            param: McpDiscoveredToolParam {
329                server_label: "counter".to_owned(),
330                tool_name: tool_name.to_owned(),
331                internal_name: format!("mcp__counter__{tool_name}"),
332                tool: serde_json::from_value(serde_json::json!({
333                    "name": tool_name,
334                    "inputSchema": {"type": "object"}
335                }))
336                .unwrap(),
337            },
338            handler: Arc::new(McpHandler::discovered_tool_spec_only()),
339        }
340    }
341
342    #[test]
343    fn mcp_execution_allows_explicit_never_approval_policy() {
344        let param = mcp_param(serde_json::json!({
345            "server_label": "counter",
346            "server_url": "http://localhost:8000/mcp",
347            "require_approval": "never"
348        }));
349
350        validate_mcp_execution_options(&param, false).unwrap();
351    }
352
353    #[test]
354    fn mcp_execution_uses_configured_never_approval_policy() {
355        let param = mcp_param(serde_json::json!({
356            "server_label": "counter"
357        }));
358
359        validate_mcp_execution_options(&param, true).unwrap();
360    }
361
362    #[test]
363    fn mcp_execution_rejects_omitted_approval_policy() {
364        let param = mcp_param(serde_json::json!({
365            "server_label": "counter",
366            "server_url": "http://localhost:8000/mcp"
367        }));
368
369        let error = validate_mcp_execution_options(&param, false).unwrap_err();
370        assert!(error.to_string().contains("gateway configuration or the request"));
371    }
372
373    #[test]
374    fn mcp_execution_rejects_unsupported_approval_policy() {
375        let param = mcp_param(serde_json::json!({
376            "server_label": "counter",
377            "server_url": "http://localhost:8000/mcp",
378            "require_approval": "always"
379        }));
380
381        let error = validate_mcp_execution_options(&param, false).unwrap_err();
382        assert!(error.to_string().contains("approval gating is not yet supported"));
383    }
384
385    #[test]
386    fn mcp_execution_rejects_connector_id() {
387        let param = mcp_param(serde_json::json!({
388            "server_label": "counter",
389            "connector_id": "connector_dropbox"
390        }));
391
392        let error = validate_mcp_execution_options(&param, false).unwrap_err();
393        assert!(error.to_string().contains("connector_id is not supported"));
394    }
395
396    #[tokio::test]
397    async fn configured_mcp_server_rejects_request_connection_override() {
398        let mut executors = GatewayExecutors::default();
399        executors.insert(GatewayExecutorRegistration::Mcp {
400            server_label: "counter".to_owned(),
401            handlers: vec![discovered_handler("read")],
402        });
403        let param = mcp_param(serde_json::json!({
404            "server_label": "counter",
405            "server_url": "http://localhost:8000/mcp",
406            "require_approval": "never"
407        }));
408
409        let Err(error) = executors.mcp_server_tools(&param).await else {
410            panic!("request connection override must fail");
411        };
412        assert!(error.to_string().contains("configured by the gateway"));
413        assert!(error.to_string().contains("omit server_url"));
414    }
415
416    #[tokio::test]
417    async fn unavailable_configured_mcp_server_does_not_block_startup() {
418        let mut servers = HashMap::new();
419        servers.insert(
420            "unavailable".to_owned(),
421            McpServerEntry::Http {
422                url: "http://127.0.0.1:1/mcp".to_owned(),
423                headers: None,
424                allowed_tools: Some(vec!["read".to_owned()]),
425                require_approval: Some("never".to_owned()),
426            },
427        );
428        let config = ToolRuntimeConfig {
429            mcp_servers: servers,
430            ..ToolRuntimeConfig::default()
431        };
432
433        let executors = GatewayExecutors::from_config(Arc::new(reqwest::Client::new()), &config);
434
435        assert!(executors.is_ok());
436    }
437
438    #[tokio::test]
439    async fn configured_allowed_tools_cannot_be_expanded_by_request() {
440        let mut executors = GatewayExecutors::default();
441        executors.insert(GatewayExecutorRegistration::Mcp {
442            server_label: "counter".to_owned(),
443            handlers: vec![discovered_handler("read")],
444        });
445        let param = mcp_param(serde_json::json!({
446            "server_label": "counter",
447            "allowed_tools": ["read", "delete"]
448        }));
449
450        let tools = executors.mcp_server_tools(&param).await.unwrap();
451
452        assert_eq!(tools.discovered_handlers.len(), 1);
453        assert_eq!(tools.discovered_handlers[0].param.tool_name, "read");
454    }
455
456    #[tokio::test]
457    async fn cached_mcp_server_tools_apply_request_allowed_tools_with_fresh_output_id() {
458        let mut executors = GatewayExecutors::default();
459        executors.insert(GatewayExecutorRegistration::Mcp {
460            server_label: "counter".to_owned(),
461            handlers: vec![discovered_handler("read"), discovered_handler("delete")],
462        });
463        let param = mcp_param(serde_json::json!({
464            "server_label": "counter",
465            "allowed_tools": ["read"],
466            "require_approval": "never"
467        }));
468
469        let first = executors.mcp_server_tools(&param).await.unwrap();
470        let first_output_id = first.list_tools_item.id.clone();
471        let second = executors.mcp_server_tools(&param).await.unwrap();
472
473        assert_eq!(first.discovered_handlers.len(), 1);
474        assert_eq!(first.discovered_handlers[0].param.tool_name, "read");
475        assert_eq!(first.list_tools_item.tools.len(), 1);
476        assert_eq!(first.list_tools_item.tools[0].name, "read");
477        assert_ne!(first_output_id, second.list_tools_item.id);
478    }
479
480    #[tokio::test]
481    async fn cached_mcp_handlers_reject_empty_final_allowed_set() {
482        let mut executors = GatewayExecutors::default();
483        executors.insert(GatewayExecutorRegistration::Mcp {
484            server_label: "counter".to_owned(),
485            handlers: vec![discovered_handler("delete")],
486        });
487        let param = mcp_param(serde_json::json!({
488            "server_label": "counter",
489            "allowed_tools": ["read"],
490            "require_approval": "never"
491        }));
492
493        let Err(error) = executors.mcp_server_tools(&param).await else {
494            panic!("expected empty allowed set to be rejected");
495        };
496
497        assert!(error.to_string().contains("empty final allowed tool set"));
498    }
499}