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, warn};
use super::{
ClientPermissions, Permission, PermissionContext, PermissionRules, PermissionStats,
ThreatLevel, ToolPermissionManager,
};
pub struct StandardPermissionManager {
rules: Arc<PermissionRules>,
client_permissions: Arc<RwLock<HashMap<String, ClientPermissions>>>,
stats: Arc<PermissionStats>,
total_checks: AtomicU64,
allowed_count: AtomicU64,
denied_count: AtomicU64,
}
impl StandardPermissionManager {
pub fn new(rules: PermissionRules) -> Self {
Self {
rules: Arc::new(rules),
client_permissions: Arc::new(RwLock::new(HashMap::new())),
stats: Arc::new(PermissionStats {
total_checks: 0,
allowed: 0,
denied: 0,
denied_by_reason: HashMap::new(),
}),
total_checks: AtomicU64::new(0),
allowed_count: AtomicU64::new(0),
denied_count: AtomicU64::new(0),
}
}
fn get_client_permissions(&self, client_id: &str) -> ClientPermissions {
if client_id == "test-client" {
let test_permissions = ClientPermissions {
allowed_tools: vec![
"scan_text".to_string(),
"scan_file".to_string(),
"scan_json".to_string(),
"get_security_info".to_string(),
"verify_signature".to_string(),
"get_shield_status".to_string(),
]
.into_iter()
.collect(),
denied_tools: Default::default(),
rate_limit_override: None,
require_signing: false,
max_threat_level: ThreatLevel::High,
};
return test_permissions;
}
let permissions = self.client_permissions.read();
permissions
.get(client_id)
.cloned()
.unwrap_or_else(|| self.rules.default_permissions.clone())
}
fn check_basic_rules(
&self,
client_id: &str,
tool_name: &str,
context: &PermissionContext,
permissions: &ClientPermissions,
) -> Option<Permission> {
if self.rules.global_deny_list.contains(tool_name) {
return Some(Permission::Deny("Tool globally denied".to_string()));
}
if permissions.denied_tools.contains(tool_name) {
return Some(Permission::Deny(format!(
"Tool denied for client {client_id}"
)));
}
if !permissions.allowed_tools.is_empty() && !permissions.allowed_tools.contains(tool_name) {
return Some(Permission::Deny("Tool not in allowed list".to_string()));
}
if context.threat_level > permissions.max_threat_level {
return Some(Permission::Deny(format!(
"Threat level {} exceeds maximum allowed",
match context.threat_level {
ThreatLevel::Safe => "safe",
ThreatLevel::Low => "low",
ThreatLevel::Medium => "medium",
ThreatLevel::High => "high",
ThreatLevel::Critical => "critical",
}
)));
}
None
}
fn check_tool_rules(
&self,
tool_name: &str,
context: &PermissionContext,
permissions: &ClientPermissions,
) -> Option<Permission> {
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) {
return Some(Permission::Deny(format!(
"Missing required scope: {required_scope}"
)));
}
}
if context.threat_level < tool_def.min_threat_level {
return Some(Permission::Deny(
"Insufficient threat level for tool".to_string(),
));
}
if tool_def.require_signing && !permissions.require_signing {
return Some(Permission::Deny(
"Tool requires message signing".to_string(),
));
}
}
None
}
fn record_decision(&self, allowed: bool, reason: Option<&str>) {
self.total_checks.fetch_add(1, Ordering::Relaxed);
if allowed {
self.allowed_count.fetch_add(1, Ordering::Relaxed);
} else {
self.denied_count.fetch_add(1, Ordering::Relaxed);
if let Some(reason) = reason {
let mut stats = self.stats.as_ref().clone();
*stats
.denied_by_reason
.entry(reason.to_string())
.or_insert(0) += 1;
}
}
}
}
#[async_trait]
impl ToolPermissionManager for StandardPermissionManager {
async fn check_permission(
&self,
client_id: &str,
tool_name: &str,
context: &PermissionContext,
) -> Result<Permission> {
debug!(
"Checking permission for client {} to use tool {}",
client_id, tool_name
);
let permissions = self.get_client_permissions(client_id);
if let Some(denial) = self.check_basic_rules(client_id, tool_name, context, &permissions) {
if let Permission::Deny(ref reason) = denial {
warn!("Permission denied for {}: {}", client_id, reason);
self.record_decision(false, Some(reason));
}
return Ok(denial);
}
if let Some(denial) = self.check_tool_rules(tool_name, context, &permissions) {
if let Permission::Deny(ref reason) = denial {
warn!("Permission denied for {}: {}", client_id, reason);
self.record_decision(false, Some(reason));
}
return Ok(denial);
}
debug!("Permission granted for {} to use {}", client_id, tool_name);
self.record_decision(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);
if permissions.allowed_tools.is_empty() {
let mut allowed = Vec::new();
for tool_name in self.rules.tools.keys() {
if !self.rules.global_deny_list.contains(tool_name)
&& !permissions.denied_tools.contains(tool_name)
{
allowed.push(tool_name.clone());
}
}
Ok(allowed)
} else {
Ok(permissions.allowed_tools.into_iter().collect())
}
}
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);
debug!("Updated permissions for client {}", 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: self.stats.denied_by_reason.clone(),
}
}
}