cli: avoid data-dir initialization for version; create db parent dirs; redact generated passwords in CLI output

- Prevent 'nx9-wg version' from creating data directories by avoiding database initialization.
- Create parent directories when an explicit --database path is provided.
- Redact printed generated administrator passwords; announce file path or redact instead.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
thakaresandCopilot committed 2026-08-16 16:26:24 +05:30
commit 2ac6c81dfe
140 files changed
+31342

No files matched your search

+20
View File
@@ -0,0 +1,20 @@
[package]
name = "nx9-wg-db"
description = "SQLite persistence layer for nx9-wg"
version.workspace = true
edition.workspace = true
[dependencies]
nx9-wg-core.workspace = true
sqlx = { workspace = true, features = ["runtime-tokio", "sqlite", "macros", "migrate", "chrono", "uuid"] }
tokio.workspace = true
chrono.workspace = true
uuid.workspace = true
ipnet.workspace = true
thiserror.workspace = true
tracing.workspace = true
serde.workspace = true
serde_json.workspace = true
[dev-dependencies]
tempfile.workspace = true
+67
View File
@@ -0,0 +1,67 @@
# nx9-db — SQLite Persistence Layer
`nx9-db` provides the authoritative SQLite persistence layer for the `nx9-wg` native Rust WireGuard management system.
## Architectural Boundaries
- **Authoritative State**: SQLite is the authoritative persistent store for `nx9-wg` desired state. It stores what the system intends the network, interfaces, peers, routes, firewall rules, administrator credentials, sessions, tokens, and settings to be.
- **Separation of Concerns**: SQLite records desired configuration only. Live kernel state (WireGuard interface status, handshake counters, packet counters, live nftables rules, live kernel routes) is queried directly from Linux kernel subsystems in later phases.
- **SQL Encapsulation**: All SQL queries, SQLite connection lifecycle, migrations, and row conversions are strictly encapsulated inside `nx9-db`. Neither `nx9-core`, `nx9-api`, `nx9-ui`, `nx9-wireguard`, nor `nx9-network` issue SQL directly.
## SQLite Configuration
Every connection opened by `Store` enforces:
- `PRAGMA journal_mode = WAL` — Write-Ahead Logging for high-concurrency read/write operations.
- `PRAGMA foreign_keys = ON` — Strict relational integrity across all tables.
- `PRAGMA busy_timeout = 5000` — 5-second busy timeout to avoid contention errors.
- `PRAGMA synchronous = NORMAL` — Optimal reliability and performance in WAL mode.
## Database Schema (12 Tables)
1. `admin` — Single administrator identity (`CHECK (id = 1)`), Argon2id password hash, TOTP secrets, and login timestamp.
2. `sessions` — Admin web sessions (`ON DELETE CASCADE`).
3. `login_attempts` — IP-based login attempt tracking for brute-force rate limiting.
4. `api_tokens` — Hashed API tokens for automation (`ON DELETE CASCADE`).
5. `interfaces` — Desired WireGuard interfaces (`wg0`, `wg1`, etc.), private/public keys, listen port, IPv4/IPv6 CIDRs, MTU, DNS.
6. `peers` — Desired WireGuard peer definitions, classifications (`road_warrior`, `site_gateway`, `server`, `relay`), states (`active`, `disabled`, `revoked`, `expired`), profiles (`full_tunnel`, `split_tunnel`, `custom`), public/private/preshared keys, AllowedIPs, endpoints, and persistent keepalives (`ON DELETE CASCADE`).
7. `networks` — Named network CIDRs for routing and organization.
8. `routes` — Desired kernel routing rules (`ON DELETE SET NULL`).
9. `firewall_rules` — Desired firewall policy rules with priorities and directions (`in`, `out`, `forward`).
10. `settings` — Key-value system settings with secret redaction support.
11. `audit_events` — Append-only operational audit log with event filtering and pagination.
12. `backups` — Backup metadata and manifest checksum records.
## Migration Strategy
- Migrations are defined in `crates/nx9-db/migrations/` and embedded at compile time via `sqlx::migrate!("./migrations")`.
- Migrations are executed automatically via `store.migrate().await?`.
- Migrations are tracked in the `_sqlx_migrations` table for idempotency.
## Usage in Code
```rust
use nx9_db::Store;
use std::path::Path;
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
// Connect and auto-migrate
let store = Store::connect_path(Path::new("/var/lib/nx9-wg/nx9-wg.db")).await?;
store.migrate().await?;
// Create single admin if not initialized
if !store.admin_exists().await? {
store.create_admin("admin", "$argon2id$...").await?;
}
Ok(())
}
```
## Running Tests
Tests use isolated in-memory or temporary file SQLite instances:
```bash
cargo test -p nx9-db
```
@@ -0,0 +1,219 @@
------------------------------------------------------------------------
-- nx9-wg SQLite Initial Migration (0001_initial.sql)
------------------------------------------------------------------------
------------------------------------------------------------------------
-- 1. Admin (Exactly one row, id=1 enforced by CHECK)
------------------------------------------------------------------------
CREATE TABLE admin (
id INTEGER PRIMARY KEY CHECK (id = 1),
username TEXT NOT NULL UNIQUE,
password_hash TEXT NOT NULL,
totp_secret TEXT,
totp_enabled INTEGER NOT NULL DEFAULT 0,
last_login_at TEXT,
last_login_ip TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
);
CREATE INDEX idx_admin_username ON admin(username);
------------------------------------------------------------------------
-- 2. Sessions
------------------------------------------------------------------------
CREATE TABLE sessions (
id TEXT PRIMARY KEY,
admin_id INTEGER NOT NULL DEFAULT 1 REFERENCES admin(id) ON DELETE CASCADE,
ip_address TEXT,
user_agent TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
expires_at TEXT NOT NULL,
last_seen_at TEXT
);
CREATE INDEX idx_sessions_expires_at ON sessions(expires_at);
CREATE INDEX idx_sessions_admin_id ON sessions(admin_id);
------------------------------------------------------------------------
-- 3. Login Attempts (Brute force protection)
------------------------------------------------------------------------
CREATE TABLE login_attempts (
id INTEGER PRIMARY KEY AUTOINCREMENT,
ip_address TEXT NOT NULL,
attempted_at TEXT NOT NULL DEFAULT (datetime('now')),
success INTEGER NOT NULL DEFAULT 0
);
CREATE INDEX idx_login_attempts_ip ON login_attempts(ip_address, attempted_at);
------------------------------------------------------------------------
-- 4. API Tokens
------------------------------------------------------------------------
CREATE TABLE api_tokens (
id TEXT PRIMARY KEY,
admin_id INTEGER NOT NULL DEFAULT 1 REFERENCES admin(id) ON DELETE CASCADE,
name TEXT NOT NULL,
token_hash TEXT NOT NULL UNIQUE,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
expires_at TEXT,
last_used_at TEXT,
revoked_at TEXT
);
CREATE INDEX idx_api_tokens_token_hash ON api_tokens(token_hash);
CREATE INDEX idx_api_tokens_expires_at ON api_tokens(expires_at);
------------------------------------------------------------------------
-- 5. WireGuard Interfaces
------------------------------------------------------------------------
CREATE TABLE interfaces (
id TEXT PRIMARY KEY,
name TEXT NOT NULL UNIQUE,
private_key TEXT NOT NULL,
public_key TEXT NOT NULL,
listen_port INTEGER NOT NULL DEFAULT 51820,
ipv4_cidr TEXT NOT NULL,
ipv6_cidr TEXT,
mtu INTEGER,
dns TEXT,
enabled INTEGER NOT NULL DEFAULT 1,
pre_up TEXT,
post_up TEXT,
pre_down TEXT,
post_down TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
);
CREATE INDEX idx_interfaces_name ON interfaces(name);
------------------------------------------------------------------------
-- 6. Peers / Clients
------------------------------------------------------------------------
CREATE TABLE peers (
id TEXT PRIMARY KEY,
interface_id TEXT NOT NULL REFERENCES interfaces(id) ON DELETE CASCADE,
name TEXT NOT NULL,
peer_type TEXT NOT NULL DEFAULT 'road_warrior'
CHECK (peer_type IN ('road_warrior', 'site_gateway', 'server', 'relay')),
state TEXT NOT NULL DEFAULT 'active'
CHECK (state IN ('active', 'disabled', 'revoked', 'expired')),
profile TEXT NOT NULL DEFAULT 'full_tunnel'
CHECK (profile IN ('full_tunnel', 'split_tunnel', 'custom')),
public_key TEXT NOT NULL,
private_key TEXT,
preshared_key TEXT,
endpoint TEXT,
allowed_ips TEXT NOT NULL,
server_allowed_ips TEXT,
address_ipv4 TEXT,
address_ipv6 TEXT,
dns TEXT,
mtu INTEGER,
persistent_keepalive INTEGER,
expires_at TEXT,
last_handshake_at TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
updated_at TEXT NOT NULL DEFAULT (datetime('now')),
UNIQUE(interface_id, name),
UNIQUE(interface_id, public_key)
);
CREATE INDEX idx_peers_interface_id ON peers(interface_id);
CREATE INDEX idx_peers_name ON peers(interface_id, name);
CREATE INDEX idx_peers_state ON peers(state);
------------------------------------------------------------------------
-- 7. Networks
------------------------------------------------------------------------
CREATE TABLE networks (
id TEXT PRIMARY KEY,
name TEXT NOT NULL UNIQUE,
cidr TEXT NOT NULL,
enabled INTEGER NOT NULL DEFAULT 1,
description TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
);
CREATE INDEX idx_networks_name ON networks(name);
------------------------------------------------------------------------
-- 8. Routes
------------------------------------------------------------------------
CREATE TABLE routes (
id TEXT PRIMARY KEY,
network_id TEXT REFERENCES networks(id) ON DELETE SET NULL,
interface_id TEXT REFERENCES interfaces(id) ON DELETE SET NULL,
destination TEXT NOT NULL,
gateway TEXT,
metric INTEGER,
enabled INTEGER NOT NULL DEFAULT 1,
description TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
);
CREATE INDEX idx_routes_network_id ON routes(network_id);
CREATE INDEX idx_routes_interface_id ON routes(interface_id);
------------------------------------------------------------------------
-- 9. Firewall Rules
------------------------------------------------------------------------
CREATE TABLE firewall_rules (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
interface_id TEXT REFERENCES interfaces(id) ON DELETE SET NULL,
direction TEXT NOT NULL DEFAULT 'in'
CHECK (direction IN ('in', 'out', 'forward')),
action TEXT NOT NULL DEFAULT 'accept'
CHECK (action IN ('accept', 'drop', 'reject')),
protocol TEXT NOT NULL DEFAULT 'any'
CHECK (protocol IN ('tcp', 'udp', 'tcp_udp', 'icmp', 'any')),
source TEXT,
destination TEXT,
source_port INTEGER,
destination_port INTEGER,
priority INTEGER NOT NULL DEFAULT 100,
enabled INTEGER NOT NULL DEFAULT 1,
description TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
);
CREATE INDEX idx_firewall_interface_priority ON firewall_rules(interface_id, priority);
------------------------------------------------------------------------
-- 10. Settings (Key-Value)
------------------------------------------------------------------------
CREATE TABLE settings (
key TEXT PRIMARY KEY NOT NULL,
value TEXT NOT NULL,
is_secret INTEGER NOT NULL DEFAULT 0,
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
);
------------------------------------------------------------------------
-- 11. Audit Events (Append-only)
------------------------------------------------------------------------
CREATE TABLE audit_events (
id INTEGER PRIMARY KEY AUTOINCREMENT,
event_type TEXT NOT NULL,
actor TEXT NOT NULL DEFAULT 'admin',
resource_type TEXT,
resource_id TEXT,
message TEXT,
metadata TEXT,
ip_address TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now'))
);
CREATE INDEX idx_audit_created_at ON audit_events(created_at);
CREATE INDEX idx_audit_event_type ON audit_events(event_type);
CREATE INDEX idx_audit_resource ON audit_events(resource_type, resource_id);
------------------------------------------------------------------------
-- 12. Backups (Metadata)
------------------------------------------------------------------------
CREATE TABLE backups (
id TEXT PRIMARY KEY,
filename TEXT NOT NULL,
size INTEGER NOT NULL,
checksum TEXT NOT NULL,
encrypted INTEGER NOT NULL DEFAULT 0,
schema_version TEXT NOT NULL,
description TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now'))
);
CREATE INDEX idx_backups_created_at ON backups(created_at);
@@ -0,0 +1,7 @@
-- 0002_wiregui_capabilities.sql
-- Add peer-specific firewall association and structured port semantics
ALTER TABLE firewall_rules ADD COLUMN peer_id TEXT REFERENCES peers(id) ON DELETE CASCADE;
ALTER TABLE firewall_rules ADD COLUMN port_range TEXT;
CREATE INDEX IF NOT EXISTS idx_firewall_peer_id ON firewall_rules(peer_id);
@@ -0,0 +1,38 @@
-- 0003_client_profiles.sql
-- Client Environment and MTU Profile System
CREATE TABLE IF NOT EXISTS client_profiles (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
provider TEXT,
device TEXT CHECK (device IS NULL OR device IN ('android', 'ios', 'linux', 'windows', 'macos', 'other')),
connection_type TEXT NOT NULL CHECK (connection_type IN ('web', 'mobile', 'wifi', 'wired', 'other')),
nat_type TEXT NOT NULL DEFAULT 'unknown' CHECK (nat_type IN ('direct', 'cgnat', 'unknown')),
mtu INTEGER NOT NULL CHECK (mtu >= 1280 AND mtu <= 9000),
dns TEXT,
persistent_keepalive INTEGER CHECK (persistent_keepalive IS NULL OR (persistent_keepalive >= 0 AND persistent_keepalive <= 65535)),
is_builtin BOOLEAN NOT NULL DEFAULT 0,
description TEXT,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_client_profiles_provider ON client_profiles(provider);
CREATE INDEX IF NOT EXISTS idx_client_profiles_device ON client_profiles(device);
CREATE INDEX IF NOT EXISTS idx_client_profiles_connection ON client_profiles(connection_type);
CREATE INDEX IF NOT EXISTS idx_client_profiles_nat ON client_profiles(nat_type);
-- Insert authoritative built-in client profiles
INSERT OR IGNORE INTO client_profiles (id, name, provider, device, connection_type, nat_type, mtu, dns, persistent_keepalive, is_builtin, description, created_at, updated_at)
VALUES
('default-mobile', 'Default Mobile', NULL, NULL, 'mobile', 'unknown', 1280, NULL, 25, 1, 'Standard mobile carrier profile with 1280 MTU and 25s keepalive', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('default-cgnat', 'Default CGNAT', NULL, NULL, 'other', 'cgnat', 1360, NULL, 25, 1, 'Carrier-grade NAT environment profile with 1360 MTU and 25s keepalive', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('default-wifi', 'Default Wi-Fi', NULL, NULL, 'wifi', 'unknown', 1420, NULL, 25, 1, 'Standard Wi-Fi wireless profile with 1420 MTU and 25s keepalive', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('default-web', 'Default Web', NULL, NULL, 'web', 'unknown', 1420, NULL, 25, 1, 'Standard Web client profile with 1420 MTU and 25s keepalive', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('default-wired', 'Default Wired', NULL, NULL, 'wired', 'direct', 1420, NULL, 25, 1, 'High-throughput wired Ethernet profile with 1420 MTU and 25s keepalive', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('android-mobile', 'Android Mobile', NULL, 'android', 'mobile', 'unknown', 1280, NULL, 25, 1, 'Android cellular client profile with 1280 MTU and 25s keepalive', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('ios-mobile', 'iOS Mobile', NULL, 'ios', 'mobile', 'unknown', 1280, NULL, 25, 1, 'Apple iOS cellular profile with 1280 MTU and 25s keepalive', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('tmobile-mobile', 'T-Mobile Mobile', 'tmobile', NULL, 'mobile', 'cgnat', 1280, NULL, 25, 1, 'T-Mobile US IPv6/CGNAT mobile profile with 1280 MTU and 25s keepalive', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('verizon-mobile', 'Verizon Mobile', 'verizon', NULL, 'mobile', 'cgnat', 1280, NULL, 25, 1, 'Verizon Wireless mobile profile with 1280 MTU and 25s keepalive', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('jio-mobile', 'Jio Mobile', 'jio', NULL, 'mobile', 'cgnat', 1280, NULL, 25, 1, 'Reliance Jio 4G/5G mobile profile with 1280 MTU and 25s keepalive', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
('starlink-cgnat', 'Starlink CGNAT', 'starlink', NULL, 'other', 'cgnat', 1360, NULL, 25, 1, 'Starlink satellite CGNAT profile with 1360 MTU and 25s keepalive', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP);
+215
View File
@@ -0,0 +1,215 @@
//! Administrator repository operations.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::Utc;
use nx9_wg_core::types::auth::Admin;
use sqlx::{Row, SqlitePool};
/// Retrieve the single administrator record, if initialized.
pub async fn get_admin(pool: &SqlitePool) -> Result<Option<Admin>> {
let row = sqlx::query(
r#"
SELECT id, username, password_hash, totp_secret, totp_enabled,
last_login_at, last_login_ip, created_at, updated_at
FROM admin
WHERE id = 1
"#,
)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => {
let id: i64 = r.try_get("id")?;
let username: String = r.try_get("username")?;
let password_hash: String = r.try_get("password_hash")?;
let totp_secret: Option<String> = r.try_get("totp_secret")?;
let totp_enabled_int: i64 = r.try_get("totp_enabled")?;
let last_login_at_str: Option<String> = r.try_get("last_login_at")?;
let last_login_ip: Option<String> = r.try_get("last_login_ip")?;
let created_at_str: String = r.try_get("created_at")?;
let updated_at_str: String = r.try_get("updated_at")?;
let last_login_at = match last_login_at_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
Ok(Some(Admin {
id,
username,
password_hash,
totp_secret,
totp_enabled: totp_enabled_int != 0,
last_login_at,
last_login_ip,
created_at: parse_datetime(&created_at_str)?,
updated_at: parse_datetime(&updated_at_str)?,
}))
}
None => Ok(None),
}
}
/// Retrieve the administrator record by username.
pub async fn get_admin_by_username(pool: &SqlitePool, username: &str) -> Result<Option<Admin>> {
let admin = get_admin(pool).await?;
match admin {
Some(a) if a.username == username => Ok(Some(a)),
_ => Ok(None),
}
}
/// Check whether the single administrator has already been initialized.
pub async fn admin_exists(pool: &SqlitePool) -> Result<bool> {
let row = sqlx::query("SELECT COUNT(*) as count FROM admin WHERE id = 1")
.fetch_one(pool)
.await
.map_err(DbError::Sqlx)?;
let count: i64 = row.try_get("count")?;
Ok(count > 0)
}
/// Create the single administrator record.
///
/// Fails if an administrator already exists.
pub async fn create_admin(pool: &SqlitePool, username: &str, password_hash: &str) -> Result<Admin> {
if admin_exists(pool).await? {
return Err(DbError::Conflict(
"Administrator has already been initialized".to_string(),
));
}
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
sqlx::query(
r#"
INSERT INTO admin (id, username, password_hash, totp_secret, totp_enabled, created_at, updated_at)
VALUES (1, ?, ?, NULL, 0, ?, ?)
"#,
)
.bind(username)
.bind(password_hash)
.bind(&now_str)
.bind(&now_str)
.execute(pool)
.await
.map_err(|e| match &e {
sqlx::Error::Database(dbe) if dbe.is_unique_violation() => {
DbError::Conflict("Administrator already exists or username conflict".to_string())
}
_ => DbError::Sqlx(e),
})?;
Ok(Admin {
id: 1,
username: username.to_string(),
password_hash: password_hash.to_string(),
totp_secret: None,
totp_enabled: false,
last_login_at: None,
last_login_ip: None,
created_at: now,
updated_at: now,
})
}
/// Update the administrator's password hash.
pub async fn update_admin_password(pool: &SqlitePool, new_password_hash: &str) -> Result<()> {
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE admin
SET password_hash = ?, updated_at = ?
WHERE id = 1
"#,
)
.bind(new_password_hash)
.bind(&now_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(
"Administrator record does not exist".to_string(),
));
}
Ok(())
}
/// Update administrator TOTP configuration.
pub async fn update_admin_totp(
pool: &SqlitePool,
totp_secret: Option<&str>,
totp_enabled: bool,
) -> Result<()> {
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE admin
SET totp_secret = ?, totp_enabled = ?, updated_at = ?
WHERE id = 1
"#,
)
.bind(totp_secret)
.bind(if totp_enabled { 1 } else { 0 })
.bind(&now_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(
"Administrator record does not exist".to_string(),
));
}
Ok(())
}
/// Record a successful administrator login timestamp and IP address.
pub async fn record_admin_login(pool: &SqlitePool, ip_address: Option<&str>) -> Result<()> {
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE admin
SET last_login_at = ?, last_login_ip = ?, updated_at = ?
WHERE id = 1
"#,
)
.bind(&now_str)
.bind(ip_address)
.bind(&now_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(
"Administrator record does not exist".to_string(),
));
}
Ok(())
}
/// Delete administrator record (if explicitly supported).
pub async fn delete_admin(pool: &SqlitePool) -> Result<()> {
sqlx::query("DELETE FROM admin WHERE id = 1")
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(())
}
+213
View File
@@ -0,0 +1,213 @@
//! Operational Audit log repository operations (append-only).
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::{NaiveDateTime, Utc};
use nx9_wg_core::types::audit::{AuditEvent, AuditEventType};
use sqlx::{Row, SqlitePool};
use std::str::FromStr;
/// Filter options for querying audit records.
#[derive(Debug, Default, Clone)]
pub struct AuditFilter {
pub event_type: Option<AuditEventType>,
pub resource_type: Option<String>,
pub resource_id: Option<String>,
pub since: Option<NaiveDateTime>,
pub until: Option<NaiveDateTime>,
}
/// Append a new audit event to the log.
pub async fn create_audit_event(pool: &SqlitePool, event: &AuditEvent) -> Result<i64> {
let created_at_str = format_datetime(&event.created_at);
let result = sqlx::query(
r#"
INSERT INTO audit_events (
event_type, actor, resource_type, resource_id,
message, metadata, ip_address, created_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(event.event_type.as_str())
.bind(&event.actor)
.bind(&event.resource_type)
.bind(&event.resource_id)
.bind(&event.message)
.bind(&event.metadata)
.bind(&event.ip_address)
.bind(&created_at_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(result.last_insert_rowid())
}
/// Convenience function to record an audit entry.
#[allow(clippy::too_many_arguments)]
pub async fn record_audit(
pool: &SqlitePool,
event_type: AuditEventType,
actor: &str,
resource_type: Option<&str>,
resource_id: Option<&str>,
message: Option<&str>,
metadata: Option<&str>,
ip_address: Option<&str>,
) -> Result<i64> {
let now = Utc::now().naive_utc();
let event = AuditEvent {
id: 0,
event_type,
actor: actor.to_string(),
resource_type: resource_type.map(|s| s.to_string()),
resource_id: resource_id.map(|s| s.to_string()),
message: message.map(|s| s.to_string()),
metadata: metadata.map(|s| s.to_string()),
ip_address: ip_address.map(|s| s.to_string()),
created_at: now,
};
create_audit_event(pool, &event).await
}
/// Query audit events with filtering and pagination.
pub async fn list_audit_events(
pool: &SqlitePool,
filter: &AuditFilter,
limit: u32,
offset: u32,
) -> Result<Vec<AuditEvent>> {
let event_type_str = filter.event_type.map(|et| et.as_str().to_string());
let since_str = filter.since.as_ref().map(format_datetime);
let until_str = filter.until.as_ref().map(format_datetime);
let rows = sqlx::query(
r#"
SELECT id, event_type, actor, resource_type, resource_id,
message, metadata, ip_address, created_at
FROM audit_events
WHERE (?1 IS NULL OR event_type = ?1)
AND (?2 IS NULL OR resource_type = ?2)
AND (?3 IS NULL OR resource_id = ?3)
AND (?4 IS NULL OR created_at >= ?4)
AND (?5 IS NULL OR created_at <= ?5)
ORDER BY id DESC
LIMIT ?6 OFFSET ?7
"#,
)
.bind(event_type_str)
.bind(&filter.resource_type)
.bind(&filter.resource_id)
.bind(since_str)
.bind(until_str)
.bind(limit as i64)
.bind(offset as i64)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut events = Vec::with_capacity(rows.len());
for r in rows {
let id: i64 = r.try_get("id")?;
let event_type_str: String = r.try_get("event_type")?;
let actor: String = r.try_get("actor")?;
let resource_type: Option<String> = r.try_get("resource_type")?;
let resource_id: Option<String> = r.try_get("resource_id")?;
let message: Option<String> = r.try_get("message")?;
let metadata: Option<String> = r.try_get("metadata")?;
let ip_address: Option<String> = r.try_get("ip_address")?;
let created_at_str: String = r.try_get("created_at")?;
let event_type = AuditEventType::from_str(&event_type_str)?;
events.push(AuditEvent {
id,
event_type,
actor,
resource_type,
resource_id,
message,
metadata,
ip_address,
created_at: parse_datetime(&created_at_str)?,
});
}
Ok(events)
}
/// Retrieve a single audit event by ID.
pub async fn get_audit_event(pool: &SqlitePool, id: i64) -> Result<Option<AuditEvent>> {
let row = sqlx::query(
r#"
SELECT id, event_type, actor, resource_type, resource_id,
message, metadata, ip_address, created_at
FROM audit_events
WHERE id = ?
"#,
)
.bind(id)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => {
let event_type_str: String = r.try_get("event_type")?;
let actor: String = r.try_get("actor")?;
let resource_type: Option<String> = r.try_get("resource_type")?;
let resource_id: Option<String> = r.try_get("resource_id")?;
let message: Option<String> = r.try_get("message")?;
let metadata: Option<String> = r.try_get("metadata")?;
let ip_address: Option<String> = r.try_get("ip_address")?;
let created_at_str: String = r.try_get("created_at")?;
let event_type = AuditEventType::from_str(&event_type_str)?;
Ok(Some(AuditEvent {
id,
event_type,
actor,
resource_type,
resource_id,
message,
metadata,
ip_address,
created_at: parse_datetime(&created_at_str)?,
}))
}
None => Ok(None),
}
}
/// Count total audit events matching a filter.
pub async fn count_audit_events(pool: &SqlitePool, filter: &AuditFilter) -> Result<i64> {
let event_type_str = filter.event_type.map(|et| et.as_str().to_string());
let since_str = filter.since.as_ref().map(format_datetime);
let until_str = filter.until.as_ref().map(format_datetime);
let row = sqlx::query(
r#"
SELECT COUNT(*) as count
FROM audit_events
WHERE (?1 IS NULL OR event_type = ?1)
AND (?2 IS NULL OR resource_type = ?2)
AND (?3 IS NULL OR resource_id = ?3)
AND (?4 IS NULL OR created_at >= ?4)
AND (?5 IS NULL OR created_at <= ?5)
"#,
)
.bind(event_type_str)
.bind(&filter.resource_type)
.bind(&filter.resource_id)
.bind(since_str)
.bind(until_str)
.fetch_one(pool)
.await
.map_err(DbError::Sqlx)?;
let count: i64 = row.try_get("count")?;
Ok(count)
}
+129
View File
@@ -0,0 +1,129 @@
//! Backup metadata repository operations.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use nx9_wg_core::types::backup::BackupMeta;
use sqlx::{Row, SqlitePool};
use uuid::Uuid;
/// Helper to convert a database row into a `BackupMeta` domain struct.
fn row_to_backup_meta(r: &sqlx::sqlite::SqliteRow) -> Result<BackupMeta> {
let id_str: String = r.try_get("id")?;
let filename: String = r.try_get("filename")?;
let size_i64: i64 = r.try_get("size")?;
let checksum: String = r.try_get("checksum")?;
let encrypted_i64: i64 = r.try_get("encrypted")?;
let schema_version: String = r.try_get("schema_version")?;
let description: Option<String> = r.try_get("description")?;
let created_at_str: String = r.try_get("created_at")?;
let id = Uuid::parse_str(&id_str)
.map_err(|e| DbError::Validation(format!("invalid backup UUID '{id_str}': {e}")))?;
Ok(BackupMeta {
id,
filename,
size_bytes: size_i64,
checksum,
schema_version,
encrypted: encrypted_i64 != 0,
description,
created_at: parse_datetime(&created_at_str)?,
})
}
/// Record metadata for a new backup file.
pub async fn create_backup_meta(pool: &SqlitePool, meta: &BackupMeta) -> Result<()> {
let id_str = meta.id.to_string();
let created_at_str = format_datetime(&meta.created_at);
sqlx::query(
r#"
INSERT INTO backups (id, filename, size, checksum, encrypted, schema_version, description, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id_str)
.bind(&meta.filename)
.bind(meta.size_bytes)
.bind(&meta.checksum)
.bind(if meta.encrypted { 1 } else { 0 })
.bind(&meta.schema_version)
.bind(&meta.description)
.bind(&created_at_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(())
}
/// Retrieve backup metadata by UUID.
pub async fn get_backup_meta(pool: &SqlitePool, id: Uuid) -> Result<Option<BackupMeta>> {
let id_str = id.to_string();
let row = sqlx::query("SELECT * FROM backups WHERE id = ?")
.bind(&id_str)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => Ok(Some(row_to_backup_meta(&r)?)),
None => Ok(None),
}
}
/// List all backup records ordered by creation date descending.
pub async fn list_backups(pool: &SqlitePool) -> Result<Vec<BackupMeta>> {
let rows = sqlx::query("SELECT * FROM backups ORDER BY created_at DESC")
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_backup_meta(&r)?);
}
Ok(list)
}
/// Delete a backup record by UUID.
pub async fn delete_backup_meta(pool: &SqlitePool, id: Uuid) -> Result<()> {
let id_str = id.to_string();
let result = sqlx::query("DELETE FROM backups WHERE id = ?")
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!(
"Backup record '{id_str}' not found"
)));
}
Ok(())
}
/// Create a consistent, atomic file snapshot of the database using SQLite VACUUM INTO.
pub async fn vacuum_into(pool: &SqlitePool, target_file_path: &str) -> Result<()> {
// Check if target file already exists, remove it if so since VACUUM INTO fails if target exists
let path = std::path::Path::new(target_file_path);
if path.exists() {
let _ = std::fs::remove_file(path);
}
if let Some(parent) = path.parent().filter(|p| !p.exists()) {
std::fs::create_dir_all(parent)
.map_err(|e| DbError::Internal(format!("Failed to create backup directory: {e}")))?;
}
// SQLite VACUUM INTO requires a string literal filename
let escaped_path = target_file_path.replace('\'', "''");
let query_str = format!("VACUUM INTO '{escaped_path}'");
sqlx::query(&query_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(())
}
+254
View File
@@ -0,0 +1,254 @@
//! Client Profile repository operations.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::Utc;
use nx9_wg_core::types::client_profile::{ClientProfile, ConnectionType, DeviceCategory, NatType};
use sqlx::{Row, SqlitePool};
use std::str::FromStr;
/// Helper to convert a database row into a `ClientProfile` domain struct.
fn row_to_profile(r: &sqlx::sqlite::SqliteRow) -> Result<ClientProfile> {
let id: String = r.try_get("id")?;
let name: String = r.try_get("name")?;
let provider: Option<String> = r.try_get("provider")?;
let device_str: Option<String> = r.try_get("device")?;
let connection_type_str: String = r.try_get("connection_type")?;
let nat_type_str: String = r.try_get("nat_type")?;
let mtu_i64: i64 = r.try_get("mtu")?;
let dns: Option<String> = r.try_get("dns")?;
let keepalive_i64: Option<i64> = r.try_get("persistent_keepalive")?;
let is_builtin_i64: i64 = r.try_get("is_builtin")?;
let description: Option<String> = r.try_get("description")?;
let created_at_str: String = r.try_get("created_at")?;
let updated_at_str: String = r.try_get("updated_at")?;
let device = match device_str {
Some(s) if !s.trim().is_empty() => Some(
DeviceCategory::from_str(&s)
.map_err(|e| DbError::Validation(format!("invalid device category '{s}': {e}")))?,
),
_ => None,
};
let connection_type = ConnectionType::from_str(&connection_type_str).map_err(|e| {
DbError::Validation(format!(
"invalid connection type '{connection_type_str}': {e}"
))
})?;
let nat_type = NatType::from_str(&nat_type_str)
.map_err(|e| DbError::Validation(format!("invalid nat type '{nat_type_str}': {e}")))?;
Ok(ClientProfile {
id,
name,
provider,
device,
connection_type,
nat_type,
mtu: mtu_i64 as u16,
dns,
persistent_keepalive: keepalive_i64.map(|k| k as u16),
is_builtin: is_builtin_i64 != 0,
description,
created_at: parse_datetime(&created_at_str)?,
updated_at: parse_datetime(&updated_at_str)?,
})
}
/// Create a new client profile.
pub async fn create_client_profile(pool: &SqlitePool, profile: &ClientProfile) -> Result<()> {
let now = format_datetime(&Utc::now().naive_utc());
let device_str = profile.device.map(|d| d.as_str().to_string());
sqlx::query(
r#"
INSERT INTO client_profiles (
id, name, provider, device, connection_type, nat_type,
mtu, dns, persistent_keepalive, is_builtin, description,
created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&profile.id)
.bind(&profile.name)
.bind(&profile.provider)
.bind(&device_str)
.bind(profile.connection_type.as_str())
.bind(profile.nat_type.as_str())
.bind(profile.mtu as i64)
.bind(&profile.dns)
.bind(profile.persistent_keepalive.map(|k| k as i64))
.bind(if profile.is_builtin { 1i64 } else { 0i64 })
.bind(&profile.description)
.bind(&now)
.bind(&now)
.execute(pool)
.await
.map_err(|e| match e {
sqlx::Error::Database(ref db_err) if db_err.is_unique_violation() => {
DbError::Conflict(format!("client profile '{}' already exists", profile.id))
}
other => DbError::Sqlx(other),
})?;
Ok(())
}
/// Fetch a client profile by ID.
pub async fn get_client_profile(pool: &SqlitePool, id: &str) -> Result<Option<ClientProfile>> {
let row = sqlx::query("SELECT * FROM client_profiles WHERE id = ?")
.bind(id)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
row.map(|r| row_to_profile(&r)).transpose()
}
/// List all client profiles ordered by built-in status (built-in first) then name.
pub async fn list_client_profiles(pool: &SqlitePool) -> Result<Vec<ClientProfile>> {
let rows = sqlx::query("SELECT * FROM client_profiles ORDER BY is_builtin DESC, name ASC")
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
rows.iter().map(row_to_profile).collect()
}
/// Update a custom client profile. Built-in profiles cannot be modified.
pub async fn update_client_profile(pool: &SqlitePool, profile: &ClientProfile) -> Result<()> {
let existing = get_client_profile(pool, &profile.id)
.await?
.ok_or_else(|| DbError::NotFound(format!("client profile '{}' not found", profile.id)))?;
if existing.is_builtin {
return Err(DbError::Validation(format!(
"built-in client profile '{}' cannot be modified",
profile.id
)));
}
let now = format_datetime(&Utc::now().naive_utc());
let device_str = profile.device.map(|d| d.as_str().to_string());
let result = sqlx::query(
r#"
UPDATE client_profiles SET
name = ?,
provider = ?,
device = ?,
connection_type = ?,
nat_type = ?,
mtu = ?,
dns = ?,
persistent_keepalive = ?,
description = ?,
updated_at = ?
WHERE id = ? AND is_builtin = 0
"#,
)
.bind(&profile.name)
.bind(&profile.provider)
.bind(&device_str)
.bind(profile.connection_type.as_str())
.bind(profile.nat_type.as_str())
.bind(profile.mtu as i64)
.bind(&profile.dns)
.bind(profile.persistent_keepalive.map(|k| k as i64))
.bind(&profile.description)
.bind(&now)
.bind(&profile.id)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!(
"client profile '{}' not found or is built-in",
profile.id
)));
}
Ok(())
}
/// Delete a custom client profile. Built-in profiles cannot be deleted.
pub async fn delete_client_profile(pool: &SqlitePool, id: &str) -> Result<()> {
let existing = get_client_profile(pool, id)
.await?
.ok_or_else(|| DbError::NotFound(format!("client profile '{id}' not found")))?;
if existing.is_builtin {
return Err(DbError::Validation(format!(
"built-in client profile '{id}' cannot be deleted"
)));
}
let result = sqlx::query("DELETE FROM client_profiles WHERE id = ? AND is_builtin = 0")
.bind(id)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!(
"client profile '{id}' not found or is built-in"
)));
}
Ok(())
}
/// List all distinct non-null provider identifiers.
pub async fn list_distinct_providers(pool: &SqlitePool) -> Result<Vec<String>> {
let rows = sqlx::query(
"SELECT DISTINCT provider FROM client_profiles WHERE provider IS NOT NULL AND provider != '' ORDER BY provider ASC",
)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut providers = Vec::new();
for r in rows {
let p: Option<String> = r.try_get("provider")?;
if let Some(name) = p.filter(|s| !s.trim().is_empty() && !providers.contains(s)) {
providers.push(name);
}
}
Ok(providers)
}
/// Find matching profiles from database given criteria.
pub async fn find_matching_profiles(
pool: &SqlitePool,
provider: Option<&str>,
device: Option<DeviceCategory>,
connection: Option<ConnectionType>,
nat: Option<NatType>,
) -> Result<Vec<ClientProfile>> {
let all = list_client_profiles(pool).await?;
let mut filtered = Vec::new();
for p in all {
if let Some(req_p) = provider {
match p.provider {
Some(ref prof_p) if prof_p.eq_ignore_ascii_case(req_p) => {}
_ => continue,
}
}
if device.is_some_and(|req_d| p.device.is_some_and(|d| d != req_d)) {
continue;
}
if connection.is_some_and(|req_c| p.connection_type != req_c) {
continue;
}
if nat.is_some_and(|req_n| p.nat_type != NatType::Unknown && p.nat_type != req_n) {
continue;
}
filtered.push(p);
}
Ok(filtered)
}
+44
View File
@@ -0,0 +1,44 @@
//! Database error types.
use thiserror::Error;
/// Result type for database operations.
pub type Result<T> = std::result::Result<T, DbError>;
/// Database-specific errors.
#[derive(Debug, Error)]
pub enum DbError {
/// Entity was not found.
#[error("entity not found: {0}")]
NotFound(String),
/// Unique or foreign key constraint violation.
#[error("constraint violation: {0}")]
ConstraintViolation(String),
/// Conflict, e.g. entity already exists.
#[error("conflict: {0}")]
Conflict(String),
/// Validation error when converting from raw database values.
#[error("validation error: {0}")]
Validation(String),
/// SQLx database error.
#[error("database error: {0}")]
Sqlx(#[from] sqlx::Error),
/// Migration failure.
#[error("migration error: {0}")]
Migration(String),
/// Internal or unexpected error.
#[error("internal database error: {0}")]
Internal(String),
}
impl From<nx9_wg_core::error::Nx9Error> for DbError {
fn from(err: nx9_wg_core::error::Nx9Error) -> Self {
Self::Validation(err.to_string())
}
}
+278
View File
@@ -0,0 +1,278 @@
//! Firewall Rule repository operations.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::Utc;
use nx9_wg_core::types::firewall::{
FirewallAction, FirewallDirection, FirewallProtocol, FirewallRule,
};
use sqlx::{Row, SqlitePool};
use std::str::FromStr;
use uuid::Uuid;
/// Helper to convert a database row into a `FirewallRule` domain struct.
fn row_to_rule(r: &sqlx::sqlite::SqliteRow) -> Result<FirewallRule> {
let id_str: String = r.try_get("id")?;
let name: String = r.try_get("name")?;
let interface_id_str: Option<String> = r.try_get("interface_id")?;
let peer_id_str: Option<String> = r.try_get("peer_id").unwrap_or(None);
let direction_str: String = r.try_get("direction")?;
let action_str: String = r.try_get("action")?;
let protocol_str: String = r.try_get("protocol")?;
let source: Option<String> = r.try_get("source")?;
let destination: Option<String> = r.try_get("destination")?;
let source_port_i64: Option<i64> = r.try_get("source_port")?;
let destination_port_i64: Option<i64> = r.try_get("destination_port")?;
let port_range: Option<String> = r.try_get("port_range").unwrap_or(None);
let priority_i64: i64 = r.try_get("priority")?;
let enabled_i64: i64 = r.try_get("enabled")?;
let description: Option<String> = r.try_get("description")?;
let created_at_str: String = r.try_get("created_at")?;
let updated_at_str: String = r.try_get("updated_at")?;
let id = Uuid::parse_str(&id_str)
.map_err(|e| DbError::Validation(format!("invalid firewall rule UUID '{id_str}': {e}")))?;
let interface_id = match interface_id_str {
Some(s) => Some(
Uuid::parse_str(&s)
.map_err(|e| DbError::Validation(format!("invalid interface UUID '{s}': {e}")))?,
),
None => None,
};
let peer_id = match peer_id_str {
Some(s) => Some(
Uuid::parse_str(&s)
.map_err(|e| DbError::Validation(format!("invalid peer UUID '{s}': {e}")))?,
),
None => None,
};
let direction = FirewallDirection::from_str(&direction_str)?;
let action = FirewallAction::from_str(&action_str)?;
let protocol = FirewallProtocol::from_str(&protocol_str)?;
Ok(FirewallRule {
id,
name,
interface_id,
peer_id,
direction,
action,
protocol,
source,
destination,
source_port: source_port_i64.map(|p| p as u16),
destination_port: destination_port_i64.map(|p| p as u16),
port_range,
priority: priority_i64 as i32,
enabled: enabled_i64 != 0,
description,
created_at: parse_datetime(&created_at_str)?,
updated_at: parse_datetime(&updated_at_str)?,
})
}
/// Create a new firewall rule record.
pub async fn create_rule(pool: &SqlitePool, rule: &FirewallRule) -> Result<()> {
let id_str = rule.id.to_string();
let interface_id_str = rule.interface_id.map(|id| id.to_string());
let peer_id_str = rule.peer_id.map(|id| id.to_string());
let created_at_str = format_datetime(&rule.created_at);
let updated_at_str = format_datetime(&rule.updated_at);
sqlx::query(
r#"
INSERT INTO firewall_rules (
id, name, interface_id, peer_id, direction, action, protocol,
source, destination, source_port, destination_port, port_range,
priority, enabled, description, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id_str)
.bind(&rule.name)
.bind(interface_id_str)
.bind(peer_id_str)
.bind(rule.direction.as_str())
.bind(rule.action.as_str())
.bind(rule.protocol.as_str())
.bind(&rule.source)
.bind(&rule.destination)
.bind(rule.source_port.map(|p| p as i64))
.bind(rule.destination_port.map(|p| p as i64))
.bind(&rule.port_range)
.bind(rule.priority as i64)
.bind(if rule.enabled { 1 } else { 0 })
.bind(&rule.description)
.bind(&created_at_str)
.bind(&updated_at_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(())
}
/// Retrieve a firewall rule by UUID.
pub async fn get_rule(pool: &SqlitePool, id: Uuid) -> Result<Option<FirewallRule>> {
let id_str = id.to_string();
let row = sqlx::query("SELECT * FROM firewall_rules WHERE id = ?")
.bind(&id_str)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => Ok(Some(row_to_rule(&r)?)),
None => Ok(None),
}
}
/// List all firewall rules ordered by priority ascending.
pub async fn list_rules(pool: &SqlitePool) -> Result<Vec<FirewallRule>> {
let rows = sqlx::query("SELECT * FROM firewall_rules ORDER BY priority ASC, name ASC")
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_rule(&r)?);
}
Ok(list)
}
/// List firewall rules for a given interface ordered by priority ascending.
pub async fn list_rules_for_interface(
pool: &SqlitePool,
interface_id: Uuid,
) -> Result<Vec<FirewallRule>> {
let iface_id_str = interface_id.to_string();
let rows = sqlx::query(
"SELECT * FROM firewall_rules WHERE interface_id = ? ORDER BY priority ASC, name ASC",
)
.bind(&iface_id_str)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_rule(&r)?);
}
Ok(list)
}
/// List firewall rules for a given peer ordered by priority ascending.
pub async fn list_rules_for_peer(pool: &SqlitePool, peer_id: Uuid) -> Result<Vec<FirewallRule>> {
let peer_id_str = peer_id.to_string();
let rows = sqlx::query(
"SELECT * FROM firewall_rules WHERE peer_id = ? ORDER BY priority ASC, name ASC",
)
.bind(&peer_id_str)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_rule(&r)?);
}
Ok(list)
}
/// Update a firewall rule record.
pub async fn update_rule(pool: &SqlitePool, rule: &FirewallRule) -> Result<()> {
let id_str = rule.id.to_string();
let interface_id_str = rule.interface_id.map(|id| id.to_string());
let peer_id_str = rule.peer_id.map(|id| id.to_string());
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE firewall_rules
SET name = ?, interface_id = ?, peer_id = ?, direction = ?, action = ?, protocol = ?,
source = ?, destination = ?, source_port = ?, destination_port = ?, port_range = ?,
priority = ?, enabled = ?, description = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(&rule.name)
.bind(interface_id_str)
.bind(peer_id_str)
.bind(rule.direction.as_str())
.bind(rule.action.as_str())
.bind(rule.protocol.as_str())
.bind(&rule.source)
.bind(&rule.destination)
.bind(rule.source_port.map(|p| p as i64))
.bind(rule.destination_port.map(|p| p as i64))
.bind(&rule.port_range)
.bind(rule.priority as i64)
.bind(if rule.enabled { 1 } else { 0 })
.bind(&rule.description)
.bind(&now_str)
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!(
"Firewall rule '{id_str}' not found"
)));
}
Ok(())
}
/// Delete a firewall rule by UUID.
pub async fn delete_rule(pool: &SqlitePool, id: Uuid) -> Result<()> {
let id_str = id.to_string();
let result = sqlx::query("DELETE FROM firewall_rules WHERE id = ?")
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!(
"Firewall rule '{id_str}' not found"
)));
}
Ok(())
}
/// Enable or disable a firewall rule.
pub async fn set_rule_enabled(pool: &SqlitePool, id: Uuid, enabled: bool) -> Result<()> {
let id_str = id.to_string();
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE firewall_rules
SET enabled = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(if enabled { 1 } else { 0 })
.bind(&now_str)
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!(
"Firewall rule '{id_str}' not found"
)));
}
Ok(())
}
+238
View File
@@ -0,0 +1,238 @@
//! WireGuard Interface repository operations.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::Utc;
use ipnet::IpNet;
use nx9_wg_core::types::wireguard::{Interface, WireGuardPrivateKey, WireGuardPublicKey};
use sqlx::{Row, SqlitePool};
use std::str::FromStr;
use uuid::Uuid;
/// Helper to convert a database row into an `Interface` domain struct.
fn row_to_interface(r: &sqlx::sqlite::SqliteRow) -> Result<Interface> {
let id_str: String = r.try_get("id")?;
let name: String = r.try_get("name")?;
let private_key_str: String = r.try_get("private_key")?;
let public_key_str: String = r.try_get("public_key")?;
let listen_port_i64: i64 = r.try_get("listen_port")?;
let ipv4_cidr_str: String = r.try_get("ipv4_cidr")?;
let ipv6_cidr_str: Option<String> = r.try_get("ipv6_cidr")?;
let mtu_i64: Option<i64> = r.try_get("mtu")?;
let dns: Option<String> = r.try_get("dns")?;
let enabled_i64: i64 = r.try_get("enabled")?;
let pre_up: Option<String> = r.try_get("pre_up")?;
let post_up: Option<String> = r.try_get("post_up")?;
let pre_down: Option<String> = r.try_get("pre_down")?;
let post_down: Option<String> = r.try_get("post_down")?;
let created_at_str: String = r.try_get("created_at")?;
let updated_at_str: String = r.try_get("updated_at")?;
let id = Uuid::parse_str(&id_str)
.map_err(|e| DbError::Validation(format!("invalid interface UUID '{id_str}': {e}")))?;
let address_v4 = IpNet::from_str(&ipv4_cidr_str)
.map_err(|e| DbError::Validation(format!("invalid ipv4_cidr '{ipv4_cidr_str}': {e}")))?;
let address_v6 = match ipv6_cidr_str {
Some(s) => Some(
IpNet::from_str(&s)
.map_err(|e| DbError::Validation(format!("invalid ipv6_cidr '{s}': {e}")))?,
),
None => None,
};
Ok(Interface {
id,
name,
private_key: WireGuardPrivateKey::new(private_key_str),
public_key: WireGuardPublicKey::new(public_key_str),
listen_port: listen_port_i64 as u16,
address_v4,
address_v6,
mtu: mtu_i64.map(|m| m as u16),
dns,
enabled: enabled_i64 != 0,
pre_up,
post_up,
pre_down,
post_down,
created_at: parse_datetime(&created_at_str)?,
updated_at: parse_datetime(&updated_at_str)?,
})
}
/// Create a new WireGuard interface desired configuration record.
pub async fn create_interface(pool: &SqlitePool, iface: &Interface) -> Result<()> {
let id_str = iface.id.to_string();
let ipv4_str = iface.address_v4.to_string();
let ipv6_str = iface.address_v6.as_ref().map(|ip| ip.to_string());
let created_at_str = format_datetime(&iface.created_at);
let updated_at_str = format_datetime(&iface.updated_at);
sqlx::query(
r#"
INSERT INTO interfaces (
id, name, private_key, public_key, listen_port, ipv4_cidr, ipv6_cidr,
mtu, dns, enabled, pre_up, post_up, pre_down, post_down, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id_str)
.bind(&iface.name)
.bind(iface.private_key.as_str())
.bind(iface.public_key.as_str())
.bind(iface.listen_port as i64)
.bind(&ipv4_str)
.bind(ipv6_str)
.bind(iface.mtu.map(|m| m as i64))
.bind(&iface.dns)
.bind(if iface.enabled { 1 } else { 0 })
.bind(&iface.pre_up)
.bind(&iface.post_up)
.bind(&iface.pre_down)
.bind(&iface.post_down)
.bind(&created_at_str)
.bind(&updated_at_str)
.execute(pool)
.await
.map_err(|e| match &e {
sqlx::Error::Database(dbe) if dbe.is_unique_violation() => DbError::Conflict(format!(
"Interface with name '{}' already exists",
iface.name
)),
_ => DbError::Sqlx(e),
})?;
Ok(())
}
/// Retrieve an interface by its UUID.
pub async fn get_interface(pool: &SqlitePool, id: Uuid) -> Result<Option<Interface>> {
let id_str = id.to_string();
let row = sqlx::query("SELECT * FROM interfaces WHERE id = ?")
.bind(&id_str)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => Ok(Some(row_to_interface(&r)?)),
None => Ok(None),
}
}
/// Retrieve an interface by its name.
pub async fn get_interface_by_name(pool: &SqlitePool, name: &str) -> Result<Option<Interface>> {
let row = sqlx::query("SELECT * FROM interfaces WHERE name = ?")
.bind(name)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => Ok(Some(row_to_interface(&r)?)),
None => Ok(None),
}
}
/// List all interfaces.
pub async fn list_interfaces(pool: &SqlitePool) -> Result<Vec<Interface>> {
let rows = sqlx::query("SELECT * FROM interfaces ORDER BY name ASC")
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_interface(&r)?);
}
Ok(list)
}
/// Update an interface record.
pub async fn update_interface(pool: &SqlitePool, iface: &Interface) -> Result<()> {
let id_str = iface.id.to_string();
let ipv4_str = iface.address_v4.to_string();
let ipv6_str = iface.address_v6.as_ref().map(|ip| ip.to_string());
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE interfaces
SET name = ?, private_key = ?, public_key = ?, listen_port = ?,
ipv4_cidr = ?, ipv6_cidr = ?, mtu = ?, dns = ?, enabled = ?,
pre_up = ?, post_up = ?, pre_down = ?, post_down = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(&iface.name)
.bind(iface.private_key.as_str())
.bind(iface.public_key.as_str())
.bind(iface.listen_port as i64)
.bind(&ipv4_str)
.bind(ipv6_str)
.bind(iface.mtu.map(|m| m as i64))
.bind(&iface.dns)
.bind(if iface.enabled { 1 } else { 0 })
.bind(&iface.pre_up)
.bind(&iface.post_up)
.bind(&iface.pre_down)
.bind(&iface.post_down)
.bind(&now_str)
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Interface '{id_str}' not found")));
}
Ok(())
}
/// Delete an interface by UUID. Peers are deleted automatically via ON DELETE CASCADE.
pub async fn delete_interface(pool: &SqlitePool, id: Uuid) -> Result<()> {
let id_str = id.to_string();
let result = sqlx::query("DELETE FROM interfaces WHERE id = ?")
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Interface '{id_str}' not found")));
}
Ok(())
}
/// Enable or disable an interface.
pub async fn set_interface_enabled(pool: &SqlitePool, id: Uuid, enabled: bool) -> Result<()> {
let id_str = id.to_string();
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE interfaces
SET enabled = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(if enabled { 1 } else { 0 })
.bind(&now_str)
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Interface '{id_str}' not found")));
}
Ok(())
}
+28
View File
@@ -0,0 +1,28 @@
//! SQLite persistence layer for nx9-wg.
//!
//! Provides the authoritative desired-state storage, administrator identity,
//! authentication state, interfaces, peers, networks, routes, firewall rules,
//! system settings, backup metadata, and audit records.
pub mod admin;
pub mod audit;
pub mod backups;
pub mod client_profiles;
pub mod error;
pub mod firewall;
pub mod interfaces;
pub mod login_attempts;
pub mod migrations;
pub mod models;
pub mod networks;
pub mod peers;
pub mod routes;
pub mod sessions;
pub mod settings;
pub mod store;
pub mod tokens;
pub use audit::AuditFilter;
pub use error::{DbError, Result};
pub use migrations::run_migrations;
pub use store::Store;
+119
View File
@@ -0,0 +1,119 @@
//! Login attempt tracking repository for brute-force protection.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::{Duration, Utc};
use nx9_wg_core::types::auth::LoginAttempt;
use sqlx::{Row, SqlitePool};
/// Record a login attempt (successful or failed).
pub async fn record_login_attempt(
pool: &SqlitePool,
ip_address: &str,
success: bool,
) -> Result<i64> {
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
INSERT INTO login_attempts (ip_address, attempted_at, success)
VALUES (?, ?, ?)
"#,
)
.bind(ip_address)
.bind(&now_str)
.bind(if success { 1 } else { 0 })
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(result.last_insert_rowid())
}
/// Count failed login attempts from a given IP address within the last `window_minutes`.
pub async fn count_recent_failed_attempts(
pool: &SqlitePool,
ip_address: &str,
window_minutes: i64,
) -> Result<i64> {
let cutoff = Utc::now().naive_utc() - Duration::minutes(window_minutes);
let cutoff_str = format_datetime(&cutoff);
let row = sqlx::query(
r#"
SELECT COUNT(*) as count
FROM login_attempts
WHERE ip_address = ? AND success = 0 AND attempted_at >= ?
"#,
)
.bind(ip_address)
.bind(&cutoff_str)
.fetch_one(pool)
.await
.map_err(DbError::Sqlx)?;
let count: i64 = row.try_get("count")?;
Ok(count)
}
/// Clear login attempts for an IP (e.g. after successful login).
pub async fn clear_login_attempts(pool: &SqlitePool, ip_address: &str) -> Result<u64> {
let result = sqlx::query("DELETE FROM login_attempts WHERE ip_address = ?")
.bind(ip_address)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(result.rows_affected())
}
/// Purge old login attempts older than `retention_hours`.
pub async fn purge_old_login_attempts(pool: &SqlitePool, retention_hours: i64) -> Result<u64> {
let cutoff = Utc::now().naive_utc() - Duration::hours(retention_hours);
let cutoff_str = format_datetime(&cutoff);
let result = sqlx::query("DELETE FROM login_attempts WHERE attempted_at < ?")
.bind(&cutoff_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(result.rows_affected())
}
/// List recent login attempts for diagnostics.
pub async fn list_recent_login_attempts(
pool: &SqlitePool,
limit: u32,
) -> Result<Vec<LoginAttempt>> {
let rows = sqlx::query(
r#"
SELECT id, ip_address, attempted_at, success
FROM login_attempts
ORDER BY id DESC
LIMIT ?
"#,
)
.bind(limit as i64)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
let id: i64 = r.try_get("id")?;
let ip_address: String = r.try_get("ip_address")?;
let attempted_at_str: String = r.try_get("attempted_at")?;
let success_i64: i64 = r.try_get("success")?;
list.push(LoginAttempt {
id,
ip_address,
attempted_at: parse_datetime(&attempted_at_str)?,
success: success_i64 != 0,
});
}
Ok(list)
}
+16
View File
@@ -0,0 +1,16 @@
//! Database migration runner.
use crate::error::{DbError, Result};
use sqlx::SqlitePool;
/// Embed migrations from the `migrations` directory.
pub static MIGRATOR: sqlx::migrate::Migrator = sqlx::migrate!("./migrations");
/// Run all pending SQLite database migrations.
pub async fn run_migrations(pool: &SqlitePool) -> Result<()> {
MIGRATOR
.run(pool)
.await
.map_err(|e| DbError::Migration(e.to_string()))?;
Ok(())
}
+32
View File
@@ -0,0 +1,32 @@
//! Database row models and conversion utilities.
use crate::error::{DbError, Result};
use chrono::NaiveDateTime;
use std::str::FromStr;
/// Parse a string into a `NaiveDateTime` supporting multiple common SQLite date formats.
pub fn parse_datetime(s: &str) -> Result<NaiveDateTime> {
// Try standard formats: "YYYY-MM-DD HH:MM:SS", "YYYY-MM-DDTHH:MM:SS", RFC3339
if let Ok(dt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") {
return Ok(dt);
}
if let Ok(dt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%dT%H:%M:%S") {
return Ok(dt);
}
if let Ok(dt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f") {
return Ok(dt);
}
if let Ok(dt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%dT%H:%M:%S%.f") {
return Ok(dt);
}
if let Ok(dt) = chrono::DateTime::parse_from_rfc3339(s) {
return Ok(dt.naive_utc());
}
NaiveDateTime::from_str(s)
.map_err(|e| DbError::Validation(format!("invalid datetime string '{s}': {e}")))
}
/// Format a `NaiveDateTime` to standard SQLite string format: "YYYY-MM-DD HH:MM:SS".
pub fn format_datetime(dt: &NaiveDateTime) -> String {
dt.format("%Y-%m-%d %H:%M:%S").to_string()
}
+159
View File
@@ -0,0 +1,159 @@
//! Network repository operations.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::Utc;
use ipnet::IpNet;
use nx9_wg_core::types::network::Network;
use sqlx::{Row, SqlitePool};
use std::str::FromStr;
use uuid::Uuid;
/// Helper to convert a database row into a `Network` domain struct.
fn row_to_network(r: &sqlx::sqlite::SqliteRow) -> Result<Network> {
let id_str: String = r.try_get("id")?;
let name: String = r.try_get("name")?;
let cidr_str: String = r.try_get("cidr")?;
let enabled_i64: i64 = r.try_get("enabled")?;
let description: Option<String> = r.try_get("description")?;
let created_at_str: String = r.try_get("created_at")?;
let updated_at_str: String = r.try_get("updated_at")?;
let id = Uuid::parse_str(&id_str)
.map_err(|e| DbError::Validation(format!("invalid network UUID '{id_str}': {e}")))?;
let cidr = IpNet::from_str(&cidr_str)
.map_err(|e| DbError::Validation(format!("invalid network CIDR '{cidr_str}': {e}")))?;
Ok(Network {
id,
name,
cidr,
enabled: enabled_i64 != 0,
description,
created_at: parse_datetime(&created_at_str)?,
updated_at: parse_datetime(&updated_at_str)?,
})
}
/// Create a new network record.
pub async fn create_network(pool: &SqlitePool, net: &Network) -> Result<()> {
let id_str = net.id.to_string();
let cidr_str = net.cidr.to_string();
let created_at_str = format_datetime(&net.created_at);
let updated_at_str = format_datetime(&net.updated_at);
sqlx::query(
r#"
INSERT INTO networks (id, name, cidr, enabled, description, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id_str)
.bind(&net.name)
.bind(&cidr_str)
.bind(if net.enabled { 1 } else { 0 })
.bind(&net.description)
.bind(&created_at_str)
.bind(&updated_at_str)
.execute(pool)
.await
.map_err(|e| match &e {
sqlx::Error::Database(dbe) if dbe.is_unique_violation() => {
DbError::Conflict(format!("Network with name '{}' already exists", net.name))
}
_ => DbError::Sqlx(e),
})?;
Ok(())
}
/// Retrieve a network by UUID.
pub async fn get_network(pool: &SqlitePool, id: Uuid) -> Result<Option<Network>> {
let id_str = id.to_string();
let row = sqlx::query("SELECT * FROM networks WHERE id = ?")
.bind(&id_str)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => Ok(Some(row_to_network(&r)?)),
None => Ok(None),
}
}
/// Retrieve a network by name.
pub async fn get_network_by_name(pool: &SqlitePool, name: &str) -> Result<Option<Network>> {
let row = sqlx::query("SELECT * FROM networks WHERE name = ?")
.bind(name)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => Ok(Some(row_to_network(&r)?)),
None => Ok(None),
}
}
/// List all networks.
pub async fn list_networks(pool: &SqlitePool) -> Result<Vec<Network>> {
let rows = sqlx::query("SELECT * FROM networks ORDER BY name ASC")
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_network(&r)?);
}
Ok(list)
}
/// Update a network record.
pub async fn update_network(pool: &SqlitePool, net: &Network) -> Result<()> {
let id_str = net.id.to_string();
let cidr_str = net.cidr.to_string();
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE networks
SET name = ?, cidr = ?, enabled = ?, description = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(&net.name)
.bind(&cidr_str)
.bind(if net.enabled { 1 } else { 0 })
.bind(&net.description)
.bind(&now_str)
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Network '{id_str}' not found")));
}
Ok(())
}
/// Delete a network by UUID.
pub async fn delete_network(pool: &SqlitePool, id: Uuid) -> Result<()> {
let id_str = id.to_string();
let result = sqlx::query("DELETE FROM networks WHERE id = ?")
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Network '{id_str}' not found")));
}
Ok(())
}
+438
View File
@@ -0,0 +1,438 @@
//! WireGuard Peer repository operations.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::{NaiveDateTime, Utc};
use ipnet::IpNet;
use nx9_wg_core::types::wireguard::{
Peer, PeerProfile, PeerState, PeerType, WireGuardPresharedKey, WireGuardPrivateKey,
WireGuardPublicKey,
};
use sqlx::{Row, SqlitePool};
use std::str::FromStr;
use uuid::Uuid;
/// Helper to convert a database row into a `Peer` domain struct.
fn row_to_peer(r: &sqlx::sqlite::SqliteRow) -> Result<Peer> {
let id_str: String = r.try_get("id")?;
let interface_id_str: String = r.try_get("interface_id")?;
let name: String = r.try_get("name")?;
let peer_type_str: String = r.try_get("peer_type")?;
let state_str: String = r.try_get("state")?;
let profile_str: String = r.try_get("profile")?;
let public_key_str: String = r.try_get("public_key")?;
let private_key_str: Option<String> = r.try_get("private_key")?;
let preshared_key_str: Option<String> = r.try_get("preshared_key")?;
let endpoint: Option<String> = r.try_get("endpoint")?;
let allowed_ips: String = r.try_get("allowed_ips")?;
let server_allowed_ips: Option<String> = r.try_get("server_allowed_ips")?;
let address_ipv4_str: Option<String> = r.try_get("address_ipv4")?;
let address_ipv6_str: Option<String> = r.try_get("address_ipv6")?;
let dns: Option<String> = r.try_get("dns")?;
let mtu_i64: Option<i64> = r.try_get("mtu")?;
let persistent_keepalive_i64: Option<i64> = r.try_get("persistent_keepalive")?;
let expires_at_str: Option<String> = r.try_get("expires_at")?;
let last_handshake_at_str: Option<String> = r.try_get("last_handshake_at")?;
let created_at_str: String = r.try_get("created_at")?;
let updated_at_str: String = r.try_get("updated_at")?;
let id = Uuid::parse_str(&id_str)
.map_err(|e| DbError::Validation(format!("invalid peer UUID '{id_str}': {e}")))?;
let interface_id = Uuid::parse_str(&interface_id_str).map_err(|e| {
DbError::Validation(format!("invalid interface UUID '{interface_id_str}': {e}"))
})?;
let peer_type = PeerType::from_str(&peer_type_str)?;
let state = PeerState::from_str(&state_str)?;
let profile = PeerProfile::from_str(&profile_str)?;
let address_v4 =
match address_ipv4_str {
Some(s) => Some(IpNet::from_str(&s).map_err(|e| {
DbError::Validation(format!("invalid peer address_ipv4 '{s}': {e}"))
})?),
None => None,
};
let address_v6 =
match address_ipv6_str {
Some(s) => Some(IpNet::from_str(&s).map_err(|e| {
DbError::Validation(format!("invalid peer address_ipv6 '{s}': {e}"))
})?),
None => None,
};
let expires_at = match expires_at_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
let last_handshake_at = match last_handshake_at_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
Ok(Peer {
id,
interface_id,
name,
peer_type,
state,
public_key: WireGuardPublicKey::new(public_key_str),
private_key: private_key_str.map(WireGuardPrivateKey::new),
preshared_key: preshared_key_str.map(WireGuardPresharedKey::new),
endpoint,
allowed_ips,
server_allowed_ips,
address_v4,
address_v6,
dns,
mtu: mtu_i64.map(|m| m as u16),
persistent_keepalive: persistent_keepalive_i64.map(|k| k as u16),
profile,
expires_at,
last_handshake_at,
created_at: parse_datetime(&created_at_str)?,
updated_at: parse_datetime(&updated_at_str)?,
})
}
/// Create a new WireGuard peer record.
pub async fn create_peer(pool: &SqlitePool, peer: &Peer) -> Result<()> {
let id_str = peer.id.to_string();
let interface_id_str = peer.interface_id.to_string();
let ipv4_str = peer.address_v4.as_ref().map(|ip| ip.to_string());
let ipv6_str = peer.address_v6.as_ref().map(|ip| ip.to_string());
let expires_at_str = peer.expires_at.as_ref().map(format_datetime);
let last_handshake_str = peer.last_handshake_at.as_ref().map(format_datetime);
let created_at_str = format_datetime(&peer.created_at);
let updated_at_str = format_datetime(&peer.updated_at);
sqlx::query(
r#"
INSERT INTO peers (
id, interface_id, name, peer_type, state, profile, public_key, private_key, preshared_key,
endpoint, allowed_ips, server_allowed_ips, address_ipv4, address_ipv6, dns, mtu,
persistent_keepalive, expires_at, last_handshake_at, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id_str)
.bind(&interface_id_str)
.bind(&peer.name)
.bind(peer.peer_type.as_str())
.bind(peer.state.as_str())
.bind(peer.profile.as_str())
.bind(peer.public_key.as_str())
.bind(peer.private_key.as_ref().map(|pk| pk.as_str()))
.bind(peer.preshared_key.as_ref().map(|psk| psk.as_str()))
.bind(&peer.endpoint)
.bind(&peer.allowed_ips)
.bind(&peer.server_allowed_ips)
.bind(ipv4_str)
.bind(ipv6_str)
.bind(&peer.dns)
.bind(peer.mtu.map(|m| m as i64))
.bind(peer.persistent_keepalive.map(|k| k as i64))
.bind(expires_at_str)
.bind(last_handshake_str)
.bind(&created_at_str)
.bind(&updated_at_str)
.execute(pool)
.await
.map_err(|e| match &e {
sqlx::Error::Database(dbe) if dbe.is_unique_violation() => {
DbError::Conflict(format!("Peer with name '{}' or public key already exists for this interface", peer.name))
}
sqlx::Error::Database(dbe) if dbe.is_foreign_key_violation() => {
DbError::ConstraintViolation(format!("Referenced interface '{}' does not exist", peer.interface_id))
}
_ => DbError::Sqlx(e),
})?;
Ok(())
}
/// Retrieve a peer by its UUID.
pub async fn get_peer(pool: &SqlitePool, id: Uuid) -> Result<Option<Peer>> {
let id_str = id.to_string();
let row = sqlx::query("SELECT * FROM peers WHERE id = ?")
.bind(&id_str)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => Ok(Some(row_to_peer(&r)?)),
None => Ok(None),
}
}
/// Retrieve a peer by name within an interface.
pub async fn get_peer_by_name(
pool: &SqlitePool,
interface_id: Uuid,
name: &str,
) -> Result<Option<Peer>> {
let iface_id_str = interface_id.to_string();
let row = sqlx::query("SELECT * FROM peers WHERE interface_id = ? AND name = ?")
.bind(&iface_id_str)
.bind(name)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => Ok(Some(row_to_peer(&r)?)),
None => Ok(None),
}
}
/// Retrieve a peer by public key within an interface.
pub async fn get_peer_by_public_key(
pool: &SqlitePool,
interface_id: Uuid,
public_key: &str,
) -> Result<Option<Peer>> {
let iface_id_str = interface_id.to_string();
let row = sqlx::query("SELECT * FROM peers WHERE interface_id = ? AND public_key = ?")
.bind(&iface_id_str)
.bind(public_key)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => Ok(Some(row_to_peer(&r)?)),
None => Ok(None),
}
}
/// List all peers for a given interface.
pub async fn list_peers_for_interface(pool: &SqlitePool, interface_id: Uuid) -> Result<Vec<Peer>> {
let iface_id_str = interface_id.to_string();
let rows = sqlx::query("SELECT * FROM peers WHERE interface_id = ? ORDER BY name ASC")
.bind(&iface_id_str)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_peer(&r)?);
}
Ok(list)
}
/// List all peers across all interfaces.
pub async fn list_all_peers(pool: &SqlitePool) -> Result<Vec<Peer>> {
let rows = sqlx::query("SELECT * FROM peers ORDER BY name ASC")
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_peer(&r)?);
}
Ok(list)
}
/// Update a peer record.
pub async fn update_peer(pool: &SqlitePool, peer: &Peer) -> Result<()> {
let id_str = peer.id.to_string();
let interface_id_str = peer.interface_id.to_string();
let ipv4_str = peer.address_v4.as_ref().map(|ip| ip.to_string());
let ipv6_str = peer.address_v6.as_ref().map(|ip| ip.to_string());
let expires_at_str = peer.expires_at.as_ref().map(format_datetime);
let last_handshake_str = peer.last_handshake_at.as_ref().map(format_datetime);
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE peers
SET interface_id = ?, name = ?, peer_type = ?, state = ?, profile = ?,
public_key = ?, private_key = ?, preshared_key = ?, endpoint = ?,
allowed_ips = ?, server_allowed_ips = ?, address_ipv4 = ?, address_ipv6 = ?,
dns = ?, mtu = ?, persistent_keepalive = ?, expires_at = ?,
last_handshake_at = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(&interface_id_str)
.bind(&peer.name)
.bind(peer.peer_type.as_str())
.bind(peer.state.as_str())
.bind(peer.profile.as_str())
.bind(peer.public_key.as_str())
.bind(peer.private_key.as_ref().map(|pk| pk.as_str()))
.bind(peer.preshared_key.as_ref().map(|psk| psk.as_str()))
.bind(&peer.endpoint)
.bind(&peer.allowed_ips)
.bind(&peer.server_allowed_ips)
.bind(ipv4_str)
.bind(ipv6_str)
.bind(&peer.dns)
.bind(peer.mtu.map(|m| m as i64))
.bind(peer.persistent_keepalive.map(|k| k as i64))
.bind(expires_at_str)
.bind(last_handshake_str)
.bind(&now_str)
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Peer '{id_str}' not found")));
}
Ok(())
}
/// Update peer state (active, disabled, revoked, expired).
pub async fn set_peer_state(pool: &SqlitePool, id: Uuid, state: PeerState) -> Result<()> {
let id_str = id.to_string();
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE peers
SET state = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(state.as_str())
.bind(&now_str)
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Peer '{id_str}' not found")));
}
Ok(())
}
/// Update operational last_handshake_at timestamp.
pub async fn update_peer_handshake(
pool: &SqlitePool,
id: Uuid,
handshake_at: NaiveDateTime,
) -> Result<()> {
let id_str = id.to_string();
let handshake_str = format_datetime(&handshake_at);
let result = sqlx::query(
r#"
UPDATE peers
SET last_handshake_at = ?
WHERE id = ?
"#,
)
.bind(&handshake_str)
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Peer '{id_str}' not found")));
}
Ok(())
}
/// Delete a peer by UUID.
pub async fn delete_peer(pool: &SqlitePool, id: Uuid) -> Result<()> {
let id_str = id.to_string();
let result = sqlx::query("DELETE FROM peers WHERE id = ?")
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Peer '{id_str}' not found")));
}
Ok(())
}
/// Retrieve all allocated IP addresses (CIDR strings) for an interface or across all interfaces.
pub async fn get_allocated_ips(
pool: &SqlitePool,
interface_id: Option<Uuid>,
) -> Result<Vec<String>> {
let rows = match interface_id {
Some(iface_id) => {
let iface_id_str = iface_id.to_string();
sqlx::query(
r#"
SELECT address_ipv4, address_ipv6
FROM peers
WHERE interface_id = ? AND state != 'revoked'
"#,
)
.bind(&iface_id_str)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?
}
None => sqlx::query(
r#"
SELECT address_ipv4, address_ipv6
FROM peers
WHERE state != 'revoked'
"#,
)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?,
};
let mut allocated = Vec::new();
for r in rows {
let v4: Option<String> = r.try_get("address_ipv4")?;
let v6: Option<String> = r.try_get("address_ipv6")?;
if let Some(ip) = v4.as_ref().filter(|s| !s.trim().is_empty()) {
allocated.push(ip.clone());
}
if let Some(ip) = v6.as_ref().filter(|s| !s.trim().is_empty()) {
allocated.push(ip.clone());
}
}
Ok(allocated)
}
/// Find active peers whose expiration timestamp has passed.
pub async fn get_expired_active_peers(pool: &SqlitePool, now: NaiveDateTime) -> Result<Vec<Peer>> {
let now_str = format_datetime(&now);
let rows = sqlx::query(
r#"
SELECT * FROM peers
WHERE state = 'active' AND expires_at IS NOT NULL AND expires_at <= ?
"#,
)
.bind(&now_str)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_peer(&r)?);
}
Ok(list)
}
/// Mark a peer as expired.
pub async fn mark_peer_expired(pool: &SqlitePool, id: Uuid) -> Result<()> {
set_peer_state(pool, id, PeerState::Expired).await
}
+251
View File
@@ -0,0 +1,251 @@
//! Route repository operations.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::Utc;
use ipnet::IpNet;
use nx9_wg_core::types::network::Route;
use sqlx::{Row, SqlitePool};
use std::net::IpAddr;
use std::str::FromStr;
use uuid::Uuid;
/// Helper to convert a database row into a `Route` domain struct.
fn row_to_route(r: &sqlx::sqlite::SqliteRow) -> Result<Route> {
let id_str: String = r.try_get("id")?;
let network_id_str: Option<String> = r.try_get("network_id")?;
let interface_id_str: Option<String> = r.try_get("interface_id")?;
let destination_str: String = r.try_get("destination")?;
let gateway_str: Option<String> = r.try_get("gateway")?;
let metric_i64: Option<i64> = r.try_get("metric")?;
let enabled_i64: i64 = r.try_get("enabled")?;
let description: Option<String> = r.try_get("description")?;
let created_at_str: String = r.try_get("created_at")?;
let updated_at_str: String = r.try_get("updated_at")?;
let id = Uuid::parse_str(&id_str)
.map_err(|e| DbError::Validation(format!("invalid route UUID '{id_str}': {e}")))?;
let network_id = match network_id_str {
Some(s) => Some(
Uuid::parse_str(&s)
.map_err(|e| DbError::Validation(format!("invalid network UUID '{s}': {e}")))?,
),
None => None,
};
let interface_id = match interface_id_str {
Some(s) => Some(
Uuid::parse_str(&s)
.map_err(|e| DbError::Validation(format!("invalid interface UUID '{s}': {e}")))?,
),
None => None,
};
let destination = IpNet::from_str(&destination_str).map_err(|e| {
DbError::Validation(format!("invalid destination CIDR '{destination_str}': {e}"))
})?;
let gateway = match gateway_str {
Some(s) => Some(
IpAddr::from_str(&s)
.map_err(|e| DbError::Validation(format!("invalid gateway IP '{s}': {e}")))?,
),
None => None,
};
Ok(Route {
id,
network_id,
interface_id,
destination,
gateway,
interface_name: None,
metric: metric_i64.map(|m| m as u32),
enabled: enabled_i64 != 0,
description,
created_at: parse_datetime(&created_at_str)?,
updated_at: parse_datetime(&updated_at_str)?,
})
}
/// Create a new route record.
pub async fn create_route(pool: &SqlitePool, route: &Route) -> Result<()> {
let id_str = route.id.to_string();
let network_id_str = route.network_id.map(|id| id.to_string());
let interface_id_str = route.interface_id.map(|id| id.to_string());
let dest_str = route.destination.to_string();
let gateway_str = route.gateway.map(|g| g.to_string());
let created_at_str = format_datetime(&route.created_at);
let updated_at_str = format_datetime(&route.updated_at);
sqlx::query(
r#"
INSERT INTO routes (
id, network_id, interface_id, destination, gateway,
metric, enabled, description, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id_str)
.bind(network_id_str)
.bind(interface_id_str)
.bind(&dest_str)
.bind(gateway_str)
.bind(route.metric.map(|m| m as i64))
.bind(if route.enabled { 1 } else { 0 })
.bind(&route.description)
.bind(&created_at_str)
.bind(&updated_at_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(())
}
/// Retrieve a route by UUID.
pub async fn get_route(pool: &SqlitePool, id: Uuid) -> Result<Option<Route>> {
let id_str = id.to_string();
let row = sqlx::query("SELECT * FROM routes WHERE id = ?")
.bind(&id_str)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => Ok(Some(row_to_route(&r)?)),
None => Ok(None),
}
}
/// List all routes.
pub async fn list_routes(pool: &SqlitePool) -> Result<Vec<Route>> {
let rows = sqlx::query("SELECT * FROM routes ORDER BY destination ASC")
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_route(&r)?);
}
Ok(list)
}
/// List routes referencing a given network.
pub async fn list_routes_for_network(pool: &SqlitePool, network_id: Uuid) -> Result<Vec<Route>> {
let net_id_str = network_id.to_string();
let rows = sqlx::query("SELECT * FROM routes WHERE network_id = ? ORDER BY destination ASC")
.bind(&net_id_str)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_route(&r)?);
}
Ok(list)
}
/// List routes referencing a given interface.
pub async fn list_routes_for_interface(
pool: &SqlitePool,
interface_id: Uuid,
) -> Result<Vec<Route>> {
let iface_id_str = interface_id.to_string();
let rows = sqlx::query("SELECT * FROM routes WHERE interface_id = ? ORDER BY destination ASC")
.bind(&iface_id_str)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
list.push(row_to_route(&r)?);
}
Ok(list)
}
/// Update a route record.
pub async fn update_route(pool: &SqlitePool, route: &Route) -> Result<()> {
let id_str = route.id.to_string();
let network_id_str = route.network_id.map(|id| id.to_string());
let interface_id_str = route.interface_id.map(|id| id.to_string());
let dest_str = route.destination.to_string();
let gateway_str = route.gateway.map(|g| g.to_string());
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE routes
SET network_id = ?, interface_id = ?, destination = ?, gateway = ?,
metric = ?, enabled = ?, description = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(network_id_str)
.bind(interface_id_str)
.bind(&dest_str)
.bind(gateway_str)
.bind(route.metric.map(|m| m as i64))
.bind(if route.enabled { 1 } else { 0 })
.bind(&route.description)
.bind(&now_str)
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Route '{id_str}' not found")));
}
Ok(())
}
/// Delete a route by UUID.
pub async fn delete_route(pool: &SqlitePool, id: Uuid) -> Result<()> {
let id_str = id.to_string();
let result = sqlx::query("DELETE FROM routes WHERE id = ?")
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Route '{id_str}' not found")));
}
Ok(())
}
/// Enable or disable a route.
pub async fn set_route_enabled(pool: &SqlitePool, id: Uuid, enabled: bool) -> Result<()> {
let id_str = id.to_string();
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE routes
SET enabled = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(if enabled { 1 } else { 0 })
.bind(&now_str)
.bind(&id_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Route '{id_str}' not found")));
}
Ok(())
}
+180
View File
@@ -0,0 +1,180 @@
//! Session repository operations.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::Utc;
use nx9_wg_core::types::auth::Session;
use sqlx::{Row, SqlitePool};
/// Create a new session.
pub async fn create_session(pool: &SqlitePool, session: &Session) -> Result<()> {
let created_at_str = format_datetime(&session.created_at);
let expires_at_str = format_datetime(&session.expires_at);
let last_seen_str = session.last_seen_at.as_ref().map(format_datetime);
sqlx::query(
r#"
INSERT INTO sessions (id, admin_id, ip_address, user_agent, created_at, expires_at, last_seen_at)
VALUES (?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&session.id)
.bind(session.admin_id)
.bind(&session.ip_address)
.bind(&session.user_agent)
.bind(&created_at_str)
.bind(&expires_at_str)
.bind(last_seen_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(())
}
/// Retrieve a session by its ID.
pub async fn get_session(pool: &SqlitePool, id: &str) -> Result<Option<Session>> {
let row = sqlx::query(
r#"
SELECT id, admin_id, ip_address, user_agent, created_at, expires_at, last_seen_at
FROM sessions
WHERE id = ?
"#,
)
.bind(id)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => {
let id: String = r.try_get("id")?;
let admin_id: i64 = r.try_get("admin_id")?;
let ip_address: Option<String> = r.try_get("ip_address")?;
let user_agent: Option<String> = r.try_get("user_agent")?;
let created_at_str: String = r.try_get("created_at")?;
let expires_at_str: String = r.try_get("expires_at")?;
let last_seen_str: Option<String> = r.try_get("last_seen_at")?;
let last_seen_at = match last_seen_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
Ok(Some(Session {
id,
admin_id,
created_at: parse_datetime(&created_at_str)?,
expires_at: parse_datetime(&expires_at_str)?,
last_seen_at,
ip_address,
user_agent,
}))
}
None => Ok(None),
}
}
/// Touch a session by updating its `last_seen_at` to the current time.
pub async fn touch_session(pool: &SqlitePool, id: &str) -> Result<()> {
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE sessions
SET last_seen_at = ?
WHERE id = ?
"#,
)
.bind(&now_str)
.bind(id)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("Session '{id}' not found")));
}
Ok(())
}
/// Delete a session by ID (logout).
pub async fn delete_session(pool: &SqlitePool, id: &str) -> Result<()> {
sqlx::query("DELETE FROM sessions WHERE id = ?")
.bind(id)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(())
}
/// Delete all expired sessions. Returns the count of deleted sessions.
pub async fn delete_expired_sessions(pool: &SqlitePool) -> Result<u64> {
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query("DELETE FROM sessions WHERE expires_at < ?")
.bind(&now_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(result.rows_affected())
}
/// Delete all sessions for the given administrator (e.g. after password reset).
pub async fn delete_all_admin_sessions(pool: &SqlitePool, admin_id: i64) -> Result<u64> {
let result = sqlx::query("DELETE FROM sessions WHERE admin_id = ?")
.bind(admin_id)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(result.rows_affected())
}
/// List all active and unexpired sessions.
pub async fn list_sessions(pool: &SqlitePool) -> Result<Vec<Session>> {
let rows = sqlx::query("SELECT id, admin_id, ip_address, user_agent, created_at, expires_at, last_seen_at FROM sessions ORDER BY created_at DESC")
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut sessions = Vec::with_capacity(rows.len());
for r in rows {
let id: String = r.try_get("id")?;
let admin_id: i64 = r.try_get("admin_id")?;
let ip_address: Option<String> = r.try_get("ip_address")?;
let user_agent: Option<String> = r.try_get("user_agent")?;
let created_at_str: String = r.try_get("created_at")?;
let expires_at_str: String = r.try_get("expires_at")?;
let last_seen_str: Option<String> = r.try_get("last_seen_at")?;
let last_seen_at = match last_seen_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
sessions.push(Session {
id,
admin_id,
created_at: parse_datetime(&created_at_str)?,
expires_at: parse_datetime(&expires_at_str)?,
last_seen_at,
ip_address,
user_agent,
});
}
Ok(sessions)
}
/// Delete all sessions unconditionally.
pub async fn delete_all_sessions(pool: &SqlitePool) -> Result<u64> {
let result = sqlx::query("DELETE FROM sessions")
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(result.rows_affected())
}
+101
View File
@@ -0,0 +1,101 @@
//! Settings repository operations.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::Utc;
use nx9_wg_core::types::settings::Setting;
use sqlx::{Row, SqlitePool};
/// Retrieve a setting by its key.
pub async fn get_setting(pool: &SqlitePool, key: &str) -> Result<Option<Setting>> {
let row = sqlx::query("SELECT key, value, is_secret, updated_at FROM settings WHERE key = ?")
.bind(key)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => {
let key: String = r.try_get("key")?;
let value: String = r.try_get("value")?;
let is_secret_i64: i64 = r.try_get("is_secret")?;
let updated_at_str: String = r.try_get("updated_at")?;
Ok(Some(Setting {
key,
value,
is_secret: is_secret_i64 != 0,
updated_at: parse_datetime(&updated_at_str)?,
}))
}
None => Ok(None),
}
}
/// Retrieve only the string value of a setting, if present.
pub async fn get_setting_value(pool: &SqlitePool, key: &str) -> Result<Option<String>> {
let setting = get_setting(pool, key).await?;
Ok(setting.map(|s| s.value))
}
/// Upsert a setting key-value pair.
pub async fn set_setting(pool: &SqlitePool, key: &str, value: &str, is_secret: bool) -> Result<()> {
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
sqlx::query(
r#"
INSERT INTO settings (key, value, is_secret, updated_at)
VALUES (?, ?, ?, ?)
ON CONFLICT(key) DO UPDATE SET
value = excluded.value,
is_secret = excluded.is_secret,
updated_at = excluded.updated_at
"#,
)
.bind(key)
.bind(value)
.bind(if is_secret { 1 } else { 0 })
.bind(&now_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(())
}
/// Delete a setting by key.
pub async fn delete_setting(pool: &SqlitePool, key: &str) -> Result<()> {
sqlx::query("DELETE FROM settings WHERE key = ?")
.bind(key)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(())
}
/// List all settings.
pub async fn list_settings(pool: &SqlitePool) -> Result<Vec<Setting>> {
let rows =
sqlx::query("SELECT key, value, is_secret, updated_at FROM settings ORDER BY key ASC")
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut list = Vec::with_capacity(rows.len());
for r in rows {
let key: String = r.try_get("key")?;
let value: String = r.try_get("value")?;
let is_secret_i64: i64 = r.try_get("is_secret")?;
let updated_at_str: String = r.try_get("updated_at")?;
list.push(Setting {
key,
value,
is_secret: is_secret_i64 != 0,
updated_at: parse_datetime(&updated_at_str)?,
});
}
Ok(list)
}
+645
View File
@@ -0,0 +1,645 @@
//! Central database `Store` providing connection lifecycle, migrations, and repository access.
use crate::error::{DbError, Result};
use crate::migrations::run_migrations;
use sqlx::SqlitePool;
use sqlx::sqlite::{SqliteConnectOptions, SqliteJournalMode, SqlitePoolOptions, SqliteSynchronous};
use std::path::Path;
use std::str::FromStr;
use std::time::Duration;
/// Central database store handle wrapping the SQLite connection pool.
#[derive(Debug, Clone)]
pub struct Store {
pool: SqlitePool,
}
impl Store {
/// Connect to a SQLite database by path, ensuring directory creation and setting WAL/foreign keys.
pub async fn connect_path<P: AsRef<Path>>(path: P) -> Result<Self> {
let path = path.as_ref();
if let Some(parent) = path
.parent()
.filter(|p| !p.as_os_str().is_empty() && !p.exists())
{
std::fs::create_dir_all(parent).map_err(|e| {
DbError::Internal(format!(
"failed to create database parent directory '{}': {e}",
parent.display()
))
})?;
}
let opts = SqliteConnectOptions::new()
.filename(path)
.create_if_missing(true)
.journal_mode(SqliteJournalMode::Wal)
.synchronous(SqliteSynchronous::Normal)
.foreign_keys(true)
.busy_timeout(Duration::from_millis(5000));
let pool = SqlitePoolOptions::new()
.max_connections(10)
.min_connections(1)
.acquire_timeout(Duration::from_secs(10))
.connect_with(opts)
.await
.map_err(DbError::Sqlx)?;
Ok(Self { pool })
}
/// Connect to a SQLite database using a connection string URL (e.g. `sqlite:///var/lib/nx9-wg/nx9-wg.db`).
pub async fn connect(database_url: &str) -> Result<Self> {
let opts = SqliteConnectOptions::from_str(database_url)
.map_err(|e| {
DbError::Validation(format!("invalid database URL '{database_url}': {e}"))
})?
.create_if_missing(true)
.journal_mode(SqliteJournalMode::Wal)
.synchronous(SqliteSynchronous::Normal)
.foreign_keys(true)
.busy_timeout(Duration::from_millis(5000));
let pool = SqlitePoolOptions::new()
.max_connections(10)
.min_connections(1)
.acquire_timeout(Duration::from_secs(10))
.connect_with(opts)
.await
.map_err(DbError::Sqlx)?;
Ok(Self { pool })
}
/// Create an in-memory SQLite database store (useful for tests).
pub async fn connect_in_memory() -> Result<Self> {
let opts = SqliteConnectOptions::new()
.filename(":memory:")
.foreign_keys(true)
.busy_timeout(Duration::from_millis(5000));
// In-memory SQLite databases require max_connections=1 so the same DB is shared across queries
let pool = SqlitePoolOptions::new()
.max_connections(1)
.connect_with(opts)
.await
.map_err(DbError::Sqlx)?;
Ok(Self { pool })
}
/// Run all pending SQLx migrations.
pub async fn migrate(&self) -> Result<()> {
run_migrations(&self.pool).await
}
/// Get a reference to the underlying `SqlitePool`.
pub fn pool(&self) -> &SqlitePool {
&self.pool
}
/// Check database connectivity with a simple SELECT 1 query.
pub async fn health_check(&self) -> Result<()> {
sqlx::query("SELECT 1")
.execute(&self.pool)
.await
.map_err(DbError::Sqlx)?;
Ok(())
}
/// Close the connection pool gracefully.
pub async fn close(&self) {
self.pool.close().await;
}
// ── Repository convenience accessors ───────────────────────────────
// Administrator
pub async fn get_admin(&self) -> Result<Option<nx9_wg_core::types::auth::Admin>> {
crate::admin::get_admin(&self.pool).await
}
pub async fn get_admin_by_username(
&self,
username: &str,
) -> Result<Option<nx9_wg_core::types::auth::Admin>> {
crate::admin::get_admin_by_username(&self.pool, username).await
}
pub async fn admin_exists(&self) -> Result<bool> {
crate::admin::admin_exists(&self.pool).await
}
pub async fn create_admin(
&self,
username: &str,
password_hash: &str,
) -> Result<nx9_wg_core::types::auth::Admin> {
crate::admin::create_admin(&self.pool, username, password_hash).await
}
pub async fn update_admin_password(&self, new_password_hash: &str) -> Result<()> {
crate::admin::update_admin_password(&self.pool, new_password_hash).await
}
pub async fn update_admin_totp(&self, secret: Option<&str>, enabled: bool) -> Result<()> {
crate::admin::update_admin_totp(&self.pool, secret, enabled).await
}
pub async fn record_admin_login(&self, ip_address: Option<&str>) -> Result<()> {
crate::admin::record_admin_login(&self.pool, ip_address).await
}
// Sessions
pub async fn create_session(&self, session: &nx9_wg_core::types::auth::Session) -> Result<()> {
crate::sessions::create_session(&self.pool, session).await
}
pub async fn get_session(&self, id: &str) -> Result<Option<nx9_wg_core::types::auth::Session>> {
crate::sessions::get_session(&self.pool, id).await
}
pub async fn touch_session(&self, id: &str) -> Result<()> {
crate::sessions::touch_session(&self.pool, id).await
}
pub async fn delete_session(&self, id: &str) -> Result<()> {
crate::sessions::delete_session(&self.pool, id).await
}
pub async fn delete_expired_sessions(&self) -> Result<u64> {
crate::sessions::delete_expired_sessions(&self.pool).await
}
pub async fn delete_all_admin_sessions(&self, admin_id: i64) -> Result<u64> {
crate::sessions::delete_all_admin_sessions(&self.pool, admin_id).await
}
pub async fn list_sessions(&self) -> Result<Vec<nx9_wg_core::types::auth::Session>> {
crate::sessions::list_sessions(&self.pool).await
}
pub async fn delete_all_sessions(&self) -> Result<u64> {
crate::sessions::delete_all_sessions(&self.pool).await
}
// Login Attempts (Rate Limiting)
pub async fn record_login_attempt(&self, ip_address: &str, success: bool) -> Result<i64> {
crate::login_attempts::record_login_attempt(&self.pool, ip_address, success).await
}
pub async fn count_recent_failed_attempts(
&self,
ip_address: &str,
window_minutes: i64,
) -> Result<i64> {
crate::login_attempts::count_recent_failed_attempts(&self.pool, ip_address, window_minutes)
.await
}
pub async fn clear_login_attempts(&self, ip_address: &str) -> Result<u64> {
crate::login_attempts::clear_login_attempts(&self.pool, ip_address).await
}
pub async fn purge_old_login_attempts(&self, retention_hours: i64) -> Result<u64> {
crate::login_attempts::purge_old_login_attempts(&self.pool, retention_hours).await
}
pub async fn list_recent_login_attempts(
&self,
limit: u32,
) -> Result<Vec<nx9_wg_core::types::auth::LoginAttempt>> {
crate::login_attempts::list_recent_login_attempts(&self.pool, limit).await
}
// API Tokens
pub async fn create_token(&self, token: &nx9_wg_core::types::auth::ApiToken) -> Result<()> {
crate::tokens::create_token(&self.pool, token).await
}
pub async fn list_tokens(&self) -> Result<Vec<nx9_wg_core::types::auth::ApiToken>> {
crate::tokens::list_tokens(&self.pool).await
}
pub async fn get_token(&self, id: &str) -> Result<Option<nx9_wg_core::types::auth::ApiToken>> {
crate::tokens::get_token(&self.pool, id).await
}
pub async fn find_token_by_hash(
&self,
hash: &str,
) -> Result<Option<nx9_wg_core::types::auth::ApiToken>> {
crate::tokens::find_token_by_hash(&self.pool, hash).await
}
pub async fn mark_token_used(&self, id: &str) -> Result<()> {
crate::tokens::mark_token_used(&self.pool, id).await
}
pub async fn revoke_token(&self, id: &str) -> Result<()> {
crate::tokens::revoke_token(&self.pool, id).await
}
pub async fn delete_token(&self, id: &str) -> Result<()> {
crate::tokens::delete_token(&self.pool, id).await
}
pub async fn delete_expired_tokens(&self) -> Result<u64> {
crate::tokens::delete_expired_tokens(&self.pool).await
}
// Interfaces
pub async fn create_interface(
&self,
iface: &nx9_wg_core::types::wireguard::Interface,
) -> Result<()> {
crate::interfaces::create_interface(&self.pool, iface).await
}
pub async fn get_interface(
&self,
id: uuid::Uuid,
) -> Result<Option<nx9_wg_core::types::wireguard::Interface>> {
crate::interfaces::get_interface(&self.pool, id).await
}
pub async fn get_interface_by_name(
&self,
name: &str,
) -> Result<Option<nx9_wg_core::types::wireguard::Interface>> {
crate::interfaces::get_interface_by_name(&self.pool, name).await
}
pub async fn list_interfaces(&self) -> Result<Vec<nx9_wg_core::types::wireguard::Interface>> {
crate::interfaces::list_interfaces(&self.pool).await
}
pub async fn update_interface(
&self,
iface: &nx9_wg_core::types::wireguard::Interface,
) -> Result<()> {
crate::interfaces::update_interface(&self.pool, iface).await
}
pub async fn delete_interface(&self, id: uuid::Uuid) -> Result<()> {
crate::interfaces::delete_interface(&self.pool, id).await
}
pub async fn set_interface_enabled(&self, id: uuid::Uuid, enabled: bool) -> Result<()> {
crate::interfaces::set_interface_enabled(&self.pool, id, enabled).await
}
// Peers
pub async fn create_peer(&self, peer: &nx9_wg_core::types::wireguard::Peer) -> Result<()> {
crate::peers::create_peer(&self.pool, peer).await
}
pub async fn get_peer(
&self,
id: uuid::Uuid,
) -> Result<Option<nx9_wg_core::types::wireguard::Peer>> {
crate::peers::get_peer(&self.pool, id).await
}
pub async fn get_peer_by_name(
&self,
iface_id: uuid::Uuid,
name: &str,
) -> Result<Option<nx9_wg_core::types::wireguard::Peer>> {
crate::peers::get_peer_by_name(&self.pool, iface_id, name).await
}
pub async fn get_peer_by_public_key(
&self,
iface_id: uuid::Uuid,
pub_key: &str,
) -> Result<Option<nx9_wg_core::types::wireguard::Peer>> {
crate::peers::get_peer_by_public_key(&self.pool, iface_id, pub_key).await
}
pub async fn list_peers_for_interface(
&self,
iface_id: uuid::Uuid,
) -> Result<Vec<nx9_wg_core::types::wireguard::Peer>> {
crate::peers::list_peers_for_interface(&self.pool, iface_id).await
}
pub async fn list_all_peers(&self) -> Result<Vec<nx9_wg_core::types::wireguard::Peer>> {
crate::peers::list_all_peers(&self.pool).await
}
pub async fn update_peer(&self, peer: &nx9_wg_core::types::wireguard::Peer) -> Result<()> {
crate::peers::update_peer(&self.pool, peer).await
}
pub async fn set_peer_state(
&self,
id: uuid::Uuid,
state: nx9_wg_core::types::wireguard::PeerState,
) -> Result<()> {
crate::peers::set_peer_state(&self.pool, id, state).await
}
pub async fn update_peer_handshake(
&self,
id: uuid::Uuid,
handshake_at: chrono::NaiveDateTime,
) -> Result<()> {
crate::peers::update_peer_handshake(&self.pool, id, handshake_at).await
}
pub async fn delete_peer(&self, id: uuid::Uuid) -> Result<()> {
crate::peers::delete_peer(&self.pool, id).await
}
pub async fn get_allocated_ips(&self, interface_id: Option<uuid::Uuid>) -> Result<Vec<String>> {
crate::peers::get_allocated_ips(&self.pool, interface_id).await
}
pub async fn get_expired_active_peers(
&self,
now: chrono::NaiveDateTime,
) -> Result<Vec<nx9_wg_core::types::wireguard::Peer>> {
crate::peers::get_expired_active_peers(&self.pool, now).await
}
pub async fn mark_peer_expired(&self, id: uuid::Uuid) -> Result<()> {
crate::peers::mark_peer_expired(&self.pool, id).await
}
// Networks
pub async fn create_network(&self, net: &nx9_wg_core::types::network::Network) -> Result<()> {
crate::networks::create_network(&self.pool, net).await
}
pub async fn get_network(
&self,
id: uuid::Uuid,
) -> Result<Option<nx9_wg_core::types::network::Network>> {
crate::networks::get_network(&self.pool, id).await
}
pub async fn get_network_by_name(
&self,
name: &str,
) -> Result<Option<nx9_wg_core::types::network::Network>> {
crate::networks::get_network_by_name(&self.pool, name).await
}
pub async fn list_networks(&self) -> Result<Vec<nx9_wg_core::types::network::Network>> {
crate::networks::list_networks(&self.pool).await
}
pub async fn update_network(&self, net: &nx9_wg_core::types::network::Network) -> Result<()> {
crate::networks::update_network(&self.pool, net).await
}
pub async fn delete_network(&self, id: uuid::Uuid) -> Result<()> {
crate::networks::delete_network(&self.pool, id).await
}
// Routes
pub async fn create_route(&self, route: &nx9_wg_core::types::network::Route) -> Result<()> {
crate::routes::create_route(&self.pool, route).await
}
pub async fn get_route(
&self,
id: uuid::Uuid,
) -> Result<Option<nx9_wg_core::types::network::Route>> {
crate::routes::get_route(&self.pool, id).await
}
pub async fn list_routes(&self) -> Result<Vec<nx9_wg_core::types::network::Route>> {
crate::routes::list_routes(&self.pool).await
}
pub async fn list_routes_for_network(
&self,
network_id: uuid::Uuid,
) -> Result<Vec<nx9_wg_core::types::network::Route>> {
crate::routes::list_routes_for_network(&self.pool, network_id).await
}
pub async fn list_routes_for_interface(
&self,
interface_id: uuid::Uuid,
) -> Result<Vec<nx9_wg_core::types::network::Route>> {
crate::routes::list_routes_for_interface(&self.pool, interface_id).await
}
pub async fn update_route(&self, route: &nx9_wg_core::types::network::Route) -> Result<()> {
crate::routes::update_route(&self.pool, route).await
}
pub async fn delete_route(&self, id: uuid::Uuid) -> Result<()> {
crate::routes::delete_route(&self.pool, id).await
}
pub async fn set_route_enabled(&self, id: uuid::Uuid, enabled: bool) -> Result<()> {
crate::routes::set_route_enabled(&self.pool, id, enabled).await
}
// Firewall
pub async fn create_firewall_rule(
&self,
rule: &nx9_wg_core::types::firewall::FirewallRule,
) -> Result<()> {
crate::firewall::create_rule(&self.pool, rule).await
}
pub async fn get_firewall_rule(
&self,
id: uuid::Uuid,
) -> Result<Option<nx9_wg_core::types::firewall::FirewallRule>> {
crate::firewall::get_rule(&self.pool, id).await
}
pub async fn list_firewall_rules(
&self,
) -> Result<Vec<nx9_wg_core::types::firewall::FirewallRule>> {
crate::firewall::list_rules(&self.pool).await
}
pub async fn list_firewall_rules_for_interface(
&self,
interface_id: uuid::Uuid,
) -> Result<Vec<nx9_wg_core::types::firewall::FirewallRule>> {
crate::firewall::list_rules_for_interface(&self.pool, interface_id).await
}
pub async fn list_firewall_rules_for_peer(
&self,
peer_id: uuid::Uuid,
) -> Result<Vec<nx9_wg_core::types::firewall::FirewallRule>> {
crate::firewall::list_rules_for_peer(&self.pool, peer_id).await
}
pub async fn update_firewall_rule(
&self,
rule: &nx9_wg_core::types::firewall::FirewallRule,
) -> Result<()> {
crate::firewall::update_rule(&self.pool, rule).await
}
pub async fn delete_firewall_rule(&self, id: uuid::Uuid) -> Result<()> {
crate::firewall::delete_rule(&self.pool, id).await
}
pub async fn set_firewall_rule_enabled(&self, id: uuid::Uuid, enabled: bool) -> Result<()> {
crate::firewall::set_rule_enabled(&self.pool, id, enabled).await
}
// Settings
pub async fn get_setting(
&self,
key: &str,
) -> Result<Option<nx9_wg_core::types::settings::Setting>> {
crate::settings::get_setting(&self.pool, key).await
}
pub async fn get_setting_value(&self, key: &str) -> Result<Option<String>> {
crate::settings::get_setting_value(&self.pool, key).await
}
pub async fn set_setting(&self, key: &str, value: &str, is_secret: bool) -> Result<()> {
crate::settings::set_setting(&self.pool, key, value, is_secret).await
}
pub async fn delete_setting(&self, key: &str) -> Result<()> {
crate::settings::delete_setting(&self.pool, key).await
}
pub async fn list_settings(&self) -> Result<Vec<nx9_wg_core::types::settings::Setting>> {
crate::settings::list_settings(&self.pool).await
}
// Audit
pub async fn create_audit_event(
&self,
event: &nx9_wg_core::types::audit::AuditEvent,
) -> Result<i64> {
crate::audit::create_audit_event(&self.pool, event).await
}
#[allow(clippy::too_many_arguments)]
pub async fn record_audit(
&self,
event_type: nx9_wg_core::types::audit::AuditEventType,
actor: &str,
resource_type: Option<&str>,
resource_id: Option<&str>,
message: Option<&str>,
metadata: Option<&str>,
ip_address: Option<&str>,
) -> Result<i64> {
crate::audit::record_audit(
&self.pool,
event_type,
actor,
resource_type,
resource_id,
message,
metadata,
ip_address,
)
.await
}
pub async fn list_audit_events(
&self,
filter: &crate::audit::AuditFilter,
limit: u32,
offset: u32,
) -> Result<Vec<nx9_wg_core::types::audit::AuditEvent>> {
crate::audit::list_audit_events(&self.pool, filter, limit, offset).await
}
pub async fn get_audit_event(
&self,
id: i64,
) -> Result<Option<nx9_wg_core::types::audit::AuditEvent>> {
crate::audit::get_audit_event(&self.pool, id).await
}
pub async fn count_audit_events(&self, filter: &crate::audit::AuditFilter) -> Result<i64> {
crate::audit::count_audit_events(&self.pool, filter).await
}
// Backups
pub async fn create_backup_meta(
&self,
meta: &nx9_wg_core::types::backup::BackupMeta,
) -> Result<()> {
crate::backups::create_backup_meta(&self.pool, meta).await
}
pub async fn get_backup_meta(
&self,
id: uuid::Uuid,
) -> Result<Option<nx9_wg_core::types::backup::BackupMeta>> {
crate::backups::get_backup_meta(&self.pool, id).await
}
pub async fn list_backups(&self) -> Result<Vec<nx9_wg_core::types::backup::BackupMeta>> {
crate::backups::list_backups(&self.pool).await
}
pub async fn delete_backup_meta(&self, id: uuid::Uuid) -> Result<()> {
crate::backups::delete_backup_meta(&self.pool, id).await
}
pub async fn vacuum_into(&self, target_file_path: &str) -> Result<()> {
crate::backups::vacuum_into(&self.pool, target_file_path).await
}
// Client Profiles
pub async fn create_client_profile(
&self,
profile: &nx9_wg_core::types::client_profile::ClientProfile,
) -> Result<()> {
crate::client_profiles::create_client_profile(&self.pool, profile).await
}
pub async fn get_client_profile(
&self,
id: &str,
) -> Result<Option<nx9_wg_core::types::client_profile::ClientProfile>> {
crate::client_profiles::get_client_profile(&self.pool, id).await
}
pub async fn list_client_profiles(
&self,
) -> Result<Vec<nx9_wg_core::types::client_profile::ClientProfile>> {
crate::client_profiles::list_client_profiles(&self.pool).await
}
pub async fn update_client_profile(
&self,
profile: &nx9_wg_core::types::client_profile::ClientProfile,
) -> Result<()> {
crate::client_profiles::update_client_profile(&self.pool, profile).await
}
pub async fn delete_client_profile(&self, id: &str) -> Result<()> {
crate::client_profiles::delete_client_profile(&self.pool, id).await
}
pub async fn list_distinct_providers(&self) -> Result<Vec<String>> {
crate::client_profiles::list_distinct_providers(&self.pool).await
}
pub async fn find_matching_client_profiles(
&self,
provider: Option<&str>,
device: Option<nx9_wg_core::types::client_profile::DeviceCategory>,
connection: Option<nx9_wg_core::types::client_profile::ConnectionType>,
nat: Option<nx9_wg_core::types::client_profile::NatType>,
) -> Result<Vec<nx9_wg_core::types::client_profile::ClientProfile>> {
crate::client_profiles::find_matching_profiles(
&self.pool, provider, device, connection, nat,
)
.await
}
}
+278
View File
@@ -0,0 +1,278 @@
//! API Token repository operations.
use crate::error::{DbError, Result};
use crate::models::{format_datetime, parse_datetime};
use chrono::Utc;
use nx9_wg_core::types::auth::ApiToken;
use sqlx::{Row, SqlitePool};
/// Create a new API token record. Only the token hash is stored.
pub async fn create_token(pool: &SqlitePool, token: &ApiToken) -> Result<()> {
let created_at_str = format_datetime(&token.created_at);
let expires_at_str = token.expires_at.as_ref().map(format_datetime);
let last_used_str = token.last_used_at.as_ref().map(format_datetime);
let revoked_str = token.revoked_at.as_ref().map(format_datetime);
sqlx::query(
r#"
INSERT INTO api_tokens (id, admin_id, name, token_hash, created_at, expires_at, last_used_at, revoked_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&token.id)
.bind(token.admin_id)
.bind(&token.name)
.bind(&token.token_hash)
.bind(&created_at_str)
.bind(expires_at_str)
.bind(last_used_str)
.bind(revoked_str)
.execute(pool)
.await
.map_err(|e| match &e {
sqlx::Error::Database(dbe) if dbe.is_unique_violation() => {
DbError::Conflict("API token with this hash already exists".to_string())
}
_ => DbError::Sqlx(e),
})?;
Ok(())
}
/// Retrieve all API tokens.
pub async fn list_tokens(pool: &SqlitePool) -> Result<Vec<ApiToken>> {
let rows = sqlx::query(
r#"
SELECT id, admin_id, name, token_hash, created_at, expires_at, last_used_at, revoked_at
FROM api_tokens
ORDER BY created_at DESC
"#,
)
.fetch_all(pool)
.await
.map_err(DbError::Sqlx)?;
let mut tokens = Vec::with_capacity(rows.len());
for r in rows {
let id: String = r.try_get("id")?;
let admin_id: i64 = r.try_get("admin_id")?;
let name: String = r.try_get("name")?;
let token_hash: String = r.try_get("token_hash")?;
let created_at_str: String = r.try_get("created_at")?;
let expires_at_str: Option<String> = r.try_get("expires_at")?;
let last_used_str: Option<String> = r.try_get("last_used_at")?;
let revoked_str: Option<String> = r.try_get("revoked_at")?;
let expires_at = match expires_at_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
let last_used_at = match last_used_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
let revoked_at = match revoked_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
tokens.push(ApiToken {
id,
admin_id,
name,
token_hash,
created_at: parse_datetime(&created_at_str)?,
expires_at,
last_used_at,
revoked_at,
revoked: revoked_at.is_some(),
});
}
Ok(tokens)
}
/// Retrieve an API token by ID.
pub async fn get_token(pool: &SqlitePool, id: &str) -> Result<Option<ApiToken>> {
let row = sqlx::query(
r#"
SELECT id, admin_id, name, token_hash, created_at, expires_at, last_used_at, revoked_at
FROM api_tokens
WHERE id = ?
"#,
)
.bind(id)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => {
let id: String = r.try_get("id")?;
let admin_id: i64 = r.try_get("admin_id")?;
let name: String = r.try_get("name")?;
let token_hash: String = r.try_get("token_hash")?;
let created_at_str: String = r.try_get("created_at")?;
let expires_at_str: Option<String> = r.try_get("expires_at")?;
let last_used_str: Option<String> = r.try_get("last_used_at")?;
let revoked_str: Option<String> = r.try_get("revoked_at")?;
let expires_at = match expires_at_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
let last_used_at = match last_used_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
let revoked_at = match revoked_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
Ok(Some(ApiToken {
id,
admin_id,
name,
token_hash,
created_at: parse_datetime(&created_at_str)?,
expires_at,
last_used_at,
revoked_at,
revoked: revoked_at.is_some(),
}))
}
None => Ok(None),
}
}
/// Look up an active (non-revoked) API token by its SHA-256 hash.
pub async fn find_token_by_hash(pool: &SqlitePool, token_hash: &str) -> Result<Option<ApiToken>> {
let row = sqlx::query(
r#"
SELECT id, admin_id, name, token_hash, created_at, expires_at, last_used_at, revoked_at
FROM api_tokens
WHERE token_hash = ?
"#,
)
.bind(token_hash)
.fetch_optional(pool)
.await
.map_err(DbError::Sqlx)?;
match row {
Some(r) => {
let id: String = r.try_get("id")?;
let admin_id: i64 = r.try_get("admin_id")?;
let name: String = r.try_get("name")?;
let token_hash: String = r.try_get("token_hash")?;
let created_at_str: String = r.try_get("created_at")?;
let expires_at_str: Option<String> = r.try_get("expires_at")?;
let last_used_str: Option<String> = r.try_get("last_used_at")?;
let revoked_str: Option<String> = r.try_get("revoked_at")?;
let expires_at = match expires_at_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
let last_used_at = match last_used_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
let revoked_at = match revoked_str {
Some(s) => Some(parse_datetime(&s)?),
None => None,
};
Ok(Some(ApiToken {
id,
admin_id,
name,
token_hash,
created_at: parse_datetime(&created_at_str)?,
expires_at,
last_used_at,
revoked_at,
revoked: revoked_at.is_some(),
}))
}
None => Ok(None),
}
}
/// Mark a token as used at current timestamp.
pub async fn mark_token_used(pool: &SqlitePool, id: &str) -> Result<()> {
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE api_tokens
SET last_used_at = ?
WHERE id = ?
"#,
)
.bind(&now_str)
.bind(id)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!("API token '{id}' not found")));
}
Ok(())
}
/// Revoke an API token.
pub async fn revoke_token(pool: &SqlitePool, id: &str) -> Result<()> {
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result = sqlx::query(
r#"
UPDATE api_tokens
SET revoked_at = ?
WHERE id = ? AND revoked_at IS NULL
"#,
)
.bind(&now_str)
.bind(id)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
if result.rows_affected() == 0 {
return Err(DbError::NotFound(format!(
"API token '{id}' not found or already revoked"
)));
}
Ok(())
}
/// Delete an API token.
pub async fn delete_token(pool: &SqlitePool, id: &str) -> Result<()> {
sqlx::query("DELETE FROM api_tokens WHERE id = ?")
.bind(id)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(())
}
/// Delete expired tokens. Returns count deleted.
pub async fn delete_expired_tokens(pool: &SqlitePool) -> Result<u64> {
let now = Utc::now().naive_utc();
let now_str = format_datetime(&now);
let result =
sqlx::query("DELETE FROM api_tokens WHERE expires_at IS NOT NULL AND expires_at < ?")
.bind(&now_str)
.execute(pool)
.await
.map_err(DbError::Sqlx)?;
Ok(result.rows_affected())
}
@@ -0,0 +1,80 @@
//! Tests for Administrator repository operations and security invariants.
use nx9_wg_core::crypto::{hash_password, verify_password};
use nx9_wg_db::Store;
#[tokio::test]
async fn test_admin_single_identity_and_crud() {
let store = Store::connect_in_memory().await.expect("connect");
store.migrate().await.expect("migrate");
// Initially no admin exists
assert!(!store.admin_exists().await.expect("admin_exists"));
assert!(store.get_admin().await.expect("get_admin").is_none());
// Create single admin with Argon2id hash
let password = "CorrectHorseBatteryStaple123!";
let password_hash = hash_password(password).expect("hash password");
let admin = store
.create_admin("admin", &password_hash)
.await
.expect("create_admin");
assert_eq!(admin.id, 1);
assert_eq!(admin.username, "admin");
assert!(!admin.totp_enabled);
assert!(admin.last_login_at.is_none());
// Verify admin_exists returns true
assert!(store.admin_exists().await.expect("admin_exists"));
// Verify lookup by username
let fetched = store
.get_admin_by_username("admin")
.await
.expect("get_admin_by_username")
.expect("admin found");
assert_eq!(fetched.id, 1);
assert!(verify_password(password, &fetched.password_hash).expect("verify password"));
// Reject second admin creation
let second_res = store.create_admin("admin2", "hash2").await;
assert!(second_res.is_err(), "second admin must be rejected");
// Test password change
let new_password = "NewSuperSecurePassword456!";
let new_hash = hash_password(new_password).expect("new hash");
store
.update_admin_password(&new_hash)
.await
.expect("update_admin_password");
let updated = store.get_admin().await.expect("get_admin").expect("admin");
assert!(verify_password(new_password, &updated.password_hash).expect("verify new"));
assert!(!verify_password(password, &updated.password_hash).expect("old password fails"));
// Test TOTP update
store
.update_admin_totp(Some("JBSWY3DPEHPK3PXP"), true)
.await
.expect("update_admin_totp");
let totp_admin = store.get_admin().await.expect("get_admin").expect("admin");
assert!(totp_admin.totp_enabled);
assert_eq!(totp_admin.totp_secret.as_deref(), Some("JBSWY3DPEHPK3PXP"));
// Test recording login
store
.record_admin_login(Some("192.168.1.100"))
.await
.expect("record_admin_login");
let login_admin = store.get_admin().await.expect("get_admin").expect("admin");
assert!(login_admin.last_login_at.is_some());
assert_eq!(login_admin.last_login_ip.as_deref(), Some("192.168.1.100"));
// Verify Debug formatting redacts password_hash and totp_secret
let debug_str = format!("{:?}", login_admin);
assert!(debug_str.contains("[REDACTED]"));
assert!(!debug_str.contains(password));
assert!(!debug_str.contains(new_password));
assert!(!debug_str.contains("JBSWY3DPEHPK3PXP"));
}
@@ -0,0 +1,177 @@
//! Tests for Session and API Token repository operations.
use chrono::{Duration, Utc};
use nx9_wg_core::crypto::{generate_api_token, hash_password};
use nx9_wg_core::types::auth::{ApiToken, Session};
use nx9_wg_db::Store;
use uuid::Uuid;
#[tokio::test]
async fn test_session_lifecycle() {
let store = Store::connect_in_memory().await.expect("connect");
store.migrate().await.expect("migrate");
let pw_hash = hash_password("AdminPass123!").expect("hash");
store.create_admin("admin", &pw_hash).await.expect("admin");
let now = Utc::now().naive_utc();
let session_id = Uuid::new_v4().to_string();
let session = Session {
id: session_id.clone(),
admin_id: 1,
created_at: now,
expires_at: now + Duration::hours(24),
last_seen_at: Some(now),
ip_address: Some("10.0.0.5".to_string()),
user_agent: Some("TestAgent/1.0".to_string()),
};
store
.create_session(&session)
.await
.expect("create_session");
let fetched = store
.get_session(&session_id)
.await
.expect("get_session")
.expect("session found");
assert_eq!(fetched.id, session_id);
assert_eq!(fetched.admin_id, 1);
assert_eq!(fetched.ip_address.as_deref(), Some("10.0.0.5"));
// Touch session
store
.touch_session(&session_id)
.await
.expect("touch_session");
// Test delete expired sessions
let expired_id = Uuid::new_v4().to_string();
let expired_session = Session {
id: expired_id.clone(),
admin_id: 1,
created_at: now - Duration::hours(48),
expires_at: now - Duration::hours(24),
last_seen_at: None,
ip_address: None,
user_agent: None,
};
store
.create_session(&expired_session)
.await
.expect("expired session");
let deleted = store
.delete_expired_sessions()
.await
.expect("delete expired");
assert_eq!(deleted, 1);
assert!(store.get_session(&expired_id).await.expect("get").is_none());
assert!(store.get_session(&session_id).await.expect("get").is_some());
// Delete single session
store
.delete_session(&session_id)
.await
.expect("delete session");
assert!(store.get_session(&session_id).await.expect("get").is_none());
// Test delete_all_admin_sessions
let s1 = Session {
id: "s1".to_string(),
admin_id: 1,
created_at: now,
expires_at: now + Duration::hours(1),
last_seen_at: None,
ip_address: None,
user_agent: None,
};
let s2 = Session {
id: "s2".to_string(),
admin_id: 1,
created_at: now,
expires_at: now + Duration::hours(1),
last_seen_at: None,
ip_address: None,
user_agent: None,
};
store.create_session(&s1).await.expect("s1");
store.create_session(&s2).await.expect("s2");
let deleted_all = store
.delete_all_admin_sessions(1)
.await
.expect("delete all");
assert_eq!(deleted_all, 2);
}
#[tokio::test]
async fn test_api_token_lifecycle() {
let store = Store::connect_in_memory().await.expect("connect");
store.migrate().await.expect("migrate");
let pw_hash = hash_password("AdminPass123!").expect("hash");
store.create_admin("admin", &pw_hash).await.expect("admin");
let (raw_token, token_hash) = generate_api_token();
let token_id = Uuid::new_v4().to_string();
let now = Utc::now().naive_utc();
let token = ApiToken {
id: token_id.clone(),
admin_id: 1,
name: "CI/CD Deployment Token".to_string(),
token_hash: token_hash.clone(),
created_at: now,
expires_at: Some(now + Duration::days(30)),
last_used_at: None,
revoked_at: None,
revoked: false,
};
store.create_token(&token).await.expect("create_token");
// Lookup by hash
let found = store
.find_token_by_hash(&token_hash)
.await
.expect("find by hash")
.expect("token found");
assert_eq!(found.id, token_id);
assert_eq!(found.name, "CI/CD Deployment Token");
assert!(!found.revoked);
// Verify raw token is never in the stored record
let debug_out = format!("{:?}", found);
assert!(debug_out.contains("[REDACTED]"));
assert!(!debug_out.contains(&raw_token));
// Mark token used
store.mark_token_used(&token_id).await.expect("mark used");
let after_use = store
.get_token(&token_id)
.await
.expect("get")
.expect("token");
assert!(after_use.last_used_at.is_some());
// List tokens
let tokens = store.list_tokens().await.expect("list tokens");
assert_eq!(tokens.len(), 1);
// Revoke token
store.revoke_token(&token_id).await.expect("revoke token");
let revoked = store
.get_token(&token_id)
.await
.expect("get")
.expect("token");
assert!(revoked.revoked);
assert!(revoked.revoked_at.is_some());
// Delete token
store.delete_token(&token_id).await.expect("delete token");
assert!(store.get_token(&token_id).await.expect("get").is_none());
}
@@ -0,0 +1,106 @@
//! Tests for client profile repository operations and built-in profiles.
use nx9_wg_core::types::client_profile::{ClientProfile, ConnectionType, DeviceCategory, NatType};
use nx9_wg_db::Store;
#[tokio::test]
async fn test_client_profiles_crud_and_builtin_protection() {
let store = Store::connect_in_memory().await.unwrap();
store.migrate().await.unwrap();
// Verify built-in profiles pre-populated by migration
let profiles = store.list_client_profiles().await.unwrap();
assert!(
profiles.len() >= 10,
"expected at least 10 built-in profiles"
);
// Check specific built-ins
let mobile = store.get_client_profile("default-mobile").await.unwrap();
assert!(mobile.is_some());
let mobile = mobile.unwrap();
assert_eq!(mobile.connection_type, ConnectionType::Mobile);
assert_eq!(mobile.mtu, 1280);
assert_eq!(mobile.persistent_keepalive, Some(25));
assert!(mobile.is_builtin);
let cgnat = store.get_client_profile("default-cgnat").await.unwrap();
assert!(cgnat.is_some());
let cgnat = cgnat.unwrap();
assert_eq!(cgnat.nat_type, NatType::Cgnat);
assert_eq!(cgnat.mtu, 1360);
// Verify built-in cannot be modified or deleted
let mut modified_builtin = mobile.clone();
modified_builtin.mtu = 1400;
assert!(
store
.update_client_profile(&modified_builtin)
.await
.is_err()
);
assert!(store.delete_client_profile("default-mobile").await.is_err());
// Create a custom profile
let now = chrono::Utc::now().naive_utc();
let custom = ClientProfile {
id: "office-fiber".to_string(),
name: "Office Fiber Direct".to_string(),
provider: Some("att".to_string()),
device: Some(DeviceCategory::Linux),
connection_type: ConnectionType::Wired,
nat_type: NatType::Direct,
mtu: 1420,
dns: Some("1.1.1.1, 1.0.0.1".to_string()),
persistent_keepalive: Some(15),
is_builtin: false,
description: Some("Direct fiber connection at headquarters".to_string()),
created_at: now,
updated_at: now,
};
store.create_client_profile(&custom).await.unwrap();
let fetched = store
.get_client_profile("office-fiber")
.await
.unwrap()
.unwrap();
assert_eq!(fetched.name, "Office Fiber Direct");
assert_eq!(fetched.provider.as_deref(), Some("att"));
assert_eq!(fetched.mtu, 1420);
assert!(!fetched.is_builtin);
// Update custom profile
let mut updated = fetched.clone();
updated.description = Some("Updated headquarters fiber".to_string());
updated.mtu = 1440;
store.update_client_profile(&updated).await.unwrap();
let fetched_updated = store
.get_client_profile("office-fiber")
.await
.unwrap()
.unwrap();
assert_eq!(fetched_updated.mtu, 1440);
assert_eq!(
fetched_updated.description.as_deref(),
Some("Updated headquarters fiber")
);
// List distinct providers
let providers = store.list_distinct_providers().await.unwrap();
assert!(providers.contains(&"att".to_string()));
assert!(providers.contains(&"tmobile".to_string()));
assert!(providers.contains(&"starlink".to_string()));
// Delete custom profile
store.delete_client_profile("office-fiber").await.unwrap();
assert!(
store
.get_client_profile("office-fiber")
.await
.unwrap()
.is_none()
);
}
@@ -0,0 +1,288 @@
//! Tests for Network, Route, and Firewall Rule repositories.
use chrono::Utc;
use ipnet::IpNet;
use nx9_wg_core::crypto::generate_keypair;
use nx9_wg_core::types::firewall::{
FirewallAction, FirewallDirection, FirewallProtocol, FirewallRule,
};
use nx9_wg_core::types::network::{Network, Route};
use nx9_wg_core::types::wireguard::Interface;
use nx9_wg_db::Store;
use std::net::IpAddr;
use std::str::FromStr;
use uuid::Uuid;
#[tokio::test]
async fn test_network_and_route_crud() {
let store = Store::connect_in_memory().await.expect("connect");
store.migrate().await.expect("migrate");
let now = Utc::now().naive_utc();
let net_id = Uuid::new_v4();
let net = Network {
id: net_id,
name: "Home Lab".to_string(),
cidr: IpNet::from_str("192.168.10.0/24").expect("cidr"),
enabled: true,
description: Some("Internal lab subnet".to_string()),
created_at: now,
updated_at: now,
};
store.create_network(&net).await.expect("create_network");
let fetched_net = store
.get_network(net_id)
.await
.expect("get")
.expect("found");
assert_eq!(fetched_net.name, "Home Lab");
assert_eq!(fetched_net.cidr.to_string(), "192.168.10.0/24");
assert!(fetched_net.enabled);
// Test routes
let route_id = Uuid::new_v4();
let route = Route {
id: route_id,
network_id: Some(net_id),
interface_id: None,
destination: IpNet::from_str("192.168.10.0/24").expect("dest cidr"),
gateway: Some(IpAddr::from_str("10.0.0.1").expect("gateway")),
interface_name: None,
metric: Some(100),
enabled: true,
description: Some("Lab route via wg gateway".to_string()),
created_at: now,
updated_at: now,
};
store.create_route(&route).await.expect("create_route");
let fetched_route = store
.get_route(route_id)
.await
.expect("get")
.expect("route found");
assert_eq!(fetched_route.network_id, Some(net_id));
assert_eq!(
fetched_route.gateway,
Some(IpAddr::from_str("10.0.0.1").unwrap())
);
assert_eq!(fetched_route.metric, Some(100));
// Enable/disable route
store
.set_route_enabled(route_id, false)
.await
.expect("disable");
let disabled_route = store
.get_route(route_id)
.await
.expect("get")
.expect("route");
assert!(!disabled_route.enabled);
// List routes for network
let net_routes = store
.list_routes_for_network(net_id)
.await
.expect("list net routes");
assert_eq!(net_routes.len(), 1);
// Deleting network sets route's network_id to NULL (ON DELETE SET NULL)
store.delete_network(net_id).await.expect("delete network");
let route_after_net_delete = store
.get_route(route_id)
.await
.expect("get")
.expect("route");
assert!(
route_after_net_delete.network_id.is_none(),
"network_id must be SET NULL when network is deleted"
);
}
#[tokio::test]
async fn test_firewall_rule_crud_and_priority_ordering() {
let store = Store::connect_in_memory().await.expect("connect");
store.migrate().await.expect("migrate");
let now = Utc::now().naive_utc();
let iface_id = Uuid::new_v4();
let (priv_k, pub_k) = generate_keypair();
let iface = Interface {
id: iface_id,
name: "wg0".to_string(),
private_key: priv_k,
public_key: pub_k,
listen_port: 51820,
address_v4: IpNet::from_str("10.0.0.1/24").unwrap(),
address_v6: None,
mtu: None,
dns: None,
enabled: true,
pre_up: None,
post_up: None,
pre_down: None,
post_down: None,
created_at: now,
updated_at: now,
};
store
.create_interface(&iface)
.await
.expect("create interface");
let rule1_id = Uuid::new_v4();
let rule1 = FirewallRule {
id: rule1_id,
name: "Allow SSH".to_string(),
interface_id: Some(iface_id),
peer_id: None,
direction: FirewallDirection::In,
action: FirewallAction::Accept,
protocol: FirewallProtocol::Tcp,
source: None,
destination: None,
source_port: None,
destination_port: Some(22),
port_range: None,
priority: 50,
enabled: true,
description: Some("SSH access".to_string()),
created_at: now,
updated_at: now,
};
let rule2_id = Uuid::new_v4();
let rule2 = FirewallRule {
id: rule2_id,
name: "Drop All Other".to_string(),
interface_id: Some(iface_id),
peer_id: None,
direction: FirewallDirection::In,
action: FirewallAction::Drop,
protocol: FirewallProtocol::Any,
source: None,
destination: None,
source_port: None,
destination_port: None,
port_range: None,
priority: 100,
enabled: true,
description: Some("Default drop".to_string()),
created_at: now,
updated_at: now,
};
store
.create_firewall_rule(&rule2)
.await
.expect("create rule2");
store
.create_firewall_rule(&rule1)
.await
.expect("create rule1");
// List rules should order by priority ASC (rule1 priority 50 comes before rule2 priority 100)
let rules = store.list_firewall_rules().await.expect("list rules");
assert_eq!(rules.len(), 2);
assert_eq!(rules[0].id, rule1_id);
assert_eq!(rules[0].priority, 50);
assert_eq!(rules[1].id, rule2_id);
assert_eq!(rules[1].priority, 100);
// List rules for interface
let iface_rules = store
.list_firewall_rules_for_interface(iface_id)
.await
.expect("list iface rules");
assert_eq!(iface_rules.len(), 2);
// Enable/disable rule
store
.set_firewall_rule_enabled(rule1_id, false)
.await
.expect("disable");
let disabled = store
.get_firewall_rule(rule1_id)
.await
.expect("get")
.expect("rule");
assert!(!disabled.enabled);
// Delete rule
store.delete_firewall_rule(rule1_id).await.expect("delete");
assert!(
store
.get_firewall_rule(rule1_id)
.await
.expect("get")
.is_none()
);
// Peer-specific rule with port range
let peer_id = Uuid::new_v4();
let peer = nx9_wg_core::types::wireguard::Peer {
id: peer_id,
interface_id: iface_id,
name: "test-peer-fw".to_string(),
peer_type: nx9_wg_core::types::wireguard::PeerType::RoadWarrior,
state: nx9_wg_core::types::wireguard::PeerState::Active,
public_key: nx9_wg_core::types::wireguard::WireGuardPublicKey::new(
"testpubkey12345678901234567890123456789012=".to_string(),
),
private_key: None,
preshared_key: None,
endpoint: None,
allowed_ips: "10.0.0.2/32".to_string(),
server_allowed_ips: None,
address_v4: Some("10.0.0.2/32".parse().unwrap()),
address_v6: None,
dns: None,
mtu: None,
persistent_keepalive: None,
profile: nx9_wg_core::types::wireguard::PeerProfile::FullTunnel,
expires_at: None,
last_handshake_at: None,
created_at: now,
updated_at: now,
};
store.create_peer(&peer).await.expect("create peer");
let peer_rule_id = Uuid::new_v4();
let peer_rule = FirewallRule {
id: peer_rule_id,
name: "Peer Port Range Rule".to_string(),
interface_id: Some(iface_id),
peer_id: Some(peer_id),
direction: FirewallDirection::In,
action: FirewallAction::Accept,
protocol: FirewallProtocol::TcpUdp,
source: None,
destination: None,
source_port: None,
destination_port: None,
port_range: Some("8000-8100".to_string()),
priority: 25,
enabled: true,
description: Some("Custom peer range".to_string()),
created_at: now,
updated_at: now,
};
store
.create_firewall_rule(&peer_rule)
.await
.expect("create peer rule");
let peer_rules = store
.list_firewall_rules_for_peer(peer_id)
.await
.expect("list peer rules");
assert_eq!(peer_rules.len(), 1);
assert_eq!(peer_rules[0].port_range.as_deref(), Some("8000-8100"));
assert_eq!(peer_rules[0].protocol, FirewallProtocol::TcpUdp);
}
@@ -0,0 +1,52 @@
//! Tests for Store initialization, WAL configuration, migrations, and SQLite invariants.
use nx9_wg_db::Store;
use sqlx::Row;
use tempfile::NamedTempFile;
#[tokio::test]
async fn test_in_memory_store_lifecycle() {
let store = Store::connect_in_memory().await.expect("connect in-memory");
store.migrate().await.expect("run migrations");
// Verify foreign keys are enabled
let row = sqlx::query("PRAGMA foreign_keys")
.fetch_one(store.pool())
.await
.expect("pragma foreign_keys");
let fk: i64 = row.get(0);
assert_eq!(fk, 1, "foreign keys must be enabled");
}
#[tokio::test]
async fn test_temp_file_store_wal_mode() {
let tmp = NamedTempFile::new().expect("temp file");
let path = tmp.path();
let store = Store::connect_path(path).await.expect("connect path");
store.migrate().await.expect("run migrations");
// Verify WAL mode is configured
let row = sqlx::query("PRAGMA journal_mode")
.fetch_one(store.pool())
.await
.expect("pragma journal_mode");
let mode: String = row.get(0);
assert_eq!(mode.to_lowercase(), "wal", "WAL mode must be active");
// Verify migrations table exists and records the initial migration
let migration_count_row = sqlx::query("SELECT COUNT(*) FROM _sqlx_migrations")
.fetch_one(store.pool())
.await
.expect("query migrations");
let count: i64 = migration_count_row.get(0);
assert!(count >= 1, "at least one migration should be recorded");
}
#[tokio::test]
async fn test_migration_idempotence() {
let store = Store::connect_in_memory().await.expect("connect in-memory");
store.migrate().await.expect("first migration run");
// Running migrate a second time must succeed idempotently
store.migrate().await.expect("second migration run");
}
@@ -0,0 +1,227 @@
//! Tests for Settings, Audit Log, and Backup repositories.
use chrono::Utc;
use nx9_wg_core::types::audit::AuditEventType;
use nx9_wg_core::types::backup::BackupMeta;
use nx9_wg_db::{AuditFilter, Store};
use uuid::Uuid;
#[tokio::test]
async fn test_settings_repository() {
let store = Store::connect_in_memory().await.expect("connect");
store.migrate().await.expect("migrate");
// Initially missing key returns None
assert!(
store
.get_setting("non_existent")
.await
.expect("get")
.is_none()
);
assert!(
store
.get_setting_value("non_existent")
.await
.expect("get val")
.is_none()
);
// Set normal setting
store
.set_setting("server_endpoint", "vpn.example.com:51820", false)
.await
.expect("set");
let ep = store
.get_setting("server_endpoint")
.await
.expect("get")
.expect("setting found");
assert_eq!(ep.value, "vpn.example.com:51820");
assert!(!ep.is_secret);
// Set secret setting
store
.set_setting("session_secret", "SuperSecretKey999", true)
.await
.expect("set secret");
let sec = store
.get_setting("session_secret")
.await
.expect("get")
.expect("setting found");
assert_eq!(sec.value, "SuperSecretKey999");
assert!(sec.is_secret);
// Verify Debug formatting of secret setting redacts value
let sec_debug = format!("{:?}", sec);
assert!(sec_debug.contains("[REDACTED]"));
assert!(!sec_debug.contains("SuperSecretKey999"));
// Upsert existing setting
store
.set_setting("server_endpoint", "vpn2.example.com:51820", false)
.await
.expect("upsert");
let ep2 = store
.get_setting_value("server_endpoint")
.await
.expect("get")
.expect("value found");
assert_eq!(ep2, "vpn2.example.com:51820");
// List settings
let all = store.list_settings().await.expect("list");
assert_eq!(all.len(), 2);
// Delete setting
store
.delete_setting("server_endpoint")
.await
.expect("delete");
assert!(
store
.get_setting("server_endpoint")
.await
.expect("get")
.is_none()
);
}
#[tokio::test]
async fn test_audit_log_append_only_and_filtering() {
let store = Store::connect_in_memory().await.expect("connect");
store.migrate().await.expect("migrate");
// Record various events
store
.record_audit(
AuditEventType::Login,
"admin",
Some("session"),
Some("sess-1"),
Some("Admin login succeeded"),
None,
Some("192.168.1.50"),
)
.await
.expect("record login");
store
.record_audit(
AuditEventType::InterfaceCreate,
"admin",
Some("interface"),
Some("wg0"),
Some("Interface wg0 created"),
None,
Some("192.168.1.50"),
)
.await
.expect("record iface create");
store
.record_audit(
AuditEventType::PeerCreate,
"admin",
Some("peer"),
Some("peer-alice"),
Some("Peer alice created"),
None,
Some("192.168.1.50"),
)
.await
.expect("record peer create");
// Total count
let total = store
.count_audit_events(&AuditFilter::default())
.await
.expect("count");
assert_eq!(total, 3);
// Filter by event_type
let login_filter = AuditFilter {
event_type: Some(AuditEventType::Login),
..Default::default()
};
let login_events = store
.list_audit_events(&login_filter, 10, 0)
.await
.expect("list login");
assert_eq!(login_events.len(), 1);
assert_eq!(login_events[0].event_type, AuditEventType::Login);
// Filter by resource_type
let peer_filter = AuditFilter {
resource_type: Some("peer".to_string()),
..Default::default()
};
let peer_events = store
.list_audit_events(&peer_filter, 10, 0)
.await
.expect("list peer events");
assert_eq!(peer_events.len(), 1);
assert_eq!(peer_events[0].resource_id.as_deref(), Some("peer-alice"));
// Pagination test: limit 2, offset 0 -> 2 items; offset 2 -> 1 item
let page1 = store
.list_audit_events(&AuditFilter::default(), 2, 0)
.await
.expect("page1");
assert_eq!(page1.len(), 2);
let page2 = store
.list_audit_events(&AuditFilter::default(), 2, 2)
.await
.expect("page2");
assert_eq!(page2.len(), 1);
}
#[tokio::test]
async fn test_backup_metadata_crud() {
let store = Store::connect_in_memory().await.expect("connect");
store.migrate().await.expect("migrate");
let now = Utc::now().naive_utc();
let backup_id = Uuid::new_v4();
let meta = BackupMeta {
id: backup_id,
filename: "nx9-wg-backup-20260816.tar.gz".to_string(),
size_bytes: 1048576,
checksum: "sha256:e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"
.to_string(),
schema_version: "1".to_string(),
encrypted: true,
description: Some("Automated nightly backup".to_string()),
created_at: now,
};
store
.create_backup_meta(&meta)
.await
.expect("create_backup_meta");
let fetched = store
.get_backup_meta(backup_id)
.await
.expect("get")
.expect("backup found");
assert_eq!(fetched.filename, "nx9-wg-backup-20260816.tar.gz");
assert_eq!(fetched.size_bytes, 1048576);
assert!(fetched.encrypted);
assert_eq!(fetched.schema_version, "1");
let list = store.list_backups().await.expect("list");
assert_eq!(list.len(), 1);
store.delete_backup_meta(backup_id).await.expect("delete");
assert!(
store
.get_backup_meta(backup_id)
.await
.expect("get")
.is_none()
);
}
@@ -0,0 +1,246 @@
//! Tests for WireGuard Interface and Peer repository operations.
use chrono::Utc;
use ipnet::IpNet;
use nx9_wg_core::crypto::{generate_keypair, generate_preshared_key};
use nx9_wg_core::types::wireguard::{
Interface, Peer, PeerProfile, PeerState, PeerType, WireGuardPublicKey,
};
use nx9_wg_db::Store;
use std::str::FromStr;
use uuid::Uuid;
#[tokio::test]
async fn test_interface_and_peer_crud_and_cascade() {
let store = Store::connect_in_memory().await.expect("connect");
store.migrate().await.expect("migrate");
let now = Utc::now().naive_utc();
let iface_id = Uuid::new_v4();
let (priv_k, pub_k) = generate_keypair();
let iface = Interface {
id: iface_id,
name: "wg0".to_string(),
private_key: priv_k.clone(),
public_key: pub_k.clone(),
listen_port: 51820,
address_v4: IpNet::from_str("10.0.0.1/24").expect("valid cidr"),
address_v6: Some(IpNet::from_str("fd00::1/64").expect("valid cidr")),
mtu: Some(1420),
dns: Some("1.1.1.1, 8.8.8.8".to_string()),
enabled: true,
pre_up: None,
post_up: Some("iptables -t nat -A POSTROUTING -o eth0 -j MASQUERADE".to_string()),
pre_down: None,
post_down: None,
created_at: now,
updated_at: now,
};
store
.create_interface(&iface)
.await
.expect("create_interface");
// Lookup interface by ID and name
let fetched = store
.get_interface(iface_id)
.await
.expect("get_interface")
.expect("iface found");
assert_eq!(fetched.name, "wg0");
assert_eq!(fetched.listen_port, 51820);
assert_eq!(fetched.address_v4.to_string(), "10.0.0.1/24");
assert_eq!(fetched.mtu, Some(1420));
let by_name = store
.get_interface_by_name("wg0")
.await
.expect("get_by_name")
.expect("found");
assert_eq!(by_name.id, iface_id);
// Reject duplicate interface name
let dup_iface = Interface {
id: Uuid::new_v4(),
name: "wg0".to_string(),
private_key: priv_k.clone(),
public_key: pub_k.clone(),
listen_port: 51821,
address_v4: IpNet::from_str("10.0.1.1/24").unwrap(),
address_v6: None,
mtu: None,
dns: None,
enabled: true,
pre_up: None,
post_up: None,
pre_down: None,
post_down: None,
created_at: now,
updated_at: now,
};
assert!(
store.create_interface(&dup_iface).await.is_err(),
"duplicate interface name must fail"
);
// Create a peer
let peer_id = Uuid::new_v4();
let (peer_priv, peer_pub) = generate_keypair();
let psk = generate_preshared_key();
let peer = Peer {
id: peer_id,
interface_id: iface_id,
name: "phone-alice".to_string(),
peer_type: PeerType::RoadWarrior,
state: PeerState::Active,
public_key: peer_pub.clone(),
private_key: Some(peer_priv.clone()),
preshared_key: Some(psk.clone()),
endpoint: None,
allowed_ips: "10.0.0.2/32".to_string(),
server_allowed_ips: Some("10.0.0.2/32".to_string()),
address_v4: Some(IpNet::from_str("10.0.0.2/32").unwrap()),
address_v6: None,
dns: Some("10.0.0.1".to_string()),
mtu: Some(1420),
persistent_keepalive: Some(25),
profile: PeerProfile::FullTunnel,
expires_at: None,
last_handshake_at: None,
created_at: now,
updated_at: now,
};
store.create_peer(&peer).await.expect("create_peer");
// Fetch peer
let fetched_peer = store
.get_peer(peer_id)
.await
.expect("get_peer")
.expect("peer found");
assert_eq!(fetched_peer.name, "phone-alice");
assert_eq!(fetched_peer.peer_type, PeerType::RoadWarrior);
assert_eq!(fetched_peer.state, PeerState::Active);
assert_eq!(fetched_peer.profile, PeerProfile::FullTunnel);
assert_eq!(fetched_peer.allowed_ips, "10.0.0.2/32");
assert_eq!(fetched_peer.persistent_keepalive, Some(25));
// Lookup peer by name and by public key
let by_pname = store
.get_peer_by_name(iface_id, "phone-alice")
.await
.expect("by name")
.expect("found");
assert_eq!(by_pname.id, peer_id);
let by_pubk = store
.get_peer_by_public_key(iface_id, peer_pub.as_str())
.await
.expect("by pubk")
.expect("found");
assert_eq!(by_pubk.id, peer_id);
// Reject duplicate peer name on same interface
let dup_pname = Peer {
id: Uuid::new_v4(),
interface_id: iface_id,
name: "phone-alice".to_string(),
peer_type: PeerType::RoadWarrior,
state: PeerState::Active,
public_key: WireGuardPublicKey::new("different_key_123=".to_string()),
private_key: None,
preshared_key: None,
endpoint: None,
allowed_ips: "10.0.0.3/32".to_string(),
server_allowed_ips: None,
address_v4: None,
address_v6: None,
dns: None,
mtu: None,
persistent_keepalive: None,
profile: PeerProfile::SplitTunnel,
expires_at: None,
last_handshake_at: None,
created_at: now,
updated_at: now,
};
assert!(
store.create_peer(&dup_pname).await.is_err(),
"duplicate peer name on same interface must fail"
);
// Reject peer for non-existent interface (foreign key violation)
let non_existent_iface_peer = Peer {
id: Uuid::new_v4(),
interface_id: Uuid::new_v4(),
name: "orphan-peer".to_string(),
peer_type: PeerType::RoadWarrior,
state: PeerState::Active,
public_key: WireGuardPublicKey::new("orphan_key_123=".to_string()),
private_key: None,
preshared_key: None,
endpoint: None,
allowed_ips: "10.0.0.4/32".to_string(),
server_allowed_ips: None,
address_v4: None,
address_v6: None,
dns: None,
mtu: None,
persistent_keepalive: None,
profile: PeerProfile::SplitTunnel,
expires_at: None,
last_handshake_at: None,
created_at: now,
updated_at: now,
};
assert!(
store.create_peer(&non_existent_iface_peer).await.is_err(),
"peer for non-existent interface must fail foreign key constraint"
);
// Test peer state transition: active -> disabled -> revoked
store
.set_peer_state(peer_id, PeerState::Disabled)
.await
.expect("set disabled");
let disabled = store.get_peer(peer_id).await.expect("get").expect("peer");
assert_eq!(disabled.state, PeerState::Disabled);
store
.set_peer_state(peer_id, PeerState::Revoked)
.await
.expect("set revoked");
let revoked = store.get_peer(peer_id).await.expect("get").expect("peer");
assert_eq!(revoked.state, PeerState::Revoked);
// Test update_peer_handshake
let handshake_time = Utc::now().naive_utc();
store
.update_peer_handshake(peer_id, handshake_time)
.await
.expect("update handshake");
let after_hs = store.get_peer(peer_id).await.expect("get").expect("peer");
assert!(after_hs.last_handshake_at.is_some());
// Test list_peers_for_interface
let peer_list = store
.list_peers_for_interface(iface_id)
.await
.expect("list peers");
assert_eq!(peer_list.len(), 1);
// Test cascade delete: deleting interface must cascade and delete its peers
store
.delete_interface(iface_id)
.await
.expect("delete interface");
assert!(store.get_interface(iface_id).await.expect("get").is_none());
assert!(
store.get_peer(peer_id).await.expect("get").is_none(),
"peer must be cascade-deleted with interface"
);
}