Skip to main content

llm_manager/backend/
ws_server.rs

1use 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}