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