1use anyhow::Result;
17use async_trait::async_trait;
18use serde_json::Value;
19use std::sync::Arc;
20use std::time::Duration;
21use tokio::sync::Mutex;
22use tracing::{debug, error, info, warn};
23
24use crate::auth::{AuthContext, AuthManager};
25use crate::component_selector::ComponentManager;
26use crate::config::Config;
27use crate::protocol::{
28 error_codes, error_response, success_response, ClientInfo, InitializeParams, InitializeResult,
29 JsonRpcError, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, LoggingCapability, Prompt,
30 PromptArgument, PromptsCapability, PromptsListResult, Resource, ResourceContent,
31 ResourceContentType, ResourceReadParams, ResourcesCapability, ResourcesListResult,
32 ServerCapabilities, ServerInfo, Tool, ToolCallParams, ToolsCapability, ToolsListResult,
33 PROTOCOL_VERSION,
34};
35use crate::scanner::{SecurityScanner, Threat};
36use crate::shield::Shield;
37use crate::signing::{MessageSignature, SignedMessage, SigningManager};
38use crate::telemetry::{MetricValue, TelemetryMetric};
39use crate::traits::{RateLimitKey, RateLimiter, SecurityEvent, SecurityEventProcessor};
40use crate::transport::{
41 DefaultTransportFactory, MessageHandler, TransportConnection, TransportFactory,
42 TransportManager, TransportMessage,
43};
44use crate::versioning::{add_version_metadata, ApiRegistry};
45
46struct SessionInfo {
48 #[allow(dead_code)]
49 id: String,
50 client_info: Option<ClientInfo>,
51 #[allow(dead_code)]
52 created_at: std::time::Instant,
53 #[allow(dead_code)]
54 threats_blocked: u64,
55 #[allow(dead_code)]
56 last_activity: std::time::Instant,
57}
58
59struct SessionStore {
61 sessions: std::collections::HashMap<String, SessionInfo>,
62}
63
64pub struct McpServer {
66 scanner: Arc<SecurityScanner>,
67 pub shield: Arc<Shield>,
68 config: Arc<Config>,
69 auth_manager: Arc<AuthManager>,
70 signing_manager: Arc<SigningManager>,
71 rate_limiter: Arc<dyn RateLimiter>,
72 event_processor: Arc<dyn SecurityEventProcessor>,
73 session_store: Arc<Mutex<SessionStore>>,
74 server_info: ServerInfo,
75 capabilities: ServerCapabilities,
76 pub component_manager: Arc<ComponentManager>,
77}
78
79#[derive(Debug, thiserror::Error)]
81pub enum ServerError {
82 #[error("Invalid request: {0}")]
83 InvalidRequest(String),
84
85 #[error("Method not found: {0}")]
86 MethodNotFound(String),
87
88 #[error("Invalid parameters: {0}")]
89 InvalidParams(String),
90
91 #[error("Internal error: {0}")]
92 InternalError(String),
93
94 #[error("Threat detected: {threats:?}")]
95 ThreatDetected { threats: Vec<Threat> },
96
97 #[error("Unauthorized")]
98 Unauthorized,
99
100 #[error("Rate limited")]
101 RateLimited,
102
103 #[error("Timeout")]
104 Timeout,
105}
106
107impl ServerError {
108 fn to_json_rpc_error(&self) -> JsonRpcError {
110 match self {
111 Self::InvalidRequest(msg) => JsonRpcError {
112 code: error_codes::INVALID_REQUEST,
113 message: msg.clone(),
114 data: None,
115 },
116 Self::MethodNotFound(method) => JsonRpcError {
117 code: error_codes::METHOD_NOT_FOUND,
118 message: format!("Method not found: {method}"),
119 data: None,
120 },
121 Self::InvalidParams(msg) => JsonRpcError {
122 code: error_codes::INVALID_PARAMS,
123 message: msg.clone(),
124 data: None,
125 },
126 Self::InternalError(msg) => JsonRpcError {
127 code: error_codes::INTERNAL_ERROR,
128 message: msg.clone(),
129 data: None,
130 },
131 Self::ThreatDetected { threats } => JsonRpcError {
132 code: error_codes::THREAT_DETECTED,
133 message: "Security threat detected".to_string(),
134 data: Some(serde_json::to_value(threats).unwrap_or(Value::Null)),
135 },
136 Self::Unauthorized => JsonRpcError {
137 code: error_codes::UNAUTHORIZED,
138 message: "Unauthorized".to_string(),
139 data: None,
140 },
141 Self::RateLimited => JsonRpcError {
142 code: error_codes::RATE_LIMITED,
143 message: "Rate limit exceeded".to_string(),
144 data: None,
145 },
146 Self::Timeout => JsonRpcError {
147 code: error_codes::INTERNAL_ERROR,
148 message: "Request timeout".to_string(),
149 data: None,
150 },
151 }
152 }
153}
154
155impl McpServer {
156 async fn track_security_event(
158 &self,
159 event_type: &str,
160 client_id: &str,
161 metadata: serde_json::Value,
162 ) {
163 if self.config.audit.enabled {
165 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
166
167 let audit_event_type = match event_type {
168 "auth.success" => AuditEventType::AuthSuccess {
169 user_id: client_id.to_string(),
170 },
171 "auth.failure" => AuditEventType::AuthFailure {
172 user_id: Some(client_id.to_string()),
173 reason: metadata
174 .get("reason")
175 .and_then(|v| v.as_str())
176 .unwrap_or("Unknown")
177 .to_string(),
178 },
179 "threat.detected" => AuditEventType::ThreatDetected {
180 client_id: client_id.to_string(),
181 threat_count: 1,
182 },
183 "rate_limit.exceeded" => AuditEventType::RateLimitTriggered {
184 client_id: client_id.to_string(),
185 limit_type: "request".to_string(),
186 },
187 _ => AuditEventType::Custom {
188 event_type: event_type.to_string(),
189 data: metadata.clone(),
190 },
191 };
192
193 let severity = match event_type {
194 "auth.failure" | "threat.detected" => AuditSeverity::Warning,
195 "rate_limit.exceeded" => AuditSeverity::Warning,
196 "auth.success" => AuditSeverity::Info,
197 _ => AuditSeverity::Info,
198 };
199
200 let audit_event =
201 AuditEvent::new(audit_event_type, severity).with_client_id(client_id.to_string());
202
203 let audit_logger = self.component_manager.audit_logger();
204 let retry_strategy = self.component_manager.retry_strategy();
205
206 let audit_event_json = serde_json::to_value(&audit_event).unwrap_or(Value::Null);
208 match retry_strategy
209 .execute_json("audit.log", audit_event_json.clone())
210 .await
211 {
212 Ok(_) => {
213 if let Err(e) = audit_logger.log(audit_event).await {
215 warn!("Failed to log audit event: {}", e);
216 }
217 },
218 Err(e) => {
219 error!("Failed to log audit event after retries: {}", e);
220 },
221 }
222 }
223
224 if self.config.is_event_processor_enabled() {
226 let event = SecurityEvent {
227 event_type: event_type.to_string(),
228 client_id: client_id.to_string(),
229 timestamp: std::time::SystemTime::now()
230 .duration_since(std::time::UNIX_EPOCH)
231 .unwrap_or_default()
232 .as_secs(),
233 metadata,
234 };
235 let _ = self.event_processor.process_event(event).await;
236 }
237 }
238 pub fn new(config: Config) -> Result<Self> {
240 let component_manager = Arc::new(ComponentManager::new(&config)?);
242
243 let mut scanner = SecurityScanner::new(config.scanner.clone())?;
244
245 if config.plugins.enabled {
247 scanner.set_plugin_manager(component_manager.plugin_manager().clone());
248 }
249
250 let scanner = Arc::new(scanner);
251 let shield = Arc::new(Shield::with_config(config.shield.clone()));
252
253 shield.set_event_processor_enabled(config.is_event_processor_enabled());
255
256 let server_resource_id = format!("kindlyguard:{}", env!("CARGO_PKG_VERSION"));
258 let auth_manager = Arc::new(AuthManager::new(config.auth.clone(), server_resource_id));
259
260 let signing_manager = Arc::new(SigningManager::new(config.signing.clone())?);
262
263 let rate_limiter = component_manager.rate_limiter().clone();
265 let event_processor = component_manager.event_processor().clone();
266
267 if config.is_event_processor_enabled() {
269 shield.set_event_processor(&event_processor);
270 }
271
272 let server_info = ServerInfo {
273 name: "kindly-guard".to_string(),
274 version: env!("CARGO_PKG_VERSION").to_string(),
275 };
276
277 let capabilities = ServerCapabilities {
278 tools: Some(ToolsCapability {}),
279 resources: Some(ResourcesCapability {}),
280 prompts: Some(PromptsCapability {}),
281 logging: Some(LoggingCapability {}),
282 };
283
284 Ok(Self {
285 scanner,
286 shield,
287 config: Arc::new(config),
288 auth_manager,
289 signing_manager,
290 rate_limiter,
291 event_processor,
292 session_store: Arc::new(Mutex::new(SessionStore {
293 sessions: std::collections::HashMap::new(),
294 })),
295 server_info,
296 capabilities,
297 component_manager,
298 })
299 }
300
301 pub fn scanner(&self) -> &Arc<SecurityScanner> {
304 &self.scanner
305 }
306
307 pub async fn handle_message(&self, message: &str) -> Option<String> {
308 match serde_json::from_str::<Value>(message) {
310 Ok(Value::Array(requests)) => {
311 let mut responses = Vec::new();
313 for req in requests {
314 if let Some(response) = self.handle_value(req).await {
315 if let Ok(resp_value) = serde_json::from_str::<Value>(&response) {
316 responses.push(resp_value);
317 }
318 }
319 }
320
321 if responses.is_empty() {
322 None
323 } else {
324 match serde_json::to_string(&responses) {
325 Ok(json) => Some(json),
326 Err(e) => {
327 error!("Failed to serialize batch response: {}", e);
328 None
329 },
330 }
331 }
332 },
333 Ok(value) => self.handle_value(value).await,
334 Err(_) => {
335 let error_response = error_response(
337 Value::Null,
338 error_codes::PARSE_ERROR,
339 "Parse error".to_string(),
340 None,
341 );
342 match serde_json::to_string(&error_response) {
343 Ok(json) => Some(json),
344 Err(e) => {
345 error!("Failed to serialize error response: {}", e);
346 None
347 },
348 }
349 },
350 }
351 }
352
353 async fn handle_value(&self, mut value: Value) -> Option<String> {
355 if let Ok(signed) = serde_json::from_value::<SignedMessage>(value.clone()) {
357 if self.config.signing.enabled {
359 let verification_result = self.signing_manager.verify_message(&signed);
360
361 if self.config.is_event_processor_enabled() {
363 let client_id = "anonymous"; let event = SecurityEvent {
365 event_type: if verification_result.is_ok() {
366 "signature.verified".to_string()
367 } else {
368 "signature.failed".to_string()
369 },
370 client_id: client_id.to_string(),
371 timestamp: std::time::SystemTime::now()
372 .duration_since(std::time::UNIX_EPOCH)
373 .map(|d| d.as_secs())
374 .unwrap_or(0),
375 metadata: serde_json::json!({
376 "signature": &signed.signature.signature,
377 }),
378 };
379 let _ = self.event_processor.process_event(event).await;
380 }
381
382 if let Err(e) = verification_result {
383 error!("Message signature verification failed: {}", e);
384 let error_response = error_response(
385 Value::Null,
386 error_codes::UNAUTHORIZED,
387 "Invalid message signature".to_string(),
388 None,
389 );
390 return match serde_json::to_string(&error_response) {
391 Ok(json) => Some(json),
392 Err(e) => {
393 error!("Failed to serialize error response: {}", e);
394 None
395 },
396 };
397 }
398 }
399
400 value = signed.message;
402 }
403
404 let has_id = value.get("id").is_some();
406
407 if has_id {
408 if let Ok(request) = serde_json::from_value::<JsonRpcRequest>(value.clone()) {
410 let response = self.handle_request(request).await;
411 return self.maybe_sign_response(response).await;
412 }
413 } else {
414 if let Ok(notification) = serde_json::from_value::<JsonRpcNotification>(value.clone()) {
416 self.handle_notification(notification).await;
417 return None; }
419 }
420
421 if let Some(obj) = value.as_object() {
423 if !obj.contains_key("method") {
424 let id = obj.get("id").cloned().unwrap_or(Value::Null);
426 let response_id = id;
427
428 let error_response = error_response(
429 response_id,
430 error_codes::INVALID_REQUEST,
431 "Missing method field".to_string(),
432 None,
433 );
434 match serde_json::to_string(&error_response) {
435 Ok(json) => return Some(json),
436 Err(e) => {
437 error!("Failed to serialize error response: {}", e);
438 return None;
439 },
440 }
441 }
442 }
443
444 let error_response = error_response(
446 Value::Null,
447 error_codes::INVALID_REQUEST,
448 "Invalid request".to_string(),
449 None,
450 );
451
452 match serde_json::to_string(&error_response) {
453 Ok(json) => Some(json),
454 Err(e) => {
455 error!("Failed to serialize error response: {}", e);
456 None
457 },
458 }
459 }
460
461 pub async fn handle_request(&self, request: JsonRpcRequest) -> JsonRpcResponse {
466 let telemetry = self.component_manager.telemetry_provider();
468 let request_span = telemetry.start_span(&format!("mcp.request.{}", request.method));
469
470 if request.jsonrpc != "2.0" {
472 telemetry.set_status(&request_span, true, Some("Invalid JSON-RPC version"));
473 telemetry.end_span(request_span);
474 return error_response(
475 request.id.clone(),
476 error_codes::INVALID_REQUEST,
477 "Invalid JSON-RPC version".to_string(),
478 None,
479 );
480 }
481
482 let mut authorization: Option<String> = None;
484
485 let params = &request.params;
487 if let Some(meta) = params.get("_meta") {
488 if let Some(auth_token) = meta.get("authToken").and_then(|v| v.as_str()) {
489 authorization = Some(format!("Bearer {auth_token}"));
490 }
491 } else if let Some(arguments) = params.get("arguments") {
492 if let Some(meta) = arguments.get("_meta") {
494 if let Some(auth_token) = meta.get("authToken").and_then(|v| v.as_str()) {
495 authorization = Some(format!("Bearer {auth_token}"));
496 }
497 }
498 }
499
500 let auth_context = match self
502 .auth_manager
503 .authenticate(authorization.as_deref())
504 .await
505 {
506 Ok(ctx) => {
507 let client_id = ctx.client_id.as_deref().unwrap_or("anonymous");
509 self.track_security_event("auth.success", client_id, serde_json::json!({}))
510 .await;
511
512 telemetry.record_metric(TelemetryMetric {
514 name: "auth.attempts".to_string(),
515 value: MetricValue::Counter(1),
516 labels: vec![
517 ("status".to_string(), "success".to_string()),
518 ("client_id".to_string(), client_id.to_string()),
519 ],
520 });
521 ctx
522 },
523 Err(e) => {
524 warn!("Authentication failed: {}", e);
525
526 self.track_security_event(
528 "auth.failure",
529 "anonymous",
530 serde_json::json!({
531 "reason": e.to_string()
532 }),
533 )
534 .await;
535
536 telemetry.record_metric(TelemetryMetric {
538 name: "auth.attempts".to_string(),
539 value: MetricValue::Counter(1),
540 labels: vec![
541 ("status".to_string(), "failure".to_string()),
542 ("reason".to_string(), e.to_string()),
543 ],
544 });
545
546 telemetry.set_status(&request_span, true, Some("Authentication failed"));
547 telemetry.end_span(request_span);
548 return error_response(
549 request.id.clone(),
550 error_codes::UNAUTHORIZED,
551 "Authentication required".to_string(),
552 None,
553 );
554 },
555 };
556
557 let client_id = auth_context.client_id.as_deref().unwrap_or("anonymous");
559 let rate_limit_key = RateLimitKey {
560 client_id: client_id.to_string(),
561 method: Some(request.method.clone()),
562 };
563 let rate_limit_decision = match self.rate_limiter.check_rate_limit(&rate_limit_key).await {
564 Ok(decision) => decision,
565 Err(e) => {
566 error!("Rate limiter error: {}", e);
567 crate::traits::RateLimitDecision {
569 allowed: true,
570 tokens_remaining: 0.0,
571 reset_after: Duration::ZERO,
572 }
573 },
574 };
575
576 self.track_security_event(
578 if rate_limit_decision.allowed {
579 "rate_limit.allowed"
580 } else {
581 "rate_limit.exceeded"
582 },
583 client_id,
584 serde_json::json!({
585 "method": &request.method,
586 "tokens_remaining": rate_limit_decision.tokens_remaining,
587 }),
588 )
589 .await;
590
591 if !rate_limit_decision.allowed {
592 warn!(
593 "Rate limit exceeded for client {} on method {}",
594 client_id, request.method
595 );
596
597 telemetry.record_metric(TelemetryMetric {
599 name: "rate_limit.exceeded".to_string(),
600 value: MetricValue::Counter(1),
601 labels: vec![
602 ("client_id".to_string(), client_id.to_string()),
603 ("method".to_string(), request.method.clone()),
604 ],
605 });
606
607 if self.event_processor.is_monitored(client_id) {
609 error!(
610 "Client {} is under attack monitoring - circuit breaker may activate",
611 client_id
612 );
613 }
614
615 telemetry.set_status(&request_span, true, Some("Rate limit exceeded"));
616 telemetry.end_span(request_span);
617 return error_response(
618 request.id.clone(),
619 error_codes::RATE_LIMITED,
620 format!(
621 "Rate limit exceeded. Try again in {} seconds",
622 rate_limit_decision.reset_after.as_secs()
623 ),
624 Some(serde_json::json!({
625 "retry_after": rate_limit_decision.reset_after.as_secs(),
626 "tokens_remaining": rate_limit_decision.tokens_remaining,
627 })),
628 );
629 }
630
631 let should_scan_request = match request.method.as_str() {
633 "tools/call" => {
634 let params = &request.params;
636 if let Some(tool_name) = params.get("name").and_then(|n| n.as_str()) {
637 !matches!(tool_name, "scan_text" | "scan_file" | "scan_json")
638 } else {
639 true
640 }
641 },
642 _ => true,
643 };
644
645 if should_scan_request {
646 match self.scan_request(&request).await {
647 Ok(threats) if !threats.is_empty() => {
648 self.shield.record_threats(&threats);
649 error!("Threats detected in request: {:?}", threats);
650
651 for threat in &threats {
653 self.track_security_event(
654 "threat.detected",
655 client_id,
656 serde_json::json!({
657 "threat_type": format!("{:?}", threat.threat_type),
658 "severity": format!("{:?}", threat.severity),
659 "threat": threat,
660 }),
661 )
662 .await;
663
664 telemetry.record_metric(TelemetryMetric {
666 name: "threats.detected".to_string(),
667 value: MetricValue::Counter(1),
668 labels: vec![
669 (
670 "threat_type".to_string(),
671 format!("{:?}", threat.threat_type),
672 ),
673 ("severity".to_string(), format!("{:?}", threat.severity)),
674 ("client_id".to_string(), client_id.to_string()),
675 ],
676 });
677 }
678
679 if let Err(e) = self
681 .rate_limiter
682 .apply_penalty(client_id, self.config.rate_limit.threat_penalty_multiplier)
683 .await
684 {
685 error!("Failed to apply rate limit penalty: {}", e);
686 }
687
688 telemetry.set_status(&request_span, true, Some("Threat detected"));
689 telemetry.end_span(request_span);
690 return error_response(
691 request.id.clone(),
692 error_codes::THREAT_DETECTED,
693 "Security threat detected".to_string(),
694 Some(serde_json::to_value(&threats).unwrap_or(Value::Null)),
695 );
696 },
697 Err(e) => {
698 error!("Failed to scan request: {}", e);
699 },
701 _ => {},
702 }
703 }
704
705 let request_id = match &request.id {
707 Value::String(s) => s.clone(),
708 Value::Number(n) => n.to_string(),
709 Value::Null => "null".to_string(),
710 _ => serde_json::to_string(&request.id).unwrap_or_else(|_| "unknown".to_string()),
711 };
712
713 self.track_security_event(
714 "request.received",
715 client_id,
716 serde_json::json!({
717 "method": &request.method,
718 "request_id": &request_id,
719 }),
720 )
721 .await;
722
723 let start_time = std::time::Instant::now();
724
725 if let Some(stability) = ApiRegistry::get_stability(&request.method) {
727 use crate::versioning::ApiStability;
728 if stability == ApiStability::Experimental && !ApiRegistry::experimental_enabled() {
729 return error_response(
730 request.id.clone(),
731 -32601,
732 format!(
733 "Method '{}' is experimental and not enabled",
734 request.method
735 ),
736 None,
737 );
738 }
739 }
740
741 let result = match request.method.as_str() {
743 "initialize" => self.handle_initialize(Some(request.params)).await,
744 "initialized" => self.handle_initialized(Some(request.params)).await,
745 "shutdown" => self.handle_shutdown(Some(request.params)).await,
746
747 "tools/list" => {
748 self.handle_tools_list(Some(request.params), &auth_context)
749 .await
750 },
751 "tools/call" => {
752 self.handle_tools_call(Some(request.params), &auth_context)
753 .await
754 },
755
756 "resources/list" => self.handle_resources_list(Some(request.params)).await,
757 "resources/read" => {
758 self.handle_resources_read(Some(request.params), &auth_context)
759 .await
760 },
761
762 "prompts/list" => self.handle_prompts_list(Some(request.params)).await,
763 "prompts/get" => self.handle_prompts_get(Some(request.params)).await,
764
765 "logging/setLevel" => self.handle_logging_set_level(Some(request.params)).await,
766
767 "security/status" => self.handle_security_status(Some(request.params)).await,
769 "security/threats" => self.handle_security_threats(Some(request.params)).await,
770 "security/rate_limit_status" => {
771 self.handle_rate_limit_status(Some(request.params), &auth_context)
772 .await
773 },
774
775 "$/cancelRequest" => self.handle_cancel_request(Some(request.params)).await,
777
778 method => Err(ServerError::MethodNotFound(method.to_string())),
779 };
780
781 let success = result.is_ok();
782 let mut response = match result {
783 Ok(value) => success_response(request.id.clone(), value),
784 Err(error) => {
785 telemetry.set_status(&request_span, true, Some(&error.to_string()));
786 error_response(
787 request.id.clone(),
788 error.to_json_rpc_error().code,
789 error.to_json_rpc_error().message,
790 error.to_json_rpc_error().data,
791 )
792 },
793 };
794
795 if success {
797 if let Some(ref mut result_value) = response.result {
798 add_version_metadata(result_value);
799 }
800 }
801
802 let duration_ms = start_time.elapsed().as_millis() as u64;
804 self.track_security_event(
805 "response.sent",
806 client_id,
807 serde_json::json!({
808 "method": &request.method,
809 "request_id": &request_id,
810 "duration_ms": duration_ms,
811 "success": success,
812 }),
813 )
814 .await;
815
816 telemetry.record_metric(TelemetryMetric {
818 name: "mcp.request.duration".to_string(),
819 value: MetricValue::Histogram(duration_ms as f64),
820 labels: vec![
821 ("method".to_string(), request.method.clone()),
822 ("success".to_string(), success.to_string()),
823 ],
824 });
825
826 telemetry.record_metric(TelemetryMetric {
827 name: "mcp.request.count".to_string(),
828 value: MetricValue::Counter(1),
829 labels: vec![
830 ("method".to_string(), request.method.clone()),
831 (
832 "status".to_string(),
833 if success { "success" } else { "error" }.to_string(),
834 ),
835 ],
836 });
837
838 telemetry.end_span(request_span);
840
841 response
842 }
843
844 pub async fn handle_notification(&self, notification: JsonRpcNotification) {
849 debug!("Received notification: {}", notification.method);
850
851 match notification.method.as_str() {
852 "initialized" => {
853 info!("Client sent initialized notification");
854 self.shield.set_active(true);
855 },
856 "$/cancelRequest" => {
857 if !notification.params.is_null() {
858 if let Some(id) = notification.params.get("id") {
859 debug!("Cancel request for id: {:?}", id);
860 }
862 }
863 },
864 _ => {
865 debug!("Unknown notification: {}", notification.method);
866 },
867 }
868 }
869
870 async fn scan_request(&self, request: &JsonRpcRequest) -> Result<Vec<Threat>, ServerError> {
872 let mut all_threats = Vec::new();
873 let circuit_breaker = self.component_manager.circuit_breaker();
874
875 match circuit_breaker
877 .call_json(
878 "scanner.scan_text",
879 serde_json::json!({
880 "text": &request.method
881 }),
882 )
883 .await
884 {
885 Ok(_) => {
886 match self.scanner.scan_text(&request.method) {
888 Ok(threats) => all_threats.extend(threats),
889 Err(e) => {
890 error!("Failed to scan method: {}", e);
891 },
892 }
893 },
894 Err(e) => {
895 warn!("Circuit breaker open for scanner.scan_text: {}", e);
896 },
898 }
899
900 let params = &request.params;
902 match circuit_breaker
903 .call_json(
904 "scanner.scan_json",
905 serde_json::json!({
906 "params": params
907 }),
908 )
909 .await
910 {
911 Ok(_) => {
912 match self.scanner.scan_json(params) {
914 Ok(threats) => all_threats.extend(threats),
915 Err(e) => {
916 error!("Failed to scan params: {}", e);
917 },
918 }
919 },
920 Err(e) => {
921 warn!("Circuit breaker open for scanner.scan_json: {}", e);
922 },
924 }
925
926 Ok(all_threats)
927 }
928
929 async fn handle_initialize(&self, params: Option<Value>) -> Result<Value, ServerError> {
931 let params: InitializeParams = if let Some(p) = params {
932 serde_json::from_value(p).map_err(|e| {
933 ServerError::InvalidParams(format!("Invalid initialize params: {e}"))
934 })?
935 } else {
936 return Err(ServerError::InvalidParams(
937 "Missing initialize params".to_string(),
938 ));
939 };
940
941 info!(
942 "Initialize request from {} v{}",
943 params.client_info.name, params.client_info.version
944 );
945
946 let session_id = uuid::Uuid::new_v4().to_string();
948 let mut store = self.session_store.lock().await;
949 store.sessions.insert(
950 session_id.clone(),
951 SessionInfo {
952 id: session_id,
953 client_info: Some(params.client_info),
954 created_at: std::time::Instant::now(),
955 threats_blocked: 0,
956 last_activity: std::time::Instant::now(),
957 },
958 );
959
960 if params.protocol_version != PROTOCOL_VERSION {
962 warn!(
963 "Client requested protocol version {}, we support {}",
964 params.protocol_version, PROTOCOL_VERSION
965 );
966 return Err(ServerError::InvalidParams(format!(
967 "Unsupported protocol version: {}. Supported version: {}",
968 params.protocol_version, PROTOCOL_VERSION
969 )));
970 }
971
972 let protocol_version = PROTOCOL_VERSION.to_string();
973
974 let result = InitializeResult {
975 protocol_version,
976 capabilities: self.capabilities.clone(),
977 server_info: self.server_info.clone(),
978 };
979
980 serde_json::to_value(result).map_err(|e| ServerError::InternalError(e.to_string()))
981 }
982
983 async fn handle_initialized(&self, _params: Option<Value>) -> Result<Value, ServerError> {
985 Ok(Value::Null)
987 }
988
989 async fn handle_shutdown(&self, _params: Option<Value>) -> Result<Value, ServerError> {
991 info!("Shutdown requested");
992 self.shield.set_active(false);
993 Ok(Value::Null)
994 }
995
996 async fn handle_tools_list(
998 &self,
999 _params: Option<Value>,
1000 auth: &AuthContext,
1001 ) -> Result<Value, ServerError> {
1002 let all_tools = vec![
1004 Tool {
1005 name: "scan_text".to_string(),
1006 description: "Scan text for security threats including unicode attacks and injection attempts".to_string(),
1007 input_schema: serde_json::json!({
1008 "type": "object",
1009 "properties": {
1010 "text": {
1011 "type": "string",
1012 "description": "Text to scan for threats"
1013 }
1014 },
1015 "required": ["text"]
1016 }),
1017 },
1018 Tool {
1019 name: "scan_file".to_string(),
1020 description: "Scan a file for security threats".to_string(),
1021 input_schema: serde_json::json!({
1022 "type": "object",
1023 "properties": {
1024 "path": {
1025 "type": "string",
1026 "description": "File path to scan"
1027 }
1028 },
1029 "required": ["path"]
1030 }),
1031 },
1032 Tool {
1033 name: "scan_json".to_string(),
1034 description: "Scan JSON data for security threats".to_string(),
1035 input_schema: serde_json::json!({
1036 "type": "object",
1037 "properties": {
1038 "data": {
1039 "type": "object",
1040 "description": "JSON data to scan"
1041 }
1042 },
1043 "required": ["data"]
1044 }),
1045 },
1046 Tool {
1047 name: "get_security_info".to_string(),
1048 description: "Get current security information and statistics".to_string(),
1049 input_schema: serde_json::json!({
1050 "type": "object",
1051 "properties": {},
1052 }),
1053 },
1054 Tool {
1055 name: "verify_signature".to_string(),
1056 description: "Verify message signature".to_string(),
1057 input_schema: serde_json::json!({
1058 "type": "object",
1059 "properties": {
1060 "message": {
1061 "type": "string",
1062 "description": "Message to verify"
1063 },
1064 "signature": {
1065 "type": "string",
1066 "description": "Signature to verify"
1067 }
1068 },
1069 "required": ["message", "signature"]
1070 }),
1071 },
1072 Tool {
1073 name: "get_shield_status".to_string(),
1074 description: "Get current shield status and protection level".to_string(),
1075 input_schema: serde_json::json!({
1076 "type": "object",
1077 "properties": {},
1078 }),
1079 },
1080 ];
1081
1082 let client_id = auth.client_id.as_deref().unwrap_or("anonymous");
1084 let tools = if self.config.auth.enabled {
1085 let allowed_tools = self
1086 .component_manager
1087 .permission_manager()
1088 .get_allowed_tools(client_id)
1089 .await
1090 .map_err(|e| {
1091 ServerError::InternalError(format!("Failed to get allowed tools: {e}"))
1092 })?;
1093
1094 all_tools
1096 .into_iter()
1097 .filter(|tool| allowed_tools.contains(&tool.name))
1098 .collect()
1099 } else {
1100 all_tools
1102 };
1103
1104 info!("Client {} has access to {} tools", client_id, tools.len());
1105
1106 let result = ToolsListResult { tools };
1107 serde_json::to_value(result).map_err(|e| ServerError::InternalError(e.to_string()))
1108 }
1109
1110 async fn handle_tools_call(
1112 &self,
1113 params: Option<Value>,
1114 auth: &AuthContext,
1115 ) -> Result<Value, ServerError> {
1116 let params: ToolCallParams = if let Some(p) = params {
1117 serde_json::from_value(p)
1118 .map_err(|e| ServerError::InvalidParams(format!("Invalid tool call params: {e}")))?
1119 } else {
1120 return Err(ServerError::InvalidParams(
1121 "Missing tool call params".to_string(),
1122 ));
1123 };
1124
1125 if self.config.auth.enabled {
1127 self.auth_manager
1128 .authorize_tool(auth, ¶ms.name)
1129 .map_err(|_e| ServerError::Unauthorized)?;
1130 }
1131
1132 if self.config.auth.enabled {
1134 let permission_context = crate::permissions::PermissionContext {
1135 auth_token: None, scopes: auth.scopes.clone(),
1137 threat_level: self.get_current_threat_level(),
1138 request_metadata: std::collections::HashMap::new(),
1139 };
1140
1141 let client_id = auth.client_id.as_deref().unwrap_or("anonymous");
1142 let permission = self
1143 .component_manager
1144 .permission_manager()
1145 .check_permission(client_id, ¶ms.name, &permission_context)
1146 .await
1147 .map_err(|e| ServerError::InternalError(format!("Permission check failed: {e}")))?;
1148
1149 if let crate::permissions::Permission::Deny(reason) = permission {
1150 warn!("Tool access denied for {}: {}", client_id, reason);
1151 return Err(ServerError::Unauthorized);
1152 }
1153 }
1154
1155 let bulkhead = self.component_manager.bulkhead();
1157 let tool_name = params.name.clone();
1158 let arguments = params.arguments;
1159
1160 let _bulkhead_result = bulkhead
1161 .execute_json(
1162 &format!("tool.{}", tool_name),
1163 serde_json::json!({
1164 "tool": &tool_name,
1165 "arguments": &arguments
1166 }),
1167 )
1168 .await
1169 .map_err(|e| {
1170 warn!("Bulkhead rejected tool execution for {}: {}", tool_name, e);
1171 ServerError::InternalError(format!("Tool execution rejected: {}", e))
1172 })?;
1173
1174 let result = tokio::time::timeout(
1176 Duration::from_secs(self.config.server.request_timeout_secs),
1177 self.execute_tool(&tool_name, arguments),
1178 )
1179 .await
1180 .map_err(|_| ServerError::Timeout)??;
1181
1182 Ok(result)
1183 }
1184
1185 async fn execute_tool(&self, name: &str, arguments: Value) -> Result<Value, ServerError> {
1187 match name {
1188 "scan_text" => {
1189 let text = arguments
1190 .get("text")
1191 .and_then(|v| v.as_str())
1192 .ok_or_else(|| {
1193 ServerError::InvalidParams("Missing 'text' argument".to_string())
1194 })?;
1195
1196 let threats = self
1197 .scanner
1198 .scan_text(text)
1199 .map_err(|e| ServerError::InternalError(e.to_string()))?;
1200
1201 if !threats.is_empty() {
1202 self.shield.record_threats(&threats);
1203 }
1204
1205 let neutralizer = self.component_manager.threat_neutralizer();
1207 let neutralization_mode = self.config.neutralization.mode;
1208
1209 let (neutralization_results, final_content) = if !threats.is_empty()
1210 && neutralization_mode != crate::neutralizer::NeutralizationMode::ReportOnly
1211 {
1212 let mut results = Vec::new();
1213 let mut current_content = text.to_string();
1214
1215 let client_id = {
1217 let store = self.session_store.lock().await;
1218 store
1219 .sessions
1220 .values()
1221 .next()
1222 .and_then(|s| s.client_info.as_ref())
1223 .map_or_else(|| "anonymous".to_string(), |ci| ci.name.clone())
1224 };
1225
1226 let audit_logger = if self.config.audit.enabled {
1228 Some(self.component_manager.audit_logger().clone())
1229 } else {
1230 None
1231 };
1232
1233 let telemetry = if self.config.telemetry.enabled {
1235 Some(crate::telemetry::SecureTelemetry::new(
1236 self.component_manager.telemetry_provider().clone(),
1237 ))
1238 } else {
1239 None
1240 };
1241
1242 let neutralization_metrics =
1244 crate::neutralizer::metrics::NeutralizationMetrics::new(Arc::new(
1245 crate::telemetry::metrics::MetricsCollector::new(),
1246 ));
1247
1248 let validator = crate::neutralizer::validation::NeutralizationValidator::new(
1250 crate::neutralizer::validation::ValidationConfig::default(),
1251 );
1252
1253 for threat in &threats {
1254 if let Some(ref logger) = audit_logger {
1256 let event = crate::audit::AuditEvent::new(
1257 crate::audit::AuditEventType::NeutralizationStarted {
1258 client_id: client_id.clone(),
1259 threat_id: format!(
1260 "threat-{:?}-{}",
1261 threat.threat_type,
1262 match &threat.location {
1263 crate::scanner::Location::Text { offset, .. } =>
1264 *offset,
1265 crate::scanner::Location::Json { path } => path.len(),
1266 crate::scanner::Location::Binary { offset } => *offset,
1267 }
1268 ),
1269 threat_type: format!("{:?}", threat.threat_type),
1270 },
1271 crate::audit::AuditSeverity::Info,
1272 )
1273 .with_client_id(client_id.clone())
1274 .with_tags(vec!["neutralization".to_string(), "security".to_string()]);
1275
1276 if let Err(e) = logger.log(event).await {
1277 tracing::warn!("Failed to log neutralization start: {}", e);
1278 }
1279 }
1280
1281 if let Err(e) = validator.validate_input(threat, ¤t_content) {
1283 tracing::warn!("Input validation failed: {}", e);
1284 neutralization_metrics.record_validation_failure(&e.to_string());
1285 continue; }
1287
1288 neutralization_metrics
1290 .record_content_size(current_content.len(), &threat.threat_type);
1291
1292 let start_time = std::time::Instant::now();
1293
1294 match neutralizer.neutralize(threat, ¤t_content).await {
1295 Ok(result) => {
1296 if let Err(e) =
1298 validator.validate_output(threat, ¤t_content, &result)
1299 {
1300 tracing::error!("Output validation failed: {}", e);
1301 neutralization_metrics
1302 .record_validation_failure(&e.to_string());
1303 continue; }
1305 let duration = start_time.elapsed();
1307 let action_str = match result.action_taken {
1308 crate::neutralizer::NeutralizeAction::Sanitized => "sanitized",
1309 crate::neutralizer::NeutralizeAction::Parameterized => {
1310 "parameterized"
1311 },
1312 crate::neutralizer::NeutralizeAction::Normalized => {
1313 "normalized"
1314 },
1315 crate::neutralizer::NeutralizeAction::Escaped => "escaped",
1316 crate::neutralizer::NeutralizeAction::Removed => "removed",
1317 crate::neutralizer::NeutralizeAction::Quarantined => {
1318 "quarantined"
1319 },
1320 crate::neutralizer::NeutralizeAction::NoAction => "no_action",
1321 };
1322
1323 if let Some(ref logger) = audit_logger {
1324 let event = crate::audit::AuditEvent::new(
1325 crate::audit::AuditEventType::NeutralizationCompleted {
1326 client_id: client_id.clone(),
1327 threat_id: format!(
1328 "threat-{:?}-{}",
1329 threat.threat_type,
1330 match &threat.location {
1331 crate::scanner::Location::Text {
1332 offset,
1333 ..
1334 } => *offset,
1335 crate::scanner::Location::Json { path } =>
1336 path.len(),
1337 crate::scanner::Location::Binary { offset } =>
1338 *offset,
1339 }
1340 ),
1341 action: action_str.to_string(),
1342 duration_ms: duration.as_millis() as u64,
1343 },
1344 crate::audit::AuditSeverity::Info,
1345 )
1346 .with_client_id(client_id.clone())
1347 .with_context(
1348 "confidence".to_string(),
1349 serde_json::to_value(result.confidence_score)
1350 .unwrap_or(serde_json::Value::Null),
1351 )
1352 .with_tags(vec![
1353 "neutralization".to_string(),
1354 "security".to_string(),
1355 "success".to_string(),
1356 ]);
1357
1358 if let Err(e) = logger.log(event).await {
1359 tracing::warn!(
1360 "Failed to log neutralization completion: {}",
1361 e
1362 );
1363 }
1364 }
1365
1366 if let Some(ref telemetry) = telemetry {
1368 telemetry.record_neutralization(
1369 &format!("{:?}", threat.threat_type),
1370 action_str,
1371 duration.as_millis() as f64,
1372 true,
1373 );
1374 }
1375
1376 neutralization_metrics.record_neutralization(
1378 &threat.threat_type,
1379 &result.action_taken,
1380 true,
1381 duration,
1382 neutralization_mode,
1383 );
1384
1385 if let Some(ref sanitized) = result.sanitized_content {
1386 current_content = sanitized.clone();
1387 }
1388 results.push(serde_json::json!({
1389 "threat_type": format!("{:?}", threat.threat_type),
1390 "action": format!("{}", result.action_taken),
1391 "confidence": result.confidence_score,
1392 "time_us": duration.as_micros() as u64,
1393 }));
1394 },
1395 Err(e) => {
1396 if let Some(ref logger) = audit_logger {
1398 let event = crate::audit::AuditEvent::new(
1399 crate::audit::AuditEventType::NeutralizationFailed {
1400 client_id: client_id.clone(),
1401 threat_id: format!(
1402 "threat-{:?}-{}",
1403 threat.threat_type,
1404 match &threat.location {
1405 crate::scanner::Location::Text {
1406 offset,
1407 ..
1408 } => *offset,
1409 crate::scanner::Location::Json { path } =>
1410 path.len(),
1411 crate::scanner::Location::Binary { offset } =>
1412 *offset,
1413 }
1414 ),
1415 error: e.to_string(),
1416 },
1417 crate::audit::AuditSeverity::Error,
1418 )
1419 .with_client_id(client_id.clone())
1420 .with_tags(vec![
1421 "neutralization".to_string(),
1422 "security".to_string(),
1423 "failure".to_string(),
1424 ]);
1425
1426 if let Err(e) = logger.log(event).await {
1427 tracing::warn!(
1428 "Failed to log neutralization failure: {}",
1429 e
1430 );
1431 }
1432 }
1433
1434 if let Some(ref telemetry) = telemetry {
1436 telemetry.record_neutralization(
1437 &format!("{:?}", threat.threat_type),
1438 "failed",
1439 start_time.elapsed().as_millis() as f64,
1440 false,
1441 );
1442 }
1443
1444 neutralization_metrics.record_neutralization(
1446 &threat.threat_type,
1447 &crate::neutralizer::NeutralizeAction::NoAction,
1448 false,
1449 start_time.elapsed(),
1450 neutralization_mode,
1451 );
1452
1453 tracing::warn!("Neutralization failed for threat: {}", e);
1454 },
1455 }
1456 }
1457
1458 if let Some(ref telemetry) = telemetry {
1460 let batch_start = std::time::Instant::now();
1461 let neutralized_count = results.len();
1462 telemetry.record_neutralization_batch(
1463 threats.len(),
1464 neutralized_count,
1465 batch_start.elapsed().as_millis() as f64,
1466 );
1467 }
1468
1469 let neutralized_count = results.len();
1471 let failed_count = threats.len().saturating_sub(neutralized_count);
1472 neutralization_metrics.record_batch_neutralization(
1473 threats.len(),
1474 neutralized_count,
1475 failed_count,
1476 std::time::Duration::from_millis(100), );
1478
1479 (Some(results), Some(current_content))
1480 } else {
1481 (None, None)
1482 };
1483
1484 let threat_data = threats.iter().map(|t| {
1486 serde_json::json!({
1487 "type": match &t.threat_type {
1488 crate::scanner::ThreatType::UnicodeInvisible => "unicode_invisible".to_string(),
1489 crate::scanner::ThreatType::UnicodeBiDi => "unicode_bidi".to_string(),
1490 crate::scanner::ThreatType::UnicodeHomograph => "unicode_homograph".to_string(),
1491 crate::scanner::ThreatType::UnicodeControl => "unicode_control".to_string(),
1492 crate::scanner::ThreatType::PromptInjection => "prompt_injection".to_string(),
1493 crate::scanner::ThreatType::CommandInjection => "command_injection".to_string(),
1494 crate::scanner::ThreatType::PathTraversal => "path_traversal".to_string(),
1495 crate::scanner::ThreatType::SqlInjection => "sql_injection".to_string(),
1496 crate::scanner::ThreatType::CrossSiteScripting => "cross_site_scripting".to_string(),
1497 crate::scanner::ThreatType::LdapInjection => "ldap_injection".to_string(),
1498 crate::scanner::ThreatType::XmlInjection => "xml_injection".to_string(),
1499 crate::scanner::ThreatType::NoSqlInjection => "nosql_injection".to_string(),
1500 crate::scanner::ThreatType::SessionIdExposure => "session_id_exposure".to_string(),
1501 crate::scanner::ThreatType::ToolPoisoning => "tool_poisoning".to_string(),
1502 crate::scanner::ThreatType::TokenTheft => "token_theft".to_string(),
1503 crate::scanner::ThreatType::DosPotential => "dos_potential".to_string(),
1504 crate::scanner::ThreatType::Custom(s) => s.to_lowercase().replace(' ', "_"),
1505 },
1506 "severity": format!("{:?}", t.severity).to_lowercase(),
1507 "description": &t.description,
1508 "location": t.location,
1509 })
1510 }).collect::<Vec<_>>();
1511
1512 let mut response_json = serde_json::json!({
1513 "safe": threats.is_empty(),
1514 "threats": threat_data,
1515 "scan_info": {
1516 "text_length": text.len(),
1517 "threats_found": threats.len(),
1518 }
1519 });
1520
1521 if let Some(neutralization) = neutralization_results {
1523 response_json["neutralization"] = serde_json::json!({
1524 "mode": format!("{:?}", neutralization_mode),
1525 "results": neutralization,
1526 "neutralized": true,
1527 });
1528 }
1529
1530 if let Some(sanitized) = final_content {
1531 response_json["sanitized_text"] = serde_json::Value::String(sanitized);
1532 }
1533
1534 Ok(serde_json::json!({
1535 "content": [{
1536 "type": "text",
1537 "text": serde_json::to_string(&response_json)
1538 .unwrap_or_else(|_| r#"{"error": "Failed to serialize response"}"#.to_string())
1539 }]
1540 }))
1541 },
1542
1543 "scan_file" => {
1544 let path = arguments
1545 .get("path")
1546 .and_then(|v| v.as_str())
1547 .ok_or_else(|| {
1548 ServerError::InvalidParams("Missing 'path' argument".to_string())
1549 })?;
1550
1551 if path.contains("..") || path.starts_with('/') {
1553 return Err(ServerError::InvalidParams("Invalid file path".to_string()));
1554 }
1555
1556 let content = tokio::fs::read_to_string(path)
1557 .await
1558 .map_err(|e| ServerError::InternalError(format!("Failed to read file: {e}")))?;
1559
1560 let threats = self
1561 .scanner
1562 .scan_text(&content)
1563 .map_err(|e| ServerError::InternalError(e.to_string()))?;
1564
1565 if !threats.is_empty() {
1566 self.shield.record_threats(&threats);
1567 }
1568
1569 let threat_data = threats.iter().map(|t| {
1571 serde_json::json!({
1572 "type": match &t.threat_type {
1573 crate::scanner::ThreatType::UnicodeInvisible => "unicode_invisible".to_string(),
1574 crate::scanner::ThreatType::UnicodeBiDi => "unicode_bidi".to_string(),
1575 crate::scanner::ThreatType::UnicodeHomograph => "unicode_homograph".to_string(),
1576 crate::scanner::ThreatType::UnicodeControl => "unicode_control".to_string(),
1577 crate::scanner::ThreatType::PromptInjection => "prompt_injection".to_string(),
1578 crate::scanner::ThreatType::CommandInjection => "command_injection".to_string(),
1579 crate::scanner::ThreatType::PathTraversal => "path_traversal".to_string(),
1580 crate::scanner::ThreatType::SqlInjection => "sql_injection".to_string(),
1581 crate::scanner::ThreatType::CrossSiteScripting => "cross_site_scripting".to_string(),
1582 crate::scanner::ThreatType::LdapInjection => "ldap_injection".to_string(),
1583 crate::scanner::ThreatType::XmlInjection => "xml_injection".to_string(),
1584 crate::scanner::ThreatType::NoSqlInjection => "nosql_injection".to_string(),
1585 crate::scanner::ThreatType::SessionIdExposure => "session_id_exposure".to_string(),
1586 crate::scanner::ThreatType::ToolPoisoning => "tool_poisoning".to_string(),
1587 crate::scanner::ThreatType::TokenTheft => "token_theft".to_string(),
1588 crate::scanner::ThreatType::DosPotential => "dos_potential".to_string(),
1589 crate::scanner::ThreatType::Custom(s) => s.to_lowercase().replace(' ', "_"),
1590 },
1591 "severity": format!("{:?}", t.severity).to_lowercase(),
1592 "description": &t.description,
1593 "location": t.location,
1594 })
1595 }).collect::<Vec<_>>();
1596
1597 let response_json = serde_json::json!({
1598 "safe": threats.is_empty(),
1599 "threats": threat_data,
1600 "scan_info": {
1601 "file_path": path,
1602 "file_size": content.len(),
1603 "threats_found": threats.len(),
1604 }
1605 });
1606
1607 Ok(serde_json::json!({
1608 "content": [{
1609 "type": "text",
1610 "text": serde_json::to_string(&response_json)
1611 .unwrap_or_else(|_| r#"{"error": "Failed to serialize response"}"#.to_string())
1612 }]
1613 }))
1614 },
1615
1616 "scan_json" => {
1617 let data = arguments.get("data").ok_or_else(|| {
1618 ServerError::InvalidParams("Missing 'data' argument".to_string())
1619 })?;
1620
1621 let threats = self
1622 .scanner
1623 .scan_json(data)
1624 .map_err(|e| ServerError::InternalError(e.to_string()))?;
1625
1626 if !threats.is_empty() {
1627 self.shield.record_threats(&threats);
1628 }
1629
1630 let threat_data = threats.iter().map(|t| {
1632 serde_json::json!({
1633 "type": match &t.threat_type {
1634 crate::scanner::ThreatType::UnicodeInvisible => "unicode_invisible".to_string(),
1635 crate::scanner::ThreatType::UnicodeBiDi => "unicode_bidi".to_string(),
1636 crate::scanner::ThreatType::UnicodeHomograph => "unicode_homograph".to_string(),
1637 crate::scanner::ThreatType::UnicodeControl => "unicode_control".to_string(),
1638 crate::scanner::ThreatType::PromptInjection => "prompt_injection".to_string(),
1639 crate::scanner::ThreatType::CommandInjection => "command_injection".to_string(),
1640 crate::scanner::ThreatType::PathTraversal => "path_traversal".to_string(),
1641 crate::scanner::ThreatType::SqlInjection => "sql_injection".to_string(),
1642 crate::scanner::ThreatType::CrossSiteScripting => "cross_site_scripting".to_string(),
1643 crate::scanner::ThreatType::LdapInjection => "ldap_injection".to_string(),
1644 crate::scanner::ThreatType::XmlInjection => "xml_injection".to_string(),
1645 crate::scanner::ThreatType::NoSqlInjection => "nosql_injection".to_string(),
1646 crate::scanner::ThreatType::SessionIdExposure => "session_id_exposure".to_string(),
1647 crate::scanner::ThreatType::ToolPoisoning => "tool_poisoning".to_string(),
1648 crate::scanner::ThreatType::TokenTheft => "token_theft".to_string(),
1649 crate::scanner::ThreatType::DosPotential => "dos_potential".to_string(),
1650 crate::scanner::ThreatType::Custom(s) => s.to_lowercase().replace(' ', "_"),
1651 },
1652 "severity": format!("{:?}", t.severity).to_lowercase(),
1653 "description": &t.description,
1654 "location": t.location,
1655 })
1656 }).collect::<Vec<_>>();
1657
1658 let response_json = serde_json::json!({
1659 "safe": threats.is_empty(),
1660 "threats": threat_data,
1661 "scan_info": {
1662 "data_type": "json",
1663 "threats_found": threats.len(),
1664 }
1665 });
1666
1667 Ok(serde_json::json!({
1668 "content": [{
1669 "type": "text",
1670 "text": serde_json::to_string(&response_json)
1671 .unwrap_or_else(|_| r#"{"error": "Failed to serialize response"}"#.to_string())
1672 }]
1673 }))
1674 },
1675
1676 "get_security_info" => {
1677 let shield_stats = self.shield.stats();
1679 let event_stats = self.event_processor.get_stats();
1680 let rate_limiter_stats = self.rate_limiter.get_stats();
1681 let permission_stats = self.component_manager.permission_manager().get_stats();
1682
1683 let security_info = serde_json::json!({
1684 "status": "active",
1685 "enhanced_mode": self.component_manager.is_enhanced_mode(),
1686 "shield": {
1687 "threats_blocked": shield_stats.threats_blocked,
1688 "active": shield_stats.active,
1689 },
1690 "event_processor": {
1691 "events_processed": event_stats.events_processed,
1692 "events_per_second": event_stats.events_per_second,
1693 "buffer_utilization": event_stats.buffer_utilization,
1694 },
1695 "rate_limiter": {
1696 "requests_allowed": rate_limiter_stats.requests_allowed,
1697 "requests_denied": rate_limiter_stats.requests_denied,
1698 },
1699 "permissions": {
1700 "total_checks": permission_stats.total_checks,
1701 "allowed": permission_stats.allowed,
1702 "denied": permission_stats.denied,
1703 }
1704 });
1705
1706 Ok(serde_json::json!({
1708 "content": [{
1709 "type": "text",
1710 "text": serde_json::to_string_pretty(&security_info)
1711 .unwrap_or_else(|_| "Failed to serialize security info".to_string())
1712 }]
1713 }))
1714 },
1715
1716 "verify_signature" => {
1717 let message = arguments
1718 .get("message")
1719 .and_then(|v| v.as_str())
1720 .ok_or_else(|| {
1721 ServerError::InvalidParams("Missing 'message' argument".to_string())
1722 })?;
1723
1724 let signature = arguments
1725 .get("signature")
1726 .and_then(|v| v.as_str())
1727 .ok_or_else(|| {
1728 ServerError::InvalidParams("Missing 'signature' argument".to_string())
1729 })?;
1730
1731 let message_value = serde_json::from_str(message)
1733 .unwrap_or_else(|_| serde_json::json!({"raw": message}));
1734
1735 let signed_message = SignedMessage {
1737 message: message_value,
1738 signature: MessageSignature {
1739 algorithm: self.config.signing.algorithm.clone(),
1740 signature: signature.to_string(),
1741 timestamp: None,
1742 key_id: None,
1743 },
1744 };
1745
1746 let verification_result = match self.signing_manager.verify_message(&signed_message)
1748 {
1749 Ok(()) => serde_json::json!({
1750 "valid": true,
1751 "algorithm": self.config.signing.algorithm.to_string(),
1752 "message": message,
1753 "error": null
1754 }),
1755 Err(e) => serde_json::json!({
1756 "valid": false,
1757 "algorithm": self.config.signing.algorithm.to_string(),
1758 "message": message,
1759 "error": e.to_string()
1760 }),
1761 };
1762
1763 Ok(serde_json::json!({
1765 "content": [{
1766 "type": "text",
1767 "text": serde_json::to_string_pretty(&verification_result)
1768 .unwrap_or_else(|_| "Failed to serialize verification result".to_string())
1769 }]
1770 }))
1771 },
1772
1773 "get_shield_status" => {
1774 let shield_info = self.shield.get_info();
1775
1776 let shield_status = serde_json::json!({
1777 "active": shield_info.active,
1778 "protection_level": if shield_info.threats_blocked > 100 { "high" } else if shield_info.threats_blocked > 10 { "medium" } else { "low" },
1779 "threats_blocked": shield_info.threats_blocked,
1780 "uptime_seconds": shield_info.uptime.as_secs(),
1781 "recent_threat_rate": shield_info.recent_threat_rate,
1782 });
1783
1784 Ok(serde_json::json!({
1786 "content": [{
1787 "type": "text",
1788 "text": serde_json::to_string_pretty(&shield_status)
1789 .unwrap_or_else(|_| "Failed to serialize shield status".to_string())
1790 }]
1791 }))
1792 },
1793
1794 _ => Err(ServerError::InvalidParams(format!("Unknown tool: {name}"))),
1795 }
1796 }
1797
1798 async fn handle_resources_list(&self, _params: Option<Value>) -> Result<Value, ServerError> {
1800 let resources = vec![
1801 Resource {
1802 uri: "threat-patterns://default".to_string(),
1803 name: "Default Threat Patterns".to_string(),
1804 description: Some("Built-in threat detection patterns".to_string()),
1805 mime_type: Some("application/json".to_string()),
1806 },
1807 Resource {
1808 uri: "security-report://latest".to_string(),
1809 name: "Latest Security Report".to_string(),
1810 description: Some("Current security status and recent threats".to_string()),
1811 mime_type: Some("application/json".to_string()),
1812 },
1813 Resource {
1814 uri: "config://security".to_string(),
1815 name: "security-config".to_string(),
1816 description: Some("Security configuration and settings".to_string()),
1817 mime_type: Some("application/json".to_string()),
1818 },
1819 Resource {
1820 uri: "threat-db://current".to_string(),
1821 name: "threat-database".to_string(),
1822 description: Some("Current threat database and patterns".to_string()),
1823 mime_type: Some("application/json".to_string()),
1824 },
1825 ];
1826
1827 let result = ResourcesListResult { resources };
1828 serde_json::to_value(result).map_err(|e| ServerError::InternalError(e.to_string()))
1829 }
1830
1831 async fn handle_resources_read(
1833 &self,
1834 params: Option<Value>,
1835 auth: &AuthContext,
1836 ) -> Result<Value, ServerError> {
1837 let params: ResourceReadParams = if let Some(p) = params {
1838 serde_json::from_value(p).map_err(|e| {
1839 ServerError::InvalidParams(format!("Invalid resource read params: {e}"))
1840 })?
1841 } else {
1842 return Err(ServerError::InvalidParams(
1843 "Missing resource read params".to_string(),
1844 ));
1845 };
1846
1847 self.auth_manager
1849 .authorize_resource(auth, ¶ms.uri)
1850 .map_err(|_e| ServerError::Unauthorized)?;
1851
1852 match params.uri.as_str() {
1853 "threat-patterns://default" => {
1854 let patterns = self.scanner.patterns.get_all_patterns();
1855 let content = ResourceContent {
1856 uri: params.uri,
1857 mime_type: Some("application/json".to_string()),
1858 content: ResourceContentType::Text {
1859 text: serde_json::to_string_pretty(&patterns)
1860 .unwrap_or_else(|_| "{}".to_string()),
1861 },
1862 };
1863 Ok(serde_json::to_value(content)
1864 .map_err(|e| ServerError::InternalError(e.to_string()))?)
1865 },
1866
1867 "security-report://latest" => {
1868 let shield_info = self.shield.get_info();
1869 let stats = self.scanner.stats();
1870 let recent_threats = self.shield.get_recent_threats(10);
1871
1872 let report = serde_json::json!({
1873 "generated_at": chrono::Utc::now().to_rfc3339(),
1874 "status": {
1875 "active": shield_info.active,
1876 "uptime_seconds": shield_info.uptime.as_secs(),
1877 "threats_blocked": shield_info.threats_blocked,
1878 "threat_rate_per_minute": shield_info.recent_threat_rate,
1879 },
1880 "scanner_stats": {
1881 "unicode_threats_detected": stats.unicode_threats_detected,
1882 "injection_threats_detected": stats.injection_threats_detected,
1883 "total_scans": stats.total_scans,
1884 },
1885 "recent_threats": recent_threats,
1886 });
1887
1888 let content = ResourceContent {
1889 uri: params.uri,
1890 mime_type: Some("application/json".to_string()),
1891 content: ResourceContentType::Text {
1892 text: serde_json::to_string_pretty(&report)
1893 .unwrap_or_else(|_| "{}".to_string()),
1894 },
1895 };
1896 Ok(serde_json::to_value(content)
1897 .map_err(|e| ServerError::InternalError(e.to_string()))?)
1898 },
1899
1900 _ => Err(ServerError::InvalidParams(format!(
1901 "Unknown resource URI: {}",
1902 params.uri
1903 ))),
1904 }
1905 }
1906
1907 async fn handle_logging_set_level(&self, params: Option<Value>) -> Result<Value, ServerError> {
1909 if let Some(params) = params {
1910 if let Some(level) = params.get("level").and_then(|v| v.as_str()) {
1911 info!("Setting log level to: {}", level);
1912 Ok(Value::Null)
1914 } else {
1915 Err(ServerError::InvalidParams(
1916 "Missing 'level' parameter".to_string(),
1917 ))
1918 }
1919 } else {
1920 Err(ServerError::InvalidParams("Missing parameters".to_string()))
1921 }
1922 }
1923
1924 async fn handle_security_status(&self, _params: Option<Value>) -> Result<Value, ServerError> {
1926 let shield_info = self.shield.get_info();
1927 let stats = self.scanner.stats();
1928
1929 Ok(serde_json::json!({
1930 "active": shield_info.active,
1931 "uptime_seconds": shield_info.uptime.as_secs(),
1932 "threats_blocked": shield_info.threats_blocked,
1933 "scanner_stats": {
1934 "unicode_threats": stats.unicode_threats_detected,
1935 "injection_threats": stats.injection_threats_detected,
1936 "total_scans": stats.total_scans,
1937 }
1938 }))
1939 }
1940
1941 async fn handle_security_threats(&self, params: Option<Value>) -> Result<Value, ServerError> {
1943 let limit = params
1944 .as_ref()
1945 .and_then(|p| p.get("limit"))
1946 .and_then(serde_json::Value::as_u64)
1947 .unwrap_or(100) as usize;
1948
1949 let recent_threats = self.shield.get_recent_threats(limit);
1950
1951 Ok(serde_json::json!({
1952 "threats": recent_threats,
1953 "count": recent_threats.len(),
1954 }))
1955 }
1956
1957 async fn handle_rate_limit_status(
1959 &self,
1960 _params: Option<Value>,
1961 auth: &AuthContext,
1962 ) -> Result<Value, ServerError> {
1963 let client_id = auth.client_id.as_deref().unwrap_or("anonymous");
1964
1965 let stats = self.rate_limiter.get_stats();
1967
1968 Ok(serde_json::json!({
1969 "client_id": client_id,
1970 "requests_allowed": stats.requests_allowed,
1971 "requests_denied": stats.requests_denied,
1972 "active_buckets": stats.active_buckets,
1973 "enabled": self.config.rate_limit.enabled,
1974 }))
1975 }
1976
1977 async fn handle_cancel_request(&self, params: Option<Value>) -> Result<Value, ServerError> {
1979 if let Some(params) = params {
1980 if let Some(id) = params.get("id") {
1981 debug!("Cancel request for id: {:?}", id);
1982 Ok(Value::Null)
1984 } else {
1985 Err(ServerError::InvalidParams(
1986 "Missing 'id' parameter".to_string(),
1987 ))
1988 }
1989 } else {
1990 Err(ServerError::InvalidParams("Missing parameters".to_string()))
1991 }
1992 }
1993
1994 async fn handle_prompts_list(&self, _params: Option<Value>) -> Result<Value, ServerError> {
1996 let prompts = vec![
1997 Prompt {
1998 name: "analyze-security".to_string(),
1999 description: "Analyze security implications of the given input".to_string(),
2000 arguments: vec![PromptArgument {
2001 name: "target".to_string(),
2002 description: "The target to analyze (code, text, or data)".to_string(),
2003 required: true,
2004 }],
2005 },
2006 Prompt {
2007 name: "threat-report".to_string(),
2008 description: "Generate a detailed threat report".to_string(),
2009 arguments: vec![PromptArgument {
2010 name: "scope".to_string(),
2011 description: "The scope of the report (recent, all, specific-type)".to_string(),
2012 required: false,
2013 }],
2014 },
2015 Prompt {
2016 name: "security-best-practices".to_string(),
2017 description: "Provide security best practices for a given context".to_string(),
2018 arguments: vec![PromptArgument {
2019 name: "context".to_string(),
2020 description: "The context (web, api, database, etc.)".to_string(),
2021 required: true,
2022 }],
2023 },
2024 ];
2025
2026 let result = PromptsListResult { prompts };
2027 serde_json::to_value(result).map_err(|e| ServerError::InternalError(e.to_string()))
2028 }
2029
2030 async fn handle_prompts_get(&self, params: Option<Value>) -> Result<Value, ServerError> {
2032 let name = params
2033 .as_ref()
2034 .and_then(|p| p.get("name"))
2035 .and_then(|v| v.as_str())
2036 .ok_or_else(|| ServerError::InvalidParams("Missing 'name' parameter".to_string()))?;
2037
2038 let arguments = params
2039 .as_ref()
2040 .and_then(|p| p.get("arguments"))
2041 .cloned()
2042 .unwrap_or_else(|| serde_json::json!({}));
2043
2044 match name {
2045 "analyze-security" => {
2046 let target = arguments
2047 .get("target")
2048 .and_then(|v| v.as_str())
2049 .ok_or_else(|| {
2050 ServerError::InvalidParams("Missing 'target' argument".to_string())
2051 })?;
2052
2053 let messages = vec![serde_json::json!({
2054 "role": "user",
2055 "content": {
2056 "type": "text",
2057 "text": format!("Please analyze the security implications of the following:\n\n{}", target)
2058 }
2059 })];
2060
2061 Ok(serde_json::json!({
2062 "messages": messages
2063 }))
2064 },
2065
2066 "threat-report" => {
2067 let scope = arguments
2068 .get("scope")
2069 .and_then(|v| v.as_str())
2070 .unwrap_or("recent");
2071
2072 let messages = vec![serde_json::json!({
2073 "role": "user",
2074 "content": {
2075 "type": "text",
2076 "text": format!("Generate a threat report with scope: {}", scope)
2077 }
2078 })];
2079
2080 Ok(serde_json::json!({
2081 "messages": messages
2082 }))
2083 },
2084
2085 "security-best-practices" => {
2086 let context = arguments
2087 .get("context")
2088 .and_then(|v| v.as_str())
2089 .ok_or_else(|| {
2090 ServerError::InvalidParams("Missing 'context' argument".to_string())
2091 })?;
2092
2093 let messages = vec![serde_json::json!({
2094 "role": "user",
2095 "content": {
2096 "type": "text",
2097 "text": format!("What are the security best practices for {}?", context)
2098 }
2099 })];
2100
2101 Ok(serde_json::json!({
2102 "messages": messages
2103 }))
2104 },
2105
2106 _ => Err(ServerError::InvalidParams(format!(
2107 "Unknown prompt: {name}"
2108 ))),
2109 }
2110 }
2111
2112 pub async fn run_with_transport(self: Arc<Self>) -> Result<()> {
2114 info!("Starting KindlyGuard with transport layer");
2115 self.shield.set_active(true);
2116
2117 let handler = Arc::new(ServerMessageHandler {
2119 server: self.clone(),
2120 });
2121
2122 let mut transport_manager = TransportManager::new(self.config.transport.clone(), handler)?;
2124
2125 let factory = DefaultTransportFactory;
2127 for transport_config in &self.config.transport.transports {
2128 if transport_config.enabled {
2129 match factory.create(transport_config) {
2130 Ok(transport) => {
2131 info!("Adding transport: {:?}", transport_config.transport_type);
2132 transport_manager.add_transport(transport)?;
2133 },
2134 Err(e) => {
2135 warn!(
2136 "Failed to create transport {:?}: {}",
2137 transport_config.transport_type, e
2138 );
2139 },
2140 }
2141 }
2142 }
2143
2144 transport_manager.start().await?;
2146
2147 let mut config_watcher = None;
2149 if let Ok(config_path) = std::env::var("KINDLY_GUARD_CONFIG") {
2150 let path = std::path::PathBuf::from(&config_path);
2151 if path.exists() {
2152 info!("Setting up config hot-reload for {:?}", path);
2153
2154 let mut watcher =
2155 crate::config::reload::ConfigWatcher::new(path, (*self.config).clone())?;
2156
2157 let reload_handler = Arc::new(ServerConfigReloadHandler {
2159 server: self.clone(),
2160 });
2161 watcher.add_handler(reload_handler).await;
2162
2163 watcher.start().await?;
2164 config_watcher = Some(watcher);
2165 }
2166 }
2167
2168 if self.config.audit.enabled {
2170 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2171 let audit_event = AuditEvent::new(
2172 AuditEventType::ServerStarted {
2173 version: env!("CARGO_PKG_VERSION").to_string(),
2174 },
2175 AuditSeverity::Info,
2176 );
2177 let _ = self.component_manager.audit_logger().log(audit_event).await;
2178 }
2179
2180 tokio::signal::ctrl_c().await?;
2182
2183 if let Some(mut watcher) = config_watcher {
2185 watcher.stop().await?;
2186 }
2187
2188 info!("Shutting down transport layer");
2189 transport_manager.stop().await?;
2190
2191 if self.config.audit.enabled {
2193 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2194 let audit_event = AuditEvent::new(
2195 AuditEventType::ServerStopped {
2196 reason: "Signal received".to_string(),
2197 },
2198 AuditSeverity::Info,
2199 );
2200 let _ = self.component_manager.audit_logger().log(audit_event).await;
2201 }
2202
2203 self.shield.set_active(false);
2204 Ok(())
2205 }
2206
2207 pub async fn run_http(self: Arc<Self>, bind_addr: &str) -> Result<()> {
2209 use crate::transport::{HttpTransport, Transport};
2210
2211 info!(
2212 "Starting KindlyGuard MCP server in HTTP mode at {}",
2213 bind_addr
2214 );
2215 self.shield.set_active(true);
2216
2217 if self.config.audit.enabled {
2219 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2220
2221 let audit_event = AuditEvent::new(
2222 AuditEventType::ServerStarted {
2223 version: env!("CARGO_PKG_VERSION").to_string(),
2224 },
2225 AuditSeverity::Info,
2226 );
2227
2228 let audit_logger = self.component_manager.audit_logger();
2229 if let Err(e) = audit_logger.log(audit_event).await {
2230 warn!("Failed to log server startup audit event: {}", e);
2231 }
2232 }
2233
2234 let http_config = serde_json::json!({
2236 "bind_addr": bind_addr,
2237 "tls": false,
2238 "max_body_size": 10 * 1024 * 1024,
2239 "request_timeout_ms": 30000
2240 });
2241
2242 let mut transport = HttpTransport::new(http_config)?;
2244 transport.start().await?;
2245
2246 info!("HTTP server started on {}", bind_addr);
2247
2248 tokio::signal::ctrl_c().await?;
2250
2251 transport.stop().await?;
2253
2254 if self.config.audit.enabled {
2256 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2257
2258 let audit_event = AuditEvent::new(
2259 AuditEventType::ServerStopped {
2260 reason: "Shutdown signal received".to_string(),
2261 },
2262 AuditSeverity::Info,
2263 );
2264
2265 let audit_logger = self.component_manager.audit_logger();
2266 if let Err(e) = audit_logger.log(audit_event).await {
2267 warn!("Failed to log server shutdown audit event: {}", e);
2268 }
2269 }
2270
2271 self.shield.set_active(false);
2272 Ok(())
2273 }
2274
2275 pub async fn run_proxy(self: Arc<Self>, bind_addr: &str) -> Result<()> {
2277 use crate::transport::ProxyTransport;
2278
2279 info!("Starting KindlyGuard as HTTPS proxy at {}", bind_addr);
2280 self.shield.set_active(true);
2281
2282 if self.config.audit.enabled {
2284 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2285
2286 let audit_event = AuditEvent::new(
2287 AuditEventType::ServerStarted {
2288 version: env!("CARGO_PKG_VERSION").to_string(),
2289 },
2290 AuditSeverity::Info,
2291 );
2292
2293 let audit_logger = self.component_manager.audit_logger();
2294 if let Err(e) = audit_logger.log(audit_event).await {
2295 warn!("Failed to log server startup audit event: {}", e);
2296 }
2297 }
2298
2299 let proxy_config = serde_json::json!({
2301 "bind_addr": bind_addr,
2302 "intercept_https": true,
2303 "ai_services": [
2304 "api.anthropic.com",
2305 "api.openai.com",
2306 "generativelanguage.googleapis.com",
2307 "api.cohere.ai",
2308 "api.mistral.ai"
2309 ]
2310 });
2311
2312 let transport = ProxyTransport::new(proxy_config)?;
2314 transport.serve(self.clone()).await?;
2315
2316 if self.config.audit.enabled {
2318 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2319
2320 let audit_event = AuditEvent::new(
2321 AuditEventType::ServerStopped {
2322 reason: "Normal shutdown".to_string(),
2323 },
2324 AuditSeverity::Info,
2325 );
2326
2327 let audit_logger = self.component_manager.audit_logger();
2328 if let Err(e) = audit_logger.log(audit_event).await {
2329 warn!("Failed to log server shutdown audit event: {}", e);
2330 }
2331 }
2332
2333 self.shield.set_active(false);
2334 Ok(())
2335 }
2336
2337 pub async fn run_stdio(self: Arc<Self>) -> Result<()> {
2339 use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
2340
2341 info!("Starting KindlyGuard MCP server in stdio mode");
2342 self.shield.set_active(true);
2343
2344 if self.config.audit.enabled {
2346 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2347
2348 let audit_event = AuditEvent::new(
2349 AuditEventType::ServerStarted {
2350 version: env!("CARGO_PKG_VERSION").to_string(),
2351 },
2352 AuditSeverity::Info,
2353 );
2354
2355 let audit_logger = self.component_manager.audit_logger();
2356 if let Err(e) = audit_logger.log(audit_event).await {
2357 warn!("Failed to log server startup audit event: {}", e);
2358 }
2359 }
2360
2361 let stdin = tokio::io::stdin();
2362 let mut stdout = tokio::io::stdout();
2363 let mut reader = BufReader::new(stdin);
2364 let mut line = String::new();
2365
2366 loop {
2367 line.clear();
2368 match reader.read_line(&mut line).await {
2369 Ok(0) => break, Ok(_) => {
2371 let line = line.trim();
2372 if line.is_empty() {
2373 continue;
2374 }
2375
2376 debug!("Received: {}", line);
2377
2378 if let Some(response) = self.handle_message(line).await {
2379 stdout.write_all(response.as_bytes()).await?;
2380 stdout.write_all(b"\n").await?;
2381 stdout.flush().await?;
2382 debug!("Sent: {}", response);
2383 }
2384 },
2385 Err(e) => {
2386 error!("Error reading from stdin: {}", e);
2387 break;
2388 },
2389 }
2390 }
2391
2392 info!("MCP server shutting down");
2393 self.shield.set_active(false);
2394
2395 if self.config.audit.enabled {
2397 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2398
2399 let audit_event = AuditEvent::new(
2400 AuditEventType::ServerStopped {
2401 reason: "Normal shutdown".to_string(),
2402 },
2403 AuditSeverity::Info,
2404 );
2405
2406 let audit_logger = self.component_manager.audit_logger();
2407 if let Err(e) = audit_logger.log(audit_event).await {
2408 warn!("Failed to log server shutdown audit event: {}", e);
2409 }
2410 }
2411
2412 let telemetry = self.component_manager.telemetry_provider();
2414 if let Err(e) = telemetry.flush().await {
2415 error!("Failed to flush telemetry: {}", e);
2416 }
2417 if let Err(e) = telemetry.shutdown().await {
2418 error!("Failed to shutdown telemetry: {}", e);
2419 }
2420
2421 Ok(())
2422 }
2423
2424 async fn maybe_sign_response(&self, response: JsonRpcResponse) -> Option<String> {
2426 if self.config.signing.enabled {
2427 let response_value = serde_json::to_value(&response).ok()?;
2429
2430 match self.signing_manager.sign_message(&response_value) {
2431 Ok(signed) => {
2432 Some(serde_json::to_string(&signed).unwrap_or_else(|e| {
2433 error!("Failed to serialize signed response: {}", e);
2434 serde_json::to_string(&response).unwrap_or_else(|_| {
2436 r#"{"jsonrpc":"2.0","error":{"code":-32603,"message":"Internal error"},"id":null}"#.to_string()
2437 })
2438 }))
2439 },
2440 Err(e) => {
2441 error!("Failed to sign response: {}", e);
2442 Some(serde_json::to_string(&response).unwrap_or_else(|_| {
2444 r#"{"jsonrpc":"2.0","error":{"code":-32603,"message":"Internal error"},"id":null}"#.to_string()
2445 }))
2446 },
2447 }
2448 } else {
2449 Some(serde_json::to_string(&response).unwrap_or_else(|e| {
2450 error!("Failed to serialize response: {}", e);
2451 r#"{"jsonrpc":"2.0","error":{"code":-32603,"message":"Internal error"},"id":null}"#
2452 .to_string()
2453 }))
2454 }
2455 }
2456
2457 fn get_current_threat_level(&self) -> crate::permissions::ThreatLevel {
2459 let shield_stats = self.shield.stats();
2461 let threats_blocked = shield_stats.threats_blocked;
2462
2463 if threats_blocked == 0 {
2465 crate::permissions::ThreatLevel::Safe
2466 } else if threats_blocked < 5 {
2467 crate::permissions::ThreatLevel::Low
2468 } else if threats_blocked < 20 {
2469 crate::permissions::ThreatLevel::Medium
2470 } else if threats_blocked < 50 {
2471 crate::permissions::ThreatLevel::High
2472 } else {
2473 crate::permissions::ThreatLevel::Critical
2474 }
2475 }
2476}
2477
2478struct ServerMessageHandler {
2480 server: Arc<McpServer>,
2481}
2482
2483#[async_trait]
2484impl MessageHandler for ServerMessageHandler {
2485 async fn handle_message(
2486 &self,
2487 message: TransportMessage,
2488 connection: &dyn TransportConnection,
2489 ) -> Result<Option<TransportMessage>> {
2490 let conn_info = connection.connection_info();
2491 let client_id = conn_info.client_id.as_deref().unwrap_or("unknown");
2492
2493 debug!("Handling message {} from client {}", message.id, client_id);
2494
2495 let json_str = serde_json::to_string(&message.payload)?;
2497
2498 if let Some(response_str) = self.server.handle_message(&json_str).await {
2500 let response_value: Value = serde_json::from_str(&response_str)?;
2502
2503 let response = TransportMessage {
2505 id: uuid::Uuid::new_v4().to_string(),
2506 payload: response_value,
2507 metadata: crate::transport::TransportMetadata {
2508 client_id: conn_info.client_id.clone(),
2509 timestamp: Some(chrono::Utc::now()),
2510 trace_id: message.metadata.trace_id.clone(),
2511 ..Default::default()
2512 },
2513 };
2514
2515 Ok(Some(response))
2516 } else {
2517 Ok(None)
2518 }
2519 }
2520
2521 async fn on_connect(&self, connection: &dyn TransportConnection) -> Result<()> {
2522 let conn_info = connection.connection_info();
2523 info!(
2524 "Client connected: {:?} via {:?}",
2525 conn_info.client_id, conn_info.transport_type
2526 );
2527
2528 if self.server.config.audit.enabled {
2530 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2531 let event = AuditEvent::new(
2532 AuditEventType::Custom {
2533 event_type: "transport.connect".to_string(),
2534 data: serde_json::json!({
2535 "transport": format!("{:?}", conn_info.transport_type),
2536 "client_id": conn_info.client_id,
2537 "remote_addr": conn_info.remote_addr,
2538 }),
2539 },
2540 AuditSeverity::Info,
2541 );
2542 let _ = self
2543 .server
2544 .component_manager
2545 .audit_logger()
2546 .log(event)
2547 .await;
2548 }
2549
2550 Ok(())
2551 }
2552
2553 async fn on_disconnect(&self, connection: &dyn TransportConnection) -> Result<()> {
2554 let conn_info = connection.connection_info();
2555 info!("Client disconnected: {:?}", conn_info.client_id);
2556
2557 if self.server.config.audit.enabled {
2559 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2560 let event = AuditEvent::new(
2561 AuditEventType::Custom {
2562 event_type: "transport.disconnect".to_string(),
2563 data: serde_json::json!({
2564 "transport": format!("{:?}", conn_info.transport_type),
2565 "client_id": conn_info.client_id,
2566 }),
2567 },
2568 AuditSeverity::Info,
2569 );
2570 let _ = self
2571 .server
2572 .component_manager
2573 .audit_logger()
2574 .log(event)
2575 .await;
2576 }
2577
2578 Ok(())
2579 }
2580}
2581
2582struct ServerConfigReloadHandler {
2584 server: Arc<McpServer>,
2585}
2586
2587#[async_trait]
2588impl crate::config::reload::ConfigChangeHandler for ServerConfigReloadHandler {
2589 async fn handle_change(&self, event: crate::config::reload::ConfigReloadEvent) -> Result<()> {
2590 use crate::config::reload::ConfigReloadEvent;
2591
2592 match event {
2593 ConfigReloadEvent::Reloaded {
2594 new_config,
2595 changed_fields,
2596 ..
2597 } => {
2598 info!("Applying configuration changes: {:?}", changed_fields);
2599
2600 for field in &changed_fields {
2602 match field.as_str() {
2603 "shield.enabled" => {
2604 self.server.shield.set_enabled(new_config.shield.enabled);
2605 info!(
2606 "Shield display {}",
2607 if new_config.shield.enabled {
2608 "enabled"
2609 } else {
2610 "disabled"
2611 }
2612 );
2613 },
2614 "shield.update_interval_ms" => {
2615 debug!("Shield update interval changed");
2617 },
2618 "rate_limit.enabled" => {
2619 info!(
2621 "Rate limiting {}",
2622 if new_config.rate_limit.enabled {
2623 "enabled"
2624 } else {
2625 "disabled"
2626 }
2627 );
2628 },
2629 "rate_limit.default_rpm" => {
2630 info!(
2632 "Rate limit updated to {} requests/minute",
2633 new_config.rate_limit.default_rpm
2634 );
2635 },
2636 field if field.starts_with("scanner.") => {
2637 warn!(
2639 "Scanner configuration changed ({}), restart required",
2640 field
2641 );
2642 },
2643 _ => {
2644 debug!("Configuration field {} changed", field);
2645 },
2646 }
2647 }
2648
2649 if self.server.config.audit.enabled {
2651 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2652 let audit_event = AuditEvent::new(
2653 AuditEventType::ConfigReloaded {
2654 success: true,
2655 error: None,
2656 },
2657 AuditSeverity::Info,
2658 );
2659 let _ = self
2660 .server
2661 .component_manager
2662 .audit_logger()
2663 .log(audit_event)
2664 .await;
2665 }
2666 },
2667 ConfigReloadEvent::Failed { error, .. } => {
2668 error!("Configuration reload failed: {}", error);
2669
2670 if self.server.config.audit.enabled {
2672 use crate::audit::{AuditEvent, AuditEventType, AuditSeverity};
2673 let audit_event = AuditEvent::new(
2674 AuditEventType::ConfigReloaded {
2675 success: false,
2676 error: Some(error),
2677 },
2678 AuditSeverity::Error,
2679 );
2680 let _ = self
2681 .server
2682 .component_manager
2683 .audit_logger()
2684 .log(audit_event)
2685 .await;
2686 }
2687 },
2688 ConfigReloadEvent::ValidationFailed { errors, .. } => {
2689 error!(
2690 "Configuration validation failed with {} errors",
2691 errors.len()
2692 );
2693 },
2694 }
2695
2696 Ok(())
2697 }
2698
2699 async fn validate_config(
2700 &self,
2701 config: &Config,
2702 ) -> Result<Vec<crate::config::reload::ValidationError>> {
2703 use crate::config::reload::{ValidationError, ValidationSeverity};
2704
2705 let mut errors = Vec::new();
2706
2707 if config.transport.transports.is_empty() {
2709 errors.push(ValidationError {
2710 field: "transport.transports".to_string(),
2711 message: "At least one transport must be configured".to_string(),
2712 severity: ValidationSeverity::Error,
2713 });
2714 }
2715
2716 if !config.auth.enabled && !config.rate_limit.enabled {
2718 errors.push(ValidationError {
2719 field: "security".to_string(),
2720 message: "Both authentication and rate limiting are disabled".to_string(),
2721 severity: ValidationSeverity::Warning,
2722 });
2723 }
2724
2725 let default_handler = crate::config::reload::DefaultConfigChangeHandler::new();
2727 let default_errors = default_handler.validate_config(config).await?;
2728 errors.extend(default_errors);
2729
2730 Ok(errors)
2731 }
2732
2733 fn get_reloadable_fields(&self) -> Vec<String> {
2734 vec![
2735 "shield.*".to_string(),
2736 "rate_limit.enabled".to_string(),
2737 "rate_limit.default_rpm".to_string(),
2738 "audit.enabled".to_string(),
2739 "telemetry.export_interval_seconds".to_string(),
2740 ]
2741 }
2742}