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 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 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}