use crate::monitor::MonitorSample;
use axum::http::HeaderValue;
use axum::{
Json, Router,
extract::{
State,
ws::{Message, WebSocket, WebSocketUpgrade},
},
response::IntoResponse,
routing::get,
};
use std::sync::{Arc, Mutex};
use tokio::sync::broadcast;
use tower_http::cors::{AllowOrigin, CorsLayer};
#[derive(Clone, Copy, Debug, serde::Serialize)]
pub struct ApiData {
pub timestamp: u64,
pub cpu_power: f64,
pub gpu_power: f64,
pub total_power: f64,
pub cpu_usage: f64,
pub pid_or_app_power: f64,
}
impl From<&MonitorSample> for ApiData {
fn from(sample: &MonitorSample) -> Self {
Self {
timestamp: sample.timestamp,
cpu_power: sample.cpu_power,
gpu_power: sample.gpu_power,
total_power: sample.total_power,
cpu_usage: sample.cpu_usage,
pid_or_app_power: sample.pid_app_power(),
}
}
}
struct AppState {
latest_data: Mutex<ApiData>,
tx: broadcast::Sender<ApiData>,
}
pub async fn start_api_server(
port: u16,
extra_allowed_origins: Vec<String>,
tx: broadcast::Sender<ApiData>,
shutdown: tokio::sync::oneshot::Receiver<()>,
listener: tokio::net::TcpListener,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let state = Arc::new(AppState {
latest_data: Mutex::new(ApiData {
timestamp: 0,
cpu_power: 0.0,
gpu_power: 0.0,
total_power: 0.0,
cpu_usage: 0.0,
pid_or_app_power: 0.0,
}),
tx: tx.clone(),
});
let state_for_loop = state.clone();
let mut rx = tx.subscribe();
tokio::spawn(async move {
loop {
match rx.recv().await {
Ok(data) => {
let mut latest = state_for_loop
.latest_data
.lock()
.unwrap_or_else(|e| e.into_inner());
*latest = data;
}
Err(broadcast::error::RecvError::Lagged(_)) => continue,
Err(broadcast::error::RecvError::Closed) => break,
}
}
});
let mut has_wildcard = false;
let mut origins: Vec<HeaderValue> = Vec::new();
for scheme_host in [
format!("http://127.0.0.1:{}", port),
format!("http://localhost:{}", port),
] {
if let Ok(v) = HeaderValue::from_str(&scheme_host) {
origins.push(v);
}
}
for origin in &extra_allowed_origins {
if origin == "*" {
has_wildcard = true;
break;
}
match HeaderValue::from_str(origin) {
Ok(v) => origins.push(v),
Err(_) => {
tracing::warn!(origin, "Ignoring invalid --api-allowed-origin value");
}
}
}
let cors = if has_wildcard {
CorsLayer::new().allow_origin(tower_http::cors::Any)
} else {
CorsLayer::new().allow_origin(AllowOrigin::list(origins))
};
let app = Router::new()
.route("/data", get(get_data))
.route("/ws", get(ws_handler))
.layer(cors)
.with_state(state);
tracing::info!(port, "API server listening");
println!(
"\x1b[1;32m✓ API Server running on http://127.0.0.1:{}/data\x1b[0m",
port
);
axum::serve(listener, app)
.with_graceful_shutdown(async {
let _ = shutdown.await;
})
.await?;
Ok(())
}
async fn get_data(State(state): State<Arc<AppState>>) -> impl IntoResponse {
let data = state.latest_data.lock().unwrap_or_else(|e| e.into_inner());
Json(*data)
}
async fn ws_handler(ws: WebSocketUpgrade, State(state): State<Arc<AppState>>) -> impl IntoResponse {
ws.on_upgrade(|socket| handle_socket(socket, state))
}
async fn handle_socket(mut socket: WebSocket, state: Arc<AppState>) {
let mut rx = state.tx.subscribe();
loop {
let data = match rx.recv().await {
Ok(data) => data,
Err(broadcast::error::RecvError::Lagged(_)) => continue,
Err(broadcast::error::RecvError::Closed) => break,
};
let msg = match serde_json::to_string(&data) {
Ok(json) => Message::Text(json.into()),
Err(_) => continue,
};
if socket.send(msg).await.is_err() {
break;
}
}
}