mirror of
https://github.com/DaLaw2/NetGuardia.git
synced 2026-08-24 14:10:28 +09:00
Compare commits
2 Commits
d36d6d8e8e
...
fd95405ad8
| Author | SHA1 | Date | |
|---|---|---|---|
| fd95405ad8 | |||
| d419e3ee29 |
1
.gitignore
vendored
1
.gitignore
vendored
@ -36,6 +36,7 @@ interfaces.txt
|
||||
traffic_log.csv
|
||||
|
||||
CLAUDE.md
|
||||
AGENT.md
|
||||
DESIGN.md
|
||||
TODOS.md
|
||||
VERSION
|
||||
|
||||
150
Cargo.lock
generated
150
Cargo.lock
generated
@ -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]]
|
||||
|
||||
29
Cargo.toml
29
Cargo.toml
@ -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"
|
||||
|
||||
@ -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 }
|
||||
|
||||
@ -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 }
|
||||
|
||||
@ -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 }
|
||||
|
||||
@ -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(())
|
||||
|
||||
@ -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 }
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -36,7 +36,7 @@ impl GeoBlock {
|
||||
let v6_trie = LpmTrie::try_from(v6_map).map_err(EbpfError::MapOperationError)?;
|
||||
|
||||
// todo read config from AppConfig, not db
|
||||
let db_path = app_config.load().acl.geoip_db_name.clone();
|
||||
let db_path = app_config.load().acl.geoip_db_path.clone();
|
||||
let reader = Reader::open_readfile(&db_path).map_err(|e| MiscError::GeoIPDatabaseError(db_path.clone(), e))?;
|
||||
|
||||
let index = Self::build_index(&reader)?;
|
||||
@ -50,7 +50,7 @@ impl GeoBlock {
|
||||
}
|
||||
|
||||
pub fn unavailable(app_config: Arc<ArcSwap<AppConfig>>) -> Self {
|
||||
let index = Reader::open_readfile(&app_config.load().acl.geoip_db_name)
|
||||
let index = Reader::open_readfile(&app_config.load().acl.geoip_db_path)
|
||||
.ok()
|
||||
.and_then(|reader| Self::build_index(&reader).ok())
|
||||
.unwrap_or(GeoIndex {
|
||||
|
||||
@ -183,6 +183,9 @@ pub struct XskPair {
|
||||
drop_monitor: Option<Arc<DropMonitor>>,
|
||||
packet_buffer_size: usize,
|
||||
buffer_pool_capacity: usize,
|
||||
completion_batch_size: usize,
|
||||
rx_batch_size: usize,
|
||||
tx_batch_size: usize,
|
||||
tx_packet_buf: Vec<Vec<u8>>,
|
||||
tx_frame_buf: Vec<FrameDesc>,
|
||||
}
|
||||
@ -255,8 +258,11 @@ impl XskPair {
|
||||
drop_monitor,
|
||||
packet_buffer_size: config.packet_buffer_size,
|
||||
buffer_pool_capacity: config.buffer_pool_capacity,
|
||||
tx_packet_buf: Vec::with_capacity(64),
|
||||
tx_frame_buf: Vec::with_capacity(64),
|
||||
completion_batch_size: config.xsk_completion_batch_size,
|
||||
rx_batch_size: config.xsk_rx_batch_size,
|
||||
tx_batch_size: config.xsk_tx_batch_size,
|
||||
tx_packet_buf: Vec::with_capacity(config.xsk_tx_batch_size),
|
||||
tx_frame_buf: Vec::with_capacity(config.xsk_tx_batch_size),
|
||||
};
|
||||
|
||||
Ok(xsk_pair)
|
||||
@ -277,8 +283,8 @@ impl XskPair {
|
||||
let mut shutdown_rx = Some(shutdown_rx);
|
||||
let mut idle_count: u32 = 0;
|
||||
let mut buffer_pool = BufferPool::new(self.buffer_pool_capacity, self.packet_buffer_size);
|
||||
let mut comp_descs = vec![FrameDesc::default(); 256];
|
||||
let mut rx_descs = vec![FrameDesc::default(); 64];
|
||||
let mut comp_descs = vec![FrameDesc::default(); self.completion_batch_size];
|
||||
let mut rx_descs = vec![FrameDesc::default(); self.rx_batch_size];
|
||||
|
||||
loop {
|
||||
if let Some(ref mut rx) = shutdown_rx {
|
||||
@ -426,7 +432,7 @@ impl XskPair {
|
||||
self.tx_packet_buf.clear();
|
||||
while let Ok(packet) = forward_rx.try_recv() {
|
||||
self.tx_packet_buf.push(packet);
|
||||
if self.tx_packet_buf.len() >= 64 {
|
||||
if self.tx_packet_buf.len() >= self.tx_batch_size {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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()})),
|
||||
}
|
||||
|
||||
@ -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)
|
||||
}
|
||||
|
||||
@ -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!({
|
||||
|
||||
@ -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();
|
||||
|
||||
@ -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()})),
|
||||
|
||||
@ -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(®.username, ®.password, ®.role, &auth.role) {
|
||||
match auth_svc
|
||||
.register(®.username, ®.password, ®.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()})),
|
||||
|
||||
@ -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);
|
||||
|
||||
@ -1,6 +1,6 @@
|
||||
use std::fs;
|
||||
use std::io::ErrorKind;
|
||||
use std::path::Path;
|
||||
use std::path::PathBuf;
|
||||
use std::time::UNIX_EPOCH;
|
||||
|
||||
use actix_web::{HttpResponse, Scope, web};
|
||||
@ -10,9 +10,6 @@ use serde::{Deserialize, Serialize};
|
||||
use crate::domain::common::config::AppConfig;
|
||||
use crate::infrastructure::log_buffer::{self, LogBuffer, LogEntry};
|
||||
|
||||
/// Hardcoded log directory — not configurable via API to prevent directory traversal.
|
||||
const LOG_DIR: &str = "logs";
|
||||
|
||||
/// Validate log filename: only alphanumeric, dots, underscores, hyphens.
|
||||
/// Prevents path traversal.
|
||||
fn is_valid_log_filename(name: &str) -> bool {
|
||||
@ -87,9 +84,9 @@ struct LogFileEntry {
|
||||
modified: Option<u64>,
|
||||
}
|
||||
|
||||
async fn list_logs() -> HttpResponse {
|
||||
let log_dir = LOG_DIR;
|
||||
let entries = match fs::read_dir(log_dir) {
|
||||
async fn list_logs(app_config: web::Data<ArcSwap<AppConfig>>) -> HttpResponse {
|
||||
let log_dir = app_config.load().system.log_dir.clone();
|
||||
let entries = match fs::read_dir(&log_dir) {
|
||||
Ok(dir) => dir
|
||||
.filter_map(|e| e.ok())
|
||||
.filter_map(|e| {
|
||||
@ -117,7 +114,9 @@ async fn list_logs() -> HttpResponse {
|
||||
}
|
||||
|
||||
async fn download_log(path: web::Path<String>, app_config: web::Data<ArcSwap<AppConfig>>) -> HttpResponse {
|
||||
let max_download_size = app_config.load().observability.log_max_download_size;
|
||||
let config = app_config.load();
|
||||
let max_download_size = config.observability.log_max_download_size;
|
||||
let log_dir = PathBuf::from(&config.system.log_dir);
|
||||
let filename = path.into_inner();
|
||||
|
||||
if !is_valid_log_filename(&filename) {
|
||||
@ -126,7 +125,7 @@ async fn download_log(path: web::Path<String>, app_config: web::Data<ArcSwap<App
|
||||
}));
|
||||
}
|
||||
|
||||
let file_path = Path::new(LOG_DIR).join(&filename);
|
||||
let file_path = log_dir.join(&filename);
|
||||
|
||||
// Canonicalize to prevent symlink traversal
|
||||
let canonical = match fs::canonicalize(&file_path) {
|
||||
@ -137,7 +136,7 @@ async fn download_log(path: web::Path<String>, app_config: web::Data<ArcSwap<App
|
||||
}));
|
||||
}
|
||||
};
|
||||
if let Ok(log_dir_canonical) = fs::canonicalize(LOG_DIR)
|
||||
if let Ok(log_dir_canonical) = fs::canonicalize(&log_dir)
|
||||
&& !canonical.starts_with(&log_dir_canonical)
|
||||
{
|
||||
return HttpResponse::Forbidden().json(serde_json::json!({
|
||||
|
||||
@ -63,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())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@ -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()})),
|
||||
}
|
||||
|
||||
@ -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!({
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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(())
|
||||
|
||||
@ -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 {
|
||||
|
||||
@ -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>))
|
||||
}
|
||||
}
|
||||
|
||||
@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@ -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()));
|
||||
}
|
||||
}
|
||||
|
||||
@ -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
|
||||
}
|
||||
}
|
||||
|
||||
153
net-guardia/src/adapter/persistence/config.rs
Normal file
153
net-guardia/src/adapter/persistence/config.rs
Normal 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
|
||||
}
|
||||
}
|
||||
@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@ -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();
|
||||
}
|
||||
}
|
||||
|
||||
51
net-guardia/src/adapter/persistence/report_snapshot.rs
Normal file
51
net-guardia/src/adapter/persistence/report_snapshot.rs
Normal 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
|
||||
}
|
||||
}
|
||||
@ -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(""));
|
||||
}
|
||||
}
|
||||
@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@ -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
|
||||
}
|
||||
}
|
||||
|
||||
49
net-guardia/src/adapter/persistence/system_state.rs
Normal file
49
net-guardia/src/adapter/persistence/system_state.rs
Normal 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
|
||||
}
|
||||
}
|
||||
@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@ -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 {
|
||||
|
||||
@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@ -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. \
|
||||
|
||||
@ -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();
|
||||
|
||||
@ -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)
|
||||
}
|
||||
|
||||
@ -25,6 +25,10 @@ impl DnsFilter {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn validate_domain(&self, domain: &str) -> Result<(), Error> {
|
||||
domain_to_wire_format(domain).map(|_| ())
|
||||
}
|
||||
|
||||
pub fn remove_domain(&self, domain: &str) -> Result<(), Error> {
|
||||
let name = domain_to_wire_format(domain)?;
|
||||
self.blacklist.remove(&name);
|
||||
@ -185,6 +189,9 @@ impl DnsFilter {
|
||||
}
|
||||
|
||||
impl DnsFilterPort for DnsFilter {
|
||||
fn validate_domain(&self, domain: &str) -> Result<(), Error> {
|
||||
self.validate_domain(domain)
|
||||
}
|
||||
fn add_domain(&self, domain: &str) -> Result<(), Error> {
|
||||
self.add_domain(domain)
|
||||
}
|
||||
|
||||
@ -5,19 +5,24 @@ use arc_swap::ArcSwap;
|
||||
use crate::domain::common::config::AppConfig;
|
||||
use crate::domain::common::error::Error;
|
||||
use crate::domain::common::error::misc::MiscError;
|
||||
use crate::interface::app_repo::AppRepo;
|
||||
use crate::interface::dns_filter_api::DnsFilterPort;
|
||||
use crate::interface::enforcement::EnforcementRepo;
|
||||
|
||||
/// Domain service that coordinates DNS filter changes between DB and in-memory service.
|
||||
/// Write order: eBPF/in-memory first, then DB — if eBPF fails, DB remains clean.
|
||||
/// Domains are validated up front, runtime state is changed first, and DB batch
|
||||
/// failure rolls runtime back so persisted and live policy do not drift.
|
||||
pub struct DnsFilterService {
|
||||
db: Arc<dyn AppRepo>,
|
||||
db: Arc<dyn EnforcementRepo>,
|
||||
dns_filter: Arc<dyn DnsFilterPort>,
|
||||
config: Arc<ArcSwap<AppConfig>>,
|
||||
}
|
||||
|
||||
impl DnsFilterService {
|
||||
pub fn new(db: Arc<dyn AppRepo>, dns_filter: Arc<dyn DnsFilterPort>, config: Arc<ArcSwap<AppConfig>>) -> Self {
|
||||
pub fn new(
|
||||
db: Arc<dyn EnforcementRepo>,
|
||||
dns_filter: Arc<dyn DnsFilterPort>,
|
||||
config: Arc<ArcSwap<AppConfig>>,
|
||||
) -> Self {
|
||||
Self { db, dns_filter, config }
|
||||
}
|
||||
|
||||
@ -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()));
|
||||
}
|
||||
}
|
||||
|
||||
@ -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(())
|
||||
|
||||
@ -34,6 +34,7 @@ impl BeaconingDetector {
|
||||
beaconing.min_observations,
|
||||
beaconing.cv_threshold,
|
||||
beaconing.max_cache_entries,
|
||||
beaconing.max_timestamps_per_flow,
|
||||
beaconing.expiry_secs,
|
||||
beaconing.alert_cooldown_secs,
|
||||
),
|
||||
@ -78,8 +79,6 @@ impl BeaconingDetector {
|
||||
|
||||
type FlowTuple = (String, String, u16);
|
||||
|
||||
const MAX_TIMESTAMPS_PER_FLOW: usize = 100;
|
||||
|
||||
struct CachedFlow {
|
||||
timestamps: Vec<Instant>,
|
||||
last_alerted: Option<Instant>,
|
||||
@ -90,6 +89,7 @@ pub struct BeaconingState {
|
||||
min_observations: usize,
|
||||
cv_threshold: f64,
|
||||
max_cache_entries: usize,
|
||||
max_timestamps_per_flow: usize,
|
||||
expiry_secs: u64,
|
||||
alert_cooldown_secs: u64,
|
||||
}
|
||||
@ -99,6 +99,7 @@ impl BeaconingState {
|
||||
min_observations: usize,
|
||||
cv_threshold: f64,
|
||||
max_cache_entries: usize,
|
||||
max_timestamps_per_flow: usize,
|
||||
expiry_secs: u64,
|
||||
alert_cooldown_secs: u64,
|
||||
) -> Self {
|
||||
@ -107,6 +108,7 @@ impl BeaconingState {
|
||||
min_observations,
|
||||
cv_threshold,
|
||||
max_cache_entries,
|
||||
max_timestamps_per_flow,
|
||||
expiry_secs,
|
||||
alert_cooldown_secs,
|
||||
}
|
||||
@ -123,8 +125,8 @@ impl BeaconingState {
|
||||
|
||||
entry.timestamps.push(now);
|
||||
|
||||
if entry.timestamps.len() > MAX_TIMESTAMPS_PER_FLOW {
|
||||
let excess = entry.timestamps.len() - MAX_TIMESTAMPS_PER_FLOW;
|
||||
if entry.timestamps.len() > self.max_timestamps_per_flow {
|
||||
let excess = entry.timestamps.len() - self.max_timestamps_per_flow;
|
||||
entry.timestamps.drain(..excess);
|
||||
}
|
||||
}
|
||||
@ -280,7 +282,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn beaconing_state_detects_periodic_flows() {
|
||||
let state = BeaconingState::new(5, 0.3, 50_000, 3600, 120);
|
||||
let state = BeaconingState::new(5, 0.3, 50_000, 100, 3600, 120);
|
||||
let base = Instant::now();
|
||||
let key = ("10.0.0.1".to_string(), "1.2.3.4".to_string(), 443_u16);
|
||||
state.flow_cache.insert(
|
||||
@ -295,4 +297,33 @@ mod tests {
|
||||
assert_eq!(events[0].source, DetectionSource::Beaconing);
|
||||
assert_eq!(events[0].attack_type, CanonicalAttackType::C2Beacon.as_str());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn record_flow_respects_configured_timestamp_cap() {
|
||||
let state = BeaconingState::new(1, 0.3, 50_000, 3, 3600, 120);
|
||||
let alert = crate::domain::detection::ml_detection::AlertMessage {
|
||||
timestamp: 0,
|
||||
flow_key: "10.0.0.1:12345-1.2.3.4:443".to_string(),
|
||||
src_ip: "10.0.0.1".to_string(),
|
||||
dst_ip: "1.2.3.4".to_string(),
|
||||
src_port: 12345,
|
||||
dst_port: 443,
|
||||
protocol: 6,
|
||||
is_attack: true,
|
||||
attack_type: Some("c2_beacon".to_string()),
|
||||
confidence: 0.9,
|
||||
packet_count: 10,
|
||||
flow_duration_us: 1000,
|
||||
ae_score: 0.0,
|
||||
anomaly_score: 0.0,
|
||||
c2_score: 0.0,
|
||||
};
|
||||
|
||||
for _ in 0..5 {
|
||||
state.record_flow(&alert);
|
||||
}
|
||||
|
||||
let key = ("10.0.0.1".to_string(), "1.2.3.4".to_string(), 443_u16);
|
||||
assert_eq!(state.flow_cache.get(&key).unwrap().timestamps.len(), 3);
|
||||
}
|
||||
}
|
||||
|
||||
@ -78,16 +78,17 @@ impl DetectionOrchestrator {
|
||||
) -> Self {
|
||||
let cfg = app_config.load();
|
||||
let fusion = &cfg.detection.fusion;
|
||||
let max_dedup = NonZero::new(fusion.max_dedup_entries.max(1)).unwrap_or(NonZero::<usize>::MIN);
|
||||
let source_count_max_entries = nonzero_cache_size(fusion.source_count_max_entries);
|
||||
let repeat_tracker_max_entries = nonzero_cache_size(fusion.repeat_tracker_max_entries);
|
||||
let max_dedup = nonzero_cache_size(fusion.max_dedup_entries);
|
||||
Self {
|
||||
rx,
|
||||
threat_tx,
|
||||
audit_tx,
|
||||
geoip,
|
||||
metrics,
|
||||
// SAFETY: NonZero::new on non-zero literals.
|
||||
src_ip_counts: LruCache::new(NonZero::new(10_000).unwrap()),
|
||||
repeat_tracker: LruCache::new(NonZero::new(5_000).unwrap()),
|
||||
src_ip_counts: LruCache::new(source_count_max_entries),
|
||||
repeat_tracker: LruCache::new(repeat_tracker_max_entries),
|
||||
dedup: LruCache::new(max_dedup),
|
||||
dedup_window: Duration::from_secs(fusion.dedup_window_secs),
|
||||
repeat_offender_window: Duration::from_secs(fusion.repeat_offender_window_secs),
|
||||
@ -355,6 +356,10 @@ impl DetectionOrchestrator {
|
||||
}
|
||||
}
|
||||
|
||||
fn nonzero_cache_size(value: usize) -> NonZero<usize> {
|
||||
NonZero::new(value.max(1)).unwrap_or(NonZero::<usize>::MIN)
|
||||
}
|
||||
|
||||
pub async fn bridge_ml_to_detection(mut rx: broadcast::Receiver<AlertMessage>, tx: mpsc::Sender<DetectionEvent>) {
|
||||
log!(DetectionLog::MlBridgeStarted);
|
||||
|
||||
|
||||
@ -34,6 +34,7 @@ pub struct UserProfile {
|
||||
pub groups: Vec<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum LoginError {
|
||||
Locked { retry_after_secs: u64 },
|
||||
InvalidCredentials,
|
||||
@ -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()));
|
||||
}
|
||||
}
|
||||
|
||||
@ -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} — 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 — 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("<img src=x onerror=alert(1)>"));
|
||||
assert!(html.contains(""><script>alert(1)</script>"));
|
||||
}
|
||||
}
|
||||
|
||||
@ -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()));
|
||||
|
||||
8
net-guardia/src/core/reporting/html.rs
Normal file
8
net-guardia/src/core/reporting/html.rs
Normal file
@ -0,0 +1,8 @@
|
||||
pub fn escape(input: &str) -> String {
|
||||
input
|
||||
.replace('&', "&")
|
||||
.replace('<', "<")
|
||||
.replace('>', ">")
|
||||
.replace('"', """)
|
||||
.replace('\'', "'")
|
||||
}
|
||||
@ -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;
|
||||
|
||||
@ -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);
|
||||
|
||||
|
||||
@ -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('&', "&")
|
||||
.replace('<', "<")
|
||||
.replace('>', ">")
|
||||
.replace('"', """)
|
||||
}
|
||||
|
||||
/// 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())
|
||||
}
|
||||
|
||||
@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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(),
|
||||
|
||||
@ -11,14 +11,16 @@ pub struct FrequencyTracker {
|
||||
events: DashMap<FreqKey, VecDeque<Instant>>,
|
||||
max_deque_size: usize,
|
||||
max_tracked_keys: usize,
|
||||
max_retention: Duration,
|
||||
}
|
||||
|
||||
impl FrequencyTracker {
|
||||
pub fn new(max_tracked_keys: usize) -> Self {
|
||||
pub fn new(max_tracked_keys: usize, max_events_per_key: usize, retention_secs: u64) -> Self {
|
||||
Self {
|
||||
events: DashMap::new(),
|
||||
max_deque_size: 200,
|
||||
max_deque_size: max_events_per_key.max(1),
|
||||
max_tracked_keys: max_tracked_keys.max(1),
|
||||
max_retention: Duration::from_secs(retention_secs.max(1)),
|
||||
}
|
||||
}
|
||||
|
||||
@ -50,20 +52,19 @@ impl FrequencyTracker {
|
||||
deque.len() as u64
|
||||
}
|
||||
|
||||
/// Remove empty deques and entries where all timestamps are expired.
|
||||
/// Uses a conservative 2-hour max window for expiry detection.
|
||||
/// Remove empty deques and entries where all timestamps are outside the
|
||||
/// configured retention window.
|
||||
pub fn cleanup(&self) -> u32 {
|
||||
let now = Instant::now();
|
||||
let max_window = Duration::from_secs(7200); // 2 hours — conservative upper bound
|
||||
let mut removed = 0u32;
|
||||
self.events.retain(|_, deque| {
|
||||
if deque.is_empty() {
|
||||
removed += 1;
|
||||
return false;
|
||||
}
|
||||
// If all entries are older than max_window, remove the whole entry
|
||||
// If all entries are older than max_retention, remove the whole entry.
|
||||
if let Some(newest) = deque.back()
|
||||
&& now.checked_duration_since(*newest).unwrap_or(Duration::ZERO) > max_window
|
||||
&& now.checked_duration_since(*newest).unwrap_or(Duration::ZERO) > self.max_retention
|
||||
{
|
||||
removed += 1;
|
||||
return false;
|
||||
@ -84,3 +85,17 @@ impl FrequencyTracker {
|
||||
removed
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::FrequencyTracker;
|
||||
|
||||
#[test]
|
||||
fn record_and_count_respects_configured_per_key_cap() {
|
||||
let tracker = FrequencyTracker::new(10, 2, 60);
|
||||
|
||||
assert_eq!(tracker.record_and_count(1, "10.0.0.1", 60), 1);
|
||||
assert_eq!(tracker.record_and_count(1, "10.0.0.1", 60), 2);
|
||||
assert_eq!(tracker.record_and_count(1, "10.0.0.1", 60), 2);
|
||||
}
|
||||
}
|
||||
|
||||
@ -29,12 +29,21 @@ pub struct PlaybookMatcher {
|
||||
}
|
||||
|
||||
impl PlaybookMatcher {
|
||||
pub fn new(config: Arc<ArcSwap<AppConfig>>, frequency_max_tracked_keys: usize) -> Self {
|
||||
pub fn new(
|
||||
config: Arc<ArcSwap<AppConfig>>,
|
||||
frequency_max_tracked_keys: usize,
|
||||
frequency_max_events_per_key: usize,
|
||||
frequency_retention_secs: u64,
|
||||
) -> Self {
|
||||
Self {
|
||||
playbooks: ArcSwap::from_pointee(Vec::new()),
|
||||
admin_whitelist: ArcSwap::from_pointee(HashSet::new()),
|
||||
cooldowns: DashMap::new(),
|
||||
frequency_tracker: FrequencyTracker::new(frequency_max_tracked_keys),
|
||||
frequency_tracker: FrequencyTracker::new(
|
||||
frequency_max_tracked_keys,
|
||||
frequency_max_events_per_key,
|
||||
frequency_retention_secs,
|
||||
),
|
||||
active_block_count: AtomicU32::new(0),
|
||||
config,
|
||||
}
|
||||
|
||||
@ -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(())
|
||||
}
|
||||
}
|
||||
|
||||
@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@ -3,6 +3,8 @@ use macros::config_settings;
|
||||
#[config_settings(section = "misc")]
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct AclConfig {
|
||||
#[setting(key = "geoip_db_name", default = "net-guardia/static/geo/dbip-city-lite.mmdb")]
|
||||
pub geoip_db_name: String,
|
||||
#[setting(key = "geoip_db_path", default = "net-guardia/static/geo/dbip-city-lite.mmdb")]
|
||||
pub geoip_db_path: String,
|
||||
#[setting(key = "geoip_cache_capacity", default = "10000")]
|
||||
pub geoip_cache_capacity: usize,
|
||||
}
|
||||
|
||||
@ -9,6 +9,10 @@ pub struct FusionConfig {
|
||||
pub repeat_offender_window_secs: u64,
|
||||
#[setting(key = "fusion_max_dedup_entries", default = "50000")]
|
||||
pub max_dedup_entries: usize,
|
||||
#[setting(key = "fusion_source_count_max_entries", default = "10000")]
|
||||
pub source_count_max_entries: usize,
|
||||
#[setting(key = "fusion_repeat_tracker_max_entries", default = "5000")]
|
||||
pub repeat_tracker_max_entries: usize,
|
||||
}
|
||||
|
||||
#[config_settings(section = "beaconing")]
|
||||
@ -22,12 +26,23 @@ pub struct BeaconingConfig {
|
||||
pub cv_threshold: f64,
|
||||
#[setting(key = "beaconing_max_cache_entries", default = "50000")]
|
||||
pub max_cache_entries: usize,
|
||||
#[setting(key = "beaconing_max_timestamps_per_flow", default = "100")]
|
||||
pub max_timestamps_per_flow: usize,
|
||||
#[setting(key = "beaconing_expiry_secs", default = "600")]
|
||||
pub expiry_secs: u64,
|
||||
#[setting(key = "beaconing_alert_cooldown_secs", default = "300")]
|
||||
pub alert_cooldown_secs: u64,
|
||||
}
|
||||
|
||||
#[config_settings(section = "flow_stats")]
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct FlowStatsConfig {
|
||||
#[setting(key = "flow_stats_max_snapshot_entries", default = "10000")]
|
||||
pub max_snapshot_entries: usize,
|
||||
#[setting(key = "flow_stats_max_top_n", default = "10000")]
|
||||
pub max_top_n: usize,
|
||||
}
|
||||
|
||||
#[config_settings(section = "detection")]
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct DetectionConfig {
|
||||
@ -35,6 +50,8 @@ pub struct DetectionConfig {
|
||||
pub fusion: FusionConfig,
|
||||
#[setting(flatten)]
|
||||
pub beaconing: BeaconingConfig,
|
||||
#[setting(flatten)]
|
||||
pub flow_stats: FlowStatsConfig,
|
||||
#[setting(key = "detection_cleanup_interval_secs", default = "60")]
|
||||
pub cleanup_interval_secs: u64,
|
||||
}
|
||||
|
||||
@ -32,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,
|
||||
|
||||
@ -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,
|
||||
}
|
||||
|
||||
@ -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(())
|
||||
}
|
||||
|
||||
@ -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,
|
||||
|
||||
|
||||
@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@ -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,
|
||||
}
|
||||
|
||||
@ -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)]
|
||||
|
||||
@ -5,9 +5,6 @@ traceable! {
|
||||
#[error("Database error: {err}")]
|
||||
QueryFailed => tracing::Level::ERROR,
|
||||
|
||||
#[error("Database connection failed")]
|
||||
ConnectionFailed => tracing::Level::ERROR,
|
||||
|
||||
#[no_source]
|
||||
#[error("User '{username}' already exists")]
|
||||
UserAlreadyExists { username: String } => tracing::Level::WARN,
|
||||
|
||||
@ -1,25 +0,0 @@
|
||||
use macros::traceable;
|
||||
|
||||
traceable! {
|
||||
McpError {
|
||||
#[no_source]
|
||||
#[error("MCP tool not found: {tool_name}")]
|
||||
ToolNotFound { tool_name: String } => tracing::Level::WARN,
|
||||
|
||||
#[no_source]
|
||||
#[error("MCP parameter validation failed: {detail}")]
|
||||
InvalidParams { detail: String } => tracing::Level::WARN,
|
||||
|
||||
#[no_source]
|
||||
#[error("MCP permission denied: key has '{key_level}' but tool requires '{required_level}'")]
|
||||
PermissionDenied { key_level: String, required_level: String } => tracing::Level::WARN,
|
||||
|
||||
#[no_source]
|
||||
#[error("MCP API key invalid or revoked")]
|
||||
InvalidApiKey => tracing::Level::WARN,
|
||||
|
||||
#[no_source]
|
||||
#[error("MCP proxy error: {detail}")]
|
||||
ProxyError { detail: String } => tracing::Level::ERROR,
|
||||
}
|
||||
}
|
||||
@ -6,19 +6,9 @@ traceable! {
|
||||
#[error("Failed to remove limit on locked memory, ret is: {ret}")]
|
||||
RamLimitUnlockError { ret: i32 } => tracing::Level::ERROR,
|
||||
|
||||
#[error("Failed to send message to receiver")]
|
||||
SendMessageError => tracing::Level::ERROR,
|
||||
|
||||
#[error("Failed to serialize data")]
|
||||
SerializeError => tracing::Level::ERROR,
|
||||
|
||||
#[error("Failed to deserialize data")]
|
||||
DeserializeError => tracing::Level::ERROR,
|
||||
|
||||
#[no_source]
|
||||
#[error("Network interface '{interface}' not found")]
|
||||
NetworkInterfaceNotFound { interface: String } => tracing::Level::ERROR,
|
||||
|
||||
#[error("Failed to open GeoIP database '{path}': {err}")]
|
||||
GeoIPDatabaseError { path: String } => tracing::Level::ERROR,
|
||||
|
||||
@ -33,18 +23,6 @@ traceable! {
|
||||
#[error("DNS domain name too long: '{domain}'")]
|
||||
DnsDomainTooLong { domain: String } => tracing::Level::WARN,
|
||||
|
||||
#[no_source]
|
||||
#[error("Type mismatch during message dispatch")]
|
||||
TypeMismatch => tracing::Level::ERROR,
|
||||
|
||||
#[no_source]
|
||||
#[error("No handler registered for this message type")]
|
||||
HandlerNotFound => tracing::Level::ERROR,
|
||||
|
||||
#[no_source]
|
||||
#[error("Event type not registered with communication manager")]
|
||||
TypeNotRegistered => tracing::Level::ERROR,
|
||||
|
||||
#[no_source]
|
||||
#[error("Validation error: {message}")]
|
||||
ValidationError { message: String } => tracing::Level::WARN,
|
||||
|
||||
@ -2,7 +2,6 @@ pub mod crypto;
|
||||
pub mod database;
|
||||
pub mod http;
|
||||
pub mod io;
|
||||
pub mod mcp;
|
||||
pub mod misc;
|
||||
pub mod notification;
|
||||
pub mod system;
|
||||
@ -13,7 +12,6 @@ use crate::domain::common::error::crypto::CryptoError;
|
||||
use crate::domain::common::error::database::DatabaseError;
|
||||
use crate::domain::common::error::http::HttpError;
|
||||
use crate::domain::common::error::io::IOError;
|
||||
use crate::domain::common::error::mcp::McpError;
|
||||
use crate::domain::common::error::misc::MiscError;
|
||||
use crate::domain::common::error::notification::NotificationError;
|
||||
use crate::domain::common::error::system::SystemError;
|
||||
@ -40,8 +38,6 @@ pub enum Error {
|
||||
#[error("{0}")]
|
||||
IO(#[from] IOError),
|
||||
#[error("{0}")]
|
||||
Mcp(#[from] McpError),
|
||||
#[error("{0}")]
|
||||
Misc(#[from] MiscError),
|
||||
#[error("{0}")]
|
||||
Notification(#[from] NotificationError),
|
||||
|
||||
@ -2,10 +2,6 @@ use macros::traceable;
|
||||
|
||||
traceable! {
|
||||
SystemError {
|
||||
#[no_source]
|
||||
#[error("Unable to run as administrator")]
|
||||
RunAsAdminFailed => tracing::Level::ERROR,
|
||||
|
||||
#[no_source]
|
||||
#[error("Invalid configuration")]
|
||||
InvalidConfig => tracing::Level::ERROR,
|
||||
@ -17,38 +13,24 @@ traceable! {
|
||||
#[error("Configuration file not found")]
|
||||
ConfigNotFound => tracing::Level::ERROR,
|
||||
|
||||
#[error("Failed to terminate instance")]
|
||||
TerminateError => tracing::Level::ERROR,
|
||||
|
||||
#[no_source]
|
||||
#[error("Failed to send shutdown signal")]
|
||||
ShutdownSignalFailed => tracing::Level::ERROR,
|
||||
|
||||
#[error("Unexpected thread panic")]
|
||||
ThreadPanic => tracing::Level::ERROR,
|
||||
|
||||
#[error("Unexpected error")]
|
||||
UnexpectedError => tracing::Level::ERROR,
|
||||
|
||||
#[error("Failed to reload config after setup")]
|
||||
ConfigReloadFailed => tracing::Level::ERROR,
|
||||
|
||||
#[error("HTTP server error")]
|
||||
HttpServerError => tracing::Level::ERROR,
|
||||
|
||||
#[error("Failed to set user groups")]
|
||||
SetUserGroupsFailed => tracing::Level::WARN,
|
||||
|
||||
#[error("Failed to store XDP mode")]
|
||||
XdpModeStoreFailed => tracing::Level::WARN,
|
||||
|
||||
#[error("Failed to update admin password during setup: {err}")]
|
||||
SetupPasswordUpdateFailed => tracing::Level::ERROR,
|
||||
|
||||
#[error("Failed to mark setup as complete: {err}")]
|
||||
SetupCompleteFlagFailed => tracing::Level::ERROR,
|
||||
|
||||
#[error("Failed to publish drift detected event")]
|
||||
DriftEventPublishFailed => tracing::Level::WARN,
|
||||
}
|
||||
}
|
||||
|
||||
@ -3,9 +3,6 @@ use tracing;
|
||||
|
||||
loggable! {
|
||||
SystemLog {
|
||||
#[error("Online now")]
|
||||
Online => tracing::Level::INFO,
|
||||
|
||||
#[error("Initializing")]
|
||||
Initializing => tracing::Level::INFO,
|
||||
|
||||
@ -27,9 +24,6 @@ loggable! {
|
||||
#[error("Setup wizard completed — starting full system initialization")]
|
||||
SetupCompleted => tracing::Level::INFO,
|
||||
|
||||
#[error("Config reloaded from DB: ingress={ingress}, egress={egress}")]
|
||||
ConfigReloaded { ingress: String, egress: String } => tracing::Level::INFO,
|
||||
|
||||
#[error("Full system initialization complete — all services running")]
|
||||
FullInitComplete => tracing::Level::INFO,
|
||||
|
||||
|
||||
@ -63,8 +63,8 @@ pub enum EbpfFailCategory {
|
||||
InterfaceNotFound,
|
||||
/// Interface exists but XDP native/SKB attach refused by driver.
|
||||
XdpUnsupported,
|
||||
/// AF_XDP bind rejected — driver does not implement AF_XDP on this kernel.
|
||||
/// Common case: Intel i350 (igb) on kernel < 6.17.
|
||||
/// AF_XDP bind rejected — driver or netdev capabilities do not support
|
||||
/// the requested AF_XDP socket mode on this interface.
|
||||
AfXdpUnsupported,
|
||||
/// ENOMEM / RLIMIT_MEMLOCK exhausted.
|
||||
MemlockExhausted,
|
||||
|
||||
@ -2,6 +2,28 @@ use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::domain::data_plane::direction::Direction;
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct FlowStatsLimits {
|
||||
max_snapshot_entries: usize,
|
||||
max_top_n: usize,
|
||||
}
|
||||
|
||||
impl FlowStatsLimits {
|
||||
pub fn new(max_snapshot_entries: usize, max_top_n: usize) -> Self {
|
||||
Self {
|
||||
max_snapshot_entries: max_snapshot_entries.max(1),
|
||||
max_top_n: max_top_n.max(1),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn result_limit(self, requested_top_n: Option<usize>) -> usize {
|
||||
requested_top_n
|
||||
.unwrap_or(self.max_snapshot_entries)
|
||||
.max(1)
|
||||
.min(self.max_top_n)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub struct FlowStatsEntry {
|
||||
pub direction: Direction,
|
||||
@ -61,3 +83,29 @@ pub struct FlowSubscription {
|
||||
/// Push interval in seconds (default 5)
|
||||
pub interval_secs: Option<u64>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::FlowStatsLimits;
|
||||
|
||||
#[test]
|
||||
fn result_limit_uses_snapshot_default_when_top_n_absent() {
|
||||
let limits = FlowStatsLimits::new(128, 1024);
|
||||
|
||||
assert_eq!(limits.result_limit(None), 128);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn result_limit_caps_requested_top_n() {
|
||||
let limits = FlowStatsLimits::new(128, 256);
|
||||
|
||||
assert_eq!(limits.result_limit(Some(1_000)), 256);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn result_limit_never_returns_zero() {
|
||||
let limits = FlowStatsLimits::new(0, 0);
|
||||
|
||||
assert_eq!(limits.result_limit(Some(0)), 1);
|
||||
}
|
||||
}
|
||||
|
||||
@ -36,7 +36,8 @@ impl PermissionLevel {
|
||||
pub fn permissions(self) -> &'static [&'static str] {
|
||||
match self {
|
||||
Self::ReadOnly => API_KEY_READ_ONLY_PERMISSIONS,
|
||||
Self::ReadWrite | Self::FullAccess => API_KEY_READ_WRITE_PERMISSIONS,
|
||||
Self::ReadWrite => API_KEY_READ_WRITE_PERMISSIONS,
|
||||
Self::FullAccess => API_KEY_FULL_ACCESS_PERMISSIONS,
|
||||
}
|
||||
}
|
||||
}
|
||||
@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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()));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@ -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(())
|
||||
|
||||
@ -10,7 +10,17 @@
|
||||
//! host kernel release, and the NIC driver where those are obtainable,
|
||||
//! so the operator can diagnose directly from the UI without shelling in.
|
||||
|
||||
use std::ffi::OsStr;
|
||||
use std::fmt;
|
||||
use std::fs;
|
||||
use std::mem;
|
||||
use std::os::fd::RawFd;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use libc::{
|
||||
AF_NETLINK, NETLINK_GENERIC, SOCK_CLOEXEC, SOCK_RAW, bind, close, genlmsghdr, nlmsghdr, recv, sendto, sockaddr,
|
||||
sockaddr_nl, socket,
|
||||
};
|
||||
|
||||
use crate::domain::common::error::Error;
|
||||
use crate::domain::common::system::health::{EbpfFailCategory, EbpfFailStage, EbpfHealth};
|
||||
@ -24,19 +34,92 @@ pub fn kernel_release() -> String {
|
||||
.unwrap_or_else(|| "unknown".to_string())
|
||||
}
|
||||
|
||||
/// Look up the driver name bound to a network interface via
|
||||
/// `/sys/class/net/<ifname>/device/driver`. Returns the basename of
|
||||
/// the symlink target, or `"unknown"` if the interface has no driver
|
||||
/// (e.g., virtual or renamed) or the path is not readable.
|
||||
pub fn interface_driver(ifname: &str) -> String {
|
||||
let link = format!("/sys/class/net/{}/device/driver", ifname);
|
||||
match fs::read_link(&link) {
|
||||
Ok(target) => target
|
||||
.file_name()
|
||||
.and_then(|s| s.to_str())
|
||||
.map(|s| s.to_string())
|
||||
.unwrap_or_else(|| "unknown".to_string()),
|
||||
Err(_) => "unknown".to_string(),
|
||||
const SYS_CLASS_NET: &str = "/sys/class/net";
|
||||
const NETDEV_FAMILY_NAME: &str = "netdev";
|
||||
const GENL_ID_CTRL: u16 = 0x10;
|
||||
const CTRL_CMD_GETFAMILY: u8 = 3;
|
||||
const CTRL_ATTR_FAMILY_ID: u16 = 1;
|
||||
const CTRL_ATTR_FAMILY_NAME: u16 = 2;
|
||||
const NETDEV_CMD_DEV_GET: u8 = 1;
|
||||
const NETDEV_A_DEV_IFINDEX: u16 = 1;
|
||||
const NETDEV_A_DEV_XDP_FEATURES: u16 = 3;
|
||||
const NETDEV_A_DEV_XDP_ZC_MAX_SEGS: u16 = 4;
|
||||
const NETDEV_A_DEV_XSK_FEATURES: u16 = 6;
|
||||
const NETDEV_XDP_ACT_XSK_ZEROCOPY: u64 = 8;
|
||||
const NLMSG_ERROR: u16 = 2;
|
||||
const NLM_F_REQUEST: u16 = 1;
|
||||
const NETLINK_SEQUENCE: u32 = 1;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct NetdevCapabilities {
|
||||
pub xdp_features: Option<u64>,
|
||||
pub xdp_zc_max_segs: Option<u32>,
|
||||
pub xsk_features: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct InterfaceCapabilities {
|
||||
pub driver: String,
|
||||
pub mtu: Option<u32>,
|
||||
pub rx_queues: Option<usize>,
|
||||
pub tx_queues: Option<usize>,
|
||||
pub netdev: Option<NetdevCapabilities>,
|
||||
}
|
||||
|
||||
impl InterfaceCapabilities {
|
||||
fn load(ifname: &str) -> Self {
|
||||
Self::load_from(ifname, Path::new(SYS_CLASS_NET))
|
||||
}
|
||||
|
||||
fn load_from(ifname: &str, sys_class_net: &Path) -> Self {
|
||||
let iface_path = sys_class_net.join(ifname);
|
||||
|
||||
Self {
|
||||
driver: interface_driver_from(&iface_path),
|
||||
mtu: read_u32(iface_path.join("mtu")),
|
||||
rx_queues: count_queue_dirs(&iface_path, "rx-"),
|
||||
tx_queues: count_queue_dirs(&iface_path, "tx-"),
|
||||
netdev: read_u32(iface_path.join("ifindex")).and_then(load_netdev_capabilities),
|
||||
}
|
||||
}
|
||||
|
||||
fn xsk_zerocopy_supported(&self) -> Option<bool> {
|
||||
self.netdev
|
||||
.as_ref()
|
||||
.and_then(|netdev| netdev.xdp_features)
|
||||
.map(|features| features & NETDEV_XDP_ACT_XSK_ZEROCOPY != 0)
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for InterfaceCapabilities {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
write!(f, "driver {}", self.driver)?;
|
||||
|
||||
if let Some(mtu) = self.mtu {
|
||||
write!(f, ", mtu {}", mtu)?;
|
||||
}
|
||||
if let Some(rx_queues) = self.rx_queues {
|
||||
write!(f, ", rx_queues {}", rx_queues)?;
|
||||
}
|
||||
if let Some(tx_queues) = self.tx_queues {
|
||||
write!(f, ", tx_queues {}", tx_queues)?;
|
||||
}
|
||||
match self.netdev.as_ref() {
|
||||
Some(netdev) => {
|
||||
if let Some(features) = netdev.xdp_features {
|
||||
write!(f, ", xdp_features 0x{features:x}")?;
|
||||
}
|
||||
if let Some(max_segs) = netdev.xdp_zc_max_segs {
|
||||
write!(f, ", xdp_zc_max_segs {}", max_segs)?;
|
||||
}
|
||||
if let Some(features) = netdev.xsk_features {
|
||||
write!(f, ", xsk_features 0x{features:x}")?;
|
||||
}
|
||||
}
|
||||
None => write!(f, ", netdev_features unavailable")?,
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
@ -48,24 +131,24 @@ pub fn classify(stage: EbpfFailStage, err: &Error, ifname: Option<&str>) -> Ebpf
|
||||
let raw = err.to_string();
|
||||
let category = categorize(&raw);
|
||||
let kernel = kernel_release();
|
||||
let driver = ifname.map(interface_driver);
|
||||
let capabilities = ifname.map(InterfaceCapabilities::load);
|
||||
|
||||
let mut reason = format!("stage={:?}: {}", stage, raw);
|
||||
reason.push_str(&format!(" (kernel {}", kernel));
|
||||
if let Some(iface) = ifname {
|
||||
reason.push_str(&format!(", interface {}", iface));
|
||||
if let Some(drv) = driver.as_deref() {
|
||||
reason.push_str(&format!(", driver {}", drv));
|
||||
if let Some(caps) = capabilities.as_ref() {
|
||||
reason.push_str(&format!(", {}", caps));
|
||||
}
|
||||
}
|
||||
reason.push(')');
|
||||
|
||||
// For the igb-before-6.17 case, augment the reason with a targeted hint.
|
||||
if matches!(category, EbpfFailCategory::AfXdpUnsupported)
|
||||
&& driver.as_deref() == Some("igb")
|
||||
&& !kernel_meets_igb_af_xdp(&kernel)
|
||||
{
|
||||
reason.push_str(". The igb driver supports AF_XDP only on kernel 6.17 or newer.");
|
||||
if matches!(category, EbpfFailCategory::AfXdpUnsupported) {
|
||||
if let Some(Some(false)) = capabilities.as_ref().map(InterfaceCapabilities::xsk_zerocopy_supported) {
|
||||
reason.push_str(". The interface reports no AF_XDP zero-copy capability in xdp_features.");
|
||||
} else if capabilities.as_ref().and_then(|caps| caps.netdev.as_ref()).is_none() {
|
||||
reason.push_str(". Kernel did not expose netdev XDP feature data for this interface.");
|
||||
}
|
||||
}
|
||||
|
||||
EbpfHealth::Unavailable {
|
||||
@ -118,41 +201,263 @@ fn categorize(raw: &str) -> EbpfFailCategory {
|
||||
EbpfFailCategory::Unknown
|
||||
}
|
||||
|
||||
/// Parse a kernel release string like "6.17.4-generic" and return true
|
||||
/// if it is >= 6.17. We only care about the first two numeric components.
|
||||
fn kernel_meets_igb_af_xdp(release: &str) -> bool {
|
||||
// Pull leading "MAJOR.MINOR" out of strings like "6.12.0-124.45.1.el10_1.x86_64".
|
||||
let mut parts = release.split(|c: char| !c.is_ascii_digit()).filter(|s| !s.is_empty());
|
||||
let Some(major_str) = parts.next() else {
|
||||
return false;
|
||||
fn interface_driver_from(iface_path: &Path) -> String {
|
||||
match fs::read_link(iface_path.join("device").join("driver")) {
|
||||
Ok(target) => basename_to_string(&target),
|
||||
Err(_) => "unknown".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
fn basename_to_string(path: &Path) -> String {
|
||||
path.file_name()
|
||||
.and_then(OsStr::to_str)
|
||||
.map(str::to_string)
|
||||
.unwrap_or_else(|| "unknown".to_string())
|
||||
}
|
||||
|
||||
fn read_u32(path: PathBuf) -> Option<u32> {
|
||||
fs::read_to_string(path).ok()?.trim().parse().ok()
|
||||
}
|
||||
|
||||
fn count_queue_dirs(iface_path: &Path, prefix: &str) -> Option<usize> {
|
||||
let queues_path = iface_path.join("queues");
|
||||
let entries = fs::read_dir(queues_path).ok()?;
|
||||
let count = entries
|
||||
.filter_map(Result::ok)
|
||||
.filter(|entry| entry.file_type().map(|ty| ty.is_dir()).unwrap_or(false))
|
||||
.filter(|entry| {
|
||||
entry
|
||||
.file_name()
|
||||
.to_str()
|
||||
.map(|name| name.starts_with(prefix))
|
||||
.unwrap_or(false)
|
||||
})
|
||||
.count();
|
||||
|
||||
Some(count)
|
||||
}
|
||||
|
||||
fn load_netdev_capabilities(ifindex: u32) -> Option<NetdevCapabilities> {
|
||||
let family_id = resolve_genl_family_id(NETDEV_FAMILY_NAME)?;
|
||||
let payload = genl_request(
|
||||
family_id,
|
||||
NETDEV_CMD_DEV_GET,
|
||||
&[netlink_attr_u32(NETDEV_A_DEV_IFINDEX, ifindex)],
|
||||
);
|
||||
let response = netlink_round_trip(&payload)?;
|
||||
let attrs = first_genl_attrs(&response)?;
|
||||
|
||||
Some(NetdevCapabilities {
|
||||
xdp_features: attr_u64(attrs, NETDEV_A_DEV_XDP_FEATURES),
|
||||
xdp_zc_max_segs: attr_u32(attrs, NETDEV_A_DEV_XDP_ZC_MAX_SEGS),
|
||||
xsk_features: attr_u64(attrs, NETDEV_A_DEV_XSK_FEATURES),
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_genl_family_id(name: &str) -> Option<u16> {
|
||||
let payload = genl_request(
|
||||
GENL_ID_CTRL,
|
||||
CTRL_CMD_GETFAMILY,
|
||||
&[netlink_attr_string(CTRL_ATTR_FAMILY_NAME, name)],
|
||||
);
|
||||
let response = netlink_round_trip(&payload)?;
|
||||
let attrs = first_genl_attrs(&response)?;
|
||||
attr_u16(attrs, CTRL_ATTR_FAMILY_ID)
|
||||
}
|
||||
|
||||
fn netlink_round_trip(payload: &[u8]) -> Option<Vec<u8>> {
|
||||
let fd = open_netlink_socket()?;
|
||||
let sent = send_netlink(fd, payload).is_some();
|
||||
let response = if sent { recv_netlink(fd) } else { None };
|
||||
close_fd(fd);
|
||||
response
|
||||
}
|
||||
|
||||
fn open_netlink_socket() -> Option<RawFd> {
|
||||
// SAFETY: socket and bind are called with a valid sockaddr_nl and length.
|
||||
let fd = unsafe { socket(AF_NETLINK, SOCK_RAW | SOCK_CLOEXEC, NETLINK_GENERIC) };
|
||||
if fd < 0 {
|
||||
return None;
|
||||
}
|
||||
|
||||
// SAFETY: zeroed sockaddr_nl is immediately initialized before bind.
|
||||
let mut addr: sockaddr_nl = unsafe { mem::zeroed() };
|
||||
addr.nl_family = AF_NETLINK as u16;
|
||||
|
||||
// SAFETY: fd is a netlink socket, addr points to initialized storage.
|
||||
let rc = unsafe {
|
||||
bind(
|
||||
fd,
|
||||
&addr as *const sockaddr_nl as *const sockaddr,
|
||||
mem::size_of::<sockaddr_nl>() as u32,
|
||||
)
|
||||
};
|
||||
let Some(minor_str) = parts.next() else {
|
||||
return false;
|
||||
if rc < 0 {
|
||||
close_fd(fd);
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(fd)
|
||||
}
|
||||
|
||||
fn close_fd(fd: RawFd) {
|
||||
// SAFETY: closing an owned file descriptor; failures are irrelevant here.
|
||||
unsafe {
|
||||
close(fd);
|
||||
}
|
||||
}
|
||||
|
||||
fn send_netlink(fd: RawFd, payload: &[u8]) -> Option<()> {
|
||||
// SAFETY: zeroed sockaddr_nl is immediately initialized before sendto.
|
||||
let mut kernel: sockaddr_nl = unsafe { mem::zeroed() };
|
||||
kernel.nl_family = AF_NETLINK as u16;
|
||||
|
||||
// SAFETY: payload is a valid byte slice and kernel points to initialized storage.
|
||||
let sent = unsafe {
|
||||
sendto(
|
||||
fd,
|
||||
payload.as_ptr().cast(),
|
||||
payload.len(),
|
||||
0,
|
||||
&kernel as *const sockaddr_nl as *const sockaddr,
|
||||
mem::size_of::<sockaddr_nl>() as u32,
|
||||
)
|
||||
};
|
||||
let Ok(major): Result<u32, _> = major_str.parse() else {
|
||||
return false;
|
||||
};
|
||||
let Ok(minor): Result<u32, _> = minor_str.parse() else {
|
||||
return false;
|
||||
};
|
||||
major > 6 || (major == 6 && minor >= 17)
|
||||
if sent == payload.len() as isize { Some(()) } else { None }
|
||||
}
|
||||
|
||||
fn recv_netlink(fd: RawFd) -> Option<Vec<u8>> {
|
||||
let mut buf = vec![0_u8; 8192];
|
||||
// SAFETY: buf is valid writable storage for recv.
|
||||
let received = unsafe { recv(fd, buf.as_mut_ptr().cast(), buf.len(), 0) };
|
||||
if received <= 0 {
|
||||
return None;
|
||||
}
|
||||
buf.truncate(received as usize);
|
||||
Some(buf)
|
||||
}
|
||||
|
||||
fn genl_request(nlmsg_type: u16, cmd: u8, attrs: &[Vec<u8>]) -> Vec<u8> {
|
||||
let header_len = align4(mem::size_of::<nlmsghdr>());
|
||||
let genl_len = mem::size_of::<genlmsghdr>();
|
||||
let mut buf = vec![0_u8; header_len + genl_len];
|
||||
|
||||
write_u32(&mut buf, 0, 0);
|
||||
write_u16(&mut buf, 4, nlmsg_type);
|
||||
write_u16(&mut buf, 6, NLM_F_REQUEST);
|
||||
write_u32(&mut buf, 8, NETLINK_SEQUENCE);
|
||||
write_u32(&mut buf, 12, 0);
|
||||
buf[header_len] = cmd;
|
||||
buf[header_len + 1] = 1;
|
||||
|
||||
for attr in attrs {
|
||||
buf.extend_from_slice(attr);
|
||||
}
|
||||
|
||||
let nlmsg_len = buf.len() as u32;
|
||||
write_u32(&mut buf, 0, nlmsg_len);
|
||||
buf
|
||||
}
|
||||
|
||||
fn netlink_attr_string(attr_type: u16, value: &str) -> Vec<u8> {
|
||||
let mut payload = value.as_bytes().to_vec();
|
||||
payload.push(0);
|
||||
netlink_attr(attr_type, &payload)
|
||||
}
|
||||
|
||||
fn netlink_attr_u32(attr_type: u16, value: u32) -> Vec<u8> {
|
||||
netlink_attr(attr_type, &value.to_ne_bytes())
|
||||
}
|
||||
|
||||
fn netlink_attr(attr_type: u16, payload: &[u8]) -> Vec<u8> {
|
||||
let len = 4 + payload.len();
|
||||
let mut attr = vec![0_u8; align4(len)];
|
||||
write_u16(&mut attr, 0, len as u16);
|
||||
write_u16(&mut attr, 2, attr_type);
|
||||
attr[4..4 + payload.len()].copy_from_slice(payload);
|
||||
attr
|
||||
}
|
||||
|
||||
fn first_genl_attrs(response: &[u8]) -> Option<&[u8]> {
|
||||
let header_len = align4(mem::size_of::<nlmsghdr>());
|
||||
let genl_len = mem::size_of::<genlmsghdr>();
|
||||
if response.len() < header_len + genl_len {
|
||||
return None;
|
||||
}
|
||||
|
||||
let nlmsg_len = read_u32_ne(response, 0)? as usize;
|
||||
let nlmsg_type = read_u16_ne(response, 4)?;
|
||||
if nlmsg_type == NLMSG_ERROR || nlmsg_len > response.len() || nlmsg_len < header_len + genl_len {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(&response[header_len + genl_len..nlmsg_len])
|
||||
}
|
||||
|
||||
fn attr_u16(attrs: &[u8], attr_type: u16) -> Option<u16> {
|
||||
find_attr(attrs, attr_type).and_then(|payload| read_u16_ne(payload, 0))
|
||||
}
|
||||
|
||||
fn attr_u32(attrs: &[u8], attr_type: u16) -> Option<u32> {
|
||||
find_attr(attrs, attr_type).and_then(|payload| read_u32_ne(payload, 0))
|
||||
}
|
||||
|
||||
fn attr_u64(attrs: &[u8], attr_type: u16) -> Option<u64> {
|
||||
find_attr(attrs, attr_type).and_then(|payload| read_u64_ne(payload, 0))
|
||||
}
|
||||
|
||||
fn find_attr(attrs: &[u8], attr_type: u16) -> Option<&[u8]> {
|
||||
let mut offset = 0;
|
||||
while offset + 4 <= attrs.len() {
|
||||
let len = read_u16_ne(attrs, offset)? as usize;
|
||||
let current_type = read_u16_ne(attrs, offset + 2)?;
|
||||
if len < 4 || offset + len > attrs.len() {
|
||||
return None;
|
||||
}
|
||||
|
||||
if current_type == attr_type {
|
||||
return Some(&attrs[offset + 4..offset + len]);
|
||||
}
|
||||
|
||||
offset += align4(len);
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
fn align4(value: usize) -> usize {
|
||||
(value + 3) & !3
|
||||
}
|
||||
|
||||
fn write_u16(buf: &mut [u8], offset: usize, value: u16) {
|
||||
buf[offset..offset + 2].copy_from_slice(&value.to_ne_bytes());
|
||||
}
|
||||
|
||||
fn write_u32(buf: &mut [u8], offset: usize, value: u32) {
|
||||
buf[offset..offset + 4].copy_from_slice(&value.to_ne_bytes());
|
||||
}
|
||||
|
||||
fn read_u16_ne(buf: &[u8], offset: usize) -> Option<u16> {
|
||||
let bytes = buf.get(offset..offset + 2)?;
|
||||
Some(u16::from_ne_bytes([bytes[0], bytes[1]]))
|
||||
}
|
||||
|
||||
fn read_u32_ne(buf: &[u8], offset: usize) -> Option<u32> {
|
||||
let bytes = buf.get(offset..offset + 4)?;
|
||||
Some(u32::from_ne_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]))
|
||||
}
|
||||
|
||||
fn read_u64_ne(buf: &[u8], offset: usize) -> Option<u64> {
|
||||
let bytes = buf.get(offset..offset + 8)?;
|
||||
Some(u64::from_ne_bytes([
|
||||
bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7],
|
||||
]))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn kernel_version_matrix() {
|
||||
assert!(kernel_meets_igb_af_xdp("6.17.0-generic"));
|
||||
assert!(kernel_meets_igb_af_xdp("6.18.1"));
|
||||
assert!(kernel_meets_igb_af_xdp("7.0.0"));
|
||||
assert!(!kernel_meets_igb_af_xdp("6.16.9-generic"));
|
||||
assert!(!kernel_meets_igb_af_xdp("6.12.0-124.45.1.el10_1.x86_64"));
|
||||
assert!(!kernel_meets_igb_af_xdp("5.15.0"));
|
||||
assert!(!kernel_meets_igb_af_xdp("nonsense"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn categorizes_permission_errors() {
|
||||
assert!(matches!(
|
||||
@ -176,4 +481,63 @@ mod tests {
|
||||
EbpfFailCategory::XdpUnsupported
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn loads_interface_capabilities_from_sysfs_shape() {
|
||||
let root = std::env::temp_dir().join(format!("netguardia-ebpf-preflight-{}", std::process::id()));
|
||||
let iface = root.join("eth0");
|
||||
let queues = iface.join("queues");
|
||||
let device = iface.join("device");
|
||||
let driver_target = root.join("drivers").join("virtio_net");
|
||||
|
||||
let _ = fs::remove_dir_all(&root);
|
||||
fs::create_dir_all(queues.join("rx-0")).unwrap();
|
||||
fs::create_dir_all(queues.join("rx-1")).unwrap();
|
||||
fs::create_dir_all(queues.join("tx-0")).unwrap();
|
||||
fs::create_dir_all(&driver_target).unwrap();
|
||||
fs::create_dir_all(&device).unwrap();
|
||||
fs::write(iface.join("mtu"), "1500\n").unwrap();
|
||||
fs::write(iface.join("ifindex"), "not-a-number\n").unwrap();
|
||||
std::os::unix::fs::symlink(&driver_target, device.join("driver")).unwrap();
|
||||
|
||||
let caps = InterfaceCapabilities::load_from("eth0", &root);
|
||||
|
||||
assert_eq!(caps.driver, "virtio_net");
|
||||
assert_eq!(caps.mtu, Some(1500));
|
||||
assert_eq!(caps.rx_queues, Some(2));
|
||||
assert_eq!(caps.tx_queues, Some(1));
|
||||
assert_eq!(caps.netdev, None);
|
||||
|
||||
fs::remove_dir_all(root).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_netlink_attributes() {
|
||||
let attrs = [
|
||||
netlink_attr_u32(NETDEV_A_DEV_IFINDEX, 7),
|
||||
netlink_attr(NETDEV_A_DEV_XDP_FEATURES, &8_u64.to_ne_bytes()),
|
||||
]
|
||||
.concat();
|
||||
|
||||
assert_eq!(attr_u32(&attrs, NETDEV_A_DEV_IFINDEX), Some(7));
|
||||
assert_eq!(attr_u64(&attrs, NETDEV_A_DEV_XDP_FEATURES), Some(8));
|
||||
assert_eq!(attr_u32(&attrs, NETDEV_A_DEV_XSK_FEATURES), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn interprets_xsk_zerocopy_from_netdev_features() {
|
||||
let caps = InterfaceCapabilities {
|
||||
driver: "virtio_net".to_string(),
|
||||
mtu: Some(1500),
|
||||
rx_queues: Some(1),
|
||||
tx_queues: Some(1),
|
||||
netdev: Some(NetdevCapabilities {
|
||||
xdp_features: Some(NETDEV_XDP_ACT_XSK_ZEROCOPY),
|
||||
xdp_zc_max_segs: Some(1),
|
||||
xsk_features: Some(0),
|
||||
}),
|
||||
};
|
||||
|
||||
assert_eq!(caps.xsk_zerocopy_supported(), Some(true));
|
||||
}
|
||||
}
|
||||
|
||||
@ -1,5 +1,5 @@
|
||||
use std::net::IpAddr;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::path::Path;
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
@ -17,14 +17,9 @@ pub struct GeoIpService {
|
||||
}
|
||||
|
||||
impl GeoIpService {
|
||||
pub fn new(db_name: &str) -> Result<Self, MaxMindDbError> {
|
||||
let db_path = PathBuf::from("net-guardia/static/geo").join(db_name);
|
||||
Self::with_cache_size(db_path, 10000)
|
||||
}
|
||||
|
||||
pub fn with_cache_size<P: AsRef<Path>>(db_path: P, cache_size: usize) -> Result<Self, MaxMindDbError> {
|
||||
let reader = Reader::open_readfile(db_path)?;
|
||||
let capacity = if cache_size == 0 { 10_000 } else { cache_size } as u64;
|
||||
let capacity = cache_size.max(1) as u64;
|
||||
|
||||
Ok(Self {
|
||||
reader: Arc::new(reader),
|
||||
|
||||
@ -36,8 +36,8 @@ pub struct SystemHealth {
|
||||
|
||||
impl SystemHealth {
|
||||
pub fn new(config: Arc<ArcSwap<AppConfig>>, ebpf_health: Arc<ArcSwap<EbpfHealth>>) -> Result<Self, Error> {
|
||||
let (broadcast_tx, _) = broadcast::channel(100);
|
||||
let cfg = config.load();
|
||||
let (broadcast_tx, _) = broadcast::channel(cfg.health.broadcast_channel_capacity.max(1));
|
||||
let ingress_interface = cfg.ebpf.ingress_ifname.clone();
|
||||
let egress_interface = cfg.ebpf.egress_ifname.clone();
|
||||
|
||||
|
||||
@ -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>))
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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;
|
||||
|
||||
21
net-guardia/src/infrastructure/runtime_state.rs
Normal file
21
net-guardia/src/infrastructure/runtime_state.rs
Normal 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(),
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
@ -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());
|
||||
}
|
||||
|
||||
@ -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() {
|
||||
|
||||
@ -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) {
|
||||
|
||||
@ -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>;
|
||||
}
|
||||
|
||||
@ -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>;
|
||||
}
|
||||
|
||||
@ -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
Loading…
x
Reference in New Issue
Block a user