mirror of
https://github.com/DaLaw2/NetGuardia.git
synced 2026-08-24 14:10:28 +09:00
refactor: centralize runtime configuration state
This commit is contained in:
parent
d419e3ee29
commit
fd95405ad8
@ -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<ArcSwap<AppConfig>>) -> 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 {
|
||||
|
||||
@ -183,6 +183,9 @@ pub struct XskPair {
|
||||
drop_monitor: Option<Arc<DropMonitor>>,
|
||||
packet_buffer_size: usize,
|
||||
buffer_pool_capacity: usize,
|
||||
completion_batch_size: usize,
|
||||
rx_batch_size: usize,
|
||||
tx_batch_size: usize,
|
||||
tx_packet_buf: Vec<Vec<u8>>,
|
||||
tx_frame_buf: Vec<FrameDesc>,
|
||||
}
|
||||
@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
@ -584,26 +584,32 @@ async fn validate_and_promote(ctx: &PromoteContext<'_>) -> Result<PromoteReport,
|
||||
|
||||
let before_status = inference.model_source_status();
|
||||
|
||||
let _guard = promote_lock.try_acquire().ok_or(PromoteError::ConcurrentPromote)?;
|
||||
let models_dir = PathBuf::from(MODELS_DIR);
|
||||
let target_onnx = models_dir.join(&declared_onnx);
|
||||
fs::rename(&staged_onnx, &target_onnx)
|
||||
.await
|
||||
.map_err(|e| PromoteError::PromoteIo(format!("rename onnx into models/: {e}")))?;
|
||||
|
||||
if let Some(ref pp) = manifest.preprocessing {
|
||||
validate_manifest_basename("preprocessing.scaler_sidecar", &pp.scaler_sidecar)?;
|
||||
let src = staging_dir.join(&pp.scaler_sidecar);
|
||||
let dst = models_dir.join(&pp.scaler_sidecar);
|
||||
fs::rename(&src, &dst)
|
||||
.await
|
||||
.map_err(|e| PromoteError::PromoteIo(format!("rename sidecar into models/: {e}")))?;
|
||||
}
|
||||
|
||||
let target_manifest = models_dir.join(MANIFEST_FILENAME);
|
||||
fs::rename(&staging_manifest, &target_manifest)
|
||||
.await
|
||||
.map_err(|e| PromoteError::PromoteIo(format!("rename manifest into models/: {e}")))?;
|
||||
let staged_sidecar = if let Some(ref pp) = manifest.preprocessing {
|
||||
validate_manifest_basename("preprocessing.scaler_sidecar", &pp.scaler_sidecar)?;
|
||||
Some(staging_dir.join(&pp.scaler_sidecar))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let target_sidecar = manifest
|
||||
.preprocessing
|
||||
.as_ref()
|
||||
.map(|pp| models_dir.join(&pp.scaler_sidecar));
|
||||
|
||||
let _guard = promote_lock.try_acquire().ok_or(PromoteError::ConcurrentPromote)?;
|
||||
let backup_dir = models_dir.join(format!(".promote-backup-{}", Uuid::new_v4()));
|
||||
promote_files_atomically(&PromoteFileSet {
|
||||
staging_manifest: staging_manifest.clone(),
|
||||
staged_onnx,
|
||||
staged_sidecar,
|
||||
target_manifest,
|
||||
target_onnx,
|
||||
target_sidecar,
|
||||
backup_dir,
|
||||
})
|
||||
.await?;
|
||||
drop(_guard);
|
||||
|
||||
let audit_detail = serde_json::json!({
|
||||
@ -628,6 +634,105 @@ async fn validate_and_promote(ctx: &PromoteContext<'_>) -> Result<PromoteReport,
|
||||
})
|
||||
}
|
||||
|
||||
struct PromoteFileSet {
|
||||
staging_manifest: PathBuf,
|
||||
staged_onnx: PathBuf,
|
||||
staged_sidecar: Option<PathBuf>,
|
||||
target_manifest: PathBuf,
|
||||
target_onnx: PathBuf,
|
||||
target_sidecar: Option<PathBuf>,
|
||||
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<Option<PathBuf>, 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<PathBuf>,
|
||||
backup_onnx: Option<PathBuf>,
|
||||
backup_sidecar: Option<Option<PathBuf>>,
|
||||
) {
|
||||
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<PathBuf>, 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();
|
||||
|
||||
@ -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<u64>,
|
||||
}
|
||||
|
||||
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<ArcSwap<AppConfig>>) -> 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<String>, app_config: web::Data<ArcSwap<AppConfig>>) -> 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<String>, app_config: web::Data<ArcSwap<App
|
||||
}));
|
||||
}
|
||||
|
||||
let file_path = Path::new(LOG_DIR).join(&filename);
|
||||
let file_path = log_dir.join(&filename);
|
||||
|
||||
// Canonicalize to prevent symlink traversal
|
||||
let canonical = match fs::canonicalize(&file_path) {
|
||||
@ -137,7 +136,7 @@ async fn download_log(path: web::Path<String>, app_config: web::Data<ArcSwap<App
|
||||
}));
|
||||
}
|
||||
};
|
||||
if let Ok(log_dir_canonical) = fs::canonicalize(LOG_DIR)
|
||||
if let Ok(log_dir_canonical) = fs::canonicalize(&log_dir)
|
||||
&& !canonical.starts_with(&log_dir_canonical)
|
||||
{
|
||||
return HttpResponse::Forbidden().json(serde_json::json!({
|
||||
|
||||
@ -63,7 +63,7 @@ fn required_permission(path: &str, method: &Method) -> Option<String> {
|
||||
"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())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@ -162,8 +162,12 @@ async fn manual_unblock(_auth: AuthClaims, svc: web::Data<PlaybookService>, path
|
||||
ok_or_error(svc.manual_unblock(path.into_inner()).await)
|
||||
}
|
||||
|
||||
async fn list_executions(_auth: AuthClaims, svc: web::Data<PlaybookService>) -> HttpResponse {
|
||||
ok_json_or_error(svc.list_executions(100).await)
|
||||
async fn list_executions(
|
||||
_auth: AuthClaims,
|
||||
svc: web::Data<PlaybookService>,
|
||||
app_config: web::Data<ArcSwap<AppConfig>>,
|
||||
) -> 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<PlaybookService>) -> HttpResponse {
|
||||
|
||||
@ -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()));
|
||||
}
|
||||
}
|
||||
|
||||
@ -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> {
|
||||
|
||||
@ -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();
|
||||
}
|
||||
}
|
||||
|
||||
@ -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<Engine>,
|
||||
limits: FlowStatsLimits,
|
||||
}
|
||||
|
||||
impl FlowStatistics {
|
||||
pub fn new(engine: Arc<Engine>) -> Self {
|
||||
Self { engine }
|
||||
pub fn new(engine: Arc<Engine>, limits: FlowStatsLimits) -> Self {
|
||||
Self { engine, limits }
|
||||
}
|
||||
|
||||
pub fn get_all_flows(&self) -> Vec<FlowStatsEntry> {
|
||||
self.collect_flows(Some(self.limits.result_limit(None)))
|
||||
}
|
||||
|
||||
fn collect_flows(&self, limit: Option<usize>) -> Vec<FlowStatsEntry> {
|
||||
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();
|
||||
|
||||
@ -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)
|
||||
}
|
||||
|
||||
@ -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<dyn AppRepo>,
|
||||
db: Arc<dyn EnforcementRepo>,
|
||||
dns_filter: Arc<dyn DnsFilterPort>,
|
||||
config: Arc<ArcSwap<AppConfig>>,
|
||||
}
|
||||
|
||||
impl DnsFilterService {
|
||||
pub fn new(db: Arc<dyn AppRepo>, dns_filter: Arc<dyn DnsFilterPort>, config: Arc<ArcSwap<AppConfig>>) -> Self {
|
||||
pub fn new(
|
||||
db: Arc<dyn EnforcementRepo>,
|
||||
dns_filter: Arc<dyn DnsFilterPort>,
|
||||
config: Arc<ArcSwap<AppConfig>>,
|
||||
) -> 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<usize, Error> {
|
||||
// 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<Vec<String>>,
|
||||
}
|
||||
|
||||
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<ArcSwap<AppConfig>> {
|
||||
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()));
|
||||
}
|
||||
}
|
||||
|
||||
@ -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<Instant>,
|
||||
last_alerted: Option<Instant>,
|
||||
@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@ -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::<usize>::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<usize> {
|
||||
NonZero::new(value.max(1)).unwrap_or(NonZero::<usize>::MIN)
|
||||
}
|
||||
|
||||
pub async fn bridge_ml_to_detection(mut rx: broadcast::Receiver<AlertMessage>, tx: mpsc::Sender<DetectionEvent>) {
|
||||
log!(DetectionLog::MlBridgeStarted);
|
||||
|
||||
|
||||
@ -34,6 +34,7 @@ pub struct UserProfile {
|
||||
pub groups: Vec<String>,
|
||||
}
|
||||
|
||||
#[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<Database>, Arc<JwtService>, AuthService) {
|
||||
let db = Arc::new(Database::new(":memory:").await.expect("test db"));
|
||||
let secrets: Arc<dyn SecretStorePort> = Arc::new(SecretStore::new(db.clone()));
|
||||
let jwt = Arc::new(JwtService::new(&secrets, 24).expect("jwt"));
|
||||
let auth = AuthService::new(db.clone() as Arc<dyn AppRepo>, 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()));
|
||||
}
|
||||
}
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -11,14 +11,16 @@ pub struct FrequencyTracker {
|
||||
events: DashMap<FreqKey, VecDeque<Instant>>,
|
||||
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);
|
||||
}
|
||||
}
|
||||
|
||||
@ -29,12 +29,21 @@ pub struct PlaybookMatcher {
|
||||
}
|
||||
|
||||
impl PlaybookMatcher {
|
||||
pub fn new(config: Arc<ArcSwap<AppConfig>>, frequency_max_tracked_keys: usize) -> Self {
|
||||
pub fn new(
|
||||
config: Arc<ArcSwap<AppConfig>>,
|
||||
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,
|
||||
}
|
||||
|
||||
@ -116,3 +116,204 @@ async fn unblock_ip_blocking(access_control: Arc<dyn AccessControlPort>, 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<Vec<String>>,
|
||||
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<Option<Box<dyn EmailSender>>, Error> {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
|
||||
async fn test_db() -> Arc<Database> {
|
||||
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<Database>,
|
||||
access_control: Arc<MockAccessControl>,
|
||||
) -> (TtlScheduler, Arc<SoarEngine>) {
|
||||
let cfg = AppConfig::from_config_repo(&*db).await.expect("load config");
|
||||
let engine = Arc::new(
|
||||
SoarEngine::new(SoarEngineDeps {
|
||||
db: db.clone() as Arc<dyn AppRepo>,
|
||||
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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@ -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,
|
||||
}
|
||||
|
||||
@ -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,
|
||||
}
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
}
|
||||
|
||||
@ -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,
|
||||
|
||||
|
||||
@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@ -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,
|
||||
}
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
}
|
||||
}
|
||||
@ -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,
|
||||
|
||||
@ -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),
|
||||
|
||||
@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
@ -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,
|
||||
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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>) -> 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<u64>,
|
||||
}
|
||||
|
||||
#[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);
|
||||
}
|
||||
}
|
||||
|
||||
@ -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<String>,
|
||||
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());
|
||||
}
|
||||
}
|
||||
|
||||
@ -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/<ifname>/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<u64>,
|
||||
pub xdp_zc_max_segs: Option<u32>,
|
||||
pub xsk_features: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct InterfaceCapabilities {
|
||||
pub driver: String,
|
||||
pub mtu: Option<u32>,
|
||||
pub rx_queues: Option<usize>,
|
||||
pub tx_queues: Option<usize>,
|
||||
pub netdev: Option<NetdevCapabilities>,
|
||||
}
|
||||
|
||||
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<bool> {
|
||||
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<u32> {
|
||||
fs::read_to_string(path).ok()?.trim().parse().ok()
|
||||
}
|
||||
|
||||
fn count_queue_dirs(iface_path: &Path, prefix: &str) -> Option<usize> {
|
||||
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<NetdevCapabilities> {
|
||||
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<u16> {
|
||||
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<Vec<u8>> {
|
||||
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<RawFd> {
|
||||
// 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::<sockaddr_nl>() 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::<sockaddr_nl>() as u32,
|
||||
)
|
||||
};
|
||||
let Ok(major): Result<u32, _> = major_str.parse() else {
|
||||
return false;
|
||||
};
|
||||
let Ok(minor): Result<u32, _> = 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<Vec<u8>> {
|
||||
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<u8>]) -> Vec<u8> {
|
||||
let header_len = align4(mem::size_of::<nlmsghdr>());
|
||||
let genl_len = mem::size_of::<genlmsghdr>();
|
||||
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<u8> {
|
||||
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<u8> {
|
||||
netlink_attr(attr_type, &value.to_ne_bytes())
|
||||
}
|
||||
|
||||
fn netlink_attr(attr_type: u16, payload: &[u8]) -> Vec<u8> {
|
||||
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::<nlmsghdr>());
|
||||
let genl_len = mem::size_of::<genlmsghdr>();
|
||||
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<u16> {
|
||||
find_attr(attrs, attr_type).and_then(|payload| read_u16_ne(payload, 0))
|
||||
}
|
||||
|
||||
fn attr_u32(attrs: &[u8], attr_type: u16) -> Option<u32> {
|
||||
find_attr(attrs, attr_type).and_then(|payload| read_u32_ne(payload, 0))
|
||||
}
|
||||
|
||||
fn attr_u64(attrs: &[u8], attr_type: u16) -> Option<u64> {
|
||||
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<u16> {
|
||||
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<u32> {
|
||||
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<u64> {
|
||||
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));
|
||||
}
|
||||
}
|
||||
|
||||
@ -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<Self, MaxMindDbError> {
|
||||
let db_path = PathBuf::from("net-guardia/static/geo").join(db_name);
|
||||
Self::with_cache_size(db_path, 10000)
|
||||
}
|
||||
|
||||
pub fn with_cache_size<P: AsRef<Path>>(db_path: P, cache_size: usize) -> Result<Self, MaxMindDbError> {
|
||||
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),
|
||||
|
||||
@ -36,8 +36,8 @@ pub struct SystemHealth {
|
||||
|
||||
impl SystemHealth {
|
||||
pub fn new(config: Arc<ArcSwap<AppConfig>>, ebpf_health: Arc<ArcSwap<EbpfHealth>>) -> Result<Self, Error> {
|
||||
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();
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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<dyn AppRepo>,
|
||||
@ -275,16 +286,18 @@ impl ServiceFactory {
|
||||
};
|
||||
|
||||
// Try to initialize GeoIP service
|
||||
let geoip: Option<Arc<dyn GeoLookup>> = 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<Arc<dyn GeoLookup>> =
|
||||
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<dyn AccessControlPort> =
|
||||
@ -330,7 +343,7 @@ impl ServiceFactory {
|
||||
geo_block_port,
|
||||
));
|
||||
let dns_filter_service = Arc::new(DnsFilterService::new(
|
||||
db.clone() as Arc<dyn AppRepo>,
|
||||
db.clone() as Arc<dyn EnforcementRepo>,
|
||||
dns_filter_port,
|
||||
app_config.clone(),
|
||||
));
|
||||
|
||||
@ -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<dyn AuditRepo>));
|
||||
|
||||
@ -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<String>;
|
||||
|
||||
@ -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>;
|
||||
|
||||
@ -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);
|
||||
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user