Skip to main content

prism_mcp_rs/protocol/
discovery.rs

1//! RPC Discovery Module for MCP Protocol
2//!
3//! Module implements the optional `rpc.discover` mechanism that allows clients
4//! to dynamically discover available methods, their parameters, and capabilities.
5//! This enables introspection of the MCP server's capabilities at runtime.
6
7use serde::{Deserialize, Serialize};
8use std::collections::HashMap;
9
10// ============================================================================
11// Discovery Types
12// ============================================================================
13
14/// Request for discovering available RPC methods
15#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
16pub struct DiscoverRequest {
17    /// Optional filter to limit discovery to specific categories
18    #[serde(skip_serializing_if = "Option::is_none")]
19    pub filter: Option<DiscoveryFilter>,
20
21    /// Whether to include detailed parameter schemas
22    #[serde(default = "default_include_schemas")]
23    pub include_schemas: bool,
24
25    /// Whether to include capability information
26    #[serde(default = "default_include_capabilities")]
27    pub include_capabilities: bool,
28}
29
30fn default_include_schemas() -> bool {
31    true
32}
33
34fn default_include_capabilities() -> bool {
35    true
36}
37
38/// Filter for discovery requests
39#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
40#[serde(rename_all = "lowercase")]
41pub enum DiscoveryFilter {
42    /// Discover only client methods (methods the client can call on the server)
43    Client,
44    /// Discover only server methods (methods the server can call on the client)
45    Server,
46    /// Discover only notification methods
47    Notifications,
48    /// Discover methods by category
49    Category(String),
50    /// Discover all methods (default)
51    All,
52}
53
54/// Response containing discovered RPC methods and capabilities
55#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
56pub struct DiscoverResult {
57    /// Protocol version
58    pub protocol_version: String,
59
60    /// Available methods grouped by category
61    pub methods: HashMap<String, Vec<MethodInfo>>,
62
63    /// Server capabilities (if requested)
64    #[serde(skip_serializing_if = "Option::is_none")]
65    pub capabilities: Option<DiscoveredCapabilities>,
66
67    /// Additional metadata
68    #[serde(skip_serializing_if = "Option::is_none")]
69    pub metadata: Option<DiscoveryMetadata>,
70}
71
72/// Information about a single RPC method
73#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
74pub struct MethodInfo {
75    /// Method name (e.g., "tools/list")
76    pub name: String,
77
78    /// Human-readable description of the method
79    #[serde(skip_serializing_if = "Option::is_none")]
80    pub description: Option<String>,
81
82    /// Method type
83    pub method_type: MethodType,
84
85    /// Direction of the method call
86    pub direction: MethodDirection,
87
88    /// JSON Schema for request parameters (if include_schemas is true)
89    #[serde(skip_serializing_if = "Option::is_none")]
90    pub params_schema: Option<serde_json::Value>,
91
92    /// JSON Schema for response result (if include_schemas is true)
93    #[serde(skip_serializing_if = "Option::is_none")]
94    pub result_schema: Option<serde_json::Value>,
95
96    /// Whether Method requires authentication
97    #[serde(default)]
98    pub requires_auth: bool,
99
100    /// Whether Method supports progress notifications
101    #[serde(default)]
102    pub supports_progress: bool,
103
104    /// Whether Method supports cancellation
105    #[serde(default)]
106    pub supports_cancellation: bool,
107
108    /// Tags for categorization
109    #[serde(skip_serializing_if = "Option::is_none")]
110    pub tags: Option<Vec<String>>,
111}
112
113/// Type of RPC method
114#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
115#[serde(rename_all = "lowercase")]
116pub enum MethodType {
117    /// Request-response method
118    Request,
119    /// One-way notification
120    Notification,
121    /// Subscription method
122    Subscription,
123}
124
125/// Direction of method invocation
126#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
127#[serde(rename_all = "snake_case")]
128pub enum MethodDirection {
129    /// Client calls server
130    ClientToServer,
131    /// Server calls client
132    ServerToClient,
133    /// Can be called in either direction
134    Bidirectional,
135}
136
137/// Discovered capabilities information
138#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
139pub struct DiscoveredCapabilities {
140    /// Server capabilities
141    #[serde(skip_serializing_if = "Option::is_none")]
142    pub server: Option<ServerCapabilityInfo>,
143
144    /// Required client capabilities
145    #[serde(skip_serializing_if = "Option::is_none")]
146    pub required_client: Option<ClientCapabilityInfo>,
147
148    /// Optional client capabilities that enhance functionality
149    #[serde(skip_serializing_if = "Option::is_none")]
150    pub optional_client: Option<ClientCapabilityInfo>,
151}
152
153/// Server capability information
154#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
155pub struct ServerCapabilityInfo {
156    /// Whether the server supports tools
157    pub tools: bool,
158
159    /// Whether the server supports resources
160    pub resources: bool,
161
162    /// Whether the server supports prompts
163    pub prompts: bool,
164
165    /// Whether the server supports logging
166    pub logging: bool,
167
168    /// Whether the server supports completions
169    pub completions: bool,
170
171    /// List of experimental capabilities
172    #[serde(skip_serializing_if = "Option::is_none")]
173    pub experimental: Option<Vec<String>>,
174}
175
176/// Client capability information
177#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
178pub struct ClientCapabilityInfo {
179    /// Whether the client should support sampling
180    pub sampling: bool,
181
182    /// Whether the client should support roots
183    pub roots: bool,
184
185    /// Whether the client should support elicitation
186    pub elicitation: bool,
187
188    /// List of experimental capabilities
189    #[serde(skip_serializing_if = "Option::is_none")]
190    pub experimental: Option<Vec<String>>,
191}
192
193/// Additional metadata for discovery
194#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
195pub struct DiscoveryMetadata {
196    /// Server implementation name
197    #[serde(skip_serializing_if = "Option::is_none")]
198    pub server_name: Option<String>,
199
200    /// Server implementation version
201    #[serde(skip_serializing_if = "Option::is_none")]
202    pub server_version: Option<String>,
203
204    /// API documentation URL
205    #[serde(skip_serializing_if = "Option::is_none")]
206    pub documentation_url: Option<String>,
207
208    /// Support contact information
209    #[serde(skip_serializing_if = "Option::is_none")]
210    pub support_contact: Option<String>,
211
212    /// Rate limiting information
213    #[serde(skip_serializing_if = "Option::is_none")]
214    pub rate_limits: Option<HashMap<String, RateLimitInfo>>,
215}
216
217/// Rate limiting information for methods
218#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
219pub struct RateLimitInfo {
220    /// Maximum requests per time window
221    pub max_requests: u32,
222
223    /// Time window in seconds
224    pub window_seconds: u32,
225
226    /// Whether rate limiting is per-method or global
227    #[serde(default)]
228    pub per_method: bool,
229}
230
231// ============================================================================
232// Discovery Implementation
233// ============================================================================
234
235/// Registry of available RPC methods for discovery
236pub struct MethodRegistry {
237    methods: Vec<MethodInfo>,
238}
239
240impl MethodRegistry {
241    /// Create a new method registry
242    pub fn new() -> Self {
243        Self {
244            methods: Vec::new(),
245        }
246    }
247
248    /// Register a new method
249    pub fn register(&mut self, method: MethodInfo) {
250        self.methods.push(method);
251    }
252
253    /// Get all registered methods
254    pub fn get_methods(&self) -> &[MethodInfo] {
255        &self.methods
256    }
257
258    /// Filter methods by category
259    pub fn filter_by_category(&self, category: &str) -> Vec<&MethodInfo> {
260        self.methods
261            .iter()
262            .filter(|m| {
263                m.tags
264                    .as_ref()
265                    .is_some_and(|tags| tags.contains(&category.to_string()))
266            })
267            .collect()
268    }
269
270    /// Filter methods by direction
271    pub fn filter_by_direction(&self, direction: MethodDirection) -> Vec<&MethodInfo> {
272        self.methods
273            .iter()
274            .filter(|m| m.direction == direction)
275            .collect()
276    }
277
278    /// Filter methods by type
279    pub fn filter_by_type(&self, method_type: MethodType) -> Vec<&MethodInfo> {
280        self.methods
281            .iter()
282            .filter(|m| m.method_type == method_type)
283            .collect()
284    }
285
286    /// Build the standard MCP method registry
287    pub fn build_standard_registry() -> Self {
288        let mut registry = Self::new();
289
290        // Core protocol methods
291        registry.register(MethodInfo {
292            name: "initialize".to_string(),
293            description: Some("Initialize the MCP connection".to_string()),
294            method_type: MethodType::Request,
295            direction: MethodDirection::ClientToServer,
296            params_schema: None,
297            result_schema: None,
298            requires_auth: false,
299            supports_progress: false,
300            supports_cancellation: false,
301            tags: Some(vec!["core".to_string(), "initialization".to_string()]),
302        });
303
304        registry.register(MethodInfo {
305            name: "ping".to_string(),
306            description: Some("Check connection liveness".to_string()),
307            method_type: MethodType::Request,
308            direction: MethodDirection::Bidirectional,
309            params_schema: None,
310            result_schema: None,
311            requires_auth: false,
312            supports_progress: false,
313            supports_cancellation: false,
314            tags: Some(vec!["core".to_string(), "health".to_string()]),
315        });
316
317        // Tool methods
318        registry.register(MethodInfo {
319            name: "tools/list".to_string(),
320            description: Some("List available tools".to_string()),
321            method_type: MethodType::Request,
322            direction: MethodDirection::ClientToServer,
323            params_schema: None,
324            result_schema: None,
325            requires_auth: false,
326            supports_progress: false,
327            supports_cancellation: true,
328            tags: Some(vec!["tools".to_string()]),
329        });
330
331        registry.register(MethodInfo {
332            name: "tools/call".to_string(),
333            description: Some("Call a tool with arguments".to_string()),
334            method_type: MethodType::Request,
335            direction: MethodDirection::ClientToServer,
336            params_schema: None,
337            result_schema: None,
338            requires_auth: false,
339            supports_progress: true,
340            supports_cancellation: true,
341            tags: Some(vec!["tools".to_string()]),
342        });
343
344        // Resource methods
345        registry.register(MethodInfo {
346            name: "resources/list".to_string(),
347            description: Some("List available resources".to_string()),
348            method_type: MethodType::Request,
349            direction: MethodDirection::ClientToServer,
350            params_schema: None,
351            result_schema: None,
352            requires_auth: false,
353            supports_progress: false,
354            supports_cancellation: true,
355            tags: Some(vec!["resources".to_string()]),
356        });
357
358        registry.register(MethodInfo {
359            name: "resources/read".to_string(),
360            description: Some("Read a resource by URI".to_string()),
361            method_type: MethodType::Request,
362            direction: MethodDirection::ClientToServer,
363            params_schema: None,
364            result_schema: None,
365            requires_auth: false,
366            supports_progress: true,
367            supports_cancellation: true,
368            tags: Some(vec!["resources".to_string()]),
369        });
370
371        // Prompt methods
372        registry.register(MethodInfo {
373            name: "prompts/list".to_string(),
374            description: Some("List available prompts".to_string()),
375            method_type: MethodType::Request,
376            direction: MethodDirection::ClientToServer,
377            params_schema: None,
378            result_schema: None,
379            requires_auth: false,
380            supports_progress: false,
381            supports_cancellation: true,
382            tags: Some(vec!["prompts".to_string()]),
383        });
384
385        registry.register(MethodInfo {
386            name: "prompts/get".to_string(),
387            description: Some("Get a prompt by name".to_string()),
388            method_type: MethodType::Request,
389            direction: MethodDirection::ClientToServer,
390            params_schema: None,
391            result_schema: None,
392            requires_auth: false,
393            supports_progress: false,
394            supports_cancellation: true,
395            tags: Some(vec!["prompts".to_string()]),
396        });
397
398        // Sampling methods (server to client)
399        registry.register(MethodInfo {
400            name: "sampling/createMessage".to_string(),
401            description: Some("Request message generation from client's LLM".to_string()),
402            method_type: MethodType::Request,
403            direction: MethodDirection::ServerToClient,
404            params_schema: None,
405            result_schema: None,
406            requires_auth: false,
407            supports_progress: true,
408            supports_cancellation: true,
409            tags: Some(vec!["sampling".to_string(), "llm".to_string()]),
410        });
411
412        // Roots methods (server to client)
413        registry.register(MethodInfo {
414            name: "roots/list".to_string(),
415            description: Some("List client's root directories".to_string()),
416            method_type: MethodType::Request,
417            direction: MethodDirection::ServerToClient,
418            params_schema: None,
419            result_schema: None,
420            requires_auth: false,
421            supports_progress: false,
422            supports_cancellation: false,
423            tags: Some(vec!["roots".to_string(), "filesystem".to_string()]),
424        });
425
426        // Elicitation methods (server to client)
427        registry.register(MethodInfo {
428            name: "elicitation/create".to_string(),
429            description: Some("Request user input through a form".to_string()),
430            method_type: MethodType::Request,
431            direction: MethodDirection::ServerToClient,
432            params_schema: None,
433            result_schema: None,
434            requires_auth: false,
435            supports_progress: false,
436            supports_cancellation: true,
437            tags: Some(vec!["elicitation".to_string(), "user-input".to_string()]),
438        });
439
440        // Completion methods
441        registry.register(MethodInfo {
442            name: "completion/complete".to_string(),
443            description: Some("Get completion suggestions".to_string()),
444            method_type: MethodType::Request,
445            direction: MethodDirection::ClientToServer,
446            params_schema: None,
447            result_schema: None,
448            requires_auth: false,
449            supports_progress: false,
450            supports_cancellation: true,
451            tags: Some(vec!["completion".to_string(), "autocomplete".to_string()]),
452        });
453
454        // Logging methods
455        registry.register(MethodInfo {
456            name: "logging/setLevel".to_string(),
457            description: Some("Set logging level".to_string()),
458            method_type: MethodType::Request,
459            direction: MethodDirection::ClientToServer,
460            params_schema: None,
461            result_schema: None,
462            requires_auth: false,
463            supports_progress: false,
464            supports_cancellation: false,
465            tags: Some(vec!["logging".to_string()]),
466        });
467
468        // Discovery method itself
469        registry.register(MethodInfo {
470            name: "rpc.discover".to_string(),
471            description: Some("Discover available RPC methods and capabilities".to_string()),
472            method_type: MethodType::Request,
473            direction: MethodDirection::ClientToServer,
474            params_schema: None,
475            result_schema: None,
476            requires_auth: false,
477            supports_progress: false,
478            supports_cancellation: false,
479            tags: Some(vec!["discovery".to_string(), "meta".to_string()]),
480        });
481
482        // Notification methods
483        registry.register(MethodInfo {
484            name: "notifications/initialized".to_string(),
485            description: Some("Client initialization complete notification".to_string()),
486            method_type: MethodType::Notification,
487            direction: MethodDirection::ClientToServer,
488            params_schema: None,
489            result_schema: None,
490            requires_auth: false,
491            supports_progress: false,
492            supports_cancellation: false,
493            tags: Some(vec![
494                "notifications".to_string(),
495                "initialization".to_string(),
496            ]),
497        });
498
499        registry.register(MethodInfo {
500            name: "notifications/cancelled".to_string(),
501            description: Some("Request cancellation notification".to_string()),
502            method_type: MethodType::Notification,
503            direction: MethodDirection::Bidirectional,
504            params_schema: None,
505            result_schema: None,
506            requires_auth: false,
507            supports_progress: false,
508            supports_cancellation: false,
509            tags: Some(vec!["notifications".to_string(), "control".to_string()]),
510        });
511
512        registry.register(MethodInfo {
513            name: "notifications/progress".to_string(),
514            description: Some("Progress update notification".to_string()),
515            method_type: MethodType::Notification,
516            direction: MethodDirection::Bidirectional,
517            params_schema: None,
518            result_schema: None,
519            requires_auth: false,
520            supports_progress: false,
521            supports_cancellation: false,
522            tags: Some(vec!["notifications".to_string(), "progress".to_string()]),
523        });
524
525        registry
526    }
527}
528
529impl Default for MethodRegistry {
530    fn default() -> Self {
531        Self::new()
532    }
533}
534
535#[cfg(test)]
536mod tests {
537    use super::*;
538
539    #[test]
540    fn test_method_registry_creation() {
541        let registry = MethodRegistry::build_standard_registry();
542        let methods = registry.get_methods();
543
544        // Check that we have registered methods
545        assert!(!methods.is_empty());
546
547        // Check for specific core methods
548        assert!(methods.iter().any(|m| m.name == "initialize"));
549        assert!(methods.iter().any(|m| m.name == "ping"));
550        assert!(methods.iter().any(|m| m.name == "rpc.discover"));
551    }
552
553    #[test]
554    fn test_method_filtering() {
555        let registry = MethodRegistry::build_standard_registry();
556
557        // Test filtering by direction
558        let client_to_server = registry.filter_by_direction(MethodDirection::ClientToServer);
559        assert!(!client_to_server.is_empty());
560
561        let server_to_client = registry.filter_by_direction(MethodDirection::ServerToClient);
562        assert!(!server_to_client.is_empty());
563
564        // Test filtering by type
565        let requests = registry.filter_by_type(MethodType::Request);
566        assert!(!requests.is_empty());
567
568        let notifications = registry.filter_by_type(MethodType::Notification);
569        assert!(!notifications.is_empty());
570
571        // Test filtering by category
572        let tool_methods = registry.filter_by_category("tools");
573        assert!(!tool_methods.is_empty());
574        assert!(tool_methods.iter().all(|m| m.name.starts_with("tools/")));
575    }
576
577    #[test]
578    fn test_discover_request_serialization() {
579        let request = DiscoverRequest {
580            filter: Some(DiscoveryFilter::Client),
581            include_schemas: true,
582            include_capabilities: true,
583        };
584
585        let json = serde_json::to_string(&request).unwrap();
586        let deserialized: DiscoverRequest = serde_json::from_str(&json).unwrap();
587
588        assert_eq!(request, deserialized);
589    }
590}