Compare commits

...

2 Commits

115 changed files with 4951 additions and 3170 deletions

1
.gitignore vendored
View File

@ -36,6 +36,7 @@ interfaces.txt
traffic_log.csv
CLAUDE.md
AGENT.md
DESIGN.md
TODOS.md
VERSION

150
Cargo.lock generated
View File

@ -337,7 +337,7 @@ checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0"
dependencies = [
"cfg-if",
"cipher",
"cpufeatures 0.2.17",
"cpufeatures",
]
[[package]]
@ -490,7 +490,7 @@ checksum = "3c3610892ee6e0cbce8ae2700349fcf8f98adb0dbfbee85aec3c9179d29cc072"
dependencies = [
"base64ct",
"blake2",
"cpufeatures 0.2.17",
"cpufeatures",
"password-hash",
]
@ -500,6 +500,18 @@ version = "1.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9b34d609dfbaf33d6889b2b7106d3ca345eacad44200913df5ba02bfd31d2ba9"
[[package]]
name = "async-sqlite"
version = "0.5.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1025c093eea82d9b8a37bbf3065c83814654c5b57f156bd48d49c2a38f99ac18"
dependencies = [
"crossbeam-channel",
"futures-channel",
"futures-util",
"rusqlite",
]
[[package]]
name = "async-trait"
version = "0.1.89"
@ -841,17 +853,6 @@ version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
[[package]]
name = "chacha20"
version = "0.10.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6f8d983286843e49675a4b7a2d174efe136dc93a18d69130dd18198a6c167601"
dependencies = [
"cfg-if",
"cpufeatures 0.3.0",
"rand_core 0.10.0",
]
[[package]]
name = "chrono"
version = "0.4.44"
@ -990,15 +991,6 @@ dependencies = [
"libc",
]
[[package]]
name = "cpufeatures"
version = "0.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201"
dependencies = [
"libc",
]
[[package]]
name = "crc32fast"
version = "1.5.0"
@ -1302,7 +1294,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb"
dependencies = [
"libc",
"windows-sys 0.52.0",
"windows-sys 0.61.2",
]
[[package]]
@ -1492,7 +1484,6 @@ dependencies = [
"cfg-if",
"libc",
"r-efi 6.0.0",
"rand_core 0.10.0",
"wasip2",
"wasip3",
]
@ -1583,11 +1574,11 @@ checksum = "4f467dd6dccf739c208452f8014c75c18bb8301b050ad1cfb27153803edb0f51"
[[package]]
name = "hashlink"
version = "0.10.0"
version = "0.11.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7382cf6263419f2d8df38c55d7da83da5c18aef87fc7a7fc1fb1e344edfe14c1"
checksum = "ea0b22561a9c04a7cb1a302c013e0259cd3b4bb619f145b32f72b8b4bcbed230"
dependencies = [
"hashbrown 0.15.5",
"hashbrown 0.16.1",
]
[[package]]
@ -1741,7 +1732,7 @@ dependencies = [
"libc",
"percent-encoding",
"pin-project-lite",
"socket2 0.5.10",
"socket2 0.6.3",
"tokio",
"tower-service",
"tracing",
@ -2183,9 +2174,9 @@ dependencies = [
[[package]]
name = "libsqlite3-sys"
version = "0.32.0"
version = "0.37.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fbb8270bb4060bd76c6e96f20c52d80620f1d82a3470885694e41e0f81ef6fe7"
checksum = "b1f111c8c41e7c61a49cd34e44c7619462967221a6443b0ec299e0ac30cfb9b1"
dependencies = [
"cc",
"pkg-config",
@ -2489,6 +2480,7 @@ dependencies = [
"aes-gcm",
"arc-swap",
"argon2",
"async-sqlite",
"async-trait",
"aya",
"aya-log",
@ -2516,8 +2508,6 @@ dependencies = [
"network-types",
"notify",
"parking_lot",
"r2d2",
"r2d2_sqlite",
"rand 0.9.2",
"reqwest",
"rusqlite",
@ -2893,7 +2883,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25"
dependencies = [
"cfg-if",
"cpufeatures 0.2.17",
"cpufeatures",
"opaque-debug",
"universal-hash",
]
@ -3013,7 +3003,7 @@ dependencies = [
"quinn-udp",
"rustc-hash",
"rustls",
"socket2 0.5.10",
"socket2 0.6.3",
"thiserror 2.0.18",
"tokio",
"tracing",
@ -3050,7 +3040,7 @@ dependencies = [
"cfg_aliases",
"libc",
"once_cell",
"socket2 0.5.10",
"socket2 0.6.3",
"tracing",
"windows-sys 0.52.0",
]
@ -3082,28 +3072,6 @@ version = "6.0.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf"
[[package]]
name = "r2d2"
version = "0.8.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "51de85fb3fb6524929c8a2eb85e6b6d363de4e8c48f9e2c2eac4944abc181c93"
dependencies = [
"log",
"parking_lot",
"scheduled-thread-pool",
]
[[package]]
name = "r2d2_sqlite"
version = "0.27.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "180da684f0a188977d3968f139eb44260192ef8d9a5b7b7cbd01d881e0353179"
dependencies = [
"r2d2",
"rusqlite",
"uuid",
]
[[package]]
name = "rand"
version = "0.8.5"
@ -3125,17 +3093,6 @@ dependencies = [
"rand_core 0.9.5",
]
[[package]]
name = "rand"
version = "0.10.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bc266eb313df6c5c09c1c7b1fbe2510961e5bcd3add930c1e31f7ed9da0feff8"
dependencies = [
"chacha20",
"getrandom 0.4.2",
"rand_core 0.10.0",
]
[[package]]
name = "rand_chacha"
version = "0.3.1"
@ -3174,12 +3131,6 @@ dependencies = [
"getrandom 0.3.4",
]
[[package]]
name = "rand_core"
version = "0.10.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0c8d0fd677905edcbeedbf2edb6494d676f0e98d54d5cf9bda0b061cb8fb8aba"
[[package]]
name = "rand_distr"
version = "0.4.3"
@ -3302,10 +3253,20 @@ dependencies = [
]
[[package]]
name = "rusqlite"
version = "0.34.0"
name = "rsqlite-vfs"
version = "0.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "37e34486da88d8e051c7c0e23c3f15fd806ea8546260aa2fec247e97242ec143"
checksum = "a8a1f2315036ef6b1fbacd1972e8ee7688030b0a2121edfc2a6550febd41574d"
dependencies = [
"hashbrown 0.16.1",
"thiserror 2.0.18",
]
[[package]]
name = "rusqlite"
version = "0.39.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a0d2b0146dd9661bf67bb107c0bb2a55064d556eeb3fc314151b957f313bcd4e"
dependencies = [
"bitflags 2.11.0",
"fallible-iterator",
@ -3313,6 +3274,7 @@ dependencies = [
"hashlink",
"libsqlite3-sys",
"smallvec",
"sqlite-wasm-rs",
]
[[package]]
@ -3388,7 +3350,7 @@ dependencies = [
"errno",
"libc",
"linux-raw-sys",
"windows-sys 0.52.0",
"windows-sys 0.61.2",
]
[[package]]
@ -3467,15 +3429,6 @@ dependencies = [
"regex",
]
[[package]]
name = "scheduled-thread-pool"
version = "0.2.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3cbc66816425a074528352f5789333ecff06ca41b36b0b0efdfbb29edc391a19"
dependencies = [
"parking_lot",
]
[[package]]
name = "scopeguard"
version = "1.2.0"
@ -3594,7 +3547,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba"
dependencies = [
"cfg-if",
"cpufeatures 0.2.17",
"cpufeatures",
"digest",
]
@ -3605,7 +3558,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283"
dependencies = [
"cfg-if",
"cpufeatures 0.2.17",
"cpufeatures",
"digest",
]
@ -3684,6 +3637,18 @@ dependencies = [
"windows-sys 0.61.2",
]
[[package]]
name = "sqlite-wasm-rs"
version = "0.5.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1b2c760607300407ddeaee518acf28c795661b7108c75421303dbefb237d3a36"
dependencies = [
"cc",
"js-sys",
"rsqlite-vfs",
"wasm-bindgen",
]
[[package]]
name = "stable_deref_trait"
version = "1.2.1"
@ -3805,10 +3770,10 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd"
dependencies = [
"fastrand",
"getrandom 0.3.4",
"getrandom 0.4.2",
"once_cell",
"rustix",
"windows-sys 0.52.0",
"windows-sys 0.61.2",
]
[[package]]
@ -4419,7 +4384,6 @@ checksum = "a68d3c8f01c0cfa54a75291d83601161799e4a89a39e0929f4b0354d88757a37"
dependencies = [
"getrandom 0.4.2",
"js-sys",
"rand 0.10.0",
"wasm-bindgen",
]
@ -4639,7 +4603,7 @@ version = "0.1.11"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22"
dependencies = [
"windows-sys 0.52.0",
"windows-sys 0.61.2",
]
[[package]]

View File

@ -4,6 +4,10 @@ members = ["net-guardia", "common", "macros", "ingress-ebpf", "egress-ebpf", "mc
default-members = ["net-guardia", "common", "mcp-server", "cli"]
[workspace.dependencies]
# Local crates
common = { path = "common" }
macros = { path = "macros" }
# eBPF - kernel side (pinned: aya-ebpf 0.1.2 was yanked, see aya-rs/aya#1400)
aya-ebpf = { version = "=0.1.1", default-features = false }
aya-log-ebpf = { version = "=0.1.0", default-features = false }
@ -30,6 +34,9 @@ actix = "0.13.5"
actix-web = "4.13.0"
actix-cors = "0.7.1"
actix-ws = "0.4.0"
actix-multipart = "0.7"
actix-files = "0.6"
tokio-tungstenite = "0.28.0"
# Logging / tracing
tracing = "0.1.44"
@ -51,17 +58,37 @@ sysinfo = "0.38.4"
maxminddb = "0.27.3"
ipnetwork = "0.21.1"
lru = "0.16.3"
rusqlite = { version = "0.34", features = ["bundled"] }
rusqlite = { version = "0.39", features = ["bundled-sqlcipher"] }
async-sqlite = { version = "0.5.7", default-features = false, features = ["bundled-sqlcipher"] }
jsonwebtoken = "9"
argon2 = "0.5"
rand = "0.9"
ed25519-dalek = { version = "2", features = ["std", "rand_core"] }
base64 = "0.22"
clap = { version = "4", features = ["derive"] }
uuid = { version = "1", features = ["v4"] }
rust-embed = "8.11.0"
mime_guess = "2.0.5"
url = "2.5.8"
toml = "1.0.7"
lettre = { version = "0.11", default-features = false, features = ["builder", "hostname", "smtp-transport", "tokio1-rustls-tls"] }
chrono = { version = "0.4", default-features = false, features = ["clock", "std"] }
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
async-trait = "0.1"
dashmap = "6"
arc-swap = "1"
moka = { version = "0.12", features = ["sync"] }
notify = "7"
sha2 = "0.10"
hmac = "0.12"
aes-gcm = "0.10"
hkdf = "0.12"
sd-notify = "0.4"
# Build dependencies
cargo_metadata = { version = "0.23.1", default-features = false }
which = "8.0.2"
dotenvy = "0.15.7"
# Proc macro
proc-macro2 = "1.0.106"

View File

@ -4,7 +4,7 @@ version = "1.0.0"
edition = "2024"
[dependencies]
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
reqwest = { workspace = true }
serde = { workspace = true }
serde_json = { workspace = true }
tokio = { workspace = true }

View File

@ -4,7 +4,7 @@ version = "1.0.0"
edition = "2024"
[dependencies]
common = { path = "../common", features = ["kernel"] }
common = { workspace = true, features = ["kernel"] }
aya-ebpf = { workspace = true }
aya-log-ebpf = { workspace = true }
network-types = { workspace = true }

View File

@ -4,7 +4,7 @@ version = "1.0.0"
edition = "2024"
[dependencies]
common = { path = "../common", features = ["kernel"] }
common = { workspace = true, features = ["kernel"] }
aya-ebpf = { workspace = true }
aya-log-ebpf = { workspace = true }
network-types = { workspace = true }

View File

@ -267,7 +267,7 @@ fn gen_mapped_default(mp: &MappedParent) -> TokenStream2 {
quote! { #ident: #ty { #(#sub_fields,)* } }
}
// ── Code generation: from_settings() ───────────────────────────────
// ── Code generation: from_config_repo() ───────────────────────────────
fn gen_override(f: &SettingField) -> TokenStream2 {
let ident = &f.ident;
@ -278,25 +278,25 @@ fn gen_override(f: &SettingField) -> TokenStream2 {
quote! {
crate::domain::common::config::helpers::override_string_nonempty(
&mut cfg.#ident, repo, #key,
)?;
).await?;
}
} else if is_type(ty, "bool") {
quote! {
crate::domain::common::config::helpers::override_bool(
&mut cfg.#ident, repo, #key,
)?;
).await?;
}
} else if is_vec_string(ty) {
quote! {
crate::domain::common::config::helpers::override_csv(
&mut cfg.#ident, repo, #key,
)?;
).await?;
}
} else {
quote! {
crate::domain::common::config::helpers::override_parsed(
&mut cfg.#ident, repo, #key,
)?;
).await?;
}
}
}
@ -304,7 +304,7 @@ fn gen_override(f: &SettingField) -> TokenStream2 {
fn gen_flatten_override(f: &FlattenField) -> TokenStream2 {
let ident = &f.ident;
let ty = &f.ty;
quote! { cfg.#ident = #ty::from_settings(repo)?; }
quote! { cfg.#ident = #ty::from_config_repo(repo).await?; }
}
fn gen_mapped_overrides(mp: &MappedParent) -> TokenStream2 {
@ -318,14 +318,14 @@ fn gen_mapped_overrides(mp: &MappedParent) -> TokenStream2 {
quote! {
crate::domain::common::config::helpers::override_parsed(
&mut cfg.#parent.#sub, repo, #key,
)?;
).await?;
}
})
.collect();
quote! { #(#calls)* }
}
// ── Code generation: seed_defaults() ───────────────────────────────
// ── Code generation: seed_config_defaults() ───────────────────────────────
fn gen_seed(f: &SettingField) -> TokenStream2 {
let key = &f.key;
@ -336,17 +336,17 @@ fn gen_seed(f: &SettingField) -> TokenStream2 {
crate::domain::common::config::helpers::seed_key(
repo, #key,
if cfg!(debug_assertions) { #dbg } else { #default },
)?;
).await?;
},
None => quote! {
crate::domain::common::config::helpers::seed_key(repo, #key, #default)?;
crate::domain::common::config::helpers::seed_key(repo, #key, #default).await?;
},
}
}
fn gen_flatten_seed(f: &FlattenField) -> TokenStream2 {
let ty = &f.ty;
quote! { #ty::seed_defaults(repo)?; }
quote! { #ty::seed_config_defaults(repo).await?; }
}
fn gen_mapped_seeds(mp: &MappedParent) -> TokenStream2 {
@ -361,10 +361,10 @@ fn gen_mapped_seeds(mp: &MappedParent) -> TokenStream2 {
crate::domain::common::config::helpers::seed_key(
repo, #key,
if cfg!(debug_assertions) { #dbg } else { #default },
)?;
).await?;
},
None => quote! {
crate::domain::common::config::helpers::seed_key(repo, #key, #default)?;
crate::domain::common::config::helpers::seed_key(repo, #key, #default).await?;
},
}
})
@ -421,6 +421,56 @@ fn gen_keys_consts(fields: &[ConfigField]) -> TokenStream2 {
.collect()
}
// ── Code generation: api_values() ─────────────────────────────────
fn value_to_string_expr(ty: &Type, expr: TokenStream2) -> TokenStream2 {
if is_type(ty, "String") {
quote! { #expr.clone() }
} else if is_vec_string(ty) {
quote! { #expr.join(",") }
} else {
quote! { #expr.to_string() }
}
}
fn gen_api_values(fields: &[ConfigField]) -> TokenStream2 {
let entries: Vec<_> = fields
.iter()
.flat_map(|f| match f {
ConfigField::Setting(s) if s.api => {
let key = &s.key;
let ident = &s.ident;
let value = value_to_string_expr(&s.ty, quote! { self.#ident });
vec![quote! { values.push((#key, #value)); }]
}
ConfigField::Flatten(f) => {
let ident = &f.ident;
vec![quote! { values.extend(self.#ident.api_values()); }]
}
ConfigField::MappedParent(mp) => mp
.settings
.iter()
.filter(|s| s.api)
.map(|s| {
let key = &s.key;
let parent = &mp.ident;
let sub = format_ident!("{}", s.sub_field);
quote! { values.push((#key, self.#parent.#sub.to_string())); }
})
.collect(),
_ => Vec::new(),
})
.collect();
quote! {
pub fn api_values(&self) -> Vec<(&'static str, String)> {
let mut values = Vec::new();
#(#entries)*
values
}
}
}
// ── Entry point ────────────────────────────────────────────────────
pub fn config_settings_impl(attr: TokenStream, item: TokenStream) -> TokenStream {
@ -456,6 +506,7 @@ pub fn config_settings_impl(attr: TokenStream, item: TokenStream) -> TokenStream
let struct_name = &input.ident;
let keys_consts = gen_keys_consts(&config_fields);
let api_values = gen_api_values(&config_fields);
let default_fields: Vec<_> = config_fields
.iter()
@ -489,6 +540,7 @@ pub fn config_settings_impl(attr: TokenStream, item: TokenStream) -> TokenStream
impl #struct_name {
#keys_consts
#api_values
pub fn defaults() -> Self {
Self {
@ -496,16 +548,16 @@ pub fn config_settings_impl(attr: TokenStream, item: TokenStream) -> TokenStream
}
}
pub fn from_settings(
repo: &dyn crate::interface::setting::SettingRepo,
pub async fn from_config_repo(
repo: &dyn crate::interface::config_repo::ConfigRepo,
) -> Result<Self, crate::domain::common::error::Error> {
let mut cfg = Self::defaults();
#(#override_calls)*
Ok(cfg)
}
pub fn seed_defaults(
repo: &dyn crate::interface::setting::SettingRepo,
pub async fn seed_config_defaults(
repo: &dyn crate::interface::config_repo::ConfigRepo,
) -> Result<(), crate::domain::common::error::Error> {
#(#seed_calls)*
Ok(())

View File

@ -4,7 +4,7 @@ version = "1.0.0"
edition = "2024"
[dependencies]
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
reqwest = { workspace = true }
serde = { workspace = true }
serde_json = { workspace = true }
tokio = { workspace = true }

View File

@ -4,8 +4,8 @@ version = "1.0.0"
edition = "2024"
[dependencies]
common = { path = "../common", features = ["user"] }
macros = { path = "../macros" }
common = { workspace = true, features = ["user"] }
macros = { workspace = true }
# eBPF userspace
aya = { workspace = true }
@ -20,19 +20,19 @@ actix = { workspace = true }
actix-web = { workspace = true }
actix-cors = { workspace = true }
actix-ws = { workspace = true }
actix-multipart = "0.7"
actix-files = "0.6"
uuid = { version = "1", features = ["v4"] }
rust-embed = "8.11.0"
mime_guess = "2.0.5"
url = "2.5.8"
tokio-tungstenite = "0.28.0"
actix-multipart = { workspace = true }
actix-files = { workspace = true }
uuid = { workspace = true }
rust-embed = { workspace = true }
mime_guess = { workspace = true }
url = { workspace = true }
tokio-tungstenite = { workspace = true }
# Serialization
serde = { workspace = true }
serde_json = { workspace = true }
serde_yaml_ng = { workspace = true }
toml = "1.0.7"
toml = { workspace = true }
# Async
tokio = { workspace = true }
@ -48,21 +48,21 @@ tracing-subscriber = { workspace = true }
tract-onnx = { workspace = true }
# Email
lettre = { version = "0.11", default-features = false, features = ["builder", "hostname", "smtp-transport", "tokio1-rustls-tls"] }
chrono = { version = "0.4", default-features = false, features = ["clock", "std"] }
lettre = { workspace = true }
chrono = { workspace = true }
# HTTP client (Telegram, MCP proxy)
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
reqwest = { workspace = true }
# CLI
clap = { version = "4", features = ["derive", "env"] }
clap = { workspace = true, features = ["env"] }
# Architecture
async-trait = "0.1"
dashmap = "6"
arc-swap = "1"
moka = { version = "0.12", features = ["sync"] }
notify = "7"
async-trait = { workspace = true }
dashmap = { workspace = true }
arc-swap = { workspace = true }
moka = { workspace = true }
notify = { workspace = true }
# Utilities
parking_lot = { workspace = true }
@ -71,17 +71,16 @@ sysinfo = { workspace = true }
maxminddb = { workspace = true }
ipnetwork = { workspace = true }
lru = { workspace = true }
rusqlite = { version = "0.34", features = ["bundled-sqlcipher"] }
r2d2 = "0.8"
r2d2_sqlite = "0.27"
rusqlite = { workspace = true }
async-sqlite = { workspace = true }
jsonwebtoken = { workspace = true }
argon2 = { workspace = true }
sha2 = "0.10"
hmac = "0.12"
aes-gcm = "0.10"
hkdf = "0.12"
sha2 = { workspace = true }
hmac = { workspace = true }
aes-gcm = { workspace = true }
hkdf = { workspace = true }
base64 = { workspace = true }
sd-notify = "0.4"
sd-notify = { workspace = true }
rand = { workspace = true }
[features]
@ -90,7 +89,7 @@ default = []
[build-dependencies]
cargo_metadata = { workspace = true }
which = { workspace = true }
dotenvy = "0.15.7"
dotenvy = { workspace = true }
[[bin]]
name = "net-guardia"

View File

@ -36,7 +36,7 @@ impl GeoBlock {
let v6_trie = LpmTrie::try_from(v6_map).map_err(EbpfError::MapOperationError)?;
// todo read config from AppConfig, not db
let db_path = app_config.load().acl.geoip_db_name.clone();
let db_path = app_config.load().acl.geoip_db_path.clone();
let reader = Reader::open_readfile(&db_path).map_err(|e| MiscError::GeoIPDatabaseError(db_path.clone(), e))?;
let index = Self::build_index(&reader)?;
@ -50,7 +50,7 @@ impl GeoBlock {
}
pub fn unavailable(app_config: Arc<ArcSwap<AppConfig>>) -> Self {
let index = Reader::open_readfile(&app_config.load().acl.geoip_db_name)
let index = Reader::open_readfile(&app_config.load().acl.geoip_db_path)
.ok()
.and_then(|reader| Self::build_index(&reader).ok())
.unwrap_or(GeoIndex {

View File

@ -183,6 +183,9 @@ pub struct XskPair {
drop_monitor: Option<Arc<DropMonitor>>,
packet_buffer_size: usize,
buffer_pool_capacity: usize,
completion_batch_size: usize,
rx_batch_size: usize,
tx_batch_size: usize,
tx_packet_buf: Vec<Vec<u8>>,
tx_frame_buf: Vec<FrameDesc>,
}
@ -255,8 +258,11 @@ impl XskPair {
drop_monitor,
packet_buffer_size: config.packet_buffer_size,
buffer_pool_capacity: config.buffer_pool_capacity,
tx_packet_buf: Vec::with_capacity(64),
tx_frame_buf: Vec::with_capacity(64),
completion_batch_size: config.xsk_completion_batch_size,
rx_batch_size: config.xsk_rx_batch_size,
tx_batch_size: config.xsk_tx_batch_size,
tx_packet_buf: Vec::with_capacity(config.xsk_tx_batch_size),
tx_frame_buf: Vec::with_capacity(config.xsk_tx_batch_size),
};
Ok(xsk_pair)
@ -277,8 +283,8 @@ impl XskPair {
let mut shutdown_rx = Some(shutdown_rx);
let mut idle_count: u32 = 0;
let mut buffer_pool = BufferPool::new(self.buffer_pool_capacity, self.packet_buffer_size);
let mut comp_descs = vec![FrameDesc::default(); 256];
let mut rx_descs = vec![FrameDesc::default(); 64];
let mut comp_descs = vec![FrameDesc::default(); self.completion_batch_size];
let mut rx_descs = vec![FrameDesc::default(); self.rx_batch_size];
loop {
if let Some(ref mut rx) = shutdown_rx {
@ -426,7 +432,7 @@ impl XskPair {
self.tx_packet_buf.clear();
while let Ok(packet) = forward_rx.try_recv() {
self.tx_packet_buf.push(packet);
if self.tx_packet_buf.len() >= 64 {
if self.tx_packet_buf.len() >= self.tx_batch_size {
break;
}
}

View File

@ -13,7 +13,7 @@ pub fn initialize() -> Scope {
}
async fn list_audit_logs(_auth: AuthClaims, db: web::Data<Database>) -> HttpResponse {
match db.list_audit_logs() {
match db.list_audit_logs().await {
Ok(entries) => {
let json: Vec<serde_json::Value> = entries
.into_iter()
@ -40,7 +40,7 @@ async fn list_audit_logs(_auth: AuthClaims, db: web::Data<Database>) -> HttpResp
/// integrity without shell access. Any mismatch returns the offending
/// row id inside `error` so the dashboard can link straight to it.
async fn verify_chain(_auth: AuthClaims, audit: web::Data<dyn AuditRepo>) -> HttpResponse {
match audit.verify_audit_log_chain(0) {
match audit.verify_audit_log_chain(0).await {
Ok((count, _last_id)) => HttpResponse::Ok().json(serde_json::json!({
"chain_intact": true,
"verified": count,

View File

@ -44,7 +44,7 @@ async fn add_ipv4_list(
acl: web::Data<AclService>,
) -> impl Responder {
let (direction, list_type) = path.into_inner();
ok_or_error(acl.add_ipv4(direction, list_type, address.into_inner()))
ok_or_error(acl.add_ipv4(direction, list_type, address.into_inner()).await)
}
async fn add_ipv6_list(
@ -53,7 +53,7 @@ async fn add_ipv6_list(
acl: web::Data<AclService>,
) -> impl Responder {
let (direction, list_type) = path.into_inner();
ok_or_error(acl.add_ipv6(direction, list_type, address.into_inner()))
ok_or_error(acl.add_ipv6(direction, list_type, address.into_inner()).await)
}
async fn remove_ipv4_list(
@ -62,7 +62,7 @@ async fn remove_ipv4_list(
acl: web::Data<AclService>,
) -> impl Responder {
let (direction, list_type) = path.into_inner();
ok_or_error(acl.remove_ipv4(direction, list_type, address.into_inner()))
ok_or_error(acl.remove_ipv4(direction, list_type, address.into_inner()).await)
}
async fn remove_ipv6_list(
@ -71,7 +71,7 @@ async fn remove_ipv6_list(
acl: web::Data<AclService>,
) -> impl Responder {
let (direction, list_type) = path.into_inner();
ok_or_error(acl.remove_ipv6(direction, list_type, address.into_inner()))
ok_or_error(acl.remove_ipv6(direction, list_type, address.into_inner()).await)
}
async fn get_geo_blocked(acl: web::Data<AclService>) -> impl Responder {
@ -80,7 +80,7 @@ async fn get_geo_blocked(acl: web::Data<AclService>) -> impl Responder {
async fn block_geo_countries(body: web::Json<CountryCodesRequest>, acl: web::Data<AclService>) -> impl Responder {
let codes = body.into_inner().country_codes;
match acl.block_geo_countries(&codes) {
match acl.block_geo_countries(&codes).await {
Ok(total_prefixes) => HttpResponse::Ok().json(serde_json::json!({
"blocked_countries": acl.get_blocked_countries(),
"total_prefixes": total_prefixes,
@ -91,7 +91,7 @@ async fn block_geo_countries(body: web::Json<CountryCodesRequest>, acl: web::Dat
async fn unblock_geo_countries(body: web::Json<CountryCodesRequest>, acl: web::Data<AclService>) -> impl Responder {
let codes = body.into_inner().country_codes;
match acl.unblock_geo_countries(&codes) {
match acl.unblock_geo_countries(&codes).await {
Ok(total_prefixes) => HttpResponse::Ok().json(serde_json::json!({
"blocked_countries": acl.get_blocked_countries(),
"total_prefixes": total_prefixes,

View File

@ -46,7 +46,7 @@ async fn add_dns_blacklist(
service: web::Data<DnsFilterService>,
) -> impl Responder {
let domains = payload.into_inner().domains;
match service.add_domains(&domains) {
match service.add_domains(&domains).await {
Ok(count) => HttpResponse::Ok().json(serde_json::json!({"added": count})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
}
@ -57,7 +57,7 @@ async fn remove_dns_blacklist(
service: web::Data<DnsFilterService>,
) -> impl Responder {
let domains = payload.into_inner().domains;
match service.remove_domains(&domains) {
match service.remove_domains(&domains).await {
Ok(count) => HttpResponse::Ok().json(serde_json::json!({"removed": count})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
}

View File

@ -22,5 +22,5 @@ async fn get_config(service: web::Data<RateLimitService>) -> impl Responder {
}
async fn set_config(settings: web::Json<RateLimitSettings>, service: web::Data<RateLimitService>) -> impl Responder {
ok_or_error(service.update(&settings.into_inner()))
ok_or_error(service.update(&settings.into_inner()).await)
}

View File

@ -48,7 +48,10 @@ async fn explain_ip(
};
let obs = app_config.load().observability.clone();
let entries = match audit.list_audit_logs_by_action(FUSION_AUDIT_ACTION, obs.fusion_explain_scan_limit) {
let entries = match audit
.list_audit_logs_by_action(FUSION_AUDIT_ACTION, obs.fusion_explain_scan_limit)
.await
{
Ok(e) => e,
Err(e) => {
return HttpResponse::InternalServerError().json(serde_json::json!({

View File

@ -584,26 +584,32 @@ async fn validate_and_promote(ctx: &PromoteContext<'_>) -> Result<PromoteReport,
let before_status = inference.model_source_status();
let _guard = promote_lock.try_acquire().ok_or(PromoteError::ConcurrentPromote)?;
let models_dir = PathBuf::from(MODELS_DIR);
let target_onnx = models_dir.join(&declared_onnx);
fs::rename(&staged_onnx, &target_onnx)
.await
.map_err(|e| PromoteError::PromoteIo(format!("rename onnx into models/: {e}")))?;
if let Some(ref pp) = manifest.preprocessing {
validate_manifest_basename("preprocessing.scaler_sidecar", &pp.scaler_sidecar)?;
let src = staging_dir.join(&pp.scaler_sidecar);
let dst = models_dir.join(&pp.scaler_sidecar);
fs::rename(&src, &dst)
.await
.map_err(|e| PromoteError::PromoteIo(format!("rename sidecar into models/: {e}")))?;
}
let target_manifest = models_dir.join(MANIFEST_FILENAME);
fs::rename(&staging_manifest, &target_manifest)
.await
.map_err(|e| PromoteError::PromoteIo(format!("rename manifest into models/: {e}")))?;
let staged_sidecar = if let Some(ref pp) = manifest.preprocessing {
validate_manifest_basename("preprocessing.scaler_sidecar", &pp.scaler_sidecar)?;
Some(staging_dir.join(&pp.scaler_sidecar))
} else {
None
};
let target_sidecar = manifest
.preprocessing
.as_ref()
.map(|pp| models_dir.join(&pp.scaler_sidecar));
let _guard = promote_lock.try_acquire().ok_or(PromoteError::ConcurrentPromote)?;
let backup_dir = models_dir.join(format!(".promote-backup-{}", Uuid::new_v4()));
promote_files_atomically(&PromoteFileSet {
staging_manifest: staging_manifest.clone(),
staged_onnx,
staged_sidecar,
target_manifest,
target_onnx,
target_sidecar,
backup_dir,
})
.await?;
drop(_guard);
let audit_detail = serde_json::json!({
@ -628,6 +634,105 @@ async fn validate_and_promote(ctx: &PromoteContext<'_>) -> Result<PromoteReport,
})
}
struct PromoteFileSet {
staging_manifest: PathBuf,
staged_onnx: PathBuf,
staged_sidecar: Option<PathBuf>,
target_manifest: PathBuf,
target_onnx: PathBuf,
target_sidecar: Option<PathBuf>,
backup_dir: PathBuf,
}
async fn promote_files_atomically(files: &PromoteFileSet) -> Result<(), PromoteError> {
fs::create_dir_all(&files.backup_dir)
.await
.map_err(|e| PromoteError::PromoteIo(format!("create promote backup dir: {e}")))?;
let backup_manifest = backup_existing(&files.target_manifest, &files.backup_dir, "manifest.yaml").await?;
let backup_onnx = backup_existing(&files.target_onnx, &files.backup_dir, "model.onnx").await?;
let backup_sidecar = match &files.target_sidecar {
Some(target) => Some(backup_existing(target, &files.backup_dir, "sidecar").await?),
None => None,
};
let result = async {
move_file(&files.staged_onnx, &files.target_onnx, "rename onnx into models/").await?;
if let (Some(src), Some(dst)) = (&files.staged_sidecar, &files.target_sidecar) {
move_file(src, dst, "rename sidecar into models/").await?;
}
move_file(
&files.staging_manifest,
&files.target_manifest,
"rename manifest into models/",
)
.await
}
.await;
match result {
Ok(()) => {
let _ = fs::remove_dir_all(&files.backup_dir).await;
Ok(())
}
Err(err) => {
rollback_promote(files, backup_manifest, backup_onnx, backup_sidecar).await;
let _ = fs::remove_dir_all(&files.backup_dir).await;
Err(err)
}
}
}
async fn backup_existing(target: &Path, backup_dir: &Path, backup_name: &str) -> Result<Option<PathBuf>, PromoteError> {
if !target
.try_exists()
.map_err(|e| PromoteError::PromoteIo(format!("check existing target {}: {e}", target.display())))?
{
return Ok(None);
}
let backup = backup_dir.join(backup_name);
fs::rename(target, &backup)
.await
.map_err(|e| PromoteError::PromoteIo(format!("backup existing target {}: {e}", target.display())))?;
Ok(Some(backup))
}
async fn move_file(src: &Path, dst: &Path, op: &str) -> Result<(), PromoteError> {
fs::rename(src, dst)
.await
.map_err(|e| PromoteError::PromoteIo(format!("{op}: {e}")))
}
async fn rollback_promote(
files: &PromoteFileSet,
backup_manifest: Option<PathBuf>,
backup_onnx: Option<PathBuf>,
backup_sidecar: Option<Option<PathBuf>>,
) {
remove_if_exists(&files.target_manifest).await;
remove_if_exists(&files.target_onnx).await;
if let Some(target) = &files.target_sidecar {
remove_if_exists(target).await;
}
restore_backup(backup_manifest, &files.target_manifest).await;
restore_backup(backup_onnx, &files.target_onnx).await;
if let (Some(backup), Some(target)) = (backup_sidecar.flatten(), &files.target_sidecar) {
restore_backup(Some(backup), target).await;
}
}
async fn remove_if_exists(path: &Path) {
if let Ok(true) = path.try_exists() {
let _ = fs::remove_file(path).await;
}
}
async fn restore_backup(backup: Option<PathBuf>, target: &Path) {
if let Some(backup) = backup {
let _ = fs::rename(backup, target).await;
}
}
/// Metadata surfaced back to the client when the promote succeeds.
#[derive(Debug)]
struct PromoteReport {
@ -858,6 +963,48 @@ mod tests {
assert_eq!(resp.status().as_u16(), 500);
}
#[tokio::test]
async fn promote_files_rolls_back_active_files_when_sidecar_move_fails() {
let tmp = std::env::temp_dir().join(format!("nguardia-promote-rollback-{}", Uuid::new_v4()));
let staging = tmp.join("staging");
let models = tmp.join("models");
let backup = models.join(".promote-backup-test");
fs::create_dir_all(&staging).await.unwrap();
fs::create_dir_all(&models).await.unwrap();
let staging_manifest = staging.join(MANIFEST_FILENAME);
let staged_onnx = staging.join("model.onnx");
let missing_sidecar = staging.join("missing-scaler.json");
let target_manifest = models.join(MANIFEST_FILENAME);
let target_onnx = models.join("model.onnx");
let target_sidecar = models.join("scaler.json");
fs::write(&staging_manifest, b"new manifest").await.unwrap();
fs::write(&staged_onnx, b"new onnx").await.unwrap();
fs::write(&target_manifest, b"old manifest").await.unwrap();
fs::write(&target_onnx, b"old onnx").await.unwrap();
fs::write(&target_sidecar, b"old sidecar").await.unwrap();
let err = promote_files_atomically(&PromoteFileSet {
staging_manifest: staging_manifest.clone(),
staged_onnx,
staged_sidecar: Some(missing_sidecar),
target_manifest: target_manifest.clone(),
target_onnx: target_onnx.clone(),
target_sidecar: Some(target_sidecar.clone()),
backup_dir: backup,
})
.await
.expect_err("missing sidecar should fail promote");
assert!(matches!(err, PromoteError::PromoteIo(_)));
assert_eq!(fs::read(&target_manifest).await.unwrap(), b"old manifest");
assert_eq!(fs::read(&target_onnx).await.unwrap(), b"old onnx");
assert_eq!(fs::read(&target_sidecar).await.unwrap(), b"old sidecar");
fs::remove_dir_all(&tmp).await.ok();
}
#[test]
fn upload_error_scaler_too_large_maps_to_413_and_echoes_cap() {
let resp = UploadError::ScalerTooLarge(1234).into_response();

View File

@ -13,7 +13,7 @@ pub fn initialize() -> Scope {
}
async fn list_keys(_auth: AuthClaims, db: web::Data<dyn ApiKeyRepo>) -> HttpResponse {
match db.list_api_keys() {
match db.list_api_keys().await {
Ok(keys) => {
let responses: Vec<serde_json::Value> = keys
.into_iter()
@ -62,7 +62,7 @@ async fn generate_key(
}));
};
match db.insert_api_key(&key_hash, &body.name, level.as_str()) {
match db.insert_api_key(&key_hash, &body.name, level.as_str()).await {
Ok(id) => HttpResponse::Created().json(serde_json::json!({
"id": id,
"key": raw_key,
@ -75,7 +75,7 @@ async fn generate_key(
async fn delete_key(_auth: AuthClaims, db: web::Data<dyn ApiKeyRepo>, path: web::Path<i64>) -> HttpResponse {
let id = path.into_inner();
match db.delete_api_key(id) {
match db.delete_api_key(id).await {
Ok(true) => HttpResponse::Ok().json(serde_json::json!({"deleted": true})),
Ok(false) => HttpResponse::NotFound().json(serde_json::json!({"error": "Key not found"})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),

View File

@ -54,7 +54,7 @@ pub fn initialize() -> Scope {
async fn login(body: web::Json<LoginRequest>, auth_svc: web::Data<AuthService>) -> impl Responder {
let req = body.into_inner();
match auth_svc.login(&req.username, &req.password) {
match auth_svc.login(&req.username, &req.password).await {
Ok(result) => HttpResponse::Ok().json(result),
Err(LoginError::Locked { retry_after_secs }) => HttpResponse::TooManyRequests().json(serde_json::json!({
"error": "Account temporarily locked due to too many failed login attempts",
@ -75,7 +75,10 @@ async fn register(
auth_svc: web::Data<AuthService>,
) -> impl Responder {
let reg = body.into_inner();
match auth_svc.register(&reg.username, &reg.password, &reg.role, &auth.role) {
match auth_svc
.register(&reg.username, &reg.password, &reg.role, &auth.role)
.await
{
Ok(_) => HttpResponse::Created().json(serde_json::json!({"username": reg.username, "role": reg.role})),
Err(RegisterError::Validation(msg)) => HttpResponse::BadRequest().json(serde_json::json!({"error": msg})),
Err(RegisterError::InvalidRole) => {
@ -91,7 +94,7 @@ async fn register(
}
async fn me(auth: AuthClaims, auth_svc: web::Data<AuthService>) -> impl Responder {
let profile = auth_svc.user_profile(auth.sub, &auth.username);
let profile = auth_svc.user_profile(auth.sub, &auth.username).await;
HttpResponse::Ok().json(profile)
}
@ -106,7 +109,7 @@ async fn change_password(
return HttpResponse::BadRequest().json(serde_json::json!({"error": msg}));
}
let user = match db.find_user(&auth.username) {
let user = match db.find_user(&auth.username).await {
Ok(Some(u)) => u,
_ => {
return HttpResponse::InternalServerError().json(serde_json::json!({"error": "User not found"}));
@ -127,14 +130,14 @@ async fn change_password(
}
};
match db.update_user_password(auth.sub, &new_hash) {
match db.update_user_password(auth.sub, &new_hash).await {
Ok(_) => HttpResponse::Ok().json(serde_json::json!({"message": "Password changed successfully"})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
}
}
async fn list_users(_auth: AuthClaims, db: web::Data<Repo>) -> impl Responder {
match db.list_users_with_groups() {
match db.list_users_with_groups().await {
Ok(users) => {
let result: Vec<serde_json::Value> = users
.into_iter()
@ -172,7 +175,7 @@ async fn delete_user(_auth: AuthClaims, path: web::Path<i64>, db: web::Data<Repo
return HttpResponse::BadRequest().json(serde_json::json!({"error": "Cannot delete your own account"}));
}
match db.find_user_by_id(user_id) {
match db.find_user_by_id(user_id).await {
Ok(Some(ref u)) if u.username == DEFAULT_ADMIN_USERNAME => {
return HttpResponse::Forbidden()
.json(serde_json::json!({"error": "Cannot delete the built-in admin account"}));
@ -180,7 +183,7 @@ async fn delete_user(_auth: AuthClaims, path: web::Path<i64>, db: web::Data<Repo
_ => {}
}
match db.delete_user(user_id) {
match db.delete_user(user_id).await {
Ok(true) => HttpResponse::Ok().json(serde_json::json!({"message": "User deleted successfully"})),
Ok(false) => HttpResponse::NotFound().json(serde_json::json!({"error": "User not found"})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
@ -206,7 +209,7 @@ async fn update_role(
}
};
match db.find_user_by_id(user_id) {
match db.find_user_by_id(user_id).await {
Ok(Some(_)) => {}
Ok(None) => {
return HttpResponse::NotFound().json(serde_json::json!({"error": "User not found"}));
@ -216,7 +219,7 @@ async fn update_role(
}
}
match db.update_user_role(user_id, role) {
match db.update_user_role(user_id, role).await {
Ok(_) => HttpResponse::Ok().json(serde_json::json!({"message": "Role updated successfully", "role": role})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
}
@ -245,7 +248,7 @@ async fn reset_password(
return HttpResponse::BadRequest().json(serde_json::json!({"error": msg}));
}
match db.find_user_by_id(user_id) {
match db.find_user_by_id(user_id).await {
Ok(Some(_)) => {}
Ok(None) => {
return HttpResponse::NotFound().json(serde_json::json!({"error": "User not found"}));
@ -262,32 +265,31 @@ async fn reset_password(
}
};
ok_or_error(db.reset_user_password(user_id, &hash))
ok_or_error(db.reset_user_password(user_id, &hash).await)
}
async fn list_groups(_auth: AuthClaims, db: web::Data<Repo>) -> impl Responder {
match db.list_user_groups() {
match db.list_user_groups().await {
Ok(groups) => {
let result: Vec<serde_json::Value> = groups
.into_iter()
.map(|g| {
let perms: serde_json::Value = parse_permissions(&g.permissions);
let members: Vec<serde_json::Value> = db
.list_group_members(g.id)
.unwrap_or_default()
.into_iter()
.map(|m| serde_json::json!({"id": m.id, "username": m.username}))
.collect();
serde_json::json!({
"id": g.id,
"name": g.name,
"description": g.description,
"permissions": perms,
"created_at": g.created_at,
"members": members,
})
})
.collect();
let mut result = Vec::with_capacity(groups.len());
for g in groups {
let perms: serde_json::Value = parse_permissions(&g.permissions);
let members: Vec<serde_json::Value> = db
.list_group_members(g.id)
.await
.unwrap_or_default()
.into_iter()
.map(|m| serde_json::json!({"id": m.id, "username": m.username}))
.collect();
result.push(serde_json::json!({
"id": g.id,
"name": g.name,
"description": g.description,
"permissions": perms,
"created_at": g.created_at,
"members": members,
}));
}
HttpResponse::Ok().json(result)
}
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
@ -308,7 +310,7 @@ async fn create_group(_auth: AuthClaims, body: web::Json<serde_json::Value>, db:
_ => "[]".to_string(),
};
match db.create_user_group(name, description, &permissions) {
match db.create_user_group(name, description, &permissions).await {
Ok(id) => HttpResponse::Created().json(serde_json::json!({
"id": id,
"name": name,
@ -322,10 +324,10 @@ async fn create_group(_auth: AuthClaims, body: web::Json<serde_json::Value>, db:
async fn get_group(_auth: AuthClaims, path: web::Path<i64>, db: web::Data<Repo>) -> impl Responder {
let group_id = path.into_inner();
match db.get_user_group(group_id) {
match db.get_user_group(group_id).await {
Ok(Some(g)) => {
let perms: serde_json::Value = parse_permissions(&g.permissions);
let members = db.list_group_member_ids(group_id).unwrap_or_default();
let members = db.list_group_member_ids(group_id).await.unwrap_or_default();
HttpResponse::Ok().json(serde_json::json!({
"id": g.id,
"name": g.name,
@ -348,7 +350,7 @@ async fn update_group(
) -> impl Responder {
let group_id = path.into_inner();
let existing = match db.get_user_group(group_id) {
let existing = match db.get_user_group(group_id).await {
Ok(Some(g)) => {
if g.name == GROUP_ADMIN || g.name == GROUP_VIEWER {
return HttpResponse::Forbidden().json(serde_json::json!({"error": "Cannot modify built-in groups"}));
@ -373,7 +375,7 @@ async fn update_group(
_ => existing.permissions.clone(),
};
match db.update_user_group(group_id, name, description, &permissions) {
match db.update_user_group(group_id, name, description, &permissions).await {
Ok(_) => HttpResponse::Ok().json(serde_json::json!({
"id": group_id,
"name": name,
@ -387,14 +389,14 @@ async fn update_group(
async fn delete_group(_auth: AuthClaims, path: web::Path<i64>, db: web::Data<Repo>) -> impl Responder {
let group_id = path.into_inner();
match db.get_user_group(group_id) {
match db.get_user_group(group_id).await {
Ok(Some(ref g)) if g.name == GROUP_ADMIN || g.name == GROUP_VIEWER => {
return HttpResponse::Forbidden().json(serde_json::json!({"error": "Cannot delete built-in groups"}));
}
_ => {}
}
match db.delete_user_group(group_id) {
match db.delete_user_group(group_id).await {
Ok(true) => HttpResponse::Ok().json(serde_json::json!({"message": "Group deleted successfully"})),
Ok(false) => HttpResponse::NotFound().json(serde_json::json!({"error": "Group not found"})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
@ -409,7 +411,7 @@ async fn set_user_groups(
) -> impl Responder {
let user_id = path.into_inner();
match db.find_user_by_id(user_id) {
match db.find_user_by_id(user_id).await {
Ok(Some(ref u)) if u.username == DEFAULT_ADMIN_USERNAME => {
return HttpResponse::Forbidden()
.json(serde_json::json!({"error": "Cannot modify groups for the built-in admin account"}));
@ -430,7 +432,7 @@ async fn set_user_groups(
}
};
match db.set_user_groups(user_id, &group_ids) {
match db.set_user_groups(user_id, &group_ids).await {
Ok(_) => HttpResponse::Ok()
.json(serde_json::json!({"message": "User groups updated successfully", "group_ids": group_ids})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),

View File

@ -82,15 +82,15 @@ mod tests {
use crate::adapter::persistence::Database;
use crate::infrastructure::secret_store::SecretStore;
fn test_jwt_service() -> JwtService {
let db = Arc::new(Database::new(":memory:").unwrap());
async fn test_jwt_service() -> JwtService {
let db = Arc::new(Database::new(":memory:").await.unwrap());
let secrets: Arc<dyn SecretStorePort> = Arc::new(SecretStore::new(db));
JwtService::new(&secrets, 24).unwrap()
}
#[test]
fn test_create_and_validate_token() {
let jwt = test_jwt_service();
#[tokio::test]
async fn test_create_and_validate_token() {
let jwt = test_jwt_service().await;
let perms = vec!["dashboard:read".to_string()];
let token = jwt.create_token(1, "admin", "admin", perms.clone()).unwrap();
let claims = jwt.validate_token(&token).unwrap();
@ -100,16 +100,16 @@ mod tests {
assert_eq!(claims.permissions, perms);
}
#[test]
fn test_invalid_token() {
let jwt = test_jwt_service();
#[tokio::test]
async fn test_invalid_token() {
let jwt = test_jwt_service().await;
let result = jwt.validate_token("invalid.token.here");
assert!(result.is_err());
}
#[test]
fn test_expired_token() {
let db = Arc::new(Database::new(":memory:").unwrap());
#[tokio::test]
async fn test_expired_token() {
let db = Arc::new(Database::new(":memory:").await.unwrap());
let secrets: Arc<dyn SecretStorePort> = Arc::new(SecretStore::new(db));
let jwt = JwtService::new(&secrets, 0).unwrap(); // 0 hours = immediate expiry
@ -126,9 +126,9 @@ mod tests {
assert!(result.is_err());
}
#[test]
fn test_jwt_secret_changes_on_new_instance() {
let db = Arc::new(Database::new(":memory:").unwrap());
#[tokio::test]
async fn test_jwt_secret_changes_on_new_instance() {
let db = Arc::new(Database::new(":memory:").await.unwrap());
let secrets: Arc<dyn SecretStorePort> = Arc::new(SecretStore::new(db));
let jwt1 = JwtService::new(&secrets, 24).unwrap();
@ -140,10 +140,10 @@ mod tests {
assert!(result.is_err());
}
#[test]
fn test_different_secrets_reject() {
let jwt1 = test_jwt_service();
let jwt2 = test_jwt_service(); // different in-memory DB = different secret
#[tokio::test]
async fn test_different_secrets_reject() {
let jwt1 = test_jwt_service().await;
let jwt2 = test_jwt_service().await; // different in-memory DB = different secret
let token = jwt1.create_token(1, "admin", "admin", vec![]).unwrap();
let result = jwt2.validate_token(&token);

View File

@ -1,6 +1,6 @@
use std::fs;
use std::io::ErrorKind;
use std::path::Path;
use std::path::PathBuf;
use std::time::UNIX_EPOCH;
use actix_web::{HttpResponse, Scope, web};
@ -10,9 +10,6 @@ use serde::{Deserialize, Serialize};
use crate::domain::common::config::AppConfig;
use crate::infrastructure::log_buffer::{self, LogBuffer, LogEntry};
/// Hardcoded log directory — not configurable via API to prevent directory traversal.
const LOG_DIR: &str = "logs";
/// Validate log filename: only alphanumeric, dots, underscores, hyphens.
/// Prevents path traversal.
fn is_valid_log_filename(name: &str) -> bool {
@ -87,9 +84,9 @@ struct LogFileEntry {
modified: Option<u64>,
}
async fn list_logs() -> HttpResponse {
let log_dir = LOG_DIR;
let entries = match fs::read_dir(log_dir) {
async fn list_logs(app_config: web::Data<ArcSwap<AppConfig>>) -> HttpResponse {
let log_dir = app_config.load().system.log_dir.clone();
let entries = match fs::read_dir(&log_dir) {
Ok(dir) => dir
.filter_map(|e| e.ok())
.filter_map(|e| {
@ -117,7 +114,9 @@ async fn list_logs() -> HttpResponse {
}
async fn download_log(path: web::Path<String>, app_config: web::Data<ArcSwap<AppConfig>>) -> HttpResponse {
let max_download_size = app_config.load().observability.log_max_download_size;
let config = app_config.load();
let max_download_size = config.observability.log_max_download_size;
let log_dir = PathBuf::from(&config.system.log_dir);
let filename = path.into_inner();
if !is_valid_log_filename(&filename) {
@ -126,7 +125,7 @@ async fn download_log(path: web::Path<String>, app_config: web::Data<ArcSwap<App
}));
}
let file_path = Path::new(LOG_DIR).join(&filename);
let file_path = log_dir.join(&filename);
// Canonicalize to prevent symlink traversal
let canonical = match fs::canonicalize(&file_path) {
@ -137,7 +136,7 @@ async fn download_log(path: web::Path<String>, app_config: web::Data<ArcSwap<App
}));
}
};
if let Ok(log_dir_canonical) = fs::canonicalize(LOG_DIR)
if let Ok(log_dir_canonical) = fs::canonicalize(&log_dir)
&& !canonical.starts_with(&log_dir_canonical)
{
return HttpResponse::Forbidden().json(serde_json::json!({

View File

@ -63,12 +63,13 @@ fn required_permission(path: &str, method: &Method) -> Option<String> {
"rate_limit"
} else if path.starts_with("/api/system/") {
"system"
} else if path == "/api/api-keys" || path.starts_with("/api/api-keys/") {
return Some("api_keys:admin".to_string());
} else if path.contains("/soar/blocks/") && path.ends_with("/unblock") {
return Some("access_control:write".to_string());
} else if path.starts_with("/api/soar/")
|| path.starts_with("/api/notifications/")
|| path.starts_with("/api/report/")
|| path.starts_with("/api/api-keys/")
|| path.starts_with("/api/logs/")
|| path.starts_with("/api/audit/")
{
@ -165,7 +166,7 @@ where
"apikey:{}",
req.peer_addr().map(|a| a.ip().to_string()).unwrap_or_default()
);
if let Ok(Some(remaining)) = repo.check_login_locked(&rate_key) {
if let Ok(Some(remaining)) = repo.check_login_locked(&rate_key).await {
let resp = HttpResponse::TooManyRequests().json(serde_json::json!({
"error": "Too many failed API key attempts",
"retry_after_secs": remaining,
@ -173,15 +174,15 @@ where
return Ok(req.into_response(resp).map_into_right_body());
}
match api_key_port.validate_api_key(api_key) {
match api_key_port.validate_api_key(api_key).await {
Ok(Some(key_claims)) => {
if let Err(e) = repo.clear_login_failures(&rate_key) {
if let Err(e) = repo.clear_login_failures(&rate_key).await {
log!(AuthError::LoginClearError(e));
}
key_claims
}
Ok(None) => {
if let Err(e) = repo.record_login_failure(&rate_key) {
if let Err(e) = repo.record_login_failure(&rate_key).await {
log!(AuthError::LoginFailureTrackingError(e));
}
let resp = HttpResponse::Unauthorized()
@ -216,3 +217,30 @@ where
})
}
}
#[cfg(test)]
mod tests {
use actix_web::http::Method;
use super::required_permission;
#[test]
fn api_key_collection_requires_admin_permission() {
assert_eq!(
required_permission("/api/api-keys", &Method::GET),
Some("api_keys:admin".to_string())
);
}
#[test]
fn api_key_subroutes_require_admin_permission() {
assert_eq!(
required_permission("/api/api-keys/generate", &Method::POST),
Some("api_keys:admin".to_string())
);
assert_eq!(
required_permission("/api/api-keys/1", &Method::DELETE),
Some("api_keys:admin".to_string())
);
}
}

View File

@ -14,7 +14,7 @@ pub fn initialize() -> Scope {
}
async fn get_telegram_config(_auth: AuthClaims, svc: web::Data<NotificationService>) -> HttpResponse {
ok_json_or_error(svc.get_telegram_config())
ok_json_or_error(svc.get_telegram_config().await)
}
#[derive(Deserialize)]
@ -28,7 +28,7 @@ async fn set_telegram_config(
svc: web::Data<NotificationService>,
body: web::Json<TelegramConfigRequest>,
) -> HttpResponse {
ok_or_error(svc.set_telegram_config(&body.bot_token, &body.chat_id))
ok_or_error(svc.set_telegram_config(&body.bot_token, &body.chat_id).await)
}
async fn test_telegram(_auth: AuthClaims, svc: web::Data<NotificationService>) -> HttpResponse {
@ -39,7 +39,7 @@ async fn test_telegram(_auth: AuthClaims, svc: web::Data<NotificationService>) -
}
async fn test_smtp(_auth: AuthClaims, svc: web::Data<NotificationService>) -> HttpResponse {
match svc.test_smtp() {
match svc.test_smtp().await {
Ok(msg) => HttpResponse::Ok().json(serde_json::json!({"success": true, "message": msg})),
Err(e) => HttpResponse::BadRequest().json(serde_json::json!({"success": false, "error": e.to_string()})),
}

View File

@ -14,8 +14,8 @@ use crate::core::reporting::report_engine;
use crate::domain::common::config::AppConfig;
use crate::domain::common::error::misc::MiscError;
use crate::infrastructure::secret_store::SecretStore;
use crate::interface::report_snapshot::ReportSnapshotRepo;
use crate::interface::secret_store::SecretStorePort;
use crate::interface::setting::SettingRepo;
pub fn initialize() -> Scope {
web::scope("/report")
@ -36,7 +36,7 @@ async fn generate_report(
}));
}
let db_ref = db.get_ref();
match report_engine::generate_html_report(db_ref as &dyn SettingRepo, &report_dir) {
match report_engine::generate_html_report(db_ref as &dyn ReportSnapshotRepo, &report_dir).await {
Ok(path) => match fs::read(&path) {
Ok(content) => HttpResponse::Ok()
.content_type("text/html; charset=utf-8")
@ -62,7 +62,7 @@ async fn generate_report(
async fn report_data(_auth: AuthClaims, db: web::Data<Database>) -> HttpResponse {
let db_ref = db.get_ref();
ok_json_or_error(report_engine::generate_report_json(db_ref as &dyn SettingRepo))
ok_json_or_error(report_engine::generate_report_json(db_ref as &dyn ReportSnapshotRepo).await)
}
/// Manually trigger: generate the weekly report and send it via SMTP now.
@ -72,11 +72,11 @@ async fn send_report(
config: web::Data<ArcSwap<AppConfig>>,
secrets: web::Data<SecretStore>,
) -> HttpResponse {
let db_ref = db.get_ref() as &dyn SettingRepo;
let db_ref = db.get_ref() as &dyn ReportSnapshotRepo;
let secrets_ref = secrets.get_ref() as &dyn SecretStorePort;
let smtp_cfg = config.load().notification.smtp.clone();
let smtp = match SmtpClient::from_config(&smtp_cfg, Some(secrets_ref)) {
let smtp = match SmtpClient::from_config(&smtp_cfg, Some(secrets_ref)).await {
Ok(Some(client)) => client,
Ok(None) => {
return HttpResponse::BadRequest().json(serde_json::json!({
@ -100,7 +100,7 @@ async fn send_report(
}));
}
let html = match generate_weekly_report(db_ref) {
let html = match generate_weekly_report(db_ref).await {
Ok(h) => h,
Err(e) => {
return HttpResponse::InternalServerError().json(serde_json::json!({

View File

@ -96,7 +96,7 @@ pub fn initialize() -> Scope {
}
async fn list_playbooks(_auth: AuthClaims, svc: web::Data<PlaybookService>) -> HttpResponse {
ok_json_or_error(svc.list_playbooks())
ok_json_or_error(svc.list_playbooks().await)
}
async fn create_playbook(
@ -106,7 +106,7 @@ async fn create_playbook(
body: web::Json<CreatePlaybookRequest>,
) -> HttpResponse {
let input = map_request_to_input(&body, app_config.load().soar.fallback_cooldown_secs);
match svc.create_playbook(&input) {
match svc.create_playbook(&input).await {
Ok(id) => HttpResponse::Created().json(serde_json::json!({"id": id})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
}
@ -121,7 +121,7 @@ async fn update_playbook(
) -> HttpResponse {
let id = path.into_inner();
let input = map_request_to_input(&body, app_config.load().soar.fallback_cooldown_secs);
match svc.update_playbook(id, &input) {
match svc.update_playbook(id, &input).await {
Ok(true) => HttpResponse::Ok().json(serde_json::json!({"updated": true})),
Ok(false) => HttpResponse::NotFound().json(serde_json::json!({"error": "Playbook not found"})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
@ -139,7 +139,7 @@ async fn toggle_playbook(
path: web::Path<i64>,
body: web::Json<TogglePlaybookRequest>,
) -> HttpResponse {
match svc.toggle_playbook(path.into_inner(), body.enabled) {
match svc.toggle_playbook(path.into_inner(), body.enabled).await {
Ok(true) => HttpResponse::Ok().json(serde_json::json!({"updated": true})),
Ok(false) => HttpResponse::NotFound().json(serde_json::json!({"error": "Playbook not found"})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
@ -147,7 +147,7 @@ async fn toggle_playbook(
}
async fn delete_playbook(_auth: AuthClaims, svc: web::Data<PlaybookService>, path: web::Path<i64>) -> HttpResponse {
match svc.delete_playbook(path.into_inner()) {
match svc.delete_playbook(path.into_inner()).await {
Ok(true) => HttpResponse::Ok().json(serde_json::json!({"deleted": true})),
Ok(false) => HttpResponse::NotFound().json(serde_json::json!({"error": "Playbook not found"})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
@ -155,19 +155,23 @@ async fn delete_playbook(_auth: AuthClaims, svc: web::Data<PlaybookService>, pat
}
async fn list_active_blocks(_auth: AuthClaims, svc: web::Data<PlaybookService>) -> HttpResponse {
ok_json_or_error(svc.list_active_blocks())
ok_json_or_error(svc.list_active_blocks().await)
}
async fn manual_unblock(_auth: AuthClaims, svc: web::Data<PlaybookService>, path: web::Path<i64>) -> HttpResponse {
ok_or_error(svc.manual_unblock(path.into_inner()).await)
}
async fn list_executions(_auth: AuthClaims, svc: web::Data<PlaybookService>) -> HttpResponse {
ok_json_or_error(svc.list_executions(100))
async fn list_executions(
_auth: AuthClaims,
svc: web::Data<PlaybookService>,
app_config: web::Data<ArcSwap<AppConfig>>,
) -> HttpResponse {
ok_json_or_error(svc.list_executions(app_config.load().soar.execution_list_limit).await)
}
async fn list_whitelist(_auth: AuthClaims, svc: web::Data<PlaybookService>) -> HttpResponse {
ok_json_or_error(svc.list_whitelist())
ok_json_or_error(svc.list_whitelist().await)
}
#[derive(Deserialize)]
@ -180,14 +184,14 @@ async fn add_whitelist(
svc: web::Data<PlaybookService>,
body: web::Json<WhitelistRequest>,
) -> HttpResponse {
match svc.add_whitelist(&body.ip) {
match svc.add_whitelist(&body.ip).await {
Ok(()) => HttpResponse::Created().json(serde_json::json!({"added": true})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
}
}
async fn remove_whitelist(_auth: AuthClaims, svc: web::Data<PlaybookService>, path: web::Path<String>) -> HttpResponse {
ok_or_error(svc.remove_whitelist(&path.into_inner()))
ok_or_error(svc.remove_whitelist(&path.into_inner()).await)
}
/// Client shape for `POST /api/soar/dry-run`. Only the fields a SOAR

View File

@ -15,6 +15,7 @@ use crate::domain::identity::password;
use crate::infrastructure::http_server::SetupCompleteFlag;
use crate::infrastructure::secret_store::SecretStore;
use crate::interface::secret_store::SecretStorePort;
use crate::interface::system_state::SystemStateRepo;
pub fn initialize() -> Scope {
web::scope("/setup")
@ -133,7 +134,7 @@ async fn complete_setup(
}
// Save configuration to database
if let Err(e) = save_config(&db, secret_store.as_ref(), &body) {
if let Err(e) = save_config(&db, secret_store.as_ref(), &body).await {
return HttpResponse::InternalServerError().json(serde_json::json!({
"error": format!("Failed to save configuration: {}", e)
}));
@ -143,12 +144,12 @@ async fn complete_setup(
match password::hash_password(&body.admin_password) {
Ok(hash) => {
// Find admin user and update password
if let Ok(Some(user)) = db.find_user(DEFAULT_ADMIN_USERNAME) {
if let Err(e) = db.update_user_password(user.id, &hash) {
if let Ok(Some(user)) = db.find_user(DEFAULT_ADMIN_USERNAME).await {
if let Err(e) = db.update_user_password(user.id, &hash).await {
log!(SystemError::SetupPasswordUpdateFailed(e));
}
// Clear force_password_change since setup wizard set the password
if let Err(e) = db.reset_user_password(user.id, &hash) {
if let Err(e) = db.reset_user_password(user.id, &hash).await {
log!(SystemError::SetupPasswordUpdateFailed(e));
}
}
@ -161,7 +162,8 @@ async fn complete_setup(
}
// Mark setup as complete
if let Err(e) = db.set_setting("setup_complete", "true") {
let state_repo = db.get_ref() as &dyn SystemStateRepo;
if let Err(e) = state_repo.set_system_state("setup_complete", "true").await {
log!(SystemError::SetupCompleteFlagFailed(e));
}
setup_flag.0.store(true, Ordering::SeqCst);
@ -176,43 +178,43 @@ async fn complete_setup(
}))
}
fn save_config(db: &Database, secrets: &dyn SecretStorePort, req: &SetupRequest) -> Result<(), Error> {
async fn save_config(db: &Database, secrets: &dyn SecretStorePort, req: &SetupRequest) -> Result<(), Error> {
// Save network config
db.set_setting("ingress_interface", &req.ingress_interface)?;
db.set_setting("egress_interface", &req.egress_interface)?;
db.set_config_value("ingress_interface", &req.ingress_interface).await?;
db.set_config_value("egress_interface", &req.egress_interface).await?;
if let Some(port) = req.http_port {
db.set_setting("http_port", &port.to_string())?;
db.set_config_value("http_port", &port.to_string()).await?;
}
// Save SMTP config (non-secret fields go to settings)
if let Some(host) = &req.smtp_host {
db.set_setting("smtp_host", host)?;
db.set_config_value("smtp_host", host).await?;
}
if let Some(port) = req.smtp_port {
db.set_setting("smtp_port", &port.to_string())?;
db.set_config_value("smtp_port", &port.to_string()).await?;
}
if let Some(user) = &req.smtp_username {
db.set_setting("smtp_username", user)?;
db.set_config_value("smtp_username", user).await?;
}
if let Some(pass) = &req.smtp_password {
// Store password through secret store (encrypted)
secrets.set_secret("smtp_password", pass)?;
db.set_setting("smtp_password", "__encrypted__")?;
secrets.set_secret("smtp_password", pass).await?;
db.set_config_value("smtp_password", "__encrypted__").await?;
}
if let Some(recipient) = &req.smtp_recipient {
db.set_setting("smtp_recipient", recipient)?;
db.set_config_value("smtp_recipient", recipient).await?;
}
// Save Telegram config (bot_token through secret store, chat_id in JSON)
if let (Some(token), Some(chat_id)) = (&req.telegram_bot_token, &req.telegram_chat_id) {
secrets.set_secret("telegram_bot_token", token)?;
secrets.set_secret("telegram_bot_token", token).await?;
let config_json = serde_json::json!({
"bot_token": "__encrypted__",
"chat_id": chat_id,
})
.to_string();
db.set_notification_config("telegram", &config_json)?;
db.set_notification_config("telegram", &config_json).await?;
}
Ok(())

View File

@ -8,12 +8,10 @@ use crate::domain::common::config::constants::{
ENFORCE_MODE_ENFORCE, ENFORCE_MODE_ML_ONLY, ENFORCE_MODE_MONITOR, PERMISSION_SYSTEM_ADMIN,
};
use crate::infrastructure::logger::Logger;
use crate::infrastructure::runtime_state::RuntimeState;
use crate::infrastructure::system::{ShutdownHandle, ShutdownMode};
use crate::interface::app_repo::AppRepo;
use crate::utils::boot_time;
type Repo = dyn AppRepo;
#[derive(Deserialize)]
struct EnforceModeRequest {
mode: String,
@ -38,10 +36,7 @@ async fn get_boot_time() -> impl Responder {
}
async fn get_enforce_mode(handler: web::Data<EnforceModeHandler>) -> impl Responder {
match handler.get_mode() {
Ok(mode) => HttpResponse::Ok().json(serde_json::json!({"mode": mode})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
}
HttpResponse::Ok().json(serde_json::json!({"mode": handler.get_mode()}))
}
async fn set_enforce_mode(
@ -54,33 +49,23 @@ async fn set_enforce_mode(
.json(serde_json::json!({"error": "Mode must be 'monitor', 'ml_only', or 'enforce'"}));
}
match handler.change_mode(mode.clone()) {
match handler.change_mode(mode.clone()).await {
Ok(_) => HttpResponse::Ok().json(serde_json::json!({"mode": mode})),
Err(e) => HttpResponse::InternalServerError().json(serde_json::json!({"error": e.to_string()})),
}
}
async fn get_xdp_mode(db: web::Data<Repo>) -> impl Responder {
// todo get from config, not db
let ingress = db
.get_setting("xdp_ingress_mode")
.ok()
.flatten()
.unwrap_or_else(|| "unknown".to_string());
let egress = db
.get_setting("xdp_egress_mode")
.ok()
.flatten()
.unwrap_or_else(|| "unknown".to_string());
async fn get_xdp_mode(runtime_state: web::Data<arc_swap::ArcSwap<RuntimeState>>) -> impl Responder {
let xdp = runtime_state.load().xdp.clone();
HttpResponse::Ok().json(serde_json::json!({
"ingress_mode": ingress,
"egress_mode": egress,
"ingress_mode": xdp.ingress_mode,
"egress_mode": xdp.egress_mode,
}))
}
async fn get_config(svc: web::Data<ConfigService>) -> impl Responder {
HttpResponse::Ok().json(svc.get_config())
HttpResponse::Ok().json(svc.get_config().await)
}
async fn get_log_level(logging: web::Data<Logger>) -> impl Responder {
@ -112,7 +97,7 @@ async fn update_config(
svc: web::Data<ConfigService>,
handle: web::Data<ShutdownHandle>,
) -> impl Responder {
match svc.update_config(&body) {
match svc.update_config(&body).await {
Ok(updated) => {
let needs_restart = updated.iter().any(|k| HTTP_RELOAD_KEYS.contains(&k.as_str()));
if needs_restart {

View File

@ -17,12 +17,16 @@ pub struct SmtpClient {
}
impl SmtpClient {
pub fn from_config(cfg: &SmtpConfig, secrets: Option<&dyn SecretStorePort>) -> Result<Option<Self>, Error> {
pub async fn from_config(cfg: &SmtpConfig, secrets: Option<&dyn SecretStorePort>) -> Result<Option<Self>, Error> {
if cfg.host.is_empty() || cfg.username.is_empty() {
return Ok(None);
}
let password = match secrets.and_then(|ss| ss.get_secret("smtp_password").ok().flatten()) {
let secret = match secrets {
Some(ss) => ss.get_secret("smtp_password").await?,
None => None,
};
let password = match secret {
Some(pw) if !pw.is_empty() => pw,
_ => return Ok(None),
};
@ -94,12 +98,15 @@ impl EmailSender for SmtpClient {
pub struct SmtpClientFactory;
#[async_trait::async_trait]
impl EmailSenderFactory for SmtpClientFactory {
fn build_smtp_sender(
async fn build_smtp_sender(
&self,
cfg: &SmtpConfig,
secrets: Option<&dyn SecretStorePort>,
) -> Result<Option<Box<dyn EmailSender>>, Error> {
SmtpClient::from_config(cfg, secrets).map(|opt| opt.map(|c| Box::new(c) as Box<dyn EmailSender>))
SmtpClient::from_config(cfg, secrets)
.await
.map(|opt| opt.map(|c| Box::new(c) as Box<dyn EmailSender>))
}
}

View File

@ -1,3 +1,4 @@
use async_trait::async_trait;
use rusqlite::params;
use super::Database;
@ -6,7 +7,7 @@ use crate::domain::data_plane::acl_rule::AclRuleView;
use crate::interface::acl::AclRepo;
impl Database {
pub fn insert_acl_rule(
pub async fn insert_acl_rule(
&self,
ip_version: u8,
direction: &str,
@ -14,91 +15,135 @@ impl Database {
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 list_acl_rules(&self) -> Result<Vec<AclRuleView>, 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(AclRuleView {
ip_version: row.get(0)?,
direction: row.get(1)?,
list_type: row.get(2)?,
ip_address: row.get(3)?,
port: row.get::<_, i64>(4)? as u16,
let direction = direction.to_string();
let list_type = list_type.to_string();
let ip_address = ip_address.to_string();
self.pool
.conn_and_then(move |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],
)?;
if direction == "source" && list_type == "blacklist" {
conn.execute(
"UPDATE soar_block_rules
SET preserve_acl_on_unblock = 1
WHERE source_ip = ?1 AND unblocked_at IS NULL",
params![ip_address],
)?;
}
Ok(())
})
})?;
let mut results = Vec::new();
for row in rows {
results.push(row?);
}
Ok(results)
.await
}
pub fn has_manual_acl_rule(&self, ip_address: &str) -> Result<bool, Error> {
let conn = self.conn()?;
let count: i64 = conn.query_row(
"SELECT COUNT(*) FROM acl_rules
WHERE ip_address = ?1 AND list_type = 'blacklist'
AND NOT EXISTS (
SELECT 1 FROM soar_block_rules
WHERE soar_block_rules.source_ip = acl_rules.ip_address
AND soar_block_rules.unblocked_at IS NULL
)",
params![ip_address],
|row| row.get(0),
)?;
Ok(count > 0)
pub async fn delete_acl_rule(
&self,
ip_version: u8,
direction: &str,
list_type: &str,
ip_address: &str,
port: u16,
) -> Result<(), Error> {
let direction = direction.to_string();
let list_type = list_type.to_string();
let ip_address = ip_address.to_string();
self.pool
.conn_and_then(move |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(())
})
.await
}
pub fn list_admin_whitelist(&self) -> Result<Vec<String>, 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 async fn list_acl_rules(&self) -> Result<Vec<AclRuleView>, Error> {
self.pool
.conn_and_then(move |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(AclRuleView {
ip_version: row.get(0)?,
direction: row.get(1)?,
list_type: row.get(2)?,
ip_address: row.get(3)?,
port: row.get::<_, i64>(4)? as u16,
})
})?;
let mut results = Vec::new();
for row in rows {
results.push(row?);
}
Ok(results)
})
.await
}
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 async fn has_manual_acl_rule(&self, ip_address: &str) -> Result<bool, Error> {
let ip_address = ip_address.to_string();
self.pool
.conn_and_then(move |conn| {
let count: i64 = conn.query_row(
"SELECT COUNT(*) FROM acl_rules
WHERE ip_address = ?1 AND list_type = 'blacklist'
AND EXISTS (
SELECT 1 FROM soar_block_rules
WHERE soar_block_rules.source_ip = acl_rules.ip_address
AND soar_block_rules.unblocked_at IS NULL
AND (
soar_block_rules.preserve_acl_on_unblock = 1
OR soar_block_rules.created_acl_rule = 0
)
)",
params![ip_address],
|row| row.get(0),
)?;
Ok(count > 0)
})
.await
}
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(())
pub async fn list_admin_whitelist(&self) -> Result<Vec<String>, Error> {
self.pool
.conn_and_then(move |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)
})
.await
}
pub async fn insert_admin_whitelist(&self, ip: &str) -> Result<(), Error> {
let ip = ip.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute("INSERT OR IGNORE INTO admin_whitelist (ip) VALUES (?1)", params![ip])?;
Ok(())
})
.await
}
pub async fn delete_admin_whitelist(&self, ip: &str) -> Result<(), Error> {
let ip = ip.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute("DELETE FROM admin_whitelist WHERE ip = ?1", params![ip])?;
Ok(())
})
.await
}
}
#[async_trait]
impl AclRepo for Database {
fn insert_acl_rule(
async fn insert_acl_rule(
&self,
ip_version: u8,
direction: &str,
@ -107,9 +152,10 @@ impl AclRepo for Database {
port: u16,
) -> Result<(), Error> {
self.insert_acl_rule(ip_version, direction, list_type, ip_address, port)
.await
}
fn delete_acl_rule(
async fn delete_acl_rule(
&self,
ip_version: u8,
direction: &str,
@ -118,43 +164,22 @@ impl AclRepo for Database {
port: u16,
) -> Result<(), Error> {
self.delete_acl_rule(ip_version, direction, list_type, ip_address, port)
.await
}
fn has_manual_acl_rule(&self, ip_address: &str) -> Result<bool, Error> {
self.has_manual_acl_rule(ip_address)
async fn has_manual_acl_rule(&self, ip_address: &str) -> Result<bool, Error> {
self.has_manual_acl_rule(ip_address).await
}
fn list_admin_whitelist(&self) -> Result<Vec<String>, Error> {
self.list_admin_whitelist()
async fn list_admin_whitelist(&self) -> Result<Vec<String>, Error> {
self.list_admin_whitelist().await
}
fn insert_admin_whitelist(&self, ip: &str) -> Result<(), Error> {
self.insert_admin_whitelist(ip)
async fn insert_admin_whitelist(&self, ip: &str) -> Result<(), Error> {
self.insert_admin_whitelist(ip).await
}
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.list_acl_rules().unwrap();
assert_eq!(rules.len(), 1);
assert_eq!(rules[0].ip_version, 4);
assert_eq!(rules[0].direction, "source");
assert_eq!(rules[0].list_type, "blacklist");
assert_eq!(rules[0].ip_address, "192.168.1.1");
assert_eq!(rules[0].port, 80);
db.delete_acl_rule(4, "source", "blacklist", "192.168.1.1", 80).unwrap();
let rules = db.list_acl_rules().unwrap();
assert!(rules.is_empty());
async fn delete_admin_whitelist(&self, ip: &str) -> Result<(), Error> {
self.delete_admin_whitelist(ip).await
}
}

View File

@ -1,5 +1,6 @@
use std::fmt::Write;
use async_trait::async_trait;
use hmac::{Hmac, Mac};
use rusqlite::{Error as RusqliteError, params};
use sha2::Sha256;
@ -15,9 +16,6 @@ type HmacSha256 = Hmac<Sha256>;
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();
@ -29,101 +27,148 @@ impl Database {
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<Option<Claims>, Error> {
pub async fn validate_api_key(&self, api_key: &str) -> Result<Option<Claims>, 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],
self.pool
.conn_and_then(move |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)?,
))
},
);
let perm_level = PermissionLevel::from_str(&level).unwrap_or(PermissionLevel::ReadOnly);
let permissions: Vec<String> = perm_level.permissions().iter().map(|s| (*s).to_string()).collect();
match result {
Ok((id, name, level)) => {
let _ = conn.execute(
"UPDATE api_keys SET last_used_at = datetime('now') WHERE id = ?1",
params![id],
);
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)?,
}
}
let perm_level = PermissionLevel::from_str(&level).unwrap_or(PermissionLevel::ReadOnly);
let permissions: Vec<String> =
perm_level.permissions().iter().map(|s| (*s).to_string()).collect();
pub fn insert_api_key(&self, key_hash: &str, name: &str, permission_level: &str) -> Result<i64, Error> {
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())
}
pub fn list_api_keys(&self) -> Result<Vec<ApiKeyView>, 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(ApiKeyView {
id: row.get(0)?,
name: row.get(1)?,
permission_level: row.get(2)?,
created_at: row.get(3)?,
last_used_at: row.get(4)?,
Ok(Some(Claims {
sub: -id,
username: format!("api:{}", name),
role: level,
permissions,
exp: usize::MAX,
}))
}
Err(RusqliteError::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e.into()),
}
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
.await
}
pub fn delete_api_key(&self, id: i64) -> Result<bool, Error> {
let conn = self.conn()?;
let affected = conn.execute("DELETE FROM api_keys WHERE id = ?1", params![id])?;
Ok(affected > 0)
pub async fn insert_api_key(&self, key_hash: &str, name: &str, permission_level: &str) -> Result<i64, Error> {
let key_hash = key_hash.to_string();
let name = name.to_string();
let permission_level = permission_level.to_string();
self.pool
.conn_and_then(move |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())
})
.await
}
pub async fn list_api_keys(&self) -> Result<Vec<ApiKeyView>, Error> {
self.pool
.conn_and_then(move |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(ApiKeyView {
id: row.get(0)?,
name: row.get(1)?,
permission_level: row.get(2)?,
created_at: row.get(3)?,
last_used_at: row.get(4)?,
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
})
.await
}
pub async fn delete_api_key(&self, id: i64) -> Result<bool, Error> {
self.pool
.conn_and_then(move |conn| {
let affected = conn.execute("DELETE FROM api_keys WHERE id = ?1", params![id])?;
Ok(affected > 0)
})
.await
}
}
#[async_trait]
impl ApiKeyRepo for Database {
fn validate_api_key(&self, api_key: &str) -> Result<Option<Claims>, Error> {
self.validate_api_key(api_key)
async fn validate_api_key(&self, api_key: &str) -> Result<Option<Claims>, Error> {
self.validate_api_key(api_key).await
}
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<i64, Error> {
self.insert_api_key(key_hash, name, permission_level)
async fn insert_api_key(&self, key_hash: &str, name: &str, permission_level: &str) -> Result<i64, Error> {
self.insert_api_key(key_hash, name, permission_level).await
}
fn list_api_keys(&self) -> Result<Vec<ApiKeyView>, Error> {
self.list_api_keys()
async fn list_api_keys(&self) -> Result<Vec<ApiKeyView>, Error> {
self.list_api_keys().await
}
fn delete_api_key(&self, id: i64) -> Result<bool, Error> {
self.delete_api_key(id)
async fn delete_api_key(&self, id: i64) -> Result<bool, Error> {
self.delete_api_key(id).await
}
}
#[cfg(test)]
mod tests {
use super::Database;
#[tokio::test]
async fn validate_full_access_api_key_grants_admin_permissions() {
let db = Database::new(":memory:").await.expect("test db");
let raw_key = "ng-test-full-access";
let digest = db.hmac_api_key(raw_key);
db.insert_api_key(&digest, "automation", "full_access").await.unwrap();
let claims = db.validate_api_key(raw_key).await.unwrap().expect("claims");
assert!(claims.permissions.contains(&"api_keys:admin".to_string()));
assert!(claims.permissions.contains(&"system:admin".to_string()));
assert!(claims.permissions.contains(&"users:admin".to_string()));
}
#[tokio::test]
async fn validate_read_write_api_key_does_not_grant_admin_permissions() {
let db = Database::new(":memory:").await.expect("test db");
let raw_key = "ng-test-read-write";
let digest = db.hmac_api_key(raw_key);
db.insert_api_key(&digest, "automation", "read_write").await.unwrap();
let claims = db.validate_api_key(raw_key).await.unwrap().expect("claims");
assert!(!claims.permissions.contains(&"api_keys:admin".to_string()));
assert!(!claims.permissions.contains(&"system:admin".to_string()));
assert!(!claims.permissions.contains(&"users:admin".to_string()));
}
}

View File

@ -1,5 +1,6 @@
use std::fmt::Write;
use async_trait::async_trait;
use chrono::Utc;
use rusqlite::params;
use sha2::{Digest, Sha256};
@ -10,8 +11,6 @@ use crate::domain::common::error::Error;
use crate::domain::common::error::database::DatabaseError;
use crate::interface::audit::AuditRepo;
/// 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] {
@ -27,131 +26,135 @@ fn audit_row_hash(ts: &str, actor: &str, action: &str, detail: &str, prev_hash:
}
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)
pub async fn insert_audit_log(&self, actor: &str, action: &str, detail: &str) -> Result<(), Error> {
let actor = actor.to_string();
let action = action.to_string();
let detail = detail.to_string();
self.pool
.conn_mut_and_then(move |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(())
})
.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(())
.await
}
/// List recent audit log entries (most recent first, max 200).
pub fn list_audit_logs(&self) -> Result<Vec<AuditLogEntry>, 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)
pub async fn list_audit_logs(&self) -> Result<Vec<AuditLogEntry>, Error> {
self.pool
.conn_and_then(move |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)
})
.await
}
/// 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<Vec<AuditLogEntry>, 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)
pub async fn list_audit_logs_by_action(&self, action: &str, limit: i64) -> Result<Vec<AuditLogEntry>, Error> {
let action = action.to_string();
self.pool
.conn_and_then(move |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)
})
.await
}
/// Walk audit_log rows in id order, verifying the hash chain.
/// When `after_id` is 0 the entire table is checked; otherwise only
/// rows with `id > after_id` are verified (the prev_hash of the first
/// row is validated against the stored row_hash of `after_id`).
/// Returns `Ok((verified_count, last_id))` on success.
pub fn verify_audit_log_chain(&self, after_id: i64) -> Result<(usize, i64), Error> {
let conn = self.conn()?;
pub async fn verify_audit_log_chain(&self, after_id: i64) -> Result<(usize, i64), Error> {
self.pool
.conn_and_then(move |conn| {
let mut expected_prev = if after_id > 0 {
conn.query_row(
"SELECT row_hash FROM audit_log WHERE id = ?1",
params![after_id],
|row| row.get::<_, String>(0),
)
.unwrap_or_default()
} else {
String::new()
};
let mut expected_prev = if after_id > 0 {
conn.query_row(
"SELECT row_hash FROM audit_log WHERE id = ?1",
params![after_id],
|row| row.get::<_, String>(0),
)
.unwrap_or_default()
} else {
String::new()
};
let mut stmt = conn.prepare(
"SELECT id, ts, actor, action, detail, prev_hash, row_hash \
FROM audit_log WHERE id > ?1 ORDER BY id ASC",
)?;
let mut rows = stmt.query(params![after_id])?;
let mut stmt = conn.prepare(
"SELECT id, ts, actor, action, detail, prev_hash, row_hash \
FROM audit_log WHERE id > ?1 ORDER BY id ASC",
)?;
let mut rows = stmt.query(params![after_id])?;
let mut count = 0usize;
let mut last_id = after_id;
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)?;
let mut count = 0usize;
let mut last_id = after_id;
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;
last_id = id;
count += 1;
}
Ok((count, last_id))
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;
last_id = id;
count += 1;
}
Ok((count, last_id))
})
.await
}
}
#[async_trait]
impl AuditRepo for Database {
fn insert_audit_log(&self, actor: &str, action: &str, detail: &str) -> Result<(), Error> {
self.insert_audit_log(actor, action, detail)
async fn insert_audit_log(&self, actor: &str, action: &str, detail: &str) -> Result<(), Error> {
self.insert_audit_log(actor, action, detail).await
}
fn list_audit_logs_by_action(&self, action: &str, limit: i64) -> Result<Vec<AuditLogEntry>, Error> {
self.list_audit_logs_by_action(action, limit)
async fn list_audit_logs_by_action(&self, action: &str, limit: i64) -> Result<Vec<AuditLogEntry>, Error> {
self.list_audit_logs_by_action(action, limit).await
}
fn verify_audit_log_chain(&self, after_id: i64) -> Result<(usize, i64), Error> {
self.verify_audit_log_chain(after_id)
async fn verify_audit_log_chain(&self, after_id: i64) -> Result<(usize, i64), Error> {
self.verify_audit_log_chain(after_id).await
}
}

View File

@ -0,0 +1,153 @@
use async_trait::async_trait;
use rusqlite::{Error as RusqliteError, params};
use super::Database;
use crate::domain::common::error::Error;
use crate::interface::config_repo::ConfigRepo;
impl Database {
pub async fn get_config_value(&self, key: &str) -> Result<Option<String>, Error> {
let key = key.to_string();
self.pool
.conn_and_then(move |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.into()),
}
})
.await
}
pub async fn set_config_value(&self, key: &str, value: &str) -> Result<(), Error> {
let key = key.to_string();
let value = value.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute(
"INSERT OR REPLACE INTO settings (key, value) VALUES (?1, ?2)",
params![key, value],
)?;
Ok(())
})
.await
}
async fn get_app_secret(&self, key: &str) -> Result<Option<String>, Error> {
let key = key.to_string();
self.pool
.conn_and_then(move |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.into()),
}
})
.await
}
async fn set_app_secret(&self, key: &str, value: &str) -> Result<(), Error> {
let key = key.to_string();
let value = value.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute(
"INSERT OR REPLACE INTO app_secrets (key, value) VALUES (?1, ?2)",
params![key, value],
)?;
Ok(())
})
.await
}
pub async fn get_notification_config(&self, channel: &str) -> Result<Option<String>, Error> {
let channel = channel.to_string();
self.pool
.conn_and_then(move |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.into()),
}
})
.await
}
pub async fn set_notification_config(&self, channel: &str, config_json: &str) -> Result<(), Error> {
let channel = channel.to_string();
let config_json = config_json.to_string();
self.pool
.conn_and_then(move |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(())
})
.await
}
}
#[async_trait]
impl ConfigRepo for Database {
async fn get_config_value(&self, key: &str) -> Result<Option<String>, Error> {
self.get_config_value(key).await
}
async fn set_config_value(&self, key: &str, value: &str) -> Result<(), Error> {
self.set_config_value(key, value).await
}
async fn get_app_secret(&self, key: &str) -> Result<Option<String>, Error> {
self.get_app_secret(key).await
}
async fn set_app_secret(&self, key: &str, plaintext: &str) -> Result<(), Error> {
self.set_app_secret(key, plaintext).await
}
async fn get_notification_config(&self, channel: &str) -> Result<Option<String>, Error> {
self.get_notification_config(channel).await
}
async fn set_notification_config(&self, channel: &str, config_json: &str) -> Result<(), Error> {
self.set_notification_config(channel, config_json).await
}
async fn update_config_values_atomically(
&self,
config_values: Vec<(String, String)>,
secrets: Vec<(String, String)>,
) -> Result<(), Error> {
self.pool
.conn_mut_and_then(move |conn| {
let tx = conn.transaction()?;
for (key, value) in config_values {
tx.execute(
"INSERT OR REPLACE INTO settings (key, value) VALUES (?1, ?2)",
params![key, value],
)?;
}
for (key, value) in secrets {
tx.execute(
"INSERT OR REPLACE INTO app_secrets (key, value) VALUES (?1, ?2)",
params![key, value],
)?;
}
tx.commit()?;
Ok(())
})
.await
}
}

View File

@ -1,3 +1,4 @@
use async_trait::async_trait;
use rusqlite::params;
use super::Database;
@ -5,136 +6,138 @@ use crate::domain::common::error::Error;
use crate::interface::enforcement::EnforcementRepo;
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 async fn set_rate_limit(&self, key: &str, value: u64) -> Result<(), Error> {
let key = key.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute(
"INSERT OR REPLACE INTO rate_limit_config (key, value) VALUES (?1, ?2)",
params![key, value as i64],
)?;
Ok(())
})
.await
}
pub fn load_rate_limit_config(&self) -> Result<Vec<(String, u64)>, 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 async fn load_rate_limit_config(&self) -> Result<Vec<(String, u64)>, Error> {
self.pool
.conn_and_then(move |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)
})
.await
}
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 async fn insert_dns_domains(&self, domains: &[String]) -> Result<(), Error> {
let domains = domains.to_vec();
self.pool
.conn_mut_and_then(move |conn| {
let tx = conn.transaction()?;
for domain in domains {
tx.execute(
"INSERT OR IGNORE INTO dns_blacklist (domain) VALUES (?1)",
params![domain],
)?;
}
tx.commit()?;
Ok(())
})
.await
}
pub 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 async fn delete_dns_domains(&self, domains: &[String]) -> Result<(), Error> {
let domains = domains.to_vec();
self.pool
.conn_mut_and_then(move |conn| {
let tx = conn.transaction()?;
for domain in domains {
tx.execute("DELETE FROM dns_blacklist WHERE domain = ?1", params![domain])?;
}
tx.commit()?;
Ok(())
})
.await
}
pub fn load_dns_domains(&self) -> Result<Vec<String>, 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 async fn load_dns_domains(&self) -> Result<Vec<String>, Error> {
self.pool
.conn_and_then(move |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)
})
.await
}
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 async fn insert_geo_country(&self, code: &str) -> Result<(), Error> {
let code = code.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute(
"INSERT OR IGNORE INTO geo_blocked_countries (country_code) VALUES (?1)",
params![code],
)?;
Ok(())
})
.await
}
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 async fn delete_geo_country(&self, code: &str) -> Result<(), Error> {
let code = code.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute(
"DELETE FROM geo_blocked_countries WHERE country_code = ?1",
params![code],
)?;
Ok(())
})
.await
}
pub fn load_geo_countries(&self) -> Result<Vec<String>, 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)
pub async fn load_geo_countries(&self) -> Result<Vec<String>, Error> {
self.pool
.conn_and_then(move |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)
})
.await
}
}
#[async_trait]
impl EnforcementRepo for Database {
fn set_rate_limit(&self, key: &str, value: u64) -> Result<(), Error> {
self.set_rate_limit(key, value)
async fn set_rate_limit(&self, key: &str, value: u64) -> Result<(), Error> {
self.set_rate_limit(key, value).await
}
fn insert_dns_domain(&self, domain: &str) -> Result<(), Error> {
self.insert_dns_domain(domain)
async fn insert_dns_domains(&self, domains: &[String]) -> Result<(), Error> {
self.insert_dns_domains(domains).await
}
fn delete_dns_domain(&self, domain: &str) -> Result<(), Error> {
self.delete_dns_domain(domain)
async fn delete_dns_domains(&self, domains: &[String]) -> Result<(), Error> {
self.delete_dns_domains(domains).await
}
fn insert_geo_country(&self, code: &str) -> Result<(), Error> {
self.insert_geo_country(code)
async fn insert_geo_country(&self, code: &str) -> Result<(), Error> {
self.insert_geo_country(code).await
}
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));
async fn delete_geo_country(&self, code: &str) -> Result<(), Error> {
self.delete_geo_country(code).await
}
}

View File

@ -1,18 +1,19 @@
mod acl;
mod api_key;
mod audit;
mod config;
mod enforcement;
mod setting;
mod report_snapshot;
mod soar;
mod soar_block;
mod stats;
mod system_state;
mod user;
use std::env;
use async_sqlite::{Client, ClientBuilder};
use macros::log;
use r2d2::Pool;
use r2d2_sqlite::SqliteConnectionManager;
use rusqlite::{self, Connection, params};
use crate::domain::common::error::Error;
@ -30,6 +31,12 @@ impl From<rusqlite::Error> for Error {
Self::Database(DatabaseError::from(e))
}
}
impl From<async_sqlite::Error> for Error {
fn from(e: async_sqlite::Error) -> Self {
Self::Database(DatabaseError::QueryFailed(e))
}
}
use crate::domain::identity::auth::{ADMIN_PERMISSIONS, GROUP_ADMIN, GROUP_VIEWER, VIEWER_PERMISSIONS};
/// Reads the SQLCipher encryption key from the environment variable `NETGUARDIA_DB_KEY`.
@ -41,85 +48,51 @@ fn db_encryption_key() -> Option<String> {
}
}
/// 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<String>,
}
impl r2d2::CustomizeConnection<Connection, rusqlite::Error> 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<SqliteConnectionManager>,
pool: Client,
/// 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<Self, Error> {
pub async fn new(path: &str) -> Result<Self, Error> {
let encryption_key = db_encryption_key();
if path != ":memory:" && encryption_key.is_none() {
log!(MiscLog::DbEncryptionDisabled);
}
let manager = if path == ":memory:" {
SqliteConnectionManager::memory()
let builder = if path == ":memory:" {
ClientBuilder::new()
} else {
SqliteConnectionManager::file(path)
ClientBuilder::new().path(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)?;
let pool = builder.open().await.map_err(DatabaseError::QueryFailed)?;
let key_for_pragmas = encryption_key.clone();
pool.conn(move |conn| {
if let Some(ref key) = key_for_pragmas {
conn.pragma_update(None, "key", key)?;
}
conn.execute_batch(
"PRAGMA journal_mode=WAL; \
PRAGMA synchronous=NORMAL; \
PRAGMA busy_timeout=5000; \
PRAGMA foreign_keys=ON;",
)?;
Ok(())
})
.await
.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)?;
}
pool.conn(|conn| conn.query_row("SELECT count(*) FROM sqlite_master", [], |_| Ok(())))
.await
.map_err(|_| DatabaseError::EncryptionKeyInvalid)?;
let api_key_hmac = Self::derive_api_key_hmac();
let db = Self { pool, api_key_hmac };
db.create_tables()?;
db.create_tables().await?;
Ok(db)
}
@ -177,14 +150,11 @@ impl Database {
Ok(())
}
fn conn(&self) -> Result<r2d2::PooledConnection<SqliteConnectionManager>, Error> {
self.pool.get().map_err(|e| DatabaseError::QueryFailed(e).into())
}
fn create_tables(&self) -> Result<(), Error> {
let conn = self.conn()?;
conn.execute_batch(
"
async fn create_tables(&self) -> Result<(), Error> {
self.pool
.conn_mut_and_then(|conn| {
conn.execute_batch(
"
CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username TEXT UNIQUE NOT NULL,
@ -216,6 +186,19 @@ impl Database {
key TEXT PRIMARY KEY,
value TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS system_state (
key TEXT PRIMARY KEY,
value TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS report_snapshots (
key TEXT PRIMARY KEY,
value TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS login_attempts (
username TEXT PRIMARY KEY,
failure_count INTEGER NOT NULL DEFAULT 0,
locked_until INTEGER
);
CREATE TABLE IF NOT EXISTS user_groups (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT UNIQUE NOT NULL,
@ -267,7 +250,9 @@ impl Database {
playbook_id INTEGER NOT NULL,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
expires_at TEXT NOT NULL,
unblocked_at TEXT
unblocked_at TEXT,
created_acl_rule INTEGER NOT NULL DEFAULT 1,
preserve_acl_on_unblock INTEGER NOT NULL DEFAULT 0
);
CREATE TABLE IF NOT EXISTS soar_executions (
id INTEGER PRIMARY KEY AUTOINCREMENT,
@ -332,6 +317,32 @@ impl Database {
CREATE INDEX IF NOT EXISTS idx_audit_log_action
ON audit_log(action, id DESC);
CREATE INDEX IF NOT EXISTS idx_soar_block_active_expires
ON soar_block_rules(expires_at)
WHERE unblocked_at IS NULL;
CREATE INDEX IF NOT EXISTS idx_soar_block_created_at
ON soar_block_rules(created_at);
CREATE INDEX IF NOT EXISTS idx_soar_block_unblocked_at
ON soar_block_rules(unblocked_at)
WHERE unblocked_at IS NOT NULL;
CREATE INDEX IF NOT EXISTS idx_soar_block_source_created
ON soar_block_rules(source_ip, created_at);
CREATE INDEX IF NOT EXISTS idx_soar_executions_executed_at
ON soar_executions(executed_at DESC);
CREATE INDEX IF NOT EXISTS idx_soar_executions_trigger_executed
ON soar_executions(trigger_event, executed_at);
CREATE INDEX IF NOT EXISTS idx_user_group_members_group_id
ON user_group_members(group_id, user_id);
CREATE INDEX IF NOT EXISTS idx_api_keys_key_hash
ON api_keys(key_hash);
CREATE TRIGGER IF NOT EXISTS audit_log_no_update
BEFORE UPDATE ON audit_log BEGIN
SELECT RAISE(ABORT, 'audit_log is append-only (WORM)');
@ -342,26 +353,109 @@ impl Database {
SELECT RAISE(ABORT, 'audit_log is append-only (WORM)');
END;
",
)?;
)?;
let conn_ref = &*conn;
if !Self::column_exists(conn, "soar_block_rules", "created_acl_rule")? {
conn.execute(
"ALTER TABLE soar_block_rules ADD COLUMN created_acl_rule INTEGER NOT NULL DEFAULT 1",
[],
)?;
}
if !Self::column_exists(conn, "soar_block_rules", "preserve_acl_on_unblock")? {
conn.execute(
"ALTER TABLE soar_block_rules ADD COLUMN preserve_acl_on_unblock INTEGER NOT NULL DEFAULT 0",
[],
)?;
}
// 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::to_string(&ADMIN_PERMISSIONS).unwrap_or_else(|_| "[]".to_string());
let viewer_permissions = serde_json::to_string(&VIEWER_PERMISSIONS).unwrap_or_else(|_| "[]".to_string());
Self::migrate_legacy_settings_state(conn)?;
conn_ref.execute(
"INSERT INTO user_groups (name, description, permissions) VALUES (?1, ?2, ?3)",
params![GROUP_ADMIN, "Full system access with all permissions", &all_permissions],
)?;
conn_ref.execute(
"INSERT INTO user_groups (name, description, permissions) VALUES (?1, ?2, ?3)",
params![GROUP_VIEWER, "Read-only access to all modules", &viewer_permissions],
)?;
let group_count: i64 = conn.query_row("SELECT COUNT(*) FROM user_groups", [], |row| row.get(0))?;
if group_count == 0 {
let all_permissions =
serde_json::to_string(&ADMIN_PERMISSIONS).unwrap_or_else(|_| "[]".to_string());
let viewer_permissions =
serde_json::to_string(&VIEWER_PERMISSIONS).unwrap_or_else(|_| "[]".to_string());
conn.execute(
"INSERT INTO user_groups (name, description, permissions) VALUES (?1, ?2, ?3)",
params![GROUP_ADMIN, "Full system access with all permissions", &all_permissions],
)?;
conn.execute(
"INSERT INTO user_groups (name, description, permissions) VALUES (?1, ?2, ?3)",
params![GROUP_VIEWER, "Read-only access to all modules", &viewer_permissions],
)?;
}
Ok::<(), Error>(())
})
.await
}
fn column_exists(conn: &Connection, table: &str, column: &str) -> Result<bool, Error> {
let mut stmt = conn.prepare(&format!("PRAGMA table_info({})", table))?;
let rows = stmt.query_map([], |row| row.get::<_, String>(1))?;
for row in rows {
if row? == column {
return Ok(true);
}
}
Ok(false)
}
fn migrate_legacy_settings_state(conn: &Connection) -> Result<(), Error> {
conn.execute_batch(
"
INSERT OR IGNORE INTO system_state (key, value)
SELECT key, value FROM settings WHERE key = 'setup_complete';
INSERT OR IGNORE INTO report_snapshots (key, value)
SELECT key, value FROM settings
WHERE key IN (
'weekly_threats_count',
'weekly_top_ips',
'weekly_threat_breakdown',
'weekly_system_health',
'system_uptime_percent',
'active_rules_count',
'weekly_geo_distribution',
'weekly_soar_blocks',
'weekly_soar_triggers',
'weekly_soar_unblocks',
'weekly_blocked_count'
);
INSERT OR IGNORE INTO login_attempts (username, failure_count)
SELECT substr(key, length('login_failures:') + 1), CAST(value AS INTEGER)
FROM settings
WHERE key LIKE 'login_failures:%';
INSERT INTO login_attempts (username, locked_until)
SELECT substr(key, length('login_locked_until:') + 1), CAST(value AS INTEGER)
FROM settings
WHERE key LIKE 'login_locked_until:%'
ON CONFLICT(username) DO UPDATE SET
locked_until = COALESCE(login_attempts.locked_until, excluded.locked_until);
DELETE FROM settings
WHERE key = 'setup_complete'
OR key IN (
'weekly_threats_count',
'weekly_top_ips',
'weekly_threat_breakdown',
'weekly_system_health',
'system_uptime_percent',
'active_rules_count',
'weekly_geo_distribution',
'weekly_soar_blocks',
'weekly_soar_triggers',
'weekly_soar_unblocks',
'weekly_blocked_count'
)
OR key LIKE 'login_failures:%'
OR key LIKE 'login_locked_until:%';
",
)?;
Ok(())
}
}
@ -369,42 +463,141 @@ impl Database {
#[cfg(test)]
mod tests {
use super::*;
use crate::interface::acl::AclRepo;
use crate::interface::identity::UserRepo;
use crate::interface::setting::SettingRepo;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) fn test_db() -> Database {
Database::new(":memory:").expect("Failed to create test database")
use crate::interface::acl::AclRepo;
use crate::interface::config_repo::ConfigRepo;
use crate::interface::identity::UserRepo;
pub(super) async fn test_db() -> Database {
Database::new(":memory:").await.expect("Failed to create test database")
}
#[test]
fn test_create_tables() {
let _db = test_db();
#[tokio::test]
async fn test_create_tables() {
let _db = test_db().await;
}
/// Verify that Database satisfies each aggregate Repo trait contract
/// (AclRepo / SettingRepo / UserRepo). Exercises the trait-object
/// (AclRepo / ConfigRepo / UserRepo). Exercises the trait-object
/// path so callers that take `Arc<dyn XxxRepo>` compile end-to-end.
#[test]
fn test_aggregate_repo_trait_objects() {
let db = test_db();
#[tokio::test]
async fn test_aggregate_repo_trait_objects() {
let db = test_db().await;
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 setting: &dyn ConfigRepo = &db;
setting.set_config_value("test_key", "test_value").await.unwrap();
assert_eq!(
setting.get_config_value("test_key").await.unwrap(),
Some("test_value".to_string())
);
let acl: &dyn AclRepo = &db;
acl.insert_acl_rule(4, "source", "blacklist", "10.0.0.1", 443).unwrap();
acl.insert_acl_rule(4, "source", "blacklist", "10.0.0.1", 443)
.await
.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.list_acl_rules().unwrap();
let rules = db.list_acl_rules().await.unwrap();
assert_eq!(rules.len(), 1);
let identity: &dyn UserRepo = &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);
assert_eq!(db.user_count().await.unwrap(), 0);
identity.insert_user("test", "hash", "viewer", false).await.unwrap();
assert_eq!(db.user_count().await.unwrap(), 1);
}
#[tokio::test]
async fn legacy_non_config_settings_are_backfilled() {
let unique = SystemTime::now().duration_since(UNIX_EPOCH).unwrap().as_nanos();
let path = std::env::temp_dir().join(format!("netguardia-legacy-settings-{}-{unique}.db", std::process::id()));
let locked_until = SystemTime::now().duration_since(UNIX_EPOCH).unwrap().as_secs() + 3600;
{
let conn = Connection::open(&path).unwrap();
conn.execute_batch(
"
CREATE TABLE settings (
key TEXT PRIMARY KEY,
value TEXT NOT NULL
);
",
)
.unwrap();
conn.execute(
"INSERT INTO settings (key, value) VALUES (?1, ?2)",
params!["setup_complete", "true"],
)
.unwrap();
conn.execute(
"INSERT INTO settings (key, value) VALUES (?1, ?2)",
params!["weekly_threats_count", "7"],
)
.unwrap();
conn.execute(
"INSERT INTO settings (key, value) VALUES (?1, ?2)",
params!["login_failures:alice", "4"],
)
.unwrap();
conn.execute(
"INSERT INTO settings (key, value) VALUES (?1, ?2)",
params!["login_locked_until:alice", locked_until.to_string()],
)
.unwrap();
}
let db = Database::new(path.to_str().unwrap()).await.unwrap();
assert_eq!(
db.get_system_state("setup_complete").await.unwrap(),
Some("true".to_string())
);
assert_eq!(
db.get_report_snapshot("weekly_threats_count").await.unwrap(),
Some("7".to_string())
);
assert_eq!(db.get_config_value("setup_complete").await.unwrap(), None);
assert_eq!(db.get_config_value("weekly_threats_count").await.unwrap(), None);
let remaining_lock = db.check_login_locked("alice").await.unwrap();
assert!(remaining_lock.is_some_and(|remaining| remaining > 0));
assert_eq!(db.get_config_value("login_failures:alice").await.unwrap(), None);
drop(db);
let _ = std::fs::remove_file(&path);
let _ = std::fs::remove_file(path.with_extension("db-wal"));
let _ = std::fs::remove_file(path.with_extension("db-shm"));
}
#[tokio::test]
async fn create_tables_installs_query_shape_indexes() {
let db = test_db().await;
let expected = [
("soar_block_rules", "idx_soar_block_active_expires"),
("soar_block_rules", "idx_soar_block_created_at"),
("soar_block_rules", "idx_soar_block_unblocked_at"),
("soar_block_rules", "idx_soar_block_source_created"),
("soar_executions", "idx_soar_executions_executed_at"),
("soar_executions", "idx_soar_executions_trigger_executed"),
("user_group_members", "idx_user_group_members_group_id"),
("api_keys", "idx_api_keys_key_hash"),
];
db.pool
.conn_and_then(move |conn| {
for (table, index) in expected {
let count: i64 = conn.query_row(
"SELECT COUNT(*) FROM sqlite_master WHERE type = 'index' AND tbl_name = ?1 AND name = ?2",
params![table, index],
|row| row.get(0),
)?;
assert_eq!(count, 1, "missing index {index} on {table}");
}
Ok::<(), Error>(())
})
.await
.unwrap();
}
}

View File

@ -0,0 +1,51 @@
use async_trait::async_trait;
use rusqlite::{Error as RusqliteError, params};
use super::Database;
use crate::domain::common::error::Error;
use crate::interface::report_snapshot::ReportSnapshotRepo;
impl Database {
pub async fn get_report_snapshot(&self, key: &str) -> Result<Option<String>, Error> {
let key = key.to_string();
self.pool
.conn_and_then(move |conn| {
let result = conn.query_row(
"SELECT value FROM report_snapshots WHERE key = ?1",
params![key],
|row| row.get(0),
);
match result {
Ok(value) => Ok(Some(value)),
Err(RusqliteError::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e.into()),
}
})
.await
}
pub async fn set_report_snapshot(&self, key: &str, value: &str) -> Result<(), Error> {
let key = key.to_string();
let value = value.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute(
"INSERT OR REPLACE INTO report_snapshots (key, value) VALUES (?1, ?2)",
params![key, value],
)?;
Ok(())
})
.await
}
}
#[async_trait]
impl ReportSnapshotRepo for Database {
async fn get_report_snapshot(&self, key: &str) -> Result<Option<String>, Error> {
self.get_report_snapshot(key).await
}
async fn set_report_snapshot(&self, key: &str, value: &str) -> Result<(), Error> {
self.set_report_snapshot(key, value).await
}
}

View File

@ -1,261 +0,0 @@
use rusqlite::{Error as RusqliteError, Transaction, params};
use super::Database;
use crate::domain::common::error::Error;
use crate::interface::setting::SettingRepo;
impl Database {
pub fn get_setting(&self, key: &str) -> Result<Option<String>, 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(())
}
fn get_app_secret(&self, key: &str) -> Result<Option<String>, 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)?,
}
}
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<Option<String>, 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<Option<String>, 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<Option<String>, 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<Option<String>, 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)
}
fn transaction(&self, f: &mut dyn FnMut(&dyn SettingRepo) -> Result<(), Error>) -> Result<(), Error> {
let mut conn = self.conn()?;
let tx = conn.transaction()?;
let view = SettingTxView { tx: &tx };
match f(&view) {
Ok(()) => {
tx.commit()?;
Ok(())
}
Err(e) => {
let _ = tx.rollback();
Err(e)
}
}
}
}
struct SettingTxView<'a> {
tx: &'a Transaction<'a>,
}
impl SettingRepo for SettingTxView<'_> {
fn transaction(&self, f: &mut dyn FnMut(&dyn SettingRepo) -> Result<(), Error>) -> Result<(), Error> {
f(self)
}
fn get_setting(&self, key: &str) -> Result<Option<String>, Error> {
let result = self
.tx
.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)?,
}
}
fn set_setting(&self, key: &str, value: &str) -> Result<(), Error> {
self.tx.execute(
"INSERT OR REPLACE INTO settings (key, value) VALUES (?1, ?2)",
params![key, value],
)?;
Ok(())
}
fn get_app_secret(&self, key: &str) -> Result<Option<String>, Error> {
let result = self
.tx
.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)?,
}
}
fn set_app_secret(&self, key: &str, plaintext: &str) -> Result<(), Error> {
self.tx.execute(
"INSERT OR REPLACE INTO app_secrets (key, value) VALUES (?1, ?2)",
params![key, plaintext],
)?;
Ok(())
}
fn get_notification_config(&self, channel: &str) -> Result<Option<String>, Error> {
match self.tx.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)?,
}
}
fn set_notification_config(&self, channel: &str, config_json: &str) -> Result<(), Error> {
self.tx.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(())
}
}
#[cfg(test)]
mod tests {
use super::super::tests::test_db;
use crate::domain::common::error::Error;
use crate::domain::common::error::misc::MiscError;
use crate::interface::setting::SettingRepo;
#[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 tx_commits_on_ok() {
let db = test_db();
db.transaction(&mut |repo| {
repo.set_setting("a", "1")?;
repo.set_setting("b", "2")?;
Ok(())
})
.unwrap();
assert_eq!(db.get_setting("a").unwrap().as_deref(), Some("1"));
assert_eq!(db.get_setting("b").unwrap().as_deref(), Some("2"));
}
#[test]
fn tx_rolls_back_on_err() {
let db = test_db();
db.set_setting("a", "initial").unwrap();
let result: Result<(), Error> = db.transaction(&mut |repo| {
repo.set_setting("a", "changed")?;
repo.set_setting("b", "new")?;
Err(MiscError::ValidationError("bail".to_string()))?
});
assert!(result.is_err());
assert_eq!(db.get_setting("a").unwrap().as_deref(), Some("initial"));
assert_eq!(db.get_setting("b").unwrap(), None);
}
#[test]
fn tx_reads_see_pending_writes() {
let db = test_db();
db.transaction(&mut |repo| {
repo.set_setting("k", "v")?;
assert_eq!(repo.get_setting("k")?.as_deref(), Some("v"));
Ok(())
})
.unwrap();
}
#[test]
fn tx_covers_app_secrets_table() {
let db = test_db();
db.transaction(&mut |repo| {
repo.set_app_secret("smtp_password", "envelope_blob")?;
repo.set_setting("smtp_password", "")?;
Ok(())
})
.unwrap();
assert_eq!(
db.get_app_secret("smtp_password").unwrap().as_deref(),
Some("envelope_blob")
);
assert_eq!(db.get_setting("smtp_password").unwrap().as_deref(), Some(""));
}
}

View File

@ -1,5 +1,6 @@
use std::collections::HashMap;
use async_trait::async_trait;
use rusqlite::params;
use serde_json::Value;
@ -12,8 +13,6 @@ use crate::domain::response::playbook_data::{
};
use crate::interface::soar::SoarRepo;
/// Intermediate row from the playbooks LEFT JOIN playbook_actions query.
/// Private to this module; consumed only by `SoarRepo::list_playbooks`.
struct PlaybookActionRow {
pb_id: i64,
name: String,
@ -30,7 +29,7 @@ struct PlaybookActionRow {
}
impl Database {
pub fn insert_playbook(
pub async fn insert_playbook(
&self,
name: &str,
trigger_event: &str,
@ -39,79 +38,97 @@ impl Database {
window: Option<i64>,
cooldown: i64,
) -> Result<i64, Error> {
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())
let name = name.to_string();
let trigger_event = trigger_event.to_string();
self.pool
.conn_and_then(move |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())
})
.await
}
pub fn insert_playbook_action(
pub async fn insert_playbook_action(
&self,
playbook_id: i64,
action_order: i64,
action_type: &str,
params_json: &str,
) -> Result<i64, Error> {
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).
fn list_playbooks_with_actions(&self) -> Result<Vec<PlaybookActionRow>, 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(PlaybookActionRow {
pb_id: row.get::<_, i64>(0)?,
name: row.get::<_, String>(1)?,
enabled: row.get::<_, bool>(2)?,
trigger_event: row.get::<_, String>(3)?,
condition_threshold: row.get::<_, Option<f64>>(4)?,
condition_count: row.get::<_, Option<i64>>(5)?,
condition_window_secs: row.get::<_, Option<i64>>(6)?,
cooldown_secs: row.get::<_, i64>(7)?,
action_id: row.get::<_, Option<i64>>(8)?,
action_order: row.get::<_, Option<i64>>(9)?,
action_type: row.get::<_, Option<String>>(10)?,
action_params: row.get::<_, Option<String>>(11)?,
let action_type = action_type.to_string();
let params_json = params_json.to_string();
self.pool
.conn_and_then(move |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())
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
.await
}
pub fn update_playbook_enabled(&self, id: i64, enabled: bool) -> Result<bool, Error> {
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)
async fn list_playbooks_with_actions(&self) -> Result<Vec<PlaybookActionRow>, Error> {
self.pool
.conn_and_then(move |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(PlaybookActionRow {
pb_id: row.get::<_, i64>(0)?,
name: row.get::<_, String>(1)?,
enabled: row.get::<_, bool>(2)?,
trigger_event: row.get::<_, String>(3)?,
condition_threshold: row.get::<_, Option<f64>>(4)?,
condition_count: row.get::<_, Option<i64>>(5)?,
condition_window_secs: row.get::<_, Option<i64>>(6)?,
cooldown_secs: row.get::<_, i64>(7)?,
action_id: row.get::<_, Option<i64>>(8)?,
action_order: row.get::<_, Option<i64>>(9)?,
action_type: row.get::<_, Option<String>>(10)?,
action_params: row.get::<_, Option<String>>(11)?,
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
})
.await
}
pub fn delete_playbook(&self, id: i64) -> Result<bool, Error> {
let conn = self.conn()?;
let rows = conn.execute("DELETE FROM playbooks WHERE id = ?1", params![id])?;
Ok(rows > 0)
pub async fn update_playbook_enabled(&self, id: i64, enabled: bool) -> Result<bool, Error> {
self.pool
.conn_and_then(move |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)
})
.await
}
pub fn insert_playbook_condition(
pub async fn delete_playbook(&self, id: i64) -> Result<bool, Error> {
self.pool
.conn_and_then(move |conn| {
let rows = conn.execute("DELETE FROM playbooks WHERE id = ?1", params![id])?;
Ok(rows > 0)
})
.await
}
pub async fn insert_playbook_condition(
&self,
playbook_id: i64,
condition_type: &str,
@ -119,132 +136,139 @@ impl Database {
value: &str,
value2: Option<&str>,
) -> Result<i64, Error> {
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())
let condition_type = condition_type.to_string();
let operator = operator.to_string();
let value = value.to_string();
let value2 = value2.map(str::to_string);
self.pool
.conn_and_then(move |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())
})
.await
}
fn list_all_playbook_conditions(&self) -> Result<Vec<(i64, ConditionView)>, 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>(1)?,
ConditionView {
id: row.get::<_, i64>(0)?,
condition_type: row.get::<_, String>(2)?,
operator: row.get::<_, String>(3)?,
value: row.get::<_, String>(4)?,
value2: row.get::<_, Option<String>>(5)?,
},
))
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
async fn list_all_playbook_conditions(&self) -> Result<Vec<(i64, ConditionView)>, Error> {
self.pool
.conn_and_then(move |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>(1)?,
ConditionView {
id: row.get::<_, i64>(0)?,
condition_type: row.get::<_, String>(2)?,
operator: row.get::<_, String>(3)?,
value: row.get::<_, String>(4)?,
value2: row.get::<_, Option<String>>(5)?,
},
))
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
})
.await
}
pub fn insert_soar_execution(
pub async fn insert_soar_execution(
&self,
playbook_id: i64,
source_ip: Option<&str>,
trigger_event: &str,
actions_json: &str,
) -> Result<i64, Error> {
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 list_soar_executions(&self, limit: i64) -> Result<Vec<ExecutionView>, 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| {
let actions_str: String = row.get(4)?;
Ok(ExecutionView {
id: row.get::<_, i64>(0)?,
playbook_id: row.get::<_, i64>(1)?,
source_ip: row.get::<_, Option<String>>(2)?,
trigger_event: row.get::<_, String>(3)?,
actions_executed: serde_json::from_str(&actions_str).unwrap_or(Value::Null),
created_at: row.get::<_, String>(5)?,
let source_ip = source_ip.map(str::to_string);
let trigger_event = trigger_event.to_string();
let actions_json = actions_json.to_string();
self.pool
.conn_and_then(move |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())
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
.await
}
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))?;
pub async fn list_soar_executions(&self, limit: i64) -> Result<Vec<ExecutionView>, Error> {
self.pool
.conn_and_then(move |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| {
let actions_str: String = row.get(4)?;
Ok(ExecutionView {
id: row.get::<_, i64>(0)?,
playbook_id: row.get::<_, i64>(1)?,
source_ip: row.get::<_, Option<String>>(2)?,
trigger_event: row.get::<_, String>(3)?,
actions_executed: serde_json::from_str(&actions_str).unwrap_or(Value::Null),
created_at: row.get::<_, String>(5)?,
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
})
.await
}
pub async fn seed_default_playbooks(&self) -> Result<(), Error> {
let count: i64 = self
.pool
.conn_and_then(move |conn| {
Ok::<i64, Error>(conn.query_row("SELECT COUNT(*) FROM playbooks", [], |row| row.get(0))?)
})
.await?;
if count > 0 {
return Ok(());
}
drop(conn);
for def in DEFAULT_PLAYBOOKS {
let pb_id = self.insert_playbook(
def.name,
def.trigger_event,
def.threshold,
def.count,
def.window,
def.cooldown,
)?;
let pb_id = self
.insert_playbook(
def.name,
def.trigger_event,
def.threshold,
def.count,
def.window,
def.cooldown,
)
.await?;
for action in def.actions {
self.insert_playbook_action(pb_id, action.order, action.action_type, action.params)?;
self.insert_playbook_action(pb_id, action.order, action.action_type, action.params)
.await?;
}
for cond in def.conditions {
self.insert_playbook_condition(pb_id, cond.condition_type, cond.operator, cond.value, cond.value2)?;
self.insert_playbook_condition(pb_id, cond.condition_type, cond.operator, cond.value, cond.value2)
.await?;
}
}
Ok(())
}
}
#[async_trait]
impl SoarRepo for Database {
fn list_playbooks(&self) -> Result<Vec<PlaybookView>, Error> {
let rows = self.list_playbooks_with_actions()?;
async fn list_playbooks(&self) -> Result<Vec<PlaybookView>, Error> {
let rows = self.list_playbooks_with_actions().await?;
let mut result: Vec<PlaybookView> = Vec::new();
for row in rows {
let pb = if let Some(last) = result.last_mut() {
if last.id == row.pb_id {
last
} else {
result.push(PlaybookView {
id: row.pb_id,
name: row.name,
enabled: row.enabled,
trigger_event: row.trigger_event,
condition_threshold: row.condition_threshold,
condition_count: row.condition_count,
condition_window_secs: row.condition_window_secs,
cooldown_secs: row.cooldown_secs,
actions: Vec::new(),
conditions: Vec::new(),
});
// SAFETY: just pushed above, Vec cannot be empty
result.last_mut().unwrap_or_else(|| unreachable!())
}
let pb = if let Some(position) = result.iter().position(|p| p.id == row.pb_id) {
&mut result[position]
} else {
result.push(PlaybookView {
id: row.pb_id,
@ -258,10 +282,8 @@ impl SoarRepo for Database {
actions: Vec::new(),
conditions: Vec::new(),
});
// SAFETY: just pushed above, Vec cannot be empty
result.last_mut().unwrap_or_else(|| unreachable!())
};
if let (Some(aid), Some(order), Some(atype), Some(params_str)) =
(row.action_id, row.action_order, row.action_type, row.action_params)
{
@ -274,8 +296,7 @@ impl SoarRepo for Database {
}
}
// Load conditions and attach to playbooks
let cond_rows = self.list_all_playbook_conditions()?;
let cond_rows = self.list_all_playbook_conditions().await?;
let mut cond_map: HashMap<i64, Vec<ConditionView>> = HashMap::new();
for (pb_id, cond) in cond_rows {
cond_map.entry(pb_id).or_default().push(cond);
@ -285,59 +306,58 @@ impl SoarRepo for Database {
pb.conditions = conds;
}
}
Ok(result)
}
fn update_playbook_enabled(&self, id: i64, enabled: bool) -> Result<bool, Error> {
self.update_playbook_enabled(id, enabled)
async fn update_playbook_enabled(&self, id: i64, enabled: bool) -> Result<bool, Error> {
self.update_playbook_enabled(id, enabled).await
}
fn delete_playbook(&self, id: i64) -> Result<bool, Error> {
self.delete_playbook(id)
async fn delete_playbook(&self, id: i64) -> Result<bool, Error> {
self.delete_playbook(id).await
}
fn seed_default_playbooks(&self) -> Result<(), Error> {
self.seed_default_playbooks()
async fn seed_default_playbooks(&self) -> Result<(), Error> {
self.seed_default_playbooks().await
}
fn count_active_soar_blocks(&self) -> Result<u32, Error> {
self.count_active_soar_blocks()
async fn count_active_soar_blocks(&self) -> Result<u32, Error> {
self.count_active_soar_blocks().await
}
fn list_active_soar_blocks(&self) -> Result<Vec<ActiveBlockView>, Error> {
self.list_active_soar_blocks()
async fn list_active_soar_blocks(&self) -> Result<Vec<ActiveBlockView>, Error> {
self.list_active_soar_blocks().await
}
fn find_soar_block_by_id(&self, id: i64) -> Result<Option<ActiveBlockView>, Error> {
self.find_soar_block_by_id(id)
async fn find_soar_block_by_id(&self, id: i64) -> Result<Option<ActiveBlockView>, Error> {
self.find_soar_block_by_id(id).await
}
fn list_expired_soar_blocks(&self) -> Result<Vec<ActiveBlockView>, Error> {
self.list_expired_soar_blocks()
async fn list_expired_soar_blocks(&self) -> Result<Vec<ActiveBlockView>, Error> {
self.list_expired_soar_blocks().await
}
fn mark_soar_block_unblocked(&self, id: i64) -> Result<(), Error> {
self.mark_soar_block_unblocked(id)
async fn mark_soar_block_unblocked(&self, id: i64) -> Result<(), Error> {
self.mark_soar_block_unblocked(id).await
}
fn insert_pending_unblock(&self, source_ip: &str) -> Result<i64, Error> {
self.insert_pending_unblock(source_ip)
async fn insert_pending_unblock(&self, source_ip: &str) -> Result<i64, Error> {
self.insert_pending_unblock(source_ip).await
}
fn list_pending_unblocks(&self) -> Result<Vec<PendingUnblock>, Error> {
self.list_pending_unblocks()
async fn list_pending_unblocks(&self) -> Result<Vec<PendingUnblock>, Error> {
self.list_pending_unblocks().await
}
fn delete_pending_unblock(&self, id: i64) -> Result<(), Error> {
self.delete_pending_unblock(id)
async fn delete_pending_unblock(&self, id: i64) -> Result<(), Error> {
self.delete_pending_unblock(id).await
}
fn increment_pending_unblock_retry(&self, id: i64) -> Result<(), Error> {
self.increment_pending_unblock_retry(id)
async fn increment_pending_unblock_retry(&self, id: i64) -> Result<(), Error> {
self.increment_pending_unblock_retry(id).await
}
fn insert_soar_execution(
async fn insert_soar_execution(
&self,
playbook_id: i64,
source_ip: Option<&str>,
@ -345,124 +365,95 @@ impl SoarRepo for Database {
actions_json: &str,
) -> Result<i64, Error> {
self.insert_soar_execution(playbook_id, source_ip, trigger_event, actions_json)
.await
}
fn list_soar_executions(&self, limit: i64) -> Result<Vec<ExecutionView>, Error> {
self.list_soar_executions(limit)
async fn list_soar_executions(&self, limit: i64) -> Result<Vec<ExecutionView>, Error> {
self.list_soar_executions(limit).await
}
fn insert_playbook_atomic(
async fn insert_playbook_atomic(
&self,
input: &CreatePlaybookInput,
actions: &[ActionInput],
conditions: &[CreateConditionInput],
) -> Result<i64, Error> {
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![input.name, input.trigger_event, input.condition_threshold, input.condition_count, input.condition_window_secs, input.cooldown_secs],
)?;
let playbook_id = tx.last_insert_rowid();
for a in actions {
tx.execute(
"INSERT INTO playbook_actions (playbook_id, action_order, action_type, params) VALUES (?1, ?2, ?3, ?4)",
params![playbook_id, a.action_order, a.action_type, a.params_json],
)?;
}
for c in conditions {
tx.execute(
"INSERT INTO playbook_conditions (playbook_id, condition_type, operator, value, value2) VALUES (?1, ?2, ?3, ?4, ?5)",
params![playbook_id, c.condition_type, c.operator, c.value, c.value2.as_deref()],
)?;
}
tx.commit()?;
Ok(playbook_id)
let input = input.clone();
let actions = actions.to_vec();
let conditions = conditions.to_vec();
self.pool
.conn_mut_and_then(move |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![input.name, input.trigger_event, input.condition_threshold, input.condition_count, input.condition_window_secs, input.cooldown_secs],
)?;
let playbook_id = tx.last_insert_rowid();
for a in actions {
tx.execute(
"INSERT INTO playbook_actions (playbook_id, action_order, action_type, params) VALUES (?1, ?2, ?3, ?4)",
params![playbook_id, a.action_order, a.action_type, a.params_json],
)?;
}
for c in conditions {
tx.execute(
"INSERT INTO playbook_conditions (playbook_id, condition_type, operator, value, value2) VALUES (?1, ?2, ?3, ?4, ?5)",
params![playbook_id, c.condition_type, c.operator, c.value, c.value2.as_deref()],
)?;
}
tx.commit()?;
Ok(playbook_id)
})
.await
}
fn update_playbook_atomic(
async fn update_playbook_atomic(
&self,
id: i64,
row: &UpdatePlaybookInput,
actions: &[ActionInput],
conditions: &[CreateConditionInput],
) -> Result<bool, Error> {
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 a in actions {
tx.execute(
"INSERT INTO playbook_actions (playbook_id, action_order, action_type, params) VALUES (?1, ?2, ?3, ?4)",
params![id, a.action_order, a.action_type, a.params_json],
)?;
}
for c in conditions {
tx.execute(
"INSERT INTO playbook_conditions (playbook_id, condition_type, operator, value, value2) VALUES (?1, ?2, ?3, ?4, ?5)",
params![id, c.condition_type, c.operator, c.value, c.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::soar::SoarRepo;
let db = test_db();
use crate::domain::response::playbook_data::{ActionInput, CreateConditionInput, CreatePlaybookInput};
let input = CreatePlaybookInput {
name: "atom_pb".to_string(),
trigger_event: "threat".to_string(),
condition_threshold: Some(0.8),
condition_count: None,
condition_window_secs: None,
cooldown_secs: 300,
actions: vec![("block_ip".to_string(), "{}".to_string())],
conditions: vec![],
};
let actions = vec![ActionInput {
action_order: 1,
action_type: "block_ip".to_string(),
params_json: "{}".to_string(),
}];
let conditions = vec![CreateConditionInput::new(
"threshold".to_string(),
Some(">=".to_string()),
"0.8".to_string(),
None,
)];
let id = db.insert_playbook_atomic(&input, &actions, &conditions).unwrap();
assert!(id > 0);
let loaded = db.list_playbooks().unwrap();
assert!(!loaded.is_empty());
assert_eq!(loaded[0].conditions.len(), 1);
let row = row.clone();
let actions = actions.to_vec();
let conditions = conditions.to_vec();
self.pool
.conn_mut_and_then(move |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 a in actions {
tx.execute(
"INSERT INTO playbook_actions (playbook_id, action_order, action_type, params) VALUES (?1, ?2, ?3, ?4)",
params![id, a.action_order, a.action_type, a.params_json],
)?;
}
for c in conditions {
tx.execute(
"INSERT INTO playbook_conditions (playbook_id, condition_type, operator, value, value2) VALUES (?1, ?2, ?3, ?4, ?5)",
params![id, c.condition_type, c.operator, c.value, c.value2.as_deref()],
)?;
}
tx.commit()?;
Ok(true)
})
.await
}
}

View File

@ -1,3 +1,4 @@
use async_trait::async_trait;
use rusqlite::params;
use super::Database;
@ -5,236 +6,233 @@ use crate::domain::common::error::Error;
use crate::domain::response::playbook_data::{ActiveBlockView, PendingUnblock};
use crate::interface::db_admin::DbAdminRepo;
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<i64, Error> {
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<u32, Error> {
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 list_expired_soar_blocks(&self) -> Result<Vec<ActiveBlockView>, Error> {
let conn = self.conn()?;
let mut stmt = conn.prepare(
"SELECT id, source_ip, playbook_id, expires_at FROM soar_block_rules WHERE expires_at <= datetime('now') AND unblocked_at IS NULL",
)?;
let rows = stmt.query_map([], |row| {
Ok(ActiveBlockView {
id: row.get::<_, i64>(0)?,
source_ip: row.get::<_, String>(1)?,
playbook_id: row.get::<_, i64>(2)?,
expires_at: row.get::<_, String>(3)?,
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
}
/// Get a single SOAR block rule by ID.
pub fn find_soar_block_by_id(&self, id: i64) -> Result<Option<ActiveBlockView>, 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(ActiveBlockView {
id: row.get::<_, i64>(0)?,
source_ip: row.get::<_, String>(1)?,
playbook_id: row.get::<_, i64>(2)?,
expires_at: 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 list_active_soar_blocks(&self) -> Result<Vec<ActiveBlockView>, 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(ActiveBlockView {
id: row.get::<_, i64>(0)?,
source_ip: row.get::<_, String>(1)?,
playbook_id: row.get::<_, i64>(2)?,
expires_at: 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<i64, Error> {
let conn = self.conn()?;
conn.execute(
"INSERT INTO pending_unblock (source_ip) VALUES (?1)",
params![source_ip],
)?;
Ok(conn.last_insert_rowid())
}
pub fn list_pending_unblocks(&self) -> Result<Vec<PendingUnblock>, 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(PendingUnblock {
id: row.get::<_, i64>(0)?,
source_ip: row.get::<_, String>(1)?,
retry_count: 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(())
}
fn active_block_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<ActiveBlockView> {
Ok(ActiveBlockView {
id: row.get::<_, i64>(0)?,
source_ip: row.get::<_, String>(1)?,
playbook_id: row.get::<_, i64>(2)?,
expires_at: row.get::<_, String>(3)?,
})
}
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(
impl Database {
pub async fn count_active_soar_blocks(&self) -> Result<u32, Error> {
self.pool
.conn_and_then(move |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)
})
.await
}
pub async fn list_expired_soar_blocks(&self) -> Result<Vec<ActiveBlockView>, Error> {
self.pool
.conn_and_then(move |conn| {
let mut stmt = conn.prepare(
"SELECT id, source_ip, playbook_id, expires_at FROM soar_block_rules WHERE expires_at <= datetime('now') AND unblocked_at IS NULL",
)?;
let rows = stmt.query_map([], active_block_from_row)?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
})
.await
}
pub async fn find_soar_block_by_id(&self, id: i64) -> Result<Option<ActiveBlockView>, Error> {
self.pool
.conn_and_then(move |conn| {
let result = conn.query_row(
"SELECT id, source_ip, playbook_id, expires_at FROM soar_block_rules WHERE id = ?1",
params![id],
active_block_from_row,
);
match result {
Ok(block) => Ok(Some(block)),
Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e.into()),
}
})
.await
}
pub async fn mark_soar_block_unblocked(&self, id: i64) -> Result<(), Error> {
self.pool
.conn_and_then(move |conn| {
conn.execute(
"UPDATE soar_block_rules SET unblocked_at = datetime('now') WHERE id = ?1",
params![id],
)?;
Ok(())
})
.await
}
pub async fn list_active_soar_blocks(&self) -> Result<Vec<ActiveBlockView>, Error> {
self.pool
.conn_and_then(move |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([], active_block_from_row)?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
})
.await
}
pub async fn insert_pending_unblock(&self, source_ip: &str) -> Result<i64, Error> {
let source_ip = source_ip.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute(
"INSERT INTO pending_unblock (source_ip) VALUES (?1)",
params![source_ip],
)?;
Ok(conn.last_insert_rowid())
})
.await
}
pub async fn list_pending_unblocks(&self) -> Result<Vec<PendingUnblock>, Error> {
self.pool
.conn_and_then(move |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(PendingUnblock {
id: row.get::<_, i64>(0)?,
source_ip: row.get::<_, String>(1)?,
retry_count: row.get::<_, i64>(2)?,
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
})
.await
}
pub async fn delete_pending_unblock(&self, id: i64) -> Result<(), Error> {
self.pool
.conn_and_then(move |conn| {
conn.execute("DELETE FROM pending_unblock WHERE id = ?1", params![id])?;
Ok(())
})
.await
}
pub async fn increment_pending_unblock_retry(&self, id: i64) -> Result<(), Error> {
self.pool
.conn_and_then(move |conn| {
conn.execute(
"UPDATE pending_unblock SET retry_count = retry_count + 1 WHERE id = ?1",
params![id],
)?;
Ok(())
})
.await
}
pub async fn commit_soar_block_to_db(
&self,
source_ip: &str,
ip_version: u8,
playbook_id: i64,
expires_at: &str,
) -> Result<i64, Error> {
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)
let source_ip = source_ip.to_string();
let expires_at = expires_at.to_string();
self.pool
.conn_mut_and_then(move |conn| {
let tx = conn.transaction()?;
let existing_acl_count: i64 = tx.query_row(
"SELECT COUNT(*) 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],
|row| row.get(0),
)?;
let created_acl_rule = existing_acl_count == 0;
tx.execute(
"INSERT INTO soar_block_rules (source_ip, playbook_id, expires_at, created_acl_rule)
VALUES (?1, ?2, ?3, ?4)",
params![source_ip, playbook_id, expires_at, created_acl_rule as i64],
)?;
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)
})
.await
}
/// 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(())
pub async fn commit_soar_unblock_to_db(
&self,
soar_block_id: i64,
ip_version: u8,
source_ip: &str,
) -> Result<(), Error> {
let source_ip = source_ip.to_string();
self.pool
.conn_mut_and_then(move |conn| {
let tx = conn.transaction()?;
let should_delete_acl: bool = tx.query_row(
"SELECT created_acl_rule = 1 AND preserve_acl_on_unblock = 0
FROM soar_block_rules WHERE id = ?1",
params![soar_block_id],
|row| row.get(0),
)?;
if should_delete_acl {
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(())
})
.await
}
}
#[cfg(test)]
mod tests {
use super::super::tests::test_db;
use crate::interface::db_admin::DbAdminRepo;
use crate::interface::soar::SoarRepo;
/// 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 = SoarRepo::list_active_soar_blocks(&db).unwrap();
assert_eq!(active.len(), 1);
assert_eq!(active[0].source_ip, "10.0.0.99");
// acl_rules has the matching row
let rules = db.list_acl_rules().unwrap();
assert_eq!(rules.len(), 1);
assert_eq!(rules[0].ip_address, "10.0.0.99");
#[async_trait]
impl DbAdminRepo for Database {
async fn commit_soar_block_to_db(
&self,
source_ip: &str,
ip_version: u8,
playbook_id: i64,
expires_at: &str,
) -> Result<i64, Error> {
self.commit_soar_block_to_db(source_ip, ip_version, playbook_id, expires_at)
.await
}
/// 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.list_acl_rules().unwrap().is_empty());
// soar_block_rules row no longer in "active" view (unblocked_at is set)
assert!(SoarRepo::list_active_soar_blocks(&db).unwrap().is_empty());
async fn commit_soar_unblock_to_db(
&self,
soar_block_id: i64,
ip_version: u8,
source_ip: &str,
) -> Result<(), Error> {
self.commit_soar_unblock_to_db(soar_block_id, ip_version, source_ip)
.await
}
}

View File

@ -1,3 +1,4 @@
use async_trait::async_trait;
use rusqlite::params;
use super::Database;
@ -6,107 +7,120 @@ use crate::domain::report::data::{ThreatBreakdownEntry, TopIpEntry};
use crate::interface::stats::StatsRepo;
impl Database {
/// Count SOAR executions in the last N days.
pub fn count_weekly_executions(&self, days: i64) -> Result<u64, Error> {
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<u64, Error> {
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<u64, Error> {
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<Vec<ThreatBreakdownEntry>, 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(ThreatBreakdownEntry {
threat_type: row.get(0)?,
count: row.get::<_, i64>(1)? as u64,
pub async fn count_weekly_executions(&self, days: i64) -> Result<u64, Error> {
self.pool
.conn_and_then(move |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)
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
.await
}
/// Get top blocked IPs in the last N days.
pub fn weekly_top_ips(&self, days: i64, limit: i64) -> Result<Vec<TopIpEntry>, 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(TopIpEntry {
ip: row.get(0)?,
count: row.get::<_, i64>(1)? as u64,
pub async fn count_weekly_blocks(&self, days: i64) -> Result<u64, Error> {
self.pool
.conn_and_then(move |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)
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
.await
}
/// Count current ACL rules.
pub fn count_acl_rules(&self) -> Result<u64, Error> {
let conn = self.conn()?;
let count: i64 = conn.query_row("SELECT COUNT(*) FROM acl_rules", [], |row| row.get(0))?;
Ok(count as u64)
pub async fn count_weekly_unblocks(&self, days: i64) -> Result<u64, Error> {
self.pool
.conn_and_then(move |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)
})
.await
}
pub async fn weekly_threat_breakdown(&self, days: i64) -> Result<Vec<ThreatBreakdownEntry>, Error> {
self.pool
.conn_and_then(move |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(ThreatBreakdownEntry {
threat_type: row.get(0)?,
count: row.get::<_, i64>(1)? as u64,
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
})
.await
}
pub async fn weekly_top_ips(&self, days: i64, limit: i64) -> Result<Vec<TopIpEntry>, Error> {
self.pool
.conn_and_then(move |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(TopIpEntry {
ip: row.get(0)?,
count: row.get::<_, i64>(1)? as u64,
})
})?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
})
.await
}
pub async fn count_acl_rules(&self) -> Result<u64, Error> {
self.pool
.conn_and_then(move |conn| {
let count: i64 = conn.query_row("SELECT COUNT(*) FROM acl_rules", [], |row| row.get(0))?;
Ok(count as u64)
})
.await
}
}
#[async_trait]
impl StatsRepo for Database {
fn count_weekly_executions(&self, days: i64) -> Result<u64, Error> {
self.count_weekly_executions(days)
async fn count_weekly_executions(&self, days: i64) -> Result<u64, Error> {
self.count_weekly_executions(days).await
}
fn count_weekly_blocks(&self, days: i64) -> Result<u64, Error> {
self.count_weekly_blocks(days)
async fn count_weekly_blocks(&self, days: i64) -> Result<u64, Error> {
self.count_weekly_blocks(days).await
}
fn count_weekly_unblocks(&self, days: i64) -> Result<u64, Error> {
self.count_weekly_unblocks(days)
async fn count_weekly_unblocks(&self, days: i64) -> Result<u64, Error> {
self.count_weekly_unblocks(days).await
}
fn weekly_threat_breakdown(&self, days: i64) -> Result<Vec<ThreatBreakdownEntry>, Error> {
self.weekly_threat_breakdown(days)
async fn weekly_threat_breakdown(&self, days: i64) -> Result<Vec<ThreatBreakdownEntry>, Error> {
self.weekly_threat_breakdown(days).await
}
fn weekly_top_ips(&self, days: i64, limit: i64) -> Result<Vec<TopIpEntry>, Error> {
self.weekly_top_ips(days, limit)
async fn weekly_top_ips(&self, days: i64, limit: i64) -> Result<Vec<TopIpEntry>, Error> {
self.weekly_top_ips(days, limit).await
}
fn count_acl_rules(&self) -> Result<u64, Error> {
self.count_acl_rules()
async fn count_acl_rules(&self) -> Result<u64, Error> {
self.count_acl_rules().await
}
}

View File

@ -0,0 +1,49 @@
use async_trait::async_trait;
use rusqlite::{Error as RusqliteError, params};
use super::Database;
use crate::domain::common::error::Error;
use crate::interface::system_state::SystemStateRepo;
impl Database {
pub async fn get_system_state(&self, key: &str) -> Result<Option<String>, Error> {
let key = key.to_string();
self.pool
.conn_and_then(move |conn| {
let result = conn.query_row("SELECT value FROM system_state WHERE key = ?1", params![key], |row| {
row.get(0)
});
match result {
Ok(value) => Ok(Some(value)),
Err(RusqliteError::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e.into()),
}
})
.await
}
pub async fn set_system_state(&self, key: &str, value: &str) -> Result<(), Error> {
let key = key.to_string();
let value = value.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute(
"INSERT OR REPLACE INTO system_state (key, value) VALUES (?1, ?2)",
params![key, value],
)?;
Ok(())
})
.await
}
}
#[async_trait]
impl SystemStateRepo for Database {
async fn get_system_state(&self, key: &str) -> Result<Option<String>, Error> {
self.get_system_state(key).await
}
async fn set_system_state(&self, key: &str, value: &str) -> Result<(), Error> {
self.set_system_state(key, value).await
}
}

View File

@ -1,276 +1,343 @@
use std::collections::{HashMap, HashSet};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
use rusqlite::{Error as RusqliteError, params};
use super::Database;
use crate::domain::common::error::Error;
use crate::domain::common::error::database::DatabaseError;
use crate::domain::identity::auth::{LOGIN_LOCKOUT_SECS, LOGIN_MAX_FAILURES};
use crate::domain::identity::auth::{GROUP_ADMIN, GROUP_VIEWER, LOGIN_LOCKOUT_SECS, LOGIN_MAX_FAILURES, ROLE_ADMIN};
use crate::domain::identity::user::{
GroupMemberView, UserGroupMembership, UserGroupView, UserView, UserWithGroupsView,
};
use crate::interface::identity::{LoginAttemptRepo, UserGroupRepo, UserRepo};
fn user_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<UserView> {
Ok(UserView {
id: row.get(0)?,
username: row.get(1)?,
password_hash: row.get(2)?,
force_password_change: row.get::<_, i64>(3)? != 0,
})
}
fn group_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<UserGroupView> {
Ok(UserGroupView {
id: row.get(0)?,
name: row.get(1)?,
description: row.get(2)?,
permissions: row.get(3)?,
created_at: row.get(4)?,
})
}
impl Database {
pub fn find_user(&self, username: &str) -> Result<Option<UserView>, Error> {
let conn = self.conn()?;
let result = conn.query_row(
"SELECT id, username, password_hash, force_password_change FROM users WHERE username = ?1",
params![username],
|row| {
Ok(UserView {
id: row.get(0)?,
username: row.get(1)?,
password_hash: row.get(2)?,
force_password_change: row.get::<_, i64>(3)? != 0,
})
},
);
match result {
Ok(user) => Ok(Some(user)),
Err(RusqliteError::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e)?,
}
pub async fn find_user(&self, username: &str) -> Result<Option<UserView>, Error> {
let username = username.to_string();
self.pool
.conn_and_then(move |conn| {
match conn.query_row(
"SELECT id, username, password_hash, force_password_change FROM users WHERE username = ?1",
params![username],
user_from_row,
) {
Ok(user) => Ok(Some(user)),
Err(RusqliteError::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e.into()),
}
})
.await
}
pub fn insert_user(
pub async fn insert_user(
&self,
username: &str,
password_hash: &str,
role: &str,
force_password_change: bool,
) -> Result<i64, Error> {
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())
let username = username.to_string();
let password_hash = password_hash.to_string();
let role = role.to_string();
self.pool
.conn_and_then(move |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.clone()).into()
} else {
e.into()
}
})?;
Ok(conn.last_insert_rowid())
})
.await
}
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 async fn update_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> {
let password_hash = password_hash.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute(
"UPDATE users SET password_hash = ?1, force_password_change = 0 WHERE id = ?2",
params![password_hash, user_id],
)?;
Ok(())
})
.await
}
pub fn user_count(&self) -> Result<i64, Error> {
let conn = self.conn()?;
Ok(conn.query_row("SELECT COUNT(*) FROM users", [], |row| row.get(0))?)
pub async fn user_count(&self) -> Result<i64, Error> {
self.pool
.conn_and_then(move |conn| Ok(conn.query_row("SELECT COUNT(*) FROM users", [], |row| row.get(0))?))
.await
}
pub fn list_users_with_groups(&self) -> Result<Vec<UserWithGroupsView>, Error> {
let conn = self.conn()?;
let mut stmt = conn.prepare(
"SELECT u.id, u.username, 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::<_, i64>(2)? != 0,
row.get::<_, String>(3)?,
row.get::<_, Option<i64>>(4)?,
row.get::<_, Option<String>>(5)?,
))
})?;
pub async fn list_users_with_groups(&self) -> Result<Vec<UserWithGroupsView>, Error> {
self.pool
.conn_and_then(move |conn| {
let mut stmt = conn.prepare(
"SELECT u.id, u.username, 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::<_, i64>(2)? != 0,
row.get::<_, String>(3)?,
row.get::<_, Option<i64>>(4)?,
row.get::<_, Option<String>>(5)?,
))
})?;
let mut user_map: HashMap<i64, UserWithGroupsView> = HashMap::new();
let mut order: Vec<i64> = Vec::new();
for row in rows {
let (id, username, force_pw, created_at, group_id, group_name) = row?;
let entry = user_map.entry(id).or_insert_with(|| {
order.push(id);
UserWithGroupsView {
id,
username,
force_password_change: force_pw,
created_at,
groups: Vec::new(),
let mut user_map: HashMap<i64, UserWithGroupsView> = HashMap::new();
let mut order: Vec<i64> = Vec::new();
for row in rows {
let (id, username, force_pw, created_at, group_id, group_name) = row?;
let entry = user_map.entry(id).or_insert_with(|| {
order.push(id);
UserWithGroupsView {
id,
username,
force_password_change: force_pw,
created_at,
groups: Vec::new(),
}
});
if let (Some(gid), Some(gname)) = (group_id, group_name) {
entry.groups.push(UserGroupMembership {
group_id: gid,
group_name: gname,
});
}
}
});
if let (Some(gid), Some(gname)) = (group_id, group_name) {
entry.groups.push(UserGroupMembership {
group_id: gid,
group_name: gname,
});
}
}
Ok(order.into_iter().filter_map(|id| user_map.remove(&id)).collect())
}
pub fn delete_user(&self, user_id: i64) -> Result<bool, Error> {
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<Option<UserView>, Error> {
let conn = self.conn()?;
let result = conn.query_row(
"SELECT id, username, password_hash, force_password_change FROM users WHERE id = ?1",
params![user_id],
|row| {
Ok(UserView {
id: row.get(0)?,
username: row.get(1)?,
password_hash: row.get(2)?,
force_password_change: row.get::<_, i64>(3)? != 0,
})
},
);
match result {
Ok(user) => Ok(Some(user)),
Err(RusqliteError::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e)?,
}
}
pub fn list_user_groups(&self) -> Result<Vec<UserGroupView>, 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(UserGroupView {
id: row.get(0)?,
name: row.get(1)?,
description: row.get(2)?,
permissions: row.get(3)?,
created_at: row.get(4)?,
Ok(order.into_iter().filter_map(|id| user_map.remove(&id)).collect())
})
})?;
let mut results = Vec::new();
for row in rows {
results.push(row?);
}
Ok(results)
.await
}
pub fn create_user_group(&self, name: &str, description: &str, permissions: &str) -> Result<i64, Error> {
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<bool, Error> {
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<Option<UserGroupView>, 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(UserGroupView {
id: row.get(0)?,
name: row.get(1)?,
description: row.get(2)?,
permissions: row.get(3)?,
created_at: row.get(4)?,
})
},
);
match result {
Ok(group) => Ok(Some(group)),
Err(RusqliteError::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e)?,
}
}
pub fn list_groups_for_user(&self, user_id: i64) -> Result<Vec<UserGroupView>, Error> {
let conn = self.conn()?;
let mut stmt = conn.prepare(
"SELECT g.id, g.name, g.description, g.permissions, g.created_at 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(UserGroupView {
id: row.get(0)?,
name: row.get(1)?,
description: row.get(2)?,
permissions: row.get(3)?,
created_at: row.get(4)?,
pub async fn delete_user(&self, user_id: i64) -> Result<bool, Error> {
self.pool
.conn_mut_and_then(move |conn| {
let tx = conn.transaction()?;
tx.execute("DELETE FROM user_group_members WHERE user_id = ?1", params![user_id])?;
let affected = tx.execute("DELETE FROM users WHERE id = ?1", params![user_id])?;
tx.commit()?;
Ok(affected > 0)
})
})?;
let mut results = Vec::new();
for row in rows {
results.push(row?);
}
Ok(results)
.await
}
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 async fn update_user_role(&self, user_id: i64, role: &str) -> Result<(), Error> {
let role = role.to_string();
self.pool
.conn_mut_and_then(move |conn| {
let tx = conn.transaction()?;
tx.execute("UPDATE users SET role = ?1 WHERE id = ?2", params![role, user_id])?;
let group_name = if role == ROLE_ADMIN { GROUP_ADMIN } else { GROUP_VIEWER };
let group_id: i64 = tx.query_row(
"SELECT id FROM user_groups WHERE name = ?1",
params![group_name],
|row| row.get(0),
)?;
tx.execute("DELETE FROM user_group_members WHERE user_id = ?1", params![user_id])?;
tx.execute(
"INSERT INTO user_group_members (user_id, group_id) VALUES (?1, ?2)",
params![user_id, group_id],
)?;
tx.commit()?;
Ok(())
})
.await
}
pub fn list_user_permissions(&self, user_id: i64) -> Result<Vec<String>, Error> {
let groups = self.list_groups_for_user(user_id)?;
pub async fn reset_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> {
let password_hash = password_hash.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute(
"UPDATE users SET password_hash = ?1, force_password_change = 1 WHERE id = ?2",
params![password_hash, user_id],
)?;
Ok(())
})
.await
}
pub async fn find_user_by_id(&self, user_id: i64) -> Result<Option<UserView>, Error> {
self.pool
.conn_and_then(move |conn| {
match conn.query_row(
"SELECT id, username, password_hash, force_password_change FROM users WHERE id = ?1",
params![user_id],
user_from_row,
) {
Ok(user) => Ok(Some(user)),
Err(RusqliteError::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e.into()),
}
})
.await
}
pub async fn list_user_groups(&self) -> Result<Vec<UserGroupView>, Error> {
self.pool
.conn_and_then(move |conn| {
let mut stmt =
conn.prepare("SELECT id, name, description, permissions, created_at FROM user_groups ORDER BY id")?;
let rows = stmt.query_map([], group_from_row)?;
let mut results = Vec::new();
for row in rows {
results.push(row?);
}
Ok(results)
})
.await
}
pub async fn create_user_group(&self, name: &str, description: &str, permissions: &str) -> Result<i64, Error> {
let name = name.to_string();
let description = description.to_string();
let permissions = permissions.to_string();
self.pool
.conn_and_then(move |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.clone()).into()
} else {
e.into()
}
})?;
Ok(conn.last_insert_rowid())
})
.await
}
pub async fn update_user_group(
&self,
id: i64,
name: &str,
description: &str,
permissions: &str,
) -> Result<(), Error> {
let name = name.to_string();
let description = description.to_string();
let permissions = permissions.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute(
"UPDATE user_groups SET name = ?1, description = ?2, permissions = ?3 WHERE id = ?4",
params![name, description, permissions, id],
)?;
Ok(())
})
.await
}
pub async fn delete_user_group(&self, id: i64) -> Result<bool, Error> {
self.pool
.conn_mut_and_then(move |conn| {
let tx = conn.transaction()?;
tx.execute("DELETE FROM user_group_members WHERE group_id = ?1", params![id])?;
let affected = tx.execute("DELETE FROM user_groups WHERE id = ?1", params![id])?;
tx.commit()?;
Ok(affected > 0)
})
.await
}
pub async fn get_user_group(&self, id: i64) -> Result<Option<UserGroupView>, Error> {
self.pool
.conn_and_then(move |conn| {
match conn.query_row(
"SELECT id, name, description, permissions, created_at FROM user_groups WHERE id = ?1",
params![id],
group_from_row,
) {
Ok(group) => Ok(Some(group)),
Err(RusqliteError::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e.into()),
}
})
.await
}
pub async fn list_groups_for_user(&self, user_id: i64) -> Result<Vec<UserGroupView>, Error> {
self.pool
.conn_and_then(move |conn| {
let mut stmt = conn.prepare(
"SELECT g.id, g.name, g.description, g.permissions, g.created_at 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], group_from_row)?;
let mut results = Vec::new();
for row in rows {
results.push(row?);
}
Ok(results)
})
.await
}
pub async fn set_user_groups(&self, user_id: i64, group_ids: &[i64]) -> Result<(), Error> {
let group_ids = group_ids.to_vec();
self.pool
.conn_mut_and_then(move |conn| {
let tx = conn.transaction()?;
for &gid in &group_ids {
tx.query_row("SELECT id FROM user_groups WHERE id = ?1", params![gid], |row| {
row.get::<_, i64>(0)
})?;
}
tx.execute("DELETE FROM user_group_members WHERE user_id = ?1", params![user_id])?;
for &gid in &group_ids {
tx.execute(
"INSERT INTO user_group_members (user_id, group_id) VALUES (?1, ?2)",
params![user_id, gid],
)?;
}
tx.commit()?;
Ok(())
})
.await
}
pub async fn list_user_permissions(&self, user_id: i64) -> Result<Vec<String>, Error> {
let groups = self.list_groups_for_user(user_id).await?;
let mut all_perms = HashSet::new();
for g in groups {
if let Ok(perms) = serde_json::from_str::<Vec<String>>(&g.permissions) {
@ -284,68 +351,109 @@ impl Database {
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 list_group_member_ids(&self, group_id: i64) -> Result<Vec<i64>, 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 list_group_members(&self, group_id: i64) -> Result<Vec<GroupMemberView>, 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(GroupMemberView {
id: row.get(0)?,
username: row.get(1)?,
pub async fn list_group_member_ids(&self, group_id: i64) -> Result<Vec<i64>, Error> {
self.pool
.conn_and_then(move |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)
})
})?;
let mut results = Vec::new();
for row in rows {
results.push(row?);
}
Ok(results)
.await
}
pub fn record_login_failure(&self, username: &str) -> Result<(u32, Option<u64>), 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())?;
pub async fn list_group_members(&self, group_id: i64) -> Result<Vec<GroupMemberView>, Error> {
self.pool
.conn_and_then(move |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(GroupMemberView {
id: row.get(0)?,
username: row.get(1)?,
})
})?;
let mut results = Vec::new();
for row in rows {
results.push(row?);
}
Ok(results)
})
.await
}
pub async fn record_login_failure(&self, username: &str) -> Result<(u32, Option<u64>), Error> {
let username = username.to_string();
let count: u32 = self
.pool
.conn_and_then({
let username = username.clone();
move |conn| {
let count = conn
.query_row(
"SELECT failure_count FROM login_attempts WHERE username = ?1",
params![username],
|row| row.get::<_, u32>(0),
)
.unwrap_or(0)
+ 1;
conn.execute(
"INSERT INTO login_attempts (username, failure_count) VALUES (?1, ?2) \
ON CONFLICT(username) DO UPDATE SET failure_count = ?2",
params![username, count],
)?;
Ok::<u32, Error>(count)
}
})
.await?;
if count >= LOGIN_MAX_FAILURES {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or(Duration::ZERO)
.as_secs();
let locked_until = now + LOGIN_LOCKOUT_SECS;
self.set_setting(&key_locked, &locked_until.to_string())?;
let locked_until_db = locked_until as i64;
self.pool
.conn_and_then(move |conn| {
conn.execute(
"UPDATE login_attempts SET locked_until = ?2 WHERE username = ?1",
params![username, locked_until_db],
)?;
Ok::<(), Error>(())
})
.await?;
Ok((count, Some(locked_until)))
} else {
Ok((count, None))
}
}
pub fn check_login_locked(&self, username: &str) -> Result<Option<u64>, 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::<u64>()
pub async fn check_login_locked(&self, username: &str) -> Result<Option<u64>, Error> {
let username = username.to_string();
if let Some(locked_until) = self
.pool
.conn_and_then({
let username = username.clone();
move |conn| {
let result = conn.query_row(
"SELECT locked_until FROM login_attempts WHERE username = ?1",
params![username],
|row| row.get::<_, Option<i64>>(0),
);
match result {
Ok(value) => Ok::<Option<u64>, Error>(value.and_then(|v| u64::try_from(v).ok())),
Err(rusqlite::Error::QueryReturnedNoRows) => Ok::<Option<u64>, Error>(None),
Err(e) => Err(Error::from(e)),
}
}
})
.await?
{
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
@ -354,36 +462,33 @@ impl Database {
if now < locked_until {
return Ok(Some(locked_until - now));
}
// Lock expired, clear it
self.clear_login_failures(username)?;
self.clear_login_failures(&username).await?;
}
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(())
pub async fn clear_login_failures(&self, username: &str) -> Result<(), Error> {
let username = username.to_string();
self.pool
.conn_and_then(move |conn| {
conn.execute("DELETE FROM login_attempts WHERE username = ?1", params![username])?;
Ok(())
})
.await
}
}
#[async_trait]
impl UserRepo for Database {
fn find_user(&self, username: &str) -> Result<Option<UserView>, Error> {
self.find_user(username)
async fn find_user(&self, username: &str) -> Result<Option<UserView>, Error> {
self.find_user(username).await
}
fn find_user_by_id(&self, user_id: i64) -> Result<Option<UserView>, Error> {
self.find_user_by_id(user_id)
async fn find_user_by_id(&self, user_id: i64) -> Result<Option<UserView>, Error> {
self.find_user_by_id(user_id).await
}
fn insert_user(
async fn insert_user(
&self,
username: &str,
password_hash: &str,
@ -391,157 +496,84 @@ impl UserRepo for Database {
force_password_change: bool,
) -> Result<i64, Error> {
self.insert_user(username, password_hash, role, force_password_change)
.await
}
fn update_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> {
self.update_user_password(user_id, password_hash)
async fn update_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> {
self.update_user_password(user_id, password_hash).await
}
fn list_users_with_groups(&self) -> Result<Vec<UserWithGroupsView>, Error> {
self.list_users_with_groups()
async fn list_users_with_groups(&self) -> Result<Vec<UserWithGroupsView>, Error> {
self.list_users_with_groups().await
}
fn delete_user(&self, user_id: i64) -> Result<bool, Error> {
self.delete_user(user_id)
async fn delete_user(&self, user_id: i64) -> Result<bool, Error> {
self.delete_user(user_id).await
}
fn update_user_role(&self, user_id: i64, role: &str) -> Result<(), Error> {
self.update_user_role(user_id, role)
async fn update_user_role(&self, user_id: i64, role: &str) -> Result<(), Error> {
self.update_user_role(user_id, role).await
}
fn reset_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> {
self.reset_user_password(user_id, password_hash)
async fn reset_user_password(&self, user_id: i64, password_hash: &str) -> Result<(), Error> {
self.reset_user_password(user_id, password_hash).await
}
}
#[async_trait]
impl UserGroupRepo for Database {
fn list_user_groups(&self) -> Result<Vec<UserGroupView>, Error> {
self.list_user_groups()
async fn list_user_groups(&self) -> Result<Vec<UserGroupView>, Error> {
self.list_user_groups().await
}
fn create_user_group(&self, name: &str, description: &str, permissions: &str) -> Result<i64, Error> {
self.create_user_group(name, description, permissions)
async fn create_user_group(&self, name: &str, description: &str, permissions: &str) -> Result<i64, Error> {
self.create_user_group(name, description, permissions).await
}
fn update_user_group(&self, id: i64, name: &str, description: &str, permissions: &str) -> Result<(), Error> {
self.update_user_group(id, name, description, permissions)
async fn update_user_group(&self, id: i64, name: &str, description: &str, permissions: &str) -> Result<(), Error> {
self.update_user_group(id, name, description, permissions).await
}
fn delete_user_group(&self, id: i64) -> Result<bool, Error> {
self.delete_user_group(id)
async fn delete_user_group(&self, id: i64) -> Result<bool, Error> {
self.delete_user_group(id).await
}
fn get_user_group(&self, id: i64) -> Result<Option<UserGroupView>, Error> {
self.get_user_group(id)
async fn get_user_group(&self, id: i64) -> Result<Option<UserGroupView>, Error> {
self.get_user_group(id).await
}
fn list_groups_for_user(&self, user_id: i64) -> Result<Vec<UserGroupView>, Error> {
self.list_groups_for_user(user_id)
async fn list_groups_for_user(&self, user_id: i64) -> Result<Vec<UserGroupView>, Error> {
self.list_groups_for_user(user_id).await
}
fn set_user_groups(&self, user_id: i64, group_ids: &[i64]) -> Result<(), Error> {
self.set_user_groups(user_id, group_ids)
async fn set_user_groups(&self, user_id: i64, group_ids: &[i64]) -> Result<(), Error> {
self.set_user_groups(user_id, group_ids).await
}
fn list_user_permissions(&self, user_id: i64) -> Result<Vec<String>, Error> {
self.list_user_permissions(user_id)
async fn list_user_permissions(&self, user_id: i64) -> Result<Vec<String>, Error> {
self.list_user_permissions(user_id).await
}
fn list_group_member_ids(&self, group_id: i64) -> Result<Vec<i64>, Error> {
self.list_group_member_ids(group_id)
async fn list_group_member_ids(&self, group_id: i64) -> Result<Vec<i64>, Error> {
self.list_group_member_ids(group_id).await
}
fn list_group_members(&self, group_id: i64) -> Result<Vec<GroupMemberView>, Error> {
self.list_group_members(group_id)
async fn list_group_members(&self, group_id: i64) -> Result<Vec<GroupMemberView>, Error> {
self.list_group_members(group_id).await
}
}
#[async_trait]
impl LoginAttemptRepo for Database {
fn record_login_failure(&self, username: &str) -> Result<(u32, Option<u64>), Error> {
self.record_login_failure(username)
async fn record_login_failure(&self, username: &str) -> Result<(u32, Option<u64>), Error> {
self.record_login_failure(username).await
}
fn check_login_locked(&self, username: &str) -> Result<Option<u64>, Error> {
self.check_login_locked(username)
async fn check_login_locked(&self, username: &str) -> Result<Option<u64>, Error> {
self.check_login_locked(username).await
}
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.id, 1);
assert_eq!(user.username, "admin");
assert_eq!(user.password_hash, "hash123");
assert!(user.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.force_password_change);
db.update_user_password(user.id, "new_hash").unwrap();
let user = db.find_user("admin").unwrap().unwrap();
assert!(!user.force_password_change);
assert_eq!(user.password_hash, "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());
async fn clear_login_failures(&self, username: &str) -> Result<(), Error> {
self.clear_login_failures(username).await
}
}

View File

@ -13,14 +13,14 @@ use crate::domain::common::error::Error;
use crate::domain::common::error::notification::NotificationError;
use crate::domain::common::log::system::SystemLog;
use crate::domain::common::notification::AlertPayload;
use crate::interface::config_repo::ConfigRepo;
use crate::interface::notification::{AlertNotifier, AlertNotifierFactory};
use crate::interface::secret_store::SecretStorePort;
use crate::interface::setting::SettingRepo;
/// Telegram Bot API adapter implementing AlertNotifier.
pub struct TelegramAdapter {
client: Client,
notif: Arc<dyn SettingRepo + Send + Sync>,
notif: Arc<dyn ConfigRepo + Send + Sync>,
config: Arc<ArcSwap<AppConfig>>,
secrets: Option<Arc<dyn SecretStorePort>>,
/// Packed rate-limit state: high 32 bits = window-start unix seconds,
@ -31,7 +31,7 @@ pub struct TelegramAdapter {
impl TelegramAdapter {
pub fn new(
notif: Arc<dyn SettingRepo + Send + Sync>,
notif: Arc<dyn ConfigRepo + Send + Sync>,
config: Arc<ArcSwap<AppConfig>>,
secrets: Option<Arc<dyn SecretStorePort>>,
) -> Result<Self, Error> {
@ -51,8 +51,8 @@ impl TelegramAdapter {
/// Get bot token and chat ID from DB. Returns None if not configured.
/// If the bot_token in JSON is `"__encrypted__"`, reads from the secret store.
fn get_config(&self) -> Result<Option<(String, String)>, Error> {
match self.notif.get_notification_config("telegram")? {
async fn get_config(&self) -> Result<Option<(String, String)>, Error> {
match self.notif.get_notification_config("telegram").await? {
Some(json_str) => {
let config: serde_json::Value =
serde_json::from_str(&json_str).map_err(NotificationError::TelegramRequestFailed)?;
@ -60,11 +60,10 @@ impl TelegramAdapter {
let chat_id = config.get("chat_id").and_then(|v| v.as_str()).map(|s| s.to_string());
// If token is the encrypted sentinel, resolve from secret store
if token.as_deref() == Some("__encrypted__") {
token = self
.secrets
.as_ref()
.and_then(|ss| ss.get_secret("telegram_bot_token").ok().flatten());
if token.as_deref() == Some("__encrypted__")
&& let Some(ss) = self.secrets.as_ref()
{
token = ss.get_secret("telegram_bot_token").await?;
}
match (token, chat_id) {
@ -218,7 +217,7 @@ impl TelegramAdapter {
#[async_trait]
impl AlertNotifier for TelegramAdapter {
async fn send_alert(&self, payload: &AlertPayload) -> Result<(), Error> {
let (bot_token, chat_id) = match self.get_config()? {
let (bot_token, chat_id) = match self.get_config().await? {
Some(config) => config,
None => {
log!(SystemLog::TelegramNotConfiguredSkipped);
@ -240,7 +239,7 @@ impl AlertNotifier for TelegramAdapter {
}
async fn send_test_message(&self) -> Result<(), Error> {
let (bot_token, chat_id) = match self.get_config()? {
let (bot_token, chat_id) = match self.get_config().await? {
Some(config) => config,
None => Err(NotificationError::NotConfigured("telegram"))?,
};
@ -259,14 +258,14 @@ impl AlertNotifier for TelegramAdapter {
/// call instantiates a fresh `TelegramAdapter` so the test path observes
/// whatever config the user just saved.
pub struct TelegramAdapterFactory {
notif: Arc<dyn SettingRepo + Send + Sync>,
notif: Arc<dyn ConfigRepo + Send + Sync>,
config: Arc<ArcSwap<AppConfig>>,
secrets: Option<Arc<dyn SecretStorePort>>,
}
impl TelegramAdapterFactory {
pub fn new(
notif: Arc<dyn SettingRepo + Send + Sync>,
notif: Arc<dyn ConfigRepo + Send + Sync>,
config: Arc<ArcSwap<AppConfig>>,
secrets: Option<Arc<dyn SecretStorePort>>,
) -> Self {

View File

@ -1,12 +1,15 @@
use std::collections::HashMap;
use std::sync::Arc;
use arc_swap::ArcSwap;
use async_trait::async_trait;
use crate::domain::common::config::AppConfig;
use crate::domain::common::config::section::ConfigSection;
use crate::domain::common::error::Error;
use crate::domain::common::error::misc::MiscError;
use crate::interface::app_repo::AppRepo;
use crate::interface::config_repo::ConfigRepo;
use crate::interface::secret_store::SecretStorePort;
const SECRET_KEYS: &[&str] = &["smtp_password"];
@ -19,6 +22,49 @@ pub struct ConfigService {
app_config: Arc<ArcSwap<AppConfig>>,
}
struct CandidateConfigRepo {
base: Arc<dyn AppRepo>,
values: HashMap<String, String>,
}
#[async_trait]
impl ConfigRepo for CandidateConfigRepo {
async fn get_config_value(&self, key: &str) -> Result<Option<String>, Error> {
if let Some(value) = self.values.get(key) {
return Ok(Some(value.clone()));
}
self.base.get_config_value(key).await
}
async fn set_config_value(&self, key: &str, value: &str) -> Result<(), Error> {
self.base.set_config_value(key, value).await
}
async fn get_app_secret(&self, key: &str) -> Result<Option<String>, Error> {
self.base.get_app_secret(key).await
}
async fn set_app_secret(&self, key: &str, plaintext: &str) -> Result<(), Error> {
self.base.set_app_secret(key, plaintext).await
}
async fn get_notification_config(&self, channel: &str) -> Result<Option<String>, Error> {
self.base.get_notification_config(channel).await
}
async fn set_notification_config(&self, channel: &str, config_json: &str) -> Result<(), Error> {
self.base.set_notification_config(channel, config_json).await
}
async fn update_config_values_atomically(
&self,
config_values: Vec<(String, String)>,
secrets: Vec<(String, String)>,
) -> Result<(), Error> {
self.base.update_config_values_atomically(config_values, secrets).await
}
}
impl ConfigService {
pub fn new(db: Arc<dyn AppRepo>, app_config: Arc<ArcSwap<AppConfig>>) -> Self {
Self {
@ -33,89 +79,93 @@ impl ConfigService {
self
}
pub fn get_config(&self) -> serde_json::Value {
let get = |key: &str| -> String { self.db.get_setting(key).ok().flatten().unwrap_or_default() };
pub async fn get_config(&self) -> serde_json::Value {
let cfg = self.app_config.load();
let values = cfg.api_setting_values();
let mut root = serde_json::Map::new();
for section in ConfigSection::ALL {
let mut section_obj = serde_json::Map::new();
for key in section.keys() {
section_obj.insert(key.to_string(), serde_json::Value::String(get(key)));
let value = values.get(key).cloned().unwrap_or_default();
section_obj.insert(key.to_string(), serde_json::Value::String(value));
}
root.insert(section.name().to_string(), serde_json::Value::Object(section_obj));
}
root.insert(
"pipeline".to_string(),
serde_json::json!({
"ingress": get("pipeline_ingress"),
"egress": get("pipeline_egress"),
"ingress": cfg.pipeline.ingress.join(","),
"egress": cfg.pipeline.egress.join(","),
}),
);
serde_json::Value::Object(root)
}
pub fn update_config(&self, body: &serde_json::Value) -> Result<Vec<String>, Error> {
pub async fn update_config(&self, body: &serde_json::Value) -> Result<Vec<String>, Error> {
let mut updated: Vec<String> = Vec::new();
let mut new_cfg: Option<AppConfig> = None;
let mut config_values = Vec::new();
let mut secrets_to_save = Vec::new();
self.db.transaction(&mut |repo| {
for section in ConfigSection::ALL {
if let Some(section_obj) = body.get(section.name()).and_then(|v| v.as_object()) {
for key in section.keys() {
if let Some(val) = section_obj.get(*key).and_then(json_value_as_string) {
repo.set_setting(key, &val)?;
updated.push(key.to_string());
}
}
}
}
if let Some(ref secrets) = self.secrets {
for key in SECRET_KEYS {
let section = key.split('_').next().unwrap_or("");
if let Some(val) = body
.get(section)
.and_then(|v| v.as_object())
.and_then(|obj| obj.get(*key))
.and_then(json_value_as_string)
{
let envelope = secrets.encrypt_envelope(&val)?;
repo.set_app_secret(key, &envelope)?;
repo.set_setting(key, "")?;
for section in ConfigSection::ALL {
if let Some(section_obj) = body.get(section.name()).and_then(|v| v.as_object()) {
for key in section.keys() {
if let Some(val) = section_obj.get(*key).and_then(json_value_as_string) {
config_values.push((key.to_string(), val));
updated.push(key.to_string());
}
}
}
}
if let Some(pipeline_obj) = body.get("pipeline").and_then(|v| v.as_object()) {
for (field, db_key) in [("ingress", "pipeline_ingress"), ("egress", "pipeline_egress")] {
if let Some(val) = pipeline_obj.get(field).and_then(|v| v.as_str()) {
if !val.is_empty() {
let stages: Vec<&str> = val.split(',').map(|s| s.trim()).collect();
for stage in &stages {
if !stage.is_empty() && !VALID_PIPELINE_STAGES.contains(stage) {
Err(MiscError::ValidationError(format!(
"Invalid pipeline stage '{}'. Valid stages: {}",
stage,
VALID_PIPELINE_STAGES.join(", ")
)))?;
}
}
}
repo.set_setting(db_key, val)?;
updated.push(db_key.to_string());
}
if let Some(ref secrets) = self.secrets {
for key in SECRET_KEYS {
let section = key.split('_').next().unwrap_or("");
if let Some(val) = body
.get(section)
.and_then(|v| v.as_object())
.and_then(|obj| obj.get(*key))
.and_then(json_value_as_string)
{
let envelope = secrets.encrypt_envelope(&val)?;
secrets_to_save.push((key.to_string(), envelope));
config_values.push((key.to_string(), String::new()));
updated.push(key.to_string());
}
}
new_cfg = Some(AppConfig::from_settings(repo)?);
Ok(())
})?;
if let Some(cfg) = new_cfg {
self.app_config.store(Arc::new(cfg));
}
if let Some(pipeline_obj) = body.get("pipeline").and_then(|v| v.as_object()) {
for (field, db_key) in [("ingress", "pipeline_ingress"), ("egress", "pipeline_egress")] {
if let Some(val) = pipeline_obj.get(field).and_then(|v| v.as_str()) {
if !val.is_empty() {
let stages: Vec<&str> = val.split(',').map(|s| s.trim()).collect();
for stage in &stages {
if !stage.is_empty() && !VALID_PIPELINE_STAGES.contains(stage) {
Err(MiscError::ValidationError(format!(
"Invalid pipeline stage '{}'. Valid stages: {}",
stage,
VALID_PIPELINE_STAGES.join(", ")
)))?;
}
}
}
config_values.push((db_key.to_string(), val.to_string()));
updated.push(db_key.to_string());
}
}
}
let candidate_repo = CandidateConfigRepo {
base: self.db.clone(),
values: config_values.iter().cloned().collect(),
};
let new_cfg = AppConfig::from_config_repo(&candidate_repo).await?;
self.db
.update_config_values_atomically(config_values, secrets_to_save)
.await?;
self.app_config.store(Arc::new(new_cfg));
Ok(updated)
}
}
@ -128,3 +178,27 @@ fn json_value_as_string(v: &serde_json::Value) -> Option<String> {
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::adapter::persistence::Database;
#[tokio::test]
async fn invalid_config_update_is_not_persisted() {
let db = Arc::new(Database::new(":memory:").await.expect("test db"));
let app_config = Arc::new(ArcSwap::from_pointee(
AppConfig::from_config_repo(db.as_ref()).await.expect("initial config"),
));
let service = ConfigService::new(db.clone() as Arc<dyn AppRepo>, app_config);
let body = serde_json::json!({
"http": {
"http_port": 0
}
});
assert!(service.update_config(&body).await.is_err());
assert_eq!(db.get_config_value("http_port").await.unwrap(), None);
}
}

View File

@ -1,10 +1,12 @@
use std::sync::Arc;
use std::sync::atomic::{AtomicU8, Ordering};
use arc_swap::ArcSwap;
use macros::log;
use tokio::sync::broadcast;
use crate::domain::common::config::constants::{ENFORCE_MODE_MONITOR, enforce_mode_to_u8};
use crate::domain::common::config::AppConfig;
use crate::domain::common::config::constants::enforce_mode_to_u8;
use crate::domain::common::error::Error;
use crate::domain::common::event::AuditEvent;
use crate::domain::common::log::system::SystemLog;
@ -13,23 +15,33 @@ use crate::interface::app_repo::AppRepo;
/// Handles enforce-mode commands and queries by delegating to the repository.
pub struct EnforceModeHandler {
db: Arc<dyn AppRepo>,
app_config: Arc<ArcSwap<AppConfig>>,
audit_tx: broadcast::Sender<AuditEvent>,
/// Shared AtomicU8 cache: Monitor=0, MlOnly=1, Enforce=2.
enforce_cache: Arc<AtomicU8>,
}
impl EnforceModeHandler {
pub fn new(db: Arc<dyn AppRepo>, audit_tx: broadcast::Sender<AuditEvent>, enforce_cache: Arc<AtomicU8>) -> Self {
pub fn new(
db: Arc<dyn AppRepo>,
app_config: Arc<ArcSwap<AppConfig>>,
audit_tx: broadcast::Sender<AuditEvent>,
enforce_cache: Arc<AtomicU8>,
) -> Self {
Self {
db,
app_config,
audit_tx,
enforce_cache,
}
}
pub fn change_mode(&self, mode: String) -> Result<(), Error> {
self.db.set_setting("enforce_mode", &mode)?;
pub async fn change_mode(&self, mode: String) -> Result<(), Error> {
self.db.set_config_value("enforce_mode", &mode).await?;
self.enforce_cache.store(enforce_mode_to_u8(&mode), Ordering::SeqCst);
let mut cfg = (**self.app_config.load()).clone();
cfg.system.enforce_mode = mode.clone();
self.app_config.store(Arc::new(cfg));
log!(SystemLog::EnforceModeChanged(mode.clone()));
// Publish audit event for the mode change
@ -42,11 +54,8 @@ impl EnforceModeHandler {
Ok(())
}
pub fn get_mode(&self) -> Result<String, Error> {
match self.db.get_setting("enforce_mode")? {
Some(mode) => Ok(mode),
None => Ok(ENFORCE_MODE_MONITOR.to_string()),
}
pub fn get_mode(&self) -> String {
self.app_config.load().system.enforce_mode.clone()
}
}
@ -56,51 +65,54 @@ mod tests {
use crate::adapter::persistence::Database;
use crate::domain::common::config::constants::EVENT_CHANNEL_CAPACITY;
fn test_handler() -> EnforceModeHandler {
let db = Arc::new(Database::new(":memory:").unwrap()) as Arc<dyn AppRepo>;
async fn test_handler() -> EnforceModeHandler {
let db = Arc::new(Database::new(":memory:").await.unwrap());
let app_config = Arc::new(ArcSwap::from_pointee(
AppConfig::from_config_repo(db.as_ref()).await.unwrap(),
));
let cache = Arc::new(AtomicU8::new(0));
let (audit_tx, _rx) = broadcast::channel::<AuditEvent>(EVENT_CHANNEL_CAPACITY);
EnforceModeHandler::new(db, audit_tx, cache)
EnforceModeHandler::new(db as Arc<dyn AppRepo>, app_config, audit_tx, cache)
}
#[test]
fn test_default_mode_is_monitor() {
let handler = test_handler();
let mode = handler.get_mode().unwrap();
#[tokio::test]
async fn test_default_mode_is_monitor() {
let handler = test_handler().await;
let mode = handler.get_mode();
assert_eq!(mode, "monitor");
}
#[test]
fn test_change_to_enforce() {
let handler = test_handler();
handler.change_mode("enforce".into()).unwrap();
let mode = handler.get_mode().unwrap();
#[tokio::test]
async fn test_change_to_enforce() {
let handler = test_handler().await;
handler.change_mode("enforce".into()).await.unwrap();
let mode = handler.get_mode();
assert_eq!(mode, "enforce");
}
#[test]
fn test_change_back_to_monitor() {
let handler = test_handler();
handler.change_mode("enforce".into()).unwrap();
handler.change_mode("monitor".into()).unwrap();
let mode = handler.get_mode().unwrap();
#[tokio::test]
async fn test_change_back_to_monitor() {
let handler = test_handler().await;
handler.change_mode("enforce".into()).await.unwrap();
handler.change_mode("monitor".into()).await.unwrap();
let mode = handler.get_mode();
assert_eq!(mode, "monitor");
}
#[test]
fn test_change_to_ml_only() {
let handler = test_handler();
handler.change_mode("ml_only".into()).unwrap();
let mode = handler.get_mode().unwrap();
#[tokio::test]
async fn test_change_to_ml_only() {
let handler = test_handler().await;
handler.change_mode("ml_only".into()).await.unwrap();
let mode = handler.get_mode();
assert_eq!(mode, "ml_only");
}
#[test]
fn test_cycle_all_modes() {
let handler = test_handler();
#[tokio::test]
async fn test_cycle_all_modes() {
let handler = test_handler().await;
for mode_str in ["enforce", "ml_only", "monitor"] {
handler.change_mode(mode_str.into()).unwrap();
let mode = handler.get_mode().unwrap();
handler.change_mode(mode_str.into()).await.unwrap();
let mode = handler.get_mode();
assert_eq!(mode, mode_str);
}
}

View File

@ -6,15 +6,15 @@ use serde_json::Value;
use crate::domain::common::config::AppConfig;
use crate::domain::common::error::Error;
use crate::domain::common::error::misc::MiscError;
use crate::interface::config_repo::ConfigRepo;
use crate::interface::email_sender::EmailSenderFactory;
use crate::interface::notification::AlertNotifierFactory;
use crate::interface::secret_store::SecretStorePort;
use crate::interface::setting::SettingRepo;
/// Domain service for notification config (Telegram, SMTP).
/// Coordinates DB persistence and external service testing.
pub struct NotificationService {
notif: Arc<dyn SettingRepo + Send + Sync>,
notif: Arc<dyn ConfigRepo + Send + Sync>,
config: Arc<ArcSwap<AppConfig>>,
secrets: Arc<dyn SecretStorePort>,
alert_notifier_factory: Arc<dyn AlertNotifierFactory>,
@ -23,7 +23,7 @@ pub struct NotificationService {
impl NotificationService {
pub fn new(
notif: Arc<dyn SettingRepo + Send + Sync>,
notif: Arc<dyn ConfigRepo + Send + Sync>,
config: Arc<ArcSwap<AppConfig>>,
secrets: Arc<dyn SecretStorePort>,
alert_notifier_factory: Arc<dyn AlertNotifierFactory>,
@ -39,13 +39,13 @@ impl NotificationService {
}
/// Get Telegram config with redacted bot_token.
pub fn get_telegram_config(&self) -> Result<serde_json::Value, Error> {
match self.notif.get_notification_config("telegram")? {
pub async fn get_telegram_config(&self) -> Result<serde_json::Value, Error> {
match self.notif.get_notification_config("telegram").await? {
Some(json_str) => match serde_json::from_str::<serde_json::Value>(&json_str) {
Ok(mut config) => {
// Resolve the actual token for redaction display
let token = match config.get("bot_token").and_then(|t| t.as_str()) {
Some("__encrypted__") => self.secrets.get_secret("telegram_bot_token")?,
Some("__encrypted__") => self.secrets.get_secret("telegram_bot_token").await?,
Some(t) => Some(t.to_string()),
None => None,
};
@ -69,14 +69,14 @@ impl NotificationService {
/// Save Telegram bot_token + chat_id to DB.
/// The bot_token is stored encrypted in the secret store; the config JSON
/// holds the `"__encrypted__"` sentinel.
pub fn set_telegram_config(&self, bot_token: &str, chat_id: &str) -> Result<(), Error> {
self.secrets.set_secret("telegram_bot_token", bot_token)?;
pub async fn set_telegram_config(&self, bot_token: &str, chat_id: &str) -> Result<(), Error> {
self.secrets.set_secret("telegram_bot_token", bot_token).await?;
let config_json = serde_json::json!({
"bot_token": "__encrypted__",
"chat_id": chat_id,
})
.to_string();
self.notif.set_notification_config("telegram", &config_json)
self.notif.set_notification_config("telegram", &config_json).await
}
/// Send a test Telegram message using current config. The factory
@ -89,11 +89,12 @@ impl NotificationService {
}
/// Send a test email using current SMTP config.
pub fn test_smtp(&self) -> Result<String, Error> {
pub async fn test_smtp(&self) -> Result<String, Error> {
let smtp_cfg = self.config.load().notification.smtp.clone();
let smtp = self
.email_sender_factory
.build_smtp_sender(&smtp_cfg, Some(self.secrets.as_ref()))?
.build_smtp_sender(&smtp_cfg, Some(self.secrets.as_ref()))
.await?
.ok_or_else(|| {
MiscError::ValidationError(
"SMTP not configured. Set smtp_host, smtp_port, smtp_username, smtp_password first. \

View File

@ -4,7 +4,7 @@ use std::time::{SystemTime, UNIX_EPOCH};
use crate::core::inference::engine::Engine;
use crate::domain::data_plane::direction::Direction;
use crate::domain::data_plane::flow_stats::{
FlowPushPayload, FlowStatsEntry, FlowSubscription, FlowSummary, StatsSummary,
FlowPushPayload, FlowStatsEntry, FlowStatsLimits, FlowSubscription, FlowSummary, StatsSummary,
};
use crate::domain::detection::flow_tracker::FlowData;
@ -32,17 +32,28 @@ impl From<&FlowData> for FlowStatsEntry {
pub struct FlowStatistics {
engine: Arc<Engine>,
limits: FlowStatsLimits,
}
impl FlowStatistics {
pub fn new(engine: Arc<Engine>) -> Self {
Self { engine }
pub fn new(engine: Arc<Engine>, limits: FlowStatsLimits) -> Self {
Self { engine, limits }
}
pub fn get_all_flows(&self) -> Vec<FlowStatsEntry> {
self.collect_flows(Some(self.limits.result_limit(None)))
}
fn collect_flows(&self, limit: Option<usize>) -> Vec<FlowStatsEntry> {
let mut entries = Vec::new();
for tracker in self.engine.trackers() {
entries.extend(tracker.get_flow_stats(|flow| FlowStatsEntry::from(flow)));
if let Some(limit) = limit
&& entries.len() >= limit
{
entries.truncate(limit);
break;
}
}
entries
}
@ -53,7 +64,7 @@ impl FlowStatistics {
.map(|d| d.as_micros() as u64)
.unwrap_or(0);
let mut flows = self.get_all_flows();
let mut flows = self.collect_flows(None);
if let Some(dir) = &sub.direction {
flows.retain(|f| &f.direction == dir);
@ -65,10 +76,7 @@ impl FlowStatistics {
}
flows.sort_by_key(|f| std::cmp::Reverse(f.fwd_bytes + f.bwd_bytes));
if let Some(n) = sub.top_n {
flows.truncate(n.min(10000));
}
flows.truncate(self.limits.result_limit(sub.top_n));
flows
}
@ -120,7 +128,7 @@ impl FlowStatistics {
}
pub fn get_summary(&self) -> StatsSummary {
let flows = self.get_all_flows();
let flows = self.collect_flows(None);
let total_flows = flows.len();
let total_bytes: u64 = flows.iter().map(|f| f.fwd_bytes + f.bwd_bytes).sum();
let total_packets: usize = flows.iter().map(|f| f.fwd_packets + f.bwd_packets).sum();

View File

@ -32,15 +32,24 @@ impl AclService {
}
}
pub fn add_ipv4(&self, direction: FlowDirection, list_type: ListType, address: SocketAddrV4) -> Result<(), Error> {
pub async fn add_ipv4(
&self,
direction: FlowDirection,
list_type: ListType,
address: SocketAddrV4,
) -> Result<(), Error> {
self.access_control.add_ipv4_list(direction, list_type, address)?;
if let Err(e) = self.db.insert_acl_rule(
4,
direction.as_str(),
list_type.as_str(),
&address.ip().to_string(),
address.port(),
) {
if let Err(e) = self
.db
.insert_acl_rule(
4,
direction.as_str(),
list_type.as_str(),
&address.ip().to_string(),
address.port(),
)
.await
{
if let Err(rollback_err) = self.access_control.remove_ipv4_list(direction, list_type, address) {
log!(EbpfError::RollbackFailed(rollback_err));
}
@ -49,15 +58,24 @@ impl AclService {
Ok(())
}
pub fn add_ipv6(&self, direction: FlowDirection, list_type: ListType, address: SocketAddrV6) -> Result<(), Error> {
pub async fn add_ipv6(
&self,
direction: FlowDirection,
list_type: ListType,
address: SocketAddrV6,
) -> Result<(), Error> {
self.access_control.add_ipv6_list(direction, list_type, address)?;
if let Err(e) = self.db.insert_acl_rule(
6,
direction.as_str(),
list_type.as_str(),
&address.ip().to_string(),
address.port(),
) {
if let Err(e) = self
.db
.insert_acl_rule(
6,
direction.as_str(),
list_type.as_str(),
&address.ip().to_string(),
address.port(),
)
.await
{
if let Err(rollback_err) = self.access_control.remove_ipv6_list(direction, list_type, address) {
log!(EbpfError::RollbackFailed(rollback_err));
}
@ -66,20 +84,24 @@ impl AclService {
Ok(())
}
pub fn remove_ipv4(
pub async fn remove_ipv4(
&self,
direction: FlowDirection,
list_type: ListType,
address: SocketAddrV4,
) -> Result<(), Error> {
self.access_control.remove_ipv4_list(direction, list_type, address)?;
if let Err(e) = self.db.delete_acl_rule(
4,
direction.as_str(),
list_type.as_str(),
&address.ip().to_string(),
address.port(),
) {
if let Err(e) = self
.db
.delete_acl_rule(
4,
direction.as_str(),
list_type.as_str(),
&address.ip().to_string(),
address.port(),
)
.await
{
if let Err(rollback_err) = self.access_control.add_ipv4_list(direction, list_type, address) {
log!(EbpfError::RollbackFailed(rollback_err));
}
@ -88,20 +110,24 @@ impl AclService {
Ok(())
}
pub fn remove_ipv6(
pub async fn remove_ipv6(
&self,
direction: FlowDirection,
list_type: ListType,
address: SocketAddrV6,
) -> Result<(), Error> {
self.access_control.remove_ipv6_list(direction, list_type, address)?;
if let Err(e) = self.db.delete_acl_rule(
6,
direction.as_str(),
list_type.as_str(),
&address.ip().to_string(),
address.port(),
) {
if let Err(e) = self
.db
.delete_acl_rule(
6,
direction.as_str(),
list_type.as_str(),
&address.ip().to_string(),
address.port(),
)
.await
{
if let Err(rollback_err) = self.access_control.add_ipv6_list(direction, list_type, address) {
log!(EbpfError::RollbackFailed(rollback_err));
}
@ -110,20 +136,24 @@ impl AclService {
Ok(())
}
pub fn block_geo_countries(&self, codes: &[String]) -> Result<u64, Error> {
pub async fn block_geo_countries(&self, codes: &[String]) -> Result<u64, Error> {
let total = self.geo_block.block_countries(codes)?;
if let Err(e) = codes.iter().try_for_each(|code| self.db.insert_geo_country(code)) {
let _ = self.geo_block.unblock_countries(codes);
return Err(e);
for code in codes {
if let Err(e) = self.db.insert_geo_country(code).await {
let _ = self.geo_block.unblock_countries(codes);
return Err(e);
}
}
Ok(total)
}
pub fn unblock_geo_countries(&self, codes: &[String]) -> Result<u64, Error> {
pub async fn unblock_geo_countries(&self, codes: &[String]) -> Result<u64, Error> {
let total = self.geo_block.unblock_countries(codes)?;
if let Err(e) = codes.iter().try_for_each(|code| self.db.delete_geo_country(code)) {
let _ = self.geo_block.block_countries(codes);
return Err(e);
for code in codes {
if let Err(e) = self.db.delete_geo_country(code).await {
let _ = self.geo_block.block_countries(codes);
return Err(e);
}
}
Ok(total)
}

View File

@ -25,6 +25,10 @@ impl DnsFilter {
Ok(())
}
pub fn validate_domain(&self, domain: &str) -> Result<(), Error> {
domain_to_wire_format(domain).map(|_| ())
}
pub fn remove_domain(&self, domain: &str) -> Result<(), Error> {
let name = domain_to_wire_format(domain)?;
self.blacklist.remove(&name);
@ -185,6 +189,9 @@ impl DnsFilter {
}
impl DnsFilterPort for DnsFilter {
fn validate_domain(&self, domain: &str) -> Result<(), Error> {
self.validate_domain(domain)
}
fn add_domain(&self, domain: &str) -> Result<(), Error> {
self.add_domain(domain)
}

View File

@ -5,19 +5,24 @@ use arc_swap::ArcSwap;
use crate::domain::common::config::AppConfig;
use crate::domain::common::error::Error;
use crate::domain::common::error::misc::MiscError;
use crate::interface::app_repo::AppRepo;
use crate::interface::dns_filter_api::DnsFilterPort;
use crate::interface::enforcement::EnforcementRepo;
/// Domain service that coordinates DNS filter changes between DB and in-memory service.
/// Write order: eBPF/in-memory first, then DB — if eBPF fails, DB remains clean.
/// Domains are validated up front, runtime state is changed first, and DB batch
/// failure rolls runtime back so persisted and live policy do not drift.
pub struct DnsFilterService {
db: Arc<dyn AppRepo>,
db: Arc<dyn EnforcementRepo>,
dns_filter: Arc<dyn DnsFilterPort>,
config: Arc<ArcSwap<AppConfig>>,
}
impl DnsFilterService {
pub fn new(db: Arc<dyn AppRepo>, dns_filter: Arc<dyn DnsFilterPort>, config: Arc<ArcSwap<AppConfig>>) -> Self {
pub fn new(
db: Arc<dyn EnforcementRepo>,
dns_filter: Arc<dyn DnsFilterPort>,
config: Arc<ArcSwap<AppConfig>>,
) -> Self {
Self { db, dns_filter, config }
}
@ -25,7 +30,7 @@ impl DnsFilterService {
self.dns_filter.list_domains()
}
pub fn add_domains(&self, domains: &[String]) -> Result<usize, Error> {
pub async fn add_domains(&self, domains: &[String]) -> Result<usize, Error> {
let max_domains = self.config.load().dns_filter.max_domains_per_request;
if domains.len() > max_domains {
Err(MiscError::ValidationError(format!(
@ -33,26 +38,151 @@ impl DnsFilterService {
max_domains
)))?;
}
// eBPF first
self.validate_domains(domains)?;
let mut applied: Vec<&String> = Vec::new();
for domain in domains {
self.dns_filter.add_domain(domain)?;
if let Err(err) = self.dns_filter.add_domain(domain) {
self.rollback_added(&applied);
return Err(err);
}
applied.push(domain);
}
// Then DB
for domain in domains {
self.db.insert_dns_domain(domain)?;
if let Err(err) = self.db.insert_dns_domains(domains).await {
self.rollback_added(&applied);
return Err(err);
}
Ok(domains.len())
}
pub fn remove_domains(&self, domains: &[String]) -> Result<usize, Error> {
// eBPF first
pub async fn remove_domains(&self, domains: &[String]) -> Result<usize, Error> {
self.validate_domains(domains)?;
let mut applied: Vec<&String> = Vec::new();
for domain in domains {
self.dns_filter.remove_domain(domain)?;
if let Err(err) = self.dns_filter.remove_domain(domain) {
self.rollback_removed(&applied);
return Err(err);
}
applied.push(domain);
}
// Then DB
for domain in domains {
self.db.delete_dns_domain(domain)?;
if let Err(err) = self.db.delete_dns_domains(domains).await {
self.rollback_removed(&applied);
return Err(err);
}
Ok(domains.len())
}
fn validate_domains(&self, domains: &[String]) -> Result<(), Error> {
for domain in domains {
self.dns_filter.validate_domain(domain)?;
}
Ok(())
}
fn rollback_added(&self, domains: &[&String]) {
for domain in domains {
let _ = self.dns_filter.remove_domain(domain);
}
}
fn rollback_removed(&self, domains: &[&String]) {
for domain in domains {
let _ = self.dns_filter.add_domain(domain);
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Mutex;
use async_trait::async_trait;
use super::*;
use crate::adapter::persistence::Database;
use crate::core::data_plane::dns_filter::DnsFilter;
struct FailingDnsRepo {
fail_insert: bool,
fail_delete: bool,
domains: Mutex<Vec<String>>,
}
impl FailingDnsRepo {
fn new(fail_insert: bool, fail_delete: bool) -> Self {
Self {
fail_insert,
fail_delete,
domains: Mutex::new(Vec::new()),
}
}
}
#[async_trait]
impl EnforcementRepo for FailingDnsRepo {
async fn set_rate_limit(&self, _key: &str, _value: u64) -> Result<(), Error> {
Ok(())
}
async fn insert_dns_domains(&self, domains: &[String]) -> Result<(), Error> {
if self.fail_insert {
Err(MiscError::ValidationError("forced insert failure".to_string()))?;
}
self.domains.lock().unwrap().extend_from_slice(domains);
Ok(())
}
async fn delete_dns_domains(&self, domains: &[String]) -> Result<(), Error> {
if self.fail_delete {
Err(MiscError::ValidationError("forced delete failure".to_string()))?;
}
self.domains.lock().unwrap().retain(|d| !domains.contains(d));
Ok(())
}
async fn insert_geo_country(&self, _code: &str) -> Result<(), Error> {
Ok(())
}
async fn delete_geo_country(&self, _code: &str) -> Result<(), Error> {
Ok(())
}
}
async fn test_config() -> Arc<ArcSwap<AppConfig>> {
let db = Database::new(":memory:").await.expect("test db");
Arc::new(ArcSwap::from_pointee(
AppConfig::from_config_repo(&db).await.expect("test config"),
))
}
#[tokio::test]
async fn add_domains_rolls_back_runtime_when_db_batch_fails() {
let filter = Arc::new(DnsFilter::new());
let service = DnsFilterService::new(
Arc::new(FailingDnsRepo::new(true, false)),
filter.clone(),
test_config().await,
);
let result = service.add_domains(&["example.com".to_string()]).await;
assert!(result.is_err());
assert!(filter.list_domains().is_empty());
}
#[tokio::test]
async fn remove_domains_restores_runtime_when_db_batch_fails() {
let filter = Arc::new(DnsFilter::new());
filter.add_domain("example.com").unwrap();
let service = DnsFilterService::new(
Arc::new(FailingDnsRepo::new(false, true)),
filter.clone(),
test_config().await,
);
let result = service.remove_domains(&["example.com".to_string()]).await;
assert!(result.is_err());
assert!(filter.list_domains().contains(&"example.com".to_string()));
}
}

View File

@ -20,25 +20,25 @@ impl RateLimitService {
self.config.as_ref()
}
pub fn update(&self, settings: &RateLimitSettings) -> Result<(), Error> {
pub async fn update(&self, settings: &RateLimitSettings) -> Result<(), Error> {
if let Some(v) = settings.packet_rate {
self.db.set_rate_limit("packet_rate", v)?;
self.db.set_rate_limit("packet_rate", v).await?;
self.config.set_packet_rate(v)?;
}
if let Some(v) = settings.syn_rate {
self.db.set_rate_limit("syn_rate", v)?;
self.db.set_rate_limit("syn_rate", v).await?;
self.config.set_syn_rate(v)?;
}
if let Some(v) = settings.udp_rate {
self.db.set_rate_limit("udp_rate", v)?;
self.db.set_rate_limit("udp_rate", v).await?;
self.config.set_udp_rate(v)?;
}
if let Some(v) = settings.dns_rate {
self.db.set_rate_limit("dns_rate", v)?;
self.db.set_rate_limit("dns_rate", v).await?;
self.config.set_dns_rate(v)?;
}
if let Some(v) = settings.window_ns {
self.db.set_rate_limit("window_ns", v)?;
self.db.set_rate_limit("window_ns", v).await?;
self.config.set_window_ns(v)?;
}
Ok(())

View File

@ -34,6 +34,7 @@ impl BeaconingDetector {
beaconing.min_observations,
beaconing.cv_threshold,
beaconing.max_cache_entries,
beaconing.max_timestamps_per_flow,
beaconing.expiry_secs,
beaconing.alert_cooldown_secs,
),
@ -78,8 +79,6 @@ impl BeaconingDetector {
type FlowTuple = (String, String, u16);
const MAX_TIMESTAMPS_PER_FLOW: usize = 100;
struct CachedFlow {
timestamps: Vec<Instant>,
last_alerted: Option<Instant>,
@ -90,6 +89,7 @@ pub struct BeaconingState {
min_observations: usize,
cv_threshold: f64,
max_cache_entries: usize,
max_timestamps_per_flow: usize,
expiry_secs: u64,
alert_cooldown_secs: u64,
}
@ -99,6 +99,7 @@ impl BeaconingState {
min_observations: usize,
cv_threshold: f64,
max_cache_entries: usize,
max_timestamps_per_flow: usize,
expiry_secs: u64,
alert_cooldown_secs: u64,
) -> Self {
@ -107,6 +108,7 @@ impl BeaconingState {
min_observations,
cv_threshold,
max_cache_entries,
max_timestamps_per_flow,
expiry_secs,
alert_cooldown_secs,
}
@ -123,8 +125,8 @@ impl BeaconingState {
entry.timestamps.push(now);
if entry.timestamps.len() > MAX_TIMESTAMPS_PER_FLOW {
let excess = entry.timestamps.len() - MAX_TIMESTAMPS_PER_FLOW;
if entry.timestamps.len() > self.max_timestamps_per_flow {
let excess = entry.timestamps.len() - self.max_timestamps_per_flow;
entry.timestamps.drain(..excess);
}
}
@ -280,7 +282,7 @@ mod tests {
#[test]
fn beaconing_state_detects_periodic_flows() {
let state = BeaconingState::new(5, 0.3, 50_000, 3600, 120);
let state = BeaconingState::new(5, 0.3, 50_000, 100, 3600, 120);
let base = Instant::now();
let key = ("10.0.0.1".to_string(), "1.2.3.4".to_string(), 443_u16);
state.flow_cache.insert(
@ -295,4 +297,33 @@ mod tests {
assert_eq!(events[0].source, DetectionSource::Beaconing);
assert_eq!(events[0].attack_type, CanonicalAttackType::C2Beacon.as_str());
}
#[test]
fn record_flow_respects_configured_timestamp_cap() {
let state = BeaconingState::new(1, 0.3, 50_000, 3, 3600, 120);
let alert = crate::domain::detection::ml_detection::AlertMessage {
timestamp: 0,
flow_key: "10.0.0.1:12345-1.2.3.4:443".to_string(),
src_ip: "10.0.0.1".to_string(),
dst_ip: "1.2.3.4".to_string(),
src_port: 12345,
dst_port: 443,
protocol: 6,
is_attack: true,
attack_type: Some("c2_beacon".to_string()),
confidence: 0.9,
packet_count: 10,
flow_duration_us: 1000,
ae_score: 0.0,
anomaly_score: 0.0,
c2_score: 0.0,
};
for _ in 0..5 {
state.record_flow(&alert);
}
let key = ("10.0.0.1".to_string(), "1.2.3.4".to_string(), 443_u16);
assert_eq!(state.flow_cache.get(&key).unwrap().timestamps.len(), 3);
}
}

View File

@ -78,16 +78,17 @@ impl DetectionOrchestrator {
) -> Self {
let cfg = app_config.load();
let fusion = &cfg.detection.fusion;
let max_dedup = NonZero::new(fusion.max_dedup_entries.max(1)).unwrap_or(NonZero::<usize>::MIN);
let source_count_max_entries = nonzero_cache_size(fusion.source_count_max_entries);
let repeat_tracker_max_entries = nonzero_cache_size(fusion.repeat_tracker_max_entries);
let max_dedup = nonzero_cache_size(fusion.max_dedup_entries);
Self {
rx,
threat_tx,
audit_tx,
geoip,
metrics,
// SAFETY: NonZero::new on non-zero literals.
src_ip_counts: LruCache::new(NonZero::new(10_000).unwrap()),
repeat_tracker: LruCache::new(NonZero::new(5_000).unwrap()),
src_ip_counts: LruCache::new(source_count_max_entries),
repeat_tracker: LruCache::new(repeat_tracker_max_entries),
dedup: LruCache::new(max_dedup),
dedup_window: Duration::from_secs(fusion.dedup_window_secs),
repeat_offender_window: Duration::from_secs(fusion.repeat_offender_window_secs),
@ -355,6 +356,10 @@ impl DetectionOrchestrator {
}
}
fn nonzero_cache_size(value: usize) -> NonZero<usize> {
NonZero::new(value.max(1)).unwrap_or(NonZero::<usize>::MIN)
}
pub async fn bridge_ml_to_detection(mut rx: broadcast::Receiver<AlertMessage>, tx: mpsc::Sender<DetectionEvent>) {
log!(DetectionLog::MlBridgeStarted);

View File

@ -34,6 +34,7 @@ pub struct UserProfile {
pub groups: Vec<String>,
}
#[derive(Debug)]
pub enum LoginError {
Locked { retry_after_secs: u64 },
InvalidCredentials,
@ -53,18 +54,18 @@ impl AuthService {
Self { db, jwt }
}
pub fn login(&self, username: &str, raw_password: &str) -> Result<LoginResult, LoginError> {
if let Ok(Some(remaining)) = self.db.check_login_locked(username) {
pub async fn login(&self, username: &str, raw_password: &str) -> Result<LoginResult, LoginError> {
if let Ok(Some(remaining)) = self.db.check_login_locked(username).await {
return Err(LoginError::Locked {
retry_after_secs: remaining,
});
}
let user = match self.db.find_user(username) {
let user = match self.db.find_user(username).await {
Ok(Some(u)) => u,
_ => {
let _ = password::verify_password(raw_password, DUMMY_HASH);
if let Err(e) = self.db.record_login_failure(username) {
if let Err(e) = self.db.record_login_failure(username).await {
log!(AuthError::LoginFailureTrackingError(e));
}
return Err(LoginError::InvalidCredentials);
@ -74,19 +75,19 @@ impl AuthService {
match password::verify_password(raw_password, &user.password_hash) {
Ok(true) => {}
_ => {
if let Err(e) = self.db.record_login_failure(username) {
if let Err(e) = self.db.record_login_failure(username).await {
log!(AuthError::LoginFailureTrackingError(e));
}
return Err(LoginError::InvalidCredentials);
}
}
if let Err(e) = self.db.clear_login_failures(username) {
if let Err(e) = self.db.clear_login_failures(username).await {
log!(AuthError::LoginClearError(e));
}
let permissions = self.db.list_user_permissions(user.id).unwrap_or_default();
let role = self.derive_role(user.id);
let permissions = self.db.list_user_permissions(user.id).await.unwrap_or_default();
let role = self.derive_role(user.id).await;
let token = self
.jwt
@ -100,7 +101,7 @@ impl AuthService {
})
}
pub fn register(
pub async fn register(
&self,
username: &str,
raw_password: &str,
@ -122,12 +123,13 @@ impl AuthService {
let new_id = self
.db
.insert_user(username, &hash, role, false)
.await
.map_err(RegisterError::Conflict)?;
let default_group = if role == ROLE_ADMIN { GROUP_ADMIN } else { GROUP_VIEWER };
if let Ok(groups) = self.db.list_user_groups()
if let Ok(groups) = self.db.list_user_groups().await
&& let Some(g) = groups.into_iter().find(|g| g.name == default_group)
&& let Err(e) = self.db.set_user_groups(new_id, &[g.id])
&& let Err(e) = self.db.set_user_groups(new_id, &[g.id]).await
{
log!(AuthError::GroupAssignmentFailed(e));
}
@ -135,15 +137,15 @@ impl AuthService {
Ok(new_id)
}
pub fn user_profile(&self, user_id: i64, username: &str) -> UserProfile {
let groups_raw = self.db.list_groups_for_user(user_id).unwrap_or_default();
pub async fn user_profile(&self, user_id: i64, username: &str) -> UserProfile {
let groups_raw = self.db.list_groups_for_user(user_id).await.unwrap_or_default();
let group_names: Vec<String> = groups_raw.iter().map(|g| g.name.clone()).collect();
let role = if group_names.iter().any(|n| n == GROUP_ADMIN) {
ROLE_ADMIN.to_string()
} else {
ROLE_VIEWER.to_string()
};
let permissions = self.db.list_user_permissions(user_id).unwrap_or_default();
let permissions = self.db.list_user_permissions(user_id).await.unwrap_or_default();
UserProfile {
id: user_id,
@ -154,8 +156,8 @@ impl AuthService {
}
}
pub fn derive_role(&self, user_id: i64) -> String {
let groups = self.db.list_groups_for_user(user_id).unwrap_or_default();
pub async fn derive_role(&self, user_id: i64) -> String {
let groups = self.db.list_groups_for_user(user_id).await.unwrap_or_default();
if groups.iter().any(|g| g.name == GROUP_ADMIN) {
ROLE_ADMIN.to_string()
} else {
@ -163,3 +165,62 @@ impl AuthService {
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::*;
use crate::adapter::http::jwt::JwtService;
use crate::adapter::persistence::Database;
use crate::domain::identity::password;
use crate::infrastructure::secret_store::SecretStore;
use crate::interface::secret_store::SecretStorePort;
async fn auth_fixture() -> (Arc<Database>, Arc<JwtService>, AuthService) {
let db = Arc::new(Database::new(":memory:").await.expect("test db"));
let secrets: Arc<dyn SecretStorePort> = Arc::new(SecretStore::new(db.clone()));
let jwt = Arc::new(JwtService::new(&secrets, 24).expect("jwt"));
let auth = AuthService::new(db.clone() as Arc<dyn AppRepo>, jwt.clone());
(db, jwt, auth)
}
async fn create_viewer(db: &Database, username: &str, password: &str) -> i64 {
let hash = password::hash_password(password).expect("hash");
let user_id = db
.insert_user(username, &hash, ROLE_VIEWER, false)
.await
.expect("insert user");
let viewer_group = db
.list_user_groups()
.await
.expect("groups")
.into_iter()
.find(|g| g.name == GROUP_VIEWER)
.expect("viewer group");
db.set_user_groups(user_id, &[viewer_group.id])
.await
.expect("assign viewer group");
user_id
}
#[tokio::test]
async fn relogin_token_permissions_follow_role_promotion_and_demotion() {
let (db, jwt, auth) = auth_fixture().await;
let user_id = create_viewer(&db, "alice", "Correct Horse 123!").await;
db.update_user_role(user_id, ROLE_ADMIN).await.expect("promote");
let promoted = auth.login("alice", "Correct Horse 123!").await.expect("login");
let promoted_claims = jwt.validate_token(&promoted.token).expect("promoted token");
assert_eq!(promoted.role, ROLE_ADMIN);
assert_eq!(promoted_claims.role, ROLE_ADMIN);
assert!(promoted_claims.permissions.contains(&"users:admin".to_string()));
db.update_user_role(user_id, ROLE_VIEWER).await.expect("demote");
let demoted = auth.login("alice", "Correct Horse 123!").await.expect("login");
let demoted_claims = jwt.validate_token(&demoted.token).expect("demoted token");
assert_eq!(demoted.role, ROLE_VIEWER);
assert_eq!(demoted_claims.role, ROLE_VIEWER);
assert!(!demoted_claims.permissions.contains(&"users:admin".to_string()));
}
}

View File

@ -1,102 +1,55 @@
use chrono::Local;
use super::html::escape;
use super::report_data_builder::build_report_data;
use crate::domain::common::error::Error;
use crate::interface::setting::SettingRepo;
use crate::domain::report::data::ReportData;
use crate::interface::report_snapshot::ReportSnapshotRepo;
/// Generate an HTML weekly report email body.
///
/// Reads aggregated statistics from the Database settings table and formats
/// them into a self-contained HTML email. Keys consumed:
/// - `weekly_threats_count`
/// - `weekly_top_ips` (JSON array of `{ "ip": "...", "count": N }`)
/// - `weekly_threat_breakdown` (JSON object `{ "type": count, ... }`)
/// - `weekly_bandwidth_bytes`
/// - `weekly_system_health` (JSON object with cpu, memory, disk fields)
///
/// If a key is missing the report uses empty/zero defaults.
pub fn generate_weekly_report(db: &dyn SettingRepo) -> Result<String, Error> {
let threats_count = db
.get_setting("weekly_threats_count")?
.unwrap_or_else(|| "0".to_string());
let top_ips_json = db.get_setting("weekly_top_ips")?.unwrap_or_else(|| "[]".to_string());
let threat_breakdown_json = db
.get_setting("weekly_threat_breakdown")?
.unwrap_or_else(|| "{}".to_string());
let bandwidth = db
.get_setting("weekly_bandwidth_bytes")?
.unwrap_or_else(|| "0".to_string());
let health_json = db.get_setting("weekly_system_health")?.unwrap_or_else(|| {
serde_json::json!({
"cpu_percent": 0.0,
"memory_percent": 0.0,
"disk_percent": 0.0
})
.to_string()
});
// ── Parse JSON blobs ───────────────────────────────────────────────
let top_ips: Vec<serde_json::Value> = serde_json::from_str(&top_ips_json).unwrap_or_default();
let threat_breakdown: serde_json::Map<String, serde_json::Value> =
serde_json::from_str(&threat_breakdown_json).unwrap_or_default();
let health: serde_json::Value = serde_json::from_str(&health_json).unwrap_or_default();
// ── Build HTML ─────────────────────────────────────────────────────
let bandwidth_mb = bandwidth.parse::<f64>().unwrap_or(0.0) / 1_048_576.0;
/// The email path shares the typed `ReportData` builder used by the report
/// API/file renderer. This keeps parsing/defaulting in one place and makes
/// HTML escaping a renderer invariant instead of a per-field caller detail.
pub async fn generate_weekly_report(db: &dyn ReportSnapshotRepo) -> Result<String, Error> {
let data = build_report_data(db).await?;
Ok(render_weekly_email(&data))
}
fn render_weekly_email(data: &ReportData) -> String {
let mut top_ips_rows = String::new();
for (i, entry) in top_ips.iter().enumerate().take(5) {
let ip = entry["ip"].as_str().unwrap_or("unknown");
let count = entry["count"].as_u64().unwrap_or(0);
for (i, entry) in data.top_blocked_ips.iter().enumerate().take(5) {
top_ips_rows.push_str(&format!(
"<tr><td style=\"padding:6px 12px;border-bottom:1px solid #e0e0e0;\">{}</td>\
<td style=\"padding:6px 12px;border-bottom:1px solid #e0e0e0;\">{}</td>\
<td style=\"padding:6px 12px;border-bottom:1px solid #e0e0e0;text-align:right;\">{}</td></tr>",
i + 1,
ip,
count,
escape(&entry.ip),
entry.count,
));
}
let mut breakdown_rows = String::new();
for (threat_type, count) in &threat_breakdown {
let n = count.as_u64().unwrap_or(0);
for item in &data.threat_breakdown {
breakdown_rows.push_str(&format!(
"<tr><td style=\"padding:6px 12px;border-bottom:1px solid #e0e0e0;\">{}</td>\
<td style=\"padding:6px 12px;border-bottom:1px solid #e0e0e0;text-align:right;\">{}</td></tr>",
threat_type, n,
escape(&item.threat_type),
item.count,
));
}
let cpu = health["cpu_percent"].as_f64().unwrap_or(0.0);
let mem = health["memory_percent"].as_f64().unwrap_or(0.0);
let disk = health["disk_percent"].as_f64().unwrap_or(0.0);
let now = Local::now().format("%Y-%m-%d %H:%M");
let html = format!(
format!(
r#"<!DOCTYPE html>
<html>
<head><meta charset="utf-8"></head>
<body style="font-family:Arial,Helvetica,sans-serif;background:#f4f6f9;margin:0;padding:20px;">
<div style="max-width:640px;margin:0 auto;background:#ffffff;border-radius:8px;overflow:hidden;box-shadow:0 2px 8px rgba(0,0,0,0.08);">
<!-- Header -->
<div style="background:#1a237e;color:#ffffff;padding:24px 32px;">
<h1 style="margin:0;font-size:22px;">NetGuardia Weekly Report</h1>
<p style="margin:6px 0 0;font-size:13px;opacity:0.85;">Generated {now}</p>
<p style="margin:6px 0 0;font-size:13px;opacity:0.85;">Report Period: {period} &mdash; Generated: {generated_at}</p>
</div>
<div style="padding:24px 32px;">
<!-- Threats summary -->
<h2 style="font-size:16px;color:#1a237e;border-bottom:2px solid #1a237e;padding-bottom:6px;">
Threat Summary
</h2>
@ -104,7 +57,6 @@ pub fn generate_weekly_report(db: &dyn SettingRepo) -> Result<String, Error> {
<span style="font-size:14px;font-weight:normal;color:#666;"> threats detected this week</span>
</p>
<!-- Top blocked IPs -->
<h2 style="font-size:16px;color:#1a237e;border-bottom:2px solid #1a237e;padding-bottom:6px;margin-top:24px;">
Top 5 Blocked IPs
</h2>
@ -119,7 +71,6 @@ pub fn generate_weekly_report(db: &dyn SettingRepo) -> Result<String, Error> {
<tbody>{top_ips_rows}</tbody>
</table>
<!-- Threat breakdown -->
<h2 style="font-size:16px;color:#1a237e;border-bottom:2px solid #1a237e;padding-bottom:6px;margin-top:24px;">
Threat Type Breakdown
</h2>
@ -133,42 +84,73 @@ pub fn generate_weekly_report(db: &dyn SettingRepo) -> Result<String, Error> {
<tbody>{breakdown_rows}</tbody>
</table>
<!-- Bandwidth -->
<h2 style="font-size:16px;color:#1a237e;border-bottom:2px solid #1a237e;padding-bottom:6px;margin-top:24px;">
Bandwidth
</h2>
<p style="font-size:14px;">{bandwidth_mb:.2} MB processed this week</p>
<!-- System Health -->
<h2 style="font-size:16px;color:#1a237e;border-bottom:2px solid #1a237e;padding-bottom:6px;margin-top:24px;">
System Health
</h2>
<table style="width:100%;border-collapse:collapse;font-size:14px;">
<tr>
<td style="padding:6px 12px;">CPU</td>
<td style="padding:6px 12px;text-align:right;">{cpu:.1}%</td>
</tr>
<tr>
<td style="padding:6px 12px;">Memory</td>
<td style="padding:6px 12px;text-align:right;">{mem:.1}%</td>
</tr>
<tr>
<td style="padding:6px 12px;">Disk</td>
<td style="padding:6px 12px;text-align:right;">{disk:.1}%</td>
</tr>
<tr><td style="padding:6px 12px;">CPU</td><td style="padding:6px 12px;text-align:right;">{cpu:.1}%</td></tr>
<tr><td style="padding:6px 12px;">Memory</td><td style="padding:6px 12px;text-align:right;">{mem:.1}%</td></tr>
<tr><td style="padding:6px 12px;">Disk</td><td style="padding:6px 12px;text-align:right;">{disk:.1}%</td></tr>
</table>
</div>
<!-- Footer -->
<div style="background:#f0f0f0;padding:16px 32px;font-size:12px;color:#888;text-align:center;">
NetGuardia &mdash; Automated Weekly Report
</div>
</div>
</body>
</html>"#
);
Ok(html)
</html>"#,
period = escape(&data.period),
generated_at = escape(&data.generated_at),
threats_count = data.executive_summary.total_threats,
top_ips_rows = top_ips_rows,
breakdown_rows = breakdown_rows,
cpu = data.system_health.avg_cpu_percent,
mem = data.system_health.avg_memory_percent,
disk = data.system_health.disk_usage_percent,
)
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use super::*;
struct FakeSettings {
values: HashMap<String, String>,
}
#[async_trait::async_trait]
impl ReportSnapshotRepo for FakeSettings {
async fn get_report_snapshot(&self, key: &str) -> Result<Option<String>, Error> {
Ok(self.values.get(key).cloned())
}
async fn set_report_snapshot(&self, _key: &str, _value: &str) -> Result<(), Error> {
Ok(())
}
}
#[tokio::test]
async fn weekly_report_escapes_snapshot_derived_html() {
let mut values = HashMap::new();
values.insert(
"weekly_top_ips".to_string(),
r#"[{"ip":"<img src=x onerror=alert(1)>","count":3,"country":"N/A"}]"#.to_string(),
);
values.insert(
"weekly_threat_breakdown".to_string(),
r#"{"\"><script>alert(1)</script>":2}"#.to_string(),
);
let repo = FakeSettings { values };
let html = generate_weekly_report(&repo).await.unwrap();
assert!(!html.contains("<img src=x onerror=alert(1)>"));
assert!(!html.contains("<script>alert(1)</script>"));
assert!(html.contains("&lt;img src=x onerror=alert(1)&gt;"));
assert!(html.contains("&quot;&gt;&lt;script&gt;alert(1)&lt;/script&gt;"));
}
}

View File

@ -10,11 +10,11 @@ use super::email_report as report;
use crate::domain::common::config::AppConfig;
use crate::domain::common::log::system::SystemLog;
use crate::interface::email_sender::EmailSenderFactory;
use crate::interface::report_snapshot::ReportSnapshotRepo;
use crate::interface::secret_store::SecretStorePort;
use crate::interface::setting::SettingRepo;
pub struct ReportScheduler {
db: Arc<dyn SettingRepo + Send + Sync>,
db: Arc<dyn ReportSnapshotRepo>,
config: Arc<ArcSwap<AppConfig>>,
secrets: Option<Arc<dyn SecretStorePort>>,
email_sender_factory: Arc<dyn EmailSenderFactory>,
@ -22,7 +22,7 @@ pub struct ReportScheduler {
impl ReportScheduler {
pub fn new(
db: Arc<dyn SettingRepo + Send + Sync>,
db: Arc<dyn ReportSnapshotRepo>,
config: Arc<ArcSwap<AppConfig>>,
secrets: Option<Arc<dyn SecretStorePort>>,
email_sender_factory: Arc<dyn EmailSenderFactory>,
@ -55,7 +55,10 @@ impl ReportScheduler {
let smtp_cfg = config.load().notification.smtp.clone();
let smtp = match email_sender_factory.build_smtp_sender(&smtp_cfg, secrets.as_deref()) {
let smtp = match email_sender_factory
.build_smtp_sender(&smtp_cfg, secrets.as_deref())
.await
{
Ok(Some(client)) => client,
Ok(None) => {
log!(SystemLog::SmtpNotConfigured);
@ -73,7 +76,7 @@ impl ReportScheduler {
continue;
}
let html = match report::generate_weekly_report(&*db) {
let html = match report::generate_weekly_report(&*db).await {
Ok(h) => h,
Err(e) => {
log!(SystemLog::WeeklyReportGenerationFailed(e.to_string()));

View File

@ -0,0 +1,8 @@
pub fn escape(input: &str) -> String {
input
.replace('&', "&amp;")
.replace('<', "&lt;")
.replace('>', "&gt;")
.replace('"', "&quot;")
.replace('\'', "&#39;")
}

View File

@ -1,5 +1,6 @@
pub mod email_report;
pub mod email_scheduler;
pub mod html;
pub mod report_data_builder;
pub mod report_engine;
pub mod stats_aggregator;

View File

@ -4,9 +4,9 @@ use crate::domain::common::error::Error;
use crate::domain::report::data::{
BlockedIpItem, ExecutiveSummary, GeoItem, ReportData, SoarActivity, SystemHealthSummary, ThreatBreakdownItem,
};
use crate::interface::setting::SettingRepo;
use crate::interface::report_snapshot::ReportSnapshotRepo;
pub fn build_report_data(db: &dyn SettingRepo) -> Result<ReportData, Error> {
pub async fn build_report_data(db: &dyn ReportSnapshotRepo) -> Result<ReportData, Error> {
let now = Local::now();
let period = format!(
"{} — {}",
@ -15,12 +15,14 @@ pub fn build_report_data(db: &dyn SettingRepo) -> Result<ReportData, Error> {
);
let threats_count: u64 = db
.get_setting("weekly_threats_count")?
.get_report_snapshot("weekly_threats_count")
.await?
.and_then(|v| v.parse().ok())
.unwrap_or(0);
let top_ips: Vec<BlockedIpItem> = db
.get_setting("weekly_top_ips")?
.get_report_snapshot("weekly_top_ips")
.await?
.and_then(|v| serde_json::from_str(&v).ok())
.unwrap_or_else(|| {
vec![BlockedIpItem {
@ -31,7 +33,8 @@ pub fn build_report_data(db: &dyn SettingRepo) -> Result<ReportData, Error> {
});
let breakdown: Vec<ThreatBreakdownItem> = db
.get_setting("weekly_threat_breakdown")?
.get_report_snapshot("weekly_threat_breakdown")
.await?
.and_then(|v| {
let obj: serde_json::Value = serde_json::from_str(&v).ok()?;
let items = obj
@ -48,7 +51,8 @@ pub fn build_report_data(db: &dyn SettingRepo) -> Result<ReportData, Error> {
.unwrap_or_default();
let health: SystemHealthSummary = db
.get_setting("weekly_system_health")?
.get_report_snapshot("weekly_system_health")
.await?
.and_then(|v| serde_json::from_str(&v).ok())
.unwrap_or(SystemHealthSummary {
avg_cpu_percent: 0.0,
@ -69,35 +73,42 @@ pub fn build_report_data(db: &dyn SettingRepo) -> Result<ReportData, Error> {
}
let uptime_percent: f64 = db
.get_setting("system_uptime_percent")?
.get_report_snapshot("system_uptime_percent")
.await?
.and_then(|v| v.parse().ok())
.unwrap_or(0.0);
let active_rules: u64 = db
.get_setting("active_rules_count")?
.get_report_snapshot("active_rules_count")
.await?
.and_then(|v| v.parse().ok())
.unwrap_or(0);
let geo_distribution: Vec<GeoItem> = db
.get_setting("weekly_geo_distribution")?
.get_report_snapshot("weekly_geo_distribution")
.await?
.and_then(|v| serde_json::from_str(&v).ok())
.unwrap_or_default();
let auto_blocks: u64 = db
.get_setting("weekly_soar_blocks")?
.get_report_snapshot("weekly_soar_blocks")
.await?
.and_then(|v| v.parse().ok())
.unwrap_or(0);
let playbooks_triggered: u64 = db
.get_setting("weekly_soar_triggers")?
.get_report_snapshot("weekly_soar_triggers")
.await?
.and_then(|v| v.parse().ok())
.unwrap_or(0);
let auto_unblocks: u64 = db
.get_setting("weekly_soar_unblocks")?
.get_report_snapshot("weekly_soar_unblocks")
.await?
.and_then(|v| v.parse().ok())
.unwrap_or(0);
let blocked_count: u64 = db
.get_setting("weekly_blocked_count")?
.get_report_snapshot("weekly_blocked_count")
.await?
.and_then(|v| v.parse().ok())
.unwrap_or(auto_blocks);

View File

@ -10,12 +10,12 @@ use crate::domain::common::error::io::IOError;
use crate::domain::common::error::misc::MiscError;
use crate::domain::common::log::system::SystemLog;
use crate::domain::report::data::ReportData;
use crate::interface::setting::SettingRepo;
use crate::interface::report_snapshot::ReportSnapshotRepo;
/// Generate a self-contained HTML security report and write to disk.
/// Returns the path to the generated HTML file.
pub fn generate_html_report(db: &dyn SettingRepo, output_dir: &str) -> Result<PathBuf, Error> {
let data = build_report_data(db)?;
pub async fn generate_html_report(db: &dyn ReportSnapshotRepo, output_dir: &str) -> Result<PathBuf, Error> {
let data = build_report_data(db).await?;
let html = render_html_report(&data);
let html_path = PathBuf::from(output_dir).join(format!(
@ -39,9 +39,9 @@ fn render_html_report(data: &ReportData) -> String {
for item in &data.threat_breakdown {
breakdown_rows.push_str(&format!(
"<tr><td>{}</td><td class=\"num\">{}</td><td>{}</td></tr>",
html_escape(&item.threat_type),
super::html::escape(&item.threat_type),
item.count,
html_escape(&item.trend)
super::html::escape(&item.trend)
));
}
if data.threat_breakdown.is_empty() {
@ -54,9 +54,9 @@ fn render_html_report(data: &ReportData) -> String {
for ip in &data.top_blocked_ips {
ip_rows.push_str(&format!(
"<tr><td><code>{}</code></td><td class=\"num\">{}</td><td>{}</td></tr>",
html_escape(&ip.ip),
super::html::escape(&ip.ip),
ip.count,
html_escape(&ip.country)
super::html::escape(&ip.country)
));
}
if data.top_blocked_ips.is_empty() {
@ -68,7 +68,7 @@ fn render_html_report(data: &ReportData) -> String {
for geo in &data.geo_distribution {
geo_rows.push_str(&format!(
"<tr><td>{}</td><td class=\"num\">{}</td></tr>",
html_escape(&geo.country),
super::html::escape(&geo.country),
geo.threat_count
));
}
@ -79,7 +79,7 @@ fn render_html_report(data: &ReportData) -> String {
// Recommendations
let mut rec_items = String::new();
for rec in &data.recommendations {
rec_items.push_str(&format!("<li>{}</li>", html_escape(rec)));
rec_items.push_str(&format!("<li>{}</li>", super::html::escape(rec)));
}
format!(
@ -174,8 +174,8 @@ fn render_html_report(data: &ReportData) -> String {
</div>
</body>
</html>"#,
period = html_escape(&data.period),
generated_at = html_escape(&data.generated_at),
period = super::html::escape(&data.period),
generated_at = super::html::escape(&data.generated_at),
total_threats = data.executive_summary.total_threats,
total_blocked = data.executive_summary.total_blocked,
uptime = data.executive_summary.uptime_percent,
@ -189,21 +189,13 @@ fn render_html_report(data: &ReportData) -> String {
avg_cpu = data.system_health.avg_cpu_percent,
avg_mem = data.system_health.avg_memory_percent,
disk = data.system_health.disk_usage_percent,
ebpf_status = html_escape(&data.system_health.ebpf_status),
ebpf_status = super::html::escape(&data.system_health.ebpf_status),
rec_items = rec_items,
)
}
/// Basic HTML escaping for report content.
fn html_escape(s: &str) -> String {
s.replace('&', "&amp;")
.replace('<', "&lt;")
.replace('>', "&gt;")
.replace('"', "&quot;")
}
/// Generate report data and format as JSON (for API responses).
pub fn generate_report_json(db: &dyn SettingRepo) -> Result<serde_json::Value, Error> {
let data = build_report_data(db)?;
pub async fn generate_report_json(db: &dyn ReportSnapshotRepo) -> Result<serde_json::Value, Error> {
let data = build_report_data(db).await?;
serde_json::to_value(&data).map_err(|e| MiscError::SerializeError(e).into())
}

View File

@ -8,22 +8,26 @@ use tokio::time::{self, Duration};
use crate::domain::common::error::Error;
use crate::domain::common::log::system::SystemLog;
use crate::interface::health_query::HealthQuery;
use crate::interface::setting::SettingRepo;
use crate::interface::report_snapshot::ReportSnapshotRepo;
use crate::interface::stats::StatsRepo;
pub struct StatsAggregator {
stats: Arc<dyn StatsRepo>,
repo: Arc<dyn SettingRepo + Send + Sync>,
snapshots: Arc<dyn ReportSnapshotRepo>,
health: Arc<dyn HealthQuery>,
}
impl StatsAggregator {
pub fn new(
stats: Arc<dyn StatsRepo>,
repo: Arc<dyn SettingRepo + Send + Sync>,
snapshots: Arc<dyn ReportSnapshotRepo>,
health: Arc<dyn HealthQuery>,
) -> Self {
Self { stats, repo, health }
Self {
stats,
snapshots,
health,
}
}
/// Spawn a background task that runs aggregation every hour.
@ -31,51 +35,59 @@ impl StatsAggregator {
tokio::spawn(async move {
log!(SystemLog::StatsAggregatorStarted);
// Run immediately on startup
if let Err(e) = self.aggregate() {
if let Err(e) = self.aggregate().await {
log!(SystemLog::InitialStatsAggregationFailed(e.to_string()));
}
let mut interval = time::interval(Duration::from_secs(3600));
loop {
interval.tick().await;
if let Err(e) = self.aggregate() {
if let Err(e) = self.aggregate().await {
log!(SystemLog::StatsAggregationFailed(e.to_string()));
}
}
})
}
/// Aggregate all weekly statistics and write to settings.
pub fn aggregate(&self) -> Result<(), Error> {
/// Aggregate all weekly statistics and write to the report snapshot store.
pub async fn aggregate(&self) -> Result<(), Error> {
let days = 7;
// SOAR execution counts
let threats_count = self.stats.count_weekly_executions(days)?;
self.repo
.set_setting("weekly_threats_count", &threats_count.to_string())?;
let threats_count = self.stats.count_weekly_executions(days).await?;
self.snapshots
.set_report_snapshot("weekly_threats_count", &threats_count.to_string())
.await?;
let blocks_count = self.stats.count_weekly_blocks(days)?;
self.repo.set_setting("weekly_soar_blocks", &blocks_count.to_string())?;
self.repo
.set_setting("weekly_soar_triggers", &threats_count.to_string())?;
let blocks_count = self.stats.count_weekly_blocks(days).await?;
self.snapshots
.set_report_snapshot("weekly_soar_blocks", &blocks_count.to_string())
.await?;
self.snapshots
.set_report_snapshot("weekly_soar_triggers", &threats_count.to_string())
.await?;
let unblocks_count = self.stats.count_weekly_unblocks(days)?;
self.repo
.set_setting("weekly_soar_unblocks", &unblocks_count.to_string())?;
let unblocks_count = self.stats.count_weekly_unblocks(days).await?;
self.snapshots
.set_report_snapshot("weekly_soar_unblocks", &unblocks_count.to_string())
.await?;
self.repo
.set_setting("weekly_blocked_count", &blocks_count.to_string())?;
self.snapshots
.set_report_snapshot("weekly_blocked_count", &blocks_count.to_string())
.await?;
let breakdown = self.stats.weekly_threat_breakdown(days)?;
let breakdown = self.stats.weekly_threat_breakdown(days).await?;
let breakdown_json: serde_json::Map<String, serde_json::Value> = breakdown
.into_iter()
.map(|entry| (entry.threat_type, Value::Number(entry.count.into())))
.collect();
self.repo.set_setting(
"weekly_threat_breakdown",
&serde_json::to_string(&breakdown_json).unwrap_or_else(|_| "{}".to_string()),
)?;
self.snapshots
.set_report_snapshot(
"weekly_threat_breakdown",
&serde_json::to_string(&breakdown_json).unwrap_or_else(|_| "{}".to_string()),
)
.await?;
let top_ips = self.stats.weekly_top_ips(days, 10)?;
let top_ips = self.stats.weekly_top_ips(days, 10).await?;
let top_ips_json: Vec<serde_json::Value> = top_ips
.into_iter()
.map(|entry| {
@ -86,14 +98,18 @@ impl StatsAggregator {
})
})
.collect();
self.repo.set_setting(
"weekly_top_ips",
&serde_json::to_string(&top_ips_json).unwrap_or_else(|_| "[]".to_string()),
)?;
self.snapshots
.set_report_snapshot(
"weekly_top_ips",
&serde_json::to_string(&top_ips_json).unwrap_or_else(|_| "[]".to_string()),
)
.await?;
// Active rules count
let active_rules = self.stats.count_acl_rules()?;
self.repo.set_setting("active_rules_count", &active_rules.to_string())?;
let active_rules = self.stats.count_acl_rules().await?;
self.snapshots
.set_report_snapshot("active_rules_count", &active_rules.to_string())
.await?;
{
let metrics = self.health.get_current_metrics();
@ -106,10 +122,12 @@ impl StatsAggregator {
"disk_usage_percent": 0.0,
"ebpf_status": format!("{:?}", metrics.ebpf),
});
self.repo.set_setting(
"weekly_system_health",
&serde_json::to_string(&health_json).unwrap_or_else(|_| "{}".to_string()),
)?;
self.snapshots
.set_report_snapshot(
"weekly_system_health",
&serde_json::to_string(&health_json).unwrap_or_else(|_| "{}".to_string()),
)
.await?;
let uptime_secs = metrics.uptime_seconds;
let week_secs = (days as u64) * 86400;
@ -118,13 +136,21 @@ impl StatsAggregator {
} else {
(uptime_secs as f64 / week_secs as f64) * 100.0
};
self.repo
.set_setting("system_uptime_percent", &format!("{:.1}", uptime_percent))?;
self.snapshots
.set_report_snapshot("system_uptime_percent", &format!("{:.1}", uptime_percent))
.await?;
}
// Geo distribution (initialize if not present)
if self.repo.get_setting("weekly_geo_distribution")?.is_none() {
self.repo.set_setting("weekly_geo_distribution", "[]")?;
if self
.snapshots
.get_report_snapshot("weekly_geo_distribution")
.await?
.is_none()
{
self.snapshots
.set_report_snapshot("weekly_geo_distribution", "[]")
.await?;
}
log!(SystemLog::StatsAggregated(
@ -146,68 +172,80 @@ mod tests {
use crate::domain::common::system::health::EbpfHealth;
use crate::infrastructure::health::SystemHealth;
fn test_health(db: &Arc<Database>) -> Arc<dyn HealthQuery> {
async fn test_health(db: &Arc<Database>) -> Arc<dyn HealthQuery> {
use arc_swap::ArcSwap;
let config = Arc::new(ArcSwap::from_pointee(AppConfig::from_settings(db.as_ref()).unwrap()));
let config = Arc::new(ArcSwap::from_pointee(
AppConfig::from_config_repo(db.as_ref()).await.unwrap(),
));
let ebpf_health = Arc::new(ArcSwap::from_pointee(EbpfHealth::Healthy));
Arc::new(SystemHealth::new(config, ebpf_health).unwrap())
}
#[test]
fn aggregator_writes_weekly_stats() {
let db = Arc::new(Database::new(":memory:").expect("test db"));
#[tokio::test]
async fn aggregator_writes_weekly_stats() {
let db = Arc::new(Database::new(":memory:").await.expect("test db"));
// Seed some SOAR executions
db.seed_default_playbooks().ok();
db.seed_default_playbooks().await.ok();
db.insert_soar_execution(1, Some("1.2.3.4"), "threat_detected", "[]")
.await
.ok();
db.insert_soar_execution(1, Some("5.6.7.8"), "brute_force", "[]")
.await
.ok();
db.commit_soar_block_to_db("1.2.3.4", 4, 1, "2099-01-01 00:00:00")
.await
.ok();
db.insert_soar_execution(1, Some("5.6.7.8"), "brute_force", "[]").ok();
db.insert_soar_block_rule("1.2.3.4", 1, "2099-01-01 00:00:00").ok();
let health = test_health(&db);
let health = test_health(&db).await;
let aggregator = StatsAggregator::new(
db.clone() as Arc<dyn StatsRepo>,
db.clone() as Arc<dyn SettingRepo + Send + Sync>,
db.clone() as Arc<dyn ReportSnapshotRepo>,
health,
);
aggregator.aggregate().expect("aggregation should succeed");
aggregator.aggregate().await.expect("aggregation should succeed");
let threats = db.get_setting("weekly_threats_count").unwrap().unwrap();
let threats = db.get_report_snapshot("weekly_threats_count").await.unwrap().unwrap();
assert_eq!(threats, "2");
let blocks = db.get_setting("weekly_soar_blocks").unwrap().unwrap();
let blocks = db.get_report_snapshot("weekly_soar_blocks").await.unwrap().unwrap();
assert_eq!(blocks, "1");
let breakdown = db.get_setting("weekly_threat_breakdown").unwrap().unwrap();
let breakdown = db
.get_report_snapshot("weekly_threat_breakdown")
.await
.unwrap()
.unwrap();
let parsed: serde_json::Value = serde_json::from_str(&breakdown).unwrap();
assert_eq!(parsed["threat_detected"], 1);
assert_eq!(parsed["brute_force"], 1);
let active_rules = db.get_setting("active_rules_count").unwrap().unwrap();
assert_eq!(active_rules, "0");
let active_rules = db.get_report_snapshot("active_rules_count").await.unwrap().unwrap();
assert_eq!(active_rules, "1");
let uptime = db.get_setting("system_uptime_percent").unwrap().unwrap();
let uptime = db.get_report_snapshot("system_uptime_percent").await.unwrap().unwrap();
assert!(!uptime.is_empty());
let sys_health = db.get_setting("weekly_system_health").unwrap().unwrap();
let sys_health = db.get_report_snapshot("weekly_system_health").await.unwrap().unwrap();
let health_val: serde_json::Value = serde_json::from_str(&sys_health).unwrap();
assert!(health_val["avg_cpu_percent"].as_f64().is_some());
}
#[test]
fn aggregator_handles_empty_db() {
let db = Arc::new(Database::new(":memory:").expect("test db"));
let health = test_health(&db);
#[tokio::test]
async fn aggregator_handles_empty_db() {
let db = Arc::new(Database::new(":memory:").await.expect("test db"));
let health = test_health(&db).await;
let aggregator = StatsAggregator::new(
db.clone() as Arc<dyn StatsRepo>,
db.clone() as Arc<dyn SettingRepo + Send + Sync>,
db.clone() as Arc<dyn ReportSnapshotRepo>,
health,
);
aggregator
.aggregate()
.await
.expect("aggregation should succeed with empty data");
let threats = db.get_setting("weekly_threats_count").unwrap().unwrap();
let threats = db.get_report_snapshot("weekly_threats_count").await.unwrap().unwrap();
assert_eq!(threats, "0");
}
}

View File

@ -25,7 +25,6 @@ use crate::domain::response::error::SoarError;
use crate::domain::response::log::SoarLog;
use crate::domain::response::playbook::{Playbook, PlaybookAction};
use crate::interface::access_control::AccessControlPort;
use crate::interface::app_repo::AppRepo;
use crate::utils::ip_address::{ip_version_from_str, is_private_ip};
/// Lower bound on the rate-limit factor — anything below 1% of current
@ -91,7 +90,8 @@ impl SoarEngine {
// Write audit trail
let actions_json = serde_json::to_string(&action_results).unwrap_or_default();
self.db
.insert_soar_execution(playbook.id, Some(&event.source_ip), &event.attack_type, &actions_json)?;
.insert_soar_execution(playbook.id, Some(&event.source_ip), &event.attack_type, &actions_json)
.await?;
log!(SoarLog::PlaybookExecuted(
playbook.name.clone(),
@ -190,17 +190,12 @@ impl SoarEngine {
let expires_str = expires_at.format("%Y-%m-%d %H:%M:%S").to_string();
// Atomically write soar_block_rules + acl_rules. Either both commit
// or both roll back — no half-state possible. Offloaded so the
// SQLite WAL fsync can't block the tokio worker.
// or both roll back — no half-state possible.
let ip_version = ip_version_from_str(&event.source_ip);
let commit_result = commit_block_blocking(
self.db.clone(),
event.source_ip.clone(),
ip_version,
playbook_id,
expires_str.clone(),
)
.await;
let commit_result = self
.db
.commit_soar_block_to_db(&event.source_ip, ip_version, playbook_id, &expires_str)
.await;
if let Err(e) = commit_result {
// DB tx rolled back both rows; now roll back the eBPF block.
let unblock_outcome = unblock_ip_blocking(Arc::clone(&self.access_control), event.source_ip.clone()).await;
@ -210,7 +205,7 @@ impl SoarEngine {
event.source_ip, unblock_err
)));
// Write to pending_unblock table so recovery can retry later
if let Err(pend_err) = insert_pending_unblock_blocking(self.db.clone(), event.source_ip.clone()).await {
if let Err(pend_err) = self.db.insert_pending_unblock(&event.source_ip).await {
log!(SoarLog::EventHandlingFailed(format!(
"CRITICAL: Failed to queue pending unblock for IP {}: {}",
event.source_ip, pend_err
@ -298,7 +293,8 @@ impl SoarEngine {
let smtp_cfg = self.matcher.config.load().notification.smtp.clone();
match self
.email_sender_factory
.build_smtp_sender(&smtp_cfg, self.secrets.as_deref())?
.build_smtp_sender(&smtp_cfg, self.secrets.as_deref())
.await?
{
Some(smtp) => {
let subject = format!(
@ -472,37 +468,20 @@ impl SoarEngine {
self.record_cooldown(FALLBACK_PLAYBOOK_ID, &event.source_ip);
// Audit trail under the fallback synthetic playbook id
self.db.insert_soar_execution(
FALLBACK_PLAYBOOK_ID,
Some(&event.source_ip),
&event.attack_type,
&serde_json::to_string(&[result_json]).unwrap_or_default(),
)?;
self.db
.insert_soar_execution(
FALLBACK_PLAYBOOK_ID,
Some(&event.source_ip),
&event.attack_type,
&serde_json::to_string(&[result_json]).unwrap_or_default(),
)
.await?;
log!(SoarLog::FallbackExecuted(event.source_ip.clone()));
Ok(())
}
}
async fn commit_block_blocking(
db: Arc<dyn AppRepo>,
source_ip: String,
ip_version: u8,
playbook_id: i64,
expires_str: String,
) -> Result<(), Error> {
spawn_blocking(move || db.commit_soar_block_to_db(&source_ip, ip_version, playbook_id, &expires_str))
.await
.map_err(|e| SoarError::ActionFailed("commit_soar_block_to_db", e))??;
Ok(())
}
async fn insert_pending_unblock_blocking(db: Arc<dyn AppRepo>, source_ip: String) -> Result<i64, Error> {
spawn_blocking(move || db.insert_pending_unblock(&source_ip))
.await
.map_err(|e| SoarError::ActionFailed("insert_pending_unblock", e))?
}
/// Offload `AccessControlPort::block_ip` onto a blocking thread. The port is
/// synchronous because its eBPF-map critical sections are tiny (microseconds),
/// but under SOAR burst many tokio workers would contend on the same

View File

@ -50,7 +50,7 @@ pub struct SoarEngineDeps {
}
impl SoarEngine {
pub fn new(deps: SoarEngineDeps) -> Result<Self, Error> {
pub async fn new(deps: SoarEngineDeps) -> Result<Self, Error> {
let SoarEngineDeps {
db,
config,
@ -65,9 +65,11 @@ impl SoarEngine {
let soar_cfg = config.load();
let rate_limit_channel = soar_cfg.soar.rate_limit_cmd_channel_capacity;
let freq_max_keys = soar_cfg.soar.frequency_max_tracked_keys;
let freq_max_events_per_key = soar_cfg.soar.frequency_max_events_per_key;
let freq_retention_secs = soar_cfg.soar.frequency_retention_secs;
drop(soar_cfg);
let rate_limit_owner = rate_limit.map(|rl| RateLimitOwnerHandle::spawn(rl, config.clone(), rate_limit_channel));
let matcher = PlaybookMatcher::new(config, freq_max_keys);
let matcher = PlaybookMatcher::new(config, freq_max_keys, freq_max_events_per_key, freq_retention_secs);
let engine = Self {
db,
access_control,
@ -79,13 +81,13 @@ impl SoarEngine {
secrets,
email_sender_factory,
};
engine.reload_cache()?;
engine.reload_cache().await?;
Ok(engine)
}
/// Load playbooks and admin whitelist from DB into memory.
pub fn reload_cache(&self) -> Result<(), Error> {
let views = self.db.list_playbooks()?;
pub async fn reload_cache(&self) -> Result<(), Error> {
let views = self.db.list_playbooks().await?;
let mut playbooks: Vec<Playbook> = Vec::new();
for view in views {
@ -162,13 +164,13 @@ impl SoarEngine {
self.matcher.playbooks.store(Arc::new(playbooks));
// Load admin whitelist
let whitelist = self.db.list_admin_whitelist()?;
let whitelist = self.db.list_admin_whitelist().await?;
let whitelist: HashSet<String> = whitelist.into_iter().collect();
let whitelist_count = whitelist.len();
self.matcher.admin_whitelist.store(Arc::new(whitelist));
// Initialize block counter from DB
let count = self.db.count_active_soar_blocks()?;
let count = self.db.count_active_soar_blocks().await?;
self.matcher.active_block_count.store(count, Ordering::SeqCst);
log!(SoarLog::CacheLoaded(playbook_count, whitelist_count, count));
@ -244,7 +246,7 @@ impl SoarEngine {
// First, retry any pending unblocks from previous orphan failures
self.retry_pending_unblocks().await;
let active_blocks = self.db.list_active_soar_blocks()?;
let active_blocks = self.db.list_active_soar_blocks().await?;
let count = active_blocks.len();
for block in &active_blocks {
@ -263,7 +265,7 @@ impl SoarEngine {
/// Retry pending unblocks that failed during previous runs.
async fn retry_pending_unblocks(&self) {
let pending = match self.db.list_pending_unblocks() {
let pending = match self.db.list_pending_unblocks().await {
Ok(p) => p,
Err(e) => {
log!(SoarLog::EventHandlingFailed(format!(
@ -281,13 +283,13 @@ impl SoarEngine {
pu.source_ip, pu.retry_count
)));
// Remove from queue to avoid infinite retries
let _ = self.db.delete_pending_unblock(pu.id);
let _ = self.db.delete_pending_unblock(pu.id).await;
continue;
}
match self.access_control.unblock_ip(&pu.source_ip) {
Ok(()) => {
let _ = self.db.delete_pending_unblock(pu.id);
let _ = self.db.delete_pending_unblock(pu.id).await;
log!(SoarLog::EventHandlingFailed(format!(
"Successfully unblocked orphan IP {} on retry #{}",
pu.source_ip,
@ -295,7 +297,7 @@ impl SoarEngine {
)));
}
Err(e) => {
let _ = self.db.increment_pending_unblock_retry(pu.id);
let _ = self.db.increment_pending_unblock_retry(pu.id).await;
log!(SoarLog::EventHandlingFailed(format!(
"Retry #{} failed to unblock orphan IP {}: {}",
pu.retry_count + 1,
@ -363,8 +365,9 @@ mod tests {
struct NoopEmailSenderFactory;
#[async_trait::async_trait]
impl EmailSenderFactoryTrait for NoopEmailSenderFactory {
fn build_smtp_sender(
async fn build_smtp_sender(
&self,
_cfg: &SmtpConfig,
_secrets: Option<&dyn crate::interface::secret_store::SecretStorePort>,
@ -408,18 +411,20 @@ mod tests {
}
}
fn test_db() -> Arc<crate::adapter::persistence::Database> {
async fn test_db() -> Arc<crate::adapter::persistence::Database> {
use crate::adapter::persistence::Database;
Arc::new(Database::new(":memory:").expect("Failed to create test database"))
Arc::new(Database::new(":memory:").await.expect("Failed to create test database"))
}
fn test_engine(ac: Arc<dyn AccessControlPort>) -> SoarEngine {
let db = test_db();
db.seed_default_playbooks().ok();
db.set_setting("enforce_mode", "enforce").ok();
async fn test_engine(ac: Arc<dyn AccessControlPort>) -> SoarEngine {
let db = test_db().await;
db.seed_default_playbooks().await.ok();
db.set_config_value("enforce_mode", "enforce").await.ok();
let cache = Arc::new(AtomicU8::new(2));
AppConfig::seed_defaults(&*db).expect("seed config defaults");
let cfg = AppConfig::from_settings(&*db).expect("load config");
AppConfig::seed_config_defaults(&*db)
.await
.expect("seed config defaults");
let cfg = AppConfig::from_config_repo(&*db).await.expect("load config");
let config = Arc::new(ArcSwap::from_pointee(cfg));
SoarEngine::new(SoarEngineDeps {
db: db as Arc<dyn AppRepo>,
@ -432,13 +437,14 @@ mod tests {
secrets: None,
email_sender_factory: Arc::new(NoopEmailSenderFactory),
})
.await
.expect("Failed to create SOAR engine")
}
#[tokio::test]
async fn block_ip_calls_access_control_port() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock.clone());
let engine = test_engine(mock.clone()).await;
let event = ThreatDetectedEvent {
source_ip: "1.2.3.4".to_string(),
@ -475,7 +481,7 @@ mod tests {
#[tokio::test]
async fn block_ip_ipv6_calls_access_control_port() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock.clone());
let engine = test_engine(mock.clone()).await;
let event = ThreatDetectedEvent {
source_ip: "::1".to_string(),
@ -512,18 +518,20 @@ mod tests {
#[tokio::test]
async fn recover_active_blocks_uses_port() {
let mock = Arc::new(MockAccessControl::new());
let db = test_db();
db.seed_default_playbooks().ok();
let db = test_db().await;
db.seed_default_playbooks().await.ok();
// Insert a fake active block
let expires = (Utc::now() + ChronoDuration::hours(1))
.format("%Y-%m-%d %H:%M:%S")
.to_string();
db.insert_soar_block_rule("192.168.1.100", 1, &expires).ok();
db.commit_soar_block_to_db("192.168.1.100", 4, 1, &expires).await.ok();
let cache = Arc::new(AtomicU8::new(2));
AppConfig::seed_defaults(&*db).expect("seed config defaults");
let cfg = AppConfig::from_settings(&*db).expect("load config");
AppConfig::seed_config_defaults(&*db)
.await
.expect("seed config defaults");
let cfg = AppConfig::from_config_repo(&*db).await.expect("load config");
let config = Arc::new(ArcSwap::from_pointee(cfg));
let engine = SoarEngine::new(SoarEngineDeps {
db: db as Arc<dyn AppRepo>,
@ -536,6 +544,7 @@ mod tests {
secrets: None,
email_sender_factory: Arc::new(NoopEmailSenderFactory),
})
.await
.expect("Failed to create engine");
engine.recover_active_blocks().await.expect("Recovery should succeed");
@ -548,17 +557,19 @@ mod tests {
async fn recover_logs_warning_on_failure() {
let mock = Arc::new(MockAccessControl::new());
mock.should_fail.store(true, Ordering::SeqCst);
let db = test_db();
db.seed_default_playbooks().ok();
let db = test_db().await;
db.seed_default_playbooks().await.ok();
let expires = (Utc::now() + ChronoDuration::hours(1))
.format("%Y-%m-%d %H:%M:%S")
.to_string();
db.insert_soar_block_rule("10.0.0.1", 1, &expires).ok();
db.commit_soar_block_to_db("10.0.0.1", 4, 1, &expires).await.ok();
let cache = Arc::new(AtomicU8::new(2));
AppConfig::seed_defaults(&*db).expect("seed config defaults");
let cfg = AppConfig::from_settings(&*db).expect("load config");
AppConfig::seed_config_defaults(&*db)
.await
.expect("seed config defaults");
let cfg = AppConfig::from_config_repo(&*db).await.expect("load config");
let config = Arc::new(ArcSwap::from_pointee(cfg));
let engine = SoarEngine::new(SoarEngineDeps {
db: db as Arc<dyn AppRepo>,
@ -571,6 +582,7 @@ mod tests {
secrets: None,
email_sender_factory: Arc::new(NoopEmailSenderFactory),
})
.await
.expect("Failed to create engine");
// Should not panic — errors are logged, not propagated
@ -584,7 +596,7 @@ mod tests {
#[tokio::test]
async fn block_ip_respects_cap() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock.clone());
let engine = test_engine(mock.clone()).await;
// Set counter to max (default cap is 100)
engine.matcher.active_block_count.store(100, Ordering::SeqCst);
@ -624,7 +636,7 @@ mod tests {
#[tokio::test]
async fn cooldown_prevents_duplicate_execution() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock.clone());
let engine = test_engine(mock.clone()).await;
let event = ThreatDetectedEvent {
source_ip: "1.2.3.4".to_string(),
@ -667,13 +679,15 @@ mod tests {
#[tokio::test]
async fn whitelist_prevents_execution() {
let mock = Arc::new(MockAccessControl::new());
let db = test_db();
db.seed_default_playbooks().ok();
db.insert_admin_whitelist("1.2.3.4").ok();
let db = test_db().await;
db.seed_default_playbooks().await.ok();
db.insert_admin_whitelist("1.2.3.4").await.ok();
let cache = Arc::new(AtomicU8::new(2));
AppConfig::seed_defaults(&*db).expect("seed config defaults");
let cfg = AppConfig::from_settings(&*db).expect("load config");
AppConfig::seed_config_defaults(&*db)
.await
.expect("seed config defaults");
let cfg = AppConfig::from_config_repo(&*db).await.expect("load config");
let config = Arc::new(ArcSwap::from_pointee(cfg));
let engine = SoarEngine::new(SoarEngineDeps {
db: db as Arc<dyn AppRepo>,
@ -686,6 +700,7 @@ mod tests {
secrets: None,
email_sender_factory: Arc::new(NoopEmailSenderFactory),
})
.await
.expect("Failed to create engine");
let event = ThreatDetectedEvent {
@ -717,10 +732,10 @@ mod tests {
);
}
#[test]
fn decrement_block_count_no_underflow() {
#[tokio::test]
async fn decrement_block_count_no_underflow() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock);
let engine = test_engine(mock).await;
// Start at 0
assert_eq!(engine.matcher.active_block_count.load(Ordering::SeqCst), 0);
@ -741,15 +756,17 @@ mod tests {
assert_eq!(engine.matcher.active_block_count.load(Ordering::SeqCst), 0);
}
#[test]
fn reload_cache_loads_playbooks_via_join() {
#[tokio::test]
async fn reload_cache_loads_playbooks_via_join() {
let mock = Arc::new(MockAccessControl::new());
let db = test_db();
db.seed_default_playbooks().ok();
let db = test_db().await;
db.seed_default_playbooks().await.ok();
let cache = Arc::new(AtomicU8::new(0));
AppConfig::seed_defaults(&*db).expect("seed config defaults");
let cfg = AppConfig::from_settings(&*db).expect("load config");
AppConfig::seed_config_defaults(&*db)
.await
.expect("seed config defaults");
let cfg = AppConfig::from_config_repo(&*db).await.expect("load config");
let config = Arc::new(ArcSwap::from_pointee(cfg));
let engine = SoarEngine::new(SoarEngineDeps {
db: db.clone() as Arc<dyn AppRepo>,
@ -762,6 +779,7 @@ mod tests {
secrets: None,
email_sender_factory: Arc::new(NoopEmailSenderFactory),
})
.await
.expect("Failed to create engine");
// Should have loaded default playbooks
@ -769,13 +787,15 @@ mod tests {
assert!(count > 0, "Should have loaded default playbooks");
// Add a new playbook directly to DB
db.insert_playbook("test_pb", "port_scan", None, None, None, 60).ok();
db.insert_playbook("test_pb", "port_scan", None, None, None, 60)
.await
.ok();
// Cache should not have it yet
assert_eq!(engine.matcher.playbooks.load().len(), count);
// After reload, should have one more
engine.reload_cache().expect("reload should succeed");
engine.reload_cache().await.expect("reload should succeed");
assert_eq!(engine.matcher.playbooks.load().len(), count + 1);
}
@ -811,10 +831,10 @@ mod tests {
}
}
#[test]
fn condition_threshold_gte_passes() {
#[tokio::test]
async fn condition_threshold_gte_passes() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock);
let engine = test_engine(mock).await;
let pb = make_playbook(vec![PlaybookCondition {
condition_type: ConditionType::Threshold,
operator: ">=".to_string(),
@ -825,10 +845,10 @@ mod tests {
assert!(engine.evaluate_conditions(&pb, &event));
}
#[test]
fn condition_threshold_gte_fails() {
#[tokio::test]
async fn condition_threshold_gte_fails() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock);
let engine = test_engine(mock).await;
let pb = make_playbook(vec![PlaybookCondition {
condition_type: ConditionType::Threshold,
operator: ">=".to_string(),
@ -839,10 +859,10 @@ mod tests {
assert!(!engine.evaluate_conditions(&pb, &event));
}
#[test]
fn condition_threshold_lte() {
#[tokio::test]
async fn condition_threshold_lte() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock);
let engine = test_engine(mock).await;
let pb = make_playbook(vec![PlaybookCondition {
condition_type: ConditionType::Threshold,
operator: "<=".to_string(),
@ -853,10 +873,10 @@ mod tests {
assert!(!engine.evaluate_conditions(&pb, &test_event(0.7, None, "1.2.3.4", false)));
}
#[test]
fn condition_source_country_in() {
#[tokio::test]
async fn condition_source_country_in() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock);
let engine = test_engine(mock).await;
let pb = make_playbook(vec![PlaybookCondition {
condition_type: ConditionType::SourceCountry,
operator: "in".to_string(),
@ -868,10 +888,10 @@ mod tests {
assert!(!engine.evaluate_conditions(&pb, &test_event(0.9, None, "1.2.3.4", false)));
}
#[test]
fn condition_source_country_not_in() {
#[tokio::test]
async fn condition_source_country_not_in() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock);
let engine = test_engine(mock).await;
let pb = make_playbook(vec![PlaybookCondition {
condition_type: ConditionType::SourceCountry,
operator: "not_in".to_string(),
@ -882,10 +902,10 @@ mod tests {
assert!(!engine.evaluate_conditions(&pb, &test_event(0.9, Some("TW"), "1.2.3.4", false)));
}
#[test]
fn condition_ip_pattern_in() {
#[tokio::test]
async fn condition_ip_pattern_in() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock);
let engine = test_engine(mock).await;
let pb = make_playbook(vec![PlaybookCondition {
condition_type: ConditionType::IpPattern,
operator: "in".to_string(),
@ -896,10 +916,10 @@ mod tests {
assert!(!engine.evaluate_conditions(&pb, &test_event(0.9, None, "192.168.1.1", false)));
}
#[test]
fn condition_repeat_offender() {
#[tokio::test]
async fn condition_repeat_offender() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock);
let engine = test_engine(mock).await;
let pb = make_playbook(vec![PlaybookCondition {
condition_type: ConditionType::RepeatOffender,
operator: "==".to_string(),
@ -910,10 +930,10 @@ mod tests {
assert!(!engine.evaluate_conditions(&pb, &test_event(0.9, None, "1.2.3.4", false)));
}
#[test]
fn condition_and_logic_all_must_pass() {
#[tokio::test]
async fn condition_and_logic_all_must_pass() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock);
let engine = test_engine(mock).await;
let pb = make_playbook(vec![
PlaybookCondition {
condition_type: ConditionType::Threshold,
@ -936,10 +956,10 @@ mod tests {
assert!(!engine.evaluate_conditions(&pb, &test_event(0.5, Some("CN"), "1.2.3.4", false)));
}
#[test]
fn frequency_condition_counts_events() {
#[tokio::test]
async fn frequency_condition_counts_events() {
let mock = Arc::new(MockAccessControl::new());
let engine = test_engine(mock);
let engine = test_engine(mock).await;
let pb = make_playbook(vec![PlaybookCondition {
condition_type: ConditionType::Frequency,
operator: ">=".to_string(),

View File

@ -11,14 +11,16 @@ pub struct FrequencyTracker {
events: DashMap<FreqKey, VecDeque<Instant>>,
max_deque_size: usize,
max_tracked_keys: usize,
max_retention: Duration,
}
impl FrequencyTracker {
pub fn new(max_tracked_keys: usize) -> Self {
pub fn new(max_tracked_keys: usize, max_events_per_key: usize, retention_secs: u64) -> Self {
Self {
events: DashMap::new(),
max_deque_size: 200,
max_deque_size: max_events_per_key.max(1),
max_tracked_keys: max_tracked_keys.max(1),
max_retention: Duration::from_secs(retention_secs.max(1)),
}
}
@ -50,20 +52,19 @@ impl FrequencyTracker {
deque.len() as u64
}
/// Remove empty deques and entries where all timestamps are expired.
/// Uses a conservative 2-hour max window for expiry detection.
/// Remove empty deques and entries where all timestamps are outside the
/// configured retention window.
pub fn cleanup(&self) -> u32 {
let now = Instant::now();
let max_window = Duration::from_secs(7200); // 2 hours — conservative upper bound
let mut removed = 0u32;
self.events.retain(|_, deque| {
if deque.is_empty() {
removed += 1;
return false;
}
// If all entries are older than max_window, remove the whole entry
// If all entries are older than max_retention, remove the whole entry.
if let Some(newest) = deque.back()
&& now.checked_duration_since(*newest).unwrap_or(Duration::ZERO) > max_window
&& now.checked_duration_since(*newest).unwrap_or(Duration::ZERO) > self.max_retention
{
removed += 1;
return false;
@ -84,3 +85,17 @@ impl FrequencyTracker {
removed
}
}
#[cfg(test)]
mod tests {
use super::FrequencyTracker;
#[test]
fn record_and_count_respects_configured_per_key_cap() {
let tracker = FrequencyTracker::new(10, 2, 60);
assert_eq!(tracker.record_and_count(1, "10.0.0.1", 60), 1);
assert_eq!(tracker.record_and_count(1, "10.0.0.1", 60), 2);
assert_eq!(tracker.record_and_count(1, "10.0.0.1", 60), 2);
}
}

View File

@ -29,12 +29,21 @@ pub struct PlaybookMatcher {
}
impl PlaybookMatcher {
pub fn new(config: Arc<ArcSwap<AppConfig>>, frequency_max_tracked_keys: usize) -> Self {
pub fn new(
config: Arc<ArcSwap<AppConfig>>,
frequency_max_tracked_keys: usize,
frequency_max_events_per_key: usize,
frequency_retention_secs: u64,
) -> Self {
Self {
playbooks: ArcSwap::from_pointee(Vec::new()),
admin_whitelist: ArcSwap::from_pointee(HashSet::new()),
cooldowns: DashMap::new(),
frequency_tracker: FrequencyTracker::new(frequency_max_tracked_keys),
frequency_tracker: FrequencyTracker::new(
frequency_max_tracked_keys,
frequency_max_events_per_key,
frequency_retention_secs,
),
active_block_count: AtomicU32::new(0),
config,
}

View File

@ -27,11 +27,11 @@ impl PlaybookService {
}
}
pub fn list_playbooks(&self) -> Result<Vec<PlaybookView>, Error> {
self.db.list_playbooks()
pub async fn list_playbooks(&self) -> Result<Vec<PlaybookView>, Error> {
self.db.list_playbooks().await
}
pub fn create_playbook(&self, input: &CreatePlaybookInput) -> Result<i64, Error> {
pub async fn create_playbook(&self, input: &CreatePlaybookInput) -> Result<i64, Error> {
// Single atomic insert (playbook + actions + conditions).
let actions: Vec<ActionInput> = input
.actions
@ -43,12 +43,15 @@ impl PlaybookService {
params_json: params.clone(),
})
.collect();
let playbook_id = self.db.insert_playbook_atomic(input, &actions, &input.conditions)?;
self.soar_engine.reload_cache()?;
let playbook_id = self
.db
.insert_playbook_atomic(input, &actions, &input.conditions)
.await?;
self.soar_engine.reload_cache().await?;
Ok(playbook_id)
}
pub fn update_playbook(&self, id: i64, input: &CreatePlaybookInput) -> Result<bool, Error> {
pub async fn update_playbook(&self, id: i64, input: &CreatePlaybookInput) -> Result<bool, Error> {
let row = UpdatePlaybookInput {
name: input.name.clone(),
trigger_event: input.trigger_event.clone(),
@ -68,32 +71,35 @@ impl PlaybookService {
params_json: params.clone(),
})
.collect();
let updated = self.db.update_playbook_atomic(id, &row, &actions, &input.conditions)?;
let updated = self
.db
.update_playbook_atomic(id, &row, &actions, &input.conditions)
.await?;
if !updated {
return Ok(false);
}
self.soar_engine.reload_cache()?;
self.soar_engine.reload_cache().await?;
Ok(true)
}
pub fn toggle_playbook(&self, id: i64, enabled: bool) -> Result<bool, Error> {
let updated = self.db.update_playbook_enabled(id, enabled)?;
pub async fn toggle_playbook(&self, id: i64, enabled: bool) -> Result<bool, Error> {
let updated = self.db.update_playbook_enabled(id, enabled).await?;
if updated {
self.soar_engine.reload_cache()?;
self.soar_engine.reload_cache().await?;
}
Ok(updated)
}
pub fn delete_playbook(&self, id: i64) -> Result<bool, Error> {
let deleted = self.db.delete_playbook(id)?;
pub async fn delete_playbook(&self, id: i64) -> Result<bool, Error> {
let deleted = self.db.delete_playbook(id).await?;
if deleted {
self.soar_engine.reload_cache()?;
self.soar_engine.reload_cache().await?;
}
Ok(deleted)
}
pub fn list_active_blocks(&self) -> Result<Vec<ActiveBlockView>, Error> {
self.db.list_active_soar_blocks()
pub async fn list_active_blocks(&self) -> Result<Vec<ActiveBlockView>, Error> {
self.db.list_active_soar_blocks().await
}
/// Manually unblock an IP: remove from eBPF, atomically clear both DB
@ -103,7 +109,8 @@ impl PlaybookService {
// Look up the block to get source_ip
let block = self
.db
.find_soar_block_by_id(id)?
.find_soar_block_by_id(id)
.await?
.ok_or_else(|| SoarError::UnblockRuleNotFound(id))?;
let source_ip = &block.source_ip;
@ -113,7 +120,7 @@ impl PlaybookService {
// Atomically drop acl_rules entry AND mark soar_block_rules
// unblocked in one transaction.
let ip_version = ip_version_from_str(source_ip);
self.db.commit_soar_unblock_to_db(id, ip_version, source_ip)?;
self.db.commit_soar_unblock_to_db(id, ip_version, source_ip).await?;
// Decrement active block counter
self.soar_engine.decrement_block_count();
@ -121,23 +128,23 @@ impl PlaybookService {
Ok(())
}
pub fn list_executions(&self, limit: i64) -> Result<Vec<ExecutionView>, Error> {
self.db.list_soar_executions(limit)
pub async fn list_executions(&self, limit: i64) -> Result<Vec<ExecutionView>, Error> {
self.db.list_soar_executions(limit).await
}
pub fn list_whitelist(&self) -> Result<Vec<String>, Error> {
self.db.list_admin_whitelist()
pub async fn list_whitelist(&self) -> Result<Vec<String>, Error> {
self.db.list_admin_whitelist().await
}
pub fn add_whitelist(&self, ip: &str) -> Result<(), Error> {
self.db.insert_admin_whitelist(ip)?;
self.soar_engine.reload_cache()?;
pub async fn add_whitelist(&self, ip: &str) -> Result<(), Error> {
self.db.insert_admin_whitelist(ip).await?;
self.soar_engine.reload_cache().await?;
Ok(())
}
pub fn remove_whitelist(&self, ip: &str) -> Result<(), Error> {
self.db.delete_admin_whitelist(ip)?;
self.soar_engine.reload_cache()?;
pub async fn remove_whitelist(&self, ip: &str) -> Result<(), Error> {
self.db.delete_admin_whitelist(ip).await?;
self.soar_engine.reload_cache().await?;
Ok(())
}
}

View File

@ -2,6 +2,7 @@ use std::sync::Arc;
use macros::log;
use tokio::task::JoinHandle;
use tokio::task::spawn_blocking;
use tokio::time::{self, Duration};
use crate::core::response::engine::SoarEngine;
@ -56,7 +57,7 @@ impl TtlScheduler {
// Clean up expired cooldown + frequency tracker entries to prevent unbounded memory growth
self.soar_engine.cleanup_expired_cooldowns();
let expired = self.db.list_expired_soar_blocks()?;
let expired = self.db.list_expired_soar_blocks().await?;
if expired.is_empty() {
return Ok(());
@ -67,11 +68,11 @@ impl TtlScheduler {
for block in &expired {
// Check if a manual ACL rule exists for this IP
let has_manual_rule = self.db.has_manual_acl_rule(&block.source_ip)?;
let has_manual_rule = self.db.has_manual_acl_rule(&block.source_ip).await?;
if has_manual_rule {
// Only mark as unblocked in SOAR records, don't remove from eBPF
self.db.mark_soar_block_unblocked(block.id)?;
self.db.mark_soar_block_unblocked(block.id).await?;
self.soar_engine.decrement_block_count();
skipped += 1;
log!(SoarLog::WhitelistSkipped(
@ -82,18 +83,22 @@ impl TtlScheduler {
}
// Remove from eBPF ACL via AccessControlPort
if let Err(e) = self.access_control.unblock_ip(&block.source_ip) {
if let Err(e) = unblock_ip_blocking(Arc::clone(&self.access_control), block.source_ip.clone()).await {
log!(SoarLog::RecoveryFailed(
block.source_ip.clone(),
format!("unblock failed: {}", e)
));
self.db.insert_pending_unblock(&block.source_ip).await?;
skipped += 1;
continue;
}
// Atomically drop acl_rules entry AND mark soar_block_rules
// unblocked in one transaction.
let ip_version = ip_version_from_str(&block.source_ip);
self.db
.commit_soar_unblock_to_db(block.id, ip_version, &block.source_ip)?;
.commit_soar_unblock_to_db(block.id, ip_version, &block.source_ip)
.await?;
self.soar_engine.decrement_block_count();
removed += 1;
}
@ -105,3 +110,210 @@ impl TtlScheduler {
Ok(())
}
}
async fn unblock_ip_blocking(access_control: Arc<dyn AccessControlPort>, source_ip: String) -> Result<(), Error> {
spawn_blocking(move || access_control.unblock_ip(&source_ip))
.await
.map_err(|e| crate::domain::response::error::SoarError::ActionFailed("unblock_ip", e))?
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicBool, AtomicU8, Ordering};
use arc_swap::ArcSwap;
use parking_lot::Mutex;
use super::*;
use crate::adapter::persistence::Database;
use crate::core::response::engine::SoarEngineDeps;
use crate::domain::common::config::AppConfig;
use crate::domain::common::config::notification::SmtpConfig;
use crate::domain::data_plane::error::EbpfError;
use crate::interface::email_sender::{EmailSender, EmailSenderFactory};
use crate::interface::secret_store::SecretStorePort;
const EXPIRED_AT: &str = "2000-01-01 00:00:00";
const SOURCE_IP: &str = "198.51.100.10";
struct MockAccessControl {
unblocked_ips: Mutex<Vec<String>>,
should_fail: AtomicBool,
}
impl MockAccessControl {
fn new(should_fail: bool) -> Self {
Self {
unblocked_ips: Mutex::new(Vec::new()),
should_fail: AtomicBool::new(should_fail),
}
}
}
impl AccessControlPort for MockAccessControl {
fn block_ip(&self, _ip: &str) -> Result<(), Error> {
Ok(())
}
fn unblock_ip(&self, ip: &str) -> Result<(), Error> {
if self.should_fail.load(Ordering::SeqCst) {
Err(EbpfError::UnknownError)?;
}
self.unblocked_ips.lock().push(ip.to_string());
Ok(())
}
}
struct NoopEmailSenderFactory;
#[async_trait::async_trait]
impl EmailSenderFactory for NoopEmailSenderFactory {
async fn build_smtp_sender(
&self,
_cfg: &SmtpConfig,
_secrets: Option<&dyn SecretStorePort>,
) -> Result<Option<Box<dyn EmailSender>>, Error> {
Ok(None)
}
}
async fn test_db() -> Arc<Database> {
let db = Arc::new(Database::new(":memory:").await.expect("test db"));
AppConfig::seed_config_defaults(&*db)
.await
.expect("seed config defaults");
db
}
async fn test_scheduler(
db: Arc<Database>,
access_control: Arc<MockAccessControl>,
) -> (TtlScheduler, Arc<SoarEngine>) {
let cfg = AppConfig::from_config_repo(&*db).await.expect("load config");
let engine = Arc::new(
SoarEngine::new(SoarEngineDeps {
db: db.clone() as Arc<dyn AppRepo>,
config: Arc::new(ArcSwap::from_pointee(cfg)),
access_control: access_control.clone(),
alert_notifier: None,
geoip: None,
rate_limit: None,
enforce_level_cache: Arc::new(AtomicU8::new(2)),
secrets: None,
email_sender_factory: Arc::new(NoopEmailSenderFactory),
})
.await
.expect("soar engine"),
);
let scheduler = TtlScheduler::new(db, access_control, engine.clone());
(scheduler, engine)
}
fn acl_contains(rules: &[crate::domain::data_plane::acl_rule::AclRuleView], ip: &str) -> bool {
rules
.iter()
.any(|rule| rule.ip_address == ip && rule.direction == "source" && rule.list_type == "blacklist")
}
#[tokio::test]
async fn ttl_sweep_preserves_manual_acl_rule_for_expired_soar_block() {
let db = test_db().await;
let block_id = db
.commit_soar_block_to_db(SOURCE_IP, 4, 1, EXPIRED_AT)
.await
.expect("insert expired block");
db.insert_acl_rule(4, "source", "blacklist", SOURCE_IP, 0)
.await
.expect("manual ACL should preserve block");
let access_control = Arc::new(MockAccessControl::new(false));
let (scheduler, engine) = test_scheduler(db.clone(), access_control.clone()).await;
engine.matcher.active_block_count.store(1, Ordering::SeqCst);
scheduler.sweep().await.expect("ttl sweep");
assert!(
access_control.unblocked_ips.lock().is_empty(),
"manual ACL ownership must skip data-plane unblock"
);
assert!(
db.list_expired_soar_blocks().await.expect("expired blocks").is_empty(),
"expired SOAR record should be marked unblocked"
);
assert!(
acl_contains(&db.list_acl_rules().await.expect("acl rules"), SOURCE_IP),
"manual ACL row must remain after SOAR TTL expiry"
);
assert_eq!(
engine.matcher.active_block_count.load(Ordering::SeqCst),
0,
"SOAR active-block counter should decrement once"
);
assert!(db.find_soar_block_by_id(block_id).await.expect("find block").is_some());
}
#[tokio::test]
async fn ttl_sweep_records_pending_unblock_without_clearing_db_when_data_plane_unblock_fails() {
let db = test_db().await;
db.commit_soar_block_to_db(SOURCE_IP, 4, 1, EXPIRED_AT)
.await
.expect("insert expired block");
let access_control = Arc::new(MockAccessControl::new(true));
let (scheduler, engine) = test_scheduler(db.clone(), access_control.clone()).await;
engine.matcher.active_block_count.store(1, Ordering::SeqCst);
scheduler.sweep().await.expect("ttl sweep");
let pending = db.list_pending_unblocks().await.expect("pending unblocks");
assert_eq!(pending.len(), 1);
assert_eq!(pending[0].source_ip, SOURCE_IP);
assert!(
!db.list_expired_soar_blocks().await.expect("expired blocks").is_empty(),
"failed data-plane unblock must leave SOAR block active for retry"
);
assert!(
acl_contains(&db.list_acl_rules().await.expect("acl rules"), SOURCE_IP),
"ACL row must remain while durable unblock has not succeeded"
);
assert!(
access_control.unblocked_ips.lock().is_empty(),
"failing data-plane call should not record a successful unblock"
);
assert_eq!(
engine.matcher.active_block_count.load(Ordering::SeqCst),
1,
"active-block counter must not decrement until unblock succeeds"
);
}
#[tokio::test]
async fn ttl_sweep_deletes_soar_owned_acl_only_after_data_plane_unblock_succeeds() {
let db = test_db().await;
db.commit_soar_block_to_db(SOURCE_IP, 4, 1, EXPIRED_AT)
.await
.expect("insert expired block");
let access_control = Arc::new(MockAccessControl::new(false));
let (scheduler, engine) = test_scheduler(db.clone(), access_control.clone()).await;
engine.matcher.active_block_count.store(1, Ordering::SeqCst);
scheduler.sweep().await.expect("ttl sweep");
assert_eq!(access_control.unblocked_ips.lock().as_slice(), &[SOURCE_IP.to_string()]);
assert!(
db.list_expired_soar_blocks().await.expect("expired blocks").is_empty(),
"successful unblock should mark SOAR record unblocked"
);
assert!(
!acl_contains(&db.list_acl_rules().await.expect("acl rules"), SOURCE_IP),
"SOAR-owned ACL row should be removed after data-plane unblock succeeds"
);
assert!(
db.list_pending_unblocks().await.expect("pending unblocks").is_empty(),
"successful unblock should not enqueue retry work"
);
assert_eq!(
engine.matcher.active_block_count.load(Ordering::SeqCst),
0,
"active-block counter should decrement after successful unblock"
);
}
}

View File

@ -3,6 +3,8 @@ use macros::config_settings;
#[config_settings(section = "misc")]
#[derive(Debug, Clone)]
pub struct AclConfig {
#[setting(key = "geoip_db_name", default = "net-guardia/static/geo/dbip-city-lite.mmdb")]
pub geoip_db_name: String,
#[setting(key = "geoip_db_path", default = "net-guardia/static/geo/dbip-city-lite.mmdb")]
pub geoip_db_path: String,
#[setting(key = "geoip_cache_capacity", default = "10000")]
pub geoip_cache_capacity: usize,
}

View File

@ -9,6 +9,10 @@ pub struct FusionConfig {
pub repeat_offender_window_secs: u64,
#[setting(key = "fusion_max_dedup_entries", default = "50000")]
pub max_dedup_entries: usize,
#[setting(key = "fusion_source_count_max_entries", default = "10000")]
pub source_count_max_entries: usize,
#[setting(key = "fusion_repeat_tracker_max_entries", default = "5000")]
pub repeat_tracker_max_entries: usize,
}
#[config_settings(section = "beaconing")]
@ -22,12 +26,23 @@ pub struct BeaconingConfig {
pub cv_threshold: f64,
#[setting(key = "beaconing_max_cache_entries", default = "50000")]
pub max_cache_entries: usize,
#[setting(key = "beaconing_max_timestamps_per_flow", default = "100")]
pub max_timestamps_per_flow: usize,
#[setting(key = "beaconing_expiry_secs", default = "600")]
pub expiry_secs: u64,
#[setting(key = "beaconing_alert_cooldown_secs", default = "300")]
pub alert_cooldown_secs: u64,
}
#[config_settings(section = "flow_stats")]
#[derive(Debug, Clone)]
pub struct FlowStatsConfig {
#[setting(key = "flow_stats_max_snapshot_entries", default = "10000")]
pub max_snapshot_entries: usize,
#[setting(key = "flow_stats_max_top_n", default = "10000")]
pub max_top_n: usize,
}
#[config_settings(section = "detection")]
#[derive(Debug, Clone)]
pub struct DetectionConfig {
@ -35,6 +50,8 @@ pub struct DetectionConfig {
pub fusion: FusionConfig,
#[setting(flatten)]
pub beaconing: BeaconingConfig,
#[setting(flatten)]
pub flow_stats: FlowStatsConfig,
#[setting(key = "detection_cleanup_interval_secs", default = "60")]
pub cleanup_interval_secs: u64,
}

View File

@ -32,7 +32,12 @@ pub struct EbpfConfig {
pub packet_buffer_size: usize,
#[setting(section = "xdp", key = "buffer_pool_capacity", default = "1024")]
pub buffer_pool_capacity: usize,
#[setting(section = "xdp", key = "xsk_completion_batch_size", default = "256")]
pub xsk_completion_batch_size: usize,
#[setting(section = "xdp", key = "xsk_rx_batch_size", default = "64")]
pub xsk_rx_batch_size: usize,
#[setting(section = "xdp", key = "xsk_tx_batch_size", default = "64")]
pub xsk_tx_batch_size: usize,
// ── internal (not exposed in API) ──────────────────────────────
#[setting(section = "xdp", key = "default_packet_rate", default = "10000", api = false)]
pub default_packet_rate: u64,

View File

@ -21,4 +21,6 @@ pub struct HealthConfig {
pub temp_warn_celsius: f32,
#[setting(key = "health_monitoring_interval_secs", default = "5")]
pub monitoring_interval_secs: u64,
#[setting(key = "health_broadcast_channel_capacity", default = "100")]
pub broadcast_channel_capacity: usize,
}

View File

@ -1,14 +1,14 @@
use std::str::FromStr;
use crate::domain::common::error::Error;
use crate::interface::setting::SettingRepo;
use crate::interface::config_repo::ConfigRepo;
pub(in crate::domain) fn override_parsed<T: FromStr>(
pub(in crate::domain) async fn override_parsed<T: FromStr>(
target: &mut T,
repo: &dyn SettingRepo,
repo: &dyn ConfigRepo,
key: &str,
) -> Result<(), Error> {
if let Some(v) = repo.get_setting(key)?
if let Some(v) = repo.get_config_value(key).await?
&& let Ok(parsed) = v.parse()
{
*target = parsed;
@ -16,19 +16,19 @@ pub(in crate::domain) fn override_parsed<T: FromStr>(
Ok(())
}
pub(in crate::domain) fn override_bool(target: &mut bool, repo: &dyn SettingRepo, key: &str) -> Result<(), Error> {
if let Some(v) = repo.get_setting(key)? {
pub(in crate::domain) async fn override_bool(target: &mut bool, repo: &dyn ConfigRepo, key: &str) -> Result<(), Error> {
if let Some(v) = repo.get_config_value(key).await? {
*target = v == "true" || v == "1";
}
Ok(())
}
pub(in crate::domain) fn override_string_nonempty(
pub(in crate::domain) async fn override_string_nonempty(
target: &mut String,
repo: &dyn SettingRepo,
repo: &dyn ConfigRepo,
key: &str,
) -> Result<(), Error> {
if let Some(v) = repo.get_setting(key)?
if let Some(v) = repo.get_config_value(key).await?
&& !v.is_empty()
{
*target = v;
@ -36,12 +36,12 @@ pub(in crate::domain) fn override_string_nonempty(
Ok(())
}
pub(in crate::domain) fn override_csv(
pub(in crate::domain) async fn override_csv(
target: &mut Vec<String>,
repo: &dyn SettingRepo,
repo: &dyn ConfigRepo,
key: &str,
) -> Result<(), Error> {
if let Some(v) = repo.get_setting(key)? {
if let Some(v) = repo.get_config_value(key).await? {
*target = if v.is_empty() {
Vec::new()
} else {
@ -51,9 +51,9 @@ pub(in crate::domain) fn override_csv(
Ok(())
}
pub(in crate::domain) fn seed_key(repo: &dyn SettingRepo, key: &str, value: &str) -> Result<(), Error> {
if repo.get_setting(key)?.is_none() {
repo.set_setting(key, value)?;
pub(in crate::domain) async fn seed_key(repo: &dyn ConfigRepo, key: &str, value: &str) -> Result<(), Error> {
if repo.get_config_value(key).await?.is_none() {
repo.set_config_value(key, value).await?;
}
Ok(())
}

View File

@ -3,11 +3,6 @@ use macros::config_settings;
#[config_settings]
#[derive(Debug, Clone)]
pub struct MlConfig {
// ── models ─────────────────────────────────────────────────────
#[setting(section = "models", key = "deep_autoencoder_name", default = "deep_autoencoder.onnx")]
pub deep_autoencoder_name: String,
#[setting(section = "models", key = "classifier_name", default = "classifier.onnx")]
pub classifier_name: String,
#[setting(section = "models", key = "models_config_name", default = "inference_config.json")]
pub models_config_name: String,

View File

@ -16,6 +16,8 @@ pub mod soar;
pub mod suricata;
pub mod system;
use std::collections::HashMap;
use crate::domain::common::config::acl::AclConfig;
use crate::domain::common::config::correlation::CorrelationConfig;
use crate::domain::common::config::detection::DetectionConfig;
@ -32,7 +34,7 @@ use crate::domain::common::config::suricata::SuricataConfig;
use crate::domain::common::config::system::SystemConfig;
use crate::domain::common::error::Error;
use crate::domain::common::error::system::SystemError;
use crate::interface::setting::SettingRepo;
use crate::interface::config_repo::ConfigRepo;
#[derive(Debug, Clone)]
pub struct AppConfig {
@ -53,45 +55,61 @@ pub struct AppConfig {
}
impl AppConfig {
pub fn from_settings(repo: &dyn SettingRepo) -> Result<Self, Error> {
pub async fn from_config_repo(repo: &dyn ConfigRepo) -> Result<Self, Error> {
let cfg = Self {
acl: AclConfig::from_settings(repo)?,
correlation: CorrelationConfig::from_settings(repo)?,
detection: DetectionConfig::from_settings(repo)?,
dns_filter: DnsFilterConfig::from_settings(repo)?,
ebpf: EbpfConfig::from_settings(repo)?,
health: HealthConfig::from_settings(repo)?,
http_server: HttpServerConfig::from_settings(repo)?,
ml: MlConfig::from_settings(repo)?,
notification: NotificationConfig::from_settings(repo)?,
observability: ObservabilityConfig::from_settings(repo)?,
pipeline: PipelineConfig::from_settings(repo)?,
soar: SoarConfig::from_settings(repo)?,
suricata: SuricataConfig::from_settings(repo)?,
system: SystemConfig::from_settings(repo)?,
acl: AclConfig::from_config_repo(repo).await?,
correlation: CorrelationConfig::from_config_repo(repo).await?,
detection: DetectionConfig::from_config_repo(repo).await?,
dns_filter: DnsFilterConfig::from_config_repo(repo).await?,
ebpf: EbpfConfig::from_config_repo(repo).await?,
health: HealthConfig::from_config_repo(repo).await?,
http_server: HttpServerConfig::from_config_repo(repo).await?,
ml: MlConfig::from_config_repo(repo).await?,
notification: NotificationConfig::from_config_repo(repo).await?,
observability: ObservabilityConfig::from_config_repo(repo).await?,
pipeline: PipelineConfig::from_config_repo(repo).await?,
soar: SoarConfig::from_config_repo(repo).await?,
suricata: SuricataConfig::from_config_repo(repo).await?,
system: SystemConfig::from_config_repo(repo).await?,
};
cfg.validate()?;
Ok(cfg)
}
pub fn seed_defaults(repo: &dyn SettingRepo) -> Result<(), Error> {
AclConfig::seed_defaults(repo)?;
CorrelationConfig::seed_defaults(repo)?;
DetectionConfig::seed_defaults(repo)?;
DnsFilterConfig::seed_defaults(repo)?;
EbpfConfig::seed_defaults(repo)?;
HealthConfig::seed_defaults(repo)?;
HttpServerConfig::seed_defaults(repo)?;
MlConfig::seed_defaults(repo)?;
NotificationConfig::seed_defaults(repo)?;
ObservabilityConfig::seed_defaults(repo)?;
PipelineConfig::seed_defaults(repo)?;
SoarConfig::seed_defaults(repo)?;
SuricataConfig::seed_defaults(repo)?;
SystemConfig::seed_defaults(repo)?;
pub async fn seed_config_defaults(repo: &dyn ConfigRepo) -> Result<(), Error> {
AclConfig::seed_config_defaults(repo).await?;
CorrelationConfig::seed_config_defaults(repo).await?;
DetectionConfig::seed_config_defaults(repo).await?;
DnsFilterConfig::seed_config_defaults(repo).await?;
EbpfConfig::seed_config_defaults(repo).await?;
HealthConfig::seed_config_defaults(repo).await?;
HttpServerConfig::seed_config_defaults(repo).await?;
MlConfig::seed_config_defaults(repo).await?;
NotificationConfig::seed_config_defaults(repo).await?;
ObservabilityConfig::seed_config_defaults(repo).await?;
PipelineConfig::seed_config_defaults(repo).await?;
SoarConfig::seed_config_defaults(repo).await?;
SuricataConfig::seed_config_defaults(repo).await?;
SystemConfig::seed_config_defaults(repo).await?;
Ok(())
}
pub fn api_setting_values(&self) -> HashMap<&'static str, String> {
let mut values = HashMap::new();
values.extend(self.acl.api_values());
values.extend(self.correlation.api_values());
values.extend(self.detection.api_values());
values.extend(self.dns_filter.api_values());
values.extend(self.ebpf.api_values());
values.extend(self.http_server.api_values());
values.extend(self.ml.api_values());
values.extend(self.notification.api_values());
values.extend(self.observability.api_values());
values.extend(self.soar.api_values());
values.extend(self.suricata.api_values());
values
}
fn validate(&self) -> Result<(), Error> {
let require = |ok: bool, field: &str| -> Result<(), Error> {
if ok {
@ -109,6 +127,12 @@ impl AppConfig {
require(self.ebpf.rx_queue_size > 0, "ebpf.rx_queue_size")?;
require(self.ebpf.frame_size > 0, "ebpf.frame_size")?;
require(self.ebpf.frame_count > 0, "ebpf.frame_count")?;
require(
self.ebpf.xsk_completion_batch_size > 0,
"ebpf.xsk_completion_batch_size",
)?;
require(self.ebpf.xsk_rx_batch_size > 0, "ebpf.xsk_rx_batch_size")?;
require(self.ebpf.xsk_tx_batch_size > 0, "ebpf.xsk_tx_batch_size")?;
require(self.http_server.port > 0, "http_server.port")?;
require(self.ml.max_concurrent_flows > 0, "ml.max_concurrent_flows")?;
require(self.ml.min_packets_for_inference > 0, "ml.min_packets_for_inference")?;
@ -142,6 +166,12 @@ impl AppConfig {
self.soar.frequency_max_tracked_keys > 0,
"soar.frequency_max_tracked_keys",
)?;
require(self.soar.execution_list_limit > 0, "soar.execution_list_limit")?;
require(
self.soar.frequency_max_events_per_key > 0,
"soar.frequency_max_events_per_key",
)?;
require(self.soar.frequency_retention_secs > 0, "soar.frequency_retention_secs")?;
require(
(0.0..=1.0).contains(&self.soar.default_rate_limit_factor),
"soar.default_rate_limit_factor",
@ -158,6 +188,14 @@ impl AppConfig {
self.detection.fusion.max_dedup_entries > 0,
"detection.fusion.max_dedup_entries",
)?;
require(
self.detection.fusion.source_count_max_entries > 0,
"detection.fusion.source_count_max_entries",
)?;
require(
self.detection.fusion.repeat_tracker_max_entries > 0,
"detection.fusion.repeat_tracker_max_entries",
)?;
require(
self.detection.fusion.dedup_window_secs > 0,
"detection.fusion.dedup_window_secs",
@ -170,6 +208,26 @@ impl AppConfig {
self.detection.beaconing.max_cache_entries > 0,
"detection.beaconing.max_cache_entries",
)?;
require(
self.detection.beaconing.max_timestamps_per_flow > 1,
"detection.beaconing.max_timestamps_per_flow",
)?;
require(
self.detection.flow_stats.max_snapshot_entries > 0,
"detection.flow_stats.max_snapshot_entries",
)?;
require(
self.detection.flow_stats.max_top_n >= self.detection.flow_stats.max_snapshot_entries,
"detection.flow_stats.max_top_n",
)?;
require(
self.health.monitoring_interval_secs > 0,
"health.monitoring_interval_secs",
)?;
require(
self.health.broadcast_channel_capacity > 0,
"health.broadcast_channel_capacity",
)?;
require(
self.correlation.max_tracked_entries > 0,
"correlation.max_tracked_entries",
@ -218,6 +276,7 @@ impl AppConfig {
self.observability.drop_channel_capacity > 0,
"observability.drop_channel_capacity",
)?;
require(self.acl.geoip_cache_capacity > 0, "acl.geoip_cache_capacity")?;
require(self.suricata.poll_interval_ms > 0, "suricata.poll_interval_ms")?;
require(
(0.0..=1.0).contains(&self.suricata.confidence_high),
@ -245,134 +304,214 @@ mod tests {
use super::*;
use crate::adapter::persistence::Database;
fn test_db() -> Database {
async fn test_db() -> Database {
// SAFETY: in-memory SQLite open is infallible under standard library features.
Database::new(":memory:").expect("in-memory DB")
Database::new(":memory:").await.expect("in-memory DB")
}
#[test]
fn defaults_are_valid() {
let db = test_db();
let cfg = AppConfig::from_settings(&db).expect("defaults should be valid");
#[tokio::test]
async fn defaults_are_valid() {
let db = test_db().await;
let cfg = AppConfig::from_config_repo(&db)
.await
.expect("defaults should be valid");
assert_eq!(cfg.http_server.port, 8080);
assert_eq!(cfg.ebpf.ingress_ifname, "eth0");
assert_eq!(cfg.ebpf.egress_ifname, "eth1");
assert_eq!(cfg.ebpf.combined_queue_count, 1);
assert_eq!(cfg.ebpf.frame_size, 4096);
assert_eq!(cfg.ebpf.xsk_completion_batch_size, 256);
assert_eq!(cfg.ebpf.xsk_rx_batch_size, 64);
assert_eq!(cfg.ebpf.xsk_tx_batch_size, 64);
assert_eq!(cfg.http_server.jwt_expiry_hours, 24);
}
#[test]
fn db_overrides_interface_names() {
let db = test_db();
db.set_setting("ingress_interface", "ens33").unwrap();
db.set_setting("egress_interface", "ens34").unwrap();
let cfg = AppConfig::from_settings(&db).unwrap();
#[tokio::test]
async fn db_overrides_interface_names() {
let db = test_db().await;
db.set_config_value("ingress_interface", "ens33").await.unwrap();
db.set_config_value("egress_interface", "ens34").await.unwrap();
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert_eq!(cfg.ebpf.ingress_ifname, "ens33");
assert_eq!(cfg.ebpf.egress_ifname, "ens34");
}
#[test]
fn db_overrides_http_port() {
let db = test_db();
db.set_setting("http_port", "9090").unwrap();
let cfg = AppConfig::from_settings(&db).unwrap();
#[tokio::test]
async fn db_overrides_http_port() {
let db = test_db().await;
db.set_config_value("http_port", "9090").await.unwrap();
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert_eq!(cfg.http_server.port, 9090);
}
#[test]
fn db_overrides_xdp_tuning() {
let db = test_db();
db.set_setting("frame_size", "8192").unwrap();
db.set_setting("combined_queue_count", "4").unwrap();
let cfg = AppConfig::from_settings(&db).unwrap();
#[tokio::test]
async fn db_overrides_xdp_tuning() {
let db = test_db().await;
db.set_config_value("frame_size", "8192").await.unwrap();
db.set_config_value("combined_queue_count", "4").await.unwrap();
db.set_config_value("xsk_completion_batch_size", "512").await.unwrap();
db.set_config_value("xsk_rx_batch_size", "128").await.unwrap();
db.set_config_value("xsk_tx_batch_size", "32").await.unwrap();
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert_eq!(cfg.ebpf.frame_size, 8192);
assert_eq!(cfg.ebpf.combined_queue_count, 4);
assert_eq!(cfg.ebpf.xsk_completion_batch_size, 512);
assert_eq!(cfg.ebpf.xsk_rx_batch_size, 128);
assert_eq!(cfg.ebpf.xsk_tx_batch_size, 32);
}
#[test]
fn invalid_db_values_ignored() {
let db = test_db();
db.set_setting("http_port", "not_a_number").unwrap();
let cfg = AppConfig::from_settings(&db).unwrap();
#[tokio::test]
async fn invalid_db_values_ignored() {
let db = test_db().await;
db.set_config_value("http_port", "not_a_number").await.unwrap();
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert_eq!(cfg.http_server.port, 8080);
}
#[test]
fn empty_db_uses_all_defaults() {
let db = test_db();
let cfg = AppConfig::from_settings(&db).unwrap();
#[tokio::test]
async fn empty_db_uses_all_defaults() {
let db = test_db().await;
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert_eq!(cfg.ml.inference_interval_secs, 5);
assert_eq!(cfg.ml.inference_batch_size, 200);
assert_eq!(cfg.system.database_path, "net-guardia.db");
}
#[test]
fn seed_defaults_populates_empty_db() {
let db = test_db();
AppConfig::seed_defaults(&db).expect("seed should succeed");
assert_eq!(db.get_setting("http_port").unwrap(), Some("8080".to_string()));
#[tokio::test]
async fn seed_config_defaults_populates_empty_db() {
let db = test_db().await;
AppConfig::seed_config_defaults(&db).await.expect("seed should succeed");
assert_eq!(
db.get_setting("traffic_logging_mode").unwrap(),
db.get_config_value("http_port").await.unwrap(),
Some("8080".to_string())
);
assert_eq!(
db.get_config_value("traffic_logging_mode").await.unwrap(),
Some("false".to_string())
);
assert_eq!(
db.get_setting("pipeline_ingress").unwrap(),
db.get_config_value("pipeline_ingress").await.unwrap(),
Some("access_control,rate_limit,service".to_string())
);
assert_eq!(
db.get_setting("geoip_db_name").unwrap(),
db.get_config_value("geoip_db_path").await.unwrap(),
Some("net-guardia/static/geo/dbip-city-lite.mmdb".to_string())
);
assert_eq!(
db.get_setting("soar_max_auto_block_cap").unwrap(),
db.get_config_value("soar_max_auto_block_cap").await.unwrap(),
Some("100".to_string())
);
assert_eq!(db.get_setting("soar_max_ttl_secs").unwrap(), Some("86400".to_string()));
assert_eq!(
db.get_setting("ml_drift_window_secs").unwrap(),
db.get_config_value("soar_max_ttl_secs").await.unwrap(),
Some("86400".to_string())
);
assert_eq!(
db.get_config_value("soar_execution_list_limit").await.unwrap(),
Some("100".to_string())
);
assert_eq!(
db.get_config_value("ml_drift_window_secs").await.unwrap(),
Some("3600".to_string())
);
assert_eq!(
db.get_config_value("fusion_source_count_max_entries").await.unwrap(),
Some("10000".to_string())
);
assert_eq!(
db.get_config_value("fusion_repeat_tracker_max_entries").await.unwrap(),
Some("5000".to_string())
);
assert_eq!(
db.get_config_value("health_broadcast_channel_capacity").await.unwrap(),
Some("100".to_string())
);
assert_eq!(
db.get_config_value("geoip_cache_capacity").await.unwrap(),
Some("10000".to_string())
);
}
#[test]
fn seed_defaults_does_not_overwrite_existing() {
let db = test_db();
db.set_setting("http_port", "9090").unwrap();
AppConfig::seed_defaults(&db).expect("seed should succeed");
assert_eq!(db.get_setting("http_port").unwrap(), Some("9090".to_string()));
#[tokio::test]
async fn seed_config_defaults_does_not_overwrite_existing() {
let db = test_db().await;
db.set_config_value("http_port", "9090").await.unwrap();
AppConfig::seed_config_defaults(&db).await.expect("seed should succeed");
assert_eq!(
db.get_config_value("http_port").await.unwrap(),
Some("9090".to_string())
);
}
#[test]
fn db_overrides_traffic_logging_mode() {
let db = test_db();
db.set_setting("traffic_logging_mode", "true").unwrap();
let cfg = AppConfig::from_settings(&db).unwrap();
#[tokio::test]
async fn db_overrides_traffic_logging_mode() {
let db = test_db().await;
db.set_config_value("traffic_logging_mode", "true").await.unwrap();
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert!(cfg.ml.traffic_logging_mode);
}
#[test]
fn default_traffic_logging_mode_is_false() {
let db = test_db();
let cfg = AppConfig::from_settings(&db).unwrap();
#[tokio::test]
async fn default_traffic_logging_mode_is_false() {
let db = test_db().await;
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert!(!cfg.ml.traffic_logging_mode);
}
#[test]
fn db_overrides_pipeline() {
let db = test_db();
db.set_setting("pipeline_ingress", "access_control,service").unwrap();
let cfg = AppConfig::from_settings(&db).unwrap();
#[tokio::test]
async fn db_overrides_pipeline() {
let db = test_db().await;
db.set_config_value("pipeline_ingress", "access_control,service")
.await
.unwrap();
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert_eq!(cfg.pipeline.ingress, vec!["access_control", "service"]);
}
#[test]
fn db_overrides_soar_and_drift() {
let db = test_db();
db.set_setting("soar_max_ttl_secs", "3600").unwrap();
db.set_setting("ml_drift_window_secs", "900").unwrap();
let cfg = AppConfig::from_settings(&db).unwrap();
#[tokio::test]
async fn db_overrides_soar_and_drift() {
let db = test_db().await;
db.set_config_value("soar_max_ttl_secs", "3600").await.unwrap();
db.set_config_value("soar_execution_list_limit", "42").await.unwrap();
db.set_config_value("ml_drift_window_secs", "900").await.unwrap();
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert_eq!(cfg.soar.max_ttl_secs, 3600);
assert_eq!(cfg.soar.execution_list_limit, 42);
assert_eq!(cfg.ml.drift_window_secs, 900);
}
#[tokio::test]
async fn db_overrides_fusion_runtime_cache_sizes() {
let db = test_db().await;
db.set_config_value("fusion_source_count_max_entries", "1234")
.await
.unwrap();
db.set_config_value("fusion_repeat_tracker_max_entries", "567")
.await
.unwrap();
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert_eq!(cfg.detection.fusion.source_count_max_entries, 1234);
assert_eq!(cfg.detection.fusion.repeat_tracker_max_entries, 567);
}
#[tokio::test]
async fn db_overrides_health_runtime_values() {
let db = test_db().await;
db.set_config_value("health_monitoring_interval_secs", "9")
.await
.unwrap();
db.set_config_value("health_broadcast_channel_capacity", "321")
.await
.unwrap();
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert_eq!(cfg.health.monitoring_interval_secs, 9);
assert_eq!(cfg.health.broadcast_channel_capacity, 321);
}
#[tokio::test]
async fn db_overrides_geoip_runtime_cache_capacity() {
let db = test_db().await;
db.set_config_value("geoip_cache_capacity", "2048").await.unwrap();
let cfg = AppConfig::from_config_repo(&db).await.unwrap();
assert_eq!(cfg.acl.geoip_cache_capacity, 2048);
}
}

View File

@ -29,6 +29,12 @@ pub struct SoarConfig {
pub rate_limit_cmd_channel_capacity: usize,
#[setting(key = "soar_frequency_max_tracked_keys", default = "50000")]
pub frequency_max_tracked_keys: usize,
#[setting(key = "soar_frequency_max_events_per_key", default = "200")]
pub frequency_max_events_per_key: usize,
#[setting(key = "soar_frequency_retention_secs", default = "7200")]
pub frequency_retention_secs: u64,
#[setting(key = "soar_fallback_cooldown_secs", default = "300")]
pub fallback_cooldown_secs: i64,
#[setting(key = "soar_execution_list_limit", default = "100")]
pub execution_list_limit: i64,
}

View File

@ -3,6 +3,8 @@ use macros::config_settings;
#[config_settings(section = "system")]
#[derive(Debug, Clone)]
pub struct SystemConfig {
#[setting(key = "enforce_mode", default = "monitor")]
pub enforce_mode: String,
#[setting(key = "database_path", default = "net-guardia.db", api = false)]
pub database_path: String,
#[setting(key = "report_dir", default = "/var/lib/netguardia/reports", api = false)]

View File

@ -5,9 +5,6 @@ traceable! {
#[error("Database error: {err}")]
QueryFailed => tracing::Level::ERROR,
#[error("Database connection failed")]
ConnectionFailed => tracing::Level::ERROR,
#[no_source]
#[error("User '{username}' already exists")]
UserAlreadyExists { username: String } => tracing::Level::WARN,

View File

@ -1,25 +0,0 @@
use macros::traceable;
traceable! {
McpError {
#[no_source]
#[error("MCP tool not found: {tool_name}")]
ToolNotFound { tool_name: String } => tracing::Level::WARN,
#[no_source]
#[error("MCP parameter validation failed: {detail}")]
InvalidParams { detail: String } => tracing::Level::WARN,
#[no_source]
#[error("MCP permission denied: key has '{key_level}' but tool requires '{required_level}'")]
PermissionDenied { key_level: String, required_level: String } => tracing::Level::WARN,
#[no_source]
#[error("MCP API key invalid or revoked")]
InvalidApiKey => tracing::Level::WARN,
#[no_source]
#[error("MCP proxy error: {detail}")]
ProxyError { detail: String } => tracing::Level::ERROR,
}
}

View File

@ -6,19 +6,9 @@ traceable! {
#[error("Failed to remove limit on locked memory, ret is: {ret}")]
RamLimitUnlockError { ret: i32 } => tracing::Level::ERROR,
#[error("Failed to send message to receiver")]
SendMessageError => tracing::Level::ERROR,
#[error("Failed to serialize data")]
SerializeError => tracing::Level::ERROR,
#[error("Failed to deserialize data")]
DeserializeError => tracing::Level::ERROR,
#[no_source]
#[error("Network interface '{interface}' not found")]
NetworkInterfaceNotFound { interface: String } => tracing::Level::ERROR,
#[error("Failed to open GeoIP database '{path}': {err}")]
GeoIPDatabaseError { path: String } => tracing::Level::ERROR,
@ -33,18 +23,6 @@ traceable! {
#[error("DNS domain name too long: '{domain}'")]
DnsDomainTooLong { domain: String } => tracing::Level::WARN,
#[no_source]
#[error("Type mismatch during message dispatch")]
TypeMismatch => tracing::Level::ERROR,
#[no_source]
#[error("No handler registered for this message type")]
HandlerNotFound => tracing::Level::ERROR,
#[no_source]
#[error("Event type not registered with communication manager")]
TypeNotRegistered => tracing::Level::ERROR,
#[no_source]
#[error("Validation error: {message}")]
ValidationError { message: String } => tracing::Level::WARN,

View File

@ -2,7 +2,6 @@ pub mod crypto;
pub mod database;
pub mod http;
pub mod io;
pub mod mcp;
pub mod misc;
pub mod notification;
pub mod system;
@ -13,7 +12,6 @@ use crate::domain::common::error::crypto::CryptoError;
use crate::domain::common::error::database::DatabaseError;
use crate::domain::common::error::http::HttpError;
use crate::domain::common::error::io::IOError;
use crate::domain::common::error::mcp::McpError;
use crate::domain::common::error::misc::MiscError;
use crate::domain::common::error::notification::NotificationError;
use crate::domain::common::error::system::SystemError;
@ -40,8 +38,6 @@ pub enum Error {
#[error("{0}")]
IO(#[from] IOError),
#[error("{0}")]
Mcp(#[from] McpError),
#[error("{0}")]
Misc(#[from] MiscError),
#[error("{0}")]
Notification(#[from] NotificationError),

View File

@ -2,10 +2,6 @@ use macros::traceable;
traceable! {
SystemError {
#[no_source]
#[error("Unable to run as administrator")]
RunAsAdminFailed => tracing::Level::ERROR,
#[no_source]
#[error("Invalid configuration")]
InvalidConfig => tracing::Level::ERROR,
@ -17,38 +13,24 @@ traceable! {
#[error("Configuration file not found")]
ConfigNotFound => tracing::Level::ERROR,
#[error("Failed to terminate instance")]
TerminateError => tracing::Level::ERROR,
#[no_source]
#[error("Failed to send shutdown signal")]
ShutdownSignalFailed => tracing::Level::ERROR,
#[error("Unexpected thread panic")]
ThreadPanic => tracing::Level::ERROR,
#[error("Unexpected error")]
UnexpectedError => tracing::Level::ERROR,
#[error("Failed to reload config after setup")]
ConfigReloadFailed => tracing::Level::ERROR,
#[error("HTTP server error")]
HttpServerError => tracing::Level::ERROR,
#[error("Failed to set user groups")]
SetUserGroupsFailed => tracing::Level::WARN,
#[error("Failed to store XDP mode")]
XdpModeStoreFailed => tracing::Level::WARN,
#[error("Failed to update admin password during setup: {err}")]
SetupPasswordUpdateFailed => tracing::Level::ERROR,
#[error("Failed to mark setup as complete: {err}")]
SetupCompleteFlagFailed => tracing::Level::ERROR,
#[error("Failed to publish drift detected event")]
DriftEventPublishFailed => tracing::Level::WARN,
}
}

View File

@ -3,9 +3,6 @@ use tracing;
loggable! {
SystemLog {
#[error("Online now")]
Online => tracing::Level::INFO,
#[error("Initializing")]
Initializing => tracing::Level::INFO,
@ -27,9 +24,6 @@ loggable! {
#[error("Setup wizard completed — starting full system initialization")]
SetupCompleted => tracing::Level::INFO,
#[error("Config reloaded from DB: ingress={ingress}, egress={egress}")]
ConfigReloaded { ingress: String, egress: String } => tracing::Level::INFO,
#[error("Full system initialization complete — all services running")]
FullInitComplete => tracing::Level::INFO,

View File

@ -63,8 +63,8 @@ pub enum EbpfFailCategory {
InterfaceNotFound,
/// Interface exists but XDP native/SKB attach refused by driver.
XdpUnsupported,
/// AF_XDP bind rejected — driver does not implement AF_XDP on this kernel.
/// Common case: Intel i350 (igb) on kernel < 6.17.
/// AF_XDP bind rejected — driver or netdev capabilities do not support
/// the requested AF_XDP socket mode on this interface.
AfXdpUnsupported,
/// ENOMEM / RLIMIT_MEMLOCK exhausted.
MemlockExhausted,

View File

@ -2,6 +2,28 @@ use serde::{Deserialize, Serialize};
use crate::domain::data_plane::direction::Direction;
#[derive(Debug, Clone, Copy)]
pub struct FlowStatsLimits {
max_snapshot_entries: usize,
max_top_n: usize,
}
impl FlowStatsLimits {
pub fn new(max_snapshot_entries: usize, max_top_n: usize) -> Self {
Self {
max_snapshot_entries: max_snapshot_entries.max(1),
max_top_n: max_top_n.max(1),
}
}
pub fn result_limit(self, requested_top_n: Option<usize>) -> usize {
requested_top_n
.unwrap_or(self.max_snapshot_entries)
.max(1)
.min(self.max_top_n)
}
}
#[derive(Debug, Clone, Serialize)]
pub struct FlowStatsEntry {
pub direction: Direction,
@ -61,3 +83,29 @@ pub struct FlowSubscription {
/// Push interval in seconds (default 5)
pub interval_secs: Option<u64>,
}
#[cfg(test)]
mod tests {
use super::FlowStatsLimits;
#[test]
fn result_limit_uses_snapshot_default_when_top_n_absent() {
let limits = FlowStatsLimits::new(128, 1024);
assert_eq!(limits.result_limit(None), 128);
}
#[test]
fn result_limit_caps_requested_top_n() {
let limits = FlowStatsLimits::new(128, 256);
assert_eq!(limits.result_limit(Some(1_000)), 256);
}
#[test]
fn result_limit_never_returns_zero() {
let limits = FlowStatsLimits::new(0, 0);
assert_eq!(limits.result_limit(Some(0)), 1);
}
}

View File

@ -36,7 +36,8 @@ impl PermissionLevel {
pub fn permissions(self) -> &'static [&'static str] {
match self {
Self::ReadOnly => API_KEY_READ_ONLY_PERMISSIONS,
Self::ReadWrite | Self::FullAccess => API_KEY_READ_WRITE_PERMISSIONS,
Self::ReadWrite => API_KEY_READ_WRITE_PERMISSIONS,
Self::FullAccess => API_KEY_FULL_ACCESS_PERMISSIONS,
}
}
}
@ -61,6 +62,7 @@ pub const ADMIN_PERMISSIONS: &[&str] = &[
"system:read",
"system:write",
"system:admin",
"api_keys:admin",
"users:read",
"users:write",
"users:admin",
@ -114,6 +116,8 @@ pub const API_KEY_READ_ONLY_PERMISSIONS: &[&str] = &[
"system:read",
];
pub const API_KEY_FULL_ACCESS_PERMISSIONS: &[&str] = ADMIN_PERMISSIONS;
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct Claims {
pub sub: i64,
@ -122,3 +126,31 @@ pub struct Claims {
pub permissions: Vec<String>,
pub exp: usize,
}
#[cfg(test)]
mod tests {
use super::PermissionLevel;
#[test]
fn read_write_api_keys_do_not_receive_admin_permissions() {
let permissions = PermissionLevel::ReadWrite.permissions();
assert!(!permissions.contains(&"api_keys:admin"));
assert!(!permissions.contains(&"system:admin"));
assert!(!permissions.contains(&"users:admin"));
}
#[test]
fn full_access_api_keys_receive_admin_permissions() {
let permissions = PermissionLevel::FullAccess.permissions();
assert!(permissions.contains(&"api_keys:admin"));
assert!(permissions.contains(&"system:admin"));
assert!(permissions.contains(&"users:admin"));
}
#[test]
fn full_access_has_more_permissions_than_read_write() {
assert!(PermissionLevel::FullAccess.permissions().len() > PermissionLevel::ReadWrite.permissions().len());
}
}

View File

@ -1,6 +1,7 @@
use serde::Serialize;
/// Input for updating a playbook row (without actions/conditions).
#[derive(Clone)]
pub struct UpdatePlaybookInput {
pub name: String,
pub trigger_event: String,
@ -11,6 +12,7 @@ pub struct UpdatePlaybookInput {
}
/// Input for a single playbook action in atomic create/update operations.
#[derive(Clone)]
pub struct ActionInput {
pub action_order: i64,
pub action_type: String,
@ -18,6 +20,7 @@ pub struct ActionInput {
}
/// Input for creating a single playbook condition.
#[derive(Clone)]
pub struct CreateConditionInput {
pub condition_type: String,
pub operator: String,
@ -47,6 +50,7 @@ fn default_operator_for(condition_type: &str) -> &'static str {
}
/// Input for creating a new playbook.
#[derive(Clone)]
pub struct CreatePlaybookInput {
pub name: String,
pub trigger_event: String,

View File

@ -74,28 +74,12 @@ impl AuditLogger {
// Always emit a structured log line
log!(AuditLog::AuditEvent(event.actor.clone(), event.action.clone(),));
// SQLite insert via r2d2 is blocking; offload so it can't stall the
// tokio worker that drains the broadcast channel. The actor/action
// strings are cloned for the log call above; the event itself moves
// into the blocking task.
let db = self.db.clone();
let join = tokio::task::spawn_blocking(move || {
db.insert_audit_log(&event.actor, &event.action, &event.detail)
.map_err(|e| (e.to_string(), event.actor, event.action))
})
.await;
match join {
Ok(Ok(())) => {}
Ok(Err((err, actor, action))) => {
log!(AuditLog::AuditDbWriteFailed(err, actor, action));
}
Err(join_err) => {
log!(AuditLog::AuditDbWriteFailed(
format!("blocking task join failed: {join_err}"),
"<lost>".to_string(),
"<lost>".to_string(),
));
}
if let Err(e) = self
.db
.insert_audit_log(&event.actor, &event.action, &event.detail)
.await
{
log!(AuditLog::AuditDbWriteFailed(e.to_string(), event.actor, event.action));
}
}
@ -108,22 +92,8 @@ impl AuditLogger {
log!(AuditLog::AuditDriftEvent(event.drifted_features.len()));
let db = self.db.clone();
let join = tokio::task::spawn_blocking(move || {
db.insert_audit_log("system", "ml_drift_detected", &detail)
.map_err(|e| e.to_string())
})
.await;
match join {
Ok(Ok(())) => {}
Ok(Err(err)) => {
log!(AuditLog::AuditDriftDbWriteFailed(err));
}
Err(join_err) => {
log!(AuditLog::AuditDriftDbWriteFailed(format!(
"blocking task join failed: {join_err}"
)));
}
if let Err(e) = self.db.insert_audit_log("system", "ml_drift_detected", &detail).await {
log!(AuditLog::AuditDriftDbWriteFailed(e.to_string()));
}
}
}

View File

@ -36,12 +36,12 @@ pub enum Command {
VerifyAuditLog,
}
pub fn handle_subcommand(cli: &Cli) -> Result<(), Error> {
pub async 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),
Some(Command::VerifyAuditLog) => run_verify(db_path).await,
None => Ok(()),
}
}
@ -86,16 +86,16 @@ fn run_encrypt(db_path: &str, dest: &str) -> Result<(), Error> {
}
}
fn run_verify(db_path: &str) -> Result<(), Error> {
async fn run_verify(db_path: &str) -> Result<(), Error> {
log!(CliLog::VerifyStarted(db_path.to_string()));
let db = match Database::new(db_path) {
let db = match Database::new(db_path).await {
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(0) {
match db.verify_audit_log_chain(0).await {
Ok((count, _last_id)) => {
log!(CliLog::VerifyOk(count));
Ok(())

View File

@ -10,7 +10,17 @@
//! host kernel release, and the NIC driver where those are obtainable,
//! so the operator can diagnose directly from the UI without shelling in.
use std::ffi::OsStr;
use std::fmt;
use std::fs;
use std::mem;
use std::os::fd::RawFd;
use std::path::{Path, PathBuf};
use libc::{
AF_NETLINK, NETLINK_GENERIC, SOCK_CLOEXEC, SOCK_RAW, bind, close, genlmsghdr, nlmsghdr, recv, sendto, sockaddr,
sockaddr_nl, socket,
};
use crate::domain::common::error::Error;
use crate::domain::common::system::health::{EbpfFailCategory, EbpfFailStage, EbpfHealth};
@ -24,19 +34,92 @@ pub fn kernel_release() -> String {
.unwrap_or_else(|| "unknown".to_string())
}
/// Look up the driver name bound to a network interface via
/// `/sys/class/net/<ifname>/device/driver`. Returns the basename of
/// the symlink target, or `"unknown"` if the interface has no driver
/// (e.g., virtual or renamed) or the path is not readable.
pub fn interface_driver(ifname: &str) -> String {
let link = format!("/sys/class/net/{}/device/driver", ifname);
match fs::read_link(&link) {
Ok(target) => target
.file_name()
.and_then(|s| s.to_str())
.map(|s| s.to_string())
.unwrap_or_else(|| "unknown".to_string()),
Err(_) => "unknown".to_string(),
const SYS_CLASS_NET: &str = "/sys/class/net";
const NETDEV_FAMILY_NAME: &str = "netdev";
const GENL_ID_CTRL: u16 = 0x10;
const CTRL_CMD_GETFAMILY: u8 = 3;
const CTRL_ATTR_FAMILY_ID: u16 = 1;
const CTRL_ATTR_FAMILY_NAME: u16 = 2;
const NETDEV_CMD_DEV_GET: u8 = 1;
const NETDEV_A_DEV_IFINDEX: u16 = 1;
const NETDEV_A_DEV_XDP_FEATURES: u16 = 3;
const NETDEV_A_DEV_XDP_ZC_MAX_SEGS: u16 = 4;
const NETDEV_A_DEV_XSK_FEATURES: u16 = 6;
const NETDEV_XDP_ACT_XSK_ZEROCOPY: u64 = 8;
const NLMSG_ERROR: u16 = 2;
const NLM_F_REQUEST: u16 = 1;
const NETLINK_SEQUENCE: u32 = 1;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct NetdevCapabilities {
pub xdp_features: Option<u64>,
pub xdp_zc_max_segs: Option<u32>,
pub xsk_features: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct InterfaceCapabilities {
pub driver: String,
pub mtu: Option<u32>,
pub rx_queues: Option<usize>,
pub tx_queues: Option<usize>,
pub netdev: Option<NetdevCapabilities>,
}
impl InterfaceCapabilities {
fn load(ifname: &str) -> Self {
Self::load_from(ifname, Path::new(SYS_CLASS_NET))
}
fn load_from(ifname: &str, sys_class_net: &Path) -> Self {
let iface_path = sys_class_net.join(ifname);
Self {
driver: interface_driver_from(&iface_path),
mtu: read_u32(iface_path.join("mtu")),
rx_queues: count_queue_dirs(&iface_path, "rx-"),
tx_queues: count_queue_dirs(&iface_path, "tx-"),
netdev: read_u32(iface_path.join("ifindex")).and_then(load_netdev_capabilities),
}
}
fn xsk_zerocopy_supported(&self) -> Option<bool> {
self.netdev
.as_ref()
.and_then(|netdev| netdev.xdp_features)
.map(|features| features & NETDEV_XDP_ACT_XSK_ZEROCOPY != 0)
}
}
impl fmt::Display for InterfaceCapabilities {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "driver {}", self.driver)?;
if let Some(mtu) = self.mtu {
write!(f, ", mtu {}", mtu)?;
}
if let Some(rx_queues) = self.rx_queues {
write!(f, ", rx_queues {}", rx_queues)?;
}
if let Some(tx_queues) = self.tx_queues {
write!(f, ", tx_queues {}", tx_queues)?;
}
match self.netdev.as_ref() {
Some(netdev) => {
if let Some(features) = netdev.xdp_features {
write!(f, ", xdp_features 0x{features:x}")?;
}
if let Some(max_segs) = netdev.xdp_zc_max_segs {
write!(f, ", xdp_zc_max_segs {}", max_segs)?;
}
if let Some(features) = netdev.xsk_features {
write!(f, ", xsk_features 0x{features:x}")?;
}
}
None => write!(f, ", netdev_features unavailable")?,
}
Ok(())
}
}
@ -48,24 +131,24 @@ pub fn classify(stage: EbpfFailStage, err: &Error, ifname: Option<&str>) -> Ebpf
let raw = err.to_string();
let category = categorize(&raw);
let kernel = kernel_release();
let driver = ifname.map(interface_driver);
let capabilities = ifname.map(InterfaceCapabilities::load);
let mut reason = format!("stage={:?}: {}", stage, raw);
reason.push_str(&format!(" (kernel {}", kernel));
if let Some(iface) = ifname {
reason.push_str(&format!(", interface {}", iface));
if let Some(drv) = driver.as_deref() {
reason.push_str(&format!(", driver {}", drv));
if let Some(caps) = capabilities.as_ref() {
reason.push_str(&format!(", {}", caps));
}
}
reason.push(')');
// For the igb-before-6.17 case, augment the reason with a targeted hint.
if matches!(category, EbpfFailCategory::AfXdpUnsupported)
&& driver.as_deref() == Some("igb")
&& !kernel_meets_igb_af_xdp(&kernel)
{
reason.push_str(". The igb driver supports AF_XDP only on kernel 6.17 or newer.");
if matches!(category, EbpfFailCategory::AfXdpUnsupported) {
if let Some(Some(false)) = capabilities.as_ref().map(InterfaceCapabilities::xsk_zerocopy_supported) {
reason.push_str(". The interface reports no AF_XDP zero-copy capability in xdp_features.");
} else if capabilities.as_ref().and_then(|caps| caps.netdev.as_ref()).is_none() {
reason.push_str(". Kernel did not expose netdev XDP feature data for this interface.");
}
}
EbpfHealth::Unavailable {
@ -118,41 +201,263 @@ fn categorize(raw: &str) -> EbpfFailCategory {
EbpfFailCategory::Unknown
}
/// Parse a kernel release string like "6.17.4-generic" and return true
/// if it is >= 6.17. We only care about the first two numeric components.
fn kernel_meets_igb_af_xdp(release: &str) -> bool {
// Pull leading "MAJOR.MINOR" out of strings like "6.12.0-124.45.1.el10_1.x86_64".
let mut parts = release.split(|c: char| !c.is_ascii_digit()).filter(|s| !s.is_empty());
let Some(major_str) = parts.next() else {
return false;
fn interface_driver_from(iface_path: &Path) -> String {
match fs::read_link(iface_path.join("device").join("driver")) {
Ok(target) => basename_to_string(&target),
Err(_) => "unknown".to_string(),
}
}
fn basename_to_string(path: &Path) -> String {
path.file_name()
.and_then(OsStr::to_str)
.map(str::to_string)
.unwrap_or_else(|| "unknown".to_string())
}
fn read_u32(path: PathBuf) -> Option<u32> {
fs::read_to_string(path).ok()?.trim().parse().ok()
}
fn count_queue_dirs(iface_path: &Path, prefix: &str) -> Option<usize> {
let queues_path = iface_path.join("queues");
let entries = fs::read_dir(queues_path).ok()?;
let count = entries
.filter_map(Result::ok)
.filter(|entry| entry.file_type().map(|ty| ty.is_dir()).unwrap_or(false))
.filter(|entry| {
entry
.file_name()
.to_str()
.map(|name| name.starts_with(prefix))
.unwrap_or(false)
})
.count();
Some(count)
}
fn load_netdev_capabilities(ifindex: u32) -> Option<NetdevCapabilities> {
let family_id = resolve_genl_family_id(NETDEV_FAMILY_NAME)?;
let payload = genl_request(
family_id,
NETDEV_CMD_DEV_GET,
&[netlink_attr_u32(NETDEV_A_DEV_IFINDEX, ifindex)],
);
let response = netlink_round_trip(&payload)?;
let attrs = first_genl_attrs(&response)?;
Some(NetdevCapabilities {
xdp_features: attr_u64(attrs, NETDEV_A_DEV_XDP_FEATURES),
xdp_zc_max_segs: attr_u32(attrs, NETDEV_A_DEV_XDP_ZC_MAX_SEGS),
xsk_features: attr_u64(attrs, NETDEV_A_DEV_XSK_FEATURES),
})
}
fn resolve_genl_family_id(name: &str) -> Option<u16> {
let payload = genl_request(
GENL_ID_CTRL,
CTRL_CMD_GETFAMILY,
&[netlink_attr_string(CTRL_ATTR_FAMILY_NAME, name)],
);
let response = netlink_round_trip(&payload)?;
let attrs = first_genl_attrs(&response)?;
attr_u16(attrs, CTRL_ATTR_FAMILY_ID)
}
fn netlink_round_trip(payload: &[u8]) -> Option<Vec<u8>> {
let fd = open_netlink_socket()?;
let sent = send_netlink(fd, payload).is_some();
let response = if sent { recv_netlink(fd) } else { None };
close_fd(fd);
response
}
fn open_netlink_socket() -> Option<RawFd> {
// SAFETY: socket and bind are called with a valid sockaddr_nl and length.
let fd = unsafe { socket(AF_NETLINK, SOCK_RAW | SOCK_CLOEXEC, NETLINK_GENERIC) };
if fd < 0 {
return None;
}
// SAFETY: zeroed sockaddr_nl is immediately initialized before bind.
let mut addr: sockaddr_nl = unsafe { mem::zeroed() };
addr.nl_family = AF_NETLINK as u16;
// SAFETY: fd is a netlink socket, addr points to initialized storage.
let rc = unsafe {
bind(
fd,
&addr as *const sockaddr_nl as *const sockaddr,
mem::size_of::<sockaddr_nl>() as u32,
)
};
let Some(minor_str) = parts.next() else {
return false;
if rc < 0 {
close_fd(fd);
return None;
}
Some(fd)
}
fn close_fd(fd: RawFd) {
// SAFETY: closing an owned file descriptor; failures are irrelevant here.
unsafe {
close(fd);
}
}
fn send_netlink(fd: RawFd, payload: &[u8]) -> Option<()> {
// SAFETY: zeroed sockaddr_nl is immediately initialized before sendto.
let mut kernel: sockaddr_nl = unsafe { mem::zeroed() };
kernel.nl_family = AF_NETLINK as u16;
// SAFETY: payload is a valid byte slice and kernel points to initialized storage.
let sent = unsafe {
sendto(
fd,
payload.as_ptr().cast(),
payload.len(),
0,
&kernel as *const sockaddr_nl as *const sockaddr,
mem::size_of::<sockaddr_nl>() as u32,
)
};
let Ok(major): Result<u32, _> = major_str.parse() else {
return false;
};
let Ok(minor): Result<u32, _> = minor_str.parse() else {
return false;
};
major > 6 || (major == 6 && minor >= 17)
if sent == payload.len() as isize { Some(()) } else { None }
}
fn recv_netlink(fd: RawFd) -> Option<Vec<u8>> {
let mut buf = vec![0_u8; 8192];
// SAFETY: buf is valid writable storage for recv.
let received = unsafe { recv(fd, buf.as_mut_ptr().cast(), buf.len(), 0) };
if received <= 0 {
return None;
}
buf.truncate(received as usize);
Some(buf)
}
fn genl_request(nlmsg_type: u16, cmd: u8, attrs: &[Vec<u8>]) -> Vec<u8> {
let header_len = align4(mem::size_of::<nlmsghdr>());
let genl_len = mem::size_of::<genlmsghdr>();
let mut buf = vec![0_u8; header_len + genl_len];
write_u32(&mut buf, 0, 0);
write_u16(&mut buf, 4, nlmsg_type);
write_u16(&mut buf, 6, NLM_F_REQUEST);
write_u32(&mut buf, 8, NETLINK_SEQUENCE);
write_u32(&mut buf, 12, 0);
buf[header_len] = cmd;
buf[header_len + 1] = 1;
for attr in attrs {
buf.extend_from_slice(attr);
}
let nlmsg_len = buf.len() as u32;
write_u32(&mut buf, 0, nlmsg_len);
buf
}
fn netlink_attr_string(attr_type: u16, value: &str) -> Vec<u8> {
let mut payload = value.as_bytes().to_vec();
payload.push(0);
netlink_attr(attr_type, &payload)
}
fn netlink_attr_u32(attr_type: u16, value: u32) -> Vec<u8> {
netlink_attr(attr_type, &value.to_ne_bytes())
}
fn netlink_attr(attr_type: u16, payload: &[u8]) -> Vec<u8> {
let len = 4 + payload.len();
let mut attr = vec![0_u8; align4(len)];
write_u16(&mut attr, 0, len as u16);
write_u16(&mut attr, 2, attr_type);
attr[4..4 + payload.len()].copy_from_slice(payload);
attr
}
fn first_genl_attrs(response: &[u8]) -> Option<&[u8]> {
let header_len = align4(mem::size_of::<nlmsghdr>());
let genl_len = mem::size_of::<genlmsghdr>();
if response.len() < header_len + genl_len {
return None;
}
let nlmsg_len = read_u32_ne(response, 0)? as usize;
let nlmsg_type = read_u16_ne(response, 4)?;
if nlmsg_type == NLMSG_ERROR || nlmsg_len > response.len() || nlmsg_len < header_len + genl_len {
return None;
}
Some(&response[header_len + genl_len..nlmsg_len])
}
fn attr_u16(attrs: &[u8], attr_type: u16) -> Option<u16> {
find_attr(attrs, attr_type).and_then(|payload| read_u16_ne(payload, 0))
}
fn attr_u32(attrs: &[u8], attr_type: u16) -> Option<u32> {
find_attr(attrs, attr_type).and_then(|payload| read_u32_ne(payload, 0))
}
fn attr_u64(attrs: &[u8], attr_type: u16) -> Option<u64> {
find_attr(attrs, attr_type).and_then(|payload| read_u64_ne(payload, 0))
}
fn find_attr(attrs: &[u8], attr_type: u16) -> Option<&[u8]> {
let mut offset = 0;
while offset + 4 <= attrs.len() {
let len = read_u16_ne(attrs, offset)? as usize;
let current_type = read_u16_ne(attrs, offset + 2)?;
if len < 4 || offset + len > attrs.len() {
return None;
}
if current_type == attr_type {
return Some(&attrs[offset + 4..offset + len]);
}
offset += align4(len);
}
None
}
fn align4(value: usize) -> usize {
(value + 3) & !3
}
fn write_u16(buf: &mut [u8], offset: usize, value: u16) {
buf[offset..offset + 2].copy_from_slice(&value.to_ne_bytes());
}
fn write_u32(buf: &mut [u8], offset: usize, value: u32) {
buf[offset..offset + 4].copy_from_slice(&value.to_ne_bytes());
}
fn read_u16_ne(buf: &[u8], offset: usize) -> Option<u16> {
let bytes = buf.get(offset..offset + 2)?;
Some(u16::from_ne_bytes([bytes[0], bytes[1]]))
}
fn read_u32_ne(buf: &[u8], offset: usize) -> Option<u32> {
let bytes = buf.get(offset..offset + 4)?;
Some(u32::from_ne_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]))
}
fn read_u64_ne(buf: &[u8], offset: usize) -> Option<u64> {
let bytes = buf.get(offset..offset + 8)?;
Some(u64::from_ne_bytes([
bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7],
]))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn kernel_version_matrix() {
assert!(kernel_meets_igb_af_xdp("6.17.0-generic"));
assert!(kernel_meets_igb_af_xdp("6.18.1"));
assert!(kernel_meets_igb_af_xdp("7.0.0"));
assert!(!kernel_meets_igb_af_xdp("6.16.9-generic"));
assert!(!kernel_meets_igb_af_xdp("6.12.0-124.45.1.el10_1.x86_64"));
assert!(!kernel_meets_igb_af_xdp("5.15.0"));
assert!(!kernel_meets_igb_af_xdp("nonsense"));
}
#[test]
fn categorizes_permission_errors() {
assert!(matches!(
@ -176,4 +481,63 @@ mod tests {
EbpfFailCategory::XdpUnsupported
));
}
#[test]
fn loads_interface_capabilities_from_sysfs_shape() {
let root = std::env::temp_dir().join(format!("netguardia-ebpf-preflight-{}", std::process::id()));
let iface = root.join("eth0");
let queues = iface.join("queues");
let device = iface.join("device");
let driver_target = root.join("drivers").join("virtio_net");
let _ = fs::remove_dir_all(&root);
fs::create_dir_all(queues.join("rx-0")).unwrap();
fs::create_dir_all(queues.join("rx-1")).unwrap();
fs::create_dir_all(queues.join("tx-0")).unwrap();
fs::create_dir_all(&driver_target).unwrap();
fs::create_dir_all(&device).unwrap();
fs::write(iface.join("mtu"), "1500\n").unwrap();
fs::write(iface.join("ifindex"), "not-a-number\n").unwrap();
std::os::unix::fs::symlink(&driver_target, device.join("driver")).unwrap();
let caps = InterfaceCapabilities::load_from("eth0", &root);
assert_eq!(caps.driver, "virtio_net");
assert_eq!(caps.mtu, Some(1500));
assert_eq!(caps.rx_queues, Some(2));
assert_eq!(caps.tx_queues, Some(1));
assert_eq!(caps.netdev, None);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn parses_netlink_attributes() {
let attrs = [
netlink_attr_u32(NETDEV_A_DEV_IFINDEX, 7),
netlink_attr(NETDEV_A_DEV_XDP_FEATURES, &8_u64.to_ne_bytes()),
]
.concat();
assert_eq!(attr_u32(&attrs, NETDEV_A_DEV_IFINDEX), Some(7));
assert_eq!(attr_u64(&attrs, NETDEV_A_DEV_XDP_FEATURES), Some(8));
assert_eq!(attr_u32(&attrs, NETDEV_A_DEV_XSK_FEATURES), None);
}
#[test]
fn interprets_xsk_zerocopy_from_netdev_features() {
let caps = InterfaceCapabilities {
driver: "virtio_net".to_string(),
mtu: Some(1500),
rx_queues: Some(1),
tx_queues: Some(1),
netdev: Some(NetdevCapabilities {
xdp_features: Some(NETDEV_XDP_ACT_XSK_ZEROCOPY),
xdp_zc_max_segs: Some(1),
xsk_features: Some(0),
}),
};
assert_eq!(caps.xsk_zerocopy_supported(), Some(true));
}
}

View File

@ -1,5 +1,5 @@
use std::net::IpAddr;
use std::path::{Path, PathBuf};
use std::path::Path;
use std::sync::Arc;
use async_trait::async_trait;
@ -17,14 +17,9 @@ pub struct GeoIpService {
}
impl GeoIpService {
pub fn new(db_name: &str) -> Result<Self, MaxMindDbError> {
let db_path = PathBuf::from("net-guardia/static/geo").join(db_name);
Self::with_cache_size(db_path, 10000)
}
pub fn with_cache_size<P: AsRef<Path>>(db_path: P, cache_size: usize) -> Result<Self, MaxMindDbError> {
let reader = Reader::open_readfile(db_path)?;
let capacity = if cache_size == 0 { 10_000 } else { cache_size } as u64;
let capacity = cache_size.max(1) as u64;
Ok(Self {
reader: Arc::new(reader),

View File

@ -36,8 +36,8 @@ pub struct SystemHealth {
impl SystemHealth {
pub fn new(config: Arc<ArcSwap<AppConfig>>, ebpf_health: Arc<ArcSwap<EbpfHealth>>) -> Result<Self, Error> {
let (broadcast_tx, _) = broadcast::channel(100);
let cfg = config.load();
let (broadcast_tx, _) = broadcast::channel(cfg.health.broadcast_channel_capacity.max(1));
let ingress_interface = cfg.ebpf.ingress_ifname.clone();
let egress_interface = cfg.ebpf.egress_ifname.clone();

View File

@ -48,6 +48,7 @@ use crate::infrastructure::inference_runtime::InferenceRuntime;
use crate::infrastructure::log_buffer::LogBuffer;
use crate::infrastructure::logger::Logger;
use crate::infrastructure::readiness::ReadinessState;
use crate::infrastructure::runtime_state::RuntimeState;
use crate::infrastructure::secret_store::SecretStore;
use crate::infrastructure::suricata_manager::SuricataManager;
use crate::infrastructure::system::ShutdownHandle;
@ -68,6 +69,7 @@ pub struct ForceHttpsFlag(pub Arc<AtomicBool>);
pub struct HttpServerParams {
pub app_config: Arc<ArcSwap<AppConfig>>,
pub runtime_state: Arc<ArcSwap<RuntimeState>>,
pub inference_config: Arc<MLInferenceConfig>,
pub database: Arc<Database>,
@ -223,6 +225,7 @@ pub fn start_setup_server(params: SetupServerParams) -> Result<ServerHandle, Err
pub async fn run(params: HttpServerParams) -> Result<(), Error> {
let app_config = params.app_config;
let runtime_state = params.runtime_state;
let inference_config = params.inference_config;
let database = params.database;
@ -274,6 +277,7 @@ pub async fn run(params: HttpServerParams) -> Result<(), Error> {
.wrap(HttpsRedirect)
.wrap(cors(app_config.load().http_server.cors_allowed_origins.clone()))
.app_data(web::Data::from(app_config.clone()))
.app_data(web::Data::from(runtime_state.clone()))
.app_data(web::Data::from(inference_config.clone()))
.app_data(web::Data::from(database.clone() as Arc<dyn AppRepo>))
.app_data(web::Data::from(database.clone() as Arc<dyn ApiKeyRepo>))

View File

@ -11,7 +11,7 @@ use tracing_subscriber::layer::SubscriberExt;
use tracing_subscriber::util::SubscriberInitExt;
use tracing_subscriber::{Layer, filter, reload};
use crate::domain::common::config::observability::ObservabilityConfig;
use crate::domain::common::config::AppConfig;
use crate::domain::common::error::Error;
use crate::domain::common::error::io::IOError;
use crate::infrastructure::log_buffer::{LogBuffer, LogBufferLayer};
@ -26,9 +26,10 @@ pub struct Logger {
}
impl Logger {
pub fn initialize(config: &ObservabilityConfig) -> Result<(Self, LogBuffer), Error> {
let log_directory = "logs";
fs::create_dir_all(log_directory).map_err(|err| IOError::CreateDirectoryFailed(log_directory, err))?;
pub fn initialize(config: &AppConfig) -> Result<(Self, LogBuffer), Error> {
let observability = &config.observability;
let log_directory = &config.system.log_dir;
fs::create_dir_all(log_directory).map_err(|err| IOError::CreateDirectoryFailed(log_directory.clone(), err))?;
let file_appender = RollingFileAppender::new(Rotation::DAILY, log_directory, "NetGuardia");
@ -47,7 +48,7 @@ impl Logger {
.with_ansi(false)
.with_writer(file_appender);
let level: Level = config
let level: Level = observability
.log_level
.parse()
.ok()
@ -71,8 +72,10 @@ impl Logger {
filter = Self::apply_directives(filter, &preserved);
let (filter_layer, reload_handle) = reload::Layer::new(filter);
let (log_buffer_layer, log_buffer) =
LogBufferLayer::new(config.log_buffer_capacity, config.log_buffer_max_message_bytes);
let (log_buffer_layer, log_buffer) = LogBufferLayer::new(
observability.log_buffer_capacity,
observability.log_buffer_max_message_bytes,
);
tracing_subscriber::registry()
.with(filter_layer)

View File

@ -8,6 +8,7 @@ pub mod inference_runtime;
pub mod log_buffer;
pub mod logger;
pub mod readiness;
pub mod runtime_state;
pub mod secret_store;
pub mod service_factory;
pub mod suricata_manager;

View File

@ -0,0 +1,21 @@
#[derive(Debug, Clone)]
pub struct RuntimeState {
pub xdp: XdpRuntimeState,
}
#[derive(Debug, Clone)]
pub struct XdpRuntimeState {
pub ingress_mode: String,
pub egress_mode: String,
}
impl Default for RuntimeState {
fn default() -> Self {
Self {
xdp: XdpRuntimeState {
ingress_mode: "unknown".to_string(),
egress_mode: "unknown".to_string(),
},
}
}
}

View File

@ -13,8 +13,8 @@ use crate::adapter::persistence::Database;
use crate::domain::common::error::Error;
use crate::domain::common::error::crypto::{CryptoError, EnvelopeField};
use crate::domain::common::log::crypto::CryptoLog;
use crate::interface::config_repo::ConfigRepo;
use crate::interface::secret_store::SecretStorePort;
use crate::interface::setting::SettingRepo;
pub struct SecretStore {
database: Arc<Database>,
@ -119,19 +119,20 @@ impl SecretStore {
}
}
#[async_trait::async_trait]
impl SecretStorePort for SecretStore {
fn get_secret(&self, key: &str) -> Result<Option<String>, Error> {
let repo: &dyn SettingRepo = self.database.as_ref();
match repo.get_app_secret(key)? {
async fn get_secret(&self, key: &str) -> Result<Option<String>, Error> {
let repo: &dyn ConfigRepo = self.database.as_ref();
match repo.get_app_secret(key).await? {
Some(envelope_json) => Ok(Some(self.decrypt(&envelope_json)?)),
None => Ok(None),
}
}
fn set_secret(&self, key: &str, plaintext: &str) -> Result<(), Error> {
async fn set_secret(&self, key: &str, plaintext: &str) -> Result<(), Error> {
let envelope = self.encrypt(plaintext)?;
let repo: &dyn SettingRepo = self.database.as_ref();
repo.set_app_secret(key, &envelope)
let repo: &dyn ConfigRepo = self.database.as_ref();
repo.set_app_secret(key, &envelope).await
}
fn encrypt_envelope(&self, plaintext: &str) -> Result<String, Error> {
@ -143,36 +144,36 @@ impl SecretStorePort for SecretStore {
mod tests {
use super::*;
fn store_with_cipher() -> SecretStore {
async fn store_with_cipher() -> SecretStore {
let hk = Hkdf::<Sha256>::new(Some(b"netguardia-v1-salt"), b"test-key-for-unit-tests");
let mut okm = [0u8; 32];
hk.expand(b"netguardia-envelope-v1", &mut okm).unwrap();
let cipher = Aes256Gcm::new_from_slice(&okm).unwrap();
SecretStore {
database: Arc::new(Database::new(":memory:").unwrap()),
database: Arc::new(Database::new(":memory:").await.unwrap()),
cipher: Some(cipher),
}
}
fn store_without_cipher() -> SecretStore {
async fn store_without_cipher() -> SecretStore {
SecretStore {
database: Arc::new(Database::new(":memory:").unwrap()),
database: Arc::new(Database::new(":memory:").await.unwrap()),
cipher: None,
}
}
#[test]
fn encrypt_decrypt_round_trip() {
let store = store_with_cipher();
#[tokio::test]
async fn encrypt_decrypt_round_trip() {
let store = store_with_cipher().await;
let original = "my-smtp-password-123!@#";
let encrypted = store.encrypt(original).unwrap();
let decrypted = store.decrypt(&encrypted).unwrap();
assert_eq!(decrypted, original);
}
#[test]
fn encrypt_produces_different_ciphertext_each_time() {
let store = store_with_cipher();
#[tokio::test]
async fn encrypt_produces_different_ciphertext_each_time() {
let store = store_with_cipher().await;
let plaintext = "same-value";
let e1 = store.encrypt(plaintext).unwrap();
let e2 = store.encrypt(plaintext).unwrap();
@ -181,9 +182,9 @@ mod tests {
assert_eq!(store.decrypt(&e2).unwrap(), plaintext);
}
#[test]
fn wrong_key_fails_decrypt() {
let store1 = store_with_cipher();
#[tokio::test]
async fn wrong_key_fails_decrypt() {
let store1 = store_with_cipher().await;
let encrypted = store1.encrypt("secret-value").unwrap();
let hk = Hkdf::<Sha256>::new(Some(b"netguardia-v1-salt"), b"different-key");
@ -191,7 +192,7 @@ mod tests {
hk.expand(b"netguardia-envelope-v1", &mut okm).unwrap();
let cipher = Aes256Gcm::new_from_slice(&okm).unwrap();
let store2 = SecretStore {
database: Arc::new(Database::new(":memory:").unwrap()),
database: Arc::new(Database::new(":memory:").await.unwrap()),
cipher: Some(cipher),
};
@ -199,9 +200,9 @@ mod tests {
assert!(result.is_err(), "Decrypting with wrong key should fail");
}
#[test]
fn alg_none_rejected_in_production_mode() {
let store = store_with_cipher();
#[tokio::test]
async fn alg_none_rejected_in_production_mode() {
let store = store_with_cipher().await;
let fake_envelope = serde_json::json!({
"v": 1,
"alg": "none",
@ -213,30 +214,30 @@ mod tests {
assert!(result.is_err(), "alg:none should be rejected when cipher is present");
}
#[test]
fn alg_none_allowed_in_dev_mode() {
let store = store_without_cipher();
#[tokio::test]
async fn alg_none_allowed_in_dev_mode() {
let store = store_without_cipher().await;
let encrypted = store.encrypt("dev-mode-secret").unwrap();
let decrypted = store.decrypt(&encrypted).unwrap();
assert_eq!(decrypted, "dev-mode-secret");
}
#[test]
fn invalid_envelope_json_fails() {
let store = store_with_cipher();
#[tokio::test]
async fn invalid_envelope_json_fails() {
let store = store_with_cipher().await;
assert!(store.decrypt("not-json").is_err());
}
#[test]
fn unsupported_version_fails() {
let store = store_with_cipher();
#[tokio::test]
async fn unsupported_version_fails() {
let store = store_with_cipher().await;
let envelope = serde_json::json!({"v": 99, "alg": "aes-256-gcm", "ct": "abc"}).to_string();
assert!(store.decrypt(&envelope).is_err());
}
#[test]
fn unsupported_algorithm_fails() {
let store = store_with_cipher();
#[tokio::test]
async fn unsupported_algorithm_fails() {
let store = store_with_cipher().await;
let envelope = serde_json::json!({"v": 1, "alg": "chacha20", "ct": "abc"}).to_string();
assert!(store.decrypt(&envelope).is_err());
}

View File

@ -35,7 +35,7 @@ use crate::core::response::engine::{SoarEngine, SoarEngineDeps};
use crate::core::response::playbook_service::PlaybookService;
use crate::core::response::scheduler::TtlScheduler;
use crate::domain::common::config::AppConfig;
use crate::domain::common::config::constants::{ENFORCE_MODE_MONITOR, EVENT_CHANNEL_CAPACITY, enforce_mode_to_u8};
use crate::domain::common::config::constants::{EVENT_CHANNEL_CAPACITY, enforce_mode_to_u8};
use crate::domain::common::error::Error;
use crate::domain::common::error::misc::MiscError;
use crate::domain::common::event::{AuditEvent, DriftDetectedEvent, ThreatDetectedEvent};
@ -44,6 +44,7 @@ use crate::domain::common::system::health::EbpfFailStage;
use crate::domain::common::system::health::EbpfHealth;
use crate::domain::data_plane::direction::FlowDirection;
use crate::domain::data_plane::error::EbpfError;
use crate::domain::data_plane::flow_stats::FlowStatsLimits;
use crate::domain::data_plane::list_type::ListType;
use crate::domain::data_plane::log::EbpfLog;
use crate::domain::detection::drift::FeatureBaselines;
@ -54,24 +55,28 @@ use crate::infrastructure::ebpf_preflight;
use crate::infrastructure::geoip::GeoIpService;
use crate::infrastructure::health::SystemHealth;
use crate::infrastructure::inference_runtime::InferenceRuntime;
use crate::infrastructure::runtime_state::RuntimeState;
use crate::infrastructure::secret_store::SecretStore;
use crate::infrastructure::suricata_manager::SuricataManager;
use crate::interface::access_control::AccessControlPort;
use crate::interface::access_control_admin::AccessControlAdminPort;
use crate::interface::app_repo::AppRepo;
use crate::interface::config_repo::ConfigRepo;
use crate::interface::dns_filter_api::DnsFilterPort;
use crate::interface::dns_query_filter::DnsQueryFilter;
use crate::interface::email_sender::EmailSenderFactory;
use crate::interface::enforcement::EnforcementRepo;
use crate::interface::geo_block_api::GeoBlockPort;
use crate::interface::geo_lookup::GeoLookup;
use crate::interface::notification::{AlertNotifier, AlertNotifierFactory};
use crate::interface::rate_limit_api::RateLimitPort;
use crate::interface::report_snapshot::ReportSnapshotRepo;
use crate::interface::secret_store::SecretStorePort;
use crate::interface::setting::SettingRepo;
use crate::interface::soar::SoarRepo;
pub struct AppState {
pub app_config: Arc<ArcSwap<AppConfig>>,
pub runtime_state: Arc<ArcSwap<RuntimeState>>,
pub inference_config: Arc<MLInferenceConfig>,
pub database: Arc<Database>,
@ -132,8 +137,9 @@ impl ServiceFactory {
/// Only called when setup is complete — all config values are in DB.
pub async fn build(db: Arc<Database>) -> Result<AppState, Error> {
// Ensure DB has all default config keys (INSERT OR IGNORE — never overwrites)
AppConfig::seed_defaults(db.as_ref())?;
let app_config = Arc::new(ArcSwap::from_pointee(AppConfig::from_settings(db.as_ref())?));
AppConfig::seed_config_defaults(db.as_ref()).await?;
let app_config = Arc::new(ArcSwap::from_pointee(AppConfig::from_config_repo(db.as_ref()).await?));
let runtime_state = Arc::new(ArcSwap::from_pointee(RuntimeState::default()));
// Prefer `models/manifest.yaml` when present (v12 BYO-model path). The manifest
// is the user-authored source of truth for features, labels, thresholds, and
@ -191,11 +197,6 @@ impl ServiceFactory {
}
};
// Ensure enforce_mode setting exists (default: monitor)
if db.get_setting("enforce_mode")?.is_none() {
db.set_setting("enforce_mode", ENFORCE_MODE_MONITOR)?;
}
let secret_store = Arc::new(SecretStore::new(db.clone()));
let secret_store_port: Arc<dyn SecretStorePort> = secret_store.clone();
@ -222,8 +223,8 @@ impl ServiceFactory {
// Create AtomicU8 enforce-level cache (Monitor=0, MlOnly=1, Enforce=2)
let enforce_level_cache = Arc::new(AtomicU8::new({
let mode_str = db.get_setting("enforce_mode")?.unwrap_or_default();
enforce_mode_to_u8(&mode_str)
let mode = app_config.load().system.enforce_mode.clone();
enforce_mode_to_u8(&mode)
}));
// Named broadcast channels replace the old TypeId-based CommunicationManager.
@ -244,26 +245,36 @@ impl ServiceFactory {
audit_tx.clone(),
)?);
let flow_statistics = Arc::new(FlowStatistics::new(inference_runtime.ml_engine.clone()));
let flow_stats_cfg = app_config.load();
let flow_stats_limits = FlowStatsLimits::new(
flow_stats_cfg.detection.flow_stats.max_snapshot_entries,
flow_stats_cfg.detection.flow_stats.max_top_n,
);
drop(flow_stats_cfg);
let flow_statistics = Arc::new(FlowStatistics::new(
inference_runtime.ml_engine.clone(),
flow_stats_limits,
));
let enforce_handler = Arc::new(EnforceModeHandler::new(
db.clone() as Arc<dyn AppRepo>,
app_config.clone(),
audit_tx.clone(),
enforce_level_cache.clone(),
));
// Seed default SOAR playbooks if empty
(db.as_ref() as &dyn SoarRepo).seed_default_playbooks()?;
(db.as_ref() as &dyn SoarRepo).seed_default_playbooks().await?;
// Restore persisted state from database
Self::restore_dns_blacklist(&db, dns_filter_port.as_ref());
Self::restore_geo_countries(&db, &ebpf_services);
Self::restore_rate_limits(&db, &ebpf_services);
Self::restore_dns_blacklist(&db, dns_filter_port.as_ref()).await;
Self::restore_geo_countries(&db, &ebpf_services).await;
Self::restore_rate_limits(&db, &ebpf_services).await;
Self::restore_acl_rules(&db, &ebpf_services).await;
// Create TelegramAdapter as alert notifier (may fail if not configured yet)
let alert_notifier: Option<Arc<dyn AlertNotifier>> = match TelegramAdapter::new(
db.clone() as Arc<dyn SettingRepo + Send + Sync>,
db.clone() as Arc<dyn ConfigRepo + Send + Sync>,
app_config.clone(),
Some(secret_store_port.clone()),
) {
@ -275,16 +286,18 @@ impl ServiceFactory {
};
// Try to initialize GeoIP service
let geoip: Option<Arc<dyn GeoLookup>> = match GeoIpService::new(&app_config.load().acl.geoip_db_name) {
Ok(svc) => {
log!(SystemLog::GeoIpInitialized);
Some(Arc::new(svc))
}
Err(e) => {
log!(SystemLog::GeoIpUnavailable(e.to_string()));
None
}
};
let acl_cfg = app_config.load().acl.clone();
let geoip: Option<Arc<dyn GeoLookup>> =
match GeoIpService::with_cache_size(&acl_cfg.geoip_db_path, acl_cfg.geoip_cache_capacity) {
Ok(svc) => {
log!(SystemLog::GeoIpInitialized);
Some(Arc::new(svc))
}
Err(e) => {
log!(SystemLog::GeoIpUnavailable(e.to_string()));
None
}
};
// Create AccessControlPort adapter for SOAR/TTL (decoupled from eBPF)
let access_control_port: Arc<dyn AccessControlPort> =
@ -294,24 +307,27 @@ impl ServiceFactory {
// Create SOAR engine
let rate_limit_port: Arc<dyn RateLimitPort> = ebpf_services.rate_limit.clone();
let soar_engine = Arc::new(SoarEngine::new(SoarEngineDeps {
db: db.clone(),
config: app_config.clone(),
access_control: access_control_port.clone(),
alert_notifier: alert_notifier.clone(),
geoip: geoip.clone(),
rate_limit: Some(rate_limit_port.clone()),
enforce_level_cache,
secrets: Some(secret_store_port.clone()),
email_sender_factory: email_sender_factory.clone(),
})?);
let soar_engine = Arc::new(
SoarEngine::new(SoarEngineDeps {
db: db.clone(),
config: app_config.clone(),
access_control: access_control_port.clone(),
alert_notifier: alert_notifier.clone(),
geoip: geoip.clone(),
rate_limit: Some(rate_limit_port.clone()),
enforce_level_cache,
secrets: Some(secret_store_port.clone()),
email_sender_factory: email_sender_factory.clone(),
})
.await?,
);
// Create TTL scheduler
let ttl_scheduler = TtlScheduler::new(db.clone(), access_control_port.clone(), soar_engine.clone());
// Create Report scheduler
let report_scheduler = ReportScheduler::new(
db.clone() as Arc<dyn SettingRepo + Send + Sync>,
db.clone() as Arc<dyn ReportSnapshotRepo>,
app_config.clone(),
Some(secret_store_port.clone()),
email_sender_factory.clone(),
@ -327,7 +343,7 @@ impl ServiceFactory {
geo_block_port,
));
let dns_filter_service = Arc::new(DnsFilterService::new(
db.clone() as Arc<dyn AppRepo>,
db.clone() as Arc<dyn EnforcementRepo>,
dns_filter_port,
app_config.clone(),
));
@ -342,12 +358,12 @@ impl ServiceFactory {
.with_secret_store(secret_store_port.clone()),
);
let notifier_factory: Arc<dyn AlertNotifierFactory> = Arc::new(TelegramAdapterFactory::new(
db.clone() as Arc<dyn SettingRepo + Send + Sync>,
db.clone() as Arc<dyn ConfigRepo + Send + Sync>,
app_config.clone(),
Some(secret_store_port.clone()),
));
let notification_service = Arc::new(NotificationService::new(
db.clone() as Arc<dyn SettingRepo + Send + Sync>,
db.clone() as Arc<dyn ConfigRepo + Send + Sync>,
app_config.clone(),
secret_store_port,
notifier_factory,
@ -358,6 +374,7 @@ impl ServiceFactory {
Ok(AppState {
app_config,
runtime_state,
inference_config,
database: db,
@ -567,8 +584,8 @@ impl ServiceFactory {
// --- State restoration helpers ---
fn restore_dns_blacklist(db: &Database, dns_filter_port: &dyn DnsFilterPort) {
if let Ok(domains) = db.load_dns_domains() {
async fn restore_dns_blacklist(db: &Database, dns_filter_port: &dyn DnsFilterPort) {
if let Ok(domains) = db.load_dns_domains().await {
for domain in &domains {
if let Err(e) = dns_filter_port.add_domain(domain) {
log!(SystemLog::DnsRestoreFailed(domain.clone(), e.to_string()));
@ -580,8 +597,8 @@ impl ServiceFactory {
}
}
fn restore_geo_countries(db: &Database, ebpf_services: &EbpfServices) {
if let Ok(countries) = db.load_geo_countries()
async fn restore_geo_countries(db: &Database, ebpf_services: &EbpfServices) {
if let Ok(countries) = db.load_geo_countries().await
&& !countries.is_empty()
{
if let Err(e) = ebpf_services.geo_block.block_countries(&countries) {
@ -592,8 +609,8 @@ impl ServiceFactory {
}
}
fn restore_rate_limits(db: &Database, ebpf_services: &EbpfServices) {
if let Ok(configs) = db.load_rate_limit_config() {
async fn restore_rate_limits(db: &Database, ebpf_services: &EbpfServices) {
if let Ok(configs) = db.load_rate_limit_config().await {
for (key, value) in &configs {
let result = match key.as_str() {
"packet_rate" => ebpf_services.rate_limit.set_packet_rate(*value),
@ -614,7 +631,7 @@ impl ServiceFactory {
}
async fn restore_acl_rules(db: &Database, ebpf_services: &EbpfServices) {
if let Ok(rules) = db.list_acl_rules() {
if let Ok(rules) = db.list_acl_rules().await {
let mut restored = 0u32;
for rule in &rules {
let dir = match rule.direction.as_str() {

View File

@ -55,6 +55,7 @@ use crate::infrastructure::inference_runtime::InferenceRuntime;
use crate::infrastructure::log_buffer::LogBuffer;
use crate::infrastructure::logger::Logger;
use crate::infrastructure::readiness::ReadinessState;
use crate::infrastructure::runtime_state::RuntimeState;
use crate::infrastructure::secret_store::SecretStore;
use crate::infrastructure::service_factory::ServiceFactory;
use crate::infrastructure::suricata_manager::SuricataManager;
@ -62,7 +63,7 @@ use crate::infrastructure::suricata_monitor::SuricataMonitor;
use crate::interface::audit::AuditRepo;
use crate::interface::geo_lookup::GeoLookup;
use crate::interface::packet_sink::PacketSinkFactory;
use crate::interface::setting::SettingRepo;
use crate::interface::report_snapshot::ReportSnapshotRepo;
use crate::interface::stats::StatsRepo;
use crate::utils::staging;
@ -88,6 +89,7 @@ impl ShutdownHandle {
pub struct System {
pub app_config: Arc<ArcSwap<AppConfig>>,
pub runtime_state: Arc<ArcSwap<RuntimeState>>,
pub inference_config: Arc<MLInferenceConfig>,
pub database: Arc<Database>,
@ -134,6 +136,7 @@ impl System {
let state = ServiceFactory::build(database).await?;
Ok(System {
app_config: state.app_config,
runtime_state: state.runtime_state,
inference_config: state.inference_config,
database: state.database,
@ -187,7 +190,7 @@ impl System {
self.boot_response().await?;
let suricata_detection_tx = self.boot_detection();
self.boot_observability().await;
let mut shutdown_rx = self.boot_http_server(logger, log_buffer);
let mut shutdown_rx = self.boot_http_server(logger, log_buffer).await;
self.boot_external(suricata_detection_tx);
tokio::select! {
@ -203,7 +206,7 @@ impl System {
log!(SystemLog::EbpfBringupFailed(format!("{:?}", health)));
}
log!(SystemLog::InitializeComplete);
self.attach_ebpf()?;
self.attach_ebpf().await?;
} else {
log!(SystemLog::InitializeComplete);
}
@ -293,7 +296,12 @@ impl System {
}
async fn boot_observability(&mut self) {
let health_shutdown = self.health.clone().run(Duration::from_secs(3)).await;
let monitoring_interval_secs = self.app_config.load().health.monitoring_interval_secs;
let health_shutdown = self
.health
.clone()
.run(Duration::from_secs(monitoring_interval_secs))
.await;
self.health_shutdown = Some(health_shutdown);
let audit_logger = Arc::new(AuditLogger::new(self.database.clone() as Arc<dyn AuditRepo>));
@ -302,7 +310,7 @@ impl System {
let health_query: Arc<dyn crate::interface::health_query::HealthQuery> = self.health.clone();
let stats_aggregator = StatsAggregator::new(
self.database.clone() as Arc<dyn StatsRepo>,
self.database.clone() as Arc<dyn SettingRepo + Send + Sync>,
self.database.clone() as Arc<dyn ReportSnapshotRepo>,
health_query,
);
stats_aggregator.start();
@ -318,14 +326,13 @@ impl System {
}
}
fn boot_http_server(&mut self, logger: Arc<Logger>, log_buffer: Arc<LogBuffer>) -> mpsc::Receiver<ShutdownMode> {
async fn boot_http_server(
&mut self,
logger: Arc<Logger>,
log_buffer: Arc<LogBuffer>,
) -> mpsc::Receiver<ShutdownMode> {
let force_https = ForceHttpsFlag(Arc::new(AtomicBool::new(
self.database
.get_setting("force_https")
.ok()
.flatten()
.map(|v| v == "true")
.unwrap_or(false),
self.app_config.load().http_server.force_https,
)));
let readiness_state = self.readiness_state.clone();
@ -338,6 +345,7 @@ impl System {
let ready_flag_for_set = ready_flag.0.clone();
let params = HttpServerParams {
app_config: self.app_config.clone(),
runtime_state: self.runtime_state.clone(),
inference_config: self.inference_config.clone(),
database: self.database.clone(),
@ -421,7 +429,7 @@ impl System {
Ok(())
}
fn attach_ebpf(&mut self) -> Result<(), Error> {
async fn attach_ebpf(&mut self) -> Result<(), Error> {
let cfg = self.app_config.load();
let ingress_ifname = cfg.ebpf.ingress_ifname.clone();
let egress_ifname = cfg.ebpf.egress_ifname.clone();
@ -437,12 +445,10 @@ impl System {
match (ingress_result, egress_result) {
(Ok(ingress_mode), Ok(egress_mode)) => {
if let Err(e) = self.database.set_setting("xdp_ingress_mode", &ingress_mode) {
log!(SystemError::XdpModeStoreFailed(e));
}
if let Err(e) = self.database.set_setting("xdp_egress_mode", &egress_mode) {
log!(SystemError::XdpModeStoreFailed(e));
}
let mut next_state = (**self.runtime_state.load()).clone();
next_state.xdp.ingress_mode = ingress_mode;
next_state.xdp.egress_mode = egress_mode;
self.runtime_state.store(Arc::new(next_state));
}
(ingress_res, egress_res) => {
let (err, iface) = match (&ingress_res, &egress_res) {

View File

@ -1,3 +1,5 @@
use async_trait::async_trait;
use crate::domain::common::error::Error;
/// Data Plane BC — ACL aggregate repository.
@ -5,8 +7,9 @@ use crate::domain::common::error::Error;
/// Owns ACL rules (user-managed block/allow lists) and the admin whitelist that
/// SOAR must not block. Kept disjoint from `EnforcementRepo` (rate-limit / DNS /
/// geo) so policy tables can evolve independently of packet-matching tables.
#[async_trait]
pub trait AclRepo: Send + Sync {
fn insert_acl_rule(
async fn insert_acl_rule(
&self,
ip_version: u8,
direction: &str,
@ -15,7 +18,7 @@ pub trait AclRepo: Send + Sync {
port: u16,
) -> Result<(), Error>;
fn delete_acl_rule(
async fn delete_acl_rule(
&self,
ip_version: u8,
direction: &str,
@ -27,9 +30,9 @@ pub trait AclRepo: Send + Sync {
/// Returns true if a manual (non-SOAR) ACL rule exists for this IP.
/// Used by the TTL scheduler to avoid removing an eBPF block that the user
/// explicitly installed.
fn has_manual_acl_rule(&self, ip_address: &str) -> Result<bool, Error>;
async fn has_manual_acl_rule(&self, ip_address: &str) -> Result<bool, Error>;
fn list_admin_whitelist(&self) -> Result<Vec<String>, Error>;
fn insert_admin_whitelist(&self, ip: &str) -> Result<(), Error>;
fn delete_admin_whitelist(&self, ip: &str) -> Result<(), Error>;
async fn list_admin_whitelist(&self) -> Result<Vec<String>, Error>;
async fn insert_admin_whitelist(&self, ip: &str) -> Result<(), Error>;
async fn delete_admin_whitelist(&self, ip: &str) -> Result<(), Error>;
}

View File

@ -1,13 +1,16 @@
use async_trait::async_trait;
use crate::domain::common::error::Error;
use crate::domain::identity::auth::Claims;
use crate::domain::identity::user::ApiKeyView;
/// Identity BC — API key CRUD + validation (distinct from user login,
/// used by MCP / programmatic clients).
#[async_trait]
pub trait ApiKeyRepo: Send + Sync {
fn validate_api_key(&self, api_key: &str) -> Result<Option<Claims>, Error>;
async fn validate_api_key(&self, api_key: &str) -> Result<Option<Claims>, Error>;
fn hmac_api_key(&self, raw_key: &str) -> String;
fn insert_api_key(&self, key_hash: &str, name: &str, permission_level: &str) -> Result<i64, Error>;
fn list_api_keys(&self) -> Result<Vec<ApiKeyView>, Error>;
fn delete_api_key(&self, id: i64) -> Result<bool, Error>;
async fn insert_api_key(&self, key_hash: &str, name: &str, permission_level: &str) -> Result<i64, Error>;
async fn list_api_keys(&self) -> Result<Vec<ApiKeyView>, Error>;
async fn delete_api_key(&self, id: i64) -> Result<bool, Error>;
}

View File

@ -1,12 +1,14 @@
use crate::interface::acl::AclRepo;
use crate::interface::api_key::ApiKeyRepo;
use crate::interface::audit::AuditRepo;
use crate::interface::config_repo::ConfigRepo;
use crate::interface::db_admin::DbAdminRepo;
use crate::interface::enforcement::EnforcementRepo;
use crate::interface::identity::{LoginAttemptRepo, UserGroupRepo, UserRepo};
use crate::interface::setting::SettingRepo;
use crate::interface::report_snapshot::ReportSnapshotRepo;
use crate::interface::soar::SoarRepo;
use crate::interface::stats::StatsRepo;
use crate::interface::system_state::SystemStateRepo;
/// Composition-root supertrait bundling every aggregate Repo trait +
/// `DbAdminRepo`.
@ -26,9 +28,11 @@ pub trait AppRepo:
+ UserRepo
+ UserGroupRepo
+ LoginAttemptRepo
+ SettingRepo
+ ConfigRepo
+ SystemStateRepo
+ SoarRepo
+ StatsRepo
+ ReportSnapshotRepo
+ Send
+ Sync
{
@ -43,9 +47,11 @@ impl<T> AppRepo for T where
+ UserRepo
+ UserGroupRepo
+ LoginAttemptRepo
+ SettingRepo
+ ConfigRepo
+ SystemStateRepo
+ SoarRepo
+ StatsRepo
+ ReportSnapshotRepo
+ Send
+ Sync
+ ?Sized

Some files were not shown because too many files have changed in this diff Show More