use crate::QsshError;
use crate::security::{SecurityMonitor, SecurityEvent, EventType, Severity};
use warp::{Filter, Rejection, Reply};
use warp::ws::{Message, WebSocket};
use futures::{StreamExt, SinkExt};
use std::sync::Arc;
use tokio::sync::{RwLock, mpsc};
use std::collections::HashMap;
use serde_json;
use std::net::SocketAddr;
pub struct WebServerConfig {
pub listen_addr: SocketAddr,
pub static_dir: String,
pub enable_cors: bool,
}
impl Default for WebServerConfig {
fn default() -> Self {
Self {
listen_addr: ([127, 0, 0, 1], 8080).into(),
static_dir: "app/static".to_string(),
enable_cors: true,
}
}
}
pub struct SecurityWebServer {
config: WebServerConfig,
monitor: Arc<SecurityMonitor>,
clients: Arc<RwLock<HashMap<String, mpsc::UnboundedSender<SecurityEvent>>>>,
}
impl SecurityWebServer {
pub fn new(config: WebServerConfig, monitor: Arc<SecurityMonitor>) -> Self {
Self {
config,
monitor,
clients: Arc::new(RwLock::new(HashMap::new())),
}
}
pub async fn start(self: Arc<Self>) -> crate::Result<()> {
let static_route = warp::fs::dir(self.config.static_dir.clone())
.with(warp::cors().allow_any_origin());
let ws_route = warp::path!("ws" / "security")
.and(warp::ws())
.and(with_server(self.clone()))
.map(|ws: warp::ws::Ws, server: Arc<SecurityWebServer>| {
ws.on_upgrade(move |socket| handle_ws_connection(socket, server))
});
let stats_route = warp::path!("api" / "stats")
.and(warp::get())
.and(with_server(self.clone()))
.and_then(get_stats);
let events_route = warp::path!("api" / "events")
.and(warp::get())
.and(warp::query::<EventQuery>())
.and(with_server(self.clone()))
.and_then(query_events);
let test_route = warp::path!("api" / "test" / "event")
.and(warp::post())
.and(with_server(self.clone()))
.and_then(trigger_test_event);
let routes = static_route
.or(ws_route)
.or(stats_route)
.or(events_route)
.or(test_route)
.with(warp::cors().allow_any_origin());
let server_clone = self.clone();
tokio::spawn(async move {
server_clone.broadcast_events().await;
});
log::info!("Security web server listening on {}", self.config.listen_addr);
warp::serve(routes)
.run(self.config.listen_addr)
.await;
Ok(())
}
async fn broadcast_events(&self) {
let mut interval = tokio::time::interval(std::time::Duration::from_secs(1));
loop {
interval.tick().await;
let events = self.monitor.query_events(
crate::security::EventFilter::new()
).await;
let clients = self.clients.read().await;
for (client_id, sender) in clients.iter() {
if let Some(event) = events.first() {
if let Err(_) = sender.send(event.clone()) {
log::debug!("Client {} disconnected", client_id);
}
}
}
}
}
async fn register_client(&self, client_id: String, sender: mpsc::UnboundedSender<SecurityEvent>) {
let mut clients = self.clients.write().await;
clients.insert(client_id, sender);
}
async fn unregister_client(&self, client_id: &str) {
let mut clients = self.clients.write().await;
clients.remove(client_id);
}
}
async fn handle_ws_connection(ws: WebSocket, server: Arc<SecurityWebServer>) {
let client_id = uuid::Uuid::new_v4().to_string();
log::info!("New WebSocket client connected: {}", client_id);
let (mut ws_sender, mut ws_receiver) = ws.split();
let (tx, mut rx) = mpsc::unbounded_channel::<SecurityEvent>();
server.register_client(client_id.clone(), tx).await;
let client_id_clone = client_id.clone();
tokio::spawn(async move {
while let Some(event) = rx.recv().await {
if let Ok(json) = serde_json::to_string(&event) {
if ws_sender.send(Message::text(json)).await.is_err() {
break;
}
}
}
log::info!("Client {} sender task ended", client_id_clone);
});
while let Some(result) = ws_receiver.next().await {
match result {
Ok(msg) => {
if msg.is_close() {
break;
}
if msg.is_ping() {
}
}
Err(_) => break,
}
}
server.unregister_client(&client_id).await;
log::info!("Client {} disconnected", client_id);
}
async fn get_stats(server: Arc<SecurityWebServer>) -> std::result::Result<impl Reply, Rejection> {
let stats = server.monitor.get_stats().await;
Ok(warp::reply::json(&stats))
}
async fn query_events(
query: EventQuery,
server: Arc<SecurityWebServer>,
) -> std::result::Result<impl Reply, Rejection> {
let mut filter = crate::security::EventFilter::new();
if let Some(severity) = query.severity {
filter.severity_min = Some(parse_severity(&severity));
}
if let Some(session_id) = query.session_id {
filter.session_id = Some(session_id);
}
if let Some(ip) = query.source_ip {
if let Ok(parsed_ip) = ip.parse() {
filter.source_ip = Some(parsed_ip);
}
}
let events = server.monitor.query_events(filter).await;
Ok(warp::reply::json(&events))
}
async fn trigger_test_event(server: Arc<SecurityWebServer>) -> std::result::Result<impl Reply, Rejection> {
let event = SecurityEvent::new(
Severity::Info,
EventType::AuthSuccess {
username: "test_user".to_string(),
method: "test".to_string(),
},
);
let _ = server.monitor.log_event(event).await;
Ok(warp::reply::json(&serde_json::json!({
"status": "ok",
"message": "Test event triggered"
})))
}
fn with_server(
server: Arc<SecurityWebServer>,
) -> impl Filter<Extract = (Arc<SecurityWebServer>,), Error = std::convert::Infallible> + Clone {
warp::any().map(move || server.clone())
}
fn parse_severity(s: &str) -> Severity {
match s.to_lowercase().as_str() {
"debug" => Severity::Debug,
"info" => Severity::Info,
"warning" => Severity::Warning,
"error" => Severity::Error,
"critical" => Severity::Critical,
_ => Severity::Info,
}
}
#[derive(serde::Deserialize)]
struct EventQuery {
severity: Option<String>,
session_id: Option<String>,
source_ip: Option<String>,
limit: Option<usize>,
}
pub async fn start_monitoring_interface(
monitor: Arc<SecurityMonitor>,
addr: Option<SocketAddr>,
) -> crate::Result<()> {
let mut config = WebServerConfig::default();
if let Some(addr) = addr {
config.listen_addr = addr;
}
let server = Arc::new(SecurityWebServer::new(config, monitor));
server.start().await
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_severity() {
assert_eq!(parse_severity("debug"), Severity::Debug);
assert_eq!(parse_severity("INFO"), Severity::Info);
assert_eq!(parse_severity("Warning"), Severity::Warning);
assert_eq!(parse_severity("error"), Severity::Error);
assert_eq!(parse_severity("CRITICAL"), Severity::Critical);
assert_eq!(parse_severity("unknown"), Severity::Info);
}
}