use crate::config::AppConfig;
use crate::event::Event;
use anyhow::Result;
use axum::{
extract::{
ws::{Message, WebSocket, WebSocketUpgrade},
State,
},
response::{Html, IntoResponse, Response},
routing::get,
Router,
};
use futures_util::{
sink::SinkExt,
stream::{SplitSink, SplitStream, StreamExt},
};
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::sync::{
broadcast::{error::RecvError, Sender as BroadcastSender},
mpsc::Receiver as MpscReceiver,
watch::Receiver as WatchReceiver,
};
use tracing::{debug, error, info, warn};
#[derive(Clone)]
struct AppState {
event_tx: BroadcastSender<Event>,
}
async fn serve_home() -> impl IntoResponse {
Html(include_str!("../static/index.html"))
}
async fn websocket_handler(ws: WebSocketUpgrade, State(state): State<Arc<AppState>>) -> Response {
info!("New WebSocket connection request.");
ws.on_upgrade(move |socket| handle_socket(socket, state))
}
async fn handle_socket(socket: WebSocket, state: Arc<AppState>) {
info!("WebSocket client connected.");
let (mut sender, mut receiver): (SplitSink<WebSocket, Message>, SplitStream<WebSocket>) =
socket.split();
let mut rx = state.event_tx.subscribe();
let send_task = tokio::spawn(async move {
loop {
match rx.recv().await {
Ok(event) => {
match serde_json::to_string(&event) {
Ok(json_payload) => {
if sender.send(Message::Text(json_payload)).await.is_err() {
warn!("Failed to send message to WebSocket client, client disconnected?");
break;
}
debug!("Sent event to WebSocket client: {:?}", event.op);
}
Err(e) => {
error!("Failed to serialize event for WebSocket: {}", e);
}
}
}
Err(RecvError::Lagged(missed_count)) => {
warn!(
"WebSocket client lagged behind, missed {} messages.",
missed_count
);
}
Err(RecvError::Closed) => {
info!("Broadcast channel closed, WebSocket send task for client finishing.");
break;
}
}
}
info!("WebSocket send task for a client finished.");
});
let recv_task = tokio::spawn(async move {
while let Some(Ok(msg)) = receiver.next().await {
match msg {
Message::Text(t) => {
debug!("Received text from WebSocket client: {}", t);
}
Message::Binary(_) => {
debug!("Received binary from WebSocket client.");
}
Message::Ping(_) => {
debug!("Received Ping from WebSocket client, Axum handles Pong automatically.");
}
Message::Pong(_) => {
debug!("Received Pong from WebSocket client.");
}
Message::Close(_) => {
debug!("WebSocket client sent Close frame.");
break;
}
}
}
info!("WebSocket receive task for a client finished.");
});
tokio::select! {
_ = send_task => {},
_ = recv_task => {},
}
info!("WebSocket client connection handler finished.");
}
pub async fn start_server(
app_config: Arc<AppConfig>,
mut event_source_rx: MpscReceiver<Event>,
event_broadcast_tx: BroadcastSender<Event>,
shutdown_signal: WatchReceiver<bool>,
) -> Result<()> {
let web_addr_str = app_config.web_addr.clone();
let socket_addr: SocketAddr = web_addr_str.parse()?;
let app_state = Arc::new(AppState {
event_tx: event_broadcast_tx.clone(),
});
let broadcast_tx_clone = event_broadcast_tx.clone();
tokio::spawn(async move {
while let Some(event) = event_source_rx.recv().await {
if broadcast_tx_clone.receiver_count() > 0 {
if let Err(e) = broadcast_tx_clone.send(event.clone()) {
debug!(
"Failed to broadcast event to WebSocket hub (no subscribers?): {}",
e
);
}
} else {
debug!(
"No WebSocket subscribers, event not broadcasted: {:?}",
event.op
);
}
}
info!("Event source for Web UI finished.");
});
let app = Router::new()
.route("/", get(serve_home))
.route("/ws", get(websocket_handler))
.with_state(app_state);
info!("Web server starting on http://{}", socket_addr);
let mut shutdown = shutdown_signal.clone(); axum::serve(
tokio::net::TcpListener::bind(socket_addr).await?,
app.into_make_service_with_connect_info::<SocketAddr>(),
)
.with_graceful_shutdown(async move {
shutdown.changed().await.ok();
info!("Web server shutting down gracefully.");
})
.await?;
info!("Web server stopped.");
Ok(())
}