Skip to main content

agentic_core/tool/
registry.rs

1use std::collections::HashMap;
2use std::collections::hash_map::Entry;
3use std::sync::Arc;
4
5use serde::{Deserialize, Serialize};
6use serde_json::Value;
7
8use super::codex::insert_namespace_entries;
9use super::custom::{CustomHandler, CustomToolMap, insert_custom_entry};
10use super::executors::GatewayExecutors;
11use super::function::insert_function_entry;
12use super::mcp::handler::{McpToolMap, McpToolRef};
13use super::mcp::registry::insert_discovered_mcp_entry;
14use super::web_search::insert_web_search_entry;
15use super::{CodexNamespaceHandler, GatewayExecutor, McpHandler, NamespaceMap, ToolError, ToolOutput};
16use crate::events::WireEvent;
17
18use crate::types::io::OutputItem;
19use crate::types::io::output::{FunctionToolCall, McpListTools};
20use crate::types::tools::{CodeInterpreterToolParam, FileSearchToolParam, ResponsesTool};
21use crate::utils::common::serialize_to_value_or_custom_default;
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
24#[serde(rename_all = "snake_case")]
25pub enum ToolType {
26    Function,
27    Custom,
28    CodexNamespace,
29    Mcp,
30    /// Internal routing discriminant. Serializes as `"web_search"`.
31    /// Note: the corresponding `ResponsesTool` wire tag is `"web_search_preview"`.
32    /// `ToolType` is not used in wire-facing types so the names differ intentionally.
33    WebSearch,
34    FileSearch,
35    CodeInterpreter,
36}
37
38impl ToolType {
39    #[must_use]
40    pub(crate) const fn description(self) -> &'static str {
41        match self {
42            Self::Function => "function tool",
43            Self::Custom => "custom tool",
44            Self::CodexNamespace => "Codex namespace tool",
45            Self::Mcp => "MCP tool",
46            Self::WebSearch => "web search tool",
47            Self::FileSearch => "file search tool",
48            Self::CodeInterpreter => "code interpreter tool",
49        }
50    }
51
52    #[must_use]
53    pub const fn is_gateway_owned(self) -> bool {
54        !matches!(self, Self::Function | Self::Custom | Self::CodexNamespace)
55    }
56}
57
58/// Per-request routing entry keyed by the tool name the model will call.
59#[derive(Clone)]
60pub struct ToolEntry {
61    pub tool_type: ToolType,
62    /// Full serialised tool param for the executor (used during dispatch).
63    pub config: Value,
64    /// For MCP tools: which server this tool belongs to.
65    pub server_label: Option<String>,
66    pub handler: Option<Arc<dyn GatewayExecutor>>,
67}
68
69impl std::fmt::Debug for ToolEntry {
70    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
71        f.debug_struct("ToolEntry")
72            .field("tool_type", &self.tool_type)
73            .field("config", &self.config)
74            .field("server_label", &self.server_label)
75            .field("handler", &self.handler.is_some())
76            .finish()
77    }
78}
79
80fn insert_unique_tool_entries(
81    entries: &mut HashMap<String, ToolEntry>,
82    insert: impl FnOnce(&mut HashMap<String, ToolEntry>),
83) -> Result<(), ToolError> {
84    let mut resolved = HashMap::new();
85    insert(&mut resolved);
86    for (name, entry) in resolved {
87        match entries.entry(name) {
88            Entry::Occupied(existing) => {
89                return Err(ToolError::Config(format!(
90                    "{} registry name '{}' conflicts with existing {}",
91                    entry.tool_type.description(),
92                    existing.key(),
93                    existing.get().tool_type.description()
94                )));
95            }
96            Entry::Vacant(vacant) => {
97                vacant.insert(entry);
98            }
99        }
100    }
101    Ok(())
102}
103
104pub struct GatewayDispatchResult {
105    pub tool_type: ToolType,
106    pub output: Result<ToolOutput, ToolError>,
107}
108
109// TODO: move to a dedicated file_search module alongside its `ToolHandler`
110// once file_search execution is implemented.
111fn insert_file_search_entry(
112    entries: &mut HashMap<String, ToolEntry>,
113    p: &FileSearchToolParam,
114    handler: Option<Arc<dyn GatewayExecutor>>,
115) {
116    serialize_to_value_or_custom_default(
117        p,
118        "file_search tool config serialization failed",
119        |config| {
120            entries.insert(
121                "file_search".to_owned(),
122                ToolEntry {
123                    tool_type: ToolType::FileSearch,
124                    config,
125                    server_label: None,
126                    handler,
127                },
128            );
129        },
130        (),
131    );
132}
133
134// TODO: move to a dedicated code_interpreter module alongside its `ToolHandler`
135// once code_interpreter execution is implemented.
136fn insert_code_interpreter_entry(
137    entries: &mut HashMap<String, ToolEntry>,
138    p: &CodeInterpreterToolParam,
139    handler: Option<Arc<dyn GatewayExecutor>>,
140) {
141    serialize_to_value_or_custom_default(
142        p,
143        "code_interpreter tool config serialization failed",
144        |config| {
145            entries.insert(
146                "code_interpreter".to_owned(),
147                ToolEntry {
148                    tool_type: ToolType::CodeInterpreter,
149                    config,
150                    server_label: None,
151                    handler,
152                },
153            );
154        },
155        (),
156    );
157}
158
159/// Request-scoped registry built from `RequestPayload.tools`.
160/// Maps the name the LLM sees → routing metadata.
161#[derive(Debug, Default)]
162pub struct ToolRegistry {
163    entries: HashMap<String, ToolEntry>,
164
165    /// Built once from the declared tools, so final payload and streaming event
166    /// restoration don't rebuild it on every call.
167    namespace_map: Option<NamespaceMap>,
168
169    /// Maps normalized custom function names back to their public declarations
170    /// for response lifecycle metadata restoration.
171    custom_tool_map: Option<CustomToolMap>,
172
173    /// Maps model-visible MCP function names back to their public server and
174    /// tool identities without reparsing executor configuration.
175    mcp_tool_map: McpToolMap,
176
177    /// Request-scoped MCP discovery output items retained in declaration order.
178    mcp_list_tools_items: Vec<McpListTools>,
179}
180
181impl ToolRegistry {
182    /// Build a registry from declared tools and attach gateway handlers for dispatchable tool types.
183    ///
184    /// # Errors
185    ///
186    /// Returns [`ToolError::Config`] when Codex namespace member flattening
187    /// would collide with another declared tool name, or when discovered MCP
188    /// tools derive the same internal model-visible name.
189    ///
190    /// # Panics
191    ///
192    /// Panics if serialization of a tool param struct fails, which cannot happen
193    /// for the types defined in this module (`#[derive(Serialize)]` on plain structs).
194    pub async fn build_with_handlers(
195        tools: &mut [ResponsesTool],
196        executors: &mut GatewayExecutors,
197    ) -> Result<Self, ToolError> {
198        let mut entries = HashMap::with_capacity(tools.len());
199        let mut mcp_tool_map = McpToolMap::default();
200        let mut mcp_list_tools_items = Vec::new();
201        // Namespace members must be keyed by the same flat, model-visible name
202        // the model will call, so resolve them first — the same pure pass used
203        // to build the upstream request.
204        let resolved_tools = CodexNamespaceHandler.resolve_namespace_members(tools)?;
205        McpHandler::validate_server_labels(&resolved_tools)?;
206
207        for (index, tool) in resolved_tools.iter().enumerate() {
208            match tool {
209                ResponsesTool::Function(p) => {
210                    insert_unique_tool_entries(&mut entries, |resolved| insert_function_entry(resolved, p))?;
211                }
212                ResponsesTool::Mcp(p) => {
213                    let tool_set = match executors.mcp_server_tools(p).await {
214                        Ok(tool_set) => tool_set,
215                        Err(error) => {
216                            mcp_list_tools_items.push(McpHandler::failed_list_tools_item(&p.server_label, &error));
217                            continue;
218                        }
219                    };
220                    let handlers = tool_set.discovered_handlers;
221                    mcp_list_tools_items.push(tool_set.list_tools_item);
222                    if let ResponsesTool::Mcp(declaration) = &mut tools[index] {
223                        declaration.discovered_tools = handlers.iter().map(|item| item.param.clone()).collect();
224                    }
225                    for discovered in handlers {
226                        let internal_name = discovered.param.internal_name.clone();
227                        let tool_ref = McpToolRef::from(&discovered.param);
228                        insert_unique_tool_entries(&mut entries, |resolved| {
229                            insert_discovered_mcp_entry(resolved, discovered);
230                        })?;
231                        mcp_tool_map.record(internal_name, tool_ref);
232                    }
233                }
234                ResponsesTool::WebSearch(p) => {
235                    insert_unique_tool_entries(&mut entries, |resolved| {
236                        insert_web_search_entry(resolved, p, executors.web_search_handler());
237                    })?;
238                }
239                ResponsesTool::FileSearch(p) => {
240                    insert_unique_tool_entries(&mut entries, |resolved| insert_file_search_entry(resolved, p, None))?;
241                }
242                ResponsesTool::CodeInterpreter(p) => {
243                    insert_unique_tool_entries(&mut entries, |resolved| {
244                        insert_code_interpreter_entry(resolved, p, None);
245                    })?;
246                }
247                ResponsesTool::Namespace(p) => {
248                    insert_unique_tool_entries(&mut entries, |resolved| insert_namespace_entries(resolved, p))?;
249                }
250                ResponsesTool::Custom(p) => {
251                    insert_unique_tool_entries(&mut entries, |resolved| insert_custom_entry(resolved, p))?;
252                }
253                ResponsesTool::Unknown => {
254                    tracing::debug!("unknown tool declared but skipped in registry");
255                }
256            }
257        }
258
259        let namespace_map = CodexNamespaceHandler.build_namespace_map((!tools.is_empty()).then_some(tools))?;
260        let custom_tool_map = CustomHandler::build_tool_map(tools);
261
262        Ok(Self {
263            entries,
264            namespace_map,
265            custom_tool_map,
266            mcp_tool_map,
267            mcp_list_tools_items,
268        })
269    }
270
271    #[must_use]
272    pub fn lookup(&self, tool_name: &str) -> Option<&ToolEntry> {
273        self.entries.get(tool_name)
274    }
275
276    pub(crate) fn tool_type_map(&self) -> HashMap<String, ToolType> {
277        self.entries
278            .iter()
279            .map(|(name, entry)| (name.clone(), entry.tool_type))
280            .collect()
281    }
282
283    #[must_use]
284    pub fn is_empty(&self) -> bool {
285        self.entries.is_empty()
286    }
287
288    #[must_use]
289    pub fn len(&self) -> usize {
290        self.entries.len()
291    }
292
293    #[must_use]
294    pub fn contains_mcp_server_label(&self, server_label: &str) -> bool {
295        self.mcp_tool_map.contains_server_label(server_label)
296    }
297
298    pub(crate) fn mcp_tool_ref(&self, internal_name: &str) -> Option<&McpToolRef> {
299        self.mcp_tool_map.tool_ref(internal_name)
300    }
301
302    #[must_use]
303    pub(crate) fn mcp_list_tools_items(&self) -> &[McpListTools] {
304        &self.mcp_list_tools_items
305    }
306
307    pub fn restore_final_payload_output(&self, output: &mut [OutputItem]) {
308        CodexNamespaceHandler.restore_output_items(output, self.namespace_map.as_ref());
309    }
310
311    pub fn restore_stream_event_wire(&self, wire: &mut WireEvent) -> bool {
312        let custom_restored = CustomHandler::restore_response_wire(wire, self.custom_tool_map.as_ref());
313        CodexNamespaceHandler.restore_response_wire(wire, self.namespace_map.as_ref()) | custom_restored
314    }
315
316    /// Returns the subset of `calls` whose names map to gateway-owned tools.
317    #[must_use]
318    pub fn gateway_owned<'a>(&self, calls: &'a [FunctionToolCall]) -> Vec<&'a FunctionToolCall> {
319        calls
320            .iter()
321            .filter(|c| {
322                self.entries
323                    .get(&c.name)
324                    .is_some_and(|e| e.tool_type.is_gateway_owned())
325            })
326            .collect()
327    }
328
329    #[must_use]
330    pub fn is_gateway_owned_name(&self, name: &str) -> bool {
331        self.entries
332            .get(name)
333            .is_some_and(|entry| entry.tool_type.is_gateway_owned())
334    }
335
336    /// Returns the subset of `calls` whose names map to client-owned tools
337    /// (`Function`, Codex namespace members, or unknown names).
338    #[must_use]
339    pub fn client_owned<'a>(&self, calls: &'a [FunctionToolCall]) -> Vec<&'a FunctionToolCall> {
340        calls
341            .iter()
342            .filter(|c| {
343                self.entries
344                    .get(&c.name)
345                    .is_none_or(|e| !e.tool_type.is_gateway_owned())
346            })
347            .collect()
348    }
349
350    pub async fn dispatch(&self, call: &FunctionToolCall) -> Option<GatewayDispatchResult> {
351        let entry = self.entries.get(&call.name)?;
352        let handler = entry.handler.clone()?;
353        let tool_type = entry.tool_type;
354        let config = entry.config.clone();
355        Some(GatewayDispatchResult {
356            tool_type,
357            output: handler
358                .execute(&call.call_id, &call.name, &call.arguments, &config)
359                .await,
360        })
361    }
362}
363
364#[cfg(test)]
365mod tests {
366    use super::*;
367    use crate::tool::executors::GatewayExecutorRegistration;
368    use crate::tool::mcp::{McpDiscoveredHandler, McpHandler};
369    use crate::types::event::MessageStatus;
370    use crate::types::tools::McpDiscoveredToolParam;
371
372    fn declaration(server_label: &str) -> ResponsesTool {
373        serde_json::from_value(serde_json::json!({
374            "type": "mcp",
375            "server_label": server_label,
376            "server_url": "http://127.0.0.1:8000/mcp",
377            "require_approval": "never"
378        }))
379        .expect("MCP declaration")
380    }
381
382    fn discovered_handler(server_label: &str, tool_name: &str, internal_name: &str) -> McpDiscoveredHandler {
383        let param = McpDiscoveredToolParam {
384            server_label: server_label.to_owned(),
385            tool_name: tool_name.to_owned(),
386            internal_name: internal_name.to_owned(),
387            tool: serde_json::from_value(serde_json::json!({
388                "name": tool_name,
389                "description": "Discovered test tool",
390                "inputSchema": {"type": "object"}
391            }))
392            .expect("discovered MCP tool"),
393        };
394        McpDiscoveredHandler {
395            param,
396            handler: Arc::new(McpHandler::discovered_tool_spec_only()),
397        }
398    }
399
400    fn mixed_tool_declarations() -> Vec<ResponsesTool> {
401        serde_json::from_value(serde_json::json!([
402            {
403                "type": "function",
404                "name": "echo",
405                "parameters": {"type": "object"}
406            },
407            {
408                "type": "mcp",
409                "server_label": "counter",
410                "server_url": "http://127.0.0.1:8000/mcp",
411                "require_approval": "never"
412            },
413            {"type": "web_search_preview", "search_context_size": "low"},
414            {"type": "file_search", "vector_store_ids": ["vs_test"]},
415            {"type": "code_interpreter"},
416            {
417                "type": "namespace",
418                "name": "mcp__shell",
419                "tools": [{"type": "function", "name": "run"}]
420            },
421            {"type": "custom", "name": "freeform"},
422            {"type": "future_tool", "opaque": true}
423        ]))
424        .expect("mixed tool declarations")
425    }
426
427    fn assert_namespace_call_restoration(registry: &ToolRegistry) {
428        let mut output = vec![OutputItem::FunctionCall(FunctionToolCall {
429            id: "fc_1".to_owned(),
430            call_id: "call_1".to_owned(),
431            name: "agentic_ns__mcp__shell__run".to_owned(),
432            namespace: None,
433            arguments: "{}".to_owned(),
434            status: MessageStatus::Completed,
435        })];
436        registry.restore_final_payload_output(&mut output);
437        let OutputItem::FunctionCall(call) = &output[0] else {
438            panic!("expected restored function call");
439        };
440        assert_eq!(call.namespace.as_deref(), Some("mcp__shell"));
441        assert_eq!(call.name, "run");
442    }
443
444    fn assert_mcp_list_tools_metadata(registry: &ToolRegistry) {
445        let [list_tools] = registry.mcp_list_tools_items() else {
446            panic!("expected one MCP list-tools item");
447        };
448        assert!(list_tools.id.starts_with("mcpl_"));
449        assert_eq!(list_tools.server_label, "counter");
450        assert_eq!(
451            list_tools
452                .tools
453                .iter()
454                .map(|tool| tool.name.as_str())
455                .collect::<Vec<_>>(),
456            ["increment", "get_value"]
457        );
458        assert_eq!(list_tools.tools[0].description.as_deref(), Some("Discovered test tool"));
459        assert_eq!(list_tools.tools[0].input_schema, serde_json::json!({"type": "object"}));
460        assert_eq!(
461            list_tools.tools[0].annotations,
462            Some(serde_json::json!({"read_only": false}))
463        );
464    }
465
466    #[tokio::test]
467    async fn build_with_handlers_registers_mixed_tools_and_runtime_metadata() {
468        let mut executors = GatewayExecutors::from_env(Arc::new(reqwest::Client::new()));
469        executors.insert(GatewayExecutorRegistration::Mcp {
470            server_label: "counter".to_owned(),
471            handlers: vec![
472                discovered_handler("counter", "increment", "mcp__counter__increment"),
473                discovered_handler("counter", "get_value", "mcp__counter__get_value"),
474            ],
475        });
476        let mut tools = mixed_tool_declarations();
477
478        let registry = ToolRegistry::build_with_handlers(&mut tools, &mut executors)
479            .await
480            .expect("mixed registry");
481
482        assert_eq!(registry.len(), 8);
483        assert!(registry.contains_mcp_server_label("counter"));
484        assert!(!registry.contains_mcp_server_label("missing"));
485        assert_mcp_list_tools_metadata(&registry);
486
487        let expected_entries = [
488            ("echo", ToolType::Function, None, false),
489            ("freeform", ToolType::Custom, None, false),
490            ("mcp__counter__increment", ToolType::Mcp, Some("counter"), true),
491            ("mcp__counter__get_value", ToolType::Mcp, Some("counter"), true),
492            ("web_search", ToolType::WebSearch, None, true),
493            ("file_search", ToolType::FileSearch, None, false),
494            ("code_interpreter", ToolType::CodeInterpreter, None, false),
495            (
496                "agentic_ns__mcp__shell__run",
497                ToolType::CodexNamespace,
498                Some("mcp__shell"),
499                false,
500            ),
501        ];
502        for (name, tool_type, server_label, has_handler) in expected_entries {
503            let entry = registry
504                .lookup(name)
505                .unwrap_or_else(|| panic!("missing registry entry '{name}'"));
506            assert_eq!(entry.tool_type, tool_type, "unexpected type for '{name}'");
507            assert_eq!(
508                entry.server_label.as_deref(),
509                server_label,
510                "unexpected server label for '{name}'"
511            );
512            assert_eq!(entry.handler.is_some(), has_handler, "unexpected handler for '{name}'");
513        }
514        assert_eq!(registry.lookup("freeform").unwrap().config["name"], "freeform");
515        assert_eq!(registry.lookup("echo").unwrap().config["name"], "echo");
516        assert_eq!(
517            registry.lookup("mcp__counter__increment").unwrap().config["tool_name"],
518            "increment"
519        );
520        assert_eq!(
521            registry.lookup("web_search").unwrap().config["search_context_size"],
522            "low"
523        );
524        assert_eq!(
525            registry.lookup("file_search").unwrap().config["vector_store_ids"][0],
526            "vs_test"
527        );
528        assert_eq!(
529            registry.lookup("agentic_ns__mcp__shell__run").unwrap().config["tools"][0]["name"],
530            "agentic_ns__mcp__shell__run"
531        );
532        for name in [
533            "mcp__counter__increment",
534            "mcp__counter__get_value",
535            "web_search",
536            "file_search",
537            "code_interpreter",
538        ] {
539            assert!(registry.is_gateway_owned_name(name), "'{name}' should be gateway-owned");
540        }
541        for name in ["echo", "freeform", "agentic_ns__mcp__shell__run"] {
542            assert!(!registry.is_gateway_owned_name(name), "'{name}' should be client-owned");
543        }
544
545        let ResponsesTool::Mcp(declared) = &tools[1] else {
546            panic!("expected MCP declaration");
547        };
548        assert_eq!(declared.discovered_tools.len(), 2);
549        assert_eq!(
550            tools[1]
551                .to_function_tools()
552                .into_iter()
553                .map(|tool| tool.name)
554                .collect::<Vec<_>>(),
555            ["mcp__counter__increment", "mcp__counter__get_value"]
556        );
557
558        let ResponsesTool::Namespace(namespace) = &tools[5] else {
559            panic!("expected namespace declaration");
560        };
561        assert!(matches!(
562            namespace.tools.as_slice(),
563            [crate::types::tools::CodexNamespaceMember::Function(function)] if function.name.as_str() == "run"
564        ));
565        assert_namespace_call_restoration(&registry);
566    }
567
568    #[tokio::test]
569    async fn build_with_handlers_retains_mcp_discovery_failure_output() {
570        let mut tools = vec![declaration("unreachable")];
571        let mut executors = GatewayExecutors::default();
572
573        let registry = ToolRegistry::build_with_handlers(&mut tools, &mut executors)
574            .await
575            .expect("discovery failures should become response metadata");
576
577        let [list_tools] = registry.mcp_list_tools_items() else {
578            panic!("expected one MCP list-tools item");
579        };
580        assert_eq!(list_tools.server_label, "unreachable");
581        assert!(list_tools.tools.is_empty());
582        assert!(
583            list_tools
584                .error
585                .as_deref()
586                .is_some_and(|error| error.contains("failed"))
587        );
588        assert!(registry.is_empty());
589    }
590
591    #[tokio::test]
592    async fn duplicate_mcp_server_labels_are_rejected() {
593        let mut tools = vec![declaration("counter"), declaration("counter")];
594        let mut executors = GatewayExecutors::default();
595
596        let error = ToolRegistry::build_with_handlers(&mut tools, &mut executors)
597            .await
598            .expect_err("duplicate server_label must fail");
599
600        assert!(
601            matches!(error, ToolError::Config(message) if message.contains("duplicate MCP declarations") && message.contains("counter"))
602        );
603    }
604
605    #[tokio::test]
606    async fn cross_server_internal_name_collisions_are_rejected() {
607        let internal_name = "mcp__foo__bar__baz";
608        let mut executors = GatewayExecutors::default();
609        executors.insert(GatewayExecutorRegistration::Mcp {
610            server_label: "foo".to_owned(),
611            handlers: vec![discovered_handler("foo", "bar__baz", internal_name)],
612        });
613        executors.insert(GatewayExecutorRegistration::Mcp {
614            server_label: "foo__bar".to_owned(),
615            handlers: vec![discovered_handler("foo__bar", "baz", internal_name)],
616        });
617        let mut tools = vec![declaration("foo"), declaration("foo__bar")];
618
619        let error = ToolRegistry::build_with_handlers(&mut tools, &mut executors)
620            .await
621            .expect_err("colliding derived MCP names must fail");
622
623        assert!(matches!(
624            error,
625            ToolError::Config(message)
626                if message.contains(internal_name) && message.matches("MCP tool").count() == 2
627        ));
628    }
629
630    #[tokio::test]
631    async fn discovered_mcp_name_collision_with_function_is_rejected_in_any_order() {
632        let internal_name = "mcp__counter__increment";
633
634        for mcp_first in [false, true] {
635            let function = serde_json::from_value(serde_json::json!({
636                "type": "function",
637                "name": internal_name
638            }))
639            .expect("function declaration");
640            let mcp = declaration("counter");
641            let mut tools = if mcp_first {
642                vec![mcp, function]
643            } else {
644                vec![function, mcp]
645            };
646            let mut executors = GatewayExecutors::default();
647            executors.insert(GatewayExecutorRegistration::Mcp {
648                server_label: "counter".to_owned(),
649                handlers: vec![discovered_handler("counter", "increment", internal_name)],
650            });
651
652            let error = ToolRegistry::build_with_handlers(&mut tools, &mut executors)
653                .await
654                .expect_err("MCP internal name must not overwrite a function");
655
656            assert!(matches!(
657                error,
658                ToolError::Config(message)
659                    if message.contains(internal_name)
660                        && message.contains("MCP tool")
661                        && message.contains("function tool")
662            ));
663        }
664    }
665}