1use std::collections::HashMap;
6use std::sync::Mutex;
7
8use serde_json::Value;
9
10use crate::dispatch::{BinaryTrieRouter, ToolResult};
11use crate::security::{SecurityAuditAction, SecurityAuditEvent, SecurityAuditLog};
12use crate::{CapabilityManifest, DCPError, SecurityError};
13
14use super::json_rpc::{
15 JsonRpcError, JsonRpcParseError, JsonRpcParser, JsonRpcRequest, JsonRpcResponse, RequestId,
16 DEFAULT_MAX_JSONRPC_REQUEST_SIZE,
17};
18use super::request_replay::{replay_key, RequestReplayGuard};
19
20#[derive(Debug, Clone, thiserror::Error)]
22pub enum AdapterError {
23 #[error("JSON-RPC parse error: {0}")]
24 ParseError(#[from] JsonRpcParseError),
25 #[error("unknown tool")]
26 UnknownTool(String),
27 #[error("DCP error: {0}")]
28 DcpError(#[from] DCPError),
29 #[error("serialization error: {0}")]
30 SerializationError(String),
31 #[error("invalid request: {0}")]
32 InvalidRequest(String),
33 #[error("invalid params: {0}")]
34 InvalidParams(String),
35 #[error("{kind} capacity exceeded")]
36 CapacityExceeded { kind: &'static str, max: usize },
37}
38
39pub struct McpAdapter {
41 tool_cache: HashMap<String, u16>,
43 id_to_name: HashMap<u16, String>,
45 max_request_size: usize,
47 negotiated_capabilities: Option<CapabilityManifest>,
49 security_audit: SecurityAuditLog,
51 tool_call_replay_guard: Mutex<RequestReplayGuard>,
53}
54
55impl McpAdapter {
56 pub fn new() -> Self {
58 Self {
59 tool_cache: HashMap::new(),
60 id_to_name: HashMap::new(),
61 max_request_size: DEFAULT_MAX_JSONRPC_REQUEST_SIZE,
62 negotiated_capabilities: None,
63 security_audit: SecurityAuditLog::new(),
64 tool_call_replay_guard: Mutex::new(RequestReplayGuard::default()),
65 }
66 }
67
68 pub fn with_max_request_size(mut self, max_request_size: usize) -> Self {
70 self.max_request_size = max_request_size;
71 self
72 }
73
74 pub fn with_negotiated_capabilities(mut self, capabilities: CapabilityManifest) -> Self {
76 self.negotiated_capabilities = Some(capabilities);
77 self
78 }
79
80 fn request_id_for_audit(id: &RequestId) -> Option<String> {
81 match id {
82 RequestId::String(value) => Some(value.clone()),
83 RequestId::Number(value) => Some(value.to_string()),
84 RequestId::Null => Some("null".to_string()),
85 RequestId::Missing => None,
86 }
87 }
88
89 fn audit_capability_denial(&self, request: &JsonRpcRequest) {
90 let mut event =
91 SecurityAuditEvent::new(SecurityAuditAction::CapabilityDenied, "capability_denied")
92 .with_method(request.method.clone())
93 .with_field("adapter", "legacy_mcp");
94 if let Some(request_id) = Self::request_id_for_audit(&request.id) {
95 event = event.with_request_id(request_id);
96 }
97 self.security_audit.record(event);
98 }
99
100 fn audit_request_rejection(
101 &self,
102 action: SecurityAuditAction,
103 reason: &'static str,
104 request: &JsonRpcRequest,
105 ) {
106 let mut event = SecurityAuditEvent::new(action, reason)
107 .with_method(request.method.clone())
108 .with_field("adapter", "legacy_mcp");
109 if let Some(request_id) = Self::request_id_for_audit(&request.id) {
110 event = event.with_request_id(request_id);
111 }
112 self.security_audit.record(event);
113 }
114
115 fn capability_denied_response(&self, request: &JsonRpcRequest) -> Result<String, AdapterError> {
116 self.audit_capability_denial(request);
117 self.format_error_response(
118 request.id.clone(),
119 JsonRpcError::new(-32001, "Capability denied"),
120 )
121 }
122
123 fn replay_rejected_response(&self, request: &JsonRpcRequest) -> Result<String, AdapterError> {
124 self.audit_request_rejection(
125 SecurityAuditAction::ReplayRejected,
126 "request_replay",
127 request,
128 );
129 self.format_error_response(
130 request.id.clone(),
131 JsonRpcError::new(-32002, "Request replay rejected"),
132 )
133 }
134
135 fn record_tool_call_request_id(&self, request: &JsonRpcRequest) -> Result<bool, AdapterError> {
136 let Some(request_id) = Self::request_id_for_audit(&request.id) else {
137 return Ok(true);
138 };
139 let mut guard = self
140 .tool_call_replay_guard
141 .lock()
142 .map_err(|_| AdapterError::InvalidRequest("request replay guard unavailable".into()))?;
143
144 Ok(guard.check_and_record(replay_key("tools/call", &request_id)))
145 }
146
147 fn security_error_response(
148 &self,
149 request: &JsonRpcRequest,
150 error: SecurityError,
151 ) -> Result<String, AdapterError> {
152 match error {
153 SecurityError::ValidationFailed => self.validation_error_response(
154 Some(request),
155 &AdapterError::InvalidParams("schema validation failed".into()),
156 ),
157 _ => self.capability_denied_response(request),
158 }
159 }
160
161 fn audit_tool_registration_failure(&self, reason: &'static str, tool_id: u16, name: &str) {
162 self.security_audit.record(
163 SecurityAuditEvent::new(SecurityAuditAction::ValidationRejected, reason)
164 .with_field("adapter", "legacy_mcp")
165 .with_field("operation", "tool_registration")
166 .with_field("tool_id", tool_id.to_string())
167 .with_field("tool_name", name),
168 );
169 }
170
171 fn validation_error_details(error: &AdapterError) -> (&'static str, JsonRpcError) {
172 match error {
173 AdapterError::ParseError(JsonRpcParseError::InvalidJson(_)) => {
174 ("parse_error", JsonRpcError::parse_error())
175 }
176 AdapterError::ParseError(JsonRpcParseError::RequestTooLarge) => {
177 ("request_too_large", JsonRpcError::invalid_request())
178 }
179 AdapterError::ParseError(JsonRpcParseError::RequestIdTooLarge) => {
180 ("request_id_too_large", JsonRpcError::invalid_request())
181 }
182 AdapterError::ParseError(JsonRpcParseError::RequestIdSensitive) => {
183 ("request_id_sensitive", JsonRpcError::invalid_request())
184 }
185 AdapterError::ParseError(JsonRpcParseError::BatchUnsupported) => {
186 ("batch_unsupported", JsonRpcError::invalid_request())
187 }
188 AdapterError::InvalidParams(_) => ("invalid_params", JsonRpcError::invalid_params()),
189 AdapterError::ParseError(_) | AdapterError::InvalidRequest(_) => {
190 ("invalid_request", JsonRpcError::invalid_request())
191 }
192 _ => ("request_failed", JsonRpcError::internal_error()),
193 }
194 }
195
196 fn validation_error_response(
197 &self,
198 request: Option<&JsonRpcRequest>,
199 error: &AdapterError,
200 ) -> Result<String, AdapterError> {
201 let (reason, json_error) = Self::validation_error_details(error);
202 let mut event = SecurityAuditEvent::new(SecurityAuditAction::ValidationRejected, reason)
203 .with_field("adapter", "legacy_mcp");
204
205 let response_id = match request {
206 Some(request) => {
207 event = event.with_method(request.method.clone());
208 if let Some(request_id) = Self::request_id_for_audit(&request.id) {
209 event = event.with_request_id(request_id);
210 }
211 match &request.id {
212 RequestId::Missing => RequestId::Null,
213 _ => request.id.clone(),
214 }
215 }
216 None => RequestId::Null,
217 };
218
219 self.security_audit.record(event);
220 self.format_error_response(response_id, json_error)
221 }
222
223 fn require_request_method_response(
224 &self,
225 request: &JsonRpcRequest,
226 expected_method: &'static str,
227 ) -> Option<Result<String, AdapterError>> {
228 if request.method == expected_method && !request.is_notification() {
229 return None;
230 }
231
232 Some(self.validation_error_response(
233 Some(request),
234 &AdapterError::InvalidRequest(format!(
235 "{expected_method} requires a JSON-RPC request id"
236 )),
237 ))
238 }
239
240 pub fn register_tool(
242 &mut self,
243 name: impl Into<String>,
244 tool_id: u16,
245 ) -> Result<u16, AdapterError> {
246 let name = name.into();
247 if (tool_id as usize) >= CapabilityManifest::MAX_TOOLS {
248 self.audit_tool_registration_failure(
249 "tool_registration_capacity_exceeded",
250 tool_id,
251 &name,
252 );
253 return Err(AdapterError::CapacityExceeded {
254 kind: "tool",
255 max: CapabilityManifest::MAX_TOOLS,
256 });
257 }
258 if self.tool_cache.contains_key(&name) {
259 self.audit_tool_registration_failure(
260 "tool_registration_duplicate_name",
261 tool_id,
262 &name,
263 );
264 return Err(AdapterError::InvalidRequest("duplicate tool name".into()));
265 }
266 if self.id_to_name.contains_key(&tool_id) {
267 self.audit_tool_registration_failure("tool_registration_duplicate_id", tool_id, &name);
268 return Err(AdapterError::InvalidRequest("duplicate tool id".into()));
269 }
270
271 self.tool_cache.insert(name.clone(), tool_id);
272 self.id_to_name.insert(tool_id, name);
273 Ok(tool_id)
274 }
275
276 pub fn resolve_tool_name(&self, name: &str) -> Option<u16> {
278 self.tool_cache.get(name).copied()
279 }
280
281 pub fn resolve_tool_id(&self, tool_id: u16) -> Option<&str> {
283 self.id_to_name.get(&tool_id).map(|s| s.as_str())
284 }
285
286 pub fn security_audit(&self) -> SecurityAuditLog {
288 self.security_audit.clone()
289 }
290
291 pub fn parse_request(&self, json: &str) -> Result<JsonRpcRequest, AdapterError> {
293 Ok(JsonRpcParser::parse_request_with_limit(
294 json,
295 self.max_request_size,
296 )?)
297 }
298
299 pub fn translate_params(&self, params: &Option<Value>) -> Vec<u8> {
301 match params {
302 Some(value) => {
303 serde_json::to_vec(value).unwrap_or_default()
306 }
307 None => Vec::new(),
308 }
309 }
310
311 fn translate_legacy_tool_arguments(
312 &self,
313 arguments: Option<&Value>,
314 ) -> Result<Vec<u8>, AdapterError> {
315 match arguments {
316 None => Ok(Vec::new()),
317 Some(Value::Object(arguments)) if arguments.is_empty() => Ok(Vec::new()),
318 Some(Value::Object(_)) => Err(AdapterError::InvalidParams(
319 "legacy tools/call arguments must be empty or omitted".into(),
320 )),
321 Some(_) => Err(AdapterError::InvalidParams(
322 "tools/call arguments must be an object".into(),
323 )),
324 }
325 }
326
327 fn tools_call_params_object<'a>(
328 &self,
329 params: &'a Value,
330 ) -> Result<&'a serde_json::Map<String, Value>, AdapterError> {
331 let params = params.as_object().ok_or_else(|| {
332 AdapterError::InvalidParams("tools/call params must be an object".into())
333 })?;
334
335 if params
336 .keys()
337 .any(|key| key != "name" && key != "arguments" && key != "_meta")
338 {
339 return Err(AdapterError::InvalidParams(
340 "tools/call params contain unsupported fields".into(),
341 ));
342 }
343 if params.get("_meta").is_some_and(|meta| !meta.is_object()) {
344 return Err(AdapterError::InvalidParams(
345 "tools/call _meta must be an object".into(),
346 ));
347 }
348
349 Ok(params)
350 }
351
352 pub fn translate_result(&self, result: &ToolResult) -> Value {
354 match result {
355 ToolResult::Success(data) => {
356 serde_json::from_slice(data)
358 .unwrap_or_else(|_| Value::String(String::from_utf8_lossy(data).to_string()))
359 }
360 ToolResult::Empty => Value::Null,
361 ToolResult::Error(err) => {
362 serde_json::json!({
363 "error": {
364 "code": *err as i32,
365 "message": err.to_string()
366 }
367 })
368 }
369 }
370 }
371
372 pub fn format_success_response(
374 &self,
375 id: RequestId,
376 result: Value,
377 ) -> Result<String, AdapterError> {
378 let response = JsonRpcResponse::success(id, result);
379 JsonRpcParser::format_response(&response)
380 .map_err(|e| AdapterError::SerializationError(e.to_string()))
381 }
382
383 pub fn format_error_response(
385 &self,
386 id: RequestId,
387 error: JsonRpcError,
388 ) -> Result<String, AdapterError> {
389 let response = JsonRpcResponse::error(id, error);
390 JsonRpcParser::format_response(&response)
391 .map_err(|e| AdapterError::SerializationError(e.to_string()))
392 }
393
394 pub fn handle_initialize(&self, request: &JsonRpcRequest) -> Result<String, AdapterError> {
396 if let Some(response) = self.require_request_method_response(request, "initialize") {
397 return response;
398 }
399
400 let mut capabilities = serde_json::Map::new();
401 if !self.tool_cache.is_empty() {
402 capabilities.insert(
403 "tools".to_string(),
404 serde_json::json!({
405 "listChanged": false
406 }),
407 );
408 }
409
410 let result = serde_json::json!({
411 "protocolVersion": "2024-11-05",
412 "capabilities": capabilities,
413 "serverInfo": {
414 "name": "dcp-server",
415 "version": "0.1.0"
416 }
417 });
418
419 self.format_success_response(request.id.clone(), result)
420 }
421
422 pub fn handle_tools_list(&self, request: &JsonRpcRequest) -> Result<String, AdapterError> {
424 if let Some(response) = self.require_request_method_response(request, "tools/list") {
425 return response;
426 }
427
428 let capabilities = match self.negotiated_capabilities.as_ref() {
429 Some(capabilities) if capabilities.tool_count() > 0 => capabilities,
430 _ => return self.capability_denied_response(request),
431 };
432
433 let tools: Vec<Value> = self
434 .tool_cache
435 .iter()
436 .filter(|(_, tool_id)| capabilities.has_tool(**tool_id))
437 .map(|(name, _)| {
438 serde_json::json!({
439 "name": name,
440 "description": format!("Tool: {}", name),
441 "inputSchema": {
442 "type": "object",
443 "properties": {}
444 }
445 })
446 })
447 .collect();
448
449 let result = serde_json::json!({
450 "tools": tools
451 });
452
453 self.format_success_response(request.id.clone(), result)
454 }
455
456 pub fn handle_tools_call(
458 &self,
459 request: &JsonRpcRequest,
460 router: &BinaryTrieRouter,
461 ) -> Result<String, AdapterError> {
462 if let Some(response) = self.require_request_method_response(request, "tools/call") {
463 return response;
464 }
465
466 let params_value = match request.params.as_ref() {
468 Some(params) => params,
469 None => {
470 return self.validation_error_response(
471 Some(request),
472 &AdapterError::ParseError(JsonRpcParseError::InvalidStructure),
473 );
474 }
475 };
476 let params = match self.tools_call_params_object(params_value) {
477 Ok(params) => params,
478 Err(err) => return self.validation_error_response(Some(request), &err),
479 };
480
481 let tool_name = match params.get("name").and_then(|v| v.as_str()) {
482 Some(name) => name,
483 None => {
484 return self.validation_error_response(
485 Some(request),
486 &AdapterError::ParseError(JsonRpcParseError::InvalidStructure),
487 );
488 }
489 };
490
491 let arguments = params.get("arguments");
492
493 let tool_id = match self.resolve_tool_name(tool_name) {
495 Some(tool_id) => tool_id,
496 None => return self.capability_denied_response(request),
497 };
498 let capabilities = match self.negotiated_capabilities.as_ref() {
499 Some(capabilities) => capabilities,
500 None => return self.capability_denied_response(request),
501 };
502
503 let args_bytes = match self.translate_legacy_tool_arguments(arguments) {
505 Ok(args_bytes) => args_bytes,
506 Err(err) => return self.validation_error_response(Some(request), &err),
507 };
508 if !self.record_tool_call_request_id(request)? {
509 return self.replay_rejected_response(request);
510 }
511
512 let shared_args = crate::dispatch::SharedArgs::new(&args_bytes, 0);
513
514 let result = match router.execute_authorized(capabilities, tool_id, &shared_args) {
515 Ok(result) => result,
516 Err(err) => return self.security_error_response(request, err),
517 };
518
519 let result_value = self.translate_result(&result);
521
522 let response_result = serde_json::json!({
524 "content": [{
525 "type": "text",
526 "text": serde_json::to_string(&result_value).unwrap_or_default()
527 }]
528 });
529
530 self.format_success_response(request.id.clone(), response_result)
531 }
532
533 pub fn handle_request(
535 &self,
536 json: &str,
537 router: &BinaryTrieRouter,
538 ) -> Result<String, AdapterError> {
539 let request = match self.parse_request(json) {
540 Ok(request) => request,
541 Err(error) => return self.validation_error_response(None, &error),
542 };
543
544 if request.is_notification() {
545 self.audit_request_rejection(
546 SecurityAuditAction::RequestRejected,
547 "notification_not_allowed",
548 &request,
549 );
550 return self.format_error_response(RequestId::Null, JsonRpcError::invalid_request());
551 }
552
553 let response = match request.method.as_str() {
554 "initialize" => self.handle_initialize(&request),
555 "tools/list" => self.handle_tools_list(&request),
556 "tools/call" => self.handle_tools_call(&request, router),
557 _ => {
558 self.audit_request_rejection(
560 SecurityAuditAction::RequestRejected,
561 "method_not_found",
562 &request,
563 );
564 self.format_error_response(request.id.clone(), JsonRpcError::method_not_found())
565 }
566 };
567
568 match response {
569 Err(error @ AdapterError::ParseError(_))
570 | Err(error @ AdapterError::InvalidRequest(_)) => {
571 self.validation_error_response(Some(&request), &error)
572 }
573 other => other,
574 }
575 }
576
577 pub fn tool_count(&self) -> usize {
579 self.tool_cache.len()
580 }
581}
582
583impl Default for McpAdapter {
584 fn default() -> Self {
585 Self::new()
586 }
587}
588
589#[cfg(test)]
590mod tests {
591 use super::*;
592
593 #[test]
594 fn test_register_tool() {
595 let mut adapter = McpAdapter::new();
596 adapter.register_tool("read_file", 1).unwrap();
597 adapter.register_tool("write_file", 2).unwrap();
598
599 assert_eq!(adapter.resolve_tool_name("read_file"), Some(1));
600 assert_eq!(adapter.resolve_tool_name("write_file"), Some(2));
601 assert_eq!(adapter.resolve_tool_name("unknown"), None);
602
603 assert_eq!(adapter.resolve_tool_id(1), Some("read_file"));
604 assert_eq!(adapter.resolve_tool_id(2), Some("write_file"));
605 assert_eq!(adapter.resolve_tool_id(99), None);
606 }
607
608 #[test]
609 fn test_parse_request() {
610 let adapter = McpAdapter::new();
611 let json = r#"{"jsonrpc":"2.0","method":"initialize","id":1}"#;
612
613 let request = adapter.parse_request(json).unwrap();
614 assert_eq!(request.method, "initialize");
615 }
616
617 #[test]
618 fn test_translate_params() {
619 let adapter = McpAdapter::new();
620
621 let params = Some(serde_json::json!({"path": "/tmp/test.txt"}));
622 let bytes = adapter.translate_params(¶ms);
623
624 assert!(!bytes.is_empty());
625 let parsed: Value = serde_json::from_slice(&bytes).unwrap();
626 assert_eq!(parsed["path"], "/tmp/test.txt");
627 }
628
629 #[test]
630 fn test_translate_result_success() {
631 let adapter = McpAdapter::new();
632
633 let result = ToolResult::Success(b"hello world".to_vec());
634 let value = adapter.translate_result(&result);
635
636 assert_eq!(value, Value::String("hello world".to_string()));
637 }
638
639 #[test]
640 fn test_translate_result_json() {
641 let adapter = McpAdapter::new();
642
643 let json_bytes = serde_json::to_vec(&serde_json::json!({"key": "value"})).unwrap();
644 let result = ToolResult::Success(json_bytes);
645 let value = adapter.translate_result(&result);
646
647 assert_eq!(value["key"], "value");
648 }
649
650 #[test]
651 fn test_translate_result_error() {
652 let adapter = McpAdapter::new();
653
654 let result = ToolResult::Error(DCPError::ToolNotFound);
655 let value = adapter.translate_result(&result);
656
657 assert_eq!(value["error"]["code"], DCPError::ToolNotFound as i32);
658 assert!(value["error"]["message"]
659 .as_str()
660 .unwrap()
661 .contains("not found"));
662 }
663
664 #[test]
665 fn test_handle_initialize() {
666 let adapter = McpAdapter::new();
667 let request = JsonRpcRequest::new("initialize", None, RequestId::Number(1));
668
669 let response_json = adapter.handle_initialize(&request).unwrap();
670 let response = JsonRpcParser::parse_response(&response_json).unwrap();
671
672 assert!(response.is_success());
673 let result = response.result.unwrap();
674 assert!(result["capabilities"].is_object());
675 assert!(result["capabilities"]["tools"].is_null());
676 }
677
678 #[test]
679 fn test_handle_tools_list() {
680 let mut adapter = McpAdapter::new();
681 adapter.register_tool("read_file", 1).unwrap();
682 adapter.register_tool("write_file", 2).unwrap();
683 let mut capabilities = CapabilityManifest::new(1);
684 capabilities.set_tool(1);
685 capabilities.set_tool(2);
686 let adapter = adapter.with_negotiated_capabilities(capabilities);
687
688 let request = JsonRpcRequest::new("tools/list", None, RequestId::Number(1));
689 let response_json = adapter.handle_tools_list(&request).unwrap();
690 let response = JsonRpcParser::parse_response(&response_json).unwrap();
691
692 assert!(response.is_success());
693 let result = response.result.unwrap();
694 let tools = result["tools"].as_array().unwrap();
695 assert_eq!(tools.len(), 2);
696 }
697
698 #[test]
699 fn test_format_error_response() {
700 let adapter = McpAdapter::new();
701
702 let response = adapter
703 .format_error_response(RequestId::Number(1), JsonRpcError::method_not_found())
704 .unwrap();
705
706 let parsed = JsonRpcParser::parse_response(&response).unwrap();
707 assert!(parsed.is_error());
708 assert_eq!(parsed.error.unwrap().code, -32601);
709 }
710}