Skip to main content

prism_mcp_rs/server/
discovery_handler.rs

1//! Discovery handler for the RPC discovery mechanism
2//!
3//! Module provides the handler for the `rpc.discover` method, allowing
4//! clients to introspect server capabilities and available methods at runtime.
5
6use 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
14/// Handler for RPC discovery requests
15pub struct DiscoveryHandler {
16    registry: MethodRegistry,
17}
18
19impl DiscoveryHandler {
20    /// Create a new discovery handler with the standard MCP method registry
21    pub fn new() -> Self {
22        Self {
23            registry: MethodRegistry::build_standard_registry(),
24        }
25    }
26
27    /// Create a discovery handler with a custom method registry
28    pub fn with_registry(registry: MethodRegistry) -> Self {
29        Self { registry }
30    }
31
32    /// Handle an rpc.discover request
33    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        // Filter methods based on the request
50        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        // Group methods by category
65        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            // Clear schema fields if not requested
79            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        // Build capabilities information if requested
91        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, // Can be customized based on server requirements
105                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        // Build metadata
117        let metadata = Some(DiscoveryMetadata {
118            server_name: Some(server_info.name.clone()),
119            server_version: Some(server_info.version.clone()),
120            documentation_url: None, // Can be customized
121            support_contact: None,   // Can be customized
122            rate_limits: None,       // Can be customized
123        });
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        // Test with no params (defaults to all)
152        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        // Test with client filter
170        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        // All methods should be client-to-server
182        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        // Test with category filter
200        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        // Should only have tool-related methods
210        for methods in result.methods.values() {
211            for method in methods {
212                assert!(method.name.starts_with("tools/"));
213            }
214        }
215    }
216}