Skip to main content

ares_mcp/
registry.rs

1use std::{collections::HashMap, path::Path, sync::Arc};
2
3use ares_types::types::ToolDefinition;
4use rmcp::model::{CallToolResult, Tool};
5use serde::{Deserialize, Serialize};
6use serde_json::json;
7use thiserror::Error;
8
9use super::client::{McpClient, McpServerConfig};
10use super::extension::{dispatch_extensions, McpToolExtension};
11
12pub struct McpRegistry {
13    clients: HashMap<String, Arc<McpClient>>,
14}
15
16impl Default for McpRegistry {
17    fn default() -> Self {
18        Self::new()
19    }
20}
21
22impl McpRegistry {
23    pub fn new() -> Self {
24        Self {
25            clients: HashMap::new(),
26        }
27    }
28
29    /// Register (or replace) an MCP client by config name.
30    pub fn register(&mut self, config: McpServerConfig) -> Arc<McpClient> {
31        let client = McpClient::new(config);
32        let name = client.name().to_string();
33        let arc = Arc::new(client);
34        self.clients.insert(name, arc.clone());
35        arc
36    }
37
38    /// Remove a client by name. Returns true if it existed.
39    pub fn deregister(&mut self, name: &str) -> bool {
40        self.clients.remove(name).is_some()
41    }
42
43    pub fn from_dir(config_dir: &str) -> Result<Self, Box<dyn std::error::Error>> {
44        let mut clients = HashMap::new();
45        let path = Path::new(config_dir);
46
47        if !path.exists() {
48            tracing::warn!("MCP config directory not found: {}", config_dir);
49            return Ok(Self::new());
50        }
51
52        for entry in std::fs::read_dir(path)? {
53            let entry = entry?;
54            let file_path = entry.path();
55
56            if !is_mcp_config_file(&file_path) {
57                continue;
58            }
59
60            match load_mcp_config(&file_path) {
61                Ok(config) if config.enabled => {
62                    let client = McpClient::new(config);
63                    let name = client.name().to_string();
64                    tracing::info!("Registered MCP client: {}", name);
65                    clients.insert(name, Arc::new(client));
66                }
67                Ok(config) => {
68                    tracing::debug!(name = %config.name, "Skipping disabled MCP client");
69                }
70                Err(error) => {
71                    tracing::warn!(
72                        path = %file_path.display(),
73                        error = %error,
74                        "Skipping invalid MCP config"
75                    );
76                }
77            }
78        }
79
80        let mut registry = Self::new();
81        registry.clients = clients;
82        Ok(registry)
83    }
84
85    pub fn get_client(&self, name: &str) -> Option<&Arc<McpClient>> {
86        self.clients.get(name)
87    }
88
89    pub fn eruka(&self) -> Option<&Arc<McpClient>> {
90        self.clients.get("eruka")
91    }
92
93    pub fn client_names(&self) -> Vec<String> {
94        self.clients.keys().cloned().collect()
95    }
96}
97
98impl cordis::Service for McpRegistry {
99    fn name(&self) -> &'static str { "mcp_registry" }
100    fn init(&self, _ctx: &std::sync::Arc<cordis::Context>) -> cordis::ServiceInitFuture<'_> {
101        Box::pin(async { Ok(None) })
102    }
103    fn check(&self) -> bool { true }
104}
105
106fn is_mcp_config_file(path: &Path) -> bool {
107    if path.extension().and_then(|s| s.to_str()) != Some("toon") {
108        return false;
109    }
110
111    !path
112        .file_name()
113        .and_then(|s| s.to_str())
114        .map(|name| name.ends_with(".example.toon"))
115        .unwrap_or(false)
116}
117
118fn load_mcp_config(path: &Path) -> Result<McpServerConfig, Box<dyn std::error::Error>> {
119    let content = std::fs::read_to_string(path)?;
120
121    match toml::from_str::<McpServerConfig>(&content) {
122        Ok(config) => Ok(config),
123        Err(toml_error) => {
124            toon_format::decode_default::<McpServerConfig>(&content).map_err(|toon_error| {
125                format!(
126                    "failed to parse as TOML ({}) or TOON ({})",
127                    toml_error, toon_error
128                )
129                .into()
130            })
131        }
132    }
133}
134
135#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
136pub struct ToolRegistered { pub name: String }
137
138#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
139pub struct ToolUnregistered { pub name: String }
140
141#[derive(Debug, Clone, PartialEq, Eq, Error)]
142pub enum RegistryError {
143    #[error("tool already registered: {0}")]
144    Duplicate(String),
145    #[error("tool not found: {0}")]
146    NotFound(String),
147    #[error("invalid tool schema: {0}")]
148    InvalidSchema(String),
149}
150
151#[derive(Clone)]
152pub struct ToolRegistry {
153    tools: HashMap<String, Tool>,
154    extensions: Vec<Arc<dyn McpToolExtension>>,
155}
156
157impl ToolRegistry {
158    pub fn new() -> Self { Self { tools: HashMap::new(), extensions: Vec::new() } }
159    pub fn with_builtin_tools() -> Self {
160        let mut registry = Self::new();
161        for tool in builtin_ares_tools() {
162            register_tool(&mut registry.tools, tool).expect("built-in tool names are unique");
163        }
164        registry
165    }
166    pub fn register(&mut self, tool: Tool) -> Result<ToolRegistered, RegistryError> {
167        register_tool(&mut self.tools, tool)
168    }
169    pub fn get(&self, name: &str) -> Result<&Tool, RegistryError> { get_tool(&self.tools, name) }
170    pub fn unregister(&mut self, name: &str) -> Result<ToolUnregistered, RegistryError> {
171        self.tools.remove(name).ok_or_else(|| RegistryError::NotFound(name.to_string()))?;
172        Ok(ToolUnregistered { name: name.to_string() })
173    }
174    pub fn list(&self) -> Vec<Tool> { list_tools(&self.tools, &self.extensions) }
175    pub fn register_extension(&mut self, ext: Arc<dyn McpToolExtension>) { self.extensions.push(ext); }
176    pub fn remove_extension(&mut self, index: usize) -> bool {
177        if index >= self.extensions.len() { return false; }
178        self.extensions.remove(index);
179        true
180    }
181    pub fn extensions(&self) -> &[Arc<dyn McpToolExtension>] { &self.extensions }
182    pub fn tool_count(&self) -> usize { self.tools.len() }
183    pub fn extension_count(&self) -> usize { self.extensions.len() }
184}
185
186impl Default for ToolRegistry { fn default() -> Self { Self::new() } }
187
188pub fn register_tool(tools: &mut HashMap<String, Tool>, tool: Tool) -> Result<ToolRegistered, RegistryError> {
189    validate_tool_schema(&tool)?;
190    let name = tool.name.to_string();
191    if tools.contains_key(&name) { return Err(RegistryError::Duplicate(name.clone())); }
192    tools.insert(name.clone(), tool);
193    Ok(ToolRegistered { name })
194}
195
196pub fn get_tool<'a>(tools: &'a HashMap<String, Tool>, name: &str) -> Result<&'a Tool, RegistryError> {
197    tools.get(name).ok_or_else(|| RegistryError::NotFound(name.to_string()))
198}
199
200pub fn list_tools(tools: &HashMap<String, Tool>, extensions: &[Arc<dyn McpToolExtension>]) -> Vec<Tool> {
201    let mut out: Vec<Tool> = tools.values().cloned().collect();
202    for ext in extensions { out.extend(ext.tools()); }
203    out
204}
205
206pub async fn extension_dispatch(
207    extensions: &[Arc<dyn McpToolExtension>],
208    tool_name: &str,
209    arguments: serde_json::Value,
210    tenant_id: &str,
211) -> Option<Result<CallToolResult, String>> {
212    dispatch_extensions(extensions, tool_name, arguments, tenant_id).await
213}
214
215pub fn tool_to_definition(tool: &Tool) -> ToolDefinition {
216    ToolDefinition {
217        name: tool.name.to_string(),
218        description: tool.description.clone().map(|d| d.to_string()).unwrap_or_default(),
219        parameters: serde_json::to_value(&tool.input_schema).unwrap_or_else(|_| json!({})),
220    }
221}
222
223pub fn validate_tool_schema(tool: &Tool) -> Result<(), RegistryError> {
224    if tool.name.as_ref().trim().is_empty() {
225        return Err(RegistryError::InvalidSchema("tool name must not be empty".into()));
226    }
227    let schema_value = serde_json::to_value(&tool.input_schema)
228        .map_err(|e| RegistryError::InvalidSchema(format!("input_schema not serializable: {e}")))?;
229    match schema_value.get("type").and_then(|t| t.as_str()) {
230        Some("object") => Ok(()),
231        Some(other) => Err(RegistryError::InvalidSchema(format!("input_schema type must be object, got {other}"))),
232        None => Err(RegistryError::InvalidSchema("input_schema must include type: object".into())),
233    }
234}
235
236fn build_tool(name: &str, description: &str, schema: serde_json::Value, title: &str) -> Tool {
237    let input_schema: rmcp::model::JsonObject =
238        serde_json::from_value(schema).unwrap_or_default();
239    Tool::new(name.to_string(), description.to_string(), input_schema)
240        .with_title(title.to_string())
241}
242
243pub fn builtin_ares_tools() -> Vec<Tool> {
244    vec![
245        build_tool(
246            "ares_list_agents",
247            "List all agents available in your ARES account. Returns agent names, descriptions, types, and deployment status.",
248            json!({"type":"object","properties":{},"required":[]}),
249            "List ARES Agents",
250        ),
251        build_tool(
252            "ares_run_agent",
253            "Run an ARES agent with a message. Specify the agent name and your message. Optionally pass a context_id to continue a conversation.",
254            json!({"type":"object","properties":{"agent_name":{"type":"string"},"message":{"type":"string"},"context_id":{"type":"string"}},"required":["agent_name","message"]}),
255            "Run ARES Agent",
256        ),
257        build_tool(
258            "ares_get_status",
259            "Check the status of a previous agent run. Pass the context_id from an ares_run_agent call. Returns running/completed/failed status.",
260            json!({"type":"object","properties":{"context_id":{"type":"string"}},"required":["context_id"]}),
261            "Get Agent Status",
262        ),
263        build_tool(
264            "ares_deploy_agent",
265            "Deploy a new agent to ARES by providing a .toon configuration (TOML format). The agent becomes immediately available for use.",
266            json!({"type":"object","properties":{"toon_config":{"type":"string"},"name_override":{"type":"string"}},"required":["toon_config"]}),
267            "Deploy Agent",
268        ),
269        build_tool(
270            "ares_get_usage",
271            "Check your ARES account usage statistics and quota. Shows requests made, tokens consumed, and remaining quota for your tier.",
272            json!({"type":"object","properties":{"from_date":{"type":"string"},"to_date":{"type":"string"}},"required":[]}),
273            "Get Usage Stats",
274        ),
275    ]
276}
277
278#[cfg(test)]
279mod tests {
280    use super::*;
281
282    #[test]
283    fn loads_toml_and_toon_mcp_configs() {
284        let dir = tempfile::tempdir().unwrap();
285        std::fs::write(
286            dir.path().join("eruka.toon"),
287            r#"name = "eruka"
288enabled = true
289endpoint = "https://eruka.dirmacs.com/mcp"
290transport = "http"
291timeout_secs = 30
292"#,
293        )
294        .unwrap();
295        std::fs::write(
296            dir.path().join("filesystem.toon"),
297            r#"name: filesystem
298enabled: true
299command: npx
300args[2]: "-y","@modelcontextprotocol/server-filesystem"
301timeout_secs: 30
302"#,
303        )
304        .unwrap();
305        std::fs::write(
306            dir.path().join("eruka.example.toon"),
307            r#"name: eruka
308enabled: true
309command: eruka-mcp
310"#,
311        )
312        .unwrap();
313
314        let registry = McpRegistry::from_dir(dir.path().to_str().unwrap()).unwrap();
315        let mut names = registry.client_names();
316        names.sort();
317
318        assert_eq!(names, vec!["eruka".to_string(), "filesystem".to_string()]);
319        assert!(registry.eruka().is_some());
320    }
321
322    #[test]
323    fn register_and_deregister_client() {
324        let mut registry = McpRegistry::new();
325        assert!(registry.get_client("pom").is_none());
326
327        registry.register(McpServerConfig {
328            name: "pom".into(),
329            enabled: true,
330            command: None,
331            args: None,
332            timeout_secs: Some(15),
333            endpoint: Some("http://localhost:3002/mcp".into()),
334            transport: Some("http".into()),
335            api_key: None,
336        });
337
338        assert!(registry.get_client("pom").is_some());
339        assert_eq!(registry.client_names(), vec!["pom".to_string()]);
340
341        assert!(registry.deregister("pom"));
342        assert!(!registry.deregister("pom"));
343        assert!(registry.get_client("pom").is_none());
344    }
345
346    #[test]
347    fn skips_invalid_mcp_config_without_failing_registry() {
348        let dir = tempfile::tempdir().unwrap();
349        std::fs::write(
350            dir.path().join("pom.toon"),
351            r#"name: pom
352enabled: true
353transport: http
354endpoint: http://localhost:3002/mcp
355timeout_secs: 15
356"#,
357        )
358        .unwrap();
359        std::fs::write(dir.path().join("broken.toon"), "not valid =").unwrap();
360
361        let registry = McpRegistry::from_dir(dir.path().to_str().unwrap()).unwrap();
362
363        assert_eq!(registry.client_names(), vec!["pom".to_string()]);
364    }
365
366    #[test]
367    fn from_dir_missing_directory_returns_empty_registry() {
368        let path = std::env::temp_dir().join(format!(
369            "ares-mcp-missing-{}",
370            uuid::Uuid::new_v4()
371        ));
372        assert!(!path.exists());
373
374        let registry = McpRegistry::from_dir(path.to_str().unwrap()).unwrap();
375
376        assert!(registry.client_names().is_empty());
377        assert!(registry.eruka().is_none());
378    }
379
380    #[test]
381    fn from_dir_skips_disabled_configs_and_non_toon_files() {
382        let dir = tempfile::tempdir().unwrap();
383        std::fs::write(
384            dir.path().join("disabled.toon"),
385            r#"name = "disabled"
386enabled = false
387endpoint = "http://localhost/mcp"
388transport = "http"
389timeout_secs = 10
390"#,
391        )
392        .unwrap();
393        std::fs::write(dir.path().join("readme.txt"), "not an mcp config").unwrap();
394
395        let registry = McpRegistry::from_dir(dir.path().to_str().unwrap()).unwrap();
396
397        assert!(registry.client_names().is_empty());
398    }
399
400    #[test]
401    fn register_replaces_existing_client() {
402        let mut registry = McpRegistry::new();
403
404        registry.register(McpServerConfig {
405            name: "svc".into(),
406            enabled: true,
407            command: None,
408            args: None,
409            timeout_secs: Some(10),
410            endpoint: Some("http://localhost:3001/mcp".into()),
411            transport: Some("http".into()),
412            api_key: None,
413        });
414        registry.register(McpServerConfig {
415            name: "svc".into(),
416            enabled: true,
417            command: None,
418            args: None,
419            timeout_secs: Some(20),
420            endpoint: Some("http://localhost:3002/mcp".into()),
421            transport: Some("http".into()),
422            api_key: None,
423        });
424
425        assert_eq!(registry.client_names(), vec!["svc".to_string()]);
426        assert!(registry.get_client("svc").is_some());
427    }
428
429    use crate::extension::NoOpMcpExtension;
430    use async_trait::async_trait;
431    use rmcp::model::ContentBlock;
432
433    fn sample_tool(name: &str) -> Tool {
434        let input_schema: rmcp::model::JsonObject = serde_json::from_value(json!({"type":"object","properties":{},"required":[]})).unwrap_or_default();
435        Tool::new(name.to_string(), format!("{name} tool"), input_schema)
436    }
437
438    fn serde_roundtrip<T>(value: &T) -> T
439    where T: serde::Serialize + for<'de> serde::Deserialize<'de> + PartialEq + std::fmt::Debug,
440    {
441        let j = serde_json::to_string(value).unwrap();
442        let p: T = serde_json::from_str(&j).unwrap();
443        assert_eq!(*value, p);
444        p
445    }
446
447    #[test] fn tool_registry_new_is_empty() { let r = ToolRegistry::new(); assert_eq!(r.tool_count(), 0); assert!(r.list().is_empty()); }
448    #[test] fn tool_registry_default_matches_new() { assert_eq!(ToolRegistry::default().tool_count(), 0); }
449    #[test] fn tool_registry_with_builtin_has_five_unique_tools() {
450        let r = ToolRegistry::with_builtin_tools();
451        assert_eq!(r.tool_count(), 5);
452        let names: Vec<String> = r.list().into_iter().map(|t| t.name.to_string()).collect();
453        assert!(names.iter().any(|n| n == "ares_list_agents"));
454        assert_eq!(names.len(), names.iter().collect::<std::collections::HashSet<_>>().len());
455    }
456    #[test] fn register_tool_inserts_and_returns_event() {
457        let mut tools = HashMap::new();
458        assert_eq!(register_tool(&mut tools, sample_tool("custom")).unwrap().name, "custom");
459    }
460    #[test] fn register_tool_duplicate_returns_error() {
461        let mut r = ToolRegistry::new();
462        r.register(sample_tool("dup")).unwrap();
463        assert!(matches!(r.register(sample_tool("dup")).unwrap_err(), RegistryError::Duplicate(_)));
464    }
465    #[test] fn get_tool_returns_reference_when_present() {
466        assert_eq!(ToolRegistry::with_builtin_tools().get("ares_list_agents").unwrap().name.as_ref(), "ares_list_agents");
467    }
468    #[test] fn get_tool_not_found_returns_error() {
469        assert!(matches!(ToolRegistry::with_builtin_tools().get("missing").unwrap_err(), RegistryError::NotFound(_)));
470    }
471    #[test] fn unregister_tool_returns_event() {
472        let mut r = ToolRegistry::new();
473        r.register(sample_tool("temp")).unwrap();
474        assert_eq!(r.unregister("temp").unwrap().name, "temp");
475        assert_eq!(r.tool_count(), 0);
476    }
477    #[test] fn unregister_tool_not_found_returns_error() {
478        assert!(matches!(ToolRegistry::new().unregister("ghost").unwrap_err(), RegistryError::NotFound(_)));
479    }
480    #[test] fn list_tools_includes_extension_tools() {
481        struct Ext;
482        #[async_trait]
483        impl McpToolExtension for Ext {
484            fn tools(&self) -> Vec<Tool> { vec![sample_tool("ext_search")] }
485            async fn execute(&self, _tool_name: &str, _arguments: serde_json::Value, _tenant_id: &str) -> Option<Result<CallToolResult, String>> { None }
486        }
487        let mut r = ToolRegistry::with_builtin_tools();
488        r.register_extension(Arc::new(Ext));
489        assert_eq!(r.list().len(), 6);
490    }
491    #[test] fn register_and_remove_extension() {
492        let mut r = ToolRegistry::new();
493        r.register_extension(Arc::new(NoOpMcpExtension));
494        assert!(r.remove_extension(0));
495        assert!(!r.remove_extension(0));
496    }
497    #[test] fn validate_tool_schema_rejects_empty_name() {
498        assert!(matches!(validate_tool_schema(&sample_tool(" ")).unwrap_err(), RegistryError::InvalidSchema(_)));
499    }
500    #[test] fn validate_tool_schema_rejects_non_object_type() {
501        let mut t = sample_tool("bad");
502        t.input_schema = serde_json::from_value(json!({"type":"string"})).unwrap_or_default();
503        assert!(matches!(validate_tool_schema(&t).unwrap_err(), RegistryError::InvalidSchema(_)));
504    }
505    #[test] fn validate_tool_schema_accepts_builtin_tools() { for t in builtin_ares_tools() { validate_tool_schema(&t).unwrap(); } }
506    #[test] fn tool_to_definition_maps_fields() {
507        let d = tool_to_definition(&sample_tool("mapper"));
508        assert_eq!(d.name, "mapper");
509        assert_eq!(d.parameters["type"], "object");
510    }
511    #[test] fn tool_registered_serde_roundtrip() { serde_roundtrip(&ToolRegistered { name: "x".into() }); }
512    #[test] fn tool_unregistered_serde_roundtrip() { serde_roundtrip(&ToolUnregistered { name: "y".into() }); }
513    #[test] fn tool_definition_from_builtin_serde_roundtrip() {
514        let t = builtin_ares_tools().into_iter().find(|x| x.name.as_ref() == "ares_deploy_agent").unwrap();
515        let d = tool_to_definition(&t);
516        let r: ToolDefinition = serde_json::from_str(&serde_json::to_string(&d).unwrap()).unwrap();
517        assert_eq!(r.name, "ares_deploy_agent");
518    }
519    #[test] fn pure_get_tool_helper_matches_registry_get() {
520        let mut tools = HashMap::new();
521        register_tool(&mut tools, sample_tool("ares_get_status")).unwrap();
522        assert_eq!(get_tool(&tools, "ares_get_status").unwrap().name.as_ref(), ToolRegistry::with_builtin_tools().get("ares_get_status").unwrap().name.as_ref());
523    }
524    #[test] fn pure_list_tools_helper_without_extensions() {
525        let mut tools = HashMap::new();
526        for t in builtin_ares_tools() { register_tool(&mut tools, t).unwrap(); }
527        assert_eq!(list_tools(&tools, &[]).len(), 5);
528    }
529    #[tokio::test] async fn extension_dispatch_returns_none_for_unknown_tool() {
530        assert!(extension_dispatch(ToolRegistry::new().extensions(), "unknown", json!({}), "t").await.is_none());
531    }
532    #[tokio::test] async fn extension_dispatch_returns_ok_when_extension_handles_tool() {
533        struct Echo;
534        #[async_trait]
535        impl McpToolExtension for Echo {
536            fn tools(&self) -> Vec<Tool> { vec![sample_tool("echo_ext")] }
537            async fn execute(&self, n: &str, _: serde_json::Value, _: &str) -> Option<Result<CallToolResult, String>> {
538                if n == "echo_ext" { Some(Ok(CallToolResult::success(vec![ContentBlock::text("ok")]))) } else { None }
539            }
540        }
541        let mut r = ToolRegistry::new();
542        r.register_extension(Arc::new(Echo));
543        let ok = extension_dispatch(r.extensions(), "echo_ext", json!({}), "t").await.unwrap().unwrap();
544        assert!(!ok.is_error.unwrap_or(true));
545    }
546    #[test] fn register_tool_rejects_invalid_schema_before_duplicate_check() {
547        let mut tools = HashMap::new();
548        let mut bad = sample_tool("bad");
549        bad.input_schema = serde_json::from_value(json!({"type":"array"})).unwrap_or_default();
550        assert!(matches!(register_tool(&mut tools, bad).unwrap_err(), RegistryError::InvalidSchema(_)));
551        assert!(tools.is_empty());
552    }
553
554    #[test]
555    fn mcp_registry_readable_via_cordis() {
556        use cordis::Service;
557        let ctx = std::sync::Arc::new(cordis::Context::new_root());
558        ctx.provide(McpRegistry::new());
559        let got = ctx.get::<McpRegistry>().expect("provided");
560        assert_eq!(got.name(), "mcp_registry");
561        assert!(got.check());
562    }
563
564}