Skip to main content

prism_mcp_rs/server/
handlers.rs

1//! MCP server request handlers
2//!
3//! Module provides specialized handlers for different types of MCP requests,
4//! implementing the business logic for each protocol operation.
5
6use serde_json::Value;
7use std::collections::HashMap;
8
9use crate::core::error::{McpError, McpResult};
10use crate::protocol::{messages::*, methods, types::*, LEGACY_PROTOCOL_VERSION};
11
12/// Handler for initialization requests
13pub struct InitializeHandler;
14
15impl InitializeHandler {
16    /// Process an initialize request
17    pub async fn handle(
18        server_info: &ServerInfo,
19        capabilities: &ServerCapabilities,
20        params: Option<Value>,
21    ) -> McpResult<InitializeResult> {
22        let params: InitializeParams = match params {
23            Some(p) => serde_json::from_value(p)
24                .map_err(|e| McpError::Validation(format!("Invalid initialize params: {e}")))?,
25            None => {
26                return Err(McpError::Validation(
27                    "Missing initialize parameters".to_string(),
28                ));
29            }
30        };
31
32        // Validate protocol version compatibility
33        if params.protocol_version != LEGACY_PROTOCOL_VERSION {
34            let protocol_version = params.protocol_version;
35            let expected = LEGACY_PROTOCOL_VERSION;
36            return Err(McpError::Protocol(format!(
37                "Unsupported protocol version: {protocol_version}. Expected: {expected}"
38            )));
39        }
40
41        // Validate client info
42        if params.client_info.name.is_empty() {
43            return Err(McpError::Validation(
44                "Client name cannot be empty".to_string(),
45            ));
46        }
47
48        if params.client_info.version.is_empty() {
49            return Err(McpError::Validation(
50                "Client version cannot be empty".to_string(),
51            ));
52        }
53
54        Ok(InitializeResult::new(
55            LEGACY_PROTOCOL_VERSION.to_string(),
56            capabilities.clone(),
57            server_info.clone(),
58        ))
59    }
60}
61
62/// Handler for tool-related requests
63pub struct ToolHandler;
64
65impl ToolHandler {
66    /// Handle tools/list request
67    pub async fn handle_list(
68        tools: &HashMap<String, crate::core::tool::Tool>,
69        params: Option<Value>,
70    ) -> McpResult<ListToolsResult> {
71        let _params: ListToolsParams = match params {
72            Some(p) => serde_json::from_value(p)
73                .map_err(|e| McpError::Validation(format!("Invalid list tools params: {e}")))?,
74            None => ListToolsParams::default(),
75        };
76
77        // Pagination support will be added in future versions
78        let tools: Vec<ToolInfo> = tools
79            .values()
80            .filter(|tool| tool.enabled)
81            .map(|tool| {
82                // Convert from core::tool::ToolInfo to protocol::types::ToolInfo
83                ToolInfo {
84                    name: tool.info.name.clone(),
85                    description: tool.info.description.clone(),
86                    input_schema: tool.info.input_schema.clone(),
87                    output_schema: tool.info.output_schema.clone(),
88                    annotations: None,
89                    icons: None,
90                    title: None,
91                    meta: None,
92                }
93            })
94            .collect();
95
96        Ok(ListToolsResult {
97            tools,
98            next_cursor: None,
99            meta: None,
100        })
101    }
102
103    /// Handle tools/call request
104    pub async fn handle_call(
105        tools: &HashMap<String, crate::core::tool::Tool>,
106        params: Option<Value>,
107    ) -> McpResult<CallToolResult> {
108        let params: CallToolParams = match params {
109            Some(p) => serde_json::from_value(p)
110                .map_err(|e| McpError::Validation(format!("Invalid call tool params: {e}")))?,
111            None => {
112                return Err(McpError::Validation(
113                    "Missing tool call parameters".to_string(),
114                ));
115            }
116        };
117
118        if params.name.is_empty() {
119            return Err(McpError::Validation(
120                "Tool name cannot be empty".to_string(),
121            ));
122        }
123
124        let tool = tools
125            .get(&params.name)
126            .ok_or_else(|| McpError::ToolNotFound(params.name.clone()))?;
127
128        if !tool.enabled {
129            let name = &params.name;
130            return Err(McpError::ToolNotFound(format!("Tool '{name}' is disabled")));
131        }
132
133        let arguments = params.arguments.unwrap_or_default();
134        let result = tool.handler.call(arguments).await?;
135
136        Ok(CallToolResult {
137            content: result.content,
138            is_error: result.is_error,
139            structured_content: None,
140            meta: None,
141        })
142    }
143}
144
145/// Handler for resource-related requests
146pub struct ResourceHandler;
147
148impl ResourceHandler {
149    /// Handle resources/list request
150    pub async fn handle_list(
151        resources: &HashMap<String, crate::core::resource::Resource>,
152        params: Option<Value>,
153    ) -> McpResult<ListResourcesResult> {
154        let _params: ListResourcesParams = match params {
155            Some(p) => serde_json::from_value(p)
156                .map_err(|e| McpError::Validation(format!("Invalid list resources params: {e}")))?,
157            None => ListResourcesParams::default(),
158        };
159
160        // Pagination support will be added in future versions
161        let resources: Vec<ResourceInfo> = resources
162            .values()
163            .map(|resource| {
164                // Convert from core::resource::ResourceInfo to protocol::types::ResourceInfo
165                ResourceInfo {
166                    uri: resource.info.uri.clone(),
167                    name: resource.info.name.clone(),
168                    description: resource.info.description.clone(),
169                    mime_type: resource.info.mime_type.clone(),
170                    annotations: None,
171                    size: None,
172                    icons: None,
173                    title: None,
174                    meta: None,
175                }
176            })
177            .collect();
178
179        Ok(ListResourcesResult {
180            resources,
181            next_cursor: None,
182            meta: None,
183        })
184    }
185
186    /// Handle resources/read request
187    pub async fn handle_read(
188        resources: &HashMap<String, crate::core::resource::Resource>,
189        params: Option<Value>,
190    ) -> McpResult<ReadResourceResult> {
191        let params: ReadResourceParams = match params {
192            Some(p) => serde_json::from_value(p)
193                .map_err(|e| McpError::Validation(format!("Invalid read resource params: {e}")))?,
194            None => {
195                return Err(McpError::Validation(
196                    "Missing resource read parameters".to_string(),
197                ));
198            }
199        };
200
201        if params.uri.is_empty() {
202            return Err(McpError::Validation(
203                "Resource URI cannot be empty".to_string(),
204            ));
205        }
206
207        let resource = resources
208            .get(&params.uri)
209            .ok_or_else(|| McpError::ResourceNotFound(params.uri.clone()))?;
210
211        // Query parameter extraction from URI will be implemented in future versions
212        let query_params = HashMap::new();
213        let contents = resource.handler.read(&params.uri, &query_params).await?;
214
215        Ok(ReadResourceResult {
216            contents,
217            meta: None,
218        })
219    }
220
221    /// Handle resources/subscribe request
222    pub async fn handle_subscribe(
223        resources: &HashMap<String, crate::core::resource::Resource>,
224        params: Option<Value>,
225    ) -> McpResult<SubscribeResourceResult> {
226        let params: SubscribeResourceParams = match params {
227            Some(p) => serde_json::from_value(p).map_err(|e| {
228                McpError::Validation(format!("Invalid subscribe resource params: {e}"))
229            })?,
230            None => {
231                return Err(McpError::Validation(
232                    "Missing resource subscribe parameters".to_string(),
233                ));
234            }
235        };
236
237        if params.uri.is_empty() {
238            return Err(McpError::Validation(
239                "Resource URI cannot be empty".to_string(),
240            ));
241        }
242
243        let resource = resources
244            .get(&params.uri)
245            .ok_or_else(|| McpError::ResourceNotFound(params.uri.clone()))?;
246
247        resource.handler.subscribe(&params.uri).await?;
248
249        Ok(SubscribeResourceResult { meta: None })
250    }
251
252    /// Handle resources/unsubscribe request
253    pub async fn handle_unsubscribe(
254        resources: &HashMap<String, crate::core::resource::Resource>,
255        params: Option<Value>,
256    ) -> McpResult<UnsubscribeResourceResult> {
257        let params: UnsubscribeResourceParams = match params {
258            Some(p) => serde_json::from_value(p).map_err(|e| {
259                McpError::Validation(format!("Invalid unsubscribe resource params: {e}"))
260            })?,
261            None => {
262                return Err(McpError::Validation(
263                    "Missing resource unsubscribe parameters".to_string(),
264                ));
265            }
266        };
267
268        if params.uri.is_empty() {
269            return Err(McpError::Validation(
270                "Resource URI cannot be empty".to_string(),
271            ));
272        }
273
274        let resource = resources
275            .get(&params.uri)
276            .ok_or_else(|| McpError::ResourceNotFound(params.uri.clone()))?;
277
278        resource.handler.unsubscribe(&params.uri).await?;
279
280        Ok(UnsubscribeResourceResult { meta: None })
281    }
282}
283
284/// Handler for prompt-related requests
285pub struct PromptHandler;
286
287impl PromptHandler {
288    /// Handle prompts/list request
289    pub async fn handle_list(
290        prompts: &HashMap<String, crate::core::prompt::Prompt>,
291        params: Option<Value>,
292    ) -> McpResult<ListPromptsResult> {
293        let _params: ListPromptsParams = match params {
294            Some(p) => serde_json::from_value(p)
295                .map_err(|e| McpError::Validation(format!("Invalid list prompts params: {e}")))?,
296            None => ListPromptsParams::default(),
297        };
298
299        // Pagination support will be added in future versions
300        let prompts: Vec<PromptInfo> = prompts
301            .values()
302            .map(|prompt| {
303                // Convert from core::prompt::PromptInfo to protocol::types::PromptInfo
304                PromptInfo {
305                    name: prompt.info.name.clone(),
306                    description: prompt.info.description.clone(),
307                    arguments: prompt.info.arguments.as_ref().map(|args| {
308                        args.iter()
309                            .map(|arg| PromptArgument {
310                                name: arg.name.clone(),
311                                description: arg.description.clone(),
312                                required: arg.required,
313                                title: None,
314                            })
315                            .collect()
316                    }),
317                    icons: None,
318                    title: None,
319                    meta: None,
320                }
321            })
322            .collect();
323
324        Ok(ListPromptsResult {
325            prompts,
326            next_cursor: None,
327            meta: None,
328        })
329    }
330
331    /// Handle prompts/get request
332    pub async fn handle_get(
333        prompts: &HashMap<String, crate::core::prompt::Prompt>,
334        params: Option<Value>,
335    ) -> McpResult<GetPromptResult> {
336        let params: GetPromptParams = match params {
337            Some(p) => serde_json::from_value(p)
338                .map_err(|e| McpError::Validation(format!("Invalid get prompt params: {e}")))?,
339            None => {
340                return Err(McpError::Validation(
341                    "Missing prompt get parameters".to_string(),
342                ));
343            }
344        };
345
346        if params.name.is_empty() {
347            return Err(McpError::Validation(
348                "Prompt name cannot be empty".to_string(),
349            ));
350        }
351
352        let prompt = prompts
353            .get(&params.name)
354            .ok_or_else(|| McpError::PromptNotFound(params.name.clone()))?;
355
356        let arguments = params
357            .arguments
358            .unwrap_or_default()
359            .into_iter()
360            .map(|(k, v)| (k, serde_json::Value::String(v)))
361            .collect();
362        let result = prompt.handler.get(arguments).await?;
363
364        Ok(GetPromptResult {
365            description: result.description,
366            messages: result
367                .messages
368                .into_iter()
369                .map(|msg| {
370                    // Convert from core::prompt::PromptMessage to protocol::types::PromptMessage
371                    PromptMessage {
372                        role: msg.role,
373                        content: match msg.content {
374                            ContentBlock::Text { text, .. } => ContentBlock::Text {
375                                text,
376                                annotations: None,
377                                meta: None,
378                            },
379                            ContentBlock::Image {
380                                data, mime_type, ..
381                            } => ContentBlock::Image {
382                                data,
383                                mime_type,
384                                annotations: None,
385                                meta: None,
386                            },
387                            other => other,
388                        },
389                    }
390                })
391                .collect(),
392            meta: None,
393        })
394    }
395}
396
397/// Handler for sampling requests
398pub struct SamplingHandler;
399
400impl SamplingHandler {
401    /// Handle sampling/createMessage request
402    pub async fn handle_create_message(_params: Option<Value>) -> McpResult<CreateMessageResult> {
403        // Note: Sampling is typically handled by the client side (LLM),
404        // but servers can provide sampling capabilities if they have access to LLMs
405        Err(McpError::Protocol(
406            "Sampling not implemented on server side".to_string(),
407        ))
408    }
409}
410
411/// Handler for logging requests
412pub struct LoggingHandler;
413
414impl LoggingHandler {
415    /// Handle logging/setLevel request
416    pub async fn handle_set_level(params: Option<Value>) -> McpResult<SetLoggingLevelResult> {
417        let _params: SetLoggingLevelParams = match params {
418            Some(p) => serde_json::from_value(p).map_err(|e| {
419                McpError::Validation(format!("Invalid set logging level params: {e}"))
420            })?,
421            None => {
422                return Err(McpError::Validation(
423                    "Missing logging level parameters".to_string(),
424                ));
425            }
426        };
427
428        // Logging level management feature planned for future implementation
429        // This would typically integrate with a logging framework like tracing
430
431        Ok(SetLoggingLevelResult { meta: None })
432    }
433}
434
435/// Handler for ping requests
436pub struct PingHandler;
437
438impl PingHandler {
439    /// Handle ping request
440    pub async fn handle(_params: Option<Value>) -> McpResult<PingResult> {
441        Ok(PingResult { meta: None })
442    }
443}
444
445/// Helper functions for common validation patterns
446pub mod validation {
447    use super::*;
448
449    /// Validate that required parameters are present
450    pub fn require_params<T>(params: Option<Value>, error_msg: &str) -> McpResult<T>
451    where
452        T: serde::de::DeserializeOwned,
453    {
454        match params {
455            Some(p) => serde_json::from_value(p)
456                .map_err(|e| McpError::Validation(format!("{error_msg}: {e}"))),
457            None => Err(McpError::Validation(error_msg.to_string())),
458        }
459    }
460
461    /// Validate that a string parameter is not empty
462    pub fn require_non_empty_string(value: &str, field_name: &str) -> McpResult<()> {
463        if value.is_empty() {
464            Err(McpError::Validation(format!(
465                "{field_name} cannot be empty"
466            )))
467        } else {
468            Ok(())
469        }
470    }
471
472    /// Validate URI format
473    pub fn validate_uri_format(uri: &str) -> McpResult<()> {
474        if uri.is_empty() {
475            return Err(McpError::Validation("URI cannot be empty".to_string()));
476        }
477
478        // Basic URI validation - check for scheme or absolute path
479        if !uri.contains("://") && !uri.starts_with('/') && !uri.starts_with("file:") {
480            return Err(McpError::Validation(
481                "URI must have a scheme or be an absolute path".to_string(),
482            ));
483        }
484
485        Ok(())
486    }
487}
488
489/// Notification builders for common server events
490pub mod notifications {
491    use super::*;
492
493    /// Create a tools list changed notification
494    pub fn tools_list_changed() -> McpResult<JsonRpcNotification> {
495        Ok(JsonRpcNotification::new(
496            methods::TOOLS_LIST_CHANGED.to_string(),
497            Some(ToolListChangedParams { meta: None }),
498        )?)
499    }
500
501    /// Create a resources list changed notification
502    pub fn resources_list_changed() -> McpResult<JsonRpcNotification> {
503        Ok(JsonRpcNotification::new(
504            methods::RESOURCES_LIST_CHANGED.to_string(),
505            Some(ResourceListChangedParams { meta: None }),
506        )?)
507    }
508
509    /// Create a prompts list changed notification
510    pub fn prompts_list_changed() -> McpResult<JsonRpcNotification> {
511        Ok(JsonRpcNotification::new(
512            methods::PROMPTS_LIST_CHANGED.to_string(),
513            Some(PromptListChangedParams { meta: None }),
514        )?)
515    }
516
517    /// Create a resource updated notification
518    pub fn resource_updated(uri: String) -> McpResult<JsonRpcNotification> {
519        Ok(JsonRpcNotification::new(
520            methods::RESOURCES_UPDATED.to_string(),
521            Some(ResourceUpdatedParams { uri }),
522        )?)
523    }
524
525    /// Create a progress notification
526    pub fn progress(
527        progress_token: String,
528        progress: f32,
529        total: Option<f32>,
530    ) -> McpResult<JsonRpcNotification> {
531        Ok(JsonRpcNotification::new(
532            methods::PROGRESS.to_string(),
533            Some(ProgressParams {
534                progress_token: serde_json::Value::String(progress_token),
535                progress,
536                total,
537                message: None,
538            }),
539        )?)
540    }
541
542    /// Create a logging message notification
543    pub fn log_message(
544        level: LoggingLevel,
545        logger: Option<String>,
546        data: Value,
547    ) -> McpResult<JsonRpcNotification> {
548        Ok(JsonRpcNotification::new(
549            methods::LOGGING_MESSAGE.to_string(),
550            Some(LoggingMessageParams {
551                level,
552                logger,
553                data,
554            }),
555        )?)
556    }
557}
558
559#[cfg(test)]
560mod tests {
561    use super::*;
562    use serde_json::json;
563
564    #[tokio::test]
565    async fn test_initialize_handler() {
566        let server_info = ServerInfo {
567            name: "test-server".to_string(),
568            version: "1.0.0".to_string(),
569            description: None,
570            title: Some("Test Server".to_string()),
571            website_url: None,
572            icons: None,
573        };
574        let capabilities = ServerCapabilities::default();
575
576        let params = json!({
577            "clientInfo": {
578                "name": "test-client",
579                "version": "1.0.0"
580            },
581            "capabilities": {},
582            "protocolVersion": LEGACY_PROTOCOL_VERSION
583        });
584
585        let result = InitializeHandler::handle(&server_info, &capabilities, Some(params)).await;
586        assert!(result.is_ok());
587
588        let init_result = result.unwrap();
589        assert_eq!(init_result.server_info.name, "test-server");
590        assert_eq!(init_result.protocol_version, LEGACY_PROTOCOL_VERSION);
591    }
592
593    #[tokio::test]
594    async fn test_ping_handler() {
595        let result = PingHandler::handle(None).await;
596        assert!(result.is_ok());
597    }
598
599    #[test]
600    fn test_validation_helpers() {
601        // Test require_non_empty_string
602        assert!(validation::require_non_empty_string("test", "field").is_ok());
603        assert!(validation::require_non_empty_string("", "field").is_err());
604
605        // Test validate_uri_format
606        assert!(validation::validate_uri_format("https://example.com").is_ok());
607        assert!(validation::validate_uri_format("file:///path").is_ok());
608        assert!(validation::validate_uri_format("/absolute/path").is_ok());
609        assert!(validation::validate_uri_format("").is_err());
610        assert!(validation::validate_uri_format("invalid").is_err());
611    }
612
613    #[test]
614    fn test_notification_builders() {
615        assert!(notifications::tools_list_changed().is_ok());
616        assert!(notifications::resources_list_changed().is_ok());
617        assert!(notifications::prompts_list_changed().is_ok());
618        assert!(notifications::resource_updated("file:///test".to_string()).is_ok());
619        assert!(notifications::progress("token".to_string(), 0.5, Some(100.0)).is_ok());
620        assert!(notifications::log_message(
621            LoggingLevel::Info,
622            Some("test".to_string()),
623            json!({"message": "test log"})
624        )
625        .is_ok());
626    }
627}