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#[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 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 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 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(¶m, 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(¶m, 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(¶m, 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(¶m, 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(¶m, 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(¶m).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(¶m).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(¶m).await.unwrap();
470 let first_output_id = first.list_tools_item.id.clone();
471 let second = executors.mcp_server_tools(¶m).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(¶m).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}