Skip to main content

rskit_mcp/
server.rs

1//! MCP server backed by an rskit tool [`Registry`].
2//!
3//! Implements the MCP `ServerHandler` trait, delegating `tools/list`
4//! and `tools/call` to the registry while providing sensible defaults for the rest of the protocol.
5
6use 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
26/// An MCP server handler backed by an rskit [`Registry`].
27///
28/// Created via [`create_server`].
29pub 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
207/// Create an MCP [`ServerHandler`] backed by an rskit [`Registry`].
208///
209/// # Arguments
210///
211/// * `name` — server name advertised in `initialize` response
212/// * `version` — server version advertised in `initialize` response
213/// * `registry` — the tool registry to expose
214/// * `config` — optional server configuration (prefix, etc.)
215pub 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}