diff --git a/net-guardia/src/adapter/ebpf/geo_block.rs b/net-guardia/src/adapter/ebpf/geo_block.rs index 23ada33..1179ba3 100644 --- a/net-guardia/src/adapter/ebpf/geo_block.rs +++ b/net-guardia/src/adapter/ebpf/geo_block.rs @@ -36,7 +36,7 @@ impl GeoBlock { let v6_trie = LpmTrie::try_from(v6_map).map_err(EbpfError::MapOperationError)?; // todo read config from AppConfig, not db - let db_path = app_config.load().acl.geoip_db_name.clone(); + let db_path = app_config.load().acl.geoip_db_path.clone(); let reader = Reader::open_readfile(&db_path).map_err(|e| MiscError::GeoIPDatabaseError(db_path.clone(), e))?; let index = Self::build_index(&reader)?; @@ -50,7 +50,7 @@ impl GeoBlock { } pub fn unavailable(app_config: Arc>) -> Self { - let index = Reader::open_readfile(&app_config.load().acl.geoip_db_name) + let index = Reader::open_readfile(&app_config.load().acl.geoip_db_path) .ok() .and_then(|reader| Self::build_index(&reader).ok()) .unwrap_or(GeoIndex { diff --git a/net-guardia/src/adapter/ebpf/xsk_manager.rs b/net-guardia/src/adapter/ebpf/xsk_manager.rs index 8d2ce29..e5d8789 100644 --- a/net-guardia/src/adapter/ebpf/xsk_manager.rs +++ b/net-guardia/src/adapter/ebpf/xsk_manager.rs @@ -183,6 +183,9 @@ pub struct XskPair { drop_monitor: Option>, packet_buffer_size: usize, buffer_pool_capacity: usize, + completion_batch_size: usize, + rx_batch_size: usize, + tx_batch_size: usize, tx_packet_buf: Vec>, tx_frame_buf: Vec, } @@ -255,8 +258,11 @@ impl XskPair { drop_monitor, packet_buffer_size: config.packet_buffer_size, buffer_pool_capacity: config.buffer_pool_capacity, - tx_packet_buf: Vec::with_capacity(64), - tx_frame_buf: Vec::with_capacity(64), + completion_batch_size: config.xsk_completion_batch_size, + rx_batch_size: config.xsk_rx_batch_size, + tx_batch_size: config.xsk_tx_batch_size, + tx_packet_buf: Vec::with_capacity(config.xsk_tx_batch_size), + tx_frame_buf: Vec::with_capacity(config.xsk_tx_batch_size), }; Ok(xsk_pair) @@ -277,8 +283,8 @@ impl XskPair { let mut shutdown_rx = Some(shutdown_rx); let mut idle_count: u32 = 0; let mut buffer_pool = BufferPool::new(self.buffer_pool_capacity, self.packet_buffer_size); - let mut comp_descs = vec![FrameDesc::default(); 256]; - let mut rx_descs = vec![FrameDesc::default(); 64]; + let mut comp_descs = vec![FrameDesc::default(); self.completion_batch_size]; + let mut rx_descs = vec![FrameDesc::default(); self.rx_batch_size]; loop { if let Some(ref mut rx) = shutdown_rx { @@ -426,7 +432,7 @@ impl XskPair { self.tx_packet_buf.clear(); while let Ok(packet) = forward_rx.try_recv() { self.tx_packet_buf.push(packet); - if self.tx_packet_buf.len() >= 64 { + if self.tx_packet_buf.len() >= self.tx_batch_size { break; } } diff --git a/net-guardia/src/adapter/http/detection/model_upload.rs b/net-guardia/src/adapter/http/detection/model_upload.rs index ee5a6c9..545b622 100644 --- a/net-guardia/src/adapter/http/detection/model_upload.rs +++ b/net-guardia/src/adapter/http/detection/model_upload.rs @@ -584,26 +584,32 @@ async fn validate_and_promote(ctx: &PromoteContext<'_>) -> Result) -> Result, + target_manifest: PathBuf, + target_onnx: PathBuf, + target_sidecar: Option, + backup_dir: PathBuf, +} + +async fn promote_files_atomically(files: &PromoteFileSet) -> Result<(), PromoteError> { + fs::create_dir_all(&files.backup_dir) + .await + .map_err(|e| PromoteError::PromoteIo(format!("create promote backup dir: {e}")))?; + + let backup_manifest = backup_existing(&files.target_manifest, &files.backup_dir, "manifest.yaml").await?; + let backup_onnx = backup_existing(&files.target_onnx, &files.backup_dir, "model.onnx").await?; + let backup_sidecar = match &files.target_sidecar { + Some(target) => Some(backup_existing(target, &files.backup_dir, "sidecar").await?), + None => None, + }; + + let result = async { + move_file(&files.staged_onnx, &files.target_onnx, "rename onnx into models/").await?; + if let (Some(src), Some(dst)) = (&files.staged_sidecar, &files.target_sidecar) { + move_file(src, dst, "rename sidecar into models/").await?; + } + move_file( + &files.staging_manifest, + &files.target_manifest, + "rename manifest into models/", + ) + .await + } + .await; + + match result { + Ok(()) => { + let _ = fs::remove_dir_all(&files.backup_dir).await; + Ok(()) + } + Err(err) => { + rollback_promote(files, backup_manifest, backup_onnx, backup_sidecar).await; + let _ = fs::remove_dir_all(&files.backup_dir).await; + Err(err) + } + } +} + +async fn backup_existing(target: &Path, backup_dir: &Path, backup_name: &str) -> Result, PromoteError> { + if !target + .try_exists() + .map_err(|e| PromoteError::PromoteIo(format!("check existing target {}: {e}", target.display())))? + { + return Ok(None); + } + let backup = backup_dir.join(backup_name); + fs::rename(target, &backup) + .await + .map_err(|e| PromoteError::PromoteIo(format!("backup existing target {}: {e}", target.display())))?; + Ok(Some(backup)) +} + +async fn move_file(src: &Path, dst: &Path, op: &str) -> Result<(), PromoteError> { + fs::rename(src, dst) + .await + .map_err(|e| PromoteError::PromoteIo(format!("{op}: {e}"))) +} + +async fn rollback_promote( + files: &PromoteFileSet, + backup_manifest: Option, + backup_onnx: Option, + backup_sidecar: Option>, +) { + remove_if_exists(&files.target_manifest).await; + remove_if_exists(&files.target_onnx).await; + if let Some(target) = &files.target_sidecar { + remove_if_exists(target).await; + } + restore_backup(backup_manifest, &files.target_manifest).await; + restore_backup(backup_onnx, &files.target_onnx).await; + if let (Some(backup), Some(target)) = (backup_sidecar.flatten(), &files.target_sidecar) { + restore_backup(Some(backup), target).await; + } +} + +async fn remove_if_exists(path: &Path) { + if let Ok(true) = path.try_exists() { + let _ = fs::remove_file(path).await; + } +} + +async fn restore_backup(backup: Option, target: &Path) { + if let Some(backup) = backup { + let _ = fs::rename(backup, target).await; + } +} + /// Metadata surfaced back to the client when the promote succeeds. #[derive(Debug)] struct PromoteReport { @@ -858,6 +963,48 @@ mod tests { assert_eq!(resp.status().as_u16(), 500); } + #[tokio::test] + async fn promote_files_rolls_back_active_files_when_sidecar_move_fails() { + let tmp = std::env::temp_dir().join(format!("nguardia-promote-rollback-{}", Uuid::new_v4())); + let staging = tmp.join("staging"); + let models = tmp.join("models"); + let backup = models.join(".promote-backup-test"); + fs::create_dir_all(&staging).await.unwrap(); + fs::create_dir_all(&models).await.unwrap(); + + let staging_manifest = staging.join(MANIFEST_FILENAME); + let staged_onnx = staging.join("model.onnx"); + let missing_sidecar = staging.join("missing-scaler.json"); + let target_manifest = models.join(MANIFEST_FILENAME); + let target_onnx = models.join("model.onnx"); + let target_sidecar = models.join("scaler.json"); + + fs::write(&staging_manifest, b"new manifest").await.unwrap(); + fs::write(&staged_onnx, b"new onnx").await.unwrap(); + fs::write(&target_manifest, b"old manifest").await.unwrap(); + fs::write(&target_onnx, b"old onnx").await.unwrap(); + fs::write(&target_sidecar, b"old sidecar").await.unwrap(); + + let err = promote_files_atomically(&PromoteFileSet { + staging_manifest: staging_manifest.clone(), + staged_onnx, + staged_sidecar: Some(missing_sidecar), + target_manifest: target_manifest.clone(), + target_onnx: target_onnx.clone(), + target_sidecar: Some(target_sidecar.clone()), + backup_dir: backup, + }) + .await + .expect_err("missing sidecar should fail promote"); + + assert!(matches!(err, PromoteError::PromoteIo(_))); + assert_eq!(fs::read(&target_manifest).await.unwrap(), b"old manifest"); + assert_eq!(fs::read(&target_onnx).await.unwrap(), b"old onnx"); + assert_eq!(fs::read(&target_sidecar).await.unwrap(), b"old sidecar"); + + fs::remove_dir_all(&tmp).await.ok(); + } + #[test] fn upload_error_scaler_too_large_maps_to_413_and_echoes_cap() { let resp = UploadError::ScalerTooLarge(1234).into_response(); diff --git a/net-guardia/src/adapter/http/logs.rs b/net-guardia/src/adapter/http/logs.rs index b693785..2544b9f 100644 --- a/net-guardia/src/adapter/http/logs.rs +++ b/net-guardia/src/adapter/http/logs.rs @@ -1,6 +1,6 @@ use std::fs; use std::io::ErrorKind; -use std::path::Path; +use std::path::PathBuf; use std::time::UNIX_EPOCH; use actix_web::{HttpResponse, Scope, web}; @@ -10,9 +10,6 @@ use serde::{Deserialize, Serialize}; use crate::domain::common::config::AppConfig; use crate::infrastructure::log_buffer::{self, LogBuffer, LogEntry}; -/// Hardcoded log directory — not configurable via API to prevent directory traversal. -const LOG_DIR: &str = "logs"; - /// Validate log filename: only alphanumeric, dots, underscores, hyphens. /// Prevents path traversal. fn is_valid_log_filename(name: &str) -> bool { @@ -87,9 +84,9 @@ struct LogFileEntry { modified: Option, } -async fn list_logs() -> HttpResponse { - let log_dir = LOG_DIR; - let entries = match fs::read_dir(log_dir) { +async fn list_logs(app_config: web::Data>) -> HttpResponse { + let log_dir = app_config.load().system.log_dir.clone(); + let entries = match fs::read_dir(&log_dir) { Ok(dir) => dir .filter_map(|e| e.ok()) .filter_map(|e| { @@ -117,7 +114,9 @@ async fn list_logs() -> HttpResponse { } async fn download_log(path: web::Path, app_config: web::Data>) -> HttpResponse { - let max_download_size = app_config.load().observability.log_max_download_size; + let config = app_config.load(); + let max_download_size = config.observability.log_max_download_size; + let log_dir = PathBuf::from(&config.system.log_dir); let filename = path.into_inner(); if !is_valid_log_filename(&filename) { @@ -126,7 +125,7 @@ async fn download_log(path: web::Path, app_config: web::Data, app_config: web::Data Option { "rate_limit" } else if path.starts_with("/api/system/") { "system" - } else if path.starts_with("/api/api-keys/") { + } else if path == "/api/api-keys" || path.starts_with("/api/api-keys/") { return Some("api_keys:admin".to_string()); } else if path.contains("/soar/blocks/") && path.ends_with("/unblock") { return Some("access_control:write".to_string()); @@ -217,3 +217,30 @@ where }) } } + +#[cfg(test)] +mod tests { + use actix_web::http::Method; + + use super::required_permission; + + #[test] + fn api_key_collection_requires_admin_permission() { + assert_eq!( + required_permission("/api/api-keys", &Method::GET), + Some("api_keys:admin".to_string()) + ); + } + + #[test] + fn api_key_subroutes_require_admin_permission() { + assert_eq!( + required_permission("/api/api-keys/generate", &Method::POST), + Some("api_keys:admin".to_string()) + ); + assert_eq!( + required_permission("/api/api-keys/1", &Method::DELETE), + Some("api_keys:admin".to_string()) + ); + } +} diff --git a/net-guardia/src/adapter/http/response/soar.rs b/net-guardia/src/adapter/http/response/soar.rs index 759cf23..598956a 100644 --- a/net-guardia/src/adapter/http/response/soar.rs +++ b/net-guardia/src/adapter/http/response/soar.rs @@ -162,8 +162,12 @@ async fn manual_unblock(_auth: AuthClaims, svc: web::Data, path ok_or_error(svc.manual_unblock(path.into_inner()).await) } -async fn list_executions(_auth: AuthClaims, svc: web::Data) -> HttpResponse { - ok_json_or_error(svc.list_executions(100).await) +async fn list_executions( + _auth: AuthClaims, + svc: web::Data, + app_config: web::Data>, +) -> HttpResponse { + ok_json_or_error(svc.list_executions(app_config.load().soar.execution_list_limit).await) } async fn list_whitelist(_auth: AuthClaims, svc: web::Data) -> HttpResponse { diff --git a/net-guardia/src/adapter/persistence/api_key.rs b/net-guardia/src/adapter/persistence/api_key.rs index 56c6e19..a8770ea 100644 --- a/net-guardia/src/adapter/persistence/api_key.rs +++ b/net-guardia/src/adapter/persistence/api_key.rs @@ -139,3 +139,36 @@ impl ApiKeyRepo for Database { self.delete_api_key(id).await } } + +#[cfg(test)] +mod tests { + use super::Database; + + #[tokio::test] + async fn validate_full_access_api_key_grants_admin_permissions() { + let db = Database::new(":memory:").await.expect("test db"); + let raw_key = "ng-test-full-access"; + let digest = db.hmac_api_key(raw_key); + + db.insert_api_key(&digest, "automation", "full_access").await.unwrap(); + + let claims = db.validate_api_key(raw_key).await.unwrap().expect("claims"); + assert!(claims.permissions.contains(&"api_keys:admin".to_string())); + assert!(claims.permissions.contains(&"system:admin".to_string())); + assert!(claims.permissions.contains(&"users:admin".to_string())); + } + + #[tokio::test] + async fn validate_read_write_api_key_does_not_grant_admin_permissions() { + let db = Database::new(":memory:").await.expect("test db"); + let raw_key = "ng-test-read-write"; + let digest = db.hmac_api_key(raw_key); + + db.insert_api_key(&digest, "automation", "read_write").await.unwrap(); + + let claims = db.validate_api_key(raw_key).await.unwrap().expect("claims"); + assert!(!claims.permissions.contains(&"api_keys:admin".to_string())); + assert!(!claims.permissions.contains(&"system:admin".to_string())); + assert!(!claims.permissions.contains(&"users:admin".to_string())); + } +} diff --git a/net-guardia/src/adapter/persistence/enforcement.rs b/net-guardia/src/adapter/persistence/enforcement.rs index c06a5d4..a16c5a1 100644 --- a/net-guardia/src/adapter/persistence/enforcement.rs +++ b/net-guardia/src/adapter/persistence/enforcement.rs @@ -33,24 +33,32 @@ impl Database { .await } - pub async fn insert_dns_domain(&self, domain: &str) -> Result<(), Error> { - let domain = domain.to_string(); + pub async fn insert_dns_domains(&self, domains: &[String]) -> Result<(), Error> { + let domains = domains.to_vec(); self.pool - .conn_and_then(move |conn| { - conn.execute( - "INSERT OR IGNORE INTO dns_blacklist (domain) VALUES (?1)", - params![domain], - )?; + .conn_mut_and_then(move |conn| { + let tx = conn.transaction()?; + for domain in domains { + tx.execute( + "INSERT OR IGNORE INTO dns_blacklist (domain) VALUES (?1)", + params![domain], + )?; + } + tx.commit()?; Ok(()) }) .await } - pub async fn delete_dns_domain(&self, domain: &str) -> Result<(), Error> { - let domain = domain.to_string(); + pub async fn delete_dns_domains(&self, domains: &[String]) -> Result<(), Error> { + let domains = domains.to_vec(); self.pool - .conn_and_then(move |conn| { - conn.execute("DELETE FROM dns_blacklist WHERE domain = ?1", params![domain])?; + .conn_mut_and_then(move |conn| { + let tx = conn.transaction()?; + for domain in domains { + tx.execute("DELETE FROM dns_blacklist WHERE domain = ?1", params![domain])?; + } + tx.commit()?; Ok(()) }) .await @@ -117,12 +125,12 @@ impl EnforcementRepo for Database { self.set_rate_limit(key, value).await } - async fn insert_dns_domain(&self, domain: &str) -> Result<(), Error> { - self.insert_dns_domain(domain).await + async fn insert_dns_domains(&self, domains: &[String]) -> Result<(), Error> { + self.insert_dns_domains(domains).await } - async fn delete_dns_domain(&self, domain: &str) -> Result<(), Error> { - self.delete_dns_domain(domain).await + async fn delete_dns_domains(&self, domains: &[String]) -> Result<(), Error> { + self.delete_dns_domains(domains).await } async fn insert_geo_country(&self, code: &str) -> Result<(), Error> { diff --git a/net-guardia/src/adapter/persistence/mod.rs b/net-guardia/src/adapter/persistence/mod.rs index b1302ba..cd55ce1 100644 --- a/net-guardia/src/adapter/persistence/mod.rs +++ b/net-guardia/src/adapter/persistence/mod.rs @@ -317,6 +317,32 @@ impl Database { CREATE INDEX IF NOT EXISTS idx_audit_log_action ON audit_log(action, id DESC); + CREATE INDEX IF NOT EXISTS idx_soar_block_active_expires + ON soar_block_rules(expires_at) + WHERE unblocked_at IS NULL; + + CREATE INDEX IF NOT EXISTS idx_soar_block_created_at + ON soar_block_rules(created_at); + + CREATE INDEX IF NOT EXISTS idx_soar_block_unblocked_at + ON soar_block_rules(unblocked_at) + WHERE unblocked_at IS NOT NULL; + + CREATE INDEX IF NOT EXISTS idx_soar_block_source_created + ON soar_block_rules(source_ip, created_at); + + CREATE INDEX IF NOT EXISTS idx_soar_executions_executed_at + ON soar_executions(executed_at DESC); + + CREATE INDEX IF NOT EXISTS idx_soar_executions_trigger_executed + ON soar_executions(trigger_event, executed_at); + + CREATE INDEX IF NOT EXISTS idx_user_group_members_group_id + ON user_group_members(group_id, user_id); + + CREATE INDEX IF NOT EXISTS idx_api_keys_key_hash + ON api_keys(key_hash); + CREATE TRIGGER IF NOT EXISTS audit_log_no_update BEFORE UPDATE ON audit_log BEGIN SELECT RAISE(ABORT, 'audit_log is append-only (WORM)'); @@ -361,7 +387,7 @@ impl Database { )?; } - Ok(()) + Ok::<(), Error>(()) }) .await } @@ -544,4 +570,34 @@ mod tests { let _ = std::fs::remove_file(path.with_extension("db-wal")); let _ = std::fs::remove_file(path.with_extension("db-shm")); } + + #[tokio::test] + async fn create_tables_installs_query_shape_indexes() { + let db = test_db().await; + let expected = [ + ("soar_block_rules", "idx_soar_block_active_expires"), + ("soar_block_rules", "idx_soar_block_created_at"), + ("soar_block_rules", "idx_soar_block_unblocked_at"), + ("soar_block_rules", "idx_soar_block_source_created"), + ("soar_executions", "idx_soar_executions_executed_at"), + ("soar_executions", "idx_soar_executions_trigger_executed"), + ("user_group_members", "idx_user_group_members_group_id"), + ("api_keys", "idx_api_keys_key_hash"), + ]; + + db.pool + .conn_and_then(move |conn| { + for (table, index) in expected { + let count: i64 = conn.query_row( + "SELECT COUNT(*) FROM sqlite_master WHERE type = 'index' AND tbl_name = ?1 AND name = ?2", + params![table, index], + |row| row.get(0), + )?; + assert_eq!(count, 1, "missing index {index} on {table}"); + } + Ok::<(), Error>(()) + }) + .await + .unwrap(); + } } diff --git a/net-guardia/src/core/common/statistics.rs b/net-guardia/src/core/common/statistics.rs index 4df241b..033c37a 100644 --- a/net-guardia/src/core/common/statistics.rs +++ b/net-guardia/src/core/common/statistics.rs @@ -4,7 +4,7 @@ use std::time::{SystemTime, UNIX_EPOCH}; use crate::core::inference::engine::Engine; use crate::domain::data_plane::direction::Direction; use crate::domain::data_plane::flow_stats::{ - FlowPushPayload, FlowStatsEntry, FlowSubscription, FlowSummary, StatsSummary, + FlowPushPayload, FlowStatsEntry, FlowStatsLimits, FlowSubscription, FlowSummary, StatsSummary, }; use crate::domain::detection::flow_tracker::FlowData; @@ -32,17 +32,28 @@ impl From<&FlowData> for FlowStatsEntry { pub struct FlowStatistics { engine: Arc, + limits: FlowStatsLimits, } impl FlowStatistics { - pub fn new(engine: Arc) -> Self { - Self { engine } + pub fn new(engine: Arc, limits: FlowStatsLimits) -> Self { + Self { engine, limits } } pub fn get_all_flows(&self) -> Vec { + self.collect_flows(Some(self.limits.result_limit(None))) + } + + fn collect_flows(&self, limit: Option) -> Vec { let mut entries = Vec::new(); for tracker in self.engine.trackers() { entries.extend(tracker.get_flow_stats(|flow| FlowStatsEntry::from(flow))); + if let Some(limit) = limit + && entries.len() >= limit + { + entries.truncate(limit); + break; + } } entries } @@ -53,7 +64,7 @@ impl FlowStatistics { .map(|d| d.as_micros() as u64) .unwrap_or(0); - let mut flows = self.get_all_flows(); + let mut flows = self.collect_flows(None); if let Some(dir) = &sub.direction { flows.retain(|f| &f.direction == dir); @@ -65,10 +76,7 @@ impl FlowStatistics { } flows.sort_by_key(|f| std::cmp::Reverse(f.fwd_bytes + f.bwd_bytes)); - - if let Some(n) = sub.top_n { - flows.truncate(n.min(10000)); - } + flows.truncate(self.limits.result_limit(sub.top_n)); flows } @@ -120,7 +128,7 @@ impl FlowStatistics { } pub fn get_summary(&self) -> StatsSummary { - let flows = self.get_all_flows(); + let flows = self.collect_flows(None); let total_flows = flows.len(); let total_bytes: u64 = flows.iter().map(|f| f.fwd_bytes + f.bwd_bytes).sum(); let total_packets: usize = flows.iter().map(|f| f.fwd_packets + f.bwd_packets).sum(); diff --git a/net-guardia/src/core/data_plane/dns_filter.rs b/net-guardia/src/core/data_plane/dns_filter.rs index 8841297..5481fa3 100644 --- a/net-guardia/src/core/data_plane/dns_filter.rs +++ b/net-guardia/src/core/data_plane/dns_filter.rs @@ -25,6 +25,10 @@ impl DnsFilter { Ok(()) } + pub fn validate_domain(&self, domain: &str) -> Result<(), Error> { + domain_to_wire_format(domain).map(|_| ()) + } + pub fn remove_domain(&self, domain: &str) -> Result<(), Error> { let name = domain_to_wire_format(domain)?; self.blacklist.remove(&name); @@ -185,6 +189,9 @@ impl DnsFilter { } impl DnsFilterPort for DnsFilter { + fn validate_domain(&self, domain: &str) -> Result<(), Error> { + self.validate_domain(domain) + } fn add_domain(&self, domain: &str) -> Result<(), Error> { self.add_domain(domain) } diff --git a/net-guardia/src/core/data_plane/dns_filter_service.rs b/net-guardia/src/core/data_plane/dns_filter_service.rs index 1895692..8fe63b5 100644 --- a/net-guardia/src/core/data_plane/dns_filter_service.rs +++ b/net-guardia/src/core/data_plane/dns_filter_service.rs @@ -5,19 +5,24 @@ use arc_swap::ArcSwap; use crate::domain::common::config::AppConfig; use crate::domain::common::error::Error; use crate::domain::common::error::misc::MiscError; -use crate::interface::app_repo::AppRepo; use crate::interface::dns_filter_api::DnsFilterPort; +use crate::interface::enforcement::EnforcementRepo; /// Domain service that coordinates DNS filter changes between DB and in-memory service. -/// Write order: eBPF/in-memory first, then DB — if eBPF fails, DB remains clean. +/// Domains are validated up front, runtime state is changed first, and DB batch +/// failure rolls runtime back so persisted and live policy do not drift. pub struct DnsFilterService { - db: Arc, + db: Arc, dns_filter: Arc, config: Arc>, } impl DnsFilterService { - pub fn new(db: Arc, dns_filter: Arc, config: Arc>) -> Self { + pub fn new( + db: Arc, + dns_filter: Arc, + config: Arc>, + ) -> Self { Self { db, dns_filter, config } } @@ -33,26 +38,151 @@ impl DnsFilterService { max_domains )))?; } - // eBPF first + self.validate_domains(domains)?; + let mut applied: Vec<&String> = Vec::new(); for domain in domains { - self.dns_filter.add_domain(domain)?; + if let Err(err) = self.dns_filter.add_domain(domain) { + self.rollback_added(&applied); + return Err(err); + } + applied.push(domain); } - // Then DB - for domain in domains { - self.db.insert_dns_domain(domain).await?; + if let Err(err) = self.db.insert_dns_domains(domains).await { + self.rollback_added(&applied); + return Err(err); } Ok(domains.len()) } pub async fn remove_domains(&self, domains: &[String]) -> Result { - // eBPF first + self.validate_domains(domains)?; + let mut applied: Vec<&String> = Vec::new(); for domain in domains { - self.dns_filter.remove_domain(domain)?; + if let Err(err) = self.dns_filter.remove_domain(domain) { + self.rollback_removed(&applied); + return Err(err); + } + applied.push(domain); } - // Then DB - for domain in domains { - self.db.delete_dns_domain(domain).await?; + if let Err(err) = self.db.delete_dns_domains(domains).await { + self.rollback_removed(&applied); + return Err(err); } Ok(domains.len()) } + + fn validate_domains(&self, domains: &[String]) -> Result<(), Error> { + for domain in domains { + self.dns_filter.validate_domain(domain)?; + } + Ok(()) + } + + fn rollback_added(&self, domains: &[&String]) { + for domain in domains { + let _ = self.dns_filter.remove_domain(domain); + } + } + + fn rollback_removed(&self, domains: &[&String]) { + for domain in domains { + let _ = self.dns_filter.add_domain(domain); + } + } +} + +#[cfg(test)] +mod tests { + use std::sync::Mutex; + + use async_trait::async_trait; + + use super::*; + use crate::adapter::persistence::Database; + use crate::core::data_plane::dns_filter::DnsFilter; + + struct FailingDnsRepo { + fail_insert: bool, + fail_delete: bool, + domains: Mutex>, + } + + impl FailingDnsRepo { + fn new(fail_insert: bool, fail_delete: bool) -> Self { + Self { + fail_insert, + fail_delete, + domains: Mutex::new(Vec::new()), + } + } + } + + #[async_trait] + impl EnforcementRepo for FailingDnsRepo { + async fn set_rate_limit(&self, _key: &str, _value: u64) -> Result<(), Error> { + Ok(()) + } + + async fn insert_dns_domains(&self, domains: &[String]) -> Result<(), Error> { + if self.fail_insert { + Err(MiscError::ValidationError("forced insert failure".to_string()))?; + } + self.domains.lock().unwrap().extend_from_slice(domains); + Ok(()) + } + + async fn delete_dns_domains(&self, domains: &[String]) -> Result<(), Error> { + if self.fail_delete { + Err(MiscError::ValidationError("forced delete failure".to_string()))?; + } + self.domains.lock().unwrap().retain(|d| !domains.contains(d)); + Ok(()) + } + + async fn insert_geo_country(&self, _code: &str) -> Result<(), Error> { + Ok(()) + } + + async fn delete_geo_country(&self, _code: &str) -> Result<(), Error> { + Ok(()) + } + } + + async fn test_config() -> Arc> { + let db = Database::new(":memory:").await.expect("test db"); + Arc::new(ArcSwap::from_pointee( + AppConfig::from_config_repo(&db).await.expect("test config"), + )) + } + + #[tokio::test] + async fn add_domains_rolls_back_runtime_when_db_batch_fails() { + let filter = Arc::new(DnsFilter::new()); + let service = DnsFilterService::new( + Arc::new(FailingDnsRepo::new(true, false)), + filter.clone(), + test_config().await, + ); + + let result = service.add_domains(&["example.com".to_string()]).await; + + assert!(result.is_err()); + assert!(filter.list_domains().is_empty()); + } + + #[tokio::test] + async fn remove_domains_restores_runtime_when_db_batch_fails() { + let filter = Arc::new(DnsFilter::new()); + filter.add_domain("example.com").unwrap(); + let service = DnsFilterService::new( + Arc::new(FailingDnsRepo::new(false, true)), + filter.clone(), + test_config().await, + ); + + let result = service.remove_domains(&["example.com".to_string()]).await; + + assert!(result.is_err()); + assert!(filter.list_domains().contains(&"example.com".to_string())); + } } diff --git a/net-guardia/src/core/detection/beaconing.rs b/net-guardia/src/core/detection/beaconing.rs index a64b095..92e2236 100644 --- a/net-guardia/src/core/detection/beaconing.rs +++ b/net-guardia/src/core/detection/beaconing.rs @@ -34,6 +34,7 @@ impl BeaconingDetector { beaconing.min_observations, beaconing.cv_threshold, beaconing.max_cache_entries, + beaconing.max_timestamps_per_flow, beaconing.expiry_secs, beaconing.alert_cooldown_secs, ), @@ -78,8 +79,6 @@ impl BeaconingDetector { type FlowTuple = (String, String, u16); -const MAX_TIMESTAMPS_PER_FLOW: usize = 100; - struct CachedFlow { timestamps: Vec, last_alerted: Option, @@ -90,6 +89,7 @@ pub struct BeaconingState { min_observations: usize, cv_threshold: f64, max_cache_entries: usize, + max_timestamps_per_flow: usize, expiry_secs: u64, alert_cooldown_secs: u64, } @@ -99,6 +99,7 @@ impl BeaconingState { min_observations: usize, cv_threshold: f64, max_cache_entries: usize, + max_timestamps_per_flow: usize, expiry_secs: u64, alert_cooldown_secs: u64, ) -> Self { @@ -107,6 +108,7 @@ impl BeaconingState { min_observations, cv_threshold, max_cache_entries, + max_timestamps_per_flow, expiry_secs, alert_cooldown_secs, } @@ -123,8 +125,8 @@ impl BeaconingState { entry.timestamps.push(now); - if entry.timestamps.len() > MAX_TIMESTAMPS_PER_FLOW { - let excess = entry.timestamps.len() - MAX_TIMESTAMPS_PER_FLOW; + if entry.timestamps.len() > self.max_timestamps_per_flow { + let excess = entry.timestamps.len() - self.max_timestamps_per_flow; entry.timestamps.drain(..excess); } } @@ -280,7 +282,7 @@ mod tests { #[test] fn beaconing_state_detects_periodic_flows() { - let state = BeaconingState::new(5, 0.3, 50_000, 3600, 120); + let state = BeaconingState::new(5, 0.3, 50_000, 100, 3600, 120); let base = Instant::now(); let key = ("10.0.0.1".to_string(), "1.2.3.4".to_string(), 443_u16); state.flow_cache.insert( @@ -295,4 +297,33 @@ mod tests { assert_eq!(events[0].source, DetectionSource::Beaconing); assert_eq!(events[0].attack_type, CanonicalAttackType::C2Beacon.as_str()); } + + #[test] + fn record_flow_respects_configured_timestamp_cap() { + let state = BeaconingState::new(1, 0.3, 50_000, 3, 3600, 120); + let alert = crate::domain::detection::ml_detection::AlertMessage { + timestamp: 0, + flow_key: "10.0.0.1:12345-1.2.3.4:443".to_string(), + src_ip: "10.0.0.1".to_string(), + dst_ip: "1.2.3.4".to_string(), + src_port: 12345, + dst_port: 443, + protocol: 6, + is_attack: true, + attack_type: Some("c2_beacon".to_string()), + confidence: 0.9, + packet_count: 10, + flow_duration_us: 1000, + ae_score: 0.0, + anomaly_score: 0.0, + c2_score: 0.0, + }; + + for _ in 0..5 { + state.record_flow(&alert); + } + + let key = ("10.0.0.1".to_string(), "1.2.3.4".to_string(), 443_u16); + assert_eq!(state.flow_cache.get(&key).unwrap().timestamps.len(), 3); + } } diff --git a/net-guardia/src/core/detection/orchestrator.rs b/net-guardia/src/core/detection/orchestrator.rs index c8f1118..e6687c2 100644 --- a/net-guardia/src/core/detection/orchestrator.rs +++ b/net-guardia/src/core/detection/orchestrator.rs @@ -78,16 +78,17 @@ impl DetectionOrchestrator { ) -> Self { let cfg = app_config.load(); let fusion = &cfg.detection.fusion; - let max_dedup = NonZero::new(fusion.max_dedup_entries.max(1)).unwrap_or(NonZero::::MIN); + let source_count_max_entries = nonzero_cache_size(fusion.source_count_max_entries); + let repeat_tracker_max_entries = nonzero_cache_size(fusion.repeat_tracker_max_entries); + let max_dedup = nonzero_cache_size(fusion.max_dedup_entries); Self { rx, threat_tx, audit_tx, geoip, metrics, - // SAFETY: NonZero::new on non-zero literals. - src_ip_counts: LruCache::new(NonZero::new(10_000).unwrap()), - repeat_tracker: LruCache::new(NonZero::new(5_000).unwrap()), + src_ip_counts: LruCache::new(source_count_max_entries), + repeat_tracker: LruCache::new(repeat_tracker_max_entries), dedup: LruCache::new(max_dedup), dedup_window: Duration::from_secs(fusion.dedup_window_secs), repeat_offender_window: Duration::from_secs(fusion.repeat_offender_window_secs), @@ -355,6 +356,10 @@ impl DetectionOrchestrator { } } +fn nonzero_cache_size(value: usize) -> NonZero { + NonZero::new(value.max(1)).unwrap_or(NonZero::::MIN) +} + pub async fn bridge_ml_to_detection(mut rx: broadcast::Receiver, tx: mpsc::Sender) { log!(DetectionLog::MlBridgeStarted); diff --git a/net-guardia/src/core/identity/auth_service.rs b/net-guardia/src/core/identity/auth_service.rs index e8ca937..5727a1d 100644 --- a/net-guardia/src/core/identity/auth_service.rs +++ b/net-guardia/src/core/identity/auth_service.rs @@ -34,6 +34,7 @@ pub struct UserProfile { pub groups: Vec, } +#[derive(Debug)] pub enum LoginError { Locked { retry_after_secs: u64 }, InvalidCredentials, @@ -164,3 +165,62 @@ impl AuthService { } } } + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use super::*; + use crate::adapter::http::jwt::JwtService; + use crate::adapter::persistence::Database; + use crate::domain::identity::password; + use crate::infrastructure::secret_store::SecretStore; + use crate::interface::secret_store::SecretStorePort; + + async fn auth_fixture() -> (Arc, Arc, AuthService) { + let db = Arc::new(Database::new(":memory:").await.expect("test db")); + let secrets: Arc = Arc::new(SecretStore::new(db.clone())); + let jwt = Arc::new(JwtService::new(&secrets, 24).expect("jwt")); + let auth = AuthService::new(db.clone() as Arc, jwt.clone()); + (db, jwt, auth) + } + + async fn create_viewer(db: &Database, username: &str, password: &str) -> i64 { + let hash = password::hash_password(password).expect("hash"); + let user_id = db + .insert_user(username, &hash, ROLE_VIEWER, false) + .await + .expect("insert user"); + let viewer_group = db + .list_user_groups() + .await + .expect("groups") + .into_iter() + .find(|g| g.name == GROUP_VIEWER) + .expect("viewer group"); + db.set_user_groups(user_id, &[viewer_group.id]) + .await + .expect("assign viewer group"); + user_id + } + + #[tokio::test] + async fn relogin_token_permissions_follow_role_promotion_and_demotion() { + let (db, jwt, auth) = auth_fixture().await; + let user_id = create_viewer(&db, "alice", "Correct Horse 123!").await; + + db.update_user_role(user_id, ROLE_ADMIN).await.expect("promote"); + let promoted = auth.login("alice", "Correct Horse 123!").await.expect("login"); + let promoted_claims = jwt.validate_token(&promoted.token).expect("promoted token"); + assert_eq!(promoted.role, ROLE_ADMIN); + assert_eq!(promoted_claims.role, ROLE_ADMIN); + assert!(promoted_claims.permissions.contains(&"users:admin".to_string())); + + db.update_user_role(user_id, ROLE_VIEWER).await.expect("demote"); + let demoted = auth.login("alice", "Correct Horse 123!").await.expect("login"); + let demoted_claims = jwt.validate_token(&demoted.token).expect("demoted token"); + assert_eq!(demoted.role, ROLE_VIEWER); + assert_eq!(demoted_claims.role, ROLE_VIEWER); + assert!(!demoted_claims.permissions.contains(&"users:admin".to_string())); + } +} diff --git a/net-guardia/src/core/response/engine.rs b/net-guardia/src/core/response/engine.rs index 52a849f..86da25d 100644 --- a/net-guardia/src/core/response/engine.rs +++ b/net-guardia/src/core/response/engine.rs @@ -65,9 +65,11 @@ impl SoarEngine { let soar_cfg = config.load(); let rate_limit_channel = soar_cfg.soar.rate_limit_cmd_channel_capacity; let freq_max_keys = soar_cfg.soar.frequency_max_tracked_keys; + let freq_max_events_per_key = soar_cfg.soar.frequency_max_events_per_key; + let freq_retention_secs = soar_cfg.soar.frequency_retention_secs; drop(soar_cfg); let rate_limit_owner = rate_limit.map(|rl| RateLimitOwnerHandle::spawn(rl, config.clone(), rate_limit_channel)); - let matcher = PlaybookMatcher::new(config, freq_max_keys); + let matcher = PlaybookMatcher::new(config, freq_max_keys, freq_max_events_per_key, freq_retention_secs); let engine = Self { db, access_control, diff --git a/net-guardia/src/core/response/frequency.rs b/net-guardia/src/core/response/frequency.rs index 51452df..48cdef4 100644 --- a/net-guardia/src/core/response/frequency.rs +++ b/net-guardia/src/core/response/frequency.rs @@ -11,14 +11,16 @@ pub struct FrequencyTracker { events: DashMap>, max_deque_size: usize, max_tracked_keys: usize, + max_retention: Duration, } impl FrequencyTracker { - pub fn new(max_tracked_keys: usize) -> Self { + pub fn new(max_tracked_keys: usize, max_events_per_key: usize, retention_secs: u64) -> Self { Self { events: DashMap::new(), - max_deque_size: 200, + max_deque_size: max_events_per_key.max(1), max_tracked_keys: max_tracked_keys.max(1), + max_retention: Duration::from_secs(retention_secs.max(1)), } } @@ -50,20 +52,19 @@ impl FrequencyTracker { deque.len() as u64 } - /// Remove empty deques and entries where all timestamps are expired. - /// Uses a conservative 2-hour max window for expiry detection. + /// Remove empty deques and entries where all timestamps are outside the + /// configured retention window. pub fn cleanup(&self) -> u32 { let now = Instant::now(); - let max_window = Duration::from_secs(7200); // 2 hours — conservative upper bound let mut removed = 0u32; self.events.retain(|_, deque| { if deque.is_empty() { removed += 1; return false; } - // If all entries are older than max_window, remove the whole entry + // If all entries are older than max_retention, remove the whole entry. if let Some(newest) = deque.back() - && now.checked_duration_since(*newest).unwrap_or(Duration::ZERO) > max_window + && now.checked_duration_since(*newest).unwrap_or(Duration::ZERO) > self.max_retention { removed += 1; return false; @@ -84,3 +85,17 @@ impl FrequencyTracker { removed } } + +#[cfg(test)] +mod tests { + use super::FrequencyTracker; + + #[test] + fn record_and_count_respects_configured_per_key_cap() { + let tracker = FrequencyTracker::new(10, 2, 60); + + assert_eq!(tracker.record_and_count(1, "10.0.0.1", 60), 1); + assert_eq!(tracker.record_and_count(1, "10.0.0.1", 60), 2); + assert_eq!(tracker.record_and_count(1, "10.0.0.1", 60), 2); + } +} diff --git a/net-guardia/src/core/response/matcher.rs b/net-guardia/src/core/response/matcher.rs index 0ca1266..3b35dfc 100644 --- a/net-guardia/src/core/response/matcher.rs +++ b/net-guardia/src/core/response/matcher.rs @@ -29,12 +29,21 @@ pub struct PlaybookMatcher { } impl PlaybookMatcher { - pub fn new(config: Arc>, frequency_max_tracked_keys: usize) -> Self { + pub fn new( + config: Arc>, + frequency_max_tracked_keys: usize, + frequency_max_events_per_key: usize, + frequency_retention_secs: u64, + ) -> Self { Self { playbooks: ArcSwap::from_pointee(Vec::new()), admin_whitelist: ArcSwap::from_pointee(HashSet::new()), cooldowns: DashMap::new(), - frequency_tracker: FrequencyTracker::new(frequency_max_tracked_keys), + frequency_tracker: FrequencyTracker::new( + frequency_max_tracked_keys, + frequency_max_events_per_key, + frequency_retention_secs, + ), active_block_count: AtomicU32::new(0), config, } diff --git a/net-guardia/src/core/response/scheduler.rs b/net-guardia/src/core/response/scheduler.rs index 68e28eb..9dcd4db 100644 --- a/net-guardia/src/core/response/scheduler.rs +++ b/net-guardia/src/core/response/scheduler.rs @@ -116,3 +116,204 @@ async fn unblock_ip_blocking(access_control: Arc, source_ .await .map_err(|e| crate::domain::response::error::SoarError::ActionFailed("unblock_ip", e))? } + +#[cfg(test)] +mod tests { + use std::sync::atomic::{AtomicBool, AtomicU8, Ordering}; + + use arc_swap::ArcSwap; + use parking_lot::Mutex; + + use super::*; + use crate::adapter::persistence::Database; + use crate::core::response::engine::SoarEngineDeps; + use crate::domain::common::config::AppConfig; + use crate::domain::common::config::notification::SmtpConfig; + use crate::domain::data_plane::error::EbpfError; + use crate::interface::email_sender::{EmailSender, EmailSenderFactory}; + use crate::interface::secret_store::SecretStorePort; + + const EXPIRED_AT: &str = "2000-01-01 00:00:00"; + const SOURCE_IP: &str = "198.51.100.10"; + + struct MockAccessControl { + unblocked_ips: Mutex>, + should_fail: AtomicBool, + } + + impl MockAccessControl { + fn new(should_fail: bool) -> Self { + Self { + unblocked_ips: Mutex::new(Vec::new()), + should_fail: AtomicBool::new(should_fail), + } + } + } + + impl AccessControlPort for MockAccessControl { + fn block_ip(&self, _ip: &str) -> Result<(), Error> { + Ok(()) + } + + fn unblock_ip(&self, ip: &str) -> Result<(), Error> { + if self.should_fail.load(Ordering::SeqCst) { + Err(EbpfError::UnknownError)?; + } + self.unblocked_ips.lock().push(ip.to_string()); + Ok(()) + } + } + + struct NoopEmailSenderFactory; + + #[async_trait::async_trait] + impl EmailSenderFactory for NoopEmailSenderFactory { + async fn build_smtp_sender( + &self, + _cfg: &SmtpConfig, + _secrets: Option<&dyn SecretStorePort>, + ) -> Result>, Error> { + Ok(None) + } + } + + async fn test_db() -> Arc { + let db = Arc::new(Database::new(":memory:").await.expect("test db")); + AppConfig::seed_config_defaults(&*db) + .await + .expect("seed config defaults"); + db + } + + async fn test_scheduler( + db: Arc, + access_control: Arc, + ) -> (TtlScheduler, Arc) { + let cfg = AppConfig::from_config_repo(&*db).await.expect("load config"); + let engine = Arc::new( + SoarEngine::new(SoarEngineDeps { + db: db.clone() as Arc, + config: Arc::new(ArcSwap::from_pointee(cfg)), + access_control: access_control.clone(), + alert_notifier: None, + geoip: None, + rate_limit: None, + enforce_level_cache: Arc::new(AtomicU8::new(2)), + secrets: None, + email_sender_factory: Arc::new(NoopEmailSenderFactory), + }) + .await + .expect("soar engine"), + ); + let scheduler = TtlScheduler::new(db, access_control, engine.clone()); + (scheduler, engine) + } + + fn acl_contains(rules: &[crate::domain::data_plane::acl_rule::AclRuleView], ip: &str) -> bool { + rules + .iter() + .any(|rule| rule.ip_address == ip && rule.direction == "source" && rule.list_type == "blacklist") + } + + #[tokio::test] + async fn ttl_sweep_preserves_manual_acl_rule_for_expired_soar_block() { + let db = test_db().await; + let block_id = db + .commit_soar_block_to_db(SOURCE_IP, 4, 1, EXPIRED_AT) + .await + .expect("insert expired block"); + db.insert_acl_rule(4, "source", "blacklist", SOURCE_IP, 0) + .await + .expect("manual ACL should preserve block"); + let access_control = Arc::new(MockAccessControl::new(false)); + let (scheduler, engine) = test_scheduler(db.clone(), access_control.clone()).await; + engine.matcher.active_block_count.store(1, Ordering::SeqCst); + + scheduler.sweep().await.expect("ttl sweep"); + + assert!( + access_control.unblocked_ips.lock().is_empty(), + "manual ACL ownership must skip data-plane unblock" + ); + assert!( + db.list_expired_soar_blocks().await.expect("expired blocks").is_empty(), + "expired SOAR record should be marked unblocked" + ); + assert!( + acl_contains(&db.list_acl_rules().await.expect("acl rules"), SOURCE_IP), + "manual ACL row must remain after SOAR TTL expiry" + ); + assert_eq!( + engine.matcher.active_block_count.load(Ordering::SeqCst), + 0, + "SOAR active-block counter should decrement once" + ); + assert!(db.find_soar_block_by_id(block_id).await.expect("find block").is_some()); + } + + #[tokio::test] + async fn ttl_sweep_records_pending_unblock_without_clearing_db_when_data_plane_unblock_fails() { + let db = test_db().await; + db.commit_soar_block_to_db(SOURCE_IP, 4, 1, EXPIRED_AT) + .await + .expect("insert expired block"); + let access_control = Arc::new(MockAccessControl::new(true)); + let (scheduler, engine) = test_scheduler(db.clone(), access_control.clone()).await; + engine.matcher.active_block_count.store(1, Ordering::SeqCst); + + scheduler.sweep().await.expect("ttl sweep"); + + let pending = db.list_pending_unblocks().await.expect("pending unblocks"); + assert_eq!(pending.len(), 1); + assert_eq!(pending[0].source_ip, SOURCE_IP); + assert!( + !db.list_expired_soar_blocks().await.expect("expired blocks").is_empty(), + "failed data-plane unblock must leave SOAR block active for retry" + ); + assert!( + acl_contains(&db.list_acl_rules().await.expect("acl rules"), SOURCE_IP), + "ACL row must remain while durable unblock has not succeeded" + ); + assert!( + access_control.unblocked_ips.lock().is_empty(), + "failing data-plane call should not record a successful unblock" + ); + assert_eq!( + engine.matcher.active_block_count.load(Ordering::SeqCst), + 1, + "active-block counter must not decrement until unblock succeeds" + ); + } + + #[tokio::test] + async fn ttl_sweep_deletes_soar_owned_acl_only_after_data_plane_unblock_succeeds() { + let db = test_db().await; + db.commit_soar_block_to_db(SOURCE_IP, 4, 1, EXPIRED_AT) + .await + .expect("insert expired block"); + let access_control = Arc::new(MockAccessControl::new(false)); + let (scheduler, engine) = test_scheduler(db.clone(), access_control.clone()).await; + engine.matcher.active_block_count.store(1, Ordering::SeqCst); + + scheduler.sweep().await.expect("ttl sweep"); + + assert_eq!(access_control.unblocked_ips.lock().as_slice(), &[SOURCE_IP.to_string()]); + assert!( + db.list_expired_soar_blocks().await.expect("expired blocks").is_empty(), + "successful unblock should mark SOAR record unblocked" + ); + assert!( + !acl_contains(&db.list_acl_rules().await.expect("acl rules"), SOURCE_IP), + "SOAR-owned ACL row should be removed after data-plane unblock succeeds" + ); + assert!( + db.list_pending_unblocks().await.expect("pending unblocks").is_empty(), + "successful unblock should not enqueue retry work" + ); + assert_eq!( + engine.matcher.active_block_count.load(Ordering::SeqCst), + 0, + "active-block counter should decrement after successful unblock" + ); + } +} diff --git a/net-guardia/src/domain/common/config/acl.rs b/net-guardia/src/domain/common/config/acl.rs index aa43867..a22ed2c 100644 --- a/net-guardia/src/domain/common/config/acl.rs +++ b/net-guardia/src/domain/common/config/acl.rs @@ -3,6 +3,8 @@ use macros::config_settings; #[config_settings(section = "misc")] #[derive(Debug, Clone)] pub struct AclConfig { - #[setting(key = "geoip_db_name", default = "net-guardia/static/geo/dbip-city-lite.mmdb")] - pub geoip_db_name: String, + #[setting(key = "geoip_db_path", default = "net-guardia/static/geo/dbip-city-lite.mmdb")] + pub geoip_db_path: String, + #[setting(key = "geoip_cache_capacity", default = "10000")] + pub geoip_cache_capacity: usize, } diff --git a/net-guardia/src/domain/common/config/detection.rs b/net-guardia/src/domain/common/config/detection.rs index a16917b..247ee92 100644 --- a/net-guardia/src/domain/common/config/detection.rs +++ b/net-guardia/src/domain/common/config/detection.rs @@ -9,6 +9,10 @@ pub struct FusionConfig { pub repeat_offender_window_secs: u64, #[setting(key = "fusion_max_dedup_entries", default = "50000")] pub max_dedup_entries: usize, + #[setting(key = "fusion_source_count_max_entries", default = "10000")] + pub source_count_max_entries: usize, + #[setting(key = "fusion_repeat_tracker_max_entries", default = "5000")] + pub repeat_tracker_max_entries: usize, } #[config_settings(section = "beaconing")] @@ -22,12 +26,23 @@ pub struct BeaconingConfig { pub cv_threshold: f64, #[setting(key = "beaconing_max_cache_entries", default = "50000")] pub max_cache_entries: usize, + #[setting(key = "beaconing_max_timestamps_per_flow", default = "100")] + pub max_timestamps_per_flow: usize, #[setting(key = "beaconing_expiry_secs", default = "600")] pub expiry_secs: u64, #[setting(key = "beaconing_alert_cooldown_secs", default = "300")] pub alert_cooldown_secs: u64, } +#[config_settings(section = "flow_stats")] +#[derive(Debug, Clone)] +pub struct FlowStatsConfig { + #[setting(key = "flow_stats_max_snapshot_entries", default = "10000")] + pub max_snapshot_entries: usize, + #[setting(key = "flow_stats_max_top_n", default = "10000")] + pub max_top_n: usize, +} + #[config_settings(section = "detection")] #[derive(Debug, Clone)] pub struct DetectionConfig { @@ -35,6 +50,8 @@ pub struct DetectionConfig { pub fusion: FusionConfig, #[setting(flatten)] pub beaconing: BeaconingConfig, + #[setting(flatten)] + pub flow_stats: FlowStatsConfig, #[setting(key = "detection_cleanup_interval_secs", default = "60")] pub cleanup_interval_secs: u64, } diff --git a/net-guardia/src/domain/common/config/ebpf.rs b/net-guardia/src/domain/common/config/ebpf.rs index 13f5b48..3567b81 100644 --- a/net-guardia/src/domain/common/config/ebpf.rs +++ b/net-guardia/src/domain/common/config/ebpf.rs @@ -32,6 +32,12 @@ pub struct EbpfConfig { pub packet_buffer_size: usize, #[setting(section = "xdp", key = "buffer_pool_capacity", default = "1024")] pub buffer_pool_capacity: usize, + #[setting(section = "xdp", key = "xsk_completion_batch_size", default = "256")] + pub xsk_completion_batch_size: usize, + #[setting(section = "xdp", key = "xsk_rx_batch_size", default = "64")] + pub xsk_rx_batch_size: usize, + #[setting(section = "xdp", key = "xsk_tx_batch_size", default = "64")] + pub xsk_tx_batch_size: usize, // ── internal (not exposed in API) ────────────────────────────── #[setting(section = "xdp", key = "default_packet_rate", default = "10000", api = false)] pub default_packet_rate: u64, diff --git a/net-guardia/src/domain/common/config/health.rs b/net-guardia/src/domain/common/config/health.rs index 5e217be..33cfdb6 100644 --- a/net-guardia/src/domain/common/config/health.rs +++ b/net-guardia/src/domain/common/config/health.rs @@ -21,4 +21,6 @@ pub struct HealthConfig { pub temp_warn_celsius: f32, #[setting(key = "health_monitoring_interval_secs", default = "5")] pub monitoring_interval_secs: u64, + #[setting(key = "health_broadcast_channel_capacity", default = "100")] + pub broadcast_channel_capacity: usize, } diff --git a/net-guardia/src/domain/common/config/ml.rs b/net-guardia/src/domain/common/config/ml.rs index 94458e5..60a0788 100644 --- a/net-guardia/src/domain/common/config/ml.rs +++ b/net-guardia/src/domain/common/config/ml.rs @@ -3,11 +3,6 @@ use macros::config_settings; #[config_settings] #[derive(Debug, Clone)] pub struct MlConfig { - // ── models ───────────────────────────────────────────────────── - #[setting(section = "models", key = "deep_autoencoder_name", default = "deep_autoencoder.onnx")] - pub deep_autoencoder_name: String, - #[setting(section = "models", key = "classifier_name", default = "classifier.onnx")] - pub classifier_name: String, #[setting(section = "models", key = "models_config_name", default = "inference_config.json")] pub models_config_name: String, diff --git a/net-guardia/src/domain/common/config/mod.rs b/net-guardia/src/domain/common/config/mod.rs index 7dc79c3..6299d0b 100644 --- a/net-guardia/src/domain/common/config/mod.rs +++ b/net-guardia/src/domain/common/config/mod.rs @@ -127,6 +127,12 @@ impl AppConfig { require(self.ebpf.rx_queue_size > 0, "ebpf.rx_queue_size")?; require(self.ebpf.frame_size > 0, "ebpf.frame_size")?; require(self.ebpf.frame_count > 0, "ebpf.frame_count")?; + require( + self.ebpf.xsk_completion_batch_size > 0, + "ebpf.xsk_completion_batch_size", + )?; + require(self.ebpf.xsk_rx_batch_size > 0, "ebpf.xsk_rx_batch_size")?; + require(self.ebpf.xsk_tx_batch_size > 0, "ebpf.xsk_tx_batch_size")?; require(self.http_server.port > 0, "http_server.port")?; require(self.ml.max_concurrent_flows > 0, "ml.max_concurrent_flows")?; require(self.ml.min_packets_for_inference > 0, "ml.min_packets_for_inference")?; @@ -160,6 +166,12 @@ impl AppConfig { self.soar.frequency_max_tracked_keys > 0, "soar.frequency_max_tracked_keys", )?; + require(self.soar.execution_list_limit > 0, "soar.execution_list_limit")?; + require( + self.soar.frequency_max_events_per_key > 0, + "soar.frequency_max_events_per_key", + )?; + require(self.soar.frequency_retention_secs > 0, "soar.frequency_retention_secs")?; require( (0.0..=1.0).contains(&self.soar.default_rate_limit_factor), "soar.default_rate_limit_factor", @@ -176,6 +188,14 @@ impl AppConfig { self.detection.fusion.max_dedup_entries > 0, "detection.fusion.max_dedup_entries", )?; + require( + self.detection.fusion.source_count_max_entries > 0, + "detection.fusion.source_count_max_entries", + )?; + require( + self.detection.fusion.repeat_tracker_max_entries > 0, + "detection.fusion.repeat_tracker_max_entries", + )?; require( self.detection.fusion.dedup_window_secs > 0, "detection.fusion.dedup_window_secs", @@ -188,6 +208,26 @@ impl AppConfig { self.detection.beaconing.max_cache_entries > 0, "detection.beaconing.max_cache_entries", )?; + require( + self.detection.beaconing.max_timestamps_per_flow > 1, + "detection.beaconing.max_timestamps_per_flow", + )?; + require( + self.detection.flow_stats.max_snapshot_entries > 0, + "detection.flow_stats.max_snapshot_entries", + )?; + require( + self.detection.flow_stats.max_top_n >= self.detection.flow_stats.max_snapshot_entries, + "detection.flow_stats.max_top_n", + )?; + require( + self.health.monitoring_interval_secs > 0, + "health.monitoring_interval_secs", + )?; + require( + self.health.broadcast_channel_capacity > 0, + "health.broadcast_channel_capacity", + )?; require( self.correlation.max_tracked_entries > 0, "correlation.max_tracked_entries", @@ -236,6 +276,7 @@ impl AppConfig { self.observability.drop_channel_capacity > 0, "observability.drop_channel_capacity", )?; + require(self.acl.geoip_cache_capacity > 0, "acl.geoip_cache_capacity")?; require(self.suricata.poll_interval_ms > 0, "suricata.poll_interval_ms")?; require( (0.0..=1.0).contains(&self.suricata.confidence_high), @@ -279,6 +320,9 @@ mod tests { assert_eq!(cfg.ebpf.egress_ifname, "eth1"); assert_eq!(cfg.ebpf.combined_queue_count, 1); assert_eq!(cfg.ebpf.frame_size, 4096); + assert_eq!(cfg.ebpf.xsk_completion_batch_size, 256); + assert_eq!(cfg.ebpf.xsk_rx_batch_size, 64); + assert_eq!(cfg.ebpf.xsk_tx_batch_size, 64); assert_eq!(cfg.http_server.jwt_expiry_hours, 24); } @@ -305,9 +349,15 @@ mod tests { let db = test_db().await; db.set_config_value("frame_size", "8192").await.unwrap(); db.set_config_value("combined_queue_count", "4").await.unwrap(); + db.set_config_value("xsk_completion_batch_size", "512").await.unwrap(); + db.set_config_value("xsk_rx_batch_size", "128").await.unwrap(); + db.set_config_value("xsk_tx_batch_size", "32").await.unwrap(); let cfg = AppConfig::from_config_repo(&db).await.unwrap(); assert_eq!(cfg.ebpf.frame_size, 8192); assert_eq!(cfg.ebpf.combined_queue_count, 4); + assert_eq!(cfg.ebpf.xsk_completion_batch_size, 512); + assert_eq!(cfg.ebpf.xsk_rx_batch_size, 128); + assert_eq!(cfg.ebpf.xsk_tx_batch_size, 32); } #[tokio::test] @@ -344,7 +394,7 @@ mod tests { Some("access_control,rate_limit,service".to_string()) ); assert_eq!( - db.get_config_value("geoip_db_name").await.unwrap(), + db.get_config_value("geoip_db_path").await.unwrap(), Some("net-guardia/static/geo/dbip-city-lite.mmdb".to_string()) ); assert_eq!( @@ -355,10 +405,30 @@ mod tests { db.get_config_value("soar_max_ttl_secs").await.unwrap(), Some("86400".to_string()) ); + assert_eq!( + db.get_config_value("soar_execution_list_limit").await.unwrap(), + Some("100".to_string()) + ); assert_eq!( db.get_config_value("ml_drift_window_secs").await.unwrap(), Some("3600".to_string()) ); + assert_eq!( + db.get_config_value("fusion_source_count_max_entries").await.unwrap(), + Some("10000".to_string()) + ); + assert_eq!( + db.get_config_value("fusion_repeat_tracker_max_entries").await.unwrap(), + Some("5000".to_string()) + ); + assert_eq!( + db.get_config_value("health_broadcast_channel_capacity").await.unwrap(), + Some("100".to_string()) + ); + assert_eq!( + db.get_config_value("geoip_cache_capacity").await.unwrap(), + Some("10000".to_string()) + ); } #[tokio::test] @@ -401,9 +471,47 @@ mod tests { async fn db_overrides_soar_and_drift() { let db = test_db().await; db.set_config_value("soar_max_ttl_secs", "3600").await.unwrap(); + db.set_config_value("soar_execution_list_limit", "42").await.unwrap(); db.set_config_value("ml_drift_window_secs", "900").await.unwrap(); let cfg = AppConfig::from_config_repo(&db).await.unwrap(); assert_eq!(cfg.soar.max_ttl_secs, 3600); + assert_eq!(cfg.soar.execution_list_limit, 42); assert_eq!(cfg.ml.drift_window_secs, 900); } + + #[tokio::test] + async fn db_overrides_fusion_runtime_cache_sizes() { + let db = test_db().await; + db.set_config_value("fusion_source_count_max_entries", "1234") + .await + .unwrap(); + db.set_config_value("fusion_repeat_tracker_max_entries", "567") + .await + .unwrap(); + let cfg = AppConfig::from_config_repo(&db).await.unwrap(); + assert_eq!(cfg.detection.fusion.source_count_max_entries, 1234); + assert_eq!(cfg.detection.fusion.repeat_tracker_max_entries, 567); + } + + #[tokio::test] + async fn db_overrides_health_runtime_values() { + let db = test_db().await; + db.set_config_value("health_monitoring_interval_secs", "9") + .await + .unwrap(); + db.set_config_value("health_broadcast_channel_capacity", "321") + .await + .unwrap(); + let cfg = AppConfig::from_config_repo(&db).await.unwrap(); + assert_eq!(cfg.health.monitoring_interval_secs, 9); + assert_eq!(cfg.health.broadcast_channel_capacity, 321); + } + + #[tokio::test] + async fn db_overrides_geoip_runtime_cache_capacity() { + let db = test_db().await; + db.set_config_value("geoip_cache_capacity", "2048").await.unwrap(); + let cfg = AppConfig::from_config_repo(&db).await.unwrap(); + assert_eq!(cfg.acl.geoip_cache_capacity, 2048); + } } diff --git a/net-guardia/src/domain/common/config/soar.rs b/net-guardia/src/domain/common/config/soar.rs index 7560705..1803729 100644 --- a/net-guardia/src/domain/common/config/soar.rs +++ b/net-guardia/src/domain/common/config/soar.rs @@ -29,6 +29,12 @@ pub struct SoarConfig { pub rate_limit_cmd_channel_capacity: usize, #[setting(key = "soar_frequency_max_tracked_keys", default = "50000")] pub frequency_max_tracked_keys: usize, + #[setting(key = "soar_frequency_max_events_per_key", default = "200")] + pub frequency_max_events_per_key: usize, + #[setting(key = "soar_frequency_retention_secs", default = "7200")] + pub frequency_retention_secs: u64, #[setting(key = "soar_fallback_cooldown_secs", default = "300")] pub fallback_cooldown_secs: i64, + #[setting(key = "soar_execution_list_limit", default = "100")] + pub execution_list_limit: i64, } diff --git a/net-guardia/src/domain/common/error/database.rs b/net-guardia/src/domain/common/error/database.rs index db34e45..a71348b 100644 --- a/net-guardia/src/domain/common/error/database.rs +++ b/net-guardia/src/domain/common/error/database.rs @@ -5,9 +5,6 @@ traceable! { #[error("Database error: {err}")] QueryFailed => tracing::Level::ERROR, - #[error("Database connection failed")] - ConnectionFailed => tracing::Level::ERROR, - #[no_source] #[error("User '{username}' already exists")] UserAlreadyExists { username: String } => tracing::Level::WARN, diff --git a/net-guardia/src/domain/common/error/mcp.rs b/net-guardia/src/domain/common/error/mcp.rs deleted file mode 100644 index 85f3532..0000000 --- a/net-guardia/src/domain/common/error/mcp.rs +++ /dev/null @@ -1,25 +0,0 @@ -use macros::traceable; - -traceable! { - McpError { - #[no_source] - #[error("MCP tool not found: {tool_name}")] - ToolNotFound { tool_name: String } => tracing::Level::WARN, - - #[no_source] - #[error("MCP parameter validation failed: {detail}")] - InvalidParams { detail: String } => tracing::Level::WARN, - - #[no_source] - #[error("MCP permission denied: key has '{key_level}' but tool requires '{required_level}'")] - PermissionDenied { key_level: String, required_level: String } => tracing::Level::WARN, - - #[no_source] - #[error("MCP API key invalid or revoked")] - InvalidApiKey => tracing::Level::WARN, - - #[no_source] - #[error("MCP proxy error: {detail}")] - ProxyError { detail: String } => tracing::Level::ERROR, - } -} diff --git a/net-guardia/src/domain/common/error/misc.rs b/net-guardia/src/domain/common/error/misc.rs index b007c29..a57abd2 100644 --- a/net-guardia/src/domain/common/error/misc.rs +++ b/net-guardia/src/domain/common/error/misc.rs @@ -6,19 +6,9 @@ traceable! { #[error("Failed to remove limit on locked memory, ret is: {ret}")] RamLimitUnlockError { ret: i32 } => tracing::Level::ERROR, - #[error("Failed to send message to receiver")] - SendMessageError => tracing::Level::ERROR, - #[error("Failed to serialize data")] SerializeError => tracing::Level::ERROR, - #[error("Failed to deserialize data")] - DeserializeError => tracing::Level::ERROR, - - #[no_source] - #[error("Network interface '{interface}' not found")] - NetworkInterfaceNotFound { interface: String } => tracing::Level::ERROR, - #[error("Failed to open GeoIP database '{path}': {err}")] GeoIPDatabaseError { path: String } => tracing::Level::ERROR, @@ -33,18 +23,6 @@ traceable! { #[error("DNS domain name too long: '{domain}'")] DnsDomainTooLong { domain: String } => tracing::Level::WARN, - #[no_source] - #[error("Type mismatch during message dispatch")] - TypeMismatch => tracing::Level::ERROR, - - #[no_source] - #[error("No handler registered for this message type")] - HandlerNotFound => tracing::Level::ERROR, - - #[no_source] - #[error("Event type not registered with communication manager")] - TypeNotRegistered => tracing::Level::ERROR, - #[no_source] #[error("Validation error: {message}")] ValidationError { message: String } => tracing::Level::WARN, diff --git a/net-guardia/src/domain/common/error/mod.rs b/net-guardia/src/domain/common/error/mod.rs index 9d33386..a7c3b3e 100644 --- a/net-guardia/src/domain/common/error/mod.rs +++ b/net-guardia/src/domain/common/error/mod.rs @@ -2,7 +2,6 @@ pub mod crypto; pub mod database; pub mod http; pub mod io; -pub mod mcp; pub mod misc; pub mod notification; pub mod system; @@ -13,7 +12,6 @@ use crate::domain::common::error::crypto::CryptoError; use crate::domain::common::error::database::DatabaseError; use crate::domain::common::error::http::HttpError; use crate::domain::common::error::io::IOError; -use crate::domain::common::error::mcp::McpError; use crate::domain::common::error::misc::MiscError; use crate::domain::common::error::notification::NotificationError; use crate::domain::common::error::system::SystemError; @@ -40,8 +38,6 @@ pub enum Error { #[error("{0}")] IO(#[from] IOError), #[error("{0}")] - Mcp(#[from] McpError), - #[error("{0}")] Misc(#[from] MiscError), #[error("{0}")] Notification(#[from] NotificationError), diff --git a/net-guardia/src/domain/common/error/system.rs b/net-guardia/src/domain/common/error/system.rs index edf2627..860aa31 100644 --- a/net-guardia/src/domain/common/error/system.rs +++ b/net-guardia/src/domain/common/error/system.rs @@ -2,10 +2,6 @@ use macros::traceable; traceable! { SystemError { - #[no_source] - #[error("Unable to run as administrator")] - RunAsAdminFailed => tracing::Level::ERROR, - #[no_source] #[error("Invalid configuration")] InvalidConfig => tracing::Level::ERROR, @@ -17,38 +13,24 @@ traceable! { #[error("Configuration file not found")] ConfigNotFound => tracing::Level::ERROR, - #[error("Failed to terminate instance")] - TerminateError => tracing::Level::ERROR, - #[no_source] #[error("Failed to send shutdown signal")] ShutdownSignalFailed => tracing::Level::ERROR, - #[error("Unexpected thread panic")] - ThreadPanic => tracing::Level::ERROR, - #[error("Unexpected error")] UnexpectedError => tracing::Level::ERROR, - #[error("Failed to reload config after setup")] - ConfigReloadFailed => tracing::Level::ERROR, - #[error("HTTP server error")] HttpServerError => tracing::Level::ERROR, #[error("Failed to set user groups")] SetUserGroupsFailed => tracing::Level::WARN, - #[error("Failed to store XDP mode")] - XdpModeStoreFailed => tracing::Level::WARN, - #[error("Failed to update admin password during setup: {err}")] SetupPasswordUpdateFailed => tracing::Level::ERROR, #[error("Failed to mark setup as complete: {err}")] SetupCompleteFlagFailed => tracing::Level::ERROR, - #[error("Failed to publish drift detected event")] - DriftEventPublishFailed => tracing::Level::WARN, } } diff --git a/net-guardia/src/domain/common/log/system.rs b/net-guardia/src/domain/common/log/system.rs index e385b2e..914eff1 100644 --- a/net-guardia/src/domain/common/log/system.rs +++ b/net-guardia/src/domain/common/log/system.rs @@ -3,9 +3,6 @@ use tracing; loggable! { SystemLog { - #[error("Online now")] - Online => tracing::Level::INFO, - #[error("Initializing")] Initializing => tracing::Level::INFO, @@ -27,9 +24,6 @@ loggable! { #[error("Setup wizard completed — starting full system initialization")] SetupCompleted => tracing::Level::INFO, - #[error("Config reloaded from DB: ingress={ingress}, egress={egress}")] - ConfigReloaded { ingress: String, egress: String } => tracing::Level::INFO, - #[error("Full system initialization complete — all services running")] FullInitComplete => tracing::Level::INFO, diff --git a/net-guardia/src/domain/common/system/health.rs b/net-guardia/src/domain/common/system/health.rs index 4a798f1..de36b11 100644 --- a/net-guardia/src/domain/common/system/health.rs +++ b/net-guardia/src/domain/common/system/health.rs @@ -63,8 +63,8 @@ pub enum EbpfFailCategory { InterfaceNotFound, /// Interface exists but XDP native/SKB attach refused by driver. XdpUnsupported, - /// AF_XDP bind rejected — driver does not implement AF_XDP on this kernel. - /// Common case: Intel i350 (igb) on kernel < 6.17. + /// AF_XDP bind rejected — driver or netdev capabilities do not support + /// the requested AF_XDP socket mode on this interface. AfXdpUnsupported, /// ENOMEM / RLIMIT_MEMLOCK exhausted. MemlockExhausted, diff --git a/net-guardia/src/domain/data_plane/flow_stats.rs b/net-guardia/src/domain/data_plane/flow_stats.rs index 95224e5..8de9ab1 100644 --- a/net-guardia/src/domain/data_plane/flow_stats.rs +++ b/net-guardia/src/domain/data_plane/flow_stats.rs @@ -2,6 +2,28 @@ use serde::{Deserialize, Serialize}; use crate::domain::data_plane::direction::Direction; +#[derive(Debug, Clone, Copy)] +pub struct FlowStatsLimits { + max_snapshot_entries: usize, + max_top_n: usize, +} + +impl FlowStatsLimits { + pub fn new(max_snapshot_entries: usize, max_top_n: usize) -> Self { + Self { + max_snapshot_entries: max_snapshot_entries.max(1), + max_top_n: max_top_n.max(1), + } + } + + pub fn result_limit(self, requested_top_n: Option) -> usize { + requested_top_n + .unwrap_or(self.max_snapshot_entries) + .max(1) + .min(self.max_top_n) + } +} + #[derive(Debug, Clone, Serialize)] pub struct FlowStatsEntry { pub direction: Direction, @@ -61,3 +83,29 @@ pub struct FlowSubscription { /// Push interval in seconds (default 5) pub interval_secs: Option, } + +#[cfg(test)] +mod tests { + use super::FlowStatsLimits; + + #[test] + fn result_limit_uses_snapshot_default_when_top_n_absent() { + let limits = FlowStatsLimits::new(128, 1024); + + assert_eq!(limits.result_limit(None), 128); + } + + #[test] + fn result_limit_caps_requested_top_n() { + let limits = FlowStatsLimits::new(128, 256); + + assert_eq!(limits.result_limit(Some(1_000)), 256); + } + + #[test] + fn result_limit_never_returns_zero() { + let limits = FlowStatsLimits::new(0, 0); + + assert_eq!(limits.result_limit(Some(0)), 1); + } +} diff --git a/net-guardia/src/domain/identity/auth.rs b/net-guardia/src/domain/identity/auth.rs index 0565a14..7291b7b 100644 --- a/net-guardia/src/domain/identity/auth.rs +++ b/net-guardia/src/domain/identity/auth.rs @@ -36,7 +36,8 @@ impl PermissionLevel { pub fn permissions(self) -> &'static [&'static str] { match self { Self::ReadOnly => API_KEY_READ_ONLY_PERMISSIONS, - Self::ReadWrite | Self::FullAccess => API_KEY_READ_WRITE_PERMISSIONS, + Self::ReadWrite => API_KEY_READ_WRITE_PERMISSIONS, + Self::FullAccess => API_KEY_FULL_ACCESS_PERMISSIONS, } } } @@ -115,6 +116,8 @@ pub const API_KEY_READ_ONLY_PERMISSIONS: &[&str] = &[ "system:read", ]; +pub const API_KEY_FULL_ACCESS_PERMISSIONS: &[&str] = ADMIN_PERMISSIONS; + #[derive(Debug, Serialize, Deserialize, Clone)] pub struct Claims { pub sub: i64, @@ -123,3 +126,31 @@ pub struct Claims { pub permissions: Vec, pub exp: usize, } + +#[cfg(test)] +mod tests { + use super::PermissionLevel; + + #[test] + fn read_write_api_keys_do_not_receive_admin_permissions() { + let permissions = PermissionLevel::ReadWrite.permissions(); + + assert!(!permissions.contains(&"api_keys:admin")); + assert!(!permissions.contains(&"system:admin")); + assert!(!permissions.contains(&"users:admin")); + } + + #[test] + fn full_access_api_keys_receive_admin_permissions() { + let permissions = PermissionLevel::FullAccess.permissions(); + + assert!(permissions.contains(&"api_keys:admin")); + assert!(permissions.contains(&"system:admin")); + assert!(permissions.contains(&"users:admin")); + } + + #[test] + fn full_access_has_more_permissions_than_read_write() { + assert!(PermissionLevel::FullAccess.permissions().len() > PermissionLevel::ReadWrite.permissions().len()); + } +} diff --git a/net-guardia/src/infrastructure/ebpf_preflight.rs b/net-guardia/src/infrastructure/ebpf_preflight.rs index 87e4f9c..87bfc5b 100644 --- a/net-guardia/src/infrastructure/ebpf_preflight.rs +++ b/net-guardia/src/infrastructure/ebpf_preflight.rs @@ -10,7 +10,17 @@ //! host kernel release, and the NIC driver where those are obtainable, //! so the operator can diagnose directly from the UI without shelling in. +use std::ffi::OsStr; +use std::fmt; use std::fs; +use std::mem; +use std::os::fd::RawFd; +use std::path::{Path, PathBuf}; + +use libc::{ + AF_NETLINK, NETLINK_GENERIC, SOCK_CLOEXEC, SOCK_RAW, bind, close, genlmsghdr, nlmsghdr, recv, sendto, sockaddr, + sockaddr_nl, socket, +}; use crate::domain::common::error::Error; use crate::domain::common::system::health::{EbpfFailCategory, EbpfFailStage, EbpfHealth}; @@ -24,19 +34,92 @@ pub fn kernel_release() -> String { .unwrap_or_else(|| "unknown".to_string()) } -/// Look up the driver name bound to a network interface via -/// `/sys/class/net//device/driver`. Returns the basename of -/// the symlink target, or `"unknown"` if the interface has no driver -/// (e.g., virtual or renamed) or the path is not readable. -pub fn interface_driver(ifname: &str) -> String { - let link = format!("/sys/class/net/{}/device/driver", ifname); - match fs::read_link(&link) { - Ok(target) => target - .file_name() - .and_then(|s| s.to_str()) - .map(|s| s.to_string()) - .unwrap_or_else(|| "unknown".to_string()), - Err(_) => "unknown".to_string(), +const SYS_CLASS_NET: &str = "/sys/class/net"; +const NETDEV_FAMILY_NAME: &str = "netdev"; +const GENL_ID_CTRL: u16 = 0x10; +const CTRL_CMD_GETFAMILY: u8 = 3; +const CTRL_ATTR_FAMILY_ID: u16 = 1; +const CTRL_ATTR_FAMILY_NAME: u16 = 2; +const NETDEV_CMD_DEV_GET: u8 = 1; +const NETDEV_A_DEV_IFINDEX: u16 = 1; +const NETDEV_A_DEV_XDP_FEATURES: u16 = 3; +const NETDEV_A_DEV_XDP_ZC_MAX_SEGS: u16 = 4; +const NETDEV_A_DEV_XSK_FEATURES: u16 = 6; +const NETDEV_XDP_ACT_XSK_ZEROCOPY: u64 = 8; +const NLMSG_ERROR: u16 = 2; +const NLM_F_REQUEST: u16 = 1; +const NETLINK_SEQUENCE: u32 = 1; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct NetdevCapabilities { + pub xdp_features: Option, + pub xdp_zc_max_segs: Option, + pub xsk_features: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct InterfaceCapabilities { + pub driver: String, + pub mtu: Option, + pub rx_queues: Option, + pub tx_queues: Option, + pub netdev: Option, +} + +impl InterfaceCapabilities { + fn load(ifname: &str) -> Self { + Self::load_from(ifname, Path::new(SYS_CLASS_NET)) + } + + fn load_from(ifname: &str, sys_class_net: &Path) -> Self { + let iface_path = sys_class_net.join(ifname); + + Self { + driver: interface_driver_from(&iface_path), + mtu: read_u32(iface_path.join("mtu")), + rx_queues: count_queue_dirs(&iface_path, "rx-"), + tx_queues: count_queue_dirs(&iface_path, "tx-"), + netdev: read_u32(iface_path.join("ifindex")).and_then(load_netdev_capabilities), + } + } + + fn xsk_zerocopy_supported(&self) -> Option { + self.netdev + .as_ref() + .and_then(|netdev| netdev.xdp_features) + .map(|features| features & NETDEV_XDP_ACT_XSK_ZEROCOPY != 0) + } +} + +impl fmt::Display for InterfaceCapabilities { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "driver {}", self.driver)?; + + if let Some(mtu) = self.mtu { + write!(f, ", mtu {}", mtu)?; + } + if let Some(rx_queues) = self.rx_queues { + write!(f, ", rx_queues {}", rx_queues)?; + } + if let Some(tx_queues) = self.tx_queues { + write!(f, ", tx_queues {}", tx_queues)?; + } + match self.netdev.as_ref() { + Some(netdev) => { + if let Some(features) = netdev.xdp_features { + write!(f, ", xdp_features 0x{features:x}")?; + } + if let Some(max_segs) = netdev.xdp_zc_max_segs { + write!(f, ", xdp_zc_max_segs {}", max_segs)?; + } + if let Some(features) = netdev.xsk_features { + write!(f, ", xsk_features 0x{features:x}")?; + } + } + None => write!(f, ", netdev_features unavailable")?, + } + + Ok(()) } } @@ -48,24 +131,24 @@ pub fn classify(stage: EbpfFailStage, err: &Error, ifname: Option<&str>) -> Ebpf let raw = err.to_string(); let category = categorize(&raw); let kernel = kernel_release(); - let driver = ifname.map(interface_driver); + let capabilities = ifname.map(InterfaceCapabilities::load); let mut reason = format!("stage={:?}: {}", stage, raw); reason.push_str(&format!(" (kernel {}", kernel)); if let Some(iface) = ifname { reason.push_str(&format!(", interface {}", iface)); - if let Some(drv) = driver.as_deref() { - reason.push_str(&format!(", driver {}", drv)); + if let Some(caps) = capabilities.as_ref() { + reason.push_str(&format!(", {}", caps)); } } reason.push(')'); - // For the igb-before-6.17 case, augment the reason with a targeted hint. - if matches!(category, EbpfFailCategory::AfXdpUnsupported) - && driver.as_deref() == Some("igb") - && !kernel_meets_igb_af_xdp(&kernel) - { - reason.push_str(". The igb driver supports AF_XDP only on kernel 6.17 or newer."); + if matches!(category, EbpfFailCategory::AfXdpUnsupported) { + if let Some(Some(false)) = capabilities.as_ref().map(InterfaceCapabilities::xsk_zerocopy_supported) { + reason.push_str(". The interface reports no AF_XDP zero-copy capability in xdp_features."); + } else if capabilities.as_ref().and_then(|caps| caps.netdev.as_ref()).is_none() { + reason.push_str(". Kernel did not expose netdev XDP feature data for this interface."); + } } EbpfHealth::Unavailable { @@ -118,41 +201,263 @@ fn categorize(raw: &str) -> EbpfFailCategory { EbpfFailCategory::Unknown } -/// Parse a kernel release string like "6.17.4-generic" and return true -/// if it is >= 6.17. We only care about the first two numeric components. -fn kernel_meets_igb_af_xdp(release: &str) -> bool { - // Pull leading "MAJOR.MINOR" out of strings like "6.12.0-124.45.1.el10_1.x86_64". - let mut parts = release.split(|c: char| !c.is_ascii_digit()).filter(|s| !s.is_empty()); - let Some(major_str) = parts.next() else { - return false; +fn interface_driver_from(iface_path: &Path) -> String { + match fs::read_link(iface_path.join("device").join("driver")) { + Ok(target) => basename_to_string(&target), + Err(_) => "unknown".to_string(), + } +} + +fn basename_to_string(path: &Path) -> String { + path.file_name() + .and_then(OsStr::to_str) + .map(str::to_string) + .unwrap_or_else(|| "unknown".to_string()) +} + +fn read_u32(path: PathBuf) -> Option { + fs::read_to_string(path).ok()?.trim().parse().ok() +} + +fn count_queue_dirs(iface_path: &Path, prefix: &str) -> Option { + let queues_path = iface_path.join("queues"); + let entries = fs::read_dir(queues_path).ok()?; + let count = entries + .filter_map(Result::ok) + .filter(|entry| entry.file_type().map(|ty| ty.is_dir()).unwrap_or(false)) + .filter(|entry| { + entry + .file_name() + .to_str() + .map(|name| name.starts_with(prefix)) + .unwrap_or(false) + }) + .count(); + + Some(count) +} + +fn load_netdev_capabilities(ifindex: u32) -> Option { + let family_id = resolve_genl_family_id(NETDEV_FAMILY_NAME)?; + let payload = genl_request( + family_id, + NETDEV_CMD_DEV_GET, + &[netlink_attr_u32(NETDEV_A_DEV_IFINDEX, ifindex)], + ); + let response = netlink_round_trip(&payload)?; + let attrs = first_genl_attrs(&response)?; + + Some(NetdevCapabilities { + xdp_features: attr_u64(attrs, NETDEV_A_DEV_XDP_FEATURES), + xdp_zc_max_segs: attr_u32(attrs, NETDEV_A_DEV_XDP_ZC_MAX_SEGS), + xsk_features: attr_u64(attrs, NETDEV_A_DEV_XSK_FEATURES), + }) +} + +fn resolve_genl_family_id(name: &str) -> Option { + let payload = genl_request( + GENL_ID_CTRL, + CTRL_CMD_GETFAMILY, + &[netlink_attr_string(CTRL_ATTR_FAMILY_NAME, name)], + ); + let response = netlink_round_trip(&payload)?; + let attrs = first_genl_attrs(&response)?; + attr_u16(attrs, CTRL_ATTR_FAMILY_ID) +} + +fn netlink_round_trip(payload: &[u8]) -> Option> { + let fd = open_netlink_socket()?; + let sent = send_netlink(fd, payload).is_some(); + let response = if sent { recv_netlink(fd) } else { None }; + close_fd(fd); + response +} + +fn open_netlink_socket() -> Option { + // SAFETY: socket and bind are called with a valid sockaddr_nl and length. + let fd = unsafe { socket(AF_NETLINK, SOCK_RAW | SOCK_CLOEXEC, NETLINK_GENERIC) }; + if fd < 0 { + return None; + } + + // SAFETY: zeroed sockaddr_nl is immediately initialized before bind. + let mut addr: sockaddr_nl = unsafe { mem::zeroed() }; + addr.nl_family = AF_NETLINK as u16; + + // SAFETY: fd is a netlink socket, addr points to initialized storage. + let rc = unsafe { + bind( + fd, + &addr as *const sockaddr_nl as *const sockaddr, + mem::size_of::() as u32, + ) }; - let Some(minor_str) = parts.next() else { - return false; + if rc < 0 { + close_fd(fd); + return None; + } + + Some(fd) +} + +fn close_fd(fd: RawFd) { + // SAFETY: closing an owned file descriptor; failures are irrelevant here. + unsafe { + close(fd); + } +} + +fn send_netlink(fd: RawFd, payload: &[u8]) -> Option<()> { + // SAFETY: zeroed sockaddr_nl is immediately initialized before sendto. + let mut kernel: sockaddr_nl = unsafe { mem::zeroed() }; + kernel.nl_family = AF_NETLINK as u16; + + // SAFETY: payload is a valid byte slice and kernel points to initialized storage. + let sent = unsafe { + sendto( + fd, + payload.as_ptr().cast(), + payload.len(), + 0, + &kernel as *const sockaddr_nl as *const sockaddr, + mem::size_of::() as u32, + ) }; - let Ok(major): Result = major_str.parse() else { - return false; - }; - let Ok(minor): Result = minor_str.parse() else { - return false; - }; - major > 6 || (major == 6 && minor >= 17) + if sent == payload.len() as isize { Some(()) } else { None } +} + +fn recv_netlink(fd: RawFd) -> Option> { + let mut buf = vec![0_u8; 8192]; + // SAFETY: buf is valid writable storage for recv. + let received = unsafe { recv(fd, buf.as_mut_ptr().cast(), buf.len(), 0) }; + if received <= 0 { + return None; + } + buf.truncate(received as usize); + Some(buf) +} + +fn genl_request(nlmsg_type: u16, cmd: u8, attrs: &[Vec]) -> Vec { + let header_len = align4(mem::size_of::()); + let genl_len = mem::size_of::(); + let mut buf = vec![0_u8; header_len + genl_len]; + + write_u32(&mut buf, 0, 0); + write_u16(&mut buf, 4, nlmsg_type); + write_u16(&mut buf, 6, NLM_F_REQUEST); + write_u32(&mut buf, 8, NETLINK_SEQUENCE); + write_u32(&mut buf, 12, 0); + buf[header_len] = cmd; + buf[header_len + 1] = 1; + + for attr in attrs { + buf.extend_from_slice(attr); + } + + let nlmsg_len = buf.len() as u32; + write_u32(&mut buf, 0, nlmsg_len); + buf +} + +fn netlink_attr_string(attr_type: u16, value: &str) -> Vec { + let mut payload = value.as_bytes().to_vec(); + payload.push(0); + netlink_attr(attr_type, &payload) +} + +fn netlink_attr_u32(attr_type: u16, value: u32) -> Vec { + netlink_attr(attr_type, &value.to_ne_bytes()) +} + +fn netlink_attr(attr_type: u16, payload: &[u8]) -> Vec { + let len = 4 + payload.len(); + let mut attr = vec![0_u8; align4(len)]; + write_u16(&mut attr, 0, len as u16); + write_u16(&mut attr, 2, attr_type); + attr[4..4 + payload.len()].copy_from_slice(payload); + attr +} + +fn first_genl_attrs(response: &[u8]) -> Option<&[u8]> { + let header_len = align4(mem::size_of::()); + let genl_len = mem::size_of::(); + if response.len() < header_len + genl_len { + return None; + } + + let nlmsg_len = read_u32_ne(response, 0)? as usize; + let nlmsg_type = read_u16_ne(response, 4)?; + if nlmsg_type == NLMSG_ERROR || nlmsg_len > response.len() || nlmsg_len < header_len + genl_len { + return None; + } + + Some(&response[header_len + genl_len..nlmsg_len]) +} + +fn attr_u16(attrs: &[u8], attr_type: u16) -> Option { + find_attr(attrs, attr_type).and_then(|payload| read_u16_ne(payload, 0)) +} + +fn attr_u32(attrs: &[u8], attr_type: u16) -> Option { + find_attr(attrs, attr_type).and_then(|payload| read_u32_ne(payload, 0)) +} + +fn attr_u64(attrs: &[u8], attr_type: u16) -> Option { + find_attr(attrs, attr_type).and_then(|payload| read_u64_ne(payload, 0)) +} + +fn find_attr(attrs: &[u8], attr_type: u16) -> Option<&[u8]> { + let mut offset = 0; + while offset + 4 <= attrs.len() { + let len = read_u16_ne(attrs, offset)? as usize; + let current_type = read_u16_ne(attrs, offset + 2)?; + if len < 4 || offset + len > attrs.len() { + return None; + } + + if current_type == attr_type { + return Some(&attrs[offset + 4..offset + len]); + } + + offset += align4(len); + } + + None +} + +fn align4(value: usize) -> usize { + (value + 3) & !3 +} + +fn write_u16(buf: &mut [u8], offset: usize, value: u16) { + buf[offset..offset + 2].copy_from_slice(&value.to_ne_bytes()); +} + +fn write_u32(buf: &mut [u8], offset: usize, value: u32) { + buf[offset..offset + 4].copy_from_slice(&value.to_ne_bytes()); +} + +fn read_u16_ne(buf: &[u8], offset: usize) -> Option { + let bytes = buf.get(offset..offset + 2)?; + Some(u16::from_ne_bytes([bytes[0], bytes[1]])) +} + +fn read_u32_ne(buf: &[u8], offset: usize) -> Option { + let bytes = buf.get(offset..offset + 4)?; + Some(u32::from_ne_bytes([bytes[0], bytes[1], bytes[2], bytes[3]])) +} + +fn read_u64_ne(buf: &[u8], offset: usize) -> Option { + let bytes = buf.get(offset..offset + 8)?; + Some(u64::from_ne_bytes([ + bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7], + ])) } #[cfg(test)] mod tests { use super::*; - #[test] - fn kernel_version_matrix() { - assert!(kernel_meets_igb_af_xdp("6.17.0-generic")); - assert!(kernel_meets_igb_af_xdp("6.18.1")); - assert!(kernel_meets_igb_af_xdp("7.0.0")); - assert!(!kernel_meets_igb_af_xdp("6.16.9-generic")); - assert!(!kernel_meets_igb_af_xdp("6.12.0-124.45.1.el10_1.x86_64")); - assert!(!kernel_meets_igb_af_xdp("5.15.0")); - assert!(!kernel_meets_igb_af_xdp("nonsense")); - } - #[test] fn categorizes_permission_errors() { assert!(matches!( @@ -176,4 +481,63 @@ mod tests { EbpfFailCategory::XdpUnsupported )); } + + #[test] + fn loads_interface_capabilities_from_sysfs_shape() { + let root = std::env::temp_dir().join(format!("netguardia-ebpf-preflight-{}", std::process::id())); + let iface = root.join("eth0"); + let queues = iface.join("queues"); + let device = iface.join("device"); + let driver_target = root.join("drivers").join("virtio_net"); + + let _ = fs::remove_dir_all(&root); + fs::create_dir_all(queues.join("rx-0")).unwrap(); + fs::create_dir_all(queues.join("rx-1")).unwrap(); + fs::create_dir_all(queues.join("tx-0")).unwrap(); + fs::create_dir_all(&driver_target).unwrap(); + fs::create_dir_all(&device).unwrap(); + fs::write(iface.join("mtu"), "1500\n").unwrap(); + fs::write(iface.join("ifindex"), "not-a-number\n").unwrap(); + std::os::unix::fs::symlink(&driver_target, device.join("driver")).unwrap(); + + let caps = InterfaceCapabilities::load_from("eth0", &root); + + assert_eq!(caps.driver, "virtio_net"); + assert_eq!(caps.mtu, Some(1500)); + assert_eq!(caps.rx_queues, Some(2)); + assert_eq!(caps.tx_queues, Some(1)); + assert_eq!(caps.netdev, None); + + fs::remove_dir_all(root).unwrap(); + } + + #[test] + fn parses_netlink_attributes() { + let attrs = [ + netlink_attr_u32(NETDEV_A_DEV_IFINDEX, 7), + netlink_attr(NETDEV_A_DEV_XDP_FEATURES, &8_u64.to_ne_bytes()), + ] + .concat(); + + assert_eq!(attr_u32(&attrs, NETDEV_A_DEV_IFINDEX), Some(7)); + assert_eq!(attr_u64(&attrs, NETDEV_A_DEV_XDP_FEATURES), Some(8)); + assert_eq!(attr_u32(&attrs, NETDEV_A_DEV_XSK_FEATURES), None); + } + + #[test] + fn interprets_xsk_zerocopy_from_netdev_features() { + let caps = InterfaceCapabilities { + driver: "virtio_net".to_string(), + mtu: Some(1500), + rx_queues: Some(1), + tx_queues: Some(1), + netdev: Some(NetdevCapabilities { + xdp_features: Some(NETDEV_XDP_ACT_XSK_ZEROCOPY), + xdp_zc_max_segs: Some(1), + xsk_features: Some(0), + }), + }; + + assert_eq!(caps.xsk_zerocopy_supported(), Some(true)); + } } diff --git a/net-guardia/src/infrastructure/geoip.rs b/net-guardia/src/infrastructure/geoip.rs index f1cc33d..7fb61ef 100644 --- a/net-guardia/src/infrastructure/geoip.rs +++ b/net-guardia/src/infrastructure/geoip.rs @@ -1,5 +1,5 @@ use std::net::IpAddr; -use std::path::{Path, PathBuf}; +use std::path::Path; use std::sync::Arc; use async_trait::async_trait; @@ -17,14 +17,9 @@ pub struct GeoIpService { } impl GeoIpService { - pub fn new(db_name: &str) -> Result { - let db_path = PathBuf::from("net-guardia/static/geo").join(db_name); - Self::with_cache_size(db_path, 10000) - } - pub fn with_cache_size>(db_path: P, cache_size: usize) -> Result { let reader = Reader::open_readfile(db_path)?; - let capacity = if cache_size == 0 { 10_000 } else { cache_size } as u64; + let capacity = cache_size.max(1) as u64; Ok(Self { reader: Arc::new(reader), diff --git a/net-guardia/src/infrastructure/health.rs b/net-guardia/src/infrastructure/health.rs index ecc7e04..f482865 100644 --- a/net-guardia/src/infrastructure/health.rs +++ b/net-guardia/src/infrastructure/health.rs @@ -36,8 +36,8 @@ pub struct SystemHealth { impl SystemHealth { pub fn new(config: Arc>, ebpf_health: Arc>) -> Result { - let (broadcast_tx, _) = broadcast::channel(100); let cfg = config.load(); + let (broadcast_tx, _) = broadcast::channel(cfg.health.broadcast_channel_capacity.max(1)); let ingress_interface = cfg.ebpf.ingress_ifname.clone(); let egress_interface = cfg.ebpf.egress_ifname.clone(); diff --git a/net-guardia/src/infrastructure/logger.rs b/net-guardia/src/infrastructure/logger.rs index 8e2db99..9e8ede5 100644 --- a/net-guardia/src/infrastructure/logger.rs +++ b/net-guardia/src/infrastructure/logger.rs @@ -11,7 +11,7 @@ use tracing_subscriber::layer::SubscriberExt; use tracing_subscriber::util::SubscriberInitExt; use tracing_subscriber::{Layer, filter, reload}; -use crate::domain::common::config::observability::ObservabilityConfig; +use crate::domain::common::config::AppConfig; use crate::domain::common::error::Error; use crate::domain::common::error::io::IOError; use crate::infrastructure::log_buffer::{LogBuffer, LogBufferLayer}; @@ -26,9 +26,10 @@ pub struct Logger { } impl Logger { - pub fn initialize(config: &ObservabilityConfig) -> Result<(Self, LogBuffer), Error> { - let log_directory = "logs"; - fs::create_dir_all(log_directory).map_err(|err| IOError::CreateDirectoryFailed(log_directory, err))?; + pub fn initialize(config: &AppConfig) -> Result<(Self, LogBuffer), Error> { + let observability = &config.observability; + let log_directory = &config.system.log_dir; + fs::create_dir_all(log_directory).map_err(|err| IOError::CreateDirectoryFailed(log_directory.clone(), err))?; let file_appender = RollingFileAppender::new(Rotation::DAILY, log_directory, "NetGuardia"); @@ -47,7 +48,7 @@ impl Logger { .with_ansi(false) .with_writer(file_appender); - let level: Level = config + let level: Level = observability .log_level .parse() .ok() @@ -71,8 +72,10 @@ impl Logger { filter = Self::apply_directives(filter, &preserved); let (filter_layer, reload_handle) = reload::Layer::new(filter); - let (log_buffer_layer, log_buffer) = - LogBufferLayer::new(config.log_buffer_capacity, config.log_buffer_max_message_bytes); + let (log_buffer_layer, log_buffer) = LogBufferLayer::new( + observability.log_buffer_capacity, + observability.log_buffer_max_message_bytes, + ); tracing_subscriber::registry() .with(filter_layer) diff --git a/net-guardia/src/infrastructure/service_factory.rs b/net-guardia/src/infrastructure/service_factory.rs index a070c65..c5d9e15 100644 --- a/net-guardia/src/infrastructure/service_factory.rs +++ b/net-guardia/src/infrastructure/service_factory.rs @@ -44,6 +44,7 @@ use crate::domain::common::system::health::EbpfFailStage; use crate::domain::common::system::health::EbpfHealth; use crate::domain::data_plane::direction::FlowDirection; use crate::domain::data_plane::error::EbpfError; +use crate::domain::data_plane::flow_stats::FlowStatsLimits; use crate::domain::data_plane::list_type::ListType; use crate::domain::data_plane::log::EbpfLog; use crate::domain::detection::drift::FeatureBaselines; @@ -64,6 +65,7 @@ use crate::interface::config_repo::ConfigRepo; use crate::interface::dns_filter_api::DnsFilterPort; use crate::interface::dns_query_filter::DnsQueryFilter; use crate::interface::email_sender::EmailSenderFactory; +use crate::interface::enforcement::EnforcementRepo; use crate::interface::geo_block_api::GeoBlockPort; use crate::interface::geo_lookup::GeoLookup; use crate::interface::notification::{AlertNotifier, AlertNotifierFactory}; @@ -243,7 +245,16 @@ impl ServiceFactory { audit_tx.clone(), )?); - let flow_statistics = Arc::new(FlowStatistics::new(inference_runtime.ml_engine.clone())); + let flow_stats_cfg = app_config.load(); + let flow_stats_limits = FlowStatsLimits::new( + flow_stats_cfg.detection.flow_stats.max_snapshot_entries, + flow_stats_cfg.detection.flow_stats.max_top_n, + ); + drop(flow_stats_cfg); + let flow_statistics = Arc::new(FlowStatistics::new( + inference_runtime.ml_engine.clone(), + flow_stats_limits, + )); let enforce_handler = Arc::new(EnforceModeHandler::new( db.clone() as Arc, @@ -275,16 +286,18 @@ impl ServiceFactory { }; // Try to initialize GeoIP service - let geoip: Option> = match GeoIpService::new(&app_config.load().acl.geoip_db_name) { - Ok(svc) => { - log!(SystemLog::GeoIpInitialized); - Some(Arc::new(svc)) - } - Err(e) => { - log!(SystemLog::GeoIpUnavailable(e.to_string())); - None - } - }; + let acl_cfg = app_config.load().acl.clone(); + let geoip: Option> = + match GeoIpService::with_cache_size(&acl_cfg.geoip_db_path, acl_cfg.geoip_cache_capacity) { + Ok(svc) => { + log!(SystemLog::GeoIpInitialized); + Some(Arc::new(svc)) + } + Err(e) => { + log!(SystemLog::GeoIpUnavailable(e.to_string())); + None + } + }; // Create AccessControlPort adapter for SOAR/TTL (decoupled from eBPF) let access_control_port: Arc = @@ -330,7 +343,7 @@ impl ServiceFactory { geo_block_port, )); let dns_filter_service = Arc::new(DnsFilterService::new( - db.clone() as Arc, + db.clone() as Arc, dns_filter_port, app_config.clone(), )); diff --git a/net-guardia/src/infrastructure/system.rs b/net-guardia/src/infrastructure/system.rs index 81eb910..051fc8b 100644 --- a/net-guardia/src/infrastructure/system.rs +++ b/net-guardia/src/infrastructure/system.rs @@ -296,7 +296,12 @@ impl System { } async fn boot_observability(&mut self) { - let health_shutdown = self.health.clone().run(Duration::from_secs(3)).await; + let monitoring_interval_secs = self.app_config.load().health.monitoring_interval_secs; + let health_shutdown = self + .health + .clone() + .run(Duration::from_secs(monitoring_interval_secs)) + .await; self.health_shutdown = Some(health_shutdown); let audit_logger = Arc::new(AuditLogger::new(self.database.clone() as Arc)); diff --git a/net-guardia/src/interface/dns_filter_api.rs b/net-guardia/src/interface/dns_filter_api.rs index 729bed7..7f23836 100644 --- a/net-guardia/src/interface/dns_filter_api.rs +++ b/net-guardia/src/interface/dns_filter_api.rs @@ -5,6 +5,7 @@ use crate::domain::common::error::Error; /// `DnsQueryFilter` (which is the fast-path check) to reflect their distinct /// call sites and latency profiles. pub trait DnsFilterPort: Send + Sync { + fn validate_domain(&self, domain: &str) -> Result<(), Error>; fn add_domain(&self, domain: &str) -> Result<(), Error>; fn remove_domain(&self, domain: &str) -> Result<(), Error>; fn list_domains(&self) -> Vec; diff --git a/net-guardia/src/interface/enforcement.rs b/net-guardia/src/interface/enforcement.rs index 3a2e4d2..79595e3 100644 --- a/net-guardia/src/interface/enforcement.rs +++ b/net-guardia/src/interface/enforcement.rs @@ -14,8 +14,8 @@ pub trait EnforcementRepo: Send + Sync { async fn set_rate_limit(&self, key: &str, value: u64) -> Result<(), Error>; // --- DNS --- - async fn insert_dns_domain(&self, domain: &str) -> Result<(), Error>; - async fn delete_dns_domain(&self, domain: &str) -> Result<(), Error>; + async fn insert_dns_domains(&self, domains: &[String]) -> Result<(), Error>; + async fn delete_dns_domains(&self, domains: &[String]) -> Result<(), Error>; // --- Geo --- async fn insert_geo_country(&self, code: &str) -> Result<(), Error>; diff --git a/net-guardia/src/main.rs b/net-guardia/src/main.rs index 272ca9c..575571d 100644 --- a/net-guardia/src/main.rs +++ b/net-guardia/src/main.rs @@ -22,7 +22,7 @@ use tokio::{signal, time}; use crate::adapter::http::jwt::JwtService; use crate::adapter::persistence::Database; use crate::core::identity::auth_service::AuthService; -use crate::domain::common::config::observability::ObservabilityConfig; +use crate::domain::common::config::AppConfig; use crate::domain::common::error::Error; use crate::domain::common::error::system::SystemError; use crate::domain::common::log::system::SystemLog; @@ -145,8 +145,8 @@ async fn main() -> Result<(), Error> { } let database = Arc::new(Database::new(&cli.db_path).await?); - let obs_config = ObservabilityConfig::from_config_repo(database.as_ref()).await?; - let (logger, log_buffer) = Logger::initialize(&obs_config)?; + let app_config = AppConfig::from_config_repo(database.as_ref()).await?; + let (logger, log_buffer) = Logger::initialize(&app_config)?; let logger = Arc::new(logger); let log_buffer = Arc::new(log_buffer);