1use serde_json::Value;
7use std::collections::HashMap;
8
9use crate::core::error::{McpError, McpResult};
10use crate::protocol::{messages::*, methods, types::*, LEGACY_PROTOCOL_VERSION};
11
12pub struct InitializeHandler;
14
15impl InitializeHandler {
16 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 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 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
62pub struct ToolHandler;
64
65impl ToolHandler {
66 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 let tools: Vec<ToolInfo> = tools
79 .values()
80 .filter(|tool| tool.enabled)
81 .map(|tool| {
82 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 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(¶ms.name)
126 .ok_or_else(|| McpError::ToolNotFound(params.name.clone()))?;
127
128 if !tool.enabled {
129 let name = ¶ms.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
145pub struct ResourceHandler;
147
148impl ResourceHandler {
149 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 let resources: Vec<ResourceInfo> = resources
162 .values()
163 .map(|resource| {
164 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 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(¶ms.uri)
209 .ok_or_else(|| McpError::ResourceNotFound(params.uri.clone()))?;
210
211 let query_params = HashMap::new();
213 let contents = resource.handler.read(¶ms.uri, &query_params).await?;
214
215 Ok(ReadResourceResult {
216 contents,
217 meta: None,
218 })
219 }
220
221 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(¶ms.uri)
245 .ok_or_else(|| McpError::ResourceNotFound(params.uri.clone()))?;
246
247 resource.handler.subscribe(¶ms.uri).await?;
248
249 Ok(SubscribeResourceResult { meta: None })
250 }
251
252 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(¶ms.uri)
276 .ok_or_else(|| McpError::ResourceNotFound(params.uri.clone()))?;
277
278 resource.handler.unsubscribe(¶ms.uri).await?;
279
280 Ok(UnsubscribeResourceResult { meta: None })
281 }
282}
283
284pub struct PromptHandler;
286
287impl PromptHandler {
288 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 let prompts: Vec<PromptInfo> = prompts
301 .values()
302 .map(|prompt| {
303 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 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(¶ms.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 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
397pub struct SamplingHandler;
399
400impl SamplingHandler {
401 pub async fn handle_create_message(_params: Option<Value>) -> McpResult<CreateMessageResult> {
403 Err(McpError::Protocol(
406 "Sampling not implemented on server side".to_string(),
407 ))
408 }
409}
410
411pub struct LoggingHandler;
413
414impl LoggingHandler {
415 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 Ok(SetLoggingLevelResult { meta: None })
432 }
433}
434
435pub struct PingHandler;
437
438impl PingHandler {
439 pub async fn handle(_params: Option<Value>) -> McpResult<PingResult> {
441 Ok(PingResult { meta: None })
442 }
443}
444
445pub mod validation {
447 use super::*;
448
449 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 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 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 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
489pub mod notifications {
491 use super::*;
492
493 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 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 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 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 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 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 assert!(validation::require_non_empty_string("test", "field").is_ok());
603 assert!(validation::require_non_empty_string("", "field").is_err());
604
605 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}