use anyhow::Result;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{mpsc, Mutex, RwLock};
use tokio::time::interval;
use tracing::{debug, error, info, warn};
use crate::protocol::claude_code::{
create_status_notification, threat_to_severity, ClaudeCodeError, ClaudeCodeErrorCode,
LastThreatInfo, PerformanceMetrics, ShieldControlAction, ShieldControlRequest,
ShieldControlResponse, ShieldInfoParams, ShieldInfoResponse, ShieldState, ShieldStatistics,
ShieldStatusNotification, ShieldStatusParams, ThreatPattern, ThreatSeverity,
};
use crate::scanner::{SecurityScanner, Threat};
use crate::shield::Shield;
use crate::traits::SecurityEvent;
use super::{
ConnectionInfo, ConnectionStats, Transport, TransportConnection, TransportMessage,
TransportStats, TransportType,
};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClaudeCodeConfig {
pub port: u16,
pub bind_addr: String,
pub enhanced_mode: bool,
pub batch_delay_ms: u64,
pub shared_memory: bool,
pub max_connections: usize,
pub auth_token_env: Option<String>,
pub notifications: NotificationConfig,
}
impl Default for ClaudeCodeConfig {
fn default() -> Self {
Self {
port: 9955,
bind_addr: "127.0.0.1".to_string(),
enhanced_mode: false,
batch_delay_ms: 50,
shared_memory: false,
max_connections: 10,
auth_token_env: Some("CLAUDE_CODE_TOKEN".to_string()),
notifications: NotificationConfig::default(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NotificationConfig {
pub threat_alerts: bool,
pub performance_metrics: bool,
pub detailed_threats: bool,
pub min_interval_ms: u64,
}
impl Default for NotificationConfig {
fn default() -> Self {
Self {
threat_alerts: true,
performance_metrics: true,
detailed_threats: false,
min_interval_ms: 100, }
}
}
#[async_trait]
pub trait EventProcessor: Send + Sync {
async fn process_batch(&self, events: &[SecurityEvent]) -> Result<ShieldStatusParams>;
fn supports_binary(&self) -> bool;
fn supports_shared_memory(&self) -> bool;
fn get_stats(&self) -> ShieldStatistics;
}
pub struct StandardEventProcessor {
shield: Arc<Shield>,
scanner: Arc<SecurityScanner>,
stats: Arc<RwLock<ProcessorStats>>,
}
#[derive(Default)]
struct ProcessorStats {
threats_blocked: u64,
total_scans: u64,
total_scan_time_us: u64,
threats_by_type: std::collections::HashMap<String, u64>,
}
impl StandardEventProcessor {
pub fn new(shield: Arc<Shield>, scanner: Arc<SecurityScanner>) -> Self {
Self {
shield,
scanner,
stats: Arc::new(RwLock::new(ProcessorStats::default())),
}
}
}
#[async_trait]
impl EventProcessor for StandardEventProcessor {
async fn process_batch(&self, events: &[SecurityEvent]) -> Result<ShieldStatusParams> {
let start = std::time::Instant::now();
let mut stats = self.stats.write().await;
for event in events {
match event {
SecurityEvent::ThreatDetected { threat, .. } => {
stats.threats_blocked += 1;
let threat_type = format!("{:?}", threat.threat_type);
*stats.threats_by_type.entry(threat_type).or_insert(0) += 1;
}
SecurityEvent::ScanCompleted { duration_us, .. } => {
stats.total_scans += 1;
stats.total_scan_time_us += duration_us;
}
_ => {}
}
}
let scan_time_us = start.elapsed().as_micros() as u64;
let threat_rate = if stats.total_scans > 0 {
(stats.threats_blocked as f64 / stats.total_scans as f64) * 600.0
} else {
0.0
};
let shield_stats = self.shield.get_stats();
Ok(ShieldStatusParams {
active: self.shield.is_active(),
enhanced: false, threats: stats.threats_blocked,
threat_rate,
last_threat: None, performance: PerformanceMetrics {
scan_time_us,
queue_depth: 0, memory_mb: 0.0, },
})
}
fn supports_binary(&self) -> bool {
false
}
fn supports_shared_memory(&self) -> bool {
false
}
fn get_stats(&self) -> ShieldStatistics {
let stats = self.stats.blocking_read();
let avg_scan_time = if stats.total_scans > 0 {
stats.total_scan_time_us / stats.total_scans
} else {
0
};
ShieldStatistics {
threats_blocked: stats.threats_blocked,
threats_by_type: stats.threats_by_type.clone(),
total_scans: stats.total_scans,
avg_scan_time_us: avg_scan_time,
uptime_seconds: 0, memory_usage_mb: 0.0, }
}
}
#[cfg(feature = "enhanced")]
pub struct EnhancedEventProcessor {
shield: Arc<Shield>,
scanner: Arc<SecurityScanner>,
threats_blocked: AtomicU64,
total_scans: AtomicU64,
total_scan_time_us: AtomicU64,
event_buffer: Arc<dyn EventBuffer>,
}
#[cfg(feature = "enhanced")]
trait EventBuffer: Send + Sync {
fn enqueue(&self, event: &SecurityEvent) -> Result<()>;
fn dequeue_batch(&self, max_count: usize) -> Vec<SecurityEvent>;
}
pub fn create_event_processor(
config: &ClaudeCodeConfig,
shield: Arc<Shield>,
scanner: Arc<SecurityScanner>,
) -> Arc<dyn EventProcessor> {
if config.enhanced_mode {
#[cfg(feature = "enhanced")]
{
info!("Creating enhanced Claude Code event processor");
return Arc::new(EnhancedEventProcessor::new(shield, scanner));
}
warn!("Enhanced mode requested but not available, using standard processor");
}
Arc::new(StandardEventProcessor::new(shield, scanner))
}
pub struct ClaudeCodeTransport {
config: ClaudeCodeConfig,
shield: Arc<Shield>,
scanner: Arc<SecurityScanner>,
event_processor: Arc<dyn EventProcessor>,
running: AtomicBool,
stats: Arc<Mutex<TransportStats>>,
connections: Arc<RwLock<Vec<Arc<ClaudeCodeConnection>>>>,
shutdown_tx: Option<mpsc::Sender<()>>,
}
impl ClaudeCodeTransport {
pub fn new(
config: ClaudeCodeConfig,
shield: Arc<Shield>,
scanner: Arc<SecurityScanner>,
) -> Result<Self> {
let event_processor = create_event_processor(&config, shield.clone(), scanner.clone());
Ok(Self {
config,
shield,
scanner,
event_processor,
running: AtomicBool::new(false),
stats: Arc::new(Mutex::new(TransportStats::default())),
connections: Arc::new(RwLock::new(Vec::new())),
shutdown_tx: None,
})
}
async fn start_notification_loop(&self) -> Result<()> {
let connections = self.connections.clone();
let event_processor = self.event_processor.clone();
let config = self.config.clone();
let mut interval = interval(Duration::from_millis(config.batch_delay_ms));
tokio::spawn(async move {
let mut event_batch = Vec::new();
loop {
interval.tick().await;
if let Ok(status) = event_processor.process_batch(&event_batch).await {
let notification = create_status_notification(status);
let conns = connections.read().await;
for conn in conns.iter() {
if conn.is_connected() {
let msg = TransportMessage {
id: uuid::Uuid::new_v4().to_string(),
payload: serde_json::to_value(¬ification).unwrap(),
metadata: Default::default(),
};
if let Err(e) = conn.send_notification(msg).await {
warn!("Failed to send notification: {}", e);
}
}
}
}
event_batch.clear();
}
});
Ok(())
}
}
#[async_trait]
impl Transport for ClaudeCodeTransport {
fn transport_type(&self) -> TransportType {
TransportType::Custom(9955) }
async fn start(&mut self) -> Result<()> {
if self.running.load(Ordering::Relaxed) {
return Err(anyhow::anyhow!("Transport already running"));
}
let addr = format!("{}:{}", self.config.bind_addr, self.config.port);
info!("Starting Claude Code transport on {}", addr);
self.start_notification_loop().await?;
self.running.store(true, Ordering::Relaxed);
Ok(())
}
async fn stop(&mut self) -> Result<()> {
self.running.store(false, Ordering::Relaxed);
let mut connections = self.connections.write().await;
for conn in connections.iter_mut() {
let _ = conn.close().await;
}
connections.clear();
info!("Stopped Claude Code transport");
Ok(())
}
async fn accept(&mut self) -> Result<Box<dyn TransportConnection>> {
Err(anyhow::anyhow!("Not implemented yet"))
}
async fn connect(&mut self, _address: &str) -> Result<Box<dyn TransportConnection>> {
Err(anyhow::anyhow!("Claude Code transport is server-only"))
}
fn is_running(&self) -> bool {
self.running.load(Ordering::Relaxed)
}
fn get_stats(&self) -> TransportStats {
self.stats.blocking_lock().clone()
}
async fn set_option(&mut self, key: &str, value: serde_json::Value) -> Result<()> {
match key {
"enhanced_mode" => {
if let Some(enabled) = value.as_bool() {
self.config.enhanced_mode = enabled;
self.event_processor = create_event_processor(
&self.config,
self.shield.clone(),
self.scanner.clone(),
);
}
}
"batch_delay_ms" => {
if let Some(delay) = value.as_u64() {
self.config.batch_delay_ms = delay.min(1000); }
}
_ => {}
}
Ok(())
}
}
struct ClaudeCodeConnection {
info: ConnectionInfo,
connected: AtomicBool,
stats: Arc<Mutex<ConnectionStats>>,
outgoing_tx: mpsc::UnboundedSender<TransportMessage>,
}
impl ClaudeCodeConnection {
fn is_connected(&self) -> bool {
self.connected.load(Ordering::Relaxed)
}
async fn send_notification(&self, message: TransportMessage) -> Result<()> {
self.outgoing_tx.send(message)?;
Ok(())
}
async fn close(&self) -> Result<()> {
self.connected.store(false, Ordering::Relaxed);
Ok(())
}
}
#[async_trait]
impl TransportConnection for ClaudeCodeConnection {
fn connection_info(&self) -> &ConnectionInfo {
&self.info
}
async fn send(&mut self, message: TransportMessage) -> Result<()> {
self.outgoing_tx.send(message)?;
let mut stats = self.stats.lock().await;
stats.messages_sent += 1;
Ok(())
}
async fn receive(&mut self) -> Result<Option<TransportMessage>> {
Ok(None)
}
async fn close(&mut self) -> Result<()> {
self.connected.store(false, Ordering::Relaxed);
Ok(())
}
fn is_connected(&self) -> bool {
self.connected.load(Ordering::Relaxed)
}
fn get_stats(&self) -> ConnectionStats {
self.stats.blocking_lock().clone()
}
async fn set_option(&mut self, _key: &str, _value: serde_json::Value) -> Result<()> {
Ok(())
}
}