Skip to main content

zeph_tools/
registry.rs

1// SPDX-FileCopyrightText: 2026 Andrei G <bug-ops>
2// SPDX-License-Identifier: MIT OR Apache-2.0
3
4use std::borrow::Cow;
5use std::fmt::Write;
6
7#[non_exhaustive]
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum InvocationHint {
10    /// Tool invoked via ```{tag}\n...\n``` fenced block in LLM response
11    FencedBlock(&'static str),
12    /// Tool invoked via structured `ToolCall` JSON
13    ToolCall,
14}
15
16#[derive(Debug, Clone)]
17pub struct ToolDef {
18    pub id: Cow<'static, str>,
19    pub description: Cow<'static, str>,
20    pub schema: schemars::Schema,
21    pub invocation: InvocationHint,
22    /// Raw output schema from an MCP server, if present.
23    ///
24    /// DO NOT convert to `schemars::Schema` — lossy; see #2931 critique P0-1.
25    pub output_schema: Option<serde_json::Value>,
26    /// ID of the MCP server that registered this tool, or `None` for built-in tools.
27    ///
28    /// This is the authoritative way to determine whether a tool originates from an MCP
29    /// server — `id` is a sanitized, server-namespaced string with no reliable prefix to
30    /// pattern-match on (see #5712).
31    pub server_id: Option<String>,
32}
33
34impl ToolDef {
35    /// Returns `true` if this tool was registered by an MCP server.
36    #[must_use]
37    pub fn is_mcp_tool(&self) -> bool {
38        self.server_id.is_some()
39    }
40}
41
42#[derive(Debug, Default)]
43pub struct ToolRegistry {
44    tools: Vec<ToolDef>,
45}
46
47impl ToolRegistry {
48    #[must_use]
49    pub fn from_definitions(tools: Vec<ToolDef>) -> Self {
50        Self { tools }
51    }
52
53    #[must_use]
54    pub fn tools(&self) -> &[ToolDef] {
55        &self.tools
56    }
57
58    #[must_use]
59    pub fn find(&self, id: &str) -> Option<&ToolDef> {
60        self.tools.iter().find(|t| t.id.as_ref() == id)
61    }
62
63    /// Format tools for prompt, excluding tools fully denied by policy.
64    #[must_use]
65    pub fn format_for_prompt_filtered(
66        &self,
67        policy: &crate::permissions::PermissionPolicy,
68    ) -> String {
69        let mut out = String::from("<tools>\n");
70        for tool in &self.tools {
71            if policy.is_fully_denied(&tool.id) {
72                continue;
73            }
74            format_tool(&mut out, tool);
75        }
76        out.push_str("</tools>");
77        out
78    }
79}
80
81fn format_tool(out: &mut String, tool: &ToolDef) {
82    let _ = writeln!(out, "## {}", tool.id);
83    let _ = writeln!(out, "{}", tool.description);
84    match tool.invocation {
85        InvocationHint::FencedBlock(tag) => {
86            let _ = writeln!(out, "Invocation: use ```{tag} fenced block");
87        }
88        InvocationHint::ToolCall => {
89            let _ = writeln!(
90                out,
91                "Invocation: use tool_call with {{\"tool_id\": \"{}\", \"params\": {{...}}}}",
92                tool.id
93            );
94        }
95    }
96    format_schema_params(out, &tool.schema);
97    out.push('\n');
98}
99
100/// Extract the primary type when schemars renders `Option<T>` as `"type": ["T", "null"]`
101/// or `"anyOf": [{"type": "T"}, {"type": "null"}]`.
102fn extract_non_null_type(obj: &serde_json::Map<String, serde_json::Value>) -> Option<&str> {
103    if let Some(arr) = obj.get("type").and_then(|v| v.as_array()) {
104        return arr.iter().filter_map(|v| v.as_str()).find(|t| *t != "null");
105    }
106    obj.get("anyOf")?
107        .as_array()?
108        .iter()
109        .filter_map(|v| v.as_object())
110        .filter_map(|o| o.get("type")?.as_str())
111        .find(|t| *t != "null")
112}
113
114fn format_schema_params(out: &mut String, schema: &schemars::Schema) {
115    let Some(obj) = schema.as_object() else {
116        return;
117    };
118    let Some(serde_json::Value::Object(props)) = obj.get("properties") else {
119        return;
120    };
121    if props.is_empty() {
122        return;
123    }
124
125    let required: Vec<&str> = obj
126        .get("required")
127        .and_then(|v| v.as_array())
128        .map(|arr| arr.iter().filter_map(|v| v.as_str()).collect())
129        .unwrap_or_default();
130
131    let _ = writeln!(out, "Parameters:");
132    for (name, prop) in props {
133        let prop_obj = prop.as_object();
134        let ty = prop_obj
135            .and_then(|o| {
136                o.get("type")
137                    .and_then(|v| v.as_str())
138                    .or_else(|| extract_non_null_type(o))
139            })
140            .unwrap_or("string");
141        let desc = prop_obj
142            .and_then(|o| o.get("description"))
143            .and_then(|v| v.as_str())
144            .unwrap_or("");
145        let req = if required.contains(&name.as_str()) {
146            "required"
147        } else {
148            "optional"
149        };
150        let _ = writeln!(out, "  - {name}: {desc} ({ty}, {req})");
151    }
152}
153
154#[cfg(test)]
155mod tests {
156    use super::*;
157    use crate::file::ReadParams;
158    use crate::shell::BashParams;
159
160    fn sample_tools() -> Vec<ToolDef> {
161        vec![
162            ToolDef {
163                id: "bash".into(),
164                description: "Execute a shell command".into(),
165                schema: schemars::schema_for!(BashParams),
166                invocation: InvocationHint::FencedBlock("bash"),
167                output_schema: None,
168                server_id: None,
169            },
170            ToolDef {
171                id: "read".into(),
172                description: "Read file contents".into(),
173                schema: schemars::schema_for!(ReadParams),
174                invocation: InvocationHint::ToolCall,
175                output_schema: None,
176                server_id: None,
177            },
178        ]
179    }
180
181    #[test]
182    fn from_definitions_stores_tools() {
183        let reg = ToolRegistry::from_definitions(sample_tools());
184        assert_eq!(reg.tools().len(), 2);
185    }
186
187    #[test]
188    fn default_registry_is_empty() {
189        let reg = ToolRegistry::default();
190        assert!(reg.tools().is_empty());
191    }
192
193    #[test]
194    fn find_existing_tool() {
195        let reg = ToolRegistry::from_definitions(sample_tools());
196        assert!(reg.find("bash").is_some());
197        assert!(reg.find("read").is_some());
198    }
199
200    #[test]
201    fn find_nonexistent_returns_none() {
202        let reg = ToolRegistry::from_definitions(sample_tools());
203        assert!(reg.find("nonexistent").is_none());
204    }
205
206    #[test]
207    fn format_for_prompt_contains_tools() {
208        let reg = ToolRegistry::from_definitions(sample_tools());
209        let prompt =
210            reg.format_for_prompt_filtered(&crate::permissions::PermissionPolicy::default());
211        assert!(prompt.contains("<tools>"));
212        assert!(prompt.contains("</tools>"));
213        assert!(prompt.contains("## bash"));
214        assert!(prompt.contains("## read"));
215    }
216
217    #[test]
218    fn format_for_prompt_shows_invocation_fenced() {
219        let reg = ToolRegistry::from_definitions(sample_tools());
220        let prompt =
221            reg.format_for_prompt_filtered(&crate::permissions::PermissionPolicy::default());
222        assert!(prompt.contains("Invocation: use ```bash fenced block"));
223    }
224
225    #[test]
226    fn format_for_prompt_shows_invocation_tool_call() {
227        let reg = ToolRegistry::from_definitions(sample_tools());
228        let prompt =
229            reg.format_for_prompt_filtered(&crate::permissions::PermissionPolicy::default());
230        assert!(prompt.contains("Invocation: use tool_call"));
231        assert!(prompt.contains("\"tool_id\": \"read\""));
232    }
233
234    #[test]
235    fn format_for_prompt_shows_param_info() {
236        let reg = ToolRegistry::from_definitions(sample_tools());
237        let prompt =
238            reg.format_for_prompt_filtered(&crate::permissions::PermissionPolicy::default());
239        assert!(prompt.contains("command:"));
240        assert!(prompt.contains("required"));
241        assert!(prompt.contains("string"));
242    }
243
244    #[test]
245    fn format_for_prompt_shows_optional_params() {
246        let reg = ToolRegistry::from_definitions(sample_tools());
247        let prompt =
248            reg.format_for_prompt_filtered(&crate::permissions::PermissionPolicy::default());
249        assert!(prompt.contains("offset:"));
250        assert!(prompt.contains("optional"));
251        assert!(
252            prompt.contains("(integer, optional)"),
253            "Option<u32> should render as integer, not string: {prompt}"
254        );
255    }
256
257    #[test]
258    fn format_filtered_excludes_fully_denied() {
259        use crate::permissions::{PermissionAction, PermissionPolicy, PermissionRule};
260        use std::collections::HashMap;
261        let mut rules = HashMap::new();
262        rules.insert(
263            "bash".to_owned(),
264            vec![PermissionRule {
265                pattern: "*".to_owned(),
266                action: PermissionAction::Deny,
267            }],
268        );
269        let policy = PermissionPolicy::new(rules);
270        let reg = ToolRegistry::from_definitions(sample_tools());
271        let prompt = reg.format_for_prompt_filtered(&policy);
272        assert!(!prompt.contains("## bash"));
273        assert!(prompt.contains("## read"));
274    }
275
276    #[test]
277    fn format_filtered_includes_mixed_rules() {
278        use crate::permissions::{PermissionAction, PermissionPolicy, PermissionRule};
279        use std::collections::HashMap;
280        let mut rules = HashMap::new();
281        rules.insert(
282            "bash".to_owned(),
283            vec![
284                PermissionRule {
285                    pattern: "echo *".to_owned(),
286                    action: PermissionAction::Allow,
287                },
288                PermissionRule {
289                    pattern: "*".to_owned(),
290                    action: PermissionAction::Deny,
291                },
292            ],
293        );
294        let policy = PermissionPolicy::new(rules);
295        let reg = ToolRegistry::from_definitions(sample_tools());
296        let prompt = reg.format_for_prompt_filtered(&policy);
297        assert!(prompt.contains("## bash"));
298    }
299
300    #[test]
301    fn format_filtered_no_rules_includes_all() {
302        let policy = crate::permissions::PermissionPolicy::default();
303        let reg = ToolRegistry::from_definitions(sample_tools());
304        let prompt = reg.format_for_prompt_filtered(&policy);
305        assert!(prompt.contains("## bash"));
306        assert!(prompt.contains("## read"));
307    }
308}