1use std::sync::Arc;
7
8use async_trait::async_trait;
9use rmcp::handler::server::ServerHandler;
10use rmcp::model::{
11 CallToolRequestParams, CallToolResult, GetPromptRequestParams, GetPromptResult, Implementation,
12 ListPromptsResult, ListResourceTemplatesResult, ListResourcesResult, ListToolsResult,
13 PaginatedRequestParams, Prompt, ReadResourceRequestParams, ReadResourceResult, Resource,
14 ResourceTemplate, ServerCapabilities, ServerInfo, Tool,
15};
16use rmcp::service::{RequestContext, RoleServer};
17use rskit_component::{Component, Health};
18
19use rskit_tool::registry::Registry;
20
21use crate::config::ServerConfig;
22use crate::convert;
23use crate::prompts::{invalid_params_error, prompt_name};
24use crate::resources::{resource_template_matches, resource_template_uri, resource_uri};
25
26pub struct RegistryHandler {
30 name: String,
31 version: String,
32 pub(crate) registry: Arc<Registry>,
33 pub(crate) config: ServerConfig,
34}
35
36impl RegistryHandler {
37 pub(crate) fn mcp_tools(&self) -> Vec<Tool> {
38 self.registry
39 .list()
40 .iter()
41 .filter(|d| self.allows_tool(&d.name))
42 .map(|d| convert::definition_to_tool(d, &self.config.prefix))
43 .collect()
44 }
45
46 pub(crate) fn mcp_prompts(&self) -> Vec<Prompt> {
47 self.config
48 .prompts
49 .iter()
50 .map(|entry| entry.prompt.clone())
51 .collect()
52 }
53
54 pub(crate) fn mcp_resources(&self) -> Vec<Resource> {
55 self.config
56 .resources
57 .iter()
58 .map(|entry| entry.resource.clone())
59 .collect()
60 }
61
62 pub(crate) fn mcp_resource_templates(&self) -> Vec<ResourceTemplate> {
63 self.config
64 .resource_templates
65 .iter()
66 .map(|entry| entry.resource_template.clone())
67 .collect()
68 }
69
70 pub(crate) async fn handle_get_prompt(
71 &self,
72 request: GetPromptRequestParams,
73 ) -> Result<GetPromptResult, rmcp::ErrorData> {
74 let entry = self
75 .config
76 .prompts
77 .iter()
78 .find(|entry| prompt_name(&entry.prompt).as_deref() == Some(request.name.as_str()))
79 .ok_or_else(|| invalid_params_error(format!("prompt not found: {}", request.name)))?;
80 (entry.handler)(request).await
81 }
82
83 pub(crate) async fn handle_read_resource(
84 &self,
85 request: ReadResourceRequestParams,
86 ) -> Result<ReadResourceResult, rmcp::ErrorData> {
87 let uri = request.uri.clone();
88 if let Some(entry) = self
89 .config
90 .resources
91 .iter()
92 .find(|entry| resource_uri(&entry.resource).as_deref() == Some(uri.as_str()))
93 {
94 return (entry.handler)(request).await;
95 }
96 if let Some(entry) = self.config.resource_templates.iter().find(|entry| {
97 resource_template_uri(&entry.resource_template)
98 .is_some_and(|template| resource_template_matches(&template, &uri))
99 }) {
100 return (entry.handler)(request).await;
101 }
102 Err(invalid_params_error(format!("resource not found: {uri}")))
103 }
104}
105
106impl ServerHandler for RegistryHandler {
107 fn get_info(&self) -> ServerInfo {
108 let capabilities = ServerCapabilities::builder()
109 .enable_tools()
110 .enable_prompts()
111 .enable_resources()
112 .build();
113 let server_info = Implementation::new(&self.name, &self.version);
114
115 ServerInfo::new(capabilities)
116 .with_server_info(server_info)
117 .with_instructions(format!(
118 "Tool server '{}' v{} — {} tools available",
119 self.name,
120 self.version,
121 self.registry.len()
122 ))
123 }
124
125 async fn list_tools(
126 &self,
127 _request: Option<PaginatedRequestParams>,
128 _context: RequestContext<RoleServer>,
129 ) -> Result<ListToolsResult, rmcp::ErrorData> {
130 let tools = self.mcp_tools();
131 tracing::debug!(count = tools.len(), "MCP tools/list");
132 Ok(ListToolsResult {
133 tools,
134 next_cursor: None,
135 meta: None,
136 })
137 }
138
139 async fn list_prompts(
140 &self,
141 _request: Option<PaginatedRequestParams>,
142 _context: RequestContext<RoleServer>,
143 ) -> Result<ListPromptsResult, rmcp::ErrorData> {
144 Ok(ListPromptsResult {
145 prompts: self.mcp_prompts(),
146 ..Default::default()
147 })
148 }
149
150 async fn get_prompt(
151 &self,
152 request: GetPromptRequestParams,
153 _context: RequestContext<RoleServer>,
154 ) -> Result<GetPromptResult, rmcp::ErrorData> {
155 self.handle_get_prompt(request).await
156 }
157
158 async fn list_resources(
159 &self,
160 _request: Option<PaginatedRequestParams>,
161 _context: RequestContext<RoleServer>,
162 ) -> Result<ListResourcesResult, rmcp::ErrorData> {
163 Ok(ListResourcesResult {
164 resources: self.mcp_resources(),
165 ..Default::default()
166 })
167 }
168
169 async fn list_resource_templates(
170 &self,
171 _request: Option<PaginatedRequestParams>,
172 _context: RequestContext<RoleServer>,
173 ) -> Result<ListResourceTemplatesResult, rmcp::ErrorData> {
174 Ok(ListResourceTemplatesResult {
175 resource_templates: self.mcp_resource_templates(),
176 ..Default::default()
177 })
178 }
179
180 async fn read_resource(
181 &self,
182 request: ReadResourceRequestParams,
183 _context: RequestContext<RoleServer>,
184 ) -> Result<ReadResourceResult, rmcp::ErrorData> {
185 self.handle_read_resource(request).await
186 }
187
188 fn get_tool(&self, name: &str) -> Option<Tool> {
189 let registry_name = self.strip_prefix(name);
190 if !self.allows_tool(registry_name) {
191 return None;
192 }
193 self.registry
194 .get(registry_name)
195 .map(|t| convert::definition_to_tool(t.definition(), &self.config.prefix))
196 }
197
198 async fn call_tool(
199 &self,
200 request: CallToolRequestParams,
201 _context: RequestContext<RoleServer>,
202 ) -> Result<CallToolResult, rmcp::ErrorData> {
203 Ok(self.handle_call_tool(request).await)
204 }
205}
206
207pub fn create_server(
216 name: impl Into<String>,
217 version: impl Into<String>,
218 registry: Arc<Registry>,
219 config: ServerConfig,
220) -> RegistryHandler {
221 RegistryHandler {
222 name: name.into(),
223 version: version.into(),
224 registry,
225 config,
226 }
227}
228
229#[async_trait]
230impl Component for RegistryHandler {
231 fn name(&self) -> &str {
232 &self.name
233 }
234
235 async fn start(&self) -> rskit_errors::AppResult<()> {
236 Ok(())
237 }
238
239 async fn stop(&self) -> rskit_errors::AppResult<()> {
240 Ok(())
241 }
242
243 fn health(&self) -> Health {
244 Health::healthy(self.name())
245 }
246}
247
248#[cfg(test)]
249mod tests {
250 use super::*;
251 use crate::audit::{ToolAuditEvent, ToolAuditSink};
252 use crate::authz::{ToolAuthorizationDecision, ToolAuthorizationRequest, ToolAuthorizer};
253 use crate::prompts::PromptEntry;
254 use crate::resources::{ResourceEntry, ResourceTemplateEntry};
255 use parking_lot::Mutex;
256
257 use rskit_schema::ValidationResult;
258 use rskit_tool::context::Context;
259 use rskit_tool::{Callable, Definition, ToolInput, ToolResult, from_fn, text_result};
260 use schemars::JsonSchema;
261 use serde::Deserialize;
262 use serde_json::json;
263
264 #[derive(Deserialize, JsonSchema)]
265 struct EchoInput {
266 message: String,
267 }
268
269 fn test_registry() -> Arc<Registry> {
270 let registry = Registry::new();
271 registry
272 .register(
273 from_fn(
274 "echo",
275 "Echo a message back",
276 |_ctx: Context, input: EchoInput| async move {
277 Ok(text_result(&input.message))
278 },
279 )
280 .unwrap(),
281 )
282 .unwrap();
283 Arc::new(registry)
284 }
285
286 #[test]
287 fn test_get_info() {
288 let handler = create_server(
289 "test-server",
290 "0.1.0",
291 test_registry(),
292 ServerConfig::default(),
293 );
294 let info = handler.get_info();
295 assert_eq!(info.server_info.name, "test-server");
296 assert_eq!(info.server_info.version, "0.1.0");
297 }
298
299 #[test]
300 fn test_get_tool_found() {
301 let handler = create_server("test", "0.1.0", test_registry(), ServerConfig::default());
302 let tool = handler.get_tool("echo");
303 assert!(tool.is_some());
304 assert_eq!(tool.unwrap().name.as_ref(), "echo");
305 }
306
307 #[test]
308 fn test_get_tool_not_found() {
309 let handler = create_server("test", "0.1.0", test_registry(), ServerConfig::default());
310 assert!(handler.get_tool("nonexistent").is_none());
311 }
312
313 #[test]
314 fn test_get_tool_with_prefix() {
315 let config = ServerConfig {
316 prefix: "myapp_".to_string(),
317 ..Default::default()
318 };
319 let handler = create_server("test", "0.1.0", test_registry(), config);
320 let tool = handler.get_tool("myapp_echo");
321 assert!(tool.is_some());
322 assert_eq!(tool.unwrap().name.as_ref(), "myapp_echo");
323 }
324
325 #[test]
326 fn test_mcp_tools_lists_all() {
327 let handler = create_server("test", "0.1.0", test_registry(), ServerConfig::default());
328 let tools = handler.mcp_tools();
329 assert_eq!(tools.len(), 1);
330 assert_eq!(tools[0].name.as_ref(), "echo");
331 }
332
333 #[test]
334 fn test_mcp_tools_with_prefix() {
335 let config = ServerConfig {
336 prefix: "pre_".to_string(),
337 ..Default::default()
338 };
339 let handler = create_server("test", "0.1.0", test_registry(), config);
340 let tools = handler.mcp_tools();
341 assert_eq!(tools[0].name.as_ref(), "pre_echo");
342 }
343
344 #[test]
345 fn test_allowed_tools_filter_list_and_lookup() {
346 let config = ServerConfig {
347 allowed_tools: vec!["echo".to_string()],
348 ..Default::default()
349 };
350 let handler = create_server("test", "0.1.0", test_registry(), config);
351
352 let tools = handler.mcp_tools();
353 assert_eq!(tools.len(), 1);
354 assert_eq!(tools[0].name.as_ref(), "echo");
355 assert!(handler.get_tool("echo").is_some());
356 assert!(handler.get_tool("missing").is_none());
357 }
358
359 struct DenyAuthorizer;
360
361 #[async_trait]
362 impl ToolAuthorizer for DenyAuthorizer {
363 async fn authorize_tool(
364 &self,
365 request: &ToolAuthorizationRequest,
366 ) -> Result<ToolAuthorizationDecision, String> {
367 if request.tool_name == "echo" {
368 return Ok(ToolAuthorizationDecision {
369 allowed: false,
370 reason: String::from("echo disabled"),
371 });
372 }
373 Ok(ToolAuthorizationDecision {
374 allowed: true,
375 reason: String::from("allowed"),
376 })
377 }
378 }
379
380 struct RecordingAuthorizer {
381 calls: Arc<Mutex<Vec<String>>>,
382 }
383
384 #[async_trait]
385 impl ToolAuthorizer for RecordingAuthorizer {
386 async fn authorize_tool(
387 &self,
388 request: &ToolAuthorizationRequest,
389 ) -> Result<ToolAuthorizationDecision, String> {
390 self.calls.lock().push(request.tool_name.clone());
391 Ok(ToolAuthorizationDecision {
392 allowed: true,
393 reason: String::from("allowed"),
394 })
395 }
396 }
397
398 struct RecordingAuditSink {
399 events: Arc<Mutex<Vec<ToolAuditEvent>>>,
400 }
401
402 #[async_trait]
403 impl ToolAuditSink for RecordingAuditSink {
404 async fn record_tool_call(&self, event: ToolAuditEvent) {
405 self.events.lock().push(event);
406 }
407 }
408
409 #[tokio::test]
410 async fn test_tool_authorizer_and_audit_sink() {
411 let events = Arc::new(Mutex::new(Vec::new()));
412 let config = ServerConfig {
413 tool_authorizer: Some(Arc::new(DenyAuthorizer)),
414 tool_audit_sink: Some(Arc::new(RecordingAuditSink {
415 events: Arc::clone(&events),
416 })),
417 ..Default::default()
418 };
419 let handler = create_server("test", "0.1.0", test_registry(), config);
420
421 let request: CallToolRequestParams = serde_json::from_value(json!({
422 "name": "echo",
423 "arguments": {
424 "message": "hi"
425 }
426 }))
427 .unwrap();
428 let result = handler.handle_call_tool(request).await;
429
430 assert_eq!(result.is_error, Some(true));
431 assert_eq!(first_text(&result), Some("tool call denied: echo disabled"));
432
433 let captured = events.lock();
434 assert_eq!(captured.len(), 1);
435 assert_eq!(captured[0].tool_name, "echo");
436 assert_eq!(captured[0].outcome, "denied");
437 drop(captured);
438 }
439
440 #[tokio::test]
441 async fn invalid_input_is_rejected_before_authorization() {
442 let calls = Arc::new(Mutex::new(Vec::new()));
443 let events = Arc::new(Mutex::new(Vec::new()));
444 let config = ServerConfig {
445 tool_authorizer: Some(Arc::new(RecordingAuthorizer {
446 calls: Arc::clone(&calls),
447 })),
448 tool_audit_sink: Some(Arc::new(RecordingAuditSink {
449 events: Arc::clone(&events),
450 })),
451 ..Default::default()
452 };
453 let handler = create_server("test", "0.1.0", test_registry(), config);
454
455 let request: CallToolRequestParams = serde_json::from_value(json!({
456 "name": "echo",
457 "arguments": {}
458 }))
459 .unwrap();
460 let result = handler.handle_call_tool(request).await;
461
462 assert_eq!(result.is_error, Some(true));
463 assert!(
464 first_text(&result)
465 .unwrap_or_default()
466 .starts_with("invalid tool input:")
467 );
468 assert!(calls.lock().is_empty());
469 assert_eq!(events.lock()[0].outcome, "invalid_input");
470 }
471
472 #[tokio::test]
473 async fn unknown_tool_is_rejected_before_authorization() {
474 let calls = Arc::new(Mutex::new(Vec::new()));
475 let events = Arc::new(Mutex::new(Vec::new()));
476 let config = ServerConfig {
477 tool_authorizer: Some(Arc::new(RecordingAuthorizer {
478 calls: Arc::clone(&calls),
479 })),
480 tool_audit_sink: Some(Arc::new(RecordingAuditSink {
481 events: Arc::clone(&events),
482 })),
483 ..Default::default()
484 };
485 let handler = create_server("test", "0.1.0", test_registry(), config);
486
487 let request: CallToolRequestParams = serde_json::from_value(json!({
488 "name": "missing",
489 "arguments": {}
490 }))
491 .unwrap();
492 let result = handler.handle_call_tool(request).await;
493
494 assert_eq!(result.is_error, Some(true));
495 assert_eq!(first_text(&result), Some("tool not found: missing"));
496 assert!(calls.lock().is_empty());
497 assert_eq!(events.lock()[0].outcome, "not_found");
498 }
499
500 #[tokio::test]
501 async fn test_max_input_bytes() {
502 let config = ServerConfig {
503 max_input_bytes: 8,
504 ..Default::default()
505 };
506 let handler = create_server("test", "0.1.0", test_registry(), config);
507
508 let request: CallToolRequestParams = serde_json::from_value(json!({
509 "name": "echo",
510 "arguments": {
511 "message": "hello"
512 }
513 }))
514 .unwrap();
515 let result = handler.handle_call_tool(request).await;
516
517 assert_eq!(result.is_error, Some(true));
518 assert_eq!(
519 first_text(&result),
520 Some("input too large: exceeds 8 bytes")
521 );
522 }
523
524 struct InvalidOutputTool {
525 definition: Definition,
526 }
527
528 #[async_trait]
529 impl Callable for InvalidOutputTool {
530 fn definition(&self) -> &Definition {
531 &self.definition
532 }
533
534 fn validate(&self, _input: &ToolInput) -> ValidationResult {
535 ValidationResult {
536 valid: true,
537 errors: Vec::new(),
538 }
539 }
540
541 async fn call(
542 &self,
543 _ctx: &Context,
544 _input: ToolInput,
545 ) -> rskit_errors::AppResult<ToolResult> {
546 Ok(ToolResult {
547 output: Some(json!({"sum": "bad"}).into()),
548 content: String::from("{\"sum\":\"bad\"}"),
549 is_error: false,
550 metadata: rskit_tool::ToolMetadata::new(),
551 })
552 }
553 }
554
555 #[tokio::test]
556 async fn test_output_schema_validation() {
557 let registry = Registry::new();
558 registry
559 .register(Box::new(InvalidOutputTool {
560 definition: Definition {
561 name: String::from("bad_output"),
562 description: String::from("Return invalid output"),
563 input_schema: rskit_tool::ToolSchema::new(
564 json!({"type": "object", "properties": {}}),
565 )
566 .unwrap(),
567 output_schema: Some(
568 rskit_tool::ToolSchema::new(json!({
569 "type": "object",
570 "properties": {"sum": {"type": "integer"}},
571 "required": ["sum"]
572 }))
573 .unwrap(),
574 ),
575 annotations: rskit_tool::Annotations::default(),
576 envelope: rskit_tool::Envelope::default(),
577 },
578 }))
579 .unwrap();
580 let handler = create_server("test", "0.1.0", Arc::new(registry), ServerConfig::default());
581
582 let request: CallToolRequestParams = serde_json::from_value(json!({
583 "name": "bad_output",
584 "arguments": {}
585 }))
586 .unwrap();
587 let result = handler.handle_call_tool(request).await;
588
589 assert_eq!(result.is_error, Some(true));
590 assert!(
591 first_text(&result)
592 .unwrap_or_default()
593 .starts_with("output validation error:")
594 );
595 }
596
597 #[tokio::test]
598 async fn test_prompts_resources_and_templates() {
599 let prompt: Prompt = serde_json::from_value(json!({
600 "name": "greet",
601 "description": "Render a greeting prompt",
602 "arguments": [{"name": "name", "required": true}]
603 }))
604 .unwrap();
605 let resource: Resource = serde_json::from_value(json!({
606 "uri": "memo://info",
607 "name": "info",
608 "mimeType": "text/plain"
609 }))
610 .unwrap();
611 let template: ResourceTemplate = serde_json::from_value(json!({
612 "uriTemplate": "memo://items/{id}",
613 "name": "item",
614 "mimeType": "text/plain"
615 }))
616 .unwrap();
617
618 let config = ServerConfig {
619 prompts: vec![PromptEntry::new(prompt, |request| async move {
620 let name = request
621 .arguments
622 .as_ref()
623 .and_then(|arguments| arguments.get("name"))
624 .and_then(serde_json::Value::as_str)
625 .unwrap_or_default()
626 .to_owned();
627 serde_json::from_value(json!({
628 "description": "Greeting prompt",
629 "messages": [{
630 "role": "user",
631 "content": {"type": "text", "text": format!("Say hello to {name}")}
632 }]
633 }))
634 .map_err(|err| invalid_params_error(err.to_string()))
635 })],
636 resources: vec![ResourceEntry::new(resource, |request| async move {
637 serde_json::from_value(json!({
638 "contents": [{
639 "uri": request.uri.clone(),
640 "mimeType": "text/plain",
641 "text": "info"
642 }]
643 }))
644 .map_err(|err| invalid_params_error(err.to_string()))
645 })],
646 resource_templates: vec![ResourceTemplateEntry::new(template, |request| async move {
647 serde_json::from_value(json!({
648 "contents": [{
649 "uri": request.uri.clone(),
650 "mimeType": "text/plain",
651 "text": format!("templated:{}", request.uri)
652 }]
653 }))
654 .map_err(|err| invalid_params_error(err.to_string()))
655 })],
656 ..Default::default()
657 };
658 let handler = create_server("test", "0.1.0", test_registry(), config);
659
660 let prompts = handler.mcp_prompts();
661 assert_eq!(prompt_name(&prompts[0]).as_deref(), Some("greet"));
662
663 let prompt_result = handler
664 .handle_get_prompt(
665 serde_json::from_value(json!({
666 "name": "greet",
667 "arguments": {"name": "World"}
668 }))
669 .unwrap(),
670 )
671 .await
672 .unwrap();
673 let prompt_json = serde_json::to_value(&prompt_result).unwrap();
674 assert_eq!(
675 prompt_json["messages"][0]["content"]["text"].as_str(),
676 Some("Say hello to World")
677 );
678
679 let resources = handler.mcp_resources();
680 assert_eq!(resource_uri(&resources[0]).as_deref(), Some("memo://info"));
681
682 let templates = handler.mcp_resource_templates();
683 assert_eq!(
684 resource_template_uri(&templates[0]).as_deref(),
685 Some("memo://items/{id}")
686 );
687
688 let resource_result = handler
689 .handle_read_resource(serde_json::from_value(json!({"uri": "memo://info"})).unwrap())
690 .await
691 .unwrap();
692 let resource_json = serde_json::to_value(&resource_result).unwrap();
693 assert_eq!(resource_json["contents"][0]["text"].as_str(), Some("info"));
694
695 let templated_result = handler
696 .handle_read_resource(
697 serde_json::from_value(json!({"uri": "memo://items/123"})).unwrap(),
698 )
699 .await
700 .unwrap();
701 let templated_json = serde_json::to_value(&templated_result).unwrap();
702 assert_eq!(
703 templated_json["contents"][0]["text"].as_str(),
704 Some("templated:memo://items/123")
705 );
706 }
707
708 #[tokio::test]
709 async fn test_prompt_and_resource_not_found_errors() {
710 let handler = create_server("test", "0.1.0", test_registry(), ServerConfig::default());
711
712 let prompt_error = handler
713 .handle_get_prompt(serde_json::from_value(json!({"name": "missing"})).unwrap())
714 .await
715 .expect_err("missing prompt is rejected");
716 assert!(prompt_error.message.contains("prompt not found"));
717
718 let resource_error = handler
719 .handle_read_resource(serde_json::from_value(json!({"uri": "memo://missing"})).unwrap())
720 .await
721 .expect_err("missing resource is rejected");
722 assert!(resource_error.message.contains("resource not found"));
723 }
724
725 #[test]
726 fn test_resource_template_matching_edges() {
727 assert!(resource_template_matches(
728 "memo://items/{id}",
729 "memo://items/123"
730 ));
731 assert!(resource_template_matches(
732 "memo://{tenant}/items/{id}/details",
733 "memo://acme/items/123/details"
734 ));
735 assert!(!resource_template_matches(
736 "memo://items/{id}",
737 "file://items/123"
738 ));
739 assert!(!resource_template_matches(
740 "memo://items/{id}/details",
741 "memo://items/123/summary"
742 ));
743 assert!(resource_template_matches(
744 "memo://literal",
745 "memo://literal"
746 ));
747 assert!(!resource_template_matches("memo://literal", "memo://other"));
748 }
749
750 fn first_text(result: &CallToolResult) -> Option<&str> {
751 result
752 .content
753 .first()
754 .and_then(|content| match &content.raw {
755 rmcp::model::RawContent::Text(text) => Some(text.text.as_ref()),
756 _ => None,
757 })
758 }
759}