llm_manager/backend/
ws_server.rs1use std::collections::HashMap;
2use std::sync::Arc;
3
4use anyhow::{Context, Result, anyhow};
5use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
6use axum::http::StatusCode;
7use axum::response::IntoResponse;
8use axum::{Router, response::Html, routing::get};
9use futures_util::{SinkExt, StreamExt};
10use tokio::sync::broadcast;
11use tokio::task::JoinHandle;
12use tower_http::trace::TraceLayer;
13use tracing::{error, info, warn};
14
15use crate::models::WsMetrics;
16
17#[derive(Clone)]
18pub struct WsAppState {
19 pub metrics_rx: Arc<broadcast::Receiver<WsMetrics>>,
20 pub auth_key: Option<String>,
21}
22
23pub async fn start_ws_server(
24 port: u16,
25 metrics_rx: Arc<broadcast::Receiver<WsMetrics>>,
26 auth_key: Option<String>,
27 tls_config: Option<axum_server::tls_rustls::RustlsConfig>,
28 host: String,
29 mut shutdown_rx: tokio::sync::watch::Receiver<bool>,
30) -> Result<JoinHandle<()>> {
31 let state = WsAppState {
32 metrics_rx,
33 auth_key,
34 };
35
36 let app = Router::new()
37 .route("/dashboard", get(serve_dashboard))
38 .route("/ws", get(ws_handler))
39 .route("/health", get(|| async { "OK" }))
40 .layer(TraceLayer::new_for_http())
41 .with_state(state);
42
43 let addr = format!("{host}:{port}");
44
45 match tls_config {
46 Some(tls_cfg) => {
47 let socket_addr: std::net::SocketAddr = addr
48 .parse()
49 .map_err(|e| anyhow!("Invalid bind address {addr} for TLS: {e}"))?;
50 let tls_listener = axum_server::bind_rustls(socket_addr, tls_cfg);
51 let shutdown_fut = async move {
52 let _ = shutdown_rx.wait_for(|v| *v).await;
53 };
54 let handle = tokio::spawn(async move {
55 tokio::select! {
56 result = tls_listener.serve(app.into_make_service()) => {
57 if let Err(e) = result {
58 error!("WebSocket server error: {e}");
59 }
60 }
61 _ = shutdown_fut => {},
62 }
63 });
64 info!("WebSocket server listening on https://{addr}");
65 Ok(handle)
66 }
67 None => {
68 let listener = tokio::net::TcpListener::bind(&addr)
69 .await
70 .with_context(|| format!("Failed to bind WebSocket server to {addr}"))?;
71 let handle = tokio::spawn(async move {
72 let _ = axum::serve(listener, app)
73 .with_graceful_shutdown(async move {
74 let _ = shutdown_rx.wait_for(|v| *v).await;
75 })
76 .await;
77 });
78 info!("WebSocket server listening on http://{addr}");
79 Ok(handle)
80 }
81 }
82}
83
84pub fn stop_ws_server(handle: JoinHandle<()>) {
85 handle.abort();
86}
87
88async fn serve_dashboard(
89 axum::extract::State(state): axum::extract::State<WsAppState>,
90) -> Html<String> {
91 let auth_json = serde_json::to_string(&state.auth_key).unwrap_or("null".to_string());
92 let auth_script = format!("<script>window.__WS_AUTH={};</script>", auth_json);
93 let html = include_str!("../dashboard.html");
94 Html(html.replacen("</body>", &format!("{}\n</body>", auth_script), 1))
95}
96
97async fn ws_handler(
98 ws: WebSocketUpgrade,
99 axum::extract::State(state): axum::extract::State<WsAppState>,
100 axum::extract::Query(query): axum::extract::Query<HashMap<String, String>>,
101) -> impl IntoResponse {
102 if let Some(ref expected) = state.auth_key {
103 if let Some(provided) = query.get("auth").and_then(|v| urlencoding::decode(v).ok()) {
104 if provided != *expected {
105 return StatusCode::UNAUTHORIZED.into_response();
106 }
107 } else {
108 return StatusCode::UNAUTHORIZED.into_response();
109 }
110 }
111 ws.on_upgrade(move |socket| handle_socket(socket, state))
112}
113
114async fn handle_socket(socket: WebSocket, state: WsAppState) {
115 let mut rx = state.metrics_rx.resubscribe();
116 info!("WebSocket client connected");
117
118 let (mut sender, mut receiver) = socket.split();
119
120 loop {
121 tokio::select! {
122 biased;
123 _ = receiver.next() => {
124 info!("WebSocket client disconnected");
125 break;
126 }
127 metrics = rx.recv() => match metrics {
128 Ok(m) => {
129 let json = match serde_json::to_string(&m) {
130 Ok(j) => j,
131 Err(e) => {
132 error!("Failed to serialize metrics: {e}");
133 continue;
134 }
135 };
136 if sender.send(Message::Text(json.into())).await.is_err() {
137 info!("WebSocket client disconnected");
138 break;
139 }
140 }
141 Err(broadcast::error::RecvError::Lagged(n)) => {
142 warn!("WebSocket client lagged behind, skipped {n} metrics");
143 }
144 Err(broadcast::error::RecvError::Closed) => {
145 break;
146 }
147 },
148 }
149 }
150}