use anyhow::Result;
use async_trait::async_trait;
use parking_lot::RwLock;
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use tracing::{debug, info, warn};
use super::{
ClientPermissions, Permission, PermissionContext, PermissionRules, PermissionStats,
ThreatLevel, ToolPermissionManager,
};
use crate::traits::Priority;
pub struct EnhancedPermissionManager {
rules: Arc<PermissionRules>,
client_permissions: Arc<RwLock<HashMap<String, ClientPermissions>>>,
risk_scores: Arc<RwLock<HashMap<String, f32>>>,
total_checks: AtomicU64,
allowed_count: AtomicU64,
denied_count: AtomicU64,
pattern_matches: AtomicU64,
}
impl EnhancedPermissionManager {
pub fn new(rules: PermissionRules) -> Self {
Self {
rules: Arc::new(rules),
client_permissions: Arc::new(RwLock::new(HashMap::new())),
risk_scores: Arc::new(RwLock::new(HashMap::new())),
total_checks: AtomicU64::new(0),
allowed_count: AtomicU64::new(0),
denied_count: AtomicU64::new(0),
pattern_matches: AtomicU64::new(0),
}
}
fn get_client_permissions(&self, client_id: &str) -> ClientPermissions {
let permissions = self.client_permissions.read();
let mut perms = permissions
.get(client_id)
.cloned()
.unwrap_or_else(|| self.rules.default_permissions.clone());
let risk_scores = self.risk_scores.read();
if let Some(&risk_score) = risk_scores.get(client_id) {
if risk_score > 0.8 {
perms.max_threat_level = ThreatLevel::Low;
} else if risk_score > 0.6 {
perms.max_threat_level = ThreatLevel::Medium;
}
}
perms
}
fn calculate_risk_score(&self, client_id: &str) -> f32 {
let endpoint_id = self.get_or_create_endpoint_id(client_id);
let mut risk_score: f32 = 0.0;
risk_score.min(1.0f32)
}
fn get_or_create_endpoint_id(&self, client_id: &str) -> u32 {
let hash = client_id
.bytes()
.fold(0u32, |acc, b| acc.wrapping_mul(31).wrapping_add(b as u32));
hash % 1000 }
fn track_permission_event(
&self,
client_id: &str,
tool_name: &str,
allowed: bool,
reason: Option<&str>,
) {
let endpoint_id = self.get_or_create_endpoint_id(client_id);
let event = format!(
"perm:{}:{}:{}:{}",
tool_name,
if allowed { "allow" } else { "deny" },
reason.unwrap_or(""),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
);
let priority = if allowed {
Priority::Normal
} else {
Priority::Urgent
};
debug!("Permission event: {} (priority: {:?})", event, priority);
}
fn detect_suspicious_patterns(&self, client_id: &str) -> bool {
let risk_score = self
.risk_scores
.read()
.get(client_id)
.copied()
.unwrap_or(0.0);
if risk_score > 0.8 {
self.pattern_matches.fetch_add(1, Ordering::Relaxed);
return true;
}
false
}
async fn check_enhanced_permission(
&self,
client_id: &str,
_tool_name: &str,
_context: &PermissionContext,
) -> Option<Permission> {
let risk_score = self.calculate_risk_score(client_id);
{
let mut scores = self.risk_scores.write();
scores.insert(client_id.to_string(), risk_score);
}
if self.detect_suspicious_patterns(client_id) {
warn!("Suspicious pattern detected for client {}", client_id);
return Some(Permission::Deny(
"Suspicious activity pattern detected".to_string(),
));
}
if risk_score > 0.7 {
debug!("High risk score {} for client {}", risk_score, client_id);
}
None
}
}
#[async_trait]
impl ToolPermissionManager for EnhancedPermissionManager {
async fn check_permission(
&self,
client_id: &str,
tool_name: &str,
context: &PermissionContext,
) -> Result<Permission> {
debug!(
"Enhanced permission check for client {} tool {}",
client_id, tool_name
);
self.total_checks.fetch_add(1, Ordering::Relaxed);
if let Some(denial) = self
.check_enhanced_permission(client_id, tool_name, context)
.await
{
if let Permission::Deny(ref reason) = denial {
self.denied_count.fetch_add(1, Ordering::Relaxed);
self.track_permission_event(client_id, tool_name, false, Some(reason));
}
return Ok(denial);
}
let permissions = self.get_client_permissions(client_id);
if self.rules.global_deny_list.contains(tool_name) {
let reason = "Tool globally denied";
self.denied_count.fetch_add(1, Ordering::Relaxed);
self.track_permission_event(client_id, tool_name, false, Some(reason));
return Ok(Permission::Deny(reason.to_string()));
}
if permissions.denied_tools.contains(tool_name) {
let reason = format!("Tool denied for client {}", client_id);
self.denied_count.fetch_add(1, Ordering::Relaxed);
self.track_permission_event(client_id, tool_name, false, Some(&reason));
return Ok(Permission::Deny(reason));
}
if let Some(tool_def) = self.rules.tools.get(tool_name) {
for required_scope in &tool_def.required_scopes {
if !context.scopes.contains(required_scope) {
let reason = format!("Missing required scope: {}", required_scope);
self.denied_count.fetch_add(1, Ordering::Relaxed);
self.track_permission_event(client_id, tool_name, false, Some(&reason));
return Ok(Permission::Deny(reason));
}
}
}
info!(
"Enhanced permission granted for {} to use {}",
client_id, tool_name
);
self.allowed_count.fetch_add(1, Ordering::Relaxed);
self.track_permission_event(client_id, tool_name, true, None);
Ok(Permission::Allow)
}
async fn get_allowed_tools(&self, client_id: &str) -> Result<Vec<String>> {
let permissions = self.get_client_permissions(client_id);
let risk_score = self
.risk_scores
.read()
.get(client_id)
.copied()
.unwrap_or(0.0);
let mut allowed = Vec::new();
for (tool_name, tool_def) in &self.rules.tools {
if self.rules.global_deny_list.contains(tool_name) {
continue;
}
if permissions.denied_tools.contains(tool_name) {
continue;
}
if risk_score > 0.8 && matches!(tool_def.category, super::ToolCategory::Administrative)
{
continue;
}
if !permissions.allowed_tools.is_empty()
&& !permissions.allowed_tools.contains(tool_name)
{
continue;
}
allowed.push(tool_name.clone());
}
Ok(allowed)
}
async fn update_permissions(
&self,
client_id: &str,
permissions: ClientPermissions,
) -> Result<()> {
let mut client_perms = self.client_permissions.write();
client_perms.insert(client_id.to_string(), permissions);
self.track_permission_event(client_id, "update_permissions", true, None);
debug!(
"Updated permissions for client {} with pattern tracking",
client_id
);
Ok(())
}
fn get_stats(&self) -> PermissionStats {
PermissionStats {
total_checks: self.total_checks.load(Ordering::Relaxed),
allowed: self.allowed_count.load(Ordering::Relaxed),
denied: self.denied_count.load(Ordering::Relaxed),
denied_by_reason: {
let mut reasons = HashMap::new();
reasons.insert(
"pattern_detection".to_string(),
self.pattern_matches.load(Ordering::Relaxed),
);
reasons.insert(
"enhanced_checks".to_string(),
self.total_checks.load(Ordering::Relaxed),
);
reasons
},
}
}
}