From 4bfd00d12a0406a52ce08604b4bae38574f855c8 Mon Sep 17 00:00:00 2001 From: DaLaw2 Date: Tue, 21 Apr 2026 20:36:13 +0800 Subject: [PATCH] chore: Remove/Fix AI trash --- .gitignore | 16 +- Cargo.lock | 1 + cli/src/main.rs | 201 +- common/src/model/port_rule.rs | 5 +- egress-ebpf/src/main.rs | 4 +- ingress-ebpf/src/action/access_control.rs | 4 +- ingress-ebpf/src/action/mod.rs | 2 +- ingress-ebpf/src/action/rate_limit.rs | 76 +- ingress-ebpf/src/main.rs | 16 +- macros/src/log.rs | 4 +- mcp-server/src/main.rs | 91 +- net-guardia/Cargo.toml | 3 + ...s_control_adapter.rs => access_control.rs} | 6 +- .../src/adapter/ebpf/access_control.rs | 8 +- net-guardia/src/adapter/ebpf/dns_filter.rs | 5 - net-guardia/src/adapter/ebpf/drop_monitor.rs | 5 +- net-guardia/src/adapter/ebpf/geo_block.rs | 2 + net-guardia/src/adapter/ebpf/xsk_manager.rs | 3 +- net-guardia/src/adapter/mod.rs | 2 +- net-guardia/src/adapter/persistence/acl.rs | 158 ++ .../src/adapter/persistence/api_key.rs | 156 ++ net-guardia/src/adapter/persistence/audit.rs | 139 + .../src/adapter/persistence/enforcement.rs | 140 + net-guardia/src/adapter/persistence/mod.rs | 449 +++- .../src/adapter/persistence/repository.rs | 2305 ----------------- .../src/adapter/persistence/setting.rs | 115 + net-guardia/src/adapter/persistence/soar.rs | 413 +++ .../src/adapter/persistence/soar_block.rs | 229 ++ net-guardia/src/adapter/persistence/stats.rs | 105 + net-guardia/src/adapter/persistence/user.rs | 530 ++++ .../adapter/{telegram/mod.rs => telegram.rs} | 0 net-guardia/src/core/acl_service.rs | 6 +- net-guardia/src/core/auth/middleware.rs | 1 - net-guardia/src/core/ml/model_watcher.rs | 4 +- net-guardia/src/core/report/data.rs | 1 - net-guardia/src/core/report/mod.rs | 1 - net-guardia/src/core/soar/engine.rs | 8 +- net-guardia/src/infrastructure/cli.rs | 108 + .../infrastructure/communication_manager.rs | 13 +- .../infrastructure/enforce_mode_handler.rs | 2 +- net-guardia/src/infrastructure/mod.rs | 1 + .../src/infrastructure/secret_store.rs | 3 +- .../src/infrastructure/service_factory.rs | 10 +- .../src/interface/communication/command.rs | 8 +- .../src/interface/communication/event.rs | 3 +- .../src/interface/communication/query.rs | 8 +- net-guardia/src/interface/mod.rs | 1 + net-guardia/src/interface/utils/logging.rs | 6 + net-guardia/src/interface/utils/mod.rs | 1 + net-guardia/src/main.rs | 105 +- net-guardia/src/model/log/cli.rs | 38 + net-guardia/src/model/log/mod.rs | 1 + net-guardia/src/model/log/system.rs | 8 +- net-guardia/src/utils/ip_address.rs | 7 +- net-guardia/src/utils/logging.rs | 61 +- 55 files changed, 2961 insertions(+), 2637 deletions(-) rename net-guardia/src/adapter/{access_control_adapter.rs => access_control.rs} (94%) create mode 100644 net-guardia/src/adapter/persistence/acl.rs create mode 100644 net-guardia/src/adapter/persistence/api_key.rs create mode 100644 net-guardia/src/adapter/persistence/audit.rs create mode 100644 net-guardia/src/adapter/persistence/enforcement.rs delete mode 100644 net-guardia/src/adapter/persistence/repository.rs create mode 100644 net-guardia/src/adapter/persistence/setting.rs create mode 100644 net-guardia/src/adapter/persistence/soar.rs create mode 100644 net-guardia/src/adapter/persistence/soar_block.rs create mode 100644 net-guardia/src/adapter/persistence/stats.rs create mode 100644 net-guardia/src/adapter/persistence/user.rs rename net-guardia/src/adapter/{telegram/mod.rs => telegram.rs} (100%) delete mode 100644 net-guardia/src/core/report/data.rs create mode 100644 net-guardia/src/infrastructure/cli.rs create mode 100644 net-guardia/src/interface/utils/logging.rs create mode 100644 net-guardia/src/interface/utils/mod.rs create mode 100644 net-guardia/src/model/log/cli.rs diff --git a/.gitignore b/.gitignore index 306a301..29a2949 100644 --- a/.gitignore +++ b/.gitignore @@ -30,20 +30,12 @@ net-guardia/static/web *.profraw *.profdata -# License keys -license-generator/target/ *.hex -license.key -license_priv.key -license_pub.key .gstack/ interfaces.txt traffic_log.csv -# Project docs (local only) -# CLAUDE.md — tracked on dev branches; MUST be untracked before PR to master -# (see CLAUDE.md "Branch discipline" section) -# CLAUDE.md +CLAUDE.md DESIGN.md TODOS.md VERSION @@ -51,11 +43,7 @@ CHANGELOG.md # Benchmark data/results (local only) benchmark/ - -# Generated docs -# docs/ — tracked on dev branches; MUST be untracked before PR to master -# (see CLAUDE.md "Branch discipline" section) -# docs/ +docs/ # SQLite database files *.db diff --git a/Cargo.lock b/Cargo.lock index cc41b96..694e787 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2495,6 +2495,7 @@ dependencies = [ "base64", "cargo_metadata", "chrono", + "clap", "common", "crossbeam", "dashmap", diff --git a/cli/src/main.rs b/cli/src/main.rs index f1c6f67..f3656ee 100644 --- a/cli/src/main.rs +++ b/cli/src/main.rs @@ -24,9 +24,7 @@ enum Commands { /// ML engine status Ml, /// Add IP to source blacklist - Block { - ip: String, - }, + Block { ip: String }, /// Remove IP from source blacklist Unblock { ip: String }, /// List ACL rules (source blacklist by default) @@ -88,7 +86,11 @@ impl ApiClient { let token_path = dirs_next().join("token"); - Self { client, base_url, token_path } + Self { + client, + base_url, + token_path, + } } fn load_token(&self) -> Option { @@ -114,7 +116,13 @@ impl ApiClient { return Err("Session expired. Run `ng login` to re-authenticate.".into()); } let text = resp.text().await.map_err(|e| format!("Read error: {}", e))?; - serde_json::from_str(&text).map_err(|_| format!("Unexpected response (HTTP {}): {}", status, &text[..text.len().min(200)])) + serde_json::from_str(&text).map_err(|_| { + format!( + "Unexpected response (HTTP {}): {}", + status, + &text[..text.len().min(200)] + ) + }) } async fn request(&self, method: reqwest::Method, path: &str, body: Option) -> Result { @@ -138,17 +146,35 @@ impl ApiClient { } return Err(format!("Empty response (HTTP {})", status)); } - serde_json::from_str(&text).map_err(|_| format!("Unexpected response (HTTP {}): {}", status, &text[..text.len().min(200)])) + serde_json::from_str(&text).map_err(|_| { + format!( + "Unexpected response (HTTP {}): {}", + status, + &text[..text.len().min(200)] + ) + }) } async fn login(&self, username: &str, password: &str) -> Result { let url = format!("{}/api/auth/login", self.base_url); let body = serde_json::json!({"username": username, "password": password}); - let resp = self.client.post(&url).json(&body).send().await + let resp = self + .client + .post(&url) + .json(&body) + .send() + .await .map_err(|e| format!("Connection error: {}", e))?; let data: Value = resp.json().await.map_err(|e| format!("Parse error: {}", e))?; - data.get("token").and_then(|t| t.as_str()).map(|s| s.to_string()) - .ok_or_else(|| data.get("error").and_then(|e| e.as_str()).unwrap_or("Login failed").to_string()) + data.get("token") + .and_then(|t| t.as_str()) + .map(|s| s.to_string()) + .ok_or_else(|| { + data.get("error") + .and_then(|e| e.as_str()) + .unwrap_or("Login failed") + .to_string() + }) } } @@ -194,25 +220,39 @@ async fn main() { let api = ApiClient::new(cli.url); let result = match cli.command { - Commands::Status => { - api.get("/api/health/status").await.map(|d| print_json(&d)) - } - Commands::Ml => { - api.get("/api/ml/status").await.map(|d| print_json(&d)) - } + Commands::Status => api.get("/api/health/status").await.map(|d| print_json(&d)), + Commands::Ml => api.get("/api/ml/status").await.map(|d| print_json(&d)), Commands::Block { ip } => { let is_v6 = ip.contains(':'); let ip_ver = if is_v6 { "ipv6" } else { "ipv4" }; - let addr = if is_v6 { format!("[{}]:0", ip) } else { format!("{}:0", ip) }; - api.request(reqwest::Method::PUT, &format!("/api/acl/{}/source/blacklist", ip_ver), Some(Value::String(addr))) - .await.map(|_| println!("Blocked: {}", ip)) + let addr = if is_v6 { + format!("[{}]:0", ip) + } else { + format!("{}:0", ip) + }; + api.request( + reqwest::Method::PUT, + &format!("/api/acl/{}/source/blacklist", ip_ver), + Some(Value::String(addr)), + ) + .await + .map(|_| println!("Blocked: {}", ip)) } Commands::Unblock { ip } => { let is_v6 = ip.contains(':'); let ip_ver = if is_v6 { "ipv6" } else { "ipv4" }; - let addr = if is_v6 { format!("[{}]:0", ip) } else { format!("{}:0", ip) }; - api.request(reqwest::Method::DELETE, &format!("/api/acl/{}/source/blacklist", ip_ver), Some(Value::String(addr))) - .await.map(|_| println!("Unblocked: {}", ip)) + let addr = if is_v6 { + format!("[{}]:0", ip) + } else { + format!("{}:0", ip) + }; + api.request( + reqwest::Method::DELETE, + &format!("/api/acl/{}/source/blacklist", ip_ver), + Some(Value::String(addr)), + ) + .await + .map(|_| println!("Unblocked: {}", ip)) } Commands::Rules { direction, list_type } => { // Try both IPv4 and IPv6 @@ -234,18 +274,15 @@ async fn main() { // Use /api/report/data for JSON output api.get("/api/report/data").await.map(|d| print_json(&d)) } - Commands::Mode { mode } => { - match mode { - Some(m) => { - let body = serde_json::json!({"mode": m}); - api.request(reqwest::Method::PUT, "/api/system/enforce-mode", Some(body)) - .await.map(|d| print_json(&d)) - } - None => { - api.get("/api/system/enforce-mode").await.map(|d| print_json(&d)) - } + Commands::Mode { mode } => match mode { + Some(m) => { + let body = serde_json::json!({"mode": m}); + api.request(reqwest::Method::PUT, "/api/system/enforce-mode", Some(body)) + .await + .map(|d| print_json(&d)) } - } + None => api.get("/api/system/enforce-mode").await.map(|d| print_json(&d)), + }, Commands::Login => { print!("Username: "); std::io::Write::flush(&mut std::io::stdout()).unwrap(); @@ -266,65 +303,57 @@ async fn main() { Err(e) => Err(e), } } - Commands::Blocks => { - api.get("/api/soar/blocks").await.map(|d| print_json(&d)) - } - Commands::Playbooks => { - api.get("/api/soar/playbooks").await.map(|d| print_json(&d)) - } - Commands::Executions => { - api.get("/api/soar/executions").await.map(|d| print_json(&d)) - } - Commands::ApiKey { action } => { - match action { - ApiKeyAction::Generate { name, level } => { - let body = serde_json::json!({"name": name, "level": level}); - api.request(reqwest::Method::POST, "/api/api-keys/generate", Some(body)) - .await.map(|data| { - if let Some(key) = data.get("key").and_then(|k| k.as_str()) { - println!("Generated API key: {}", key); - println!("Name: {}, Level: {}", name, level); - println!("Set NETGUARDIA_API_KEY={} in your client config", key); - } else { - print_json(&data); - } - }) - } - ApiKeyAction::List => { - api.get("/api/api-keys").await.map(|data| { - if let Some(keys) = data.as_array() { - if keys.is_empty() { - println!("No API keys found."); - } else { - println!("{:<6} {:<20} {:<15} {:<22} Last Used", "ID", "Name", "Level", "Created"); - println!("{}", "-".repeat(80)); - for key in keys { - println!("{:<6} {:<20} {:<15} {:<22} {}", - key.get("id").and_then(|v| v.as_i64()).unwrap_or(0), - key.get("name").and_then(|v| v.as_str()).unwrap_or("-"), - key.get("permission_level").and_then(|v| v.as_str()).unwrap_or("-"), - key.get("created_at").and_then(|v| v.as_str()).unwrap_or("-"), - key.get("last_used_at").and_then(|v| v.as_str()).unwrap_or("never"), - ); - } - } + Commands::Blocks => api.get("/api/soar/blocks").await.map(|d| print_json(&d)), + Commands::Playbooks => api.get("/api/soar/playbooks").await.map(|d| print_json(&d)), + Commands::Executions => api.get("/api/soar/executions").await.map(|d| print_json(&d)), + Commands::ApiKey { action } => match action { + ApiKeyAction::Generate { name, level } => { + let body = serde_json::json!({"name": name, "level": level}); + api.request(reqwest::Method::POST, "/api/api-keys/generate", Some(body)) + .await + .map(|data| { + if let Some(key) = data.get("key").and_then(|k| k.as_str()) { + println!("Generated API key: {}", key); + println!("Name: {}, Level: {}", name, level); + println!("Set NETGUARDIA_API_KEY={} in your client config", key); } else { print_json(&data); } }) - } - ApiKeyAction::Revoke { id } => { - api.request(reqwest::Method::DELETE, &format!("/api/api-keys/{}", id), None) - .await.map(|data| { - if data.get("deleted").and_then(|v| v.as_bool()).unwrap_or(false) { - println!("Key #{} revoked successfully.", id); - } else { - print_json(&data); - } - }) - } } - } + ApiKeyAction::List => api.get("/api/api-keys").await.map(|data| { + if let Some(keys) = data.as_array() { + if keys.is_empty() { + println!("No API keys found."); + } else { + println!("{:<6} {:<20} {:<15} {:<22} Last Used", "ID", "Name", "Level", "Created"); + println!("{}", "-".repeat(80)); + for key in keys { + println!( + "{:<6} {:<20} {:<15} {:<22} {}", + key.get("id").and_then(|v| v.as_i64()).unwrap_or(0), + key.get("name").and_then(|v| v.as_str()).unwrap_or("-"), + key.get("permission_level").and_then(|v| v.as_str()).unwrap_or("-"), + key.get("created_at").and_then(|v| v.as_str()).unwrap_or("-"), + key.get("last_used_at").and_then(|v| v.as_str()).unwrap_or("never"), + ); + } + } + } else { + print_json(&data); + } + }), + ApiKeyAction::Revoke { id } => api + .request(reqwest::Method::DELETE, &format!("/api/api-keys/{}", id), None) + .await + .map(|data| { + if data.get("deleted").and_then(|v| v.as_bool()).unwrap_or(false) { + println!("Key #{} revoked successfully.", id); + } else { + print_json(&data); + } + }), + }, }; if let Err(e) = result { diff --git a/common/src/model/port_rule.rs b/common/src/model/port_rule.rs index aa825bd..996d641 100644 --- a/common/src/model/port_rule.rs +++ b/common/src/model/port_rule.rs @@ -1,8 +1,9 @@ #[cfg(feature = "user")] -use aya::Pod; -#[cfg(feature = "user")] use std::vec::Vec; +#[cfg(feature = "user")] +use aya::Pod; + use crate::define::setting::MAX_RULES_PORT; use crate::model::ip_address::Port; diff --git a/egress-ebpf/src/main.rs b/egress-ebpf/src/main.rs index 05393b8..1cd2d5c 100644 --- a/egress-ebpf/src/main.rs +++ b/egress-ebpf/src/main.rs @@ -19,9 +19,7 @@ static EGRESS_XSKS_MAP: XskMap = XskMap::pinned(64, 0); #[xdp] pub fn net_guardia(ctx: XdpContext) -> u32 { - let queue_id = unsafe { - compute_symmetric_queue_id(&ctx).unwrap_or((*ctx.ctx).rx_queue_index) - }; + let queue_id = unsafe { compute_symmetric_queue_id(&ctx).unwrap_or((*ctx.ctx).rx_queue_index) }; match EGRESS_XSKS_MAP.redirect(queue_id, 0) { Ok(action) => action, Err(_) => xdp_action::XDP_PASS, diff --git a/ingress-ebpf/src/action/access_control.rs b/ingress-ebpf/src/action/access_control.rs index 2c36c94..44c972c 100644 --- a/ingress-ebpf/src/action/access_control.rs +++ b/ingress-ebpf/src/action/access_control.rs @@ -2,9 +2,9 @@ use aya_ebpf::macros::map; use aya_ebpf::maps::HashMap; use aya_ebpf::maps::LpmTrie; use aya_ebpf::maps::lpm_trie::Key; -use common::define::setting::{MAX_RULES, MAX_GEO_ENTRIES}; -use common::model::parsed_packet::ParsedPacket; +use common::define::setting::{MAX_GEO_ENTRIES, MAX_RULES}; use common::model::ip_address::{IPv4, IPv6}; +use common::model::parsed_packet::ParsedPacket; use common::model::port_rule::PortRule; #[map] diff --git a/ingress-ebpf/src/action/mod.rs b/ingress-ebpf/src/action/mod.rs index 19ea3ec..1caa87a 100644 --- a/ingress-ebpf/src/action/mod.rs +++ b/ingress-ebpf/src/action/mod.rs @@ -1,3 +1,3 @@ pub mod access_control; -pub mod rate_limit; pub mod protocol_filter; +pub mod rate_limit; diff --git a/ingress-ebpf/src/action/rate_limit.rs b/ingress-ebpf/src/action/rate_limit.rs index eee298a..37d3dc2 100644 --- a/ingress-ebpf/src/action/rate_limit.rs +++ b/ingress-ebpf/src/action/rate_limit.rs @@ -50,19 +50,11 @@ fn get_config(index: u32, default: u64) -> u64 { #[inline(always)] fn is_syn_only(pkt: &ParsedPacket) -> bool { - matches!(pkt.protocol, IpProto::Tcp) - && (pkt.tcp_flags & TCP_SYN != 0) - && (pkt.tcp_flags & TCP_ACK == 0) + matches!(pkt.protocol, IpProto::Tcp) && (pkt.tcp_flags & TCP_SYN != 0) && (pkt.tcp_flags & TCP_ACK == 0) } #[inline(always)] -fn check_rate( - map: &LruHashMap, - key: &K, - now: u64, - window: u64, - limit: u64, -) -> bool { +fn check_rate(map: &LruHashMap, key: &K, now: u64, window: u64, limit: u64) -> bool { unsafe { if let Some(state) = map.get_ptr_mut(key) { if now - (*state).window_start >= window { @@ -91,24 +83,48 @@ fn ipv4_should_drop(pkt: &ParsedPacket) -> Option { let window = get_config(CFG_WINDOW_NS, DEFAULT_WINDOW_NS); let src_ip = pkt.src_ip_v4(); - if check_rate(&IPV4_PACKET_RATE_MAP, &src_ip, now, window, get_config(CFG_PACKET_RATE, DEFAULT_PACKET_RATE)) { + if check_rate( + &IPV4_PACKET_RATE_MAP, + &src_ip, + now, + window, + get_config(CFG_PACKET_RATE, DEFAULT_PACKET_RATE), + ) { return Some(DROP_REASON_RATE_LIMIT_PKT); } if is_syn_only(pkt) { - if check_rate(&IPV4_SYN_RATE_MAP, &src_ip, now, window, get_config(CFG_SYN_RATE, DEFAULT_SYN_RATE)) { + if check_rate( + &IPV4_SYN_RATE_MAP, + &src_ip, + now, + window, + get_config(CFG_SYN_RATE, DEFAULT_SYN_RATE), + ) { return Some(DROP_REASON_RATE_LIMIT_SYN); } } if matches!(pkt.protocol, IpProto::Udp) { - if check_rate(&IPV4_UDP_RATE_MAP, &src_ip, now, window, get_config(CFG_UDP_RATE, DEFAULT_UDP_RATE)) { + if check_rate( + &IPV4_UDP_RATE_MAP, + &src_ip, + now, + window, + get_config(CFG_UDP_RATE, DEFAULT_UDP_RATE), + ) { return Some(DROP_REASON_RATE_LIMIT_UDP); } } if matches!(pkt.protocol, IpProto::Udp) && pkt.dst_port == 53 { - if check_rate(&IPV4_DNS_RATE_MAP, &src_ip, now, window, get_config(CFG_DNS_RATE, DEFAULT_DNS_RATE)) { + if check_rate( + &IPV4_DNS_RATE_MAP, + &src_ip, + now, + window, + get_config(CFG_DNS_RATE, DEFAULT_DNS_RATE), + ) { return Some(DROP_REASON_RATE_LIMIT_DNS); } } @@ -122,24 +138,48 @@ fn ipv6_should_drop(pkt: &ParsedPacket) -> Option { let window = get_config(CFG_WINDOW_NS, DEFAULT_WINDOW_NS); let src_ip = pkt.src_ip_v6(); - if check_rate(&IPV6_PACKET_RATE_MAP, &src_ip, now, window, get_config(CFG_PACKET_RATE, DEFAULT_PACKET_RATE)) { + if check_rate( + &IPV6_PACKET_RATE_MAP, + &src_ip, + now, + window, + get_config(CFG_PACKET_RATE, DEFAULT_PACKET_RATE), + ) { return Some(DROP_REASON_RATE_LIMIT_PKT); } if is_syn_only(pkt) { - if check_rate(&IPV6_SYN_RATE_MAP, &src_ip, now, window, get_config(CFG_SYN_RATE, DEFAULT_SYN_RATE)) { + if check_rate( + &IPV6_SYN_RATE_MAP, + &src_ip, + now, + window, + get_config(CFG_SYN_RATE, DEFAULT_SYN_RATE), + ) { return Some(DROP_REASON_RATE_LIMIT_SYN); } } if matches!(pkt.protocol, IpProto::Udp) { - if check_rate(&IPV6_UDP_RATE_MAP, &src_ip, now, window, get_config(CFG_UDP_RATE, DEFAULT_UDP_RATE)) { + if check_rate( + &IPV6_UDP_RATE_MAP, + &src_ip, + now, + window, + get_config(CFG_UDP_RATE, DEFAULT_UDP_RATE), + ) { return Some(DROP_REASON_RATE_LIMIT_UDP); } } if matches!(pkt.protocol, IpProto::Udp) && pkt.dst_port == 53 { - if check_rate(&IPV6_DNS_RATE_MAP, &src_ip, now, window, get_config(CFG_DNS_RATE, DEFAULT_DNS_RATE)) { + if check_rate( + &IPV6_DNS_RATE_MAP, + &src_ip, + now, + window, + get_config(CFG_DNS_RATE, DEFAULT_DNS_RATE), + ) { return Some(DROP_REASON_RATE_LIMIT_DNS); } } diff --git a/ingress-ebpf/src/main.rs b/ingress-ebpf/src/main.rs index 3e5664d..f661814 100644 --- a/ingress-ebpf/src/main.rs +++ b/ingress-ebpf/src/main.rs @@ -5,17 +5,17 @@ mod action; use aya_ebpf::bindings::xdp_action; use aya_ebpf::macros::{map, xdp}; use aya_ebpf::maps::{Array, PerCpuArray, ProgramArray, RingBuf, XskMap}; -use common::ebpf::symmetric_hash::symmetric_queue_id; use aya_ebpf::programs::XdpContext; #[allow(unused_imports)] use aya_log_ebpf::info; -use common::ebpf::parsing; -use common::define::pipeline::*; use common::define::drop_reason::*; +use common::define::pipeline::*; +use common::ebpf::parsing; +use common::ebpf::symmetric_hash::symmetric_queue_id; use common::model::drop_event::DropEvent; use common::model::parsed_packet::ParsedPacket; -use crate::action::{access_control, rate_limit, protocol_filter}; +use crate::action::{access_control, protocol_filter, rate_limit}; #[map] static PROGRAM_ARRAY: ProgramArray = ProgramArray::with_max_entries(MAX_STAGES, 0); @@ -70,7 +70,9 @@ unsafe fn emit_drop_event(pkt: &ParsedPacket, reason: u8) { #[inline(always)] unsafe fn packet_intake(ctx: &XdpContext) { - let Some(ptr) = PARSED_PACKET.get_ptr_mut(0) else { return }; + let Some(ptr) = PARSED_PACKET.get_ptr_mut(0) else { + return; + }; if parsing::parse_packet(ctx.data(), ctx.data_end(), ptr).is_ok() { chain_next(ctx, STAGE_ENTRY); } @@ -209,9 +211,7 @@ unsafe fn compute_symmetric_queue_id() -> Option { #[xdp] pub fn transmission(ctx: XdpContext) -> u32 { - let queue_id = unsafe { - compute_symmetric_queue_id().unwrap_or((*ctx.ctx).rx_queue_index) - }; + let queue_id = unsafe { compute_symmetric_queue_id().unwrap_or((*ctx.ctx).rx_queue_index) }; match INGRESS_XSKS_MAP.redirect(queue_id, 0) { Ok(action) => action, Err(_) => xdp_action::XDP_PASS, diff --git a/macros/src/log.rs b/macros/src/log.rs index ff56b3f..06a698f 100644 --- a/macros/src/log.rs +++ b/macros/src/log.rs @@ -1,7 +1,7 @@ use proc_macro::TokenStream; use quote::quote; use syn::parse::{Parse, ParseStream}; -use syn::{parse_macro_input, Expr, Token}; +use syn::{Expr, Token, parse_macro_input}; struct LogInput { error: Expr, @@ -62,5 +62,5 @@ pub fn log_impl(input: TokenStream) -> TokenStream { } } } - .into() + .into() } diff --git a/mcp-server/src/main.rs b/mcp-server/src/main.rs index 6d6fa3e..41816f6 100644 --- a/mcp-server/src/main.rs +++ b/mcp-server/src/main.rs @@ -57,7 +57,11 @@ impl McpServer { .timeout(Duration::from_secs(30)) .build() .expect("Failed to create HTTP client"); - Self { client, api_url, api_key } + Self { + client, + api_url, + api_key, + } } async fn handle_request(&self, req: JsonRpcRequest) -> JsonRpcResponse { @@ -69,7 +73,10 @@ impl McpServer { jsonrpc: "2.0".into(), id: req.id, result: None, - error: Some(JsonRpcError { code: -32601, message: "Method not found".into() }), + error: Some(JsonRpcError { + code: -32601, + message: "Method not found".into(), + }), }, } } @@ -120,7 +127,10 @@ impl McpServer { async fn handle_tool_call(&self, id: Option, params: Value) -> JsonRpcResponse { let tool_name = params.get("name").and_then(|n| n.as_str()).unwrap_or(""); - let arguments = params.get("arguments").cloned().unwrap_or(Value::Object(Default::default())); + let arguments = params + .get("arguments") + .cloned() + .unwrap_or(Value::Object(Default::default())); let (method, path, body): (&str, String, Option) = match tool_name { "get_health" => ("GET", "/api/health/status".into(), None), @@ -136,34 +146,65 @@ impl McpServer { let ip = arguments.get("ip").and_then(|v| v.as_str()).unwrap_or(""); let is_v6 = ip.contains(':'); let ip_ver = if is_v6 { "ipv6" } else { "ipv4" }; - let addr = if is_v6 { format!("[{}]:0", ip) } else { format!("{}:0", ip) }; - ("PUT", format!("/api/acl/{}/source/blacklist", ip_ver), Some(Value::String(addr))) + let addr = if is_v6 { + format!("[{}]:0", ip) + } else { + format!("{}:0", ip) + }; + ( + "PUT", + format!("/api/acl/{}/source/blacklist", ip_ver), + Some(Value::String(addr)), + ) } "unblock_ip" => { let ip = arguments.get("ip").and_then(|v| v.as_str()).unwrap_or(""); let is_v6 = ip.contains(':'); let ip_ver = if is_v6 { "ipv6" } else { "ipv4" }; - let addr = if is_v6 { format!("[{}]:0", ip) } else { format!("{}:0", ip) }; - ("DELETE", format!("/api/acl/{}/source/blacklist", ip_ver), Some(Value::String(addr))) + let addr = if is_v6 { + format!("[{}]:0", ip) + } else { + format!("{}:0", ip) + }; + ( + "DELETE", + format!("/api/acl/{}/source/blacklist", ip_ver), + Some(Value::String(addr)), + ) } "set_enforce_mode" => { let mode = arguments.get("mode").and_then(|v| v.as_str()).unwrap_or("monitor"); - ("PUT", "/api/system/enforce-mode".into(), Some(serde_json::json!({"mode": mode}))) + ( + "PUT", + "/api/system/enforce-mode".into(), + Some(serde_json::json!({"mode": mode})), + ) } "add_dns_filter" => { let domain = arguments.get("domain").and_then(|v| v.as_str()).unwrap_or(""); - ("PUT", "/api/filter/dns/blacklist".into(), Some(serde_json::json!({"domains": [domain]}))) + ( + "PUT", + "/api/filter/dns/blacklist".into(), + Some(serde_json::json!({"domains": [domain]})), + ) } "add_geo_block" => { let code = arguments.get("country_code").and_then(|v| v.as_str()).unwrap_or(""); - ("PUT", "/api/acl/geo/block".into(), Some(serde_json::json!({"country_codes": [code]}))) + ( + "PUT", + "/api/acl/geo/block".into(), + Some(serde_json::json!({"country_codes": [code]})), + ) } _ => { return JsonRpcResponse { jsonrpc: "2.0".into(), id, result: None, - error: Some(JsonRpcError { code: -32602, message: format!("Unknown tool: {}", tool_name) }), + error: Some(JsonRpcError { + code: -32602, + message: format!("Unknown tool: {}", tool_name), + }), }; } }; @@ -208,17 +249,15 @@ impl McpServer { } } } - Err(e) => { - JsonRpcResponse { - jsonrpc: "2.0".into(), - id, - result: Some(serde_json::json!({ - "content": [{ "type": "text", "text": format!("Connection error: {}", e) }], - "isError": true - })), - error: None, - } - } + Err(e) => JsonRpcResponse { + jsonrpc: "2.0".into(), + id, + result: Some(serde_json::json!({ + "content": [{ "type": "text", "text": format!("Connection error: {}", e) }], + "isError": true + })), + error: None, + }, } } } @@ -227,7 +266,8 @@ impl McpServer { async fn main() { let args = Args::parse(); - let api_key = args.api_key + let api_key = args + .api_key .or_else(|| std::env::var("NETGUARDIA_API_KEY").ok()) .unwrap_or_else(|| { eprintln!("Error: No API key provided. Set NETGUARDIA_API_KEY env var or use --api-key flag."); @@ -256,7 +296,10 @@ async fn main() { jsonrpc: "2.0".into(), id: None, result: None, - error: Some(JsonRpcError { code: -32700, message: format!("Parse error: {}", e) }), + error: Some(JsonRpcError { + code: -32700, + message: format!("Parse error: {}", e), + }), }; let _ = writeln!(stdout, "{}", serde_json::to_string(&err_resp).unwrap()); let _ = stdout.flush(); diff --git a/net-guardia/Cargo.toml b/net-guardia/Cargo.toml index 33cef24..7fe5fe2 100644 --- a/net-guardia/Cargo.toml +++ b/net-guardia/Cargo.toml @@ -54,6 +54,9 @@ chrono = { version = "0.4", default-features = false, features = ["clock", "std" # HTTP client (Telegram, MCP proxy) reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] } +# CLI +clap = { version = "4", features = ["derive", "env"] } + # Architecture async-trait = "0.1" dashmap = "6" diff --git a/net-guardia/src/adapter/access_control_adapter.rs b/net-guardia/src/adapter/access_control.rs similarity index 94% rename from net-guardia/src/adapter/access_control_adapter.rs rename to net-guardia/src/adapter/access_control.rs index 7bc7968..4c776d6 100644 --- a/net-guardia/src/adapter/access_control_adapter.rs +++ b/net-guardia/src/adapter/access_control.rs @@ -9,17 +9,17 @@ use crate::model::error::ebpf::EbpfError; use crate::model::monitoring::direction::FlowDirection; /// Adapter that implements AccessControlPort by delegating to the eBPF AccessControl. -pub struct EbpfAccessControlAdapter { +pub struct AccessControlAdapter { access_control: Arc, } -impl EbpfAccessControlAdapter { +impl AccessControlAdapter { pub fn new(access_control: Arc) -> Self { Self { access_control } } } -impl AccessControlPort for EbpfAccessControlAdapter { +impl AccessControlPort for AccessControlAdapter { fn block_ip(&self, ip: &str) -> Result<(), Error> { let addr: IpAddr = ip .parse() diff --git a/net-guardia/src/adapter/ebpf/access_control.rs b/net-guardia/src/adapter/ebpf/access_control.rs index 0cc07c1..b954991 100644 --- a/net-guardia/src/adapter/ebpf/access_control.rs +++ b/net-guardia/src/adapter/ebpf/access_control.rs @@ -40,9 +40,6 @@ impl AccessControl { Ok(access_control) } - /// Construct an AccessControl backed by no eBPF maps. Used when eBPF - /// failed to load at startup; every mutating call returns `EbpfError::NotLoaded`, - /// and list queries return empty maps. pub fn unavailable() -> Self { Self { ipv4_src_whitelist: RwLock::new(MapWrapper::unavailable()), @@ -149,9 +146,11 @@ impl AccessControlAdminPort for AccessControl { fn add_ipv4_list(&self, direction: FlowDirection, list_type: ListType, address: SocketAddrV4) -> Result<(), Error> { self.add_ipv4_list(direction, list_type, address) } + fn add_ipv6_list(&self, direction: FlowDirection, list_type: ListType, address: SocketAddrV6) -> Result<(), Error> { self.add_ipv6_list(direction, list_type, address) } + fn remove_ipv4_list( &self, direction: FlowDirection, @@ -160,6 +159,7 @@ impl AccessControlAdminPort for AccessControl { ) -> Result<(), Error> { self.remove_ipv4_list(direction, list_type, address) } + fn remove_ipv6_list( &self, direction: FlowDirection, @@ -168,9 +168,11 @@ impl AccessControlAdminPort for AccessControl { ) -> Result<(), Error> { self.remove_ipv6_list(direction, list_type, address) } + fn get_ipv4_list(&self, direction: FlowDirection, list_type: ListType) -> HashMap> { self.get_ipv4_list(direction, list_type) } + fn get_ipv6_list(&self, direction: FlowDirection, list_type: ListType) -> HashMap> { self.get_ipv6_list(direction, list_type) } diff --git a/net-guardia/src/adapter/ebpf/dns_filter.rs b/net-guardia/src/adapter/ebpf/dns_filter.rs index 80f3245..8cf7ea7 100644 --- a/net-guardia/src/adapter/ebpf/dns_filter.rs +++ b/net-guardia/src/adapter/ebpf/dns_filter.rs @@ -38,8 +38,6 @@ impl DnsFilter { .collect() } - /// Fast-path helper combining `parse_query_name` + `is_blacklisted` — used - /// by the AF_XDP RX loop. pub fn is_query_blacklisted(&self, raw: &[u8]) -> bool { match Self::parse_query_name(raw) { Some((name, name_len)) => self.is_blacklisted(&name, name_len), @@ -47,7 +45,6 @@ impl DnsFilter { } } - /// Check if a DNS query name (in wire format) or any of its parent domains is blacklisted. pub fn is_blacklisted(&self, name: &DnsName, name_len: usize) -> bool { if self.blacklist.is_empty() { return false; @@ -84,8 +81,6 @@ impl DnsFilter { false } - /// Parse DNS query name from raw packet bytes. - /// Returns the DNS name in wire format and the name length, or None if not a DNS query. pub fn parse_query_name(raw: &[u8]) -> Option<(DnsName, usize)> { if raw.len() < 14 { return None; diff --git a/net-guardia/src/adapter/ebpf/drop_monitor.rs b/net-guardia/src/adapter/ebpf/drop_monitor.rs index 4277c5c..31eba34 100644 --- a/net-guardia/src/adapter/ebpf/drop_monitor.rs +++ b/net-guardia/src/adapter/ebpf/drop_monitor.rs @@ -5,11 +5,10 @@ use std::sync::atomic::Ordering; use std::time::Duration; use aya::maps::{MapData, RingBuf}; -use tokio::sync::{broadcast, oneshot}; -use tokio::time::interval; - use common::define::drop_reason::*; use common::model::drop_event::DropEvent as RawDropEvent; +use tokio::sync::{broadcast, oneshot}; +use tokio::time::interval; use crate::model::config::constants::DROP_CHANNEL_CAPACITY; use crate::model::monitoring::drop_event::{DropCounters, DropCountersAtomic, DropEventMessage}; diff --git a/net-guardia/src/adapter/ebpf/geo_block.rs b/net-guardia/src/adapter/ebpf/geo_block.rs index ca08c90..8bf5621 100644 --- a/net-guardia/src/adapter/ebpf/geo_block.rs +++ b/net-guardia/src/adapter/ebpf/geo_block.rs @@ -211,9 +211,11 @@ impl GeoBlockPort for GeoBlock { fn block_countries(&self, codes: &[String]) -> Result { self.block_countries(codes) } + fn unblock_countries(&self, codes: &[String]) -> Result { self.unblock_countries(codes) } + fn list_blocked(&self) -> Vec { self.get_blocked_countries() } diff --git a/net-guardia/src/adapter/ebpf/xsk_manager.rs b/net-guardia/src/adapter/ebpf/xsk_manager.rs index 0a2a40b..a33cd8c 100644 --- a/net-guardia/src/adapter/ebpf/xsk_manager.rs +++ b/net-guardia/src/adapter/ebpf/xsk_manager.rs @@ -8,6 +8,7 @@ use std::time::Duration; use aya::Ebpf; use aya::maps::{MapData, XskMap}; +use common::define::drop_reason::DROP_REASON_DNS_BLACKLIST; use crossbeam::channel::{Receiver, Sender, TrySendError, bounded}; use crossbeam::queue::SegQueue; use macros::log; @@ -16,8 +17,6 @@ use tokio::sync::oneshot::{self, error::TryRecvError}; use xsk_rs::config::{BindFlags, FrameSize, Interface, LibxdpFlags, QueueSize, SocketConfig, UmemConfig}; use xsk_rs::{CompQueue, FillQueue, FrameDesc, RxQueue, Socket, TxQueue, Umem}; -use common::define::drop_reason::DROP_REASON_DNS_BLACKLIST; - use crate::adapter::ebpf::drop_monitor::DropMonitor; use crate::infrastructure::app_config::AppConfig; use crate::interface::port::dns_query_filter::DnsQueryFilter; diff --git a/net-guardia/src/adapter/mod.rs b/net-guardia/src/adapter/mod.rs index 0670e9f..cd91183 100644 --- a/net-guardia/src/adapter/mod.rs +++ b/net-guardia/src/adapter/mod.rs @@ -1,4 +1,4 @@ -pub mod access_control_adapter; +pub mod access_control; pub mod ebpf; pub mod http; pub mod persistence; diff --git a/net-guardia/src/adapter/persistence/acl.rs b/net-guardia/src/adapter/persistence/acl.rs new file mode 100644 index 0000000..9c64e81 --- /dev/null +++ b/net-guardia/src/adapter/persistence/acl.rs @@ -0,0 +1,158 @@ +use rusqlite::params; + +use super::Database; +use crate::interface::port::acl::{AclRepo, AclRuleTuple}; +use crate::model::error::Error; + +impl Database { + pub fn insert_acl_rule( + &self, + ip_version: u8, + direction: &str, + list_type: &str, + ip_address: &str, + port: u16, + ) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "INSERT OR IGNORE INTO acl_rules (ip_version, direction, list_type, ip_address, port) VALUES (?1, ?2, ?3, ?4, ?5)", + params![ip_version, direction, list_type, ip_address, port as i64], + )?; + Ok(()) + } + + pub fn delete_acl_rule( + &self, + ip_version: u8, + direction: &str, + list_type: &str, + ip_address: &str, + port: u16, + ) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "DELETE FROM acl_rules WHERE ip_version = ?1 AND direction = ?2 AND list_type = ?3 AND ip_address = ?4 AND port = ?5", + params![ip_version, direction, list_type, ip_address, port as i64], + )?; + Ok(()) + } + + pub fn load_acl_rules(&self) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare("SELECT ip_version, direction, list_type, ip_address, port FROM acl_rules")?; + let rows = stmt.query_map([], |row| { + Ok(( + row.get::<_, u8>(0)?, + row.get::<_, String>(1)?, + row.get::<_, String>(2)?, + row.get::<_, String>(3)?, + row.get::<_, i64>(4)? as u16, + )) + })?; + let mut results = Vec::new(); + for row in rows { + results.push(row?); + } + Ok(results) + } + + pub fn has_manual_acl_rule(&self, ip_address: &str) -> Result { + let conn = self.conn()?; + let count: i64 = conn.query_row( + "SELECT COUNT(*) FROM acl_rules WHERE ip_address = ?1 AND list_type = 'blacklist'", + params![ip_address], + |row| row.get(0), + )?; + Ok(count > 0) + } + + pub fn load_admin_whitelist(&self) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare("SELECT ip FROM admin_whitelist")?; + let rows = stmt.query_map([], |row| row.get::<_, String>(0))?; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + + pub fn insert_admin_whitelist(&self, ip: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute("INSERT OR IGNORE INTO admin_whitelist (ip) VALUES (?1)", params![ip])?; + Ok(()) + } + + pub fn delete_admin_whitelist(&self, ip: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute("DELETE FROM admin_whitelist WHERE ip = ?1", params![ip])?; + Ok(()) + } +} + +impl AclRepo for Database { + fn insert_acl_rule( + &self, + ip_version: u8, + direction: &str, + list_type: &str, + ip_address: &str, + port: u16, + ) -> Result<(), Error> { + self.insert_acl_rule(ip_version, direction, list_type, ip_address, port) + } + + fn delete_acl_rule( + &self, + ip_version: u8, + direction: &str, + list_type: &str, + ip_address: &str, + port: u16, + ) -> Result<(), Error> { + self.delete_acl_rule(ip_version, direction, list_type, ip_address, port) + } + + fn has_manual_acl_rule(&self, ip_address: &str) -> Result { + self.has_manual_acl_rule(ip_address) + } + + fn load_admin_whitelist(&self) -> Result, Error> { + self.load_admin_whitelist() + } + + fn insert_admin_whitelist(&self, ip: &str) -> Result<(), Error> { + self.insert_admin_whitelist(ip) + } + + fn delete_admin_whitelist(&self, ip: &str) -> Result<(), Error> { + self.delete_admin_whitelist(ip) + } +} + +#[cfg(test)] +mod tests { + use super::super::tests::test_db; + + #[test] + fn test_acl_crud() { + let db = test_db(); + db.insert_acl_rule(4, "source", "blacklist", "192.168.1.1", 80).unwrap(); + let rules = db.load_acl_rules().unwrap(); + assert_eq!(rules.len(), 1); + assert_eq!( + rules[0], + ( + 4, + "source".to_string(), + "blacklist".to_string(), + "192.168.1.1".to_string(), + 80 + ) + ); + + db.delete_acl_rule(4, "source", "blacklist", "192.168.1.1", 80).unwrap(); + let rules = db.load_acl_rules().unwrap(); + assert!(rules.is_empty()); + } +} diff --git a/net-guardia/src/adapter/persistence/api_key.rs b/net-guardia/src/adapter/persistence/api_key.rs new file mode 100644 index 0000000..00ae325 --- /dev/null +++ b/net-guardia/src/adapter/persistence/api_key.rs @@ -0,0 +1,156 @@ +use std::fmt::Write; + +use hmac::{Hmac, Mac}; +use rusqlite::{Error as RusqliteError, params}; +use sha2::Sha256; + +use super::Database; +use crate::interface::port::api_key::{ApiKeyListItem, ApiKeyRepo}; +use crate::model::error::Error; +use crate::model::identity::auth::Claims; + +type HmacSha256 = Hmac; + +impl Database { + /// Compute HMAC-SHA256 of an API key using the derived secret. + pub fn hmac_api_key(&self, raw_key: &str) -> String { + // SAFETY: HMAC-SHA256 accepts keys of any length; the only error + // `new_from_slice` returns (`InvalidLength`) is unreachable for this + // algorithm. The unreachable!() is the correct sentinel. + let mut mac = HmacSha256::new_from_slice(&self.api_key_hmac).unwrap_or_else(|_| unreachable!()); + mac.update(raw_key.as_bytes()); + let result = mac.finalize().into_bytes(); + + let mut hex = String::with_capacity(64); + for byte in result { + let _ = write!(&mut hex, "{:02x}", byte); + } + hex + } + + /// Validate an API key and return Claims if valid. + /// Computes HMAC-SHA256 of the key and looks it up in api_keys table. + pub fn validate_api_key(&self, api_key: &str) -> Result, Error> { + let digest = self.hmac_api_key(api_key); + + let conn = self.conn()?; + let result = conn.query_row( + "SELECT id, name, permission_level FROM api_keys WHERE key_hash = ?1", + params![digest], + |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, String>(1)?, + row.get::<_, String>(2)?, + )) + }, + ); + + match result { + Ok((id, name, level)) => { + // Update last_used_at + let _ = conn.execute( + "UPDATE api_keys SET last_used_at = datetime('now') WHERE id = ?1", + params![id], + ); + + // Build permissions based on permission level + let permissions = match level.as_str() { + "read_write" | "full_access" => vec![ + "dashboard:read".into(), + "statistics:read".into(), + "ai_detection:read".into(), + "ai_detection:write".into(), + "access_control:read".into(), + "access_control:write".into(), + "geo_block:read".into(), + "geo_block:write".into(), + "dns_filter:read".into(), + "dns_filter:write".into(), + "rate_limit:read".into(), + "rate_limit:write".into(), + "system:read".into(), + "system:write".into(), + ], + _ => vec![ + "dashboard:read".into(), + "statistics:read".into(), + "ai_detection:read".into(), + "access_control:read".into(), + "geo_block:read".into(), + "dns_filter:read".into(), + "rate_limit:read".into(), + "system:read".into(), + ], + }; + + Ok(Some(Claims { + sub: -id, // negative ID to distinguish from user IDs + username: format!("api:{}", name), + role: level, + permissions, + exp: usize::MAX, // API keys don't expire (revocation via DB deletion) + })) + } + Err(RusqliteError::QueryReturnedNoRows) => Ok(None), + Err(e) => Err(e)?, + } + } + + pub fn insert_api_key(&self, key_hash: &str, name: &str, permission_level: &str) -> Result { + let conn = self.conn()?; + conn.execute( + "INSERT INTO api_keys (key_hash, name, permission_level) VALUES (?1, ?2, ?3)", + params![key_hash, name, permission_level], + )?; + Ok(conn.last_insert_rowid()) + } + + #[allow(clippy::type_complexity)] + pub fn list_api_keys(&self) -> Result)>, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare("SELECT id, name, permission_level, created_at, last_used_at FROM api_keys")?; + let rows = stmt.query_map([], |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, String>(1)?, + row.get::<_, String>(2)?, + row.get::<_, String>(3)?, + row.get::<_, Option>(4)?, + )) + })?; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + + pub fn delete_api_key(&self, id: i64) -> Result { + let conn = self.conn()?; + let affected = conn.execute("DELETE FROM api_keys WHERE id = ?1", params![id])?; + Ok(affected > 0) + } +} + +impl ApiKeyRepo for Database { + fn validate_api_key(&self, api_key: &str) -> Result, Error> { + self.validate_api_key(api_key) + } + + fn hmac_api_key(&self, raw_key: &str) -> String { + self.hmac_api_key(raw_key) + } + + fn insert_api_key(&self, key_hash: &str, name: &str, permission_level: &str) -> Result { + self.insert_api_key(key_hash, name, permission_level) + } + + fn list_api_keys(&self) -> Result, Error> { + self.list_api_keys() + } + + fn delete_api_key(&self, id: i64) -> Result { + self.delete_api_key(id) + } +} diff --git a/net-guardia/src/adapter/persistence/audit.rs b/net-guardia/src/adapter/persistence/audit.rs new file mode 100644 index 0000000..5c1a471 --- /dev/null +++ b/net-guardia/src/adapter/persistence/audit.rs @@ -0,0 +1,139 @@ +use std::fmt::Write; + +use chrono::Utc; +use rusqlite::params; +use sha2::{Digest, Sha256}; + +use super::Database; +use crate::interface::port::audit::{AuditLogEntry, AuditRepo}; +use crate::model::error::Error; +use crate::model::error::database::DatabaseError; + +/// Compute the row hash for an audit_log entry. +/// Formula: sha256_hex(ts || 0x00 || actor || 0x00 || action || 0x00 || detail || 0x00 || prev_hash) +fn audit_row_hash(ts: &str, actor: &str, action: &str, detail: &str, prev_hash: &str) -> String { + let mut h = Sha256::new(); + for part in [ts, actor, action, detail, prev_hash] { + h.update(part.as_bytes()); + h.update([0u8]); + } + let out = h.finalize(); + let mut hex = String::with_capacity(64); + for byte in out { + let _ = write!(&mut hex, "{:02x}", byte); + } + hex +} + +impl Database { + /// Insert an audit trail entry. Runs in a transaction so the (prev_hash + /// lookup, row_hash compute, insert) sequence is atomic and serializable. + pub fn insert_audit_log(&self, actor: &str, action: &str, detail: &str) -> Result<(), Error> { + let mut conn = self.conn()?; + let ts = Utc::now().format("%Y-%m-%d %H:%M:%S").to_string(); + + let tx = conn.transaction()?; + let prev_hash: String = tx + .query_row("SELECT row_hash FROM audit_log ORDER BY id DESC LIMIT 1", [], |row| { + row.get(0) + }) + .unwrap_or_default(); + + let row_hash = audit_row_hash(&ts, actor, action, detail, &prev_hash); + tx.execute( + "INSERT INTO audit_log (ts, actor, action, detail, prev_hash, row_hash) VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ts, actor, action, detail, prev_hash, row_hash], + )?; + tx.commit()?; + Ok(()) + } + + /// List recent audit log entries (most recent first, max 200). + pub fn list_audit_logs(&self) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = + conn.prepare("SELECT id, actor, action, detail, ts FROM audit_log ORDER BY id DESC LIMIT 200")?; + let rows = stmt + .query_map([], |row| { + Ok(AuditLogEntry { + id: row.get(0)?, + actor: row.get(1)?, + action: row.get(2)?, + detail: row.get(3)?, + created_at: row.get(4)?, + }) + })? + .filter_map(|r| r.ok()) + .collect(); + Ok(rows) + } + + /// Read audit entries whose `action` matches exactly, newest-first, + /// capped at `limit`. Drives the fusion explain endpoint. + pub fn list_audit_logs_by_action(&self, action: &str, limit: i64) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare( + "SELECT id, actor, action, detail, ts FROM audit_log WHERE action = ?1 ORDER BY id DESC LIMIT ?2", + )?; + let rows = stmt + .query_map(params![action, limit], |row| { + Ok(AuditLogEntry { + id: row.get(0)?, + actor: row.get(1)?, + action: row.get(2)?, + detail: row.get(3)?, + created_at: row.get(4)?, + }) + })? + .filter_map(|r| r.ok()) + .collect(); + Ok(rows) + } + + /// Walk the entire audit_log in id order and verify the hash chain. + /// Returns `Ok(count)` on success; returns `Err` at the first mismatch, + /// naming the offending row id and the kind of mismatch. + pub fn verify_audit_log_chain(&self) -> Result { + let conn = self.conn()?; + let mut stmt = + conn.prepare("SELECT id, ts, actor, action, detail, prev_hash, row_hash FROM audit_log ORDER BY id ASC")?; + let mut rows = stmt.query([])?; + + let mut expected_prev = String::new(); + let mut count = 0usize; + while let Some(row) = rows.next()? { + let id: i64 = row.get(0)?; + let ts: String = row.get(1)?; + let actor: String = row.get(2)?; + let action: String = row.get(3)?; + let detail: String = row.get(4)?; + let prev_hash: String = row.get(5)?; + let row_hash: String = row.get(6)?; + + if prev_hash != expected_prev { + return Err(DatabaseError::AuditPrevHashMismatch(id, expected_prev, prev_hash).into()); + } + let computed = audit_row_hash(&ts, &actor, &action, &detail, &prev_hash); + if computed != row_hash { + return Err(DatabaseError::AuditRowHashMismatch(id, computed, row_hash).into()); + } + expected_prev = row_hash; + count += 1; + } + Ok(count) + } +} + +impl AuditRepo for Database { + fn insert_audit_log(&self, actor: &str, action: &str, detail: &str) -> Result<(), Error> { + self.insert_audit_log(actor, action, detail) + } + + fn list_audit_logs_by_action(&self, action: &str, limit: i64) -> Result, Error> { + self.list_audit_logs_by_action(action, limit) + } + + fn verify_audit_log_chain(&self) -> Result { + self.verify_audit_log_chain() + } +} diff --git a/net-guardia/src/adapter/persistence/enforcement.rs b/net-guardia/src/adapter/persistence/enforcement.rs new file mode 100644 index 0000000..5bcaf48 --- /dev/null +++ b/net-guardia/src/adapter/persistence/enforcement.rs @@ -0,0 +1,140 @@ +use rusqlite::params; + +use super::Database; +use crate::interface::port::enforcement::EnforcementRepo; +use crate::model::error::Error; + +impl Database { + pub fn set_rate_limit(&self, key: &str, value: u64) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "INSERT OR REPLACE INTO rate_limit_config (key, value) VALUES (?1, ?2)", + params![key, value as i64], + )?; + Ok(()) + } + + pub fn load_rate_limit_config(&self) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare("SELECT key, value FROM rate_limit_config")?; + let rows = stmt.query_map([], |row| Ok((row.get::<_, String>(0)?, row.get::<_, i64>(1)? as u64)))?; + let mut results = Vec::new(); + for row in rows { + results.push(row?); + } + Ok(results) + } + + pub fn insert_dns_domain(&self, domain: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "INSERT OR IGNORE INTO dns_blacklist (domain) VALUES (?1)", + params![domain], + )?; + Ok(()) + } + + pub fn delete_dns_domain(&self, domain: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute("DELETE FROM dns_blacklist WHERE domain = ?1", params![domain])?; + Ok(()) + } + + pub fn load_dns_domains(&self) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare("SELECT domain FROM dns_blacklist")?; + let rows = stmt.query_map([], |row| row.get(0))?; + let mut results = Vec::new(); + for row in rows { + results.push(row?); + } + Ok(results) + } + + pub fn insert_geo_country(&self, code: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "INSERT OR IGNORE INTO geo_blocked_countries (country_code) VALUES (?1)", + params![code], + )?; + Ok(()) + } + + pub fn delete_geo_country(&self, code: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "DELETE FROM geo_blocked_countries WHERE country_code = ?1", + params![code], + )?; + Ok(()) + } + + pub fn load_geo_countries(&self) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare("SELECT country_code FROM geo_blocked_countries")?; + let rows = stmt.query_map([], |row| row.get(0))?; + let mut results = Vec::new(); + for row in rows { + results.push(row?); + } + Ok(results) + } +} + +impl EnforcementRepo for Database { + fn set_rate_limit(&self, key: &str, value: u64) -> Result<(), Error> { + self.set_rate_limit(key, value) + } + + fn insert_dns_domain(&self, domain: &str) -> Result<(), Error> { + self.insert_dns_domain(domain) + } + + fn delete_dns_domain(&self, domain: &str) -> Result<(), Error> { + self.delete_dns_domain(domain) + } + + fn insert_geo_country(&self, code: &str) -> Result<(), Error> { + self.insert_geo_country(code) + } + + fn delete_geo_country(&self, code: &str) -> Result<(), Error> { + self.delete_geo_country(code) + } +} + +#[cfg(test)] +mod tests { + use super::super::tests::test_db; + + #[test] + fn test_dns_crud() { + let db = test_db(); + db.insert_dns_domain("evil.com").unwrap(); + let domains = db.load_dns_domains().unwrap(); + assert_eq!(domains, vec!["evil.com"]); + + db.delete_dns_domain("evil.com").unwrap(); + assert!(db.load_dns_domains().unwrap().is_empty()); + } + + #[test] + fn test_geo_crud() { + let db = test_db(); + db.insert_geo_country("CN").unwrap(); + let countries = db.load_geo_countries().unwrap(); + assert_eq!(countries, vec!["CN"]); + + db.delete_geo_country("CN").unwrap(); + assert!(db.load_geo_countries().unwrap().is_empty()); + } + + #[test] + fn test_rate_limit_crud() { + let db = test_db(); + db.set_rate_limit("packet_rate", 1000).unwrap(); + let configs = db.load_rate_limit_config().unwrap(); + assert_eq!(configs.len(), 1); + assert_eq!(configs[0], ("packet_rate".to_string(), 1000)); + } +} diff --git a/net-guardia/src/adapter/persistence/mod.rs b/net-guardia/src/adapter/persistence/mod.rs index b3879f6..c64ae6a 100644 --- a/net-guardia/src/adapter/persistence/mod.rs +++ b/net-guardia/src/adapter/persistence/mod.rs @@ -1,3 +1,448 @@ -pub mod repository; +mod acl; +mod api_key; +mod audit; +mod enforcement; +mod setting; +mod soar; +mod soar_block; +mod stats; +mod user; -pub use repository::Database; +use std::env; + +use macros::log; +use r2d2::Pool; +use r2d2_sqlite::SqliteConnectionManager; +use rusqlite::{self, Connection, params}; + +use crate::model::error::Error; +use crate::model::error::database::DatabaseError; +use crate::model::log::misc::MiscLog; + +/// Reads the SQLCipher encryption key from the environment variable `NETGUARDIA_DB_KEY`. +/// Returns `Some(key)` if set and non-empty, `None` otherwise (dev / unencrypted mode). +fn db_encryption_key() -> Option { + match env::var("NETGUARDIA_DB_KEY") { + Ok(k) if !k.is_empty() => Some(k), + _ => None, + } +} + +/// Applies the SQLCipher PRAGMA key (if configured) and standard PRAGMAs +/// to every new connection obtained from the pool. +#[derive(Debug, Clone)] +struct SqlitePragmaCustomizer { + /// `None` means no encryption (dev mode). + encryption_key: Option, +} + +impl r2d2::CustomizeConnection for SqlitePragmaCustomizer { + fn on_acquire(&self, conn: &mut Connection) -> Result<(), rusqlite::Error> { + // SQLCipher: the very first statement on a connection MUST be PRAGMA key. + if let Some(ref key) = self.encryption_key { + // Use a parameterised query to avoid SQL-injection via the key value. + conn.pragma_update(None, "key", key)?; + } + // PRAGMA tuning notes: + // - `journal_mode=WAL`: many concurrent readers + one writer; the only + // journal mode that survives crashes without losing committed rows. + // - `synchronous=NORMAL`: canonical pairing with WAL — `FULL` adds an + // extra fsync per commit that buys no durability guarantees beyond + // what WAL already provides for a power-loss event. + // - `busy_timeout=5000`: WAL still serializes writers (SOAR, audit, + // drift, SQL hooks all share one DB), and the default 0ms returns + // SQLITE_BUSY immediately on any contention. 5s gives the loser + // enough time to wait out a normal commit (sub-ms) without masking + // genuine deadlocks. + // - `foreign_keys=ON`: enforce FK constraints at the connection + // level (SQLite's default is OFF for backwards compatibility). + conn.execute_batch( + "PRAGMA journal_mode=WAL; \ + PRAGMA synchronous=NORMAL; \ + PRAGMA busy_timeout=5000; \ + PRAGMA foreign_keys=ON;", + )?; + Ok(()) + } +} + +pub struct Database { + pool: Pool, + /// HMAC-SHA256 key for API key hashing, derived from NETGUARDIA_SECRETS_KEY. + api_key_hmac: [u8; 32], +} + +impl Database { + pub fn new(path: &str) -> Result { + let encryption_key = db_encryption_key(); + + if path != ":memory:" && encryption_key.is_none() { + log!(MiscLog::DbEncryptionDisabled); + } + + let manager = if path == ":memory:" { + SqliteConnectionManager::memory() + } else { + SqliteConnectionManager::file(path) + }; + + let customizer = SqlitePragmaCustomizer { + encryption_key: encryption_key.clone(), + }; + + let pool = Pool::builder() + .max_size(if path == ":memory:" { 1 } else { 6 }) + .connection_customizer(Box::new(customizer)) + .build(manager) + .map_err(DatabaseError::QueryFailed)?; + + // Verify the pool is actually usable (catches wrong key / corrupt DB early). + { + let test_conn = pool.get().map_err(DatabaseError::QueryFailed)?; + test_conn + .query_row("SELECT count(*) FROM sqlite_master", [], |_| Ok(())) + .map_err(|_| DatabaseError::EncryptionKeyInvalid)?; + } + + let api_key_hmac = Self::derive_api_key_hmac(); + let db = Self { pool, api_key_hmac }; + db.create_tables()?; + Ok(db) + } + + /// Derive HMAC-SHA256 key for API key hashing from NETGUARDIA_SECRETS_KEY. + /// Falls back to a static dev key if the env var is unset. + fn derive_api_key_hmac() -> [u8; 32] { + use hkdf::Hkdf; + use sha2::Sha256; + + let root_key = env::var("NETGUARDIA_SECRETS_KEY") + .ok() + .filter(|k| !k.is_empty()) + .or_else(|| env::var("NETGUARDIA_DB_KEY").ok().filter(|k| !k.is_empty())) + .unwrap_or_else(|| "netguardia-dev-api-key-secret".to_string()); + + let hk = Hkdf::::new(Some(b"netguardia-v1-salt"), root_key.as_bytes()); + let mut okm = [0u8; 32]; + // SAFETY: 32 bytes is a valid output length for HKDF-SHA256 + hk.expand(b"netguardia-apikey-hmac-v1", &mut okm).unwrap(); + okm + } + + /// Export an encrypted database to a plaintext copy. + /// The original file is NOT modified. + pub fn decrypt_to_file(src_path: &str, key: &str, dest_path: &str) -> Result<(), Error> { + let conn = Connection::open(src_path).map_err(DatabaseError::QueryFailed)?; + conn.pragma_update(None, "key", key) + .map_err(DatabaseError::QueryFailed)?; + // Verify we can read the encrypted DB + conn.query_row("SELECT count(*) FROM sqlite_master", [], |_| Ok(())) + .map_err(|_| DatabaseError::DatabaseNotReadable)?; + // Attach a plaintext destination (empty key = no encryption) + conn.execute_batch(&format!( + "ATTACH DATABASE '{}' AS plaintext KEY '';", + dest_path.replace('\'', "''"), + )) + .map_err(DatabaseError::QueryFailed)?; + conn.query_row("SELECT sqlcipher_export('plaintext')", [], |_| Ok(())) + .map_err(DatabaseError::QueryFailed)?; + conn.execute_batch("DETACH DATABASE plaintext;") + .map_err(DatabaseError::QueryFailed)?; + Ok(()) + } + + /// Encrypt a plaintext database to a new encrypted copy. + /// The original file is NOT modified. + pub fn encrypt_to_file(src_path: &str, key: &str, dest_path: &str) -> Result<(), Error> { + let conn = Connection::open(src_path).map_err(DatabaseError::QueryFailed)?; + // Verify it's readable as plaintext + conn.query_row("SELECT count(*) FROM sqlite_master", [], |_| Ok(())) + .map_err(|_| DatabaseError::SourceDatabaseNotReadable)?; + conn.execute_batch(&format!( + "ATTACH DATABASE '{}' AS encrypted KEY '{}';", + dest_path.replace('\'', "''"), + key.replace('\'', "''"), + )) + .map_err(DatabaseError::QueryFailed)?; + conn.query_row("SELECT sqlcipher_export('encrypted')", [], |_| Ok(())) + .map_err(DatabaseError::QueryFailed)?; + conn.execute_batch("DETACH DATABASE encrypted;") + .map_err(DatabaseError::QueryFailed)?; + Ok(()) + } + + fn conn(&self) -> Result, Error> { + self.pool.get().map_err(|e| DatabaseError::QueryFailed(e).into()) + } + + fn create_tables(&self) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute_batch( + " + CREATE TABLE IF NOT EXISTS users ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + username TEXT UNIQUE NOT NULL, + password_hash TEXT NOT NULL, + role TEXT NOT NULL DEFAULT 'viewer', + force_password_change INTEGER NOT NULL DEFAULT 0, + created_at TEXT NOT NULL DEFAULT (datetime('now')) + ); + CREATE TABLE IF NOT EXISTS acl_rules ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + ip_version INTEGER NOT NULL, + direction TEXT NOT NULL, + list_type TEXT NOT NULL, + ip_address TEXT NOT NULL, + port INTEGER NOT NULL, + UNIQUE(ip_version, direction, list_type, ip_address, port) + ); + CREATE TABLE IF NOT EXISTS rate_limit_config ( + key TEXT PRIMARY KEY, + value INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS dns_blacklist ( + domain TEXT PRIMARY KEY + ); + CREATE TABLE IF NOT EXISTS geo_blocked_countries ( + country_code TEXT PRIMARY KEY + ); + CREATE TABLE IF NOT EXISTS settings ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL + ); + CREATE TABLE IF NOT EXISTS user_groups ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT UNIQUE NOT NULL, + description TEXT NOT NULL DEFAULT '', + permissions TEXT NOT NULL DEFAULT '[]', + created_at TEXT NOT NULL DEFAULT (datetime('now')) + ); + CREATE TABLE IF NOT EXISTS user_group_members ( + user_id INTEGER NOT NULL, + group_id INTEGER NOT NULL, + PRIMARY KEY (user_id, group_id), + FOREIGN KEY (user_id) REFERENCES users(id), + FOREIGN KEY (group_id) REFERENCES user_groups(id) + ); + + -- SOAR tables + CREATE TABLE IF NOT EXISTS playbooks ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + enabled INTEGER DEFAULT 1, + trigger_event TEXT NOT NULL, + condition_threshold REAL, + condition_count INTEGER, + condition_window_secs INTEGER, + cooldown_secs INTEGER DEFAULT 300, + created_at TEXT NOT NULL DEFAULT (datetime('now')), + updated_at TEXT NOT NULL DEFAULT (datetime('now')) + ); + CREATE TABLE IF NOT EXISTS playbook_actions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + playbook_id INTEGER NOT NULL REFERENCES playbooks(id) ON DELETE CASCADE, + action_order INTEGER NOT NULL, + action_type TEXT NOT NULL, + params TEXT NOT NULL DEFAULT '{}', + UNIQUE(playbook_id, action_order) + ); + CREATE TABLE IF NOT EXISTS playbook_conditions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + playbook_id INTEGER NOT NULL REFERENCES playbooks(id) ON DELETE CASCADE, + condition_type TEXT NOT NULL, + operator TEXT NOT NULL DEFAULT '>=', + value TEXT NOT NULL, + value2 TEXT, + UNIQUE(playbook_id, condition_type) + ); + CREATE TABLE IF NOT EXISTS soar_block_rules ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + source_ip TEXT NOT NULL, + playbook_id INTEGER NOT NULL, + created_at TEXT NOT NULL DEFAULT (datetime('now')), + expires_at TEXT NOT NULL, + unblocked_at TEXT + ); + CREATE TABLE IF NOT EXISTS soar_executions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + playbook_id INTEGER NOT NULL, + source_ip TEXT, + trigger_event TEXT NOT NULL, + actions_executed TEXT NOT NULL DEFAULT '[]', + executed_at TEXT NOT NULL DEFAULT (datetime('now')) + ); + CREATE TABLE IF NOT EXISTS admin_whitelist ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + ip TEXT NOT NULL UNIQUE, + created_at TEXT NOT NULL DEFAULT (datetime('now')) + ); + + -- MCP API keys + CREATE TABLE IF NOT EXISTS api_keys ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + key_hash TEXT NOT NULL, + name TEXT NOT NULL, + permission_level TEXT NOT NULL DEFAULT 'read_only', + created_at TEXT NOT NULL DEFAULT (datetime('now')), + last_used_at TEXT + ); + + -- Notification config (Telegram bot token, etc.) + CREATE TABLE IF NOT EXISTS notification_config ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + channel TEXT NOT NULL UNIQUE, + config_json TEXT NOT NULL, + enabled INTEGER DEFAULT 1, + created_at TEXT NOT NULL DEFAULT (datetime('now')), + updated_at TEXT NOT NULL DEFAULT (datetime('now')) + ); + + -- App secrets (encryption keys for sensitive data) + CREATE TABLE IF NOT EXISTS app_secrets ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL, + created_at TEXT NOT NULL DEFAULT (datetime('now')) + ); + + -- Pending unblock queue for orphan eBPF block recovery + CREATE TABLE IF NOT EXISTS pending_unblock ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + source_ip TEXT NOT NULL, + created_at TEXT NOT NULL DEFAULT (datetime('now')), + retry_count INTEGER NOT NULL DEFAULT 0 + ); + + -- Audit trail (WORM: hash-chained, triggers block UPDATE/DELETE) + CREATE TABLE IF NOT EXISTS audit_log ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + ts TEXT NOT NULL, + actor TEXT NOT NULL, + action TEXT NOT NULL, + detail TEXT NOT NULL DEFAULT '{}', + prev_hash TEXT NOT NULL DEFAULT '', + row_hash TEXT NOT NULL DEFAULT '' + ); + + CREATE TRIGGER IF NOT EXISTS audit_log_no_update + BEFORE UPDATE ON audit_log BEGIN + SELECT RAISE(ABORT, 'audit_log is append-only (WORM)'); + END; + + CREATE TRIGGER IF NOT EXISTS audit_log_no_delete + BEFORE DELETE ON audit_log BEGIN + SELECT RAISE(ABORT, 'audit_log is append-only (WORM)'); + END; + ", + )?; + + let conn_ref = &*conn; + + // Seed default user groups on first install (empty table) + let group_count: i64 = conn_ref.query_row("SELECT COUNT(*) FROM user_groups", [], |row| row.get(0))?; + if group_count == 0 { + let all_permissions = serde_json::json!([ + "dashboard:read", + "statistics:read", + "traffic_map:read", + "drops:read", + "ai_detection:read", + "ai_detection:write", + "access_control:read", + "access_control:write", + "geo_block:read", + "geo_block:write", + "dns_filter:read", + "dns_filter:write", + "rate_limit:read", + "rate_limit:write", + "protocol_filter:read", + "protocol_filter:write", + "system:read", + "system:write", + "system:admin", + "users:read", + "users:write", + "users:admin", + "fusion:read", + "fusion:write", + "flow_trace:read", + "flow_trace:write" + ]) + .to_string(); + let viewer_permissions = serde_json::json!([ + "dashboard:read", + "statistics:read", + "traffic_map:read", + "drops:read", + "ai_detection:read", + "access_control:read", + "geo_block:read", + "dns_filter:read", + "rate_limit:read", + "protocol_filter:read", + "system:read", + "fusion:read", + "flow_trace:read" + ]) + .to_string(); + + conn_ref.execute( + "INSERT INTO user_groups (name, description, permissions) VALUES (?1, ?2, ?3)", + params![ + "Administrator", + "Full system access with all permissions", + &all_permissions + ], + )?; + conn_ref.execute( + "INSERT INTO user_groups (name, description, permissions) VALUES (?1, ?2, ?3)", + params!["Viewer", "Read-only access to all modules", &viewer_permissions], + )?; + } + + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::interface::port::acl::AclRepo; + use crate::interface::port::identity::IdentityRepo; + use crate::interface::port::setting::SettingRepo; + + pub(super) fn test_db() -> Database { + Database::new(":memory:").expect("Failed to create test database") + } + + #[test] + fn test_create_tables() { + let _db = test_db(); + } + + /// Verify that Database satisfies each aggregate Repo trait contract + /// (AclRepo / SettingRepo / IdentityRepo). Exercises the trait-object + /// path so callers that take `Arc` compile end-to-end. + #[test] + fn test_aggregate_repo_trait_objects() { + let db = test_db(); + + let setting: &dyn SettingRepo = &db; + setting.set_setting("test_key", "test_value").unwrap(); + assert_eq!(setting.get_setting("test_key").unwrap(), Some("test_value".to_string())); + + let acl: &dyn AclRepo = &db; + acl.insert_acl_rule(4, "source", "blacklist", "10.0.0.1", 443).unwrap(); + // load_acl_rules is an inherent Database method (not on AclRepo), + // so go through `&db` directly for this read-back assertion. + let rules = db.load_acl_rules().unwrap(); + assert_eq!(rules.len(), 1); + + let identity: &dyn IdentityRepo = &db; + // user_count is inherent — inserts still go through the trait so + // the vtable has something to exercise. + assert_eq!(db.user_count().unwrap(), 0); + identity.insert_user("test", "hash", "viewer", false).unwrap(); + assert_eq!(db.user_count().unwrap(), 1); + } +} diff --git a/net-guardia/src/adapter/persistence/repository.rs b/net-guardia/src/adapter/persistence/repository.rs deleted file mode 100644 index c32d895..0000000 --- a/net-guardia/src/adapter/persistence/repository.rs +++ /dev/null @@ -1,2305 +0,0 @@ -use std::collections::{HashMap, HashSet}; -use std::env; -use std::time::{Duration, SystemTime, UNIX_EPOCH}; - -use chrono::Utc; -use macros::log; -use r2d2::Pool; -use r2d2_sqlite::SqliteConnectionManager; -use rusqlite::{self, Connection, Error as RusqliteError, params}; - -use crate::interface::port::acl::{AclRepo, AclRuleTuple}; -use crate::interface::port::api_key::{ApiKeyListItem, ApiKeyRepo}; -use crate::interface::port::audit::{AuditLogEntry, AuditRepo}; -use crate::interface::port::db_admin::DbAdminRepo; -use crate::interface::port::enforcement::EnforcementRepo; -use crate::interface::port::identity::{IdentityRepo, UserGroupTuple, UserTuple, UserWithGroups}; -use crate::interface::port::setting::SettingRepo; -use crate::interface::port::soar::{PlaybookRow, SoarExecutionRow, SoarRepo}; -use crate::interface::port::stats::StatsRepo; -use crate::model::error::Error; -use crate::model::error::database::DatabaseError; -use crate::model::identity::auth::Claims; -use crate::model::log::misc::MiscLog; -use crate::model::soar::playbook_data::UpdatePlaybookRow; - -/// Reads the SQLCipher encryption key from the environment variable `NETGUARDIA_DB_KEY`. -/// Returns `Some(key)` if set and non-empty, `None` otherwise (dev / unencrypted mode). -fn db_encryption_key() -> Option { - match env::var("NETGUARDIA_DB_KEY") { - Ok(k) if !k.is_empty() => Some(k), - _ => None, - } -} - -/// Applies the SQLCipher PRAGMA key (if configured) and standard PRAGMAs -/// to every new connection obtained from the pool. -#[derive(Debug, Clone)] -struct SqlitePragmaCustomizer { - /// `None` means no encryption (dev mode). - encryption_key: Option, -} - -impl r2d2::CustomizeConnection for SqlitePragmaCustomizer { - fn on_acquire(&self, conn: &mut rusqlite::Connection) -> Result<(), rusqlite::Error> { - // SQLCipher: the very first statement on a connection MUST be PRAGMA key. - if let Some(ref key) = self.encryption_key { - // Use a parameterised query to avoid SQL-injection via the key value. - conn.pragma_update(None, "key", key)?; - } - // PRAGMA tuning notes: - // - `journal_mode=WAL`: many concurrent readers + one writer; the only - // journal mode that survives crashes without losing committed rows. - // - `synchronous=NORMAL`: canonical pairing with WAL — `FULL` adds an - // extra fsync per commit that buys no durability guarantees beyond - // what WAL already provides for a power-loss event. - // - `busy_timeout=5000`: WAL still serializes writers (SOAR, audit, - // drift, SQL hooks all share one DB), and the default 0ms returns - // SQLITE_BUSY immediately on any contention. 5s gives the loser - // enough time to wait out a normal commit (sub-ms) without masking - // genuine deadlocks. - // - `foreign_keys=ON`: enforce FK constraints at the connection - // level (SQLite's default is OFF for backwards compatibility). - conn.execute_batch( - "PRAGMA journal_mode=WAL; \ - PRAGMA synchronous=NORMAL; \ - PRAGMA busy_timeout=5000; \ - PRAGMA foreign_keys=ON;", - )?; - Ok(()) - } -} - -pub struct Database { - pool: Pool, - /// HMAC-SHA256 key for API key hashing, derived from NETGUARDIA_SECRETS_KEY. - api_key_hmac: [u8; 32], -} - -impl Database { - pub fn new(path: &str) -> Result { - let encryption_key = db_encryption_key(); - - if path != ":memory:" && encryption_key.is_none() { - log!(MiscLog::DbEncryptionDisabled); - } - - let manager = if path == ":memory:" { - SqliteConnectionManager::memory() - } else { - SqliteConnectionManager::file(path) - }; - - let customizer = SqlitePragmaCustomizer { - encryption_key: encryption_key.clone(), - }; - - let pool = Pool::builder() - .max_size(if path == ":memory:" { 1 } else { 6 }) - .connection_customizer(Box::new(customizer)) - .build(manager) - .map_err(DatabaseError::QueryFailed)?; - - // Verify the pool is actually usable (catches wrong key / corrupt DB early). - { - let test_conn = pool.get().map_err(DatabaseError::QueryFailed)?; - test_conn - .query_row("SELECT count(*) FROM sqlite_master", [], |_| Ok(())) - .map_err(|_| DatabaseError::EncryptionKeyInvalid)?; - } - - let api_key_hmac = Self::derive_api_key_hmac(); - let db = Self { pool, api_key_hmac }; - db.create_tables()?; - Ok(db) - } - - /// Derive HMAC-SHA256 key for API key hashing from NETGUARDIA_SECRETS_KEY. - /// Falls back to a static dev key if the env var is unset. - fn derive_api_key_hmac() -> [u8; 32] { - use hkdf::Hkdf; - use sha2::Sha256; - - let root_key = env::var("NETGUARDIA_SECRETS_KEY") - .ok() - .filter(|k| !k.is_empty()) - .or_else(|| env::var("NETGUARDIA_DB_KEY").ok().filter(|k| !k.is_empty())) - .unwrap_or_else(|| "netguardia-dev-api-key-secret".to_string()); - - let hk = Hkdf::::new(Some(b"netguardia-v1-salt"), root_key.as_bytes()); - let mut okm = [0u8; 32]; - // SAFETY: 32 bytes is a valid output length for HKDF-SHA256 - hk.expand(b"netguardia-apikey-hmac-v1", &mut okm).unwrap(); - okm - } - - /// Compute HMAC-SHA256 of an API key using the derived secret. - pub fn hmac_api_key(&self, raw_key: &str) -> String { - use hmac::{Hmac, Mac}; - use sha2::Sha256; - use std::fmt::Write; - - type HmacSha256 = Hmac; - // SAFETY: HMAC-SHA256 accepts keys of any length; the only error - // `new_from_slice` returns (`InvalidLength`) is unreachable for this - // algorithm. The unreachable!() is the correct sentinel. - let mut mac = HmacSha256::new_from_slice(&self.api_key_hmac).unwrap_or_else(|_| unreachable!()); - mac.update(raw_key.as_bytes()); - let result = mac.finalize().into_bytes(); - - let mut hex = String::with_capacity(64); - for byte in result { - let _ = write!(&mut hex, "{:02x}", byte); - } - hex - } - - /// Export an encrypted database to a plaintext copy. - /// The original file is NOT modified. - pub fn decrypt_to_file(src_path: &str, key: &str, dest_path: &str) -> Result<(), Error> { - let conn = Connection::open(src_path).map_err(DatabaseError::QueryFailed)?; - conn.pragma_update(None, "key", key) - .map_err(DatabaseError::QueryFailed)?; - // Verify we can read the encrypted DB - conn.query_row("SELECT count(*) FROM sqlite_master", [], |_| Ok(())) - .map_err(|_| DatabaseError::DatabaseNotReadable)?; - // Attach a plaintext destination (empty key = no encryption) - conn.execute_batch(&format!( - "ATTACH DATABASE '{}' AS plaintext KEY '';", - dest_path.replace('\'', "''"), - )) - .map_err(DatabaseError::QueryFailed)?; - conn.query_row("SELECT sqlcipher_export('plaintext')", [], |_| Ok(())) - .map_err(DatabaseError::QueryFailed)?; - conn.execute_batch("DETACH DATABASE plaintext;") - .map_err(DatabaseError::QueryFailed)?; - Ok(()) - } - - /// Encrypt a plaintext database to a new encrypted copy. - /// The original file is NOT modified. - pub fn encrypt_to_file(src_path: &str, key: &str, dest_path: &str) -> Result<(), Error> { - let conn = Connection::open(src_path).map_err(DatabaseError::QueryFailed)?; - // Verify it's readable as plaintext - conn.query_row("SELECT count(*) FROM sqlite_master", [], |_| Ok(())) - .map_err(|_| DatabaseError::SourceDatabaseNotReadable)?; - conn.execute_batch(&format!( - "ATTACH DATABASE '{}' AS encrypted KEY '{}';", - dest_path.replace('\'', "''"), - key.replace('\'', "''"), - )) - .map_err(DatabaseError::QueryFailed)?; - conn.query_row("SELECT sqlcipher_export('encrypted')", [], |_| Ok(())) - .map_err(DatabaseError::QueryFailed)?; - conn.execute_batch("DETACH DATABASE encrypted;") - .map_err(DatabaseError::QueryFailed)?; - Ok(()) - } - - fn conn(&self) -> Result, Error> { - self.pool.get().map_err(|e| DatabaseError::QueryFailed(e).into()) - } - - fn create_tables(&self) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute_batch( - " - CREATE TABLE IF NOT EXISTS users ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - username TEXT UNIQUE NOT NULL, - password_hash TEXT NOT NULL, - role TEXT NOT NULL DEFAULT 'viewer', - force_password_change INTEGER NOT NULL DEFAULT 0, - created_at TEXT NOT NULL DEFAULT (datetime('now')) - ); - CREATE TABLE IF NOT EXISTS acl_rules ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - ip_version INTEGER NOT NULL, - direction TEXT NOT NULL, - list_type TEXT NOT NULL, - ip_address TEXT NOT NULL, - port INTEGER NOT NULL, - UNIQUE(ip_version, direction, list_type, ip_address, port) - ); - CREATE TABLE IF NOT EXISTS rate_limit_config ( - key TEXT PRIMARY KEY, - value INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS dns_blacklist ( - domain TEXT PRIMARY KEY - ); - CREATE TABLE IF NOT EXISTS geo_blocked_countries ( - country_code TEXT PRIMARY KEY - ); - CREATE TABLE IF NOT EXISTS settings ( - key TEXT PRIMARY KEY, - value TEXT NOT NULL - ); - CREATE TABLE IF NOT EXISTS user_groups ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - name TEXT UNIQUE NOT NULL, - description TEXT NOT NULL DEFAULT '', - permissions TEXT NOT NULL DEFAULT '[]', - created_at TEXT NOT NULL DEFAULT (datetime('now')) - ); - CREATE TABLE IF NOT EXISTS user_group_members ( - user_id INTEGER NOT NULL, - group_id INTEGER NOT NULL, - PRIMARY KEY (user_id, group_id), - FOREIGN KEY (user_id) REFERENCES users(id), - FOREIGN KEY (group_id) REFERENCES user_groups(id) - ); - - -- SOAR tables - CREATE TABLE IF NOT EXISTS playbooks ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - name TEXT NOT NULL, - enabled INTEGER DEFAULT 1, - trigger_event TEXT NOT NULL, - condition_threshold REAL, - condition_count INTEGER, - condition_window_secs INTEGER, - cooldown_secs INTEGER DEFAULT 300, - created_at TEXT NOT NULL DEFAULT (datetime('now')), - updated_at TEXT NOT NULL DEFAULT (datetime('now')) - ); - CREATE TABLE IF NOT EXISTS playbook_actions ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - playbook_id INTEGER NOT NULL REFERENCES playbooks(id) ON DELETE CASCADE, - action_order INTEGER NOT NULL, - action_type TEXT NOT NULL, - params TEXT NOT NULL DEFAULT '{}', - UNIQUE(playbook_id, action_order) - ); - CREATE TABLE IF NOT EXISTS playbook_conditions ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - playbook_id INTEGER NOT NULL REFERENCES playbooks(id) ON DELETE CASCADE, - condition_type TEXT NOT NULL, - operator TEXT NOT NULL DEFAULT '>=', - value TEXT NOT NULL, - value2 TEXT, - UNIQUE(playbook_id, condition_type) - ); - CREATE TABLE IF NOT EXISTS soar_block_rules ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - source_ip TEXT NOT NULL, - playbook_id INTEGER NOT NULL, - created_at TEXT NOT NULL DEFAULT (datetime('now')), - expires_at TEXT NOT NULL, - unblocked_at TEXT - ); - CREATE TABLE IF NOT EXISTS soar_executions ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - playbook_id INTEGER NOT NULL, - source_ip TEXT, - trigger_event TEXT NOT NULL, - actions_executed TEXT NOT NULL DEFAULT '[]', - executed_at TEXT NOT NULL DEFAULT (datetime('now')) - ); - CREATE TABLE IF NOT EXISTS admin_whitelist ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - ip TEXT NOT NULL UNIQUE, - created_at TEXT NOT NULL DEFAULT (datetime('now')) - ); - - -- MCP API keys - CREATE TABLE IF NOT EXISTS api_keys ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - key_hash TEXT NOT NULL, - name TEXT NOT NULL, - permission_level TEXT NOT NULL DEFAULT 'read_only', - created_at TEXT NOT NULL DEFAULT (datetime('now')), - last_used_at TEXT - ); - - -- Notification config (Telegram bot token, etc.) - CREATE TABLE IF NOT EXISTS notification_config ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - channel TEXT NOT NULL UNIQUE, - config_json TEXT NOT NULL, - enabled INTEGER DEFAULT 1, - created_at TEXT NOT NULL DEFAULT (datetime('now')), - updated_at TEXT NOT NULL DEFAULT (datetime('now')) - ); - - -- App secrets (encryption keys for sensitive data) - CREATE TABLE IF NOT EXISTS app_secrets ( - key TEXT PRIMARY KEY, - value TEXT NOT NULL, - created_at TEXT NOT NULL DEFAULT (datetime('now')) - ); - - -- Pending unblock queue for orphan eBPF block recovery - CREATE TABLE IF NOT EXISTS pending_unblock ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - source_ip TEXT NOT NULL, - created_at TEXT NOT NULL DEFAULT (datetime('now')), - retry_count INTEGER NOT NULL DEFAULT 0 - ); - - -- Audit trail (WORM: hash-chained, triggers block UPDATE/DELETE) - CREATE TABLE IF NOT EXISTS audit_log ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - ts TEXT NOT NULL, - actor TEXT NOT NULL, - action TEXT NOT NULL, - detail TEXT NOT NULL DEFAULT '{}', - prev_hash TEXT NOT NULL DEFAULT '', - row_hash TEXT NOT NULL DEFAULT '' - ); - - CREATE TRIGGER IF NOT EXISTS audit_log_no_update - BEFORE UPDATE ON audit_log BEGIN - SELECT RAISE(ABORT, 'audit_log is append-only (WORM)'); - END; - - CREATE TRIGGER IF NOT EXISTS audit_log_no_delete - BEFORE DELETE ON audit_log BEGIN - SELECT RAISE(ABORT, 'audit_log is append-only (WORM)'); - END; - ", - )?; - - let conn_ref = &*conn; - - // Seed default user groups on first install (empty table) - let group_count: i64 = conn_ref.query_row("SELECT COUNT(*) FROM user_groups", [], |row| row.get(0))?; - if group_count == 0 { - let all_permissions = serde_json::json!([ - "dashboard:read", - "statistics:read", - "traffic_map:read", - "drops:read", - "ai_detection:read", - "ai_detection:write", - "access_control:read", - "access_control:write", - "geo_block:read", - "geo_block:write", - "dns_filter:read", - "dns_filter:write", - "rate_limit:read", - "rate_limit:write", - "protocol_filter:read", - "protocol_filter:write", - "system:read", - "system:write", - "system:admin", - "users:read", - "users:write", - "users:admin", - "fusion:read", - "fusion:write", - "flow_trace:read", - "flow_trace:write" - ]) - .to_string(); - let viewer_permissions = serde_json::json!([ - "dashboard:read", - "statistics:read", - "traffic_map:read", - "drops:read", - "ai_detection:read", - "access_control:read", - "geo_block:read", - "dns_filter:read", - "rate_limit:read", - "protocol_filter:read", - "system:read", - "fusion:read", - "flow_trace:read" - ]) - .to_string(); - - conn_ref.execute( - "INSERT INTO user_groups (name, description, permissions) VALUES (?1, ?2, ?3)", - params![ - "Administrator", - "Full system access with all permissions", - &all_permissions - ], - )?; - conn_ref.execute( - "INSERT INTO user_groups (name, description, permissions) VALUES (?1, ?2, ?3)", - params!["Viewer", "Read-only access to all modules", &viewer_permissions], - )?; - } - - Ok(()) - } - - // --- ACL --- - pub fn insert_acl_rule( - &self, - ip_version: u8, - direction: &str, - list_type: &str, - ip_address: &str, - port: u16, - ) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "INSERT OR IGNORE INTO acl_rules (ip_version, direction, list_type, ip_address, port) VALUES (?1, ?2, ?3, ?4, ?5)", - params![ip_version, direction, list_type, ip_address, port as i64], - )?; - Ok(()) - } - - pub fn delete_acl_rule( - &self, - ip_version: u8, - direction: &str, - list_type: &str, - ip_address: &str, - port: u16, - ) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "DELETE FROM acl_rules WHERE ip_version = ?1 AND direction = ?2 AND list_type = ?3 AND ip_address = ?4 AND port = ?5", - params![ip_version, direction, list_type, ip_address, port as i64], - )?; - Ok(()) - } - - pub fn load_acl_rules(&self) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare("SELECT ip_version, direction, list_type, ip_address, port FROM acl_rules")?; - let rows = stmt.query_map([], |row| { - Ok(( - row.get::<_, u8>(0)?, - row.get::<_, String>(1)?, - row.get::<_, String>(2)?, - row.get::<_, String>(3)?, - row.get::<_, i64>(4)? as u16, - )) - })?; - let mut results = Vec::new(); - for row in rows { - results.push(row?); - } - Ok(results) - } - - // --- Rate Limit --- - pub fn set_rate_limit(&self, key: &str, value: u64) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "INSERT OR REPLACE INTO rate_limit_config (key, value) VALUES (?1, ?2)", - params![key, value as i64], - )?; - Ok(()) - } - - pub fn load_rate_limit_config(&self) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare("SELECT key, value FROM rate_limit_config")?; - let rows = stmt.query_map([], |row| Ok((row.get::<_, String>(0)?, row.get::<_, i64>(1)? as u64)))?; - let mut results = Vec::new(); - for row in rows { - results.push(row?); - } - Ok(results) - } - - // --- DNS --- - pub fn insert_dns_domain(&self, domain: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "INSERT OR IGNORE INTO dns_blacklist (domain) VALUES (?1)", - params![domain], - )?; - Ok(()) - } - - pub fn delete_dns_domain(&self, domain: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute("DELETE FROM dns_blacklist WHERE domain = ?1", params![domain])?; - Ok(()) - } - - pub fn load_dns_domains(&self) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare("SELECT domain FROM dns_blacklist")?; - let rows = stmt.query_map([], |row| row.get(0))?; - let mut results = Vec::new(); - for row in rows { - results.push(row?); - } - Ok(results) - } - - // --- Geo --- - pub fn insert_geo_country(&self, code: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "INSERT OR IGNORE INTO geo_blocked_countries (country_code) VALUES (?1)", - params![code], - )?; - Ok(()) - } - - pub fn delete_geo_country(&self, code: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "DELETE FROM geo_blocked_countries WHERE country_code = ?1", - params![code], - )?; - Ok(()) - } - - pub fn load_geo_countries(&self) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare("SELECT country_code FROM geo_blocked_countries")?; - let rows = stmt.query_map([], |row| row.get(0))?; - let mut results = Vec::new(); - for row in rows { - results.push(row?); - } - Ok(results) - } - - // --- Settings --- - pub fn get_setting(&self, key: &str) -> Result, Error> { - let conn = self.conn()?; - let result = conn.query_row("SELECT value FROM settings WHERE key = ?1", params![key], |row| { - row.get(0) - }); - match result { - Ok(val) => Ok(Some(val)), - Err(RusqliteError::QueryReturnedNoRows) => Ok(None), - Err(e) => Err(e)?, - } - } - - pub fn set_setting(&self, key: &str, value: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES (?1, ?2)", - params![key, value], - )?; - Ok(()) - } - - // --- App Secrets --- - - pub fn get_app_secret(&self, key: &str) -> Result, Error> { - let conn = self.conn()?; - let result = conn.query_row("SELECT value FROM app_secrets WHERE key = ?1", params![key], |row| { - row.get(0) - }); - match result { - Ok(val) => Ok(Some(val)), - Err(RusqliteError::QueryReturnedNoRows) => Ok(None), - Err(e) => Err(e)?, - } - } - - pub fn set_app_secret(&self, key: &str, value: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "INSERT OR REPLACE INTO app_secrets (key, value) VALUES (?1, ?2)", - params![key, value], - )?; - Ok(()) - } - - // --- Users --- - pub fn find_user(&self, username: &str) -> Result, Error> { - let conn = self.conn()?; - let result = conn.query_row( - "SELECT id, username, password_hash, role, force_password_change FROM users WHERE username = ?1", - params![username], - |row| { - Ok(( - row.get(0)?, - row.get(1)?, - row.get(2)?, - row.get(3)?, - row.get::<_, i64>(4)? != 0, - )) - }, - ); - match result { - Ok(user) => Ok(Some(user)), - Err(RusqliteError::QueryReturnedNoRows) => Ok(None), - Err(e) => Err(e)?, - } - } - - pub fn insert_user( - &self, - username: &str, - password_hash: &str, - role: &str, - force_password_change: bool, - ) -> Result { - let conn = self.conn()?; - conn.execute( - "INSERT INTO users (username, password_hash, role, force_password_change) VALUES (?1, ?2, ?3, ?4)", - params![username, password_hash, role, force_password_change as i64], - ) - .map_err(|e| -> Error { - if e.to_string().contains("UNIQUE constraint") { - DatabaseError::UserAlreadyExists(username.to_string()).into() - } else { - e.into() - } - })?; - Ok(conn.last_insert_rowid()) - } - - pub fn update_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "UPDATE users SET password_hash = ?1, force_password_change = 0 WHERE id = ?2", - params![password_hash, user_id], - )?; - Ok(()) - } - - pub fn user_count(&self) -> Result { - let conn = self.conn()?; - Ok(conn.query_row("SELECT COUNT(*) FROM users", [], |row| row.get(0))?) - } - - pub fn list_users_with_groups(&self) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare( - "SELECT u.id, u.username, u.role, u.force_password_change, u.created_at, \ - g.id, g.name \ - FROM users u \ - LEFT JOIN user_group_members m ON u.id = m.user_id \ - LEFT JOIN user_groups g ON g.id = m.group_id \ - ORDER BY u.id, g.id", - )?; - let rows = stmt.query_map([], |row| { - Ok(( - row.get::<_, i64>(0)?, - row.get::<_, String>(1)?, - row.get::<_, String>(2)?, - row.get::<_, i64>(3)? != 0, - row.get::<_, String>(4)?, - row.get::<_, Option>(5)?, - row.get::<_, Option>(6)?, - )) - })?; - - let mut user_map: HashMap = HashMap::new(); - let mut order: Vec = Vec::new(); - - for row in rows { - let (id, username, role, force_pw, created_at, group_id, group_name) = row?; - let entry = user_map.entry(id).or_insert_with(|| { - order.push(id); - (id, username, role, force_pw, created_at, Vec::new()) - }); - if let (Some(gid), Some(gname)) = (group_id, group_name) { - entry.5.push((gid, gname)); - } - } - - Ok(order.into_iter().filter_map(|id| user_map.remove(&id)).collect()) - } - - pub fn delete_user(&self, user_id: i64) -> Result { - self.cleanup_user_memberships(user_id)?; - let conn = self.conn()?; - let affected = conn.execute("DELETE FROM users WHERE id = ?1", params![user_id])?; - Ok(affected > 0) - } - - pub fn update_user_role(&self, user_id: i64, role: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute("UPDATE users SET role = ?1 WHERE id = ?2", params![role, user_id])?; - Ok(()) - } - - pub fn reset_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "UPDATE users SET password_hash = ?1, force_password_change = 1 WHERE id = ?2", - params![password_hash, user_id], - )?; - Ok(()) - } - - pub fn find_user_by_id(&self, user_id: i64) -> Result, Error> { - let conn = self.conn()?; - let result = conn.query_row( - "SELECT id, username, password_hash, role, force_password_change FROM users WHERE id = ?1", - params![user_id], - |row| { - Ok(( - row.get(0)?, - row.get(1)?, - row.get(2)?, - row.get(3)?, - row.get::<_, i64>(4)? != 0, - )) - }, - ); - match result { - Ok(user) => Ok(Some(user)), - Err(RusqliteError::QueryReturnedNoRows) => Ok(None), - Err(e) => Err(e)?, - } - } - - // --- User Groups --- - pub fn list_user_groups(&self) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = - conn.prepare("SELECT id, name, description, permissions, created_at FROM user_groups ORDER BY id")?; - let rows = stmt.query_map([], |row| { - Ok(( - row.get::<_, i64>(0)?, - row.get::<_, String>(1)?, - row.get::<_, String>(2)?, - row.get::<_, String>(3)?, - row.get::<_, String>(4)?, - )) - })?; - let mut results = Vec::new(); - for row in rows { - results.push(row?); - } - Ok(results) - } - - pub fn create_user_group(&self, name: &str, description: &str, permissions: &str) -> Result { - let conn = self.conn()?; - conn.execute( - "INSERT INTO user_groups (name, description, permissions) VALUES (?1, ?2, ?3)", - params![name, description, permissions], - ) - .map_err(|e| -> Error { - if e.to_string().contains("UNIQUE constraint") { - DatabaseError::GroupAlreadyExists(name).into() - } else { - e.into() - } - })?; - Ok(conn.last_insert_rowid()) - } - - pub fn update_user_group(&self, id: i64, name: &str, description: &str, permissions: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "UPDATE user_groups SET name = ?1, description = ?2, permissions = ?3 WHERE id = ?4", - params![name, description, permissions, id], - )?; - Ok(()) - } - - pub fn delete_user_group(&self, id: i64) -> Result { - let conn = self.conn()?; - conn.execute("DELETE FROM user_group_members WHERE group_id = ?1", params![id])?; - let affected = conn.execute("DELETE FROM user_groups WHERE id = ?1", params![id])?; - Ok(affected > 0) - } - - pub fn get_user_group(&self, id: i64) -> Result, Error> { - let conn = self.conn()?; - let result = conn.query_row( - "SELECT id, name, description, permissions, created_at FROM user_groups WHERE id = ?1", - params![id], - |row| { - Ok(( - row.get::<_, i64>(0)?, - row.get::<_, String>(1)?, - row.get::<_, String>(2)?, - row.get::<_, String>(3)?, - row.get::<_, String>(4)?, - )) - }, - ); - match result { - Ok(group) => Ok(Some(group)), - Err(RusqliteError::QueryReturnedNoRows) => Ok(None), - Err(e) => Err(e)?, - } - } - - // --- User Group Membership --- - pub fn get_user_groups(&self, user_id: i64) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare( - "SELECT g.id, g.name, g.description, g.permissions FROM user_groups g \ - INNER JOIN user_group_members m ON g.id = m.group_id \ - WHERE m.user_id = ?1 ORDER BY g.id", - )?; - let rows = stmt.query_map(params![user_id], |row| { - Ok(( - row.get::<_, i64>(0)?, - row.get::<_, String>(1)?, - row.get::<_, String>(2)?, - row.get::<_, String>(3)?, - )) - })?; - let mut results = Vec::new(); - for row in rows { - results.push(row?); - } - Ok(results) - } - - pub fn set_user_groups(&self, user_id: i64, group_ids: &[i64]) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute("DELETE FROM user_group_members WHERE user_id = ?1", params![user_id])?; - for &gid in group_ids { - conn.execute( - "INSERT INTO user_group_members (user_id, group_id) VALUES (?1, ?2)", - params![user_id, gid], - )?; - } - Ok(()) - } - - pub fn get_user_permissions(&self, user_id: i64) -> Result, Error> { - let groups = self.get_user_groups(user_id)?; - let mut all_perms = HashSet::new(); - for (_id, _name, _desc, perms_json) in groups { - if let Ok(perms) = serde_json::from_str::>(&perms_json) { - for p in perms { - all_perms.insert(p); - } - } - } - let mut result: Vec = all_perms.into_iter().collect(); - result.sort(); - Ok(result) - } - - pub fn cleanup_user_memberships(&self, user_id: i64) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute("DELETE FROM user_group_members WHERE user_id = ?1", params![user_id])?; - Ok(()) - } - - pub fn get_group_member_ids(&self, group_id: i64) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare("SELECT user_id FROM user_group_members WHERE group_id = ?1")?; - let rows = stmt.query_map(params![group_id], |row| row.get::<_, i64>(0))?; - let mut results = Vec::new(); - for row in rows { - results.push(row?); - } - Ok(results) - } - - pub fn get_group_members(&self, group_id: i64) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare( - "SELECT u.id, u.username FROM users u \ - INNER JOIN user_group_members m ON u.id = m.user_id \ - WHERE m.group_id = ?1 ORDER BY u.username", - )?; - let rows = stmt.query_map(params![group_id], |row| { - Ok((row.get::<_, i64>(0)?, row.get::<_, String>(1)?)) - })?; - let mut results = Vec::new(); - for row in rows { - results.push(row?); - } - Ok(results) - } - - // --- Login Rate Limiting --- - pub fn record_login_failure(&self, username: &str) -> Result<(u32, Option), Error> { - let key_count = format!("login_failures:{}", username); - let key_locked = format!("login_locked_until:{}", username); - - let count: u32 = self.get_setting(&key_count)?.and_then(|v| v.parse().ok()).unwrap_or(0) + 1; - - self.set_setting(&key_count, &count.to_string())?; - - if count >= 5 { - let now = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or(Duration::ZERO) - .as_secs(); - let locked_until = now + 900; // 15 minutes - self.set_setting(&key_locked, &locked_until.to_string())?; - Ok((count, Some(locked_until))) - } else { - Ok((count, None)) - } - } - - pub fn check_login_locked(&self, username: &str) -> Result, Error> { - let key_locked = format!("login_locked_until:{}", username); - if let Some(locked_str) = self.get_setting(&key_locked)? - && let Ok(locked_until) = locked_str.parse::() - { - let now = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or(Duration::ZERO) - .as_secs(); - if now < locked_until { - return Ok(Some(locked_until - now)); - } - // Lock expired, clear it - self.clear_login_failures(username)?; - } - Ok(None) - } - - pub fn clear_login_failures(&self, username: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "DELETE FROM settings WHERE key = ?1", - params![format!("login_failures:{}", username)], - )?; - conn.execute( - "DELETE FROM settings WHERE key = ?1", - params![format!("login_locked_until:{}", username)], - )?; - Ok(()) - } - - // --- MCP API Keys --- - - /// Validate an API key and return Claims if valid. - /// Computes HMAC-SHA256 of the key and looks it up in api_keys table. - pub fn validate_api_key(&self, api_key: &str) -> Result, Error> { - let digest = self.hmac_api_key(api_key); - - let conn = self.conn()?; - let result = conn.query_row( - "SELECT id, name, permission_level FROM api_keys WHERE key_hash = ?1", - params![digest], - |row| { - Ok(( - row.get::<_, i64>(0)?, - row.get::<_, String>(1)?, - row.get::<_, String>(2)?, - )) - }, - ); - - match result { - Ok((id, name, level)) => { - // Update last_used_at - let _ = conn.execute( - "UPDATE api_keys SET last_used_at = datetime('now') WHERE id = ?1", - params![id], - ); - - // Build permissions based on permission level - let permissions = match level.as_str() { - "read_write" | "full_access" => vec![ - "dashboard:read".into(), - "statistics:read".into(), - "ai_detection:read".into(), - "ai_detection:write".into(), - "access_control:read".into(), - "access_control:write".into(), - "geo_block:read".into(), - "geo_block:write".into(), - "dns_filter:read".into(), - "dns_filter:write".into(), - "rate_limit:read".into(), - "rate_limit:write".into(), - "system:read".into(), - "system:write".into(), - ], - _ => vec![ - "dashboard:read".into(), - "statistics:read".into(), - "ai_detection:read".into(), - "access_control:read".into(), - "geo_block:read".into(), - "dns_filter:read".into(), - "rate_limit:read".into(), - "system:read".into(), - ], - }; - - Ok(Some(Claims { - sub: -id, // negative ID to distinguish from user IDs - username: format!("api:{}", name), - role: level, - permissions, - exp: usize::MAX, // API keys don't expire (revocation via DB deletion) - })) - } - Err(RusqliteError::QueryReturnedNoRows) => Ok(None), - Err(e) => Err(e)?, - } - } - - // --- SOAR --- - - pub fn insert_playbook( - &self, - name: &str, - trigger_event: &str, - threshold: Option, - count: Option, - window: Option, - cooldown: i64, - ) -> Result { - let conn = self.conn()?; - conn.execute( - "INSERT INTO playbooks (name, trigger_event, condition_threshold, condition_count, condition_window_secs, cooldown_secs) VALUES (?1, ?2, ?3, ?4, ?5, ?6)", - params![name, trigger_event, threshold, count, window, cooldown], - )?; - Ok(conn.last_insert_rowid()) - } - - pub fn insert_playbook_action( - &self, - playbook_id: i64, - action_order: i64, - action_type: &str, - params_json: &str, - ) -> Result { - let conn = self.conn()?; - conn.execute( - "INSERT INTO playbook_actions (playbook_id, action_order, action_type, params) VALUES (?1, ?2, ?3, ?4)", - params![playbook_id, action_order, action_type, params_json], - )?; - Ok(conn.last_insert_rowid()) - } - - /// Test-only helper: direct insert of a SOAR block rule row. Production - /// code goes through `commit_soar_block_to_db` which atomically writes - /// both `soar_block_rules` and `acl_rules` under a transaction. - #[cfg(test)] - pub fn insert_soar_block_rule(&self, source_ip: &str, playbook_id: i64, expires_at: &str) -> Result { - let conn = self.conn()?; - conn.execute( - "INSERT INTO soar_block_rules (source_ip, playbook_id, expires_at) VALUES (?1, ?2, ?3)", - params![source_ip, playbook_id, expires_at], - )?; - Ok(conn.last_insert_rowid()) - } - - pub fn count_active_soar_blocks(&self) -> Result { - let conn = self.conn()?; - let count: u32 = conn.query_row( - "SELECT COUNT(*) FROM soar_block_rules WHERE unblocked_at IS NULL AND expires_at > datetime('now')", - [], - |row| row.get(0), - )?; - Ok(count) - } - - pub fn get_expired_soar_blocks(&self) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare( - "SELECT id, source_ip, playbook_id FROM soar_block_rules WHERE expires_at <= datetime('now') AND unblocked_at IS NULL" - )?; - let rows = stmt.query_map([], |row| { - Ok((row.get::<_, i64>(0)?, row.get::<_, String>(1)?, row.get::<_, i64>(2)?)) - })?; - let mut result = Vec::new(); - for row in rows { - result.push(row?); - } - Ok(result) - } - - /// Get a single SOAR block rule by ID, returning (id, source_ip, playbook_id, expires_at). - pub fn get_soar_block_by_id(&self, id: i64) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = - conn.prepare("SELECT id, source_ip, playbook_id, expires_at FROM soar_block_rules WHERE id = ?1")?; - let mut rows = stmt.query_map(params![id], |row| { - Ok(( - row.get::<_, i64>(0)?, - row.get::<_, String>(1)?, - row.get::<_, i64>(2)?, - row.get::<_, String>(3)?, - )) - })?; - match rows.next() { - Some(row) => Ok(Some(row?)), - None => Ok(None), - } - } - - /// Load all playbooks with their actions in a single JOIN query (avoids N+1). - /// Returns Vec of (playbook fields..., action fields...). - #[allow(clippy::type_complexity)] - pub fn load_playbooks_with_actions( - &self, - ) -> Result< - Vec<( - i64, - String, - bool, - String, - Option, - Option, - Option, - i64, - Option, - Option, - Option, - Option, - )>, - Error, - > { - let conn = self.conn()?; - let mut stmt = conn.prepare( - "SELECT p.id, p.name, p.enabled, p.trigger_event, p.condition_threshold, \ - p.condition_count, p.condition_window_secs, p.cooldown_secs, \ - a.id, a.action_order, a.action_type, a.params \ - FROM playbooks p \ - LEFT JOIN playbook_actions a ON a.playbook_id = p.id \ - ORDER BY p.id, a.action_order", - )?; - let rows = stmt.query_map([], |row| { - Ok(( - row.get::<_, i64>(0)?, - row.get::<_, String>(1)?, - row.get::<_, bool>(2)?, - row.get::<_, String>(3)?, - row.get::<_, Option>(4)?, - row.get::<_, Option>(5)?, - row.get::<_, Option>(6)?, - row.get::<_, i64>(7)?, - row.get::<_, Option>(8)?, - row.get::<_, Option>(9)?, - row.get::<_, Option>(10)?, - row.get::<_, Option>(11)?, - )) - })?; - let mut result = Vec::new(); - for row in rows { - result.push(row?); - } - Ok(result) - } - - pub fn mark_soar_block_unblocked(&self, id: i64) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "UPDATE soar_block_rules SET unblocked_at = datetime('now') WHERE id = ?1", - params![id], - )?; - Ok(()) - } - - pub fn get_active_soar_blocks(&self) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare( - "SELECT id, source_ip, playbook_id, expires_at FROM soar_block_rules WHERE unblocked_at IS NULL AND expires_at > datetime('now')" - )?; - let rows = stmt.query_map([], |row| { - Ok(( - row.get::<_, i64>(0)?, - row.get::<_, String>(1)?, - row.get::<_, i64>(2)?, - row.get::<_, String>(3)?, - )) - })?; - let mut result = Vec::new(); - for row in rows { - result.push(row?); - } - Ok(result) - } - - // --- Pending Unblock --- - pub fn insert_pending_unblock(&self, source_ip: &str) -> Result { - let conn = self.conn()?; - conn.execute( - "INSERT INTO pending_unblock (source_ip) VALUES (?1)", - params![source_ip], - )?; - Ok(conn.last_insert_rowid()) - } - - pub fn load_pending_unblocks(&self) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare("SELECT id, source_ip, retry_count FROM pending_unblock ORDER BY id")?; - let rows = stmt.query_map([], |row| { - Ok((row.get::<_, i64>(0)?, row.get::<_, String>(1)?, row.get::<_, i64>(2)?)) - })?; - let mut result = Vec::new(); - for row in rows { - result.push(row?); - } - Ok(result) - } - - pub fn delete_pending_unblock(&self, id: i64) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute("DELETE FROM pending_unblock WHERE id = ?1", params![id])?; - Ok(()) - } - - pub fn increment_pending_unblock_retry(&self, id: i64) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "UPDATE pending_unblock SET retry_count = retry_count + 1 WHERE id = ?1", - params![id], - )?; - Ok(()) - } - - pub fn insert_soar_execution( - &self, - playbook_id: i64, - source_ip: Option<&str>, - trigger_event: &str, - actions_json: &str, - ) -> Result { - let conn = self.conn()?; - conn.execute( - "INSERT INTO soar_executions (playbook_id, source_ip, trigger_event, actions_executed) VALUES (?1, ?2, ?3, ?4)", - params![playbook_id, source_ip, trigger_event, actions_json], - )?; - Ok(conn.last_insert_rowid()) - } - - pub fn update_playbook_enabled(&self, id: i64, enabled: bool) -> Result { - let conn = self.conn()?; - let rows = conn.execute( - "UPDATE playbooks SET enabled = ?2, updated_at = datetime('now') WHERE id = ?1", - params![id, enabled as i32], - )?; - Ok(rows > 0) - } - - pub fn delete_playbook(&self, id: i64) -> Result { - let conn = self.conn()?; - let rows = conn.execute("DELETE FROM playbooks WHERE id = ?1", params![id])?; - Ok(rows > 0) - } - - pub fn insert_playbook_condition( - &self, - playbook_id: i64, - condition_type: &str, - operator: &str, - value: &str, - value2: Option<&str>, - ) -> Result { - let conn = self.conn()?; - conn.execute( - "INSERT INTO playbook_conditions (playbook_id, condition_type, operator, value, value2) \ - VALUES (?1, ?2, ?3, ?4, ?5)", - params![playbook_id, condition_type, operator, value, value2], - )?; - Ok(conn.last_insert_rowid()) - } - - #[allow(clippy::type_complexity)] - pub fn load_all_playbook_conditions( - &self, - ) -> Result)>, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare( - "SELECT id, playbook_id, condition_type, operator, value, value2 \ - FROM playbook_conditions ORDER BY playbook_id, id", - )?; - let rows = stmt.query_map([], |row| { - Ok(( - row.get::<_, i64>(0)?, - row.get::<_, i64>(1)?, - row.get::<_, String>(2)?, - row.get::<_, String>(3)?, - row.get::<_, String>(4)?, - row.get::<_, Option>(5)?, - )) - })?; - let mut result = Vec::new(); - for row in rows { - result.push(row?); - } - Ok(result) - } - - #[allow(clippy::type_complexity)] - pub fn list_soar_executions( - &self, - limit: i64, - ) -> Result, String, String, String)>, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare( - "SELECT id, playbook_id, source_ip, trigger_event, actions_executed, executed_at FROM soar_executions ORDER BY executed_at DESC LIMIT ?1" - )?; - let rows = stmt.query_map(params![limit], |row| { - Ok(( - row.get::<_, i64>(0)?, - row.get::<_, i64>(1)?, - row.get::<_, Option>(2)?, - row.get::<_, String>(3)?, - row.get::<_, String>(4)?, - row.get::<_, String>(5)?, - )) - })?; - let mut result = Vec::new(); - for row in rows { - result.push(row?); - } - Ok(result) - } - - pub fn has_manual_acl_rule(&self, ip_address: &str) -> Result { - let conn = self.conn()?; - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM acl_rules WHERE ip_address = ?1 AND list_type = 'blacklist'", - params![ip_address], - |row| row.get(0), - )?; - Ok(count > 0) - } - - // --- Admin Whitelist --- - - pub fn load_admin_whitelist(&self) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare("SELECT ip FROM admin_whitelist")?; - let rows = stmt.query_map([], |row| row.get::<_, String>(0))?; - let mut result = Vec::new(); - for row in rows { - result.push(row?); - } - Ok(result) - } - - pub fn insert_admin_whitelist(&self, ip: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute("INSERT OR IGNORE INTO admin_whitelist (ip) VALUES (?1)", params![ip])?; - Ok(()) - } - - pub fn delete_admin_whitelist(&self, ip: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute("DELETE FROM admin_whitelist WHERE ip = ?1", params![ip])?; - Ok(()) - } - - // --- Notification Config --- - - pub fn get_notification_config(&self, channel: &str) -> Result, Error> { - let conn = self.conn()?; - match conn.query_row( - "SELECT config_json FROM notification_config WHERE channel = ?1 AND enabled = 1", - params![channel], - |row| row.get::<_, String>(0), - ) { - Ok(json) => Ok(Some(json)), - Err(RusqliteError::QueryReturnedNoRows) => Ok(None), - Err(e) => Err(e)?, - } - } - - pub fn set_notification_config(&self, channel: &str, config_json: &str) -> Result<(), Error> { - let conn = self.conn()?; - conn.execute( - "INSERT INTO notification_config (channel, config_json) VALUES (?1, ?2) \ - ON CONFLICT(channel) DO UPDATE SET config_json = ?2, updated_at = datetime('now')", - params![channel, config_json], - )?; - Ok(()) - } - - // --- API Key Management --- - - pub fn insert_api_key(&self, key_hash: &str, name: &str, permission_level: &str) -> Result { - let conn = self.conn()?; - conn.execute( - "INSERT INTO api_keys (key_hash, name, permission_level) VALUES (?1, ?2, ?3)", - params![key_hash, name, permission_level], - )?; - Ok(conn.last_insert_rowid()) - } - - #[allow(clippy::type_complexity)] - pub fn list_api_keys(&self) -> Result)>, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare("SELECT id, name, permission_level, created_at, last_used_at FROM api_keys")?; - let rows = stmt.query_map([], |row| { - Ok(( - row.get::<_, i64>(0)?, - row.get::<_, String>(1)?, - row.get::<_, String>(2)?, - row.get::<_, String>(3)?, - row.get::<_, Option>(4)?, - )) - })?; - let mut result = Vec::new(); - for row in rows { - result.push(row?); - } - Ok(result) - } - - pub fn delete_api_key(&self, id: i64) -> Result { - let conn = self.conn()?; - let affected = conn.execute("DELETE FROM api_keys WHERE id = ?1", params![id])?; - Ok(affected > 0) - } - - // --- Stats Aggregation --- - - /// Count SOAR executions in the last N days. - pub fn count_weekly_executions(&self, days: i64) -> Result { - let conn = self.conn()?; - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM soar_executions WHERE executed_at >= datetime('now', ?1)", - params![format!("-{} days", days)], - |row| row.get(0), - )?; - Ok(count as u64) - } - - /// Count SOAR blocks created in the last N days. - pub fn count_weekly_blocks(&self, days: i64) -> Result { - let conn = self.conn()?; - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM soar_block_rules WHERE created_at >= datetime('now', ?1)", - params![format!("-{} days", days)], - |row| row.get(0), - )?; - Ok(count as u64) - } - - /// Count SOAR unblocks in the last N days. - pub fn count_weekly_unblocks(&self, days: i64) -> Result { - let conn = self.conn()?; - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM soar_block_rules WHERE unblocked_at IS NOT NULL AND unblocked_at >= datetime('now', ?1)", - params![format!("-{} days", days)], - |row| row.get(0), - )?; - Ok(count as u64) - } - - /// Get threat breakdown by trigger_event in the last N days. - pub fn weekly_threat_breakdown(&self, days: i64) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare( - "SELECT trigger_event, COUNT(*) FROM soar_executions WHERE executed_at >= datetime('now', ?1) GROUP BY trigger_event ORDER BY COUNT(*) DESC" - )?; - let rows = stmt.query_map(params![format!("-{} days", days)], |row| { - Ok((row.get::<_, String>(0)?, row.get::<_, i64>(1)? as u64)) - })?; - let mut result = Vec::new(); - for row in rows { - result.push(row?); - } - Ok(result) - } - - /// Get top blocked IPs in the last N days. - pub fn weekly_top_ips(&self, days: i64, limit: i64) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare( - "SELECT source_ip, COUNT(*) as cnt FROM soar_block_rules WHERE created_at >= datetime('now', ?1) GROUP BY source_ip ORDER BY cnt DESC LIMIT ?2" - )?; - let rows = stmt.query_map(params![format!("-{} days", days), limit], |row| { - Ok((row.get::<_, String>(0)?, row.get::<_, i64>(1)? as u64)) - })?; - let mut result = Vec::new(); - for row in rows { - result.push(row?); - } - Ok(result) - } - - /// Count current ACL rules. - pub fn count_acl_rules(&self) -> Result { - let conn = self.conn()?; - let count: i64 = conn.query_row("SELECT COUNT(*) FROM acl_rules", [], |row| row.get(0))?; - Ok(count as u64) - } - - // --- Default Playbooks --- - - pub fn seed_default_playbooks(&self) -> Result<(), Error> { - let conn = self.conn()?; - let count: i64 = conn.query_row("SELECT COUNT(*) FROM playbooks", [], |row| row.get(0))?; - if count > 0 { - return Ok(()); - } - drop(conn); - - // 1. brute_force_block: brute_force, count 5 in 60s → block_ip(3600s) + send_telegram + log - let pb1 = self.insert_playbook("brute_force_block", "brute_force", None, Some(5), Some(60), 600)?; - self.insert_playbook_action(pb1, 1, "block_ip", r#"{"ttl_secs": 3600}"#)?; - self.insert_playbook_action(pb1, 2, "send_telegram", "{}")?; - self.insert_playbook_action(pb1, 3, "log", r#"{"level": "warn"}"#)?; - self.insert_playbook_condition(pb1, "frequency", ">=", "5", Some("60"))?; - - // 2. port_scan_alert: port_scan, threshold 0.7 → send_telegram + log (no block) - let pb2 = self.insert_playbook("port_scan_alert", "port_scan", Some(0.7), None, None, 300)?; - self.insert_playbook_action(pb2, 1, "send_telegram", "{}")?; - self.insert_playbook_action(pb2, 2, "log", r#"{"level": "warn"}"#)?; - self.insert_playbook_condition(pb2, "threshold", ">=", "0.7", None)?; - - // 3. fusion_c2_multi_source_block — C2 beacon observed by ≥2 sources - // (e.g. Suricata trojan-activity + Beaconing CV + ML c2 class) is - // the highest-precision fusion signal we ship. Block for 1h and - // notify, no solo-source threshold so single-source C2 hits still - // require the solo playbook below to act. - let pb3 = self.insert_playbook("fusion_c2_multi_source_block", "c2_beacon", None, None, None, 600)?; - self.insert_playbook_action(pb3, 1, "block_ip", r#"{"ttl_secs": 3600}"#)?; - self.insert_playbook_action(pb3, 2, "send_telegram", "{}")?; - self.insert_playbook_action(pb3, 3, "log", r#"{"level": "warn"}"#)?; - self.insert_playbook_condition(pb3, "multi_source_min", ">=", "2", None)?; - - // 4. fusion_c2_suricata_solo_high_block — the escape hatch for - // Suricata signature hits with very high confidence (>=0.95). - // Lets known-good rules fire without waiting for agreement from a - // second source, matching how analysts intuitively treat a - // signature "dead-on" match. - let pb4 = self.insert_playbook("fusion_c2_suricata_solo_high_block", "c2_beacon", None, None, None, 600)?; - self.insert_playbook_action(pb4, 1, "block_ip", r#"{"ttl_secs": 3600}"#)?; - self.insert_playbook_action(pb4, 2, "send_telegram", "{}")?; - self.insert_playbook_action(pb4, 3, "log", r#"{"level": "warn"}"#)?; - self.insert_playbook_condition(pb4, "single_source_high", "==", "Suricata", Some("0.95"))?; - - Ok(()) - } - - // --- Audit Log (WORM, hash-chained) --- - - /// Compute the row hash for an audit_log entry. - /// Formula: sha256_hex(ts || 0x00 || actor || 0x00 || action || 0x00 || detail || 0x00 || prev_hash) - fn audit_row_hash(ts: &str, actor: &str, action: &str, detail: &str, prev_hash: &str) -> String { - use sha2::{Digest, Sha256}; - let mut h = Sha256::new(); - for part in [ts, actor, action, detail, prev_hash] { - h.update(part.as_bytes()); - h.update([0u8]); - } - let out = h.finalize(); - let mut hex = String::with_capacity(64); - for byte in out { - use std::fmt::Write; - let _ = write!(&mut hex, "{:02x}", byte); - } - hex - } - - /// Insert an audit trail entry. Runs in a transaction so the (prev_hash - /// lookup, row_hash compute, insert) sequence is atomic and serializable. - pub fn insert_audit_log(&self, actor: &str, action: &str, detail: &str) -> Result<(), Error> { - let mut conn = self.conn()?; - let ts = Utc::now().format("%Y-%m-%d %H:%M:%S").to_string(); - - let tx = conn.transaction()?; - let prev_hash: String = tx - .query_row("SELECT row_hash FROM audit_log ORDER BY id DESC LIMIT 1", [], |row| { - row.get(0) - }) - .unwrap_or_default(); - - let row_hash = Self::audit_row_hash(&ts, actor, action, detail, &prev_hash); - tx.execute( - "INSERT INTO audit_log (ts, actor, action, detail, prev_hash, row_hash) VALUES (?1, ?2, ?3, ?4, ?5, ?6)", - params![ts, actor, action, detail, prev_hash, row_hash], - )?; - tx.commit()?; - Ok(()) - } - - /// List recent audit log entries (most recent first, max 200). - pub fn list_audit_logs(&self) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = - conn.prepare("SELECT id, actor, action, detail, ts FROM audit_log ORDER BY id DESC LIMIT 200")?; - let rows = stmt - .query_map([], |row| { - Ok(AuditLogEntry { - id: row.get(0)?, - actor: row.get(1)?, - action: row.get(2)?, - detail: row.get(3)?, - created_at: row.get(4)?, - }) - })? - .filter_map(|r| r.ok()) - .collect(); - Ok(rows) - } - - /// Read audit entries whose `action` matches exactly, newest-first, - /// capped at `limit`. Drives the fusion explain endpoint. - pub fn list_audit_logs_by_action(&self, action: &str, limit: i64) -> Result, Error> { - let conn = self.conn()?; - let mut stmt = conn.prepare( - "SELECT id, actor, action, detail, ts FROM audit_log WHERE action = ?1 ORDER BY id DESC LIMIT ?2", - )?; - let rows = stmt - .query_map(params![action, limit], |row| { - Ok(AuditLogEntry { - id: row.get(0)?, - actor: row.get(1)?, - action: row.get(2)?, - detail: row.get(3)?, - created_at: row.get(4)?, - }) - })? - .filter_map(|r| r.ok()) - .collect(); - Ok(rows) - } - - /// Walk the entire audit_log in id order and verify the hash chain. - /// Returns `Ok(count)` on success; returns `Err` at the first mismatch, - /// naming the offending row id and the kind of mismatch. - pub fn verify_audit_log_chain(&self) -> Result { - let conn = self.conn()?; - let mut stmt = - conn.prepare("SELECT id, ts, actor, action, detail, prev_hash, row_hash FROM audit_log ORDER BY id ASC")?; - let mut rows = stmt.query([])?; - - let mut expected_prev = String::new(); - let mut count = 0usize; - while let Some(row) = rows.next()? { - let id: i64 = row.get(0)?; - let ts: String = row.get(1)?; - let actor: String = row.get(2)?; - let action: String = row.get(3)?; - let detail: String = row.get(4)?; - let prev_hash: String = row.get(5)?; - let row_hash: String = row.get(6)?; - - if prev_hash != expected_prev { - return Err(DatabaseError::AuditPrevHashMismatch(id, expected_prev, prev_hash).into()); - } - let computed = Self::audit_row_hash(&ts, &actor, &action, &detail, &prev_hash); - if computed != row_hash { - return Err(DatabaseError::AuditRowHashMismatch(id, computed, row_hash).into()); - } - expected_prev = row_hash; - count += 1; - } - Ok(count) - } -} - -// --- Aggregate repository trait implementations --- -// -// All trait methods are forward-only wrappers to the inherent impl above. -// The traits exist to enforce aggregate boundaries: callers take -// `Arc` instead of `Arc` so they see only the -// methods of their own aggregate. See `docs/strategy/DOMAIN_MAP.md` §2. - -impl AclRepo for Database { - fn insert_acl_rule( - &self, - ip_version: u8, - direction: &str, - list_type: &str, - ip_address: &str, - port: u16, - ) -> Result<(), Error> { - self.insert_acl_rule(ip_version, direction, list_type, ip_address, port) - } - fn delete_acl_rule( - &self, - ip_version: u8, - direction: &str, - list_type: &str, - ip_address: &str, - port: u16, - ) -> Result<(), Error> { - self.delete_acl_rule(ip_version, direction, list_type, ip_address, port) - } - fn has_manual_acl_rule(&self, ip_address: &str) -> Result { - self.has_manual_acl_rule(ip_address) - } - fn load_admin_whitelist(&self) -> Result, Error> { - self.load_admin_whitelist() - } - fn insert_admin_whitelist(&self, ip: &str) -> Result<(), Error> { - self.insert_admin_whitelist(ip) - } - fn delete_admin_whitelist(&self, ip: &str) -> Result<(), Error> { - self.delete_admin_whitelist(ip) - } -} - -impl EnforcementRepo for Database { - fn set_rate_limit(&self, key: &str, value: u64) -> Result<(), Error> { - self.set_rate_limit(key, value) - } - fn insert_dns_domain(&self, domain: &str) -> Result<(), Error> { - self.insert_dns_domain(domain) - } - fn delete_dns_domain(&self, domain: &str) -> Result<(), Error> { - self.delete_dns_domain(domain) - } - fn insert_geo_country(&self, code: &str) -> Result<(), Error> { - self.insert_geo_country(code) - } - fn delete_geo_country(&self, code: &str) -> Result<(), Error> { - self.delete_geo_country(code) - } -} - -impl SettingRepo for Database { - fn get_setting(&self, key: &str) -> Result, Error> { - self.get_setting(key) - } - fn set_setting(&self, key: &str, value: &str) -> Result<(), Error> { - self.set_setting(key, value) - } - fn get_app_secret(&self, key: &str) -> Result, Error> { - self.get_app_secret(key) - } - fn set_app_secret(&self, key: &str, plaintext: &str) -> Result<(), Error> { - self.set_app_secret(key, plaintext) - } - fn get_notification_config(&self, channel: &str) -> Result, Error> { - self.get_notification_config(channel) - } - fn set_notification_config(&self, channel: &str, config_json: &str) -> Result<(), Error> { - self.set_notification_config(channel, config_json) - } -} - -impl IdentityRepo for Database { - fn find_user(&self, username: &str) -> Result, Error> { - self.find_user(username) - } - fn find_user_by_id(&self, user_id: i64) -> Result, Error> { - self.find_user_by_id(user_id) - } - fn insert_user( - &self, - username: &str, - password_hash: &str, - role: &str, - force_password_change: bool, - ) -> Result { - self.insert_user(username, password_hash, role, force_password_change) - } - fn update_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> { - self.update_user_password(user_id, password_hash) - } - fn list_users_with_groups(&self) -> Result, Error> { - self.list_users_with_groups() - } - fn delete_user(&self, user_id: i64) -> Result { - self.delete_user(user_id) - } - fn update_user_role(&self, user_id: i64, role: &str) -> Result<(), Error> { - self.update_user_role(user_id, role) - } - fn reset_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> { - self.reset_user_password(user_id, password_hash) - } - fn list_user_groups(&self) -> Result, Error> { - self.list_user_groups() - } - fn create_user_group(&self, name: &str, description: &str, permissions: &str) -> Result { - self.create_user_group(name, description, permissions) - } - fn update_user_group(&self, id: i64, name: &str, description: &str, permissions: &str) -> Result<(), Error> { - self.update_user_group(id, name, description, permissions) - } - fn delete_user_group(&self, id: i64) -> Result { - self.delete_user_group(id) - } - fn get_user_group(&self, id: i64) -> Result, Error> { - self.get_user_group(id) - } - fn get_user_groups(&self, user_id: i64) -> Result, Error> { - self.get_user_groups(user_id) - } - fn set_user_groups(&self, user_id: i64, group_ids: &[i64]) -> Result<(), Error> { - self.set_user_groups(user_id, group_ids) - } - fn get_user_permissions(&self, user_id: i64) -> Result, Error> { - self.get_user_permissions(user_id) - } - fn get_group_member_ids(&self, group_id: i64) -> Result, Error> { - self.get_group_member_ids(group_id) - } - fn get_group_members(&self, group_id: i64) -> Result, Error> { - self.get_group_members(group_id) - } - fn record_login_failure(&self, username: &str) -> Result<(u32, Option), Error> { - self.record_login_failure(username) - } - fn check_login_locked(&self, username: &str) -> Result, Error> { - self.check_login_locked(username) - } - fn clear_login_failures(&self, username: &str) -> Result<(), Error> { - self.clear_login_failures(username) - } -} - -impl SoarRepo for Database { - fn load_playbooks_with_actions(&self) -> Result, Error> { - self.load_playbooks_with_actions() - } - fn update_playbook_enabled(&self, id: i64, enabled: bool) -> Result { - self.update_playbook_enabled(id, enabled) - } - fn delete_playbook(&self, id: i64) -> Result { - self.delete_playbook(id) - } - fn seed_default_playbooks(&self) -> Result<(), Error> { - self.seed_default_playbooks() - } - fn load_all_playbook_conditions(&self) -> Result)>, Error> { - self.load_all_playbook_conditions() - } - fn count_active_soar_blocks(&self) -> Result { - self.count_active_soar_blocks() - } - fn get_active_soar_blocks(&self) -> Result, Error> { - self.get_active_soar_blocks() - } - fn get_soar_block_by_id(&self, id: i64) -> Result, Error> { - self.get_soar_block_by_id(id) - } - fn get_expired_soar_blocks(&self) -> Result, Error> { - self.get_expired_soar_blocks() - } - fn mark_soar_block_unblocked(&self, id: i64) -> Result<(), Error> { - self.mark_soar_block_unblocked(id) - } - fn insert_pending_unblock(&self, source_ip: &str) -> Result { - self.insert_pending_unblock(source_ip) - } - fn load_pending_unblocks(&self) -> Result, Error> { - self.load_pending_unblocks() - } - fn delete_pending_unblock(&self, id: i64) -> Result<(), Error> { - self.delete_pending_unblock(id) - } - fn increment_pending_unblock_retry(&self, id: i64) -> Result<(), Error> { - self.increment_pending_unblock_retry(id) - } - fn insert_soar_execution( - &self, - playbook_id: i64, - source_ip: Option<&str>, - trigger_event: &str, - actions_json: &str, - ) -> Result { - self.insert_soar_execution(playbook_id, source_ip, trigger_event, actions_json) - } - fn list_soar_executions(&self, limit: i64) -> Result, Error> { - self.list_soar_executions(limit) - } - - // --- intra-aggregate atomic operations --- - - fn insert_playbook_atomic( - &self, - name: &str, - trigger_event: &str, - threshold: Option, - count: Option, - window: Option, - cooldown: i64, - actions: &[(i64, String, String)], - conditions: &[(String, String, String, Option)], - ) -> Result { - let mut conn = self.conn()?; - let tx = conn.transaction()?; - tx.execute( - "INSERT INTO playbooks (name, trigger_event, condition_threshold, condition_count, condition_window_secs, cooldown_secs) VALUES (?1, ?2, ?3, ?4, ?5, ?6)", - params![name, trigger_event, threshold, count, window, cooldown], - )?; - let playbook_id = tx.last_insert_rowid(); - for (action_order, action_type, params_json) in actions { - tx.execute( - "INSERT INTO playbook_actions (playbook_id, action_order, action_type, params) VALUES (?1, ?2, ?3, ?4)", - params![playbook_id, action_order, action_type, params_json], - )?; - } - for (condition_type, operator, value, value2) in conditions { - tx.execute( - "INSERT INTO playbook_conditions (playbook_id, condition_type, operator, value, value2) VALUES (?1, ?2, ?3, ?4, ?5)", - params![playbook_id, condition_type, operator, value, value2.as_deref()], - )?; - } - tx.commit()?; - Ok(playbook_id) - } - - fn update_playbook_atomic( - &self, - id: i64, - row: &UpdatePlaybookRow, - actions: &[(i64, String, String)], - conditions: &[(String, String, String, Option)], - ) -> Result { - let mut conn = self.conn()?; - let tx = conn.transaction()?; - let rows_updated = tx.execute( - "UPDATE playbooks SET name = ?2, trigger_event = ?3, condition_threshold = ?4, \ - condition_count = ?5, condition_window_secs = ?6, cooldown_secs = ?7, \ - updated_at = datetime('now') WHERE id = ?1", - params![ - id, - row.name, - row.trigger_event, - row.condition_threshold, - row.condition_count, - row.condition_window_secs, - row.cooldown_secs - ], - )?; - if rows_updated == 0 { - return Ok(false); - } - tx.execute("DELETE FROM playbook_actions WHERE playbook_id = ?1", params![id])?; - tx.execute("DELETE FROM playbook_conditions WHERE playbook_id = ?1", params![id])?; - for (action_order, action_type, params_json) in actions { - tx.execute( - "INSERT INTO playbook_actions (playbook_id, action_order, action_type, params) VALUES (?1, ?2, ?3, ?4)", - params![id, action_order, action_type, params_json], - )?; - } - for (condition_type, operator, value, value2) in conditions { - tx.execute( - "INSERT INTO playbook_conditions (playbook_id, condition_type, operator, value, value2) VALUES (?1, ?2, ?3, ?4, ?5)", - params![id, condition_type, operator, value, value2.as_deref()], - )?; - } - tx.commit()?; - Ok(true) - } -} - -impl StatsRepo for Database { - fn count_weekly_executions(&self, days: i64) -> Result { - self.count_weekly_executions(days) - } - fn count_weekly_blocks(&self, days: i64) -> Result { - self.count_weekly_blocks(days) - } - fn count_weekly_unblocks(&self, days: i64) -> Result { - self.count_weekly_unblocks(days) - } - fn weekly_threat_breakdown(&self, days: i64) -> Result, Error> { - self.weekly_threat_breakdown(days) - } - fn weekly_top_ips(&self, days: i64, limit: i64) -> Result, Error> { - self.weekly_top_ips(days, limit) - } - fn count_acl_rules(&self) -> Result { - self.count_acl_rules() - } -} - -impl AuditRepo for Database { - fn insert_audit_log(&self, actor: &str, action: &str, detail: &str) -> Result<(), Error> { - self.insert_audit_log(actor, action, detail) - } - fn list_audit_logs_by_action(&self, action: &str, limit: i64) -> Result, Error> { - self.list_audit_logs_by_action(action, limit) - } - fn verify_audit_log_chain(&self) -> Result { - self.verify_audit_log_chain() - } -} - -impl ApiKeyRepo for Database { - fn validate_api_key(&self, api_key: &str) -> Result, Error> { - self.validate_api_key(api_key) - } - fn hmac_api_key(&self, raw_key: &str) -> String { - self.hmac_api_key(raw_key) - } - fn insert_api_key(&self, key_hash: &str, name: &str, permission_level: &str) -> Result { - self.insert_api_key(key_hash, name, permission_level) - } - fn list_api_keys(&self) -> Result, Error> { - self.list_api_keys() - } - fn delete_api_key(&self, id: i64) -> Result { - self.delete_api_key(id) - } -} - -impl DbAdminRepo for Database { - /// Commit a SOAR-driven block to both `soar_block_rules` and - /// `acl_rules` in one transaction. Callers must have already installed - /// the eBPF block before calling this, and are responsible for removing - /// the eBPF block if this returns Err. - fn commit_soar_block_to_db( - &self, - source_ip: &str, - ip_version: u8, - playbook_id: i64, - expires_at: &str, - ) -> Result { - let mut conn = self.conn()?; - let tx = conn.transaction()?; - tx.execute( - "INSERT INTO soar_block_rules (source_ip, playbook_id, expires_at) VALUES (?1, ?2, ?3)", - params![source_ip, playbook_id, expires_at], - )?; - let soar_block_id = tx.last_insert_rowid(); - tx.execute( - "INSERT OR IGNORE INTO acl_rules (ip_version, direction, list_type, ip_address, port) VALUES (?1, ?2, ?3, ?4, ?5)", - params![ip_version, "source", "blacklist", source_ip, 0i64], - )?; - tx.commit()?; - Ok(soar_block_id) - } - - /// Clear a SOAR-driven block: remove the `acl_rules` entry and mark - /// the `soar_block_rules` row as unblocked in one transaction. - /// Callers handle eBPF unblock separately. - fn commit_soar_unblock_to_db(&self, soar_block_id: i64, ip_version: u8, source_ip: &str) -> Result<(), Error> { - let mut conn = self.conn()?; - let tx = conn.transaction()?; - tx.execute( - "DELETE FROM acl_rules WHERE ip_version = ?1 AND direction = ?2 AND list_type = ?3 AND ip_address = ?4 AND port = ?5", - params![ip_version, "source", "blacklist", source_ip, 0i64], - )?; - tx.execute( - "UPDATE soar_block_rules SET unblocked_at = datetime('now') WHERE id = ?1", - params![soar_block_id], - )?; - tx.commit()?; - Ok(()) - } -} - -#[cfg(test)] -mod tests { - use super::*; - - fn test_db() -> Database { - Database::new(":memory:").expect("Failed to create test database") - } - - #[test] - fn test_create_tables() { - let _db = test_db(); - } - - #[test] - fn test_user_crud() { - let db = test_db(); - assert_eq!(db.user_count().unwrap(), 0); - - db.insert_user("admin", "hash123", "admin", true).unwrap(); - assert_eq!(db.user_count().unwrap(), 1); - - let user = db.find_user("admin").unwrap().unwrap(); - assert_eq!(user.0, 1); // id - assert_eq!(user.1, "admin"); // username - assert_eq!(user.2, "hash123"); // password_hash - assert_eq!(user.3, "admin"); // role - assert!(user.4); // force_password_change - } - - #[test] - fn test_user_duplicate() { - let db = test_db(); - db.insert_user("admin", "hash", "admin", false).unwrap(); - let result = db.insert_user("admin", "hash2", "admin", false); - assert!(result.is_err()); - } - - #[test] - fn test_update_password_clears_force_change() { - let db = test_db(); - db.insert_user("admin", "old_hash", "admin", true).unwrap(); - - let user = db.find_user("admin").unwrap().unwrap(); - assert!(user.4); // force_password_change = true - - db.update_user_password(user.0, "new_hash").unwrap(); - - let user = db.find_user("admin").unwrap().unwrap(); - assert!(!user.4); // force_password_change = false - assert_eq!(user.2, "new_hash"); - } - - #[test] - fn test_settings_crud() { - let db = test_db(); - assert_eq!(db.get_setting("foo").unwrap(), None); - - db.set_setting("foo", "bar").unwrap(); - assert_eq!(db.get_setting("foo").unwrap(), Some("bar".to_string())); - - db.set_setting("foo", "baz").unwrap(); - assert_eq!(db.get_setting("foo").unwrap(), Some("baz".to_string())); - } - - #[test] - fn test_acl_crud() { - let db = test_db(); - db.insert_acl_rule(4, "source", "blacklist", "192.168.1.1", 80).unwrap(); - let rules = db.load_acl_rules().unwrap(); - assert_eq!(rules.len(), 1); - assert_eq!( - rules[0], - ( - 4, - "source".to_string(), - "blacklist".to_string(), - "192.168.1.1".to_string(), - 80 - ) - ); - - db.delete_acl_rule(4, "source", "blacklist", "192.168.1.1", 80).unwrap(); - let rules = db.load_acl_rules().unwrap(); - assert!(rules.is_empty()); - } - - #[test] - fn test_dns_crud() { - let db = test_db(); - db.insert_dns_domain("evil.com").unwrap(); - let domains = db.load_dns_domains().unwrap(); - assert_eq!(domains, vec!["evil.com"]); - - db.delete_dns_domain("evil.com").unwrap(); - assert!(db.load_dns_domains().unwrap().is_empty()); - } - - #[test] - fn test_geo_crud() { - let db = test_db(); - db.insert_geo_country("CN").unwrap(); - let countries = db.load_geo_countries().unwrap(); - assert_eq!(countries, vec!["CN"]); - - db.delete_geo_country("CN").unwrap(); - assert!(db.load_geo_countries().unwrap().is_empty()); - } - - #[test] - fn test_rate_limit_crud() { - let db = test_db(); - db.set_rate_limit("packet_rate", 1000).unwrap(); - let configs = db.load_rate_limit_config().unwrap(); - assert_eq!(configs.len(), 1); - assert_eq!(configs[0], ("packet_rate".to_string(), 1000)); - } - - #[test] - fn test_login_lockout() { - let db = test_db(); - - // First 4 failures don't lock - for i in 1..5 { - let (count, locked) = db.record_login_failure("admin").unwrap(); - assert_eq!(count, i); - assert!(locked.is_none()); - } - - // 5th failure triggers lock - let (count, locked) = db.record_login_failure("admin").unwrap(); - assert_eq!(count, 5); - assert!(locked.is_some()); - - // Check locked - let remaining = db.check_login_locked("admin").unwrap(); - assert!(remaining.is_some()); - assert!(remaining.unwrap() > 0); - - // Clear and verify - db.clear_login_failures("admin").unwrap(); - let remaining = db.check_login_locked("admin").unwrap(); - assert!(remaining.is_none()); - } - - #[test] - fn test_find_nonexistent_user() { - let db = test_db(); - assert!(db.find_user("nobody").unwrap().is_none()); - } - - /// Verify that Database satisfies each aggregate Repo trait contract - /// (AclRepo / SettingRepo / IdentityRepo). Exercises the trait-object - /// path so callers that take `Arc` compile end-to-end. - #[test] - fn test_aggregate_repo_trait_objects() { - let db = test_db(); - - let setting: &dyn SettingRepo = &db; - setting.set_setting("test_key", "test_value").unwrap(); - assert_eq!(setting.get_setting("test_key").unwrap(), Some("test_value".to_string())); - - let acl: &dyn AclRepo = &db; - acl.insert_acl_rule(4, "source", "blacklist", "10.0.0.1", 443).unwrap(); - // load_acl_rules is an inherent Database method (not on AclRepo), - // so go through `&db` directly for this read-back assertion. - let rules = db.load_acl_rules().unwrap(); - assert_eq!(rules.len(), 1); - - let identity: &dyn IdentityRepo = &db; - // user_count is inherent — inserts still go through the trait so - // the vtable has something to exercise. - assert_eq!(db.user_count().unwrap(), 0); - identity.insert_user("test", "hash", "viewer", false).unwrap(); - assert_eq!(db.user_count().unwrap(), 1); - } - - /// Happy path. Verifies `commit_soar_block_to_db` writes both - /// `soar_block_rules` and `acl_rules` atomically. - #[test] - fn test_commit_soar_block_happy_path() { - let db = test_db(); - - // Seed a playbook so the foreign-key-ish playbook_id refers to something real. - let pb_id = db - .insert_playbook("test_pb", "threat_detected", Some(0.9), None, None, 300) - .unwrap(); - - let soar_block_id = db - .commit_soar_block_to_db("10.0.0.99", 4, pb_id, "2099-01-01 00:00:00") - .unwrap(); - assert!(soar_block_id > 0); - - // soar_block_rules has the row - let active = db.get_active_soar_blocks().unwrap(); - assert_eq!(active.len(), 1); - assert_eq!(active[0].1, "10.0.0.99"); - - // acl_rules has the matching row - let rules = db.load_acl_rules().unwrap(); - assert_eq!(rules.len(), 1); - assert_eq!(rules[0].3, "10.0.0.99"); - } - - /// Verifies `commit_soar_unblock_to_db` removes the ACL row and marks - /// the SOAR row as unblocked in one transaction. - #[test] - fn test_commit_soar_unblock_clears_both_tables() { - let db = test_db(); - let pb_id = db - .insert_playbook("test_pb", "threat_detected", Some(0.9), None, None, 300) - .unwrap(); - let soar_block_id = db - .commit_soar_block_to_db("10.0.0.99", 4, pb_id, "2099-01-01 00:00:00") - .unwrap(); - - db.commit_soar_unblock_to_db(soar_block_id, 4, "10.0.0.99").unwrap(); - - // acl_rules row gone - assert!(db.load_acl_rules().unwrap().is_empty()); - // soar_block_rules row no longer in "active" view (unblocked_at is set) - assert!(db.get_active_soar_blocks().unwrap().is_empty()); - } - - /// Verifies `insert_playbook_atomic` writes playbook + actions + conditions - /// atomically. - #[test] - fn test_insert_playbook_atomic_writes_all_three_tables() { - let db = test_db(); - let actions = vec![(1i64, "block_ip".to_string(), "{}".to_string())]; - let conditions = vec![("threshold".to_string(), ">=".to_string(), "0.8".to_string(), None)]; - let id = db - .insert_playbook_atomic("atom_pb", "threat", Some(0.8), None, None, 300, &actions, &conditions) - .unwrap(); - assert!(id > 0); - let loaded = db.load_playbooks_with_actions().unwrap(); - assert!(!loaded.is_empty()); - let cond_rows = db.load_all_playbook_conditions().unwrap(); - assert_eq!(cond_rows.len(), 1); - } -} diff --git a/net-guardia/src/adapter/persistence/setting.rs b/net-guardia/src/adapter/persistence/setting.rs new file mode 100644 index 0000000..61b654b --- /dev/null +++ b/net-guardia/src/adapter/persistence/setting.rs @@ -0,0 +1,115 @@ +use rusqlite::{Error as RusqliteError, params}; + +use super::Database; +use crate::interface::port::setting::SettingRepo; +use crate::model::error::Error; + +impl Database { + pub fn get_setting(&self, key: &str) -> Result, Error> { + let conn = self.conn()?; + let result = conn.query_row("SELECT value FROM settings WHERE key = ?1", params![key], |row| { + row.get(0) + }); + match result { + Ok(val) => Ok(Some(val)), + Err(RusqliteError::QueryReturnedNoRows) => Ok(None), + Err(e) => Err(e)?, + } + } + + pub fn set_setting(&self, key: &str, value: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "INSERT OR REPLACE INTO settings (key, value) VALUES (?1, ?2)", + params![key, value], + )?; + Ok(()) + } + + pub fn get_app_secret(&self, key: &str) -> Result, Error> { + let conn = self.conn()?; + let result = conn.query_row("SELECT value FROM app_secrets WHERE key = ?1", params![key], |row| { + row.get(0) + }); + match result { + Ok(val) => Ok(Some(val)), + Err(RusqliteError::QueryReturnedNoRows) => Ok(None), + Err(e) => Err(e)?, + } + } + + pub fn set_app_secret(&self, key: &str, value: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "INSERT OR REPLACE INTO app_secrets (key, value) VALUES (?1, ?2)", + params![key, value], + )?; + Ok(()) + } + + pub fn get_notification_config(&self, channel: &str) -> Result, Error> { + let conn = self.conn()?; + match conn.query_row( + "SELECT config_json FROM notification_config WHERE channel = ?1 AND enabled = 1", + params![channel], + |row| row.get::<_, String>(0), + ) { + Ok(json) => Ok(Some(json)), + Err(RusqliteError::QueryReturnedNoRows) => Ok(None), + Err(e) => Err(e)?, + } + } + + pub fn set_notification_config(&self, channel: &str, config_json: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "INSERT INTO notification_config (channel, config_json) VALUES (?1, ?2) \ + ON CONFLICT(channel) DO UPDATE SET config_json = ?2, updated_at = datetime('now')", + params![channel, config_json], + )?; + Ok(()) + } +} + +impl SettingRepo for Database { + fn get_setting(&self, key: &str) -> Result, Error> { + self.get_setting(key) + } + + fn set_setting(&self, key: &str, value: &str) -> Result<(), Error> { + self.set_setting(key, value) + } + + fn get_app_secret(&self, key: &str) -> Result, Error> { + self.get_app_secret(key) + } + + fn set_app_secret(&self, key: &str, plaintext: &str) -> Result<(), Error> { + self.set_app_secret(key, plaintext) + } + + fn get_notification_config(&self, channel: &str) -> Result, Error> { + self.get_notification_config(channel) + } + + fn set_notification_config(&self, channel: &str, config_json: &str) -> Result<(), Error> { + self.set_notification_config(channel, config_json) + } +} + +#[cfg(test)] +mod tests { + use super::super::tests::test_db; + + #[test] + fn test_settings_crud() { + let db = test_db(); + assert_eq!(db.get_setting("foo").unwrap(), None); + + db.set_setting("foo", "bar").unwrap(); + assert_eq!(db.get_setting("foo").unwrap(), Some("bar".to_string())); + + db.set_setting("foo", "baz").unwrap(); + assert_eq!(db.get_setting("foo").unwrap(), Some("baz".to_string())); + } +} diff --git a/net-guardia/src/adapter/persistence/soar.rs b/net-guardia/src/adapter/persistence/soar.rs new file mode 100644 index 0000000..001505a --- /dev/null +++ b/net-guardia/src/adapter/persistence/soar.rs @@ -0,0 +1,413 @@ +use rusqlite::params; + +use super::Database; +use crate::interface::port::soar::{PlaybookRow, SoarExecutionRow, SoarRepo}; +use crate::model::error::Error; +use crate::model::soar::playbook_data::UpdatePlaybookRow; + +impl Database { + pub fn insert_playbook( + &self, + name: &str, + trigger_event: &str, + threshold: Option, + count: Option, + window: Option, + cooldown: i64, + ) -> Result { + let conn = self.conn()?; + conn.execute( + "INSERT INTO playbooks (name, trigger_event, condition_threshold, condition_count, condition_window_secs, cooldown_secs) VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![name, trigger_event, threshold, count, window, cooldown], + )?; + Ok(conn.last_insert_rowid()) + } + + pub fn insert_playbook_action( + &self, + playbook_id: i64, + action_order: i64, + action_type: &str, + params_json: &str, + ) -> Result { + let conn = self.conn()?; + conn.execute( + "INSERT INTO playbook_actions (playbook_id, action_order, action_type, params) VALUES (?1, ?2, ?3, ?4)", + params![playbook_id, action_order, action_type, params_json], + )?; + Ok(conn.last_insert_rowid()) + } + + /// Load all playbooks with their actions in a single JOIN query (avoids N+1). + /// Returns Vec of (playbook fields..., action fields...). + #[allow(clippy::type_complexity)] + pub fn load_playbooks_with_actions( + &self, + ) -> Result< + Vec<( + i64, + String, + bool, + String, + Option, + Option, + Option, + i64, + Option, + Option, + Option, + Option, + )>, + Error, + > { + let conn = self.conn()?; + let mut stmt = conn.prepare( + "SELECT p.id, p.name, p.enabled, p.trigger_event, p.condition_threshold, \ + p.condition_count, p.condition_window_secs, p.cooldown_secs, \ + a.id, a.action_order, a.action_type, a.params \ + FROM playbooks p \ + LEFT JOIN playbook_actions a ON a.playbook_id = p.id \ + ORDER BY p.id, a.action_order", + )?; + let rows = stmt.query_map([], |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, String>(1)?, + row.get::<_, bool>(2)?, + row.get::<_, String>(3)?, + row.get::<_, Option>(4)?, + row.get::<_, Option>(5)?, + row.get::<_, Option>(6)?, + row.get::<_, i64>(7)?, + row.get::<_, Option>(8)?, + row.get::<_, Option>(9)?, + row.get::<_, Option>(10)?, + row.get::<_, Option>(11)?, + )) + })?; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + + pub fn update_playbook_enabled(&self, id: i64, enabled: bool) -> Result { + let conn = self.conn()?; + let rows = conn.execute( + "UPDATE playbooks SET enabled = ?2, updated_at = datetime('now') WHERE id = ?1", + params![id, enabled as i32], + )?; + Ok(rows > 0) + } + + pub fn delete_playbook(&self, id: i64) -> Result { + let conn = self.conn()?; + let rows = conn.execute("DELETE FROM playbooks WHERE id = ?1", params![id])?; + Ok(rows > 0) + } + + pub fn insert_playbook_condition( + &self, + playbook_id: i64, + condition_type: &str, + operator: &str, + value: &str, + value2: Option<&str>, + ) -> Result { + let conn = self.conn()?; + conn.execute( + "INSERT INTO playbook_conditions (playbook_id, condition_type, operator, value, value2) \ + VALUES (?1, ?2, ?3, ?4, ?5)", + params![playbook_id, condition_type, operator, value, value2], + )?; + Ok(conn.last_insert_rowid()) + } + + #[allow(clippy::type_complexity)] + pub fn load_all_playbook_conditions( + &self, + ) -> Result)>, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare( + "SELECT id, playbook_id, condition_type, operator, value, value2 \ + FROM playbook_conditions ORDER BY playbook_id, id", + )?; + let rows = stmt.query_map([], |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, i64>(1)?, + row.get::<_, String>(2)?, + row.get::<_, String>(3)?, + row.get::<_, String>(4)?, + row.get::<_, Option>(5)?, + )) + })?; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + + pub fn insert_soar_execution( + &self, + playbook_id: i64, + source_ip: Option<&str>, + trigger_event: &str, + actions_json: &str, + ) -> Result { + let conn = self.conn()?; + conn.execute( + "INSERT INTO soar_executions (playbook_id, source_ip, trigger_event, actions_executed) VALUES (?1, ?2, ?3, ?4)", + params![playbook_id, source_ip, trigger_event, actions_json], + )?; + Ok(conn.last_insert_rowid()) + } + + #[allow(clippy::type_complexity)] + pub fn list_soar_executions( + &self, + limit: i64, + ) -> Result, String, String, String)>, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare( + "SELECT id, playbook_id, source_ip, trigger_event, actions_executed, executed_at FROM soar_executions ORDER BY executed_at DESC LIMIT ?1" + )?; + let rows = stmt.query_map(params![limit], |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, i64>(1)?, + row.get::<_, Option>(2)?, + row.get::<_, String>(3)?, + row.get::<_, String>(4)?, + row.get::<_, String>(5)?, + )) + })?; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + + pub fn seed_default_playbooks(&self) -> Result<(), Error> { + let conn = self.conn()?; + let count: i64 = conn.query_row("SELECT COUNT(*) FROM playbooks", [], |row| row.get(0))?; + if count > 0 { + return Ok(()); + } + drop(conn); + + // 1. brute_force_block: brute_force, count 5 in 60s → block_ip(3600s) + send_telegram + log + let pb1 = self.insert_playbook("brute_force_block", "brute_force", None, Some(5), Some(60), 600)?; + self.insert_playbook_action(pb1, 1, "block_ip", r#"{"ttl_secs": 3600}"#)?; + self.insert_playbook_action(pb1, 2, "send_telegram", "{}")?; + self.insert_playbook_action(pb1, 3, "log", r#"{"level": "warn"}"#)?; + self.insert_playbook_condition(pb1, "frequency", ">=", "5", Some("60"))?; + + // 2. port_scan_alert: port_scan, threshold 0.7 → send_telegram + log (no block) + let pb2 = self.insert_playbook("port_scan_alert", "port_scan", Some(0.7), None, None, 300)?; + self.insert_playbook_action(pb2, 1, "send_telegram", "{}")?; + self.insert_playbook_action(pb2, 2, "log", r#"{"level": "warn"}"#)?; + self.insert_playbook_condition(pb2, "threshold", ">=", "0.7", None)?; + + // 3. fusion_c2_multi_source_block — C2 beacon observed by ≥2 sources + // (e.g. Suricata trojan-activity + Beaconing CV + ML c2 class) is + // the highest-precision fusion signal we ship. Block for 1h and + // notify, no solo-source threshold so single-source C2 hits still + // require the solo playbook below to act. + let pb3 = self.insert_playbook("fusion_c2_multi_source_block", "c2_beacon", None, None, None, 600)?; + self.insert_playbook_action(pb3, 1, "block_ip", r#"{"ttl_secs": 3600}"#)?; + self.insert_playbook_action(pb3, 2, "send_telegram", "{}")?; + self.insert_playbook_action(pb3, 3, "log", r#"{"level": "warn"}"#)?; + self.insert_playbook_condition(pb3, "multi_source_min", ">=", "2", None)?; + + // 4. fusion_c2_suricata_solo_high_block — the escape hatch for + // Suricata signature hits with very high confidence (>=0.95). + // Lets known-good rules fire without waiting for agreement from a + // second source, matching how analysts intuitively treat a + // signature "dead-on" match. + let pb4 = self.insert_playbook("fusion_c2_suricata_solo_high_block", "c2_beacon", None, None, None, 600)?; + self.insert_playbook_action(pb4, 1, "block_ip", r#"{"ttl_secs": 3600}"#)?; + self.insert_playbook_action(pb4, 2, "send_telegram", "{}")?; + self.insert_playbook_action(pb4, 3, "log", r#"{"level": "warn"}"#)?; + self.insert_playbook_condition(pb4, "single_source_high", "==", "Suricata", Some("0.95"))?; + + Ok(()) + } +} + +impl SoarRepo for Database { + fn load_playbooks_with_actions(&self) -> Result, Error> { + self.load_playbooks_with_actions() + } + + fn update_playbook_enabled(&self, id: i64, enabled: bool) -> Result { + self.update_playbook_enabled(id, enabled) + } + + fn delete_playbook(&self, id: i64) -> Result { + self.delete_playbook(id) + } + + fn seed_default_playbooks(&self) -> Result<(), Error> { + self.seed_default_playbooks() + } + + fn load_all_playbook_conditions(&self) -> Result)>, Error> { + self.load_all_playbook_conditions() + } + + fn count_active_soar_blocks(&self) -> Result { + self.count_active_soar_blocks() + } + + fn get_active_soar_blocks(&self) -> Result, Error> { + self.get_active_soar_blocks() + } + + fn get_soar_block_by_id(&self, id: i64) -> Result, Error> { + self.get_soar_block_by_id(id) + } + + fn get_expired_soar_blocks(&self) -> Result, Error> { + self.get_expired_soar_blocks() + } + + fn mark_soar_block_unblocked(&self, id: i64) -> Result<(), Error> { + self.mark_soar_block_unblocked(id) + } + + fn insert_pending_unblock(&self, source_ip: &str) -> Result { + self.insert_pending_unblock(source_ip) + } + + fn load_pending_unblocks(&self) -> Result, Error> { + self.load_pending_unblocks() + } + + fn delete_pending_unblock(&self, id: i64) -> Result<(), Error> { + self.delete_pending_unblock(id) + } + + fn increment_pending_unblock_retry(&self, id: i64) -> Result<(), Error> { + self.increment_pending_unblock_retry(id) + } + + fn insert_soar_execution( + &self, + playbook_id: i64, + source_ip: Option<&str>, + trigger_event: &str, + actions_json: &str, + ) -> Result { + self.insert_soar_execution(playbook_id, source_ip, trigger_event, actions_json) + } + + fn list_soar_executions(&self, limit: i64) -> Result, Error> { + self.list_soar_executions(limit) + } + + fn insert_playbook_atomic( + &self, + name: &str, + trigger_event: &str, + threshold: Option, + count: Option, + window: Option, + cooldown: i64, + actions: &[(i64, String, String)], + conditions: &[(String, String, String, Option)], + ) -> Result { + let mut conn = self.conn()?; + let tx = conn.transaction()?; + tx.execute( + "INSERT INTO playbooks (name, trigger_event, condition_threshold, condition_count, condition_window_secs, cooldown_secs) VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![name, trigger_event, threshold, count, window, cooldown], + )?; + let playbook_id = tx.last_insert_rowid(); + for (action_order, action_type, params_json) in actions { + tx.execute( + "INSERT INTO playbook_actions (playbook_id, action_order, action_type, params) VALUES (?1, ?2, ?3, ?4)", + params![playbook_id, action_order, action_type, params_json], + )?; + } + for (condition_type, operator, value, value2) in conditions { + tx.execute( + "INSERT INTO playbook_conditions (playbook_id, condition_type, operator, value, value2) VALUES (?1, ?2, ?3, ?4, ?5)", + params![playbook_id, condition_type, operator, value, value2.as_deref()], + )?; + } + tx.commit()?; + Ok(playbook_id) + } + + fn update_playbook_atomic( + &self, + id: i64, + row: &UpdatePlaybookRow, + actions: &[(i64, String, String)], + conditions: &[(String, String, String, Option)], + ) -> Result { + let mut conn = self.conn()?; + let tx = conn.transaction()?; + let rows_updated = tx.execute( + "UPDATE playbooks SET name = ?2, trigger_event = ?3, condition_threshold = ?4, \ + condition_count = ?5, condition_window_secs = ?6, cooldown_secs = ?7, \ + updated_at = datetime('now') WHERE id = ?1", + params![ + id, + row.name, + row.trigger_event, + row.condition_threshold, + row.condition_count, + row.condition_window_secs, + row.cooldown_secs + ], + )?; + if rows_updated == 0 { + return Ok(false); + } + tx.execute("DELETE FROM playbook_actions WHERE playbook_id = ?1", params![id])?; + tx.execute("DELETE FROM playbook_conditions WHERE playbook_id = ?1", params![id])?; + for (action_order, action_type, params_json) in actions { + tx.execute( + "INSERT INTO playbook_actions (playbook_id, action_order, action_type, params) VALUES (?1, ?2, ?3, ?4)", + params![id, action_order, action_type, params_json], + )?; + } + for (condition_type, operator, value, value2) in conditions { + tx.execute( + "INSERT INTO playbook_conditions (playbook_id, condition_type, operator, value, value2) VALUES (?1, ?2, ?3, ?4, ?5)", + params![id, condition_type, operator, value, value2.as_deref()], + )?; + } + tx.commit()?; + Ok(true) + } +} + +#[cfg(test)] +mod tests { + use super::super::tests::test_db; + + /// Verifies `insert_playbook_atomic` writes playbook + actions + conditions + /// atomically. + #[test] + fn test_insert_playbook_atomic_writes_all_three_tables() { + use crate::interface::port::soar::SoarRepo; + + let db = test_db(); + let actions = vec![(1i64, "block_ip".to_string(), "{}".to_string())]; + let conditions = vec![("threshold".to_string(), ">=".to_string(), "0.8".to_string(), None)]; + let id = db + .insert_playbook_atomic("atom_pb", "threat", Some(0.8), None, None, 300, &actions, &conditions) + .unwrap(); + assert!(id > 0); + let loaded = db.load_playbooks_with_actions().unwrap(); + assert!(!loaded.is_empty()); + let cond_rows = db.load_all_playbook_conditions().unwrap(); + assert_eq!(cond_rows.len(), 1); + } +} diff --git a/net-guardia/src/adapter/persistence/soar_block.rs b/net-guardia/src/adapter/persistence/soar_block.rs new file mode 100644 index 0000000..b68652b --- /dev/null +++ b/net-guardia/src/adapter/persistence/soar_block.rs @@ -0,0 +1,229 @@ +use rusqlite::params; + +use super::Database; +use crate::interface::port::db_admin::DbAdminRepo; +use crate::model::error::Error; + +impl Database { + /// Test-only helper: direct insert of a SOAR block rule row. Production + /// code goes through `commit_soar_block_to_db` which atomically writes + /// both `soar_block_rules` and `acl_rules` under a transaction. + #[cfg(test)] + pub fn insert_soar_block_rule(&self, source_ip: &str, playbook_id: i64, expires_at: &str) -> Result { + let conn = self.conn()?; + conn.execute( + "INSERT INTO soar_block_rules (source_ip, playbook_id, expires_at) VALUES (?1, ?2, ?3)", + params![source_ip, playbook_id, expires_at], + )?; + Ok(conn.last_insert_rowid()) + } + + pub fn count_active_soar_blocks(&self) -> Result { + let conn = self.conn()?; + let count: u32 = conn.query_row( + "SELECT COUNT(*) FROM soar_block_rules WHERE unblocked_at IS NULL AND expires_at > datetime('now')", + [], + |row| row.get(0), + )?; + Ok(count) + } + + pub fn get_expired_soar_blocks(&self) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare( + "SELECT id, source_ip, playbook_id FROM soar_block_rules WHERE expires_at <= datetime('now') AND unblocked_at IS NULL" + )?; + let rows = stmt.query_map([], |row| { + Ok((row.get::<_, i64>(0)?, row.get::<_, String>(1)?, row.get::<_, i64>(2)?)) + })?; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + + /// Get a single SOAR block rule by ID, returning (id, source_ip, playbook_id, expires_at). + pub fn get_soar_block_by_id(&self, id: i64) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = + conn.prepare("SELECT id, source_ip, playbook_id, expires_at FROM soar_block_rules WHERE id = ?1")?; + let mut rows = stmt.query_map(params![id], |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, String>(1)?, + row.get::<_, i64>(2)?, + row.get::<_, String>(3)?, + )) + })?; + match rows.next() { + Some(row) => Ok(Some(row?)), + None => Ok(None), + } + } + + pub fn mark_soar_block_unblocked(&self, id: i64) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "UPDATE soar_block_rules SET unblocked_at = datetime('now') WHERE id = ?1", + params![id], + )?; + Ok(()) + } + + pub fn get_active_soar_blocks(&self) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare( + "SELECT id, source_ip, playbook_id, expires_at FROM soar_block_rules WHERE unblocked_at IS NULL AND expires_at > datetime('now')" + )?; + let rows = stmt.query_map([], |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, String>(1)?, + row.get::<_, i64>(2)?, + row.get::<_, String>(3)?, + )) + })?; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + + pub fn insert_pending_unblock(&self, source_ip: &str) -> Result { + let conn = self.conn()?; + conn.execute( + "INSERT INTO pending_unblock (source_ip) VALUES (?1)", + params![source_ip], + )?; + Ok(conn.last_insert_rowid()) + } + + pub fn load_pending_unblocks(&self) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare("SELECT id, source_ip, retry_count FROM pending_unblock ORDER BY id")?; + let rows = stmt.query_map([], |row| { + Ok((row.get::<_, i64>(0)?, row.get::<_, String>(1)?, row.get::<_, i64>(2)?)) + })?; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + + pub fn delete_pending_unblock(&self, id: i64) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute("DELETE FROM pending_unblock WHERE id = ?1", params![id])?; + Ok(()) + } + + pub fn increment_pending_unblock_retry(&self, id: i64) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "UPDATE pending_unblock SET retry_count = retry_count + 1 WHERE id = ?1", + params![id], + )?; + Ok(()) + } +} + +impl DbAdminRepo for Database { + /// Commit a SOAR-driven block to both `soar_block_rules` and + /// `acl_rules` in one transaction. Callers must have already installed + /// the eBPF block before calling this, and are responsible for removing + /// the eBPF block if this returns Err. + fn commit_soar_block_to_db( + &self, + source_ip: &str, + ip_version: u8, + playbook_id: i64, + expires_at: &str, + ) -> Result { + let mut conn = self.conn()?; + let tx = conn.transaction()?; + tx.execute( + "INSERT INTO soar_block_rules (source_ip, playbook_id, expires_at) VALUES (?1, ?2, ?3)", + params![source_ip, playbook_id, expires_at], + )?; + let soar_block_id = tx.last_insert_rowid(); + tx.execute( + "INSERT OR IGNORE INTO acl_rules (ip_version, direction, list_type, ip_address, port) VALUES (?1, ?2, ?3, ?4, ?5)", + params![ip_version, "source", "blacklist", source_ip, 0i64], + )?; + tx.commit()?; + Ok(soar_block_id) + } + + /// Clear a SOAR-driven block: remove the `acl_rules` entry and mark + /// the `soar_block_rules` row as unblocked in one transaction. + /// Callers handle eBPF unblock separately. + fn commit_soar_unblock_to_db(&self, soar_block_id: i64, ip_version: u8, source_ip: &str) -> Result<(), Error> { + let mut conn = self.conn()?; + let tx = conn.transaction()?; + tx.execute( + "DELETE FROM acl_rules WHERE ip_version = ?1 AND direction = ?2 AND list_type = ?3 AND ip_address = ?4 AND port = ?5", + params![ip_version, "source", "blacklist", source_ip, 0i64], + )?; + tx.execute( + "UPDATE soar_block_rules SET unblocked_at = datetime('now') WHERE id = ?1", + params![soar_block_id], + )?; + tx.commit()?; + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::super::tests::test_db; + use crate::interface::port::db_admin::DbAdminRepo; + + /// Happy path. Verifies `commit_soar_block_to_db` writes both + /// `soar_block_rules` and `acl_rules` atomically. + #[test] + fn test_commit_soar_block_happy_path() { + let db = test_db(); + + // Seed a playbook so the foreign-key-ish playbook_id refers to something real. + let pb_id = db + .insert_playbook("test_pb", "threat_detected", Some(0.9), None, None, 300) + .unwrap(); + + let soar_block_id = db + .commit_soar_block_to_db("10.0.0.99", 4, pb_id, "2099-01-01 00:00:00") + .unwrap(); + assert!(soar_block_id > 0); + + // soar_block_rules has the row + let active = db.get_active_soar_blocks().unwrap(); + assert_eq!(active.len(), 1); + assert_eq!(active[0].1, "10.0.0.99"); + + // acl_rules has the matching row + let rules = db.load_acl_rules().unwrap(); + assert_eq!(rules.len(), 1); + assert_eq!(rules[0].3, "10.0.0.99"); + } + + /// Verifies `commit_soar_unblock_to_db` removes the ACL row and marks + /// the SOAR row as unblocked in one transaction. + #[test] + fn test_commit_soar_unblock_clears_both_tables() { + let db = test_db(); + let pb_id = db + .insert_playbook("test_pb", "threat_detected", Some(0.9), None, None, 300) + .unwrap(); + let soar_block_id = db + .commit_soar_block_to_db("10.0.0.99", 4, pb_id, "2099-01-01 00:00:00") + .unwrap(); + + db.commit_soar_unblock_to_db(soar_block_id, 4, "10.0.0.99").unwrap(); + + // acl_rules row gone + assert!(db.load_acl_rules().unwrap().is_empty()); + // soar_block_rules row no longer in "active" view (unblocked_at is set) + assert!(db.get_active_soar_blocks().unwrap().is_empty()); + } +} diff --git a/net-guardia/src/adapter/persistence/stats.rs b/net-guardia/src/adapter/persistence/stats.rs new file mode 100644 index 0000000..d647354 --- /dev/null +++ b/net-guardia/src/adapter/persistence/stats.rs @@ -0,0 +1,105 @@ +use rusqlite::params; + +use super::Database; +use crate::interface::port::stats::StatsRepo; +use crate::model::error::Error; + +impl Database { + /// Count SOAR executions in the last N days. + pub fn count_weekly_executions(&self, days: i64) -> Result { + let conn = self.conn()?; + let count: i64 = conn.query_row( + "SELECT COUNT(*) FROM soar_executions WHERE executed_at >= datetime('now', ?1)", + params![format!("-{} days", days)], + |row| row.get(0), + )?; + Ok(count as u64) + } + + /// Count SOAR blocks created in the last N days. + pub fn count_weekly_blocks(&self, days: i64) -> Result { + let conn = self.conn()?; + let count: i64 = conn.query_row( + "SELECT COUNT(*) FROM soar_block_rules WHERE created_at >= datetime('now', ?1)", + params![format!("-{} days", days)], + |row| row.get(0), + )?; + Ok(count as u64) + } + + /// Count SOAR unblocks in the last N days. + pub fn count_weekly_unblocks(&self, days: i64) -> Result { + let conn = self.conn()?; + let count: i64 = conn.query_row( + "SELECT COUNT(*) FROM soar_block_rules WHERE unblocked_at IS NOT NULL AND unblocked_at >= datetime('now', ?1)", + params![format!("-{} days", days)], + |row| row.get(0), + )?; + Ok(count as u64) + } + + /// Get threat breakdown by trigger_event in the last N days. + pub fn weekly_threat_breakdown(&self, days: i64) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare( + "SELECT trigger_event, COUNT(*) FROM soar_executions WHERE executed_at >= datetime('now', ?1) GROUP BY trigger_event ORDER BY COUNT(*) DESC" + )?; + let rows = stmt.query_map(params![format!("-{} days", days)], |row| { + Ok((row.get::<_, String>(0)?, row.get::<_, i64>(1)? as u64)) + })?; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + + /// Get top blocked IPs in the last N days. + pub fn weekly_top_ips(&self, days: i64, limit: i64) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare( + "SELECT source_ip, COUNT(*) as cnt FROM soar_block_rules WHERE created_at >= datetime('now', ?1) GROUP BY source_ip ORDER BY cnt DESC LIMIT ?2" + )?; + let rows = stmt.query_map(params![format!("-{} days", days), limit], |row| { + Ok((row.get::<_, String>(0)?, row.get::<_, i64>(1)? as u64)) + })?; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + + /// Count current ACL rules. + pub fn count_acl_rules(&self) -> Result { + let conn = self.conn()?; + let count: i64 = conn.query_row("SELECT COUNT(*) FROM acl_rules", [], |row| row.get(0))?; + Ok(count as u64) + } +} + +impl StatsRepo for Database { + fn count_weekly_executions(&self, days: i64) -> Result { + self.count_weekly_executions(days) + } + + fn count_weekly_blocks(&self, days: i64) -> Result { + self.count_weekly_blocks(days) + } + + fn count_weekly_unblocks(&self, days: i64) -> Result { + self.count_weekly_unblocks(days) + } + + fn weekly_threat_breakdown(&self, days: i64) -> Result, Error> { + self.weekly_threat_breakdown(days) + } + + fn weekly_top_ips(&self, days: i64, limit: i64) -> Result, Error> { + self.weekly_top_ips(days, limit) + } + + fn count_acl_rules(&self) -> Result { + self.count_acl_rules() + } +} diff --git a/net-guardia/src/adapter/persistence/user.rs b/net-guardia/src/adapter/persistence/user.rs new file mode 100644 index 0000000..f4c09b2 --- /dev/null +++ b/net-guardia/src/adapter/persistence/user.rs @@ -0,0 +1,530 @@ +use std::collections::{HashMap, HashSet}; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use rusqlite::{Error as RusqliteError, params}; + +use super::Database; +use crate::interface::port::identity::{IdentityRepo, UserGroupTuple, UserTuple, UserWithGroups}; +use crate::model::error::Error; +use crate::model::error::database::DatabaseError; + +impl Database { + pub fn find_user(&self, username: &str) -> Result, Error> { + let conn = self.conn()?; + let result = conn.query_row( + "SELECT id, username, password_hash, role, force_password_change FROM users WHERE username = ?1", + params![username], + |row| { + Ok(( + row.get(0)?, + row.get(1)?, + row.get(2)?, + row.get(3)?, + row.get::<_, i64>(4)? != 0, + )) + }, + ); + match result { + Ok(user) => Ok(Some(user)), + Err(RusqliteError::QueryReturnedNoRows) => Ok(None), + Err(e) => Err(e)?, + } + } + + pub fn insert_user( + &self, + username: &str, + password_hash: &str, + role: &str, + force_password_change: bool, + ) -> Result { + let conn = self.conn()?; + conn.execute( + "INSERT INTO users (username, password_hash, role, force_password_change) VALUES (?1, ?2, ?3, ?4)", + params![username, password_hash, role, force_password_change as i64], + ) + .map_err(|e| -> Error { + if e.to_string().contains("UNIQUE constraint") { + DatabaseError::UserAlreadyExists(username.to_string()).into() + } else { + e.into() + } + })?; + Ok(conn.last_insert_rowid()) + } + + pub fn update_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "UPDATE users SET password_hash = ?1, force_password_change = 0 WHERE id = ?2", + params![password_hash, user_id], + )?; + Ok(()) + } + + pub fn user_count(&self) -> Result { + let conn = self.conn()?; + Ok(conn.query_row("SELECT COUNT(*) FROM users", [], |row| row.get(0))?) + } + + pub fn list_users_with_groups(&self) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare( + "SELECT u.id, u.username, u.role, u.force_password_change, u.created_at, \ + g.id, g.name \ + FROM users u \ + LEFT JOIN user_group_members m ON u.id = m.user_id \ + LEFT JOIN user_groups g ON g.id = m.group_id \ + ORDER BY u.id, g.id", + )?; + let rows = stmt.query_map([], |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, String>(1)?, + row.get::<_, String>(2)?, + row.get::<_, i64>(3)? != 0, + row.get::<_, String>(4)?, + row.get::<_, Option>(5)?, + row.get::<_, Option>(6)?, + )) + })?; + + let mut user_map: HashMap = HashMap::new(); + let mut order: Vec = Vec::new(); + + for row in rows { + let (id, username, role, force_pw, created_at, group_id, group_name) = row?; + let entry = user_map.entry(id).or_insert_with(|| { + order.push(id); + (id, username, role, force_pw, created_at, Vec::new()) + }); + if let (Some(gid), Some(gname)) = (group_id, group_name) { + entry.5.push((gid, gname)); + } + } + + Ok(order.into_iter().filter_map(|id| user_map.remove(&id)).collect()) + } + + pub fn delete_user(&self, user_id: i64) -> Result { + self.cleanup_user_memberships(user_id)?; + let conn = self.conn()?; + let affected = conn.execute("DELETE FROM users WHERE id = ?1", params![user_id])?; + Ok(affected > 0) + } + + pub fn update_user_role(&self, user_id: i64, role: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute("UPDATE users SET role = ?1 WHERE id = ?2", params![role, user_id])?; + Ok(()) + } + + pub fn reset_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "UPDATE users SET password_hash = ?1, force_password_change = 1 WHERE id = ?2", + params![password_hash, user_id], + )?; + Ok(()) + } + + pub fn find_user_by_id(&self, user_id: i64) -> Result, Error> { + let conn = self.conn()?; + let result = conn.query_row( + "SELECT id, username, password_hash, role, force_password_change FROM users WHERE id = ?1", + params![user_id], + |row| { + Ok(( + row.get(0)?, + row.get(1)?, + row.get(2)?, + row.get(3)?, + row.get::<_, i64>(4)? != 0, + )) + }, + ); + match result { + Ok(user) => Ok(Some(user)), + Err(RusqliteError::QueryReturnedNoRows) => Ok(None), + Err(e) => Err(e)?, + } + } + + pub fn list_user_groups(&self) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = + conn.prepare("SELECT id, name, description, permissions, created_at FROM user_groups ORDER BY id")?; + let rows = stmt.query_map([], |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, String>(1)?, + row.get::<_, String>(2)?, + row.get::<_, String>(3)?, + row.get::<_, String>(4)?, + )) + })?; + let mut results = Vec::new(); + for row in rows { + results.push(row?); + } + Ok(results) + } + + pub fn create_user_group(&self, name: &str, description: &str, permissions: &str) -> Result { + let conn = self.conn()?; + conn.execute( + "INSERT INTO user_groups (name, description, permissions) VALUES (?1, ?2, ?3)", + params![name, description, permissions], + ) + .map_err(|e| -> Error { + if e.to_string().contains("UNIQUE constraint") { + DatabaseError::GroupAlreadyExists(name).into() + } else { + e.into() + } + })?; + Ok(conn.last_insert_rowid()) + } + + pub fn update_user_group(&self, id: i64, name: &str, description: &str, permissions: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "UPDATE user_groups SET name = ?1, description = ?2, permissions = ?3 WHERE id = ?4", + params![name, description, permissions, id], + )?; + Ok(()) + } + + pub fn delete_user_group(&self, id: i64) -> Result { + let conn = self.conn()?; + conn.execute("DELETE FROM user_group_members WHERE group_id = ?1", params![id])?; + let affected = conn.execute("DELETE FROM user_groups WHERE id = ?1", params![id])?; + Ok(affected > 0) + } + + pub fn get_user_group(&self, id: i64) -> Result, Error> { + let conn = self.conn()?; + let result = conn.query_row( + "SELECT id, name, description, permissions, created_at FROM user_groups WHERE id = ?1", + params![id], + |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, String>(1)?, + row.get::<_, String>(2)?, + row.get::<_, String>(3)?, + row.get::<_, String>(4)?, + )) + }, + ); + match result { + Ok(group) => Ok(Some(group)), + Err(RusqliteError::QueryReturnedNoRows) => Ok(None), + Err(e) => Err(e)?, + } + } + + pub fn get_user_groups(&self, user_id: i64) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare( + "SELECT g.id, g.name, g.description, g.permissions FROM user_groups g \ + INNER JOIN user_group_members m ON g.id = m.group_id \ + WHERE m.user_id = ?1 ORDER BY g.id", + )?; + let rows = stmt.query_map(params![user_id], |row| { + Ok(( + row.get::<_, i64>(0)?, + row.get::<_, String>(1)?, + row.get::<_, String>(2)?, + row.get::<_, String>(3)?, + )) + })?; + let mut results = Vec::new(); + for row in rows { + results.push(row?); + } + Ok(results) + } + + pub fn set_user_groups(&self, user_id: i64, group_ids: &[i64]) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute("DELETE FROM user_group_members WHERE user_id = ?1", params![user_id])?; + for &gid in group_ids { + conn.execute( + "INSERT INTO user_group_members (user_id, group_id) VALUES (?1, ?2)", + params![user_id, gid], + )?; + } + Ok(()) + } + + pub fn get_user_permissions(&self, user_id: i64) -> Result, Error> { + let groups = self.get_user_groups(user_id)?; + let mut all_perms = HashSet::new(); + for (_id, _name, _desc, perms_json) in groups { + if let Ok(perms) = serde_json::from_str::>(&perms_json) { + for p in perms { + all_perms.insert(p); + } + } + } + let mut result: Vec = all_perms.into_iter().collect(); + result.sort(); + Ok(result) + } + + pub fn cleanup_user_memberships(&self, user_id: i64) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute("DELETE FROM user_group_members WHERE user_id = ?1", params![user_id])?; + Ok(()) + } + + pub fn get_group_member_ids(&self, group_id: i64) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare("SELECT user_id FROM user_group_members WHERE group_id = ?1")?; + let rows = stmt.query_map(params![group_id], |row| row.get::<_, i64>(0))?; + let mut results = Vec::new(); + for row in rows { + results.push(row?); + } + Ok(results) + } + + pub fn get_group_members(&self, group_id: i64) -> Result, Error> { + let conn = self.conn()?; + let mut stmt = conn.prepare( + "SELECT u.id, u.username FROM users u \ + INNER JOIN user_group_members m ON u.id = m.user_id \ + WHERE m.group_id = ?1 ORDER BY u.username", + )?; + let rows = stmt.query_map(params![group_id], |row| { + Ok((row.get::<_, i64>(0)?, row.get::<_, String>(1)?)) + })?; + let mut results = Vec::new(); + for row in rows { + results.push(row?); + } + Ok(results) + } + + pub fn record_login_failure(&self, username: &str) -> Result<(u32, Option), Error> { + let key_count = format!("login_failures:{}", username); + let key_locked = format!("login_locked_until:{}", username); + + let count: u32 = self.get_setting(&key_count)?.and_then(|v| v.parse().ok()).unwrap_or(0) + 1; + + self.set_setting(&key_count, &count.to_string())?; + + if count >= 5 { + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or(Duration::ZERO) + .as_secs(); + let locked_until = now + 900; // 15 minutes + self.set_setting(&key_locked, &locked_until.to_string())?; + Ok((count, Some(locked_until))) + } else { + Ok((count, None)) + } + } + + pub fn check_login_locked(&self, username: &str) -> Result, Error> { + let key_locked = format!("login_locked_until:{}", username); + if let Some(locked_str) = self.get_setting(&key_locked)? + && let Ok(locked_until) = locked_str.parse::() + { + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or(Duration::ZERO) + .as_secs(); + if now < locked_until { + return Ok(Some(locked_until - now)); + } + // Lock expired, clear it + self.clear_login_failures(username)?; + } + Ok(None) + } + + pub fn clear_login_failures(&self, username: &str) -> Result<(), Error> { + let conn = self.conn()?; + conn.execute( + "DELETE FROM settings WHERE key = ?1", + params![format!("login_failures:{}", username)], + )?; + conn.execute( + "DELETE FROM settings WHERE key = ?1", + params![format!("login_locked_until:{}", username)], + )?; + Ok(()) + } +} + +impl IdentityRepo for Database { + fn find_user(&self, username: &str) -> Result, Error> { + self.find_user(username) + } + + fn find_user_by_id(&self, user_id: i64) -> Result, Error> { + self.find_user_by_id(user_id) + } + + fn insert_user( + &self, + username: &str, + password_hash: &str, + role: &str, + force_password_change: bool, + ) -> Result { + self.insert_user(username, password_hash, role, force_password_change) + } + + fn update_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> { + self.update_user_password(user_id, password_hash) + } + + fn list_users_with_groups(&self) -> Result, Error> { + self.list_users_with_groups() + } + + fn delete_user(&self, user_id: i64) -> Result { + self.delete_user(user_id) + } + + fn update_user_role(&self, user_id: i64, role: &str) -> Result<(), Error> { + self.update_user_role(user_id, role) + } + + fn reset_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> { + self.reset_user_password(user_id, password_hash) + } + + fn list_user_groups(&self) -> Result, Error> { + self.list_user_groups() + } + + fn create_user_group(&self, name: &str, description: &str, permissions: &str) -> Result { + self.create_user_group(name, description, permissions) + } + + fn update_user_group(&self, id: i64, name: &str, description: &str, permissions: &str) -> Result<(), Error> { + self.update_user_group(id, name, description, permissions) + } + + fn delete_user_group(&self, id: i64) -> Result { + self.delete_user_group(id) + } + + fn get_user_group(&self, id: i64) -> Result, Error> { + self.get_user_group(id) + } + + fn get_user_groups(&self, user_id: i64) -> Result, Error> { + self.get_user_groups(user_id) + } + + fn set_user_groups(&self, user_id: i64, group_ids: &[i64]) -> Result<(), Error> { + self.set_user_groups(user_id, group_ids) + } + + fn get_user_permissions(&self, user_id: i64) -> Result, Error> { + self.get_user_permissions(user_id) + } + + fn get_group_member_ids(&self, group_id: i64) -> Result, Error> { + self.get_group_member_ids(group_id) + } + + fn get_group_members(&self, group_id: i64) -> Result, Error> { + self.get_group_members(group_id) + } + + fn record_login_failure(&self, username: &str) -> Result<(u32, Option), Error> { + self.record_login_failure(username) + } + + fn check_login_locked(&self, username: &str) -> Result, Error> { + self.check_login_locked(username) + } + + fn clear_login_failures(&self, username: &str) -> Result<(), Error> { + self.clear_login_failures(username) + } +} + +#[cfg(test)] +mod tests { + use super::super::tests::test_db; + + #[test] + fn test_user_crud() { + let db = test_db(); + assert_eq!(db.user_count().unwrap(), 0); + + db.insert_user("admin", "hash123", "admin", true).unwrap(); + assert_eq!(db.user_count().unwrap(), 1); + + let user = db.find_user("admin").unwrap().unwrap(); + assert_eq!(user.0, 1); // id + assert_eq!(user.1, "admin"); // username + assert_eq!(user.2, "hash123"); // password_hash + assert_eq!(user.3, "admin"); // role + assert!(user.4); // force_password_change + } + + #[test] + fn test_user_duplicate() { + let db = test_db(); + db.insert_user("admin", "hash", "admin", false).unwrap(); + let result = db.insert_user("admin", "hash2", "admin", false); + assert!(result.is_err()); + } + + #[test] + fn test_update_password_clears_force_change() { + let db = test_db(); + db.insert_user("admin", "old_hash", "admin", true).unwrap(); + + let user = db.find_user("admin").unwrap().unwrap(); + assert!(user.4); // force_password_change = true + + db.update_user_password(user.0, "new_hash").unwrap(); + + let user = db.find_user("admin").unwrap().unwrap(); + assert!(!user.4); // force_password_change = false + assert_eq!(user.2, "new_hash"); + } + + #[test] + fn test_find_nonexistent_user() { + let db = test_db(); + assert!(db.find_user("nobody").unwrap().is_none()); + } + + #[test] + fn test_login_lockout() { + let db = test_db(); + + // First 4 failures don't lock + for i in 1..5 { + let (count, locked) = db.record_login_failure("admin").unwrap(); + assert_eq!(count, i); + assert!(locked.is_none()); + } + + // 5th failure triggers lock + let (count, locked) = db.record_login_failure("admin").unwrap(); + assert_eq!(count, 5); + assert!(locked.is_some()); + + // Check locked + let remaining = db.check_login_locked("admin").unwrap(); + assert!(remaining.is_some()); + assert!(remaining.unwrap() > 0); + + // Clear and verify + db.clear_login_failures("admin").unwrap(); + let remaining = db.check_login_locked("admin").unwrap(); + assert!(remaining.is_none()); + } +} diff --git a/net-guardia/src/adapter/telegram/mod.rs b/net-guardia/src/adapter/telegram.rs similarity index 100% rename from net-guardia/src/adapter/telegram/mod.rs rename to net-guardia/src/adapter/telegram.rs diff --git a/net-guardia/src/core/acl_service.rs b/net-guardia/src/core/acl_service.rs index be56df0..6394c6b 100644 --- a/net-guardia/src/core/acl_service.rs +++ b/net-guardia/src/core/acl_service.rs @@ -1,15 +1,15 @@ use std::net::{SocketAddrV4, SocketAddrV6}; use std::sync::Arc; +use macros::log; + use crate::interface::port::access_control_admin::AccessControlAdminPort; use crate::interface::port::app_repo::AppRepo; use crate::interface::port::geo_block_api::GeoBlockPort; -use crate::model::monitoring::direction::FlowDirection; -use macros::log; - use crate::model::access_control::list_type::ListType; use crate::model::error::Error; use crate::model::error::ebpf::EbpfError; +use crate::model::monitoring::direction::FlowDirection; /// Domain service that coordinates ACL changes between DB persistence and eBPF data plane. /// Atomic write: eBPF first, then DB. If DB fails, rollback eBPF. diff --git a/net-guardia/src/core/auth/middleware.rs b/net-guardia/src/core/auth/middleware.rs index 3b497be..92919dc 100644 --- a/net-guardia/src/core/auth/middleware.rs +++ b/net-guardia/src/core/auth/middleware.rs @@ -7,7 +7,6 @@ use actix_web::body::EitherBody; use actix_web::dev::{Service, ServiceRequest, ServiceResponse, Transform}; use actix_web::http::Method; use actix_web::{Error as ActixError, HttpMessage, HttpResponse, web}; - use macros::log; use crate::core::auth::jwt::JwtService; diff --git a/net-guardia/src/core/ml/model_watcher.rs b/net-guardia/src/core/ml/model_watcher.rs index ceb1f99..6e17462 100644 --- a/net-guardia/src/core/ml/model_watcher.rs +++ b/net-guardia/src/core/ml/model_watcher.rs @@ -176,10 +176,10 @@ fn path_is_inside_staging(path: &Path) -> bool { #[cfg(test)] mod tests { - use super::*; - use notify::event::{CreateKind, ModifyKind}; + use super::*; + #[test] fn staging_paths_are_filtered() { assert!(path_is_inside_staging(&PathBuf::from("models/.staging/bad.onnx"))); diff --git a/net-guardia/src/core/report/data.rs b/net-guardia/src/core/report/data.rs deleted file mode 100644 index a01d2e3..0000000 --- a/net-guardia/src/core/report/data.rs +++ /dev/null @@ -1 +0,0 @@ -// Types are available via crate::model::report::data diff --git a/net-guardia/src/core/report/mod.rs b/net-guardia/src/core/report/mod.rs index e736458..702e611 100644 --- a/net-guardia/src/core/report/mod.rs +++ b/net-guardia/src/core/report/mod.rs @@ -1,2 +1 @@ -pub mod data; pub mod engine; diff --git a/net-guardia/src/core/soar/engine.rs b/net-guardia/src/core/soar/engine.rs index ae31b61..1b97018 100644 --- a/net-guardia/src/core/soar/engine.rs +++ b/net-guardia/src/core/soar/engine.rs @@ -383,12 +383,14 @@ impl SoarEngine { #[cfg(test)] mod tests { + use std::sync::atomic::AtomicBool; + + use chrono::{Duration as ChronoDuration, Utc}; + use parking_lot::Mutex; + use super::*; use crate::model::error::ebpf::EbpfError; use crate::model::event::DetectionSource; - use chrono::{Duration as ChronoDuration, Utc}; - use parking_lot::Mutex; - use std::sync::atomic::AtomicBool; /// Mock AccessControlPort that records calls. struct MockAccessControl { diff --git a/net-guardia/src/infrastructure/cli.rs b/net-guardia/src/infrastructure/cli.rs new file mode 100644 index 0000000..7449a2c --- /dev/null +++ b/net-guardia/src/infrastructure/cli.rs @@ -0,0 +1,108 @@ +use std::env; +use std::process; + +use clap::{Parser, Subcommand}; +use macros::log; + +use crate::adapter::persistence::Database; +use crate::model::error::Error; +use crate::model::log::cli::CliLog; + +const EXIT_USAGE: i32 = 1; +const EXIT_OP_FAILED: i32 = 2; +const DEFAULT_DB_PATH: &str = "net-guardia.db"; +const DEFAULT_DECRYPTED_DEST: &str = "net-guardia-decrypted.db"; +const DEFAULT_ENCRYPTED_DEST: &str = "net-guardia-encrypted.db"; + +#[derive(Parser)] +#[command(name = "net-guardia", version, about = "network defense system")] +pub struct Cli { + #[arg(long, global = true, env = "NETGUARDIA_DB_PATH", default_value = DEFAULT_DB_PATH)] + pub db_path: String, + #[command(subcommand)] + pub command: Option, +} + +#[derive(Subcommand)] +pub enum Command { + DecryptDb { + #[arg(default_value = DEFAULT_DECRYPTED_DEST)] + dest: String, + }, + EncryptDb { + #[arg(default_value = DEFAULT_ENCRYPTED_DEST)] + dest: String, + }, + VerifyAuditLog, +} + +pub fn handle_subcommand(cli: &Cli) -> Result<(), Error> { + let db_path = cli.db_path.as_str(); + match cli.command.as_ref() { + Some(Command::DecryptDb { dest }) => run_decrypt(db_path, dest), + Some(Command::EncryptDb { dest }) => run_encrypt(db_path, dest), + Some(Command::VerifyAuditLog) => run_verify(db_path), + None => Ok(()), + } +} + +fn require_db_key(op: &str) -> String { + match env::var("NETGUARDIA_DB_KEY") { + Ok(k) if !k.is_empty() => k, + _ => { + log!(CliLog::MissingDbKey(op.to_string())); + process::exit(EXIT_USAGE); + } + } +} + +fn run_decrypt(db_path: &str, dest: &str) -> Result<(), Error> { + let key = require_db_key("decrypt"); + log!(CliLog::DecryptStarted(db_path.to_string(), dest.to_string())); + match Database::decrypt_to_file(db_path, &key, dest) { + Ok(()) => { + log!(CliLog::DecryptCompleted(dest.to_string())); + Ok(()) + } + Err(e) => { + log!(CliLog::DecryptFailed(e.to_string())); + process::exit(EXIT_OP_FAILED); + } + } +} + +fn run_encrypt(db_path: &str, dest: &str) -> Result<(), Error> { + let key = require_db_key("encrypt"); + log!(CliLog::EncryptStarted(db_path.to_string(), dest.to_string())); + match Database::encrypt_to_file(db_path, &key, dest) { + Ok(()) => { + log!(CliLog::EncryptCompleted(dest.to_string())); + Ok(()) + } + Err(e) => { + log!(CliLog::EncryptFailed(e.to_string())); + process::exit(EXIT_OP_FAILED); + } + } +} + +fn run_verify(db_path: &str) -> Result<(), Error> { + log!(CliLog::VerifyStarted(db_path.to_string())); + let db = match Database::new(db_path) { + Ok(db) => db, + Err(e) => { + log!(CliLog::DbOpenFailed(db_path.to_string(), e.to_string())); + process::exit(EXIT_OP_FAILED); + } + }; + match db.verify_audit_log_chain() { + Ok(count) => { + log!(CliLog::VerifyOk(count)); + Ok(()) + } + Err(e) => { + log!(CliLog::VerifyFailed(e.to_string())); + process::exit(EXIT_OP_FAILED); + } + } +} diff --git a/net-guardia/src/infrastructure/communication_manager.rs b/net-guardia/src/infrastructure/communication_manager.rs index 526433e..e14c005 100644 --- a/net-guardia/src/infrastructure/communication_manager.rs +++ b/net-guardia/src/infrastructure/communication_manager.rs @@ -6,6 +6,12 @@ //! surface (`Event`, `Command`, `Query`, `EventBroadcaster`, `CommandHandler`) //! lives at `interface/communication/` and remains untouched. +use std::any::{Any, TypeId}; +use std::sync::Arc; + +use dashmap::DashMap; +use tokio::sync::broadcast; + use crate::interface::communication::command::*; use crate::interface::communication::event::Event; use crate::interface::communication::event::EventBroadcaster; @@ -13,10 +19,6 @@ use crate::interface::communication::query::*; use crate::model::config::constants::DEFAULT_EVENT_CHANNEL_CAPACITY; use crate::model::error::Error; use crate::model::error::misc::MiscError; -use dashmap::DashMap; -use std::any::{Any, TypeId}; -use std::sync::Arc; -use tokio::sync::broadcast; /// Inline TypedEventBroadcaster (adapted from MirrorSphere's model). pub struct TypedEventBroadcaster { @@ -189,12 +191,13 @@ impl ServiceRegistrar { mod tests { use std::sync::Mutex; + use async_trait::async_trait; + use super::*; use crate::interface::communication::command::Command; use crate::interface::communication::event::Event; use crate::interface::communication::message::Message; use crate::interface::communication::query::Query; - use async_trait::async_trait; // ── Test Command ───────────────────────────────────────────────── diff --git a/net-guardia/src/infrastructure/enforce_mode_handler.rs b/net-guardia/src/infrastructure/enforce_mode_handler.rs index 69f26e0..a4ee9f7 100644 --- a/net-guardia/src/infrastructure/enforce_mode_handler.rs +++ b/net-guardia/src/infrastructure/enforce_mode_handler.rs @@ -1,7 +1,7 @@ -use async_trait::async_trait; use std::sync::Arc; use std::sync::atomic::{AtomicU8, Ordering}; +use async_trait::async_trait; use macros::log; use crate::infrastructure::communication_manager::CommunicationManager; diff --git a/net-guardia/src/infrastructure/mod.rs b/net-guardia/src/infrastructure/mod.rs index 99f0033..30174d3 100644 --- a/net-guardia/src/infrastructure/mod.rs +++ b/net-guardia/src/infrastructure/mod.rs @@ -1,6 +1,7 @@ pub mod app_config; pub mod app_services; pub mod audit_logger; +pub mod cli; pub mod communication_manager; pub mod ebpf_preflight; pub mod enforce_mode_handler; diff --git a/net-guardia/src/infrastructure/secret_store.rs b/net-guardia/src/infrastructure/secret_store.rs index 87ec5eb..5f7dc92 100644 --- a/net-guardia/src/infrastructure/secret_store.rs +++ b/net-guardia/src/infrastructure/secret_store.rs @@ -6,9 +6,8 @@ use aes_gcm::{AeadCore, Aes256Gcm, Nonce}; use base64::Engine; use base64::engine::general_purpose::STANDARD as B64; use hkdf::Hkdf; -use sha2::Sha256; - use macros::log; +use sha2::Sha256; use crate::adapter::persistence::Database; use crate::interface::port::secret_store::SecretStorePort; diff --git a/net-guardia/src/infrastructure/service_factory.rs b/net-guardia/src/infrastructure/service_factory.rs index 990a512..65e0b0b 100644 --- a/net-guardia/src/infrastructure/service_factory.rs +++ b/net-guardia/src/infrastructure/service_factory.rs @@ -6,20 +6,19 @@ use std::sync::atomic::AtomicU8; use std::time::Duration; use arc_swap::ArcSwap; - use aya::Ebpf; use aya::maps::{Array, MapData, ProgramArray}; use aya::programs::{Xdp, XdpFlags}; use aya_log::EbpfLogger; use common::define::pipeline::*; +use macros::log; -use crate::core::auth::jwt::JwtService; - -use crate::adapter::access_control_adapter::EbpfAccessControlAdapter; +use crate::adapter::access_control::AccessControlAdapter; use crate::adapter::ebpf::EbpfServices; use crate::adapter::persistence::Database; use crate::adapter::telegram::{TelegramAdapter, TelegramAdapterFactory}; use crate::core::acl_service::AclService; +use crate::core::auth::jwt::JwtService; use crate::core::config_service::ConfigService; use crate::core::dns_filter_service::DnsFilterService; use crate::core::email::scheduler::ReportScheduler; @@ -63,7 +62,6 @@ use crate::model::monitoring::direction::FlowDirection; use crate::model::system::config::MLInferenceConfig; use crate::model::system::health::EbpfFailStage; use crate::model::system::health::EbpfHealth; -use macros::log; /// Holds all Arc-wrapped services that make up the running application. pub struct AppState { @@ -260,7 +258,7 @@ impl ServiceFactory { // Create AccessControlPort adapter for SOAR/TTL (decoupled from eBPF) let access_control_port: Arc = - Arc::new(EbpfAccessControlAdapter::new(ebpf_services.access_control.clone())); + Arc::new(AccessControlAdapter::new(ebpf_services.access_control.clone())); // Create SOAR engine let rate_limit_port: Arc = ebpf_services.rate_limit.clone(); diff --git a/net-guardia/src/interface/communication/command.rs b/net-guardia/src/interface/communication/command.rs index 31ada71..102e515 100644 --- a/net-guardia/src/interface/communication/command.rs +++ b/net-guardia/src/interface/communication/command.rs @@ -1,10 +1,12 @@ -use crate::interface::communication::message::Message; -use crate::model::error::Error; -use async_trait::async_trait; use std::any::Any; use std::future::Future; use std::pin::Pin; +use async_trait::async_trait; + +use crate::interface::communication::message::Message; +use crate::model::error::Error; + pub type CommandFuture = Pin> + Send + 'static>>; pub type CommandHandlerFn = Box) -> CommandFuture + Send + Sync>; diff --git a/net-guardia/src/interface/communication/event.rs b/net-guardia/src/interface/communication/event.rs index faeaa07..a4c25b5 100644 --- a/net-guardia/src/interface/communication/event.rs +++ b/net-guardia/src/interface/communication/event.rs @@ -1,6 +1,7 @@ -use crate::model::error::Error; use std::any::Any; +use crate::model::error::Error; + pub trait Event: Send + Clone + 'static {} pub trait EventBroadcaster: Send + Sync { diff --git a/net-guardia/src/interface/communication/query.rs b/net-guardia/src/interface/communication/query.rs index 5a5c6de..3b6dab1 100644 --- a/net-guardia/src/interface/communication/query.rs +++ b/net-guardia/src/interface/communication/query.rs @@ -1,10 +1,12 @@ -use crate::interface::communication::message::Message; -use crate::model::error::Error; -use async_trait::async_trait; use std::any::Any; use std::future::Future; use std::pin::Pin; +use async_trait::async_trait; + +use crate::interface::communication::message::Message; +use crate::model::error::Error; + pub type QueryFuture = Pin, Error>> + Send + 'static>>; pub type QueryHandlerFn = Box) -> QueryFuture + Send + Sync>; diff --git a/net-guardia/src/interface/mod.rs b/net-guardia/src/interface/mod.rs index 27daa2a..50b89bc 100644 --- a/net-guardia/src/interface/mod.rs +++ b/net-guardia/src/interface/mod.rs @@ -1,2 +1,3 @@ pub mod communication; pub mod port; +pub mod utils; diff --git a/net-guardia/src/interface/utils/logging.rs b/net-guardia/src/interface/utils/logging.rs new file mode 100644 index 0000000..37d6f22 --- /dev/null +++ b/net-guardia/src/interface/utils/logging.rs @@ -0,0 +1,6 @@ +use tracing_subscriber::EnvFilter; + +pub trait FilterControl: Send + Sync { + fn reload_filter(&self, filter: EnvFilter) -> Result<(), String>; + fn current_filter(&self) -> String; +} diff --git a/net-guardia/src/interface/utils/mod.rs b/net-guardia/src/interface/utils/mod.rs new file mode 100644 index 0000000..31348d2 --- /dev/null +++ b/net-guardia/src/interface/utils/mod.rs @@ -0,0 +1 @@ +pub mod logging; diff --git a/net-guardia/src/main.rs b/net-guardia/src/main.rs index fbdcc97..58a1f98 100644 --- a/net-guardia/src/main.rs +++ b/net-guardia/src/main.rs @@ -6,20 +6,23 @@ mod model; mod utils; use std::env; +use std::os::unix::process::CommandExt; use std::path::PathBuf; use std::process; use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; -use std::thread; use std::time::Duration; +use clap::Parser; use macros::log; use sd_notify::NotifyState; +use tokio::time::sleep; use tokio::{signal, time}; use crate::adapter::persistence::Database; use crate::core::auth::jwt::JwtService; use crate::core::auth::password; +use crate::infrastructure::cli::{Cli, handle_subcommand}; use crate::infrastructure::http_server; use crate::infrastructure::secret_store::SecretStore; use crate::infrastructure::system::{ShutdownMode, System}; @@ -29,84 +32,16 @@ use crate::model::error::system::SystemError; use crate::model::log::system::SystemLog; use crate::utils::logging::Logging; -/// Two-phase startup: -/// -/// ```text -/// ┌────────────┐ ┌──────────────────┐ ┌──────────────────┐ -/// │ Create DB │────→│ Setup complete? │─no─→│ Setup HTTP server│ -/// │ (fast) │ │ │ │ (instant start) │ -/// └────────────┘ └────────┬─────────┘ └────────┬─────────┘ -/// │yes │done -/// ▼ ▼ -/// ┌──────────────────┐ ┌──────────────────┐ -/// │ Full build │←────│ Stop setup server│ -/// │ (eBPF, ML, SOAR) │ │ + reload config │ -/// └────────┬─────────┘ └──────────────────┘ -/// ▼ -/// ┌──────────────────┐ -/// │ Full HTTP server │ -/// └──────────────────┘ -/// ``` #[actix_web::main] async fn main() -> Result<(), Error> { - Logging::initialize()?; - - // Handle DB encrypt/decrypt subcommands before full startup - let args: Vec = env::args().collect(); - if args.len() >= 2 { - let db_path = env::var("NETGUARDIA_DB_PATH").unwrap_or_else(|_| "net-guardia.db".to_string()); - match args[1].as_str() { - "--decrypt-db" => { - let key = match env::var("NETGUARDIA_DB_KEY") { - Ok(k) if !k.is_empty() => k, - _ => { - eprintln!("Error: NETGUARDIA_DB_KEY must be set for decrypt"); - process::exit(1); - } - }; - let dest = args.get(2).map(|s| s.as_str()).unwrap_or("net-guardia-decrypted.db"); - println!("Decrypting {} → {}", db_path, dest); - Database::decrypt_to_file(&db_path, &key, dest)?; - println!("Done. Decrypted database written to {}", dest); - return Ok(()); - } - "--encrypt-db" => { - let key = match env::var("NETGUARDIA_DB_KEY") { - Ok(k) if !k.is_empty() => k, - _ => { - eprintln!("Error: NETGUARDIA_DB_KEY must be set for encrypt"); - process::exit(1); - } - }; - let dest = args.get(2).map(|s| s.as_str()).unwrap_or("net-guardia-encrypted.db"); - println!("Encrypting {} → {}", db_path, dest); - Database::encrypt_to_file(&db_path, &key, dest)?; - println!("Done. Encrypted database written to {}", dest); - return Ok(()); - } - "--verify-audit-log" => { - println!("Verifying audit_log hash chain in {}", db_path); - let db = Database::new(&db_path)?; - match db.verify_audit_log_chain() { - Ok(count) => { - println!("OK: {} audit_log rows verified, chain intact.", count); - return Ok(()); - } - Err(e) => { - eprintln!("FAIL: {}", e); - process::exit(2); - } - } - } - _ => {} - } + let cli = Cli::parse(); + if cli.command.is_some() { + Logging::initialize_cli()?; + return handle_subcommand(&cli); } - // Phase 1: Create DB (fast — needed for setup check and setup server) - let db_path = env::var("NETGUARDIA_DB_PATH").unwrap_or_else(|_| "net-guardia.db".to_string()); - let db = Arc::new(Database::new(&db_path)?); - - // Seed default admin user if no users exist + Logging::initialize()?; + let db = Arc::new(Database::new(&cli.db_path)?); if db.user_count().unwrap_or(0) == 0 { let hash = password::hash_password("admin")?; let admin_user_id = db.insert_user("admin", &hash, "admin", false)?; @@ -120,8 +55,6 @@ async fn main() -> Result<(), Error> { } let setup_complete = db.get_setting("setup_complete")?.map(|v| v == "true").unwrap_or(false); - - // Phase 2: If setup not complete, run lightweight setup server immediately if !setup_complete { log!(SystemLog::SetupMode); @@ -130,10 +63,8 @@ async fn main() -> Result<(), Error> { let jwt_service = Arc::new(JwtService::new(&secrets, 24)?); let setup_flag = Arc::new(AtomicBool::new(false)); - // Start setup server — returns handle for graceful shutdown let handle = http_server::start_setup_server(db.clone(), secret_store, jwt_service, setup_flag.clone(), 8080)?; - // Wait for setup completion or shutdown signal let flag = setup_flag.clone(); let setup_done = async move { loop { @@ -155,34 +86,26 @@ async fn main() -> Result<(), Error> { } } - // Stop setup server to free the port for full server handle.stop(true).await; log!(SystemLog::SetupServerStopped); } - // Phase 3: Full system build and run (setup is complete, DB has config) let mut system = System::new(db).await?; let mode = system.run().await?; system.terminate().await?; - match mode { ShutdownMode::Restart => { - log!(SystemLog::ApiRestart); + log!(SystemLog::Restart); let _ = sd_notify::notify(false, &[NotifyState::Reloading]); - // Drop System to detach eBPF XDP programs before re-exec drop(system); - // Brief delay for kernel to release XDP/AF_XDP resources - thread::sleep(Duration::from_millis(500)); - // Re-exec self — works with or without systemd - use std::os::unix::process::CommandExt; + sleep(Duration::from_millis(500)).await; let exe = env::current_exe().unwrap_or_else(|_| PathBuf::from("net-guardia")); - let err = process::Command::new(exe).args(env::args().skip(1)).exec(); // replaces current process - // If exec fails, fall through to exit + let err = process::Command::new(exe).args(env::args().skip(1)).exec(); log!(SystemError::UnexpectedError(err)); process::exit(1); } ShutdownMode::Shutdown => { - log!(SystemLog::ApiShutdown); + log!(SystemLog::Shutdown); let _ = sd_notify::notify(false, &[NotifyState::Stopping]); } } diff --git a/net-guardia/src/model/log/cli.rs b/net-guardia/src/model/log/cli.rs new file mode 100644 index 0000000..a46adb7 --- /dev/null +++ b/net-guardia/src/model/log/cli.rs @@ -0,0 +1,38 @@ +use macros::loggable; + +loggable! { + CliLog { + #[error("Decrypting {src} → {dst}")] + DecryptStarted { src: String, dst: String } => tracing::Level::INFO, + + #[error("Done. Decrypted database written to {dst}")] + DecryptCompleted { dst: String } => tracing::Level::INFO, + + #[error("Decrypt failed: {err}")] + DecryptFailed { err: String } => tracing::Level::ERROR, + + #[error("Encrypting {src} → {dst}")] + EncryptStarted { src: String, dst: String } => tracing::Level::INFO, + + #[error("Done. Encrypted database written to {dst}")] + EncryptCompleted { dst: String } => tracing::Level::INFO, + + #[error("Encrypt failed: {err}")] + EncryptFailed { err: String } => tracing::Level::ERROR, + + #[error("Verifying audit_log hash chain in {db_path}")] + VerifyStarted { db_path: String } => tracing::Level::INFO, + + #[error("OK: {count} audit_log rows verified, chain intact.")] + VerifyOk { count: usize } => tracing::Level::INFO, + + #[error("FAIL: {err}")] + VerifyFailed { err: String } => tracing::Level::ERROR, + + #[error("Could not open database at {db_path}: {err}")] + DbOpenFailed { db_path: String, err: String } => tracing::Level::ERROR, + + #[error("NETGUARDIA_DB_KEY must be set for {op}")] + MissingDbKey { op: String } => tracing::Level::ERROR, + } +} diff --git a/net-guardia/src/model/log/mod.rs b/net-guardia/src/model/log/mod.rs index 77f3423..5804375 100644 --- a/net-guardia/src/model/log/mod.rs +++ b/net-guardia/src/model/log/mod.rs @@ -1,4 +1,5 @@ pub mod audit; +pub mod cli; pub mod crypto; pub mod detection; pub mod ebpf; diff --git a/net-guardia/src/model/log/system.rs b/net-guardia/src/model/log/system.rs index 41fdbe8..e385b2e 100644 --- a/net-guardia/src/model/log/system.rs +++ b/net-guardia/src/model/log/system.rs @@ -93,11 +93,11 @@ loggable! { #[error("Failed to restore ACL rule ({direction} {list_type} {address}:{port}): {error}")] AclRuleRestoreFailed { direction: String, list_type: String, address: String, port: u16, error: String } => tracing::Level::WARN, - #[error("API-triggered shutdown initiated")] - ApiShutdown => tracing::Level::INFO, + #[error("Triggered shutdown initiated")] + Shutdown => tracing::Level::INFO, - #[error("API-triggered restart initiated — process will exit and systemd will restart")] - ApiRestart => tracing::Level::INFO, + #[error("Triggered restart initiated — self-reexecuting")] + Restart => tracing::Level::INFO, #[error("ML drift detected: {count} features drifted, max deviation {deviation:.2}σ")] DriftDetected { count: usize, deviation: f64 } => tracing::Level::WARN, diff --git a/net-guardia/src/utils/ip_address.rs b/net-guardia/src/utils/ip_address.rs index a2ed095..563166d 100644 --- a/net-guardia/src/utils/ip_address.rs +++ b/net-guardia/src/utils/ip_address.rs @@ -3,11 +3,6 @@ use std::net::IpAddr; pub fn is_private_ip(ip: &IpAddr) -> bool { match ip { IpAddr::V4(v4) => v4.is_private() || v4.is_loopback() || v4.is_link_local() || v4.is_broadcast(), - IpAddr::V6(v6) => { - v6.is_loopback() - || v6.is_unique_local() // fc00::/7 - || v6.is_unspecified() - || v6.is_multicast() - } + IpAddr::V6(v6) => v6.is_loopback() || v6.is_unique_local() || v6.is_unspecified() || v6.is_multicast(), } } diff --git a/net-guardia/src/utils/logging.rs b/net-guardia/src/utils/logging.rs index cc6b2f5..4cb5e61 100644 --- a/net-guardia/src/utils/logging.rs +++ b/net-guardia/src/utils/logging.rs @@ -1,17 +1,19 @@ -use std::env; use std::fs; use std::sync::OnceLock; +use std::{env, io}; use tracing::Level; +use tracing::level_filters::LevelFilter; use tracing_appender::rolling::{RollingFileAppender, Rotation}; use tracing_subscriber::filter::Directive; use tracing_subscriber::filter::EnvFilter; use tracing_subscriber::fmt::layer as fmt_layer; use tracing_subscriber::layer::SubscriberExt; -use tracing_subscriber::reload; use tracing_subscriber::util::SubscriberInitExt; +use tracing_subscriber::{Layer, filter, reload}; use crate::core::observability::log_buffer::LogBufferLayer; +use crate::interface::utils::logging::FilterControl; use crate::model::error::Error; use crate::model::error::io::IOError; @@ -25,23 +27,6 @@ static FILTER_HANDLE: OnceLock> = OnceLock::new(); /// the operator configured for specific crates from being silently lost. static PRESERVED_DIRECTIVES: OnceLock> = OnceLock::new(); -/// Trait to erase the complex generic type of reload::Handle. -trait FilterControl: Send + Sync { - fn reload_filter(&self, filter: EnvFilter) -> Result<(), String>; - fn current_filter(&self) -> String; -} - -impl FilterControl for reload::Handle { - fn reload_filter(&self, filter: EnvFilter) -> Result<(), String> { - self.reload(filter).map_err(|e| e.to_string()) - } - - fn current_filter(&self) -> String { - self.with_current(|f| f.to_string()) - .unwrap_or_else(|_| "unknown".to_string()) - } -} - pub struct Logging; impl Logging { @@ -116,7 +101,32 @@ impl Logging { Ok(()) } - /// Change the global log level at runtime. + pub fn initialize_cli() -> Result<(), Error> { + let stdout_layer = fmt_layer() + .without_time() + .with_level(false) + .with_target(false) + .with_file(false) + .with_line_number(false) + .with_thread_ids(false) + .with_ansi(false) + .with_writer(io::stdout) + .with_filter(LevelFilter::INFO); + + let stderr_layer = fmt_layer() + .without_time() + .with_level(false) + .with_target(false) + .with_writer(io::stderr) + .with_filter(filter::filter_fn(|m| m.level() <= &Level::WARN)); + + tracing_subscriber::registry() + .with(stdout_layer) + .with(stderr_layer) + .init(); + Ok(()) + } + pub fn set_level(level: &str) -> Result { let handle = FILTER_HANDLE.get().ok_or("Logging not initialized")?; @@ -153,6 +163,17 @@ impl Logging { } } +impl FilterControl for reload::Handle { + fn reload_filter(&self, filter: EnvFilter) -> Result<(), String> { + self.reload(filter).map_err(|e| e.to_string()) + } + + fn current_filter(&self) -> String { + self.with_current(|f| f.to_string()) + .unwrap_or_else(|_| "unknown".to_string()) + } +} + /// Strip per-target directives out of an EnvFilter string and return the /// bare level directive in lowercase. Falls back to the raw string if no /// bare directive is present.