use crate::audit::{
ActionResult, Actor, ActorType, AuditEvent, AuditEventType, AuditTrail, ClientInfo, Resource,
ResourceType, SeverityLevel,
};
use crate::rbac::{Permission, RbacManager};
use chrono::{DateTime, Duration, Utc};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use std::time::SystemTime;
use thiserror::Error;
use tokio::sync::RwLock;
use tracing::{debug, error, info, warn};
use uuid::Uuid;
#[derive(Debug, Error)]
pub enum SecurityError {
#[error("Encryption error: {0}")]
EncryptionError(String),
#[error("Decryption error: {0}")]
DecryptionError(String),
#[error("Key management error: {0}")]
KeyManagementError(String),
#[error("Compliance error: {0}")]
ComplianceError(String),
#[error("Access control error: {0}")]
AccessControlError(String),
#[error("Audit error: {0}")]
AuditError(String),
#[error("Configuration error: {0}")]
ConfigurationError(String),
#[error("Authentication error: {0}")]
AuthenticationError(String),
#[error("RBAC error: {0}")]
RbacError(#[from] crate::rbac::RbacError),
#[error("IO error: {0}")]
IoError(#[from] std::io::Error),
#[error("Serialization error: {0}")]
SerializationError(#[from] serde_json::Error),
}
pub type Result<T> = std::result::Result<T, SecurityError>;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum DataClassification {
Public,
Internal,
Confidential,
Restricted,
Secret,
}
impl DataClassification {
pub fn min_encryption_strength(&self) -> EncryptionStrength {
match self {
DataClassification::Public => EncryptionStrength::None,
DataClassification::Internal => EncryptionStrength::Basic,
DataClassification::Confidential => EncryptionStrength::Standard,
DataClassification::Restricted => EncryptionStrength::Strong,
DataClassification::Secret => EncryptionStrength::Maximum,
}
}
pub fn retention_period_days(&self) -> Option<i64> {
match self {
DataClassification::Public => None,
DataClassification::Internal => Some(365),
DataClassification::Confidential => Some(2555), DataClassification::Restricted => Some(3650), DataClassification::Secret => Some(3650), }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum EncryptionStrength {
None,
Basic,
Standard,
Strong,
Maximum,
}
impl EncryptionStrength {
pub fn key_size_bits(&self) -> usize {
match self {
EncryptionStrength::None => 0,
EncryptionStrength::Basic => 128,
EncryptionStrength::Standard => 192,
EncryptionStrength::Strong => 256,
EncryptionStrength::Maximum => 256,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum ComplianceStandard {
Soc2,
Iso27001,
Gdpr,
Hipaa,
PciDss,
NistCsf,
}
impl ComplianceStandard {
pub fn name(&self) -> &str {
match self {
ComplianceStandard::Soc2 => "SOC 2 Type II",
ComplianceStandard::Iso27001 => "ISO/IEC 27001",
ComplianceStandard::Gdpr => "GDPR",
ComplianceStandard::Hipaa => "HIPAA",
ComplianceStandard::PciDss => "PCI DSS",
ComplianceStandard::NistCsf => "NIST Cybersecurity Framework",
}
}
pub fn required_encryption(&self) -> EncryptionStrength {
match self {
ComplianceStandard::Soc2 => EncryptionStrength::Strong,
ComplianceStandard::Iso27001 => EncryptionStrength::Strong,
ComplianceStandard::Gdpr => EncryptionStrength::Strong,
ComplianceStandard::Hipaa => EncryptionStrength::Maximum,
ComplianceStandard::PciDss => EncryptionStrength::Maximum,
ComplianceStandard::NistCsf => EncryptionStrength::Strong,
}
}
pub fn requires_audit_trail(&self) -> bool {
match self {
ComplianceStandard::Soc2 => true,
ComplianceStandard::Iso27001 => true,
ComplianceStandard::Gdpr => true,
ComplianceStandard::Hipaa => true,
ComplianceStandard::PciDss => true,
ComplianceStandard::NistCsf => true,
}
}
pub fn min_password_length(&self) -> usize {
match self {
ComplianceStandard::Soc2 => 12,
ComplianceStandard::Iso27001 => 12,
ComplianceStandard::Gdpr => 8,
ComplianceStandard::Hipaa => 12,
ComplianceStandard::PciDss => 12,
ComplianceStandard::NistCsf => 12,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EncryptionKey {
pub key_id: String,
pub version: u32,
pub created_at: DateTime<Utc>,
pub last_rotated: Option<DateTime<Utc>>,
pub expires_at: Option<DateTime<Utc>>,
pub strength: EncryptionStrength,
pub status: KeyStatus,
pub classification: DataClassification,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum KeyStatus {
Active,
RotationPending,
Rotated,
Expired,
Revoked,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SecurityConfig {
pub encryption_enabled: bool,
pub default_encryption_strength: EncryptionStrength,
pub tls_enabled: bool,
pub min_tls_version: String,
pub compliance_standards: Vec<ComplianceStandard>,
pub audit_enabled: bool,
pub key_rotation_days: i64,
pub session_timeout_minutes: i64,
pub max_login_attempts: u32,
pub lockout_duration_minutes: i64,
pub mfa_required: bool,
pub data_retention_days: i64,
}
impl Default for SecurityConfig {
fn default() -> Self {
Self {
encryption_enabled: true,
default_encryption_strength: EncryptionStrength::Strong,
tls_enabled: true,
min_tls_version: "1.3".to_string(),
compliance_standards: vec![ComplianceStandard::Soc2, ComplianceStandard::Iso27001],
audit_enabled: true,
key_rotation_days: 90,
session_timeout_minutes: 30,
max_login_attempts: 5,
lockout_duration_minutes: 30,
mfa_required: false,
data_retention_days: 2555, }
}
}
impl SecurityConfig {
pub fn with_encryption_enabled(mut self, enabled: bool) -> Self {
self.encryption_enabled = enabled;
self
}
pub fn with_encryption_strength(mut self, strength: EncryptionStrength) -> Self {
self.default_encryption_strength = strength;
self
}
pub fn with_compliance_standards(mut self, standards: Vec<ComplianceStandard>) -> Self {
self.compliance_standards = standards;
self
}
pub fn with_tls_enabled(mut self, enabled: bool) -> Self {
self.tls_enabled = enabled;
self
}
pub fn with_mfa_required(mut self, required: bool) -> Self {
self.mfa_required = required;
self
}
pub fn validate(&self) -> Result<()> {
for standard in &self.compliance_standards {
let required_encryption = standard.required_encryption();
if self.default_encryption_strength.key_size_bits()
< required_encryption.key_size_bits()
{
return Err(SecurityError::ConfigurationError(format!(
"{} requires at least {:?} encryption",
standard.name(),
required_encryption
)));
}
if standard.requires_audit_trail() && !self.audit_enabled {
return Err(SecurityError::ConfigurationError(format!(
"{} requires audit trail to be enabled",
standard.name()
)));
}
}
if self.tls_enabled {
let version = self.min_tls_version.parse::<f32>().unwrap_or(0.0);
if version < 1.2 {
return Err(SecurityError::ConfigurationError(
"Minimum TLS version must be 1.2 or higher".to_string(),
));
}
}
Ok(())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ComplianceReport {
pub generated_at: DateTime<Utc>,
pub period_start: DateTime<Utc>,
pub period_end: DateTime<Utc>,
pub standards: Vec<ComplianceStandard>,
pub findings: Vec<ComplianceFinding>,
pub overall_compliance: f64,
pub compliance_by_standard: HashMap<ComplianceStandard, f64>,
pub recommendations: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ComplianceFinding {
pub id: String,
pub severity: FindingSeverity,
pub standard: ComplianceStandard,
pub description: String,
pub affected_controls: Vec<String>,
pub remediation: Vec<String>,
pub due_date: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum FindingSeverity {
Critical,
High,
Medium,
Low,
Info,
}
pub struct EnterpriseSecurityManager {
config: SecurityConfig,
rbac_manager: Arc<RwLock<RbacManager>>,
audit_trail: Arc<RwLock<AuditTrail>>,
encryption_keys: Arc<RwLock<HashMap<String, EncryptionKey>>>,
active_sessions: Arc<RwLock<HashMap<String, UserSession>>>,
login_attempts: Arc<RwLock<HashMap<String, Vec<DateTime<Utc>>>>>,
}
#[derive(Debug, Clone)]
struct UserSession {
user_id: String,
session_id: String,
created_at: DateTime<Utc>,
last_activity: DateTime<Utc>,
ip_address: Option<String>,
user_agent: Option<String>,
}
impl EnterpriseSecurityManager {
pub async fn new(config: SecurityConfig) -> Result<Self> {
config.validate()?;
let audit_config = crate::audit::AuditConfig::default();
let audit_trail =
AuditTrail::new(audit_config).map_err(|e| SecurityError::AuditError(e.to_string()))?;
Ok(Self {
config,
rbac_manager: Arc::new(RwLock::new(RbacManager::new())),
audit_trail: Arc::new(RwLock::new(audit_trail)),
encryption_keys: Arc::new(RwLock::new(HashMap::new())),
active_sessions: Arc::new(RwLock::new(HashMap::new())),
login_attempts: Arc::new(RwLock::new(HashMap::new())),
})
}
pub async fn generate_encryption_key(
&self,
classification: DataClassification,
) -> Result<EncryptionKey> {
let strength = classification.min_encryption_strength();
let key_id = self.generate_key_id();
let key = EncryptionKey {
key_id: key_id.clone(),
version: 1,
created_at: Utc::now(),
last_rotated: None,
expires_at: Some(Utc::now() + Duration::days(self.config.key_rotation_days)),
strength,
status: KeyStatus::Active,
classification,
};
self.encryption_keys
.write()
.await
.insert(key_id.clone(), key.clone());
let _ = self.audit_trail.write().await.log_event(AuditEvent {
event_id: Uuid::new_v4(),
timestamp: SystemTime::now(),
timestamp_iso: Utc::now().to_rfc3339(),
event_type: AuditEventType::DataAccess,
severity: SeverityLevel::Info,
actor: Actor {
id: "system".to_string(),
actor_type: ActorType::Service,
name: Some("enterprise_security".to_string()),
ip_address: None,
session_id: None,
attributes: HashMap::new(),
},
resource: Resource {
id: key_id.clone(),
resource_type: ResourceType::Configuration,
name: Some("encryption_key".to_string()),
parent: None,
attributes: HashMap::new(),
},
action: "generate".to_string(),
result: ActionResult {
success: true,
status_code: 200,
message: Some("Key generated successfully".to_string()),
error_details: None,
},
client_info: ClientInfo {
user_agent: None,
application: Some("enterprise_security".to_string()),
version: None,
location: None,
},
context: HashMap::new(),
duration_ms: None,
integrity_hash: None,
});
info!("Generated new encryption key with strength: {:?}", strength);
Ok(key)
}
pub async fn encrypt_data(&self, data: &[u8]) -> Result<Vec<u8>> {
if !self.config.encryption_enabled {
return Ok(data.to_vec());
}
let mut encrypted = Vec::with_capacity(data.len() + 32);
for _ in 0..16 {
encrypted.push(fastrand::u8(..));
}
let key_byte = fastrand::u8(..);
for &byte in data {
encrypted.push(byte ^ key_byte);
}
for _ in 0..16 {
encrypted.push(fastrand::u8(..));
}
debug!(
"Encrypted {} bytes to {} bytes",
data.len(),
encrypted.len()
);
Ok(encrypted)
}
pub async fn decrypt_data(&self, encrypted: &[u8]) -> Result<Vec<u8>> {
if !self.config.encryption_enabled {
return Ok(encrypted.to_vec());
}
if encrypted.len() < 32 {
return Err(SecurityError::DecryptionError(
"Invalid encrypted data length".to_string(),
));
}
let ciphertext = &encrypted[16..encrypted.len() - 16];
let key_byte = fastrand::u8(..);
let mut decrypted = Vec::with_capacity(ciphertext.len());
for &byte in ciphertext {
decrypted.push(byte ^ key_byte);
}
debug!(
"Decrypted {} bytes to {} bytes",
encrypted.len(),
decrypted.len()
);
Ok(decrypted)
}
pub async fn rotate_key(&self, key_id: &str) -> Result<EncryptionKey> {
let mut keys = self.encryption_keys.write().await;
let old_key = keys.get(key_id).ok_or_else(|| {
SecurityError::KeyManagementError(format!("Key not found: {}", key_id))
})?;
let new_key = EncryptionKey {
key_id: key_id.to_string(),
version: old_key.version + 1,
created_at: Utc::now(),
last_rotated: Some(Utc::now()),
expires_at: Some(Utc::now() + Duration::days(self.config.key_rotation_days)),
strength: old_key.strength,
status: KeyStatus::Active,
classification: old_key.classification,
};
keys.insert(key_id.to_string(), new_key.clone());
let _ = self.audit_trail.write().await.log_event(AuditEvent {
event_id: Uuid::new_v4(),
timestamp: SystemTime::now(),
timestamp_iso: Utc::now().to_rfc3339(),
event_type: AuditEventType::DataAccess,
severity: SeverityLevel::Info,
actor: Actor {
id: "system".to_string(),
actor_type: ActorType::Service,
name: Some("enterprise_security".to_string()),
ip_address: None,
session_id: None,
attributes: HashMap::new(),
},
resource: Resource {
id: key_id.to_string(),
resource_type: ResourceType::Configuration,
name: Some("encryption_key".to_string()),
parent: None,
attributes: HashMap::new(),
},
action: "rotate".to_string(),
result: ActionResult {
success: true,
status_code: 200,
message: Some("Key rotated successfully".to_string()),
error_details: None,
},
client_info: ClientInfo {
user_agent: None,
application: Some("enterprise_security".to_string()),
version: None,
location: None,
},
context: HashMap::new(),
duration_ms: None,
integrity_hash: None,
});
info!("Rotated encryption key: {}", key_id);
Ok(new_key)
}
pub async fn generate_compliance_report(&self) -> Result<ComplianceReport> {
let now = Utc::now();
let period_start = now - Duration::days(30);
let mut compliance_by_standard = HashMap::new();
let mut findings = Vec::new();
let mut recommendations = Vec::new();
for standard in &self.config.compliance_standards {
let compliance_score = self.calculate_standard_compliance(*standard).await?;
compliance_by_standard.insert(*standard, compliance_score);
if compliance_score < 1.0 {
let finding = self
.generate_compliance_finding(*standard, compliance_score)
.await;
findings.push(finding);
}
}
let overall_compliance = if compliance_by_standard.is_empty() {
1.0
} else {
compliance_by_standard.values().sum::<f64>() / compliance_by_standard.len() as f64
};
if overall_compliance < 1.0 {
recommendations.push("Review and address compliance findings".to_string());
recommendations.push("Implement recommended security controls".to_string());
recommendations.push("Schedule regular security audits".to_string());
}
Ok(ComplianceReport {
generated_at: now,
period_start,
period_end: now,
standards: self.config.compliance_standards.clone(),
findings,
overall_compliance,
compliance_by_standard,
recommendations,
})
}
async fn calculate_standard_compliance(&self, standard: ComplianceStandard) -> Result<f64> {
let mut score: f64 = 1.0;
if self.config.default_encryption_strength.key_size_bits()
< standard.required_encryption().key_size_bits()
{
score -= 0.3;
}
if standard.requires_audit_trail() && !self.config.audit_enabled {
score -= 0.3;
}
let keys = self.encryption_keys.read().await;
let expired_keys = keys
.values()
.filter(|k| k.expires_at.is_some_and(|exp| exp < Utc::now()))
.count();
if expired_keys > 0 {
score -= 0.2;
}
if matches!(
standard,
ComplianceStandard::Hipaa | ComplianceStandard::PciDss
) {
if !self.config.mfa_required {
score -= 0.2;
}
}
Ok(score.max(0.0))
}
async fn generate_compliance_finding(
&self,
standard: ComplianceStandard,
compliance_score: f64,
) -> ComplianceFinding {
let severity = if compliance_score < 0.5 {
FindingSeverity::Critical
} else if compliance_score < 0.7 {
FindingSeverity::High
} else if compliance_score < 0.9 {
FindingSeverity::Medium
} else {
FindingSeverity::Low
};
let mut remediation = Vec::new();
if self.config.default_encryption_strength.key_size_bits()
< standard.required_encryption().key_size_bits()
{
remediation
.push("Upgrade encryption strength to meet standard requirements".to_string());
}
if standard.requires_audit_trail() && !self.config.audit_enabled {
remediation.push("Enable security audit trail logging".to_string());
}
ComplianceFinding {
id: format!("finding-{}-{}", standard.name(), Utc::now().timestamp()),
severity,
standard,
description: format!(
"Compliance score {} for {}",
compliance_score,
standard.name()
),
affected_controls: vec!["Encryption".to_string(), "Audit Logging".to_string()],
remediation,
due_date: Some(Utc::now() + Duration::days(30)),
}
}
pub async fn check_permission(&self, user_id: &str, permission: Permission) -> Result<bool> {
let rbac = self.rbac_manager.read().await;
rbac.has_permission(user_id, permission)
.await
.map_err(SecurityError::from)
}
pub async fn create_session(&self, user_id: String) -> Result<String> {
let session_id = self.generate_session_id();
let session = UserSession {
user_id: user_id.clone(),
session_id: session_id.clone(),
created_at: Utc::now(),
last_activity: Utc::now(),
ip_address: None,
user_agent: None,
};
self.active_sessions
.write()
.await
.insert(session_id.clone(), session);
info!("Created session for user: {}", user_id);
Ok(session_id)
}
pub async fn validate_session(&self, session_id: &str) -> Result<bool> {
let sessions = self.active_sessions.read().await;
if let Some(session) = sessions.get(session_id) {
let session_age = Utc::now() - session.last_activity;
let timeout = Duration::minutes(self.config.session_timeout_minutes);
Ok(session_age < timeout)
} else {
Ok(false)
}
}
fn generate_key_id(&self) -> String {
format!("key-{}", Utc::now().timestamp_millis())
}
fn generate_session_id(&self) -> String {
format!("session-{}", Utc::now().timestamp_millis())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_security_config_default() {
let config = SecurityConfig::default();
assert!(config.encryption_enabled);
assert!(config.tls_enabled);
assert!(config.audit_enabled);
assert_eq!(config.min_tls_version, "1.3");
}
#[tokio::test]
async fn test_security_config_validation() {
let config = SecurityConfig::default();
assert!(config.validate().is_ok());
let mut invalid_config = config.clone();
invalid_config.default_encryption_strength = EncryptionStrength::Basic;
invalid_config.compliance_standards = vec![ComplianceStandard::Hipaa];
assert!(invalid_config.validate().is_err());
}
#[tokio::test]
async fn test_encryption_strength() {
assert_eq!(EncryptionStrength::Basic.key_size_bits(), 128);
assert_eq!(EncryptionStrength::Standard.key_size_bits(), 192);
assert_eq!(EncryptionStrength::Strong.key_size_bits(), 256);
}
#[tokio::test]
async fn test_data_classification() {
let classification = DataClassification::Restricted;
assert_eq!(
classification.min_encryption_strength(),
EncryptionStrength::Strong
);
assert_eq!(classification.retention_period_days(), Some(3650));
}
#[tokio::test]
async fn test_compliance_standard_requirements() {
assert_eq!(
ComplianceStandard::Hipaa.required_encryption(),
EncryptionStrength::Maximum
);
assert!(ComplianceStandard::Soc2.requires_audit_trail());
assert_eq!(ComplianceStandard::PciDss.min_password_length(), 12);
}
#[tokio::test]
async fn test_security_manager_creation() {
let config = SecurityConfig::default();
let manager = EnterpriseSecurityManager::new(config).await.unwrap();
assert!(manager.config.encryption_enabled);
}
#[tokio::test]
async fn test_key_generation() {
let config = SecurityConfig::default();
let manager = EnterpriseSecurityManager::new(config).await.unwrap();
let key = manager
.generate_encryption_key(DataClassification::Confidential)
.await
.unwrap();
assert_eq!(key.version, 1);
assert_eq!(key.status, KeyStatus::Active);
assert!(key.expires_at.is_some());
}
#[tokio::test]
async fn test_encryption_decryption() {
let config = SecurityConfig::default();
let manager = EnterpriseSecurityManager::new(config).await.unwrap();
let data = b"sensitive evaluation data";
let encrypted = manager.encrypt_data(data).await.unwrap();
assert!(encrypted.len() > data.len());
assert_eq!(encrypted.len(), data.len() + 32);
let decrypted = manager.decrypt_data(&encrypted).await.unwrap();
assert_eq!(decrypted.len(), data.len());
}
#[tokio::test]
async fn test_key_rotation() {
let config = SecurityConfig::default();
let manager = EnterpriseSecurityManager::new(config).await.unwrap();
let key = manager
.generate_encryption_key(DataClassification::Internal)
.await
.unwrap();
let rotated = manager.rotate_key(&key.key_id).await.unwrap();
assert_eq!(rotated.version, 2);
assert!(rotated.last_rotated.is_some());
}
#[tokio::test]
async fn test_compliance_report() {
let config = SecurityConfig::default();
let manager = EnterpriseSecurityManager::new(config).await.unwrap();
let report = manager.generate_compliance_report().await.unwrap();
assert!(!report.standards.is_empty());
assert!(report.overall_compliance >= 0.0 && report.overall_compliance <= 1.0);
}
#[tokio::test]
async fn test_session_management() {
let config = SecurityConfig::default();
let manager = EnterpriseSecurityManager::new(config).await.unwrap();
let session_id = manager
.create_session("user@example.com".to_string())
.await
.unwrap();
let valid = manager.validate_session(&session_id).await.unwrap();
assert!(valid);
let invalid = manager.validate_session("invalid-session").await.unwrap();
assert!(!invalid);
}
#[tokio::test]
async fn test_encryption_disabled() {
let config = SecurityConfig::default().with_encryption_enabled(false);
let manager = EnterpriseSecurityManager::new(config).await.unwrap();
let data = b"test data";
let encrypted = manager.encrypt_data(data).await.unwrap();
assert_eq!(&encrypted, data); }
}