1use bon::Builder;
2use rmcp::{
3 handler::server::ServerHandler,
4 model::{
5 CallToolRequestParams, CallToolResult, ErrorData, Implementation, InitializeResult,
6 ListToolsResult, PaginatedRequestParams, ProtocolVersion, ServerCapabilities,
7 ToolsCapability,
8 },
9 service::{RequestContext, RoleServer},
10};
11use rmcp_actix_web::transport::AuthorizationHeader;
12use serde_json::Value;
13use std::sync::Arc;
14
15use reqwest::header::HeaderMap;
16use url::Url;
17
18use crate::error::Error;
19use crate::filter::ToolFilter;
20use crate::tool::{Tool, ToolCollection, ToolMetadata};
21use crate::transformer::ResponseTransformer;
22use crate::{
23 config::{Authorization, AuthorizationMode},
24 spec::Filters,
25};
26use tracing::{debug, info, info_span, warn};
27
28#[derive(Clone, Builder)]
29pub struct Server {
30 pub openapi_spec: serde_json::Value,
31 #[builder(default)]
32 pub tool_collection: ToolCollection,
33 pub base_url: Url,
34 pub default_headers: Option<HeaderMap>,
35 pub filters: Option<Filters>,
36 #[builder(default)]
37 pub authorization_mode: AuthorizationMode,
38 pub name: Option<String>,
39 pub version: Option<String>,
40 pub title: Option<String>,
41 pub instructions: Option<String>,
42 #[builder(default)]
43 pub skip_tool_descriptions: bool,
44 #[builder(default)]
45 pub skip_parameter_descriptions: bool,
46 #[builder(default)]
47 pub insecure: bool,
48 pub response_transformer: Option<Arc<dyn ResponseTransformer>>,
55 pub tool_filter: Option<Arc<dyn ToolFilter>>,
58}
59
60impl Server {
61 pub fn new(
63 openapi_spec: serde_json::Value,
64 base_url: Url,
65 default_headers: Option<HeaderMap>,
66 filters: Option<Filters>,
67 skip_tool_descriptions: bool,
68 skip_parameter_descriptions: bool,
69 insecure: bool,
70 ) -> Self {
71 Self {
72 openapi_spec,
73 tool_collection: ToolCollection::new(),
74 base_url,
75 default_headers,
76 filters,
77 authorization_mode: AuthorizationMode::default(),
78 name: None,
79 version: None,
80 title: None,
81 instructions: None,
82 skip_tool_descriptions,
83 skip_parameter_descriptions,
84 insecure,
85 response_transformer: None,
86 tool_filter: None,
87 }
88 }
89
90 pub fn load_openapi_spec(&mut self) -> Result<(), Error> {
96 let span = info_span!("tool_registration");
97 let _enter = span.enter();
98
99 let spec = crate::spec::Spec::from_value(self.openapi_spec.clone())?;
101
102 let tools = spec.to_openapi_tools(
104 self.filters.as_ref(),
105 Some(self.base_url.clone()),
106 self.default_headers.clone(),
107 self.skip_tool_descriptions,
108 self.skip_parameter_descriptions,
109 self.insecure,
110 )?;
111
112 let tools = if let Some(ref transformer) = self.response_transformer {
114 tools
115 .into_iter()
116 .map(|mut tool| {
117 if let Some(schema) = tool.metadata.output_schema.take() {
118 tool.metadata.output_schema = Some(transformer.transform_schema(schema));
119 }
120 tool
121 })
122 .collect()
123 } else {
124 tools
125 };
126
127 self.tool_collection = ToolCollection::from_tools(tools);
128
129 info!(
130 tool_count = self.tool_collection.len(),
131 "Loaded tools from OpenAPI spec"
132 );
133
134 Ok(())
135 }
136
137 pub fn set_tool_transformer(
147 &mut self,
148 tool_name: &str,
149 transformer: Arc<dyn ResponseTransformer>,
150 ) -> Result<(), Error> {
151 self.tool_collection
152 .set_tool_transformer(tool_name, transformer)
153 }
154
155 pub fn set_tool_filter(&mut self, filter: Arc<dyn ToolFilter>) {
157 self.tool_filter = Some(filter);
158 }
159
160 #[must_use]
162 pub fn tool_count(&self) -> usize {
163 self.tool_collection.len()
164 }
165
166 #[must_use]
168 pub fn get_tool_names(&self) -> Vec<String> {
169 self.tool_collection.get_tool_names()
170 }
171
172 #[must_use]
174 pub fn has_tool(&self, name: &str) -> bool {
175 self.tool_collection.has_tool(name)
176 }
177
178 #[must_use]
180 pub fn get_tool(&self, name: &str) -> Option<&Tool> {
181 self.tool_collection.get_tool(name)
182 }
183
184 #[must_use]
186 pub fn get_tool_metadata(&self, name: &str) -> Option<&ToolMetadata> {
187 self.get_tool(name).map(|tool| &tool.metadata)
188 }
189
190 pub fn set_authorization_mode(&mut self, mode: AuthorizationMode) {
192 self.authorization_mode = mode;
193 }
194
195 pub fn authorization_mode(&self) -> AuthorizationMode {
197 self.authorization_mode
198 }
199
200 #[must_use]
202 pub fn get_tool_stats(&self) -> String {
203 self.tool_collection.get_stats()
204 }
205
206 pub fn validate_registry(&self) -> Result<(), Error> {
212 if self.tool_collection.is_empty() {
213 return Err(Error::McpError("No tools loaded".to_string()));
214 }
215 Ok(())
216 }
217
218 fn extract_openapi_title(&self) -> Option<String> {
220 self.openapi_spec
221 .get("info")?
222 .get("title")?
223 .as_str()
224 .map(|s| s.to_string())
225 }
226
227 fn extract_openapi_version(&self) -> Option<String> {
229 self.openapi_spec
230 .get("info")?
231 .get("version")?
232 .as_str()
233 .map(|s| s.to_string())
234 }
235
236 fn extract_openapi_description(&self) -> Option<String> {
238 self.openapi_spec
239 .get("info")?
240 .get("description")?
241 .as_str()
242 .map(|s| s.to_string())
243 }
244
245 fn extract_openapi_display_title(&self) -> Option<String> {
248 if let Some(display_title) = self
250 .openapi_spec
251 .get("info")
252 .and_then(|info| info.get("x-display-title"))
253 .and_then(|t| t.as_str())
254 {
255 return Some(display_title.to_string());
256 }
257
258 self.extract_openapi_title().map(|title| {
260 if title.to_lowercase().contains("server") {
261 title
262 } else {
263 format!("{} Server", title)
264 }
265 })
266 }
267}
268
269impl ServerHandler for Server {
270 fn get_info(&self) -> InitializeResult {
271 let server_name = self
273 .name
274 .clone()
275 .or_else(|| self.extract_openapi_title())
276 .unwrap_or_else(|| "OpenAPI MCP Server".to_string());
277
278 let server_version = self
280 .version
281 .clone()
282 .or_else(|| self.extract_openapi_version())
283 .unwrap_or_else(|| env!("CARGO_PKG_VERSION").to_string());
284
285 let server_title = self
287 .title
288 .clone()
289 .or_else(|| self.extract_openapi_display_title());
290
291 let instructions = self
293 .instructions
294 .clone()
295 .or_else(|| self.extract_openapi_description())
296 .or_else(|| Some("Exposes OpenAPI endpoints as MCP tools".to_string()));
297
298 let mut server_info = Implementation::new(server_name, server_version);
299 server_info.title = server_title;
300 server_info.description = self.extract_openapi_description();
301
302 let mut capabilities = ServerCapabilities::default();
303 capabilities.tools = Some(ToolsCapability {
304 list_changed: Some(false),
305 });
306
307 let mut result = InitializeResult::new(capabilities)
308 .with_protocol_version(ProtocolVersion::V_2024_11_05)
309 .with_server_info(server_info);
310 result.instructions = instructions;
311 result
312 }
313
314 async fn list_tools(
315 &self,
316 _request: Option<PaginatedRequestParams>,
317 context: RequestContext<RoleServer>,
318 ) -> Result<ListToolsResult, ErrorData> {
319 let span = info_span!("list_tools", tool_count = self.tool_collection.len());
320 let _enter = span.enter();
321
322 debug!("Processing MCP list_tools request");
323
324 let mut tools = self.tool_collection.to_mcp_tools();
326
327 if let Some(filter) = &self.tool_filter {
329 let mut filtered = Vec::with_capacity(tools.len());
330 for mcp_tool in tools {
331 if let Some(tool) = self.tool_collection.get_tool(&mcp_tool.name)
332 && filter.allow(tool, &context).await
333 {
334 filtered.push(mcp_tool);
335 }
336 }
337 tools = filtered;
338 }
339
340 info!(
341 returned_tools = tools.len(),
342 "MCP list_tools request completed successfully"
343 );
344
345 Ok(ListToolsResult {
346 meta: None,
347 tools,
348 next_cursor: None,
349 })
350 }
351
352 async fn call_tool(
353 &self,
354 request: CallToolRequestParams,
355 context: RequestContext<RoleServer>,
356 ) -> Result<CallToolResult, ErrorData> {
357 use crate::error::{ToolCallError, ToolCallValidationError};
358
359 let span = info_span!(
360 "call_tool",
361 tool_name = %request.name
362 );
363 let _enter = span.enter();
364
365 debug!(
366 tool_name = %request.name,
367 has_arguments = !request.arguments.as_ref().unwrap_or(&serde_json::Map::new()).is_empty(),
368 "Processing MCP call_tool request"
369 );
370
371 let allowed_tools: Vec<&Tool> = match &self.tool_filter {
373 None => self.tool_collection.iter().collect(),
374 Some(filter) => {
375 let mut allowed = Vec::new();
376 for tool in self.tool_collection.iter() {
377 if filter.allow(tool, &context).await {
378 allowed.push(tool);
379 }
380 }
381 allowed
382 }
383 };
384
385 let tool = allowed_tools
387 .iter()
388 .find(|t| t.metadata.name == request.name);
389
390 let tool = match tool {
391 Some(t) => *t,
392 None => {
393 let available_names: Vec<&str> = allowed_tools
394 .iter()
395 .map(|t| t.metadata.name.as_str())
396 .collect();
397
398 let error = ToolCallError::Validation(ToolCallValidationError::tool_not_found(
400 request.name.to_string(),
401 &available_names,
402 ));
403
404 warn!(
405 tool_name = %request.name,
406 success = false,
407 error = %error,
408 "MCP call_tool request failed - tool not found or filtered"
409 );
410
411 return Err(error.into());
412 }
413 };
414
415 let arguments = request.arguments.unwrap_or_default();
416 let arguments_value = Value::Object(arguments);
417
418 let auth_header = context.extensions.get::<AuthorizationHeader>().cloned();
420
421 if auth_header.is_some() {
422 debug!("Authorization header is present");
423 }
424
425 let authorization = Authorization::from_mode(self.authorization_mode, auth_header);
427
428 let server_transformer = self
430 .response_transformer
431 .as_ref()
432 .map(|t| t.as_ref() as &dyn ResponseTransformer);
433
434 match tool
436 .call(&arguments_value, authorization, server_transformer)
437 .await
438 {
439 Ok(result) => {
440 info!(
441 tool_name = %request.name,
442 success = true,
443 "MCP call_tool request completed successfully"
444 );
445 Ok(result)
446 }
447 Err(e) => {
448 warn!(
449 tool_name = %request.name,
450 success = false,
451 error = %e,
452 "MCP call_tool request failed"
453 );
454 Err(e.into())
456 }
457 }
458 }
459}
460
461#[cfg(test)]
462mod tests {
463 use super::*;
464 use crate::error::ToolCallValidationError;
465 use crate::{HttpClient, ToolCallError, ToolMetadata};
466 use serde_json::json;
467
468 #[test]
469 fn server_stores_insecure_flag() {
470 let server = Server::new(
471 serde_json::Value::Null,
472 url::Url::parse("http://example.com").unwrap(),
473 None,
474 None,
475 false,
476 false,
477 true,
478 );
479 assert!(server.insecure);
480 }
481
482 #[test]
483 fn server_insecure_defaults_to_false() {
484 let server = Server::new(
485 serde_json::Value::Null,
486 url::Url::parse("http://example.com").unwrap(),
487 None,
488 None,
489 false,
490 false,
491 false,
492 );
493 assert!(!server.insecure);
494 }
495
496 #[test]
497 fn test_tool_not_found_error_with_suggestions() {
498 let tool1_metadata = ToolMetadata {
500 name: "getPetById".to_string(),
501 title: Some("Get Pet by ID".to_string()),
502 description: Some("Find pet by ID".to_string()),
503 parameters: json!({
504 "type": "object",
505 "properties": {
506 "petId": {
507 "type": "integer"
508 }
509 },
510 "required": ["petId"]
511 }),
512 output_schema: None,
513 method: "GET".to_string(),
514 path: "/pet/{petId}".to_string(),
515 security: None,
516 parameter_mappings: std::collections::HashMap::new(),
517 };
518
519 let tool2_metadata = ToolMetadata {
520 name: "getPetsByStatus".to_string(),
521 title: Some("Find Pets by Status".to_string()),
522 description: Some("Find pets by status".to_string()),
523 parameters: json!({
524 "type": "object",
525 "properties": {
526 "status": {
527 "type": "array",
528 "items": {
529 "type": "string"
530 }
531 }
532 },
533 "required": ["status"]
534 }),
535 output_schema: None,
536 method: "GET".to_string(),
537 path: "/pet/findByStatus".to_string(),
538 security: None,
539 parameter_mappings: std::collections::HashMap::new(),
540 };
541
542 let http_client = HttpClient::new();
544 let tool1 = Tool::new(tool1_metadata, http_client.clone()).unwrap();
545 let tool2 = Tool::new(tool2_metadata, http_client.clone()).unwrap();
546
547 let mut server = Server::new(
549 serde_json::Value::Null,
550 url::Url::parse("http://example.com").unwrap(),
551 None,
552 None,
553 false,
554 false,
555 false,
556 );
557 server.tool_collection = ToolCollection::from_tools(vec![tool1, tool2]);
558
559 let tool_names = server.get_tool_names();
561 let tool_name_refs: Vec<&str> = tool_names.iter().map(|s| s.as_str()).collect();
562
563 let error = ToolCallError::Validation(ToolCallValidationError::tool_not_found(
564 "getPetByID".to_string(),
565 &tool_name_refs,
566 ));
567 let error_data: ErrorData = error.into();
568 let error_json = serde_json::to_value(&error_data).unwrap();
569
570 insta::assert_json_snapshot!(error_json);
572 }
573
574 #[test]
575 fn test_tool_not_found_error_no_suggestions() {
576 let tool_metadata = ToolMetadata {
578 name: "getPetById".to_string(),
579 title: Some("Get Pet by ID".to_string()),
580 description: Some("Find pet by ID".to_string()),
581 parameters: json!({
582 "type": "object",
583 "properties": {
584 "petId": {
585 "type": "integer"
586 }
587 },
588 "required": ["petId"]
589 }),
590 output_schema: None,
591 method: "GET".to_string(),
592 path: "/pet/{petId}".to_string(),
593 security: None,
594 parameter_mappings: std::collections::HashMap::new(),
595 };
596
597 let tool = Tool::new(tool_metadata, HttpClient::new()).unwrap();
599
600 let mut server = Server::new(
602 serde_json::Value::Null,
603 url::Url::parse("http://example.com").unwrap(),
604 None,
605 None,
606 false,
607 false,
608 false,
609 );
610 server.tool_collection = ToolCollection::from_tools(vec![tool]);
611
612 let tool_names = server.get_tool_names();
614 let tool_name_refs: Vec<&str> = tool_names.iter().map(|s| s.as_str()).collect();
615
616 let error = ToolCallError::Validation(ToolCallValidationError::tool_not_found(
617 "completelyUnrelatedToolName".to_string(),
618 &tool_name_refs,
619 ));
620 let error_data: ErrorData = error.into();
621 let error_json = serde_json::to_value(&error_data).unwrap();
622
623 insta::assert_json_snapshot!(error_json);
625 }
626
627 #[test]
628 fn test_validation_error_converted_to_error_data() {
629 let error = ToolCallError::Validation(ToolCallValidationError::InvalidParameters {
631 violations: vec![crate::error::ValidationError::invalid_parameter(
632 "page".to_string(),
633 &["page_number".to_string(), "page_size".to_string()],
634 )],
635 });
636
637 let error_data: ErrorData = error.into();
638 let error_json = serde_json::to_value(&error_data).unwrap();
639
640 assert_eq!(error_json["code"], -32602); insta::assert_json_snapshot!(error_json);
645 }
646
647 #[test]
648 fn test_extract_openapi_info_with_full_spec() {
649 let openapi_spec = json!({
650 "openapi": "3.0.0",
651 "info": {
652 "title": "Pet Store API",
653 "version": "2.1.0",
654 "description": "A sample API for managing pets"
655 },
656 "paths": {}
657 });
658
659 let server = Server::new(
660 openapi_spec,
661 url::Url::parse("http://example.com").unwrap(),
662 None,
663 None,
664 false,
665 false,
666 false,
667 );
668
669 assert_eq!(
670 server.extract_openapi_title(),
671 Some("Pet Store API".to_string())
672 );
673 assert_eq!(server.extract_openapi_version(), Some("2.1.0".to_string()));
674 assert_eq!(
675 server.extract_openapi_description(),
676 Some("A sample API for managing pets".to_string())
677 );
678 }
679
680 #[test]
681 fn test_extract_openapi_info_with_minimal_spec() {
682 let openapi_spec = json!({
683 "openapi": "3.0.0",
684 "info": {
685 "title": "My API",
686 "version": "1.0.0"
687 },
688 "paths": {}
689 });
690
691 let server = Server::new(
692 openapi_spec,
693 url::Url::parse("http://example.com").unwrap(),
694 None,
695 None,
696 false,
697 false,
698 false,
699 );
700
701 assert_eq!(server.extract_openapi_title(), Some("My API".to_string()));
702 assert_eq!(server.extract_openapi_version(), Some("1.0.0".to_string()));
703 assert_eq!(server.extract_openapi_description(), None);
704 }
705
706 #[test]
707 fn test_extract_openapi_info_with_invalid_spec() {
708 let openapi_spec = json!({
709 "invalid": "spec"
710 });
711
712 let server = Server::new(
713 openapi_spec,
714 url::Url::parse("http://example.com").unwrap(),
715 None,
716 None,
717 false,
718 false,
719 false,
720 );
721
722 assert_eq!(server.extract_openapi_title(), None);
723 assert_eq!(server.extract_openapi_version(), None);
724 assert_eq!(server.extract_openapi_description(), None);
725 }
726
727 #[test]
728 fn test_get_info_fallback_hierarchy_custom_metadata() {
729 let server = Server::new(
730 serde_json::Value::Null,
731 url::Url::parse("http://example.com").unwrap(),
732 None,
733 None,
734 false,
735 false,
736 false,
737 );
738
739 let mut server = server;
741 server.name = Some("Custom Server".to_string());
742 server.version = Some("3.0.0".to_string());
743 server.instructions = Some("Custom instructions".to_string());
744
745 let result = server.get_info();
746
747 assert_eq!(result.server_info.name, "Custom Server");
748 assert_eq!(result.server_info.version, "3.0.0");
749 assert_eq!(result.instructions, Some("Custom instructions".to_string()));
750 }
751
752 #[test]
753 fn test_get_info_fallback_hierarchy_openapi_spec() {
754 let openapi_spec = json!({
755 "openapi": "3.0.0",
756 "info": {
757 "title": "OpenAPI Server",
758 "version": "1.5.0",
759 "description": "Server from OpenAPI spec"
760 },
761 "paths": {}
762 });
763
764 let server = Server::new(
765 openapi_spec,
766 url::Url::parse("http://example.com").unwrap(),
767 None,
768 None,
769 false,
770 false,
771 false,
772 );
773
774 let result = server.get_info();
775
776 assert_eq!(result.server_info.name, "OpenAPI Server");
777 assert_eq!(result.server_info.version, "1.5.0");
778 assert_eq!(
779 result.instructions,
780 Some("Server from OpenAPI spec".to_string())
781 );
782 }
783
784 #[test]
785 fn test_get_info_fallback_hierarchy_defaults() {
786 let server = Server::new(
787 serde_json::Value::Null,
788 url::Url::parse("http://example.com").unwrap(),
789 None,
790 None,
791 false,
792 false,
793 false,
794 );
795
796 let result = server.get_info();
797
798 assert_eq!(result.server_info.name, "OpenAPI MCP Server");
799 assert_eq!(result.server_info.version, env!("CARGO_PKG_VERSION"));
800 assert_eq!(
801 result.instructions,
802 Some("Exposes OpenAPI endpoints as MCP tools".to_string())
803 );
804 }
805
806 #[test]
807 fn test_get_info_fallback_hierarchy_mixed() {
808 let openapi_spec = json!({
809 "openapi": "3.0.0",
810 "info": {
811 "title": "OpenAPI Server",
812 "version": "2.5.0",
813 "description": "Server from OpenAPI spec"
814 },
815 "paths": {}
816 });
817
818 let mut server = Server::new(
819 openapi_spec,
820 url::Url::parse("http://example.com").unwrap(),
821 None,
822 None,
823 false,
824 false,
825 false,
826 );
827
828 server.name = Some("Custom Server".to_string());
830 server.instructions = Some("Custom instructions".to_string());
831
832 let result = server.get_info();
833
834 assert_eq!(result.server_info.name, "Custom Server");
836 assert_eq!(result.server_info.version, "2.5.0");
838 assert_eq!(result.instructions, Some("Custom instructions".to_string()));
840 }
841}