prism_mcp_rs/server/
discovery_handler.rs1use serde_json::Value;
7use std::collections::HashMap;
8
9use crate::core::error::{McpError, McpResult};
10use crate::protocol::discovery::*;
11use crate::protocol::types::{ServerCapabilities, ServerInfo};
12use crate::protocol::LATEST_PROTOCOL_VERSION;
13
14pub struct DiscoveryHandler {
16 registry: MethodRegistry,
17}
18
19impl DiscoveryHandler {
20 pub fn new() -> Self {
22 Self {
23 registry: MethodRegistry::build_standard_registry(),
24 }
25 }
26
27 pub fn with_registry(registry: MethodRegistry) -> Self {
29 Self { registry }
30 }
31
32 pub async fn handle(
34 &self,
35 server_info: &ServerInfo,
36 capabilities: &ServerCapabilities,
37 params: Option<Value>,
38 ) -> McpResult<DiscoverResult> {
39 let request: DiscoverRequest = match params {
40 Some(p) => serde_json::from_value(p)
41 .map_err(|e| McpError::Validation(format!("Invalid discover params: {e}")))?,
42 None => DiscoverRequest {
43 filter: Some(DiscoveryFilter::All),
44 include_schemas: true,
45 include_capabilities: true,
46 },
47 };
48
49 let filtered_methods = match &request.filter {
51 Some(DiscoveryFilter::Client) => self
52 .registry
53 .filter_by_direction(MethodDirection::ClientToServer),
54 Some(DiscoveryFilter::Server) => self
55 .registry
56 .filter_by_direction(MethodDirection::ServerToClient),
57 Some(DiscoveryFilter::Notifications) => {
58 self.registry.filter_by_type(MethodType::Notification)
59 }
60 Some(DiscoveryFilter::Category(category)) => self.registry.filter_by_category(category),
61 Some(DiscoveryFilter::All) | None => self.registry.get_methods().iter().collect(),
62 };
63
64 let mut methods_by_category: HashMap<String, Vec<MethodInfo>> = HashMap::new();
66
67 for method in filtered_methods {
68 let category = if let Some(tags) = &method.tags {
69 tags.first()
70 .cloned()
71 .unwrap_or_else(|| "uncategorized".to_string())
72 } else {
73 "uncategorized".to_string()
74 };
75
76 let mut method_info = method.clone();
77
78 if !request.include_schemas {
80 method_info.params_schema = None;
81 method_info.result_schema = None;
82 }
83
84 methods_by_category
85 .entry(category)
86 .or_default()
87 .push(method_info);
88 }
89
90 let discovered_capabilities = if request.include_capabilities {
92 Some(DiscoveredCapabilities {
93 server: Some(ServerCapabilityInfo {
94 tools: capabilities.tools.is_some(),
95 resources: capabilities.resources.is_some(),
96 prompts: capabilities.prompts.is_some(),
97 logging: capabilities.logging.is_some(),
98 completions: capabilities.completions.is_some(),
99 experimental: capabilities
100 .experimental
101 .as_ref()
102 .map(|exp| exp.keys().cloned().collect()),
103 }),
104 required_client: None, optional_client: Some(ClientCapabilityInfo {
106 sampling: true,
107 roots: true,
108 elicitation: true,
109 experimental: None,
110 }),
111 })
112 } else {
113 None
114 };
115
116 let metadata = Some(DiscoveryMetadata {
118 server_name: Some(server_info.name.clone()),
119 server_version: Some(server_info.version.clone()),
120 documentation_url: None, support_contact: None, rate_limits: None, });
124
125 Ok(DiscoverResult {
126 protocol_version: LATEST_PROTOCOL_VERSION.to_string(),
127 methods: methods_by_category,
128 capabilities: discovered_capabilities,
129 metadata,
130 })
131 }
132}
133
134impl Default for DiscoveryHandler {
135 fn default() -> Self {
136 Self::new()
137 }
138}
139
140#[cfg(test)]
141mod tests {
142 use super::*;
143 use crate::protocol::types::Implementation;
144
145 #[tokio::test]
146 async fn test_discovery_handler() {
147 let handler = DiscoveryHandler::new();
148 let server_info = Implementation::new("test-server", "1.0.0");
149 let capabilities = ServerCapabilities::default();
150
151 let result = handler
153 .handle(&server_info, &capabilities, None)
154 .await
155 .unwrap();
156
157 assert_eq!(result.protocol_version, LATEST_PROTOCOL_VERSION);
158 assert!(!result.methods.is_empty());
159 assert!(result.capabilities.is_some());
160 assert!(result.metadata.is_some());
161 }
162
163 #[tokio::test]
164 async fn test_discovery_with_filter() {
165 let handler = DiscoveryHandler::new();
166 let server_info = Implementation::new("test-server", "1.0.0");
167 let capabilities = ServerCapabilities::default();
168
169 let params = serde_json::json!({
171 "filter": "client",
172 "include_schemas": false,
173 "include_capabilities": false
174 });
175
176 let result = handler
177 .handle(&server_info, &capabilities, Some(params))
178 .await
179 .unwrap();
180
181 for methods in result.methods.values() {
183 for method in methods {
184 assert_eq!(method.direction, MethodDirection::ClientToServer);
185 assert!(method.params_schema.is_none());
186 assert!(method.result_schema.is_none());
187 }
188 }
189
190 assert!(result.capabilities.is_none());
191 }
192
193 #[tokio::test]
194 async fn test_discovery_with_category_filter() {
195 let handler = DiscoveryHandler::new();
196 let server_info = Implementation::new("test-server", "1.0.0");
197 let capabilities = ServerCapabilities::default();
198
199 let params = serde_json::json!({
201 "filter": {"category": "tools"}
202 });
203
204 let result = handler
205 .handle(&server_info, &capabilities, Some(params))
206 .await
207 .unwrap();
208
209 for methods in result.methods.values() {
211 for method in methods {
212 assert!(method.name.starts_with("tools/"));
213 }
214 }
215 }
216}