use axum::http::{HeaderName, HeaderValue, Method, header};
use axum::{Router, routing::get};
use std::time::Duration;
use tokio::net::TcpListener;
use tokio::sync::RwLock;
use tower_http::cors::{AllowOrigin, CorsLayer};
use tower_http::trace::TraceLayer;
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt};
use crate::api::shutdown::{mark_serving, shutdown_signal};
#[cfg(unix)]
use std::path::PathBuf;
#[cfg(unix)]
use tokio::net::UnixListener;
use crate::api::FrameBus;
use crate::api::collection_loop::run_collection_loop;
use crate::api::handlers::events::events_handler;
use crate::api::handlers::snapshot::snapshot_handler;
use crate::api::handlers::{SharedState, metrics_handler, ready_handler};
use crate::api::server_state::ApiState;
use crate::app_state::AppState;
use crate::cli::ApiArgs;
use crate::common::config_file::Settings;
#[cfg(unix)]
fn get_default_socket_path() -> PathBuf {
#[cfg(target_os = "linux")]
{
let var_run_path = PathBuf::from("/var/run/all-smi.sock");
if let Ok(metadata) = std::fs::metadata("/var/run")
&& metadata.is_dir()
{
let test_path = PathBuf::from("/var/run/.all-smi-test");
if std::fs::write(&test_path, b"").is_ok() {
let _ = std::fs::remove_file(&test_path);
return var_run_path;
}
}
PathBuf::from("/tmp/all-smi.sock")
}
#[cfg(target_os = "macos")]
{
PathBuf::from("/tmp/all-smi.sock")
}
#[cfg(not(any(target_os = "linux", target_os = "macos")))]
{
PathBuf::from("/tmp/all-smi.sock")
}
}
#[cfg(unix)]
fn remove_stale_socket(path: &std::path::Path) -> std::io::Result<()> {
match std::fs::remove_file(path) {
Ok(()) => {
tracing::info!("Removed stale socket file: {}", path.display());
Ok(())
}
Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
Ok(())
}
Err(e) => Err(e),
}
}
#[cfg(unix)]
fn set_socket_permissions(path: &std::path::Path) -> std::io::Result<()> {
use std::os::unix::fs::PermissionsExt;
let permissions = std::fs::Permissions::from_mode(0o600);
std::fs::set_permissions(path, permissions)
}
pub async fn run_api_mode(args: &ApiArgs, settings: &Settings) {
if tracing_subscriber::registry()
.with(
tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| "all_smi=debug,tower_http=debug".into()),
)
.with(tracing_subscriber::fmt::layer())
.try_init()
.is_err()
{
tracing::debug!("a tracing subscriber is already installed; keeping the host's");
}
println!("Starting API mode...");
let mut initial_state = AppState::with_energy_config(&settings.energy);
initial_state.display_config = settings.display.clone();
if initial_state.energy_config.wal_enabled {
match crate::metrics::energy_wal::resolve_wal_path(
initial_state.energy_config.wal_path.as_deref(),
) {
Some(wal_path) => {
let path_display = wal_path.display().to_string();
match crate::metrics::energy_wal::replay_from_path(
&wal_path,
initial_state.energy.integrator_mut(),
) {
Ok(index) => {
if !index.is_empty() {
tracing::info!(
"energy WAL: replayed {} records from {path_display}",
index.len()
);
}
initial_state.energy_wal_replay = index;
}
Err(e) => {
tracing::warn!("energy WAL: replay from {path_display} failed: {e}");
}
}
}
None => {
tracing::warn!(
"energy WAL: no cache directory available in environment; \
counters are in-memory only"
);
}
}
}
let state = SharedState::new(RwLock::new(initial_state));
let state_clone = state.clone();
let processes = args.processes.unwrap_or(false);
let interval = args.interval.unwrap_or(3);
let wal_flush_handle = {
let state = state.clone();
let state_read = state.read().await;
let cfg = state_read.energy_config.clone();
drop(state_read);
if cfg.wal_enabled {
match crate::metrics::energy_wal::resolve_wal_path(cfg.wal_path.as_deref()) {
Some(path) => Some(crate::metrics::energy_wal::spawn_wal_flush_task(
state,
path,
crate::metrics::energy_wal::DEFAULT_FLUSH_INTERVAL,
)),
None => {
tracing::warn!(
"energy WAL: no cache directory available in environment; \
skipping flush task (counters remain in-memory only)"
);
None
}
}
} else {
None
}
};
let bus = FrameBus::new(Duration::from_secs(interval));
tokio::spawn(run_collection_loop(
state_clone.clone(),
bus.clone(),
interval,
processes,
));
let api_state = ApiState::new(state, bus);
let app = Router::new()
.route("/metrics", get(metrics_handler))
.route("/-/ready", get(ready_handler))
.route("/events", get(events_handler))
.route("/snapshot", get(snapshot_handler))
.with_state(api_state)
.layer(build_cors_layer())
.layer(TraceLayer::new_for_http());
#[cfg(unix)]
{
let socket_path = args.socket.as_ref().map(|s| {
if s.is_empty() {
get_default_socket_path()
} else {
PathBuf::from(s)
}
});
let port = args.port.unwrap_or(9090);
match (port, socket_path) {
(1..=u16::MAX, Some(path)) => {
run_dual_listeners(app, port, path).await;
}
(0, Some(path)) => {
run_unix_listener(app, path).await;
}
(1..=u16::MAX, None) => {
run_tcp_listener(app, port).await;
}
(0, None) => {
tracing::error!(
"No listeners configured. Use --port or --socket to specify a listener."
);
eprintln!(
"Error: No listeners configured. Use --port or --socket to specify a listener."
);
}
}
}
#[cfg(not(unix))]
{
run_tcp_listener(app, args.port.unwrap_or(9090)).await;
}
if let Some(handle) = wal_flush_handle {
handle.shutdown().await;
}
}
async fn run_tcp_listener(app: Router, port: u16) {
let listener = match TcpListener::bind(&format!("0.0.0.0:{port}")).await {
Ok(l) => l,
Err(e) => {
tracing::error!("Failed to bind TCP listener on port {port}: {e}");
eprintln!("Error: Failed to bind TCP listener on port {port}: {e}");
return;
}
};
tracing::info!(
"API server listening on {}",
listener
.local_addr()
.unwrap_or_else(|_| "unknown".parse().unwrap())
);
mark_serving();
if let Err(e) = axum::serve(listener, app)
.with_graceful_shutdown(shutdown_signal())
.await
{
tracing::error!("TCP server error: {e}");
}
}
#[cfg(unix)]
async fn run_unix_listener(app: Router, path: PathBuf) {
if let Err(e) = remove_stale_socket(&path) {
tracing::warn!("Failed to remove stale socket file: {e}");
}
if let Some(parent) = path.parent()
&& !parent.exists()
&& let Err(e) = std::fs::create_dir_all(parent)
{
tracing::error!(
"Failed to create socket directory {}: {e}",
parent.display()
);
eprintln!(
"Error: Failed to create socket directory {}: {e}",
parent.display()
);
return;
}
let listener = match UnixListener::bind(&path) {
Ok(l) => l,
Err(e) => {
tracing::error!("Failed to bind Unix socket at {}: {e}", path.display());
eprintln!(
"Error: Failed to bind Unix socket at {}: {e}",
path.display()
);
return;
}
};
if let Err(e) = set_socket_permissions(&path) {
tracing::warn!("Failed to set socket permissions: {e}");
}
tracing::info!("API server listening on Unix socket: {}", path.display());
mark_serving();
if let Err(e) = axum::serve(listener, app)
.with_graceful_shutdown(shutdown_signal())
.await
{
tracing::error!("Unix socket server error: {e}");
}
cleanup_socket(&path);
}
#[cfg(unix)]
async fn run_dual_listeners(app: Router, port: u16, socket_path: PathBuf) {
if let Err(e) = remove_stale_socket(&socket_path) {
tracing::warn!("Failed to remove stale socket file: {e}");
}
if let Some(parent) = socket_path.parent()
&& !parent.exists()
&& let Err(e) = std::fs::create_dir_all(parent)
{
tracing::error!(
"Failed to create socket directory {}: {e}",
parent.display()
);
eprintln!(
"Error: Failed to create socket directory {}: {e}",
parent.display()
);
return;
}
let tcp_listener = match TcpListener::bind(&format!("0.0.0.0:{port}")).await {
Ok(l) => l,
Err(e) => {
tracing::error!("Failed to bind TCP listener on port {port}: {e}");
eprintln!("Error: Failed to bind TCP listener on port {port}: {e}");
return;
}
};
let unix_listener = match UnixListener::bind(&socket_path) {
Ok(l) => l,
Err(e) => {
tracing::error!(
"Failed to bind Unix socket at {}: {e}",
socket_path.display()
);
eprintln!(
"Error: Failed to bind Unix socket at {}: {e}",
socket_path.display()
);
return;
}
};
if let Err(e) = set_socket_permissions(&socket_path) {
tracing::warn!("Failed to set socket permissions: {e}");
}
tracing::info!(
"API server listening on TCP {} and Unix socket {}",
tcp_listener
.local_addr()
.unwrap_or_else(|_| "unknown".parse().unwrap()),
socket_path.display()
);
mark_serving();
let app_clone = app.clone();
tokio::select! {
result = axum::serve(tcp_listener, app)
.with_graceful_shutdown(shutdown_signal()) => {
if let Err(e) = result {
tracing::error!("TCP server error: {e}");
}
}
result = axum::serve(unix_listener, app_clone)
.with_graceful_shutdown(shutdown_signal()) => {
if let Err(e) = result {
tracing::error!("Unix socket server error: {e}");
}
}
}
cleanup_socket(&socket_path);
}
#[cfg(unix)]
fn cleanup_socket(path: &std::path::Path) {
match std::fs::remove_file(path) {
Ok(()) => {
tracing::info!("Cleaned up socket file: {}", path.display());
}
Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
}
Err(e) => {
tracing::warn!("Failed to remove socket file on shutdown: {e}");
}
}
}
fn build_cors_layer() -> CorsLayer {
let raw = std::env::var("ALL_SMI_API_CORS_ALLOWED_ORIGINS").ok();
let trimmed = raw.as_deref().map(str::trim).unwrap_or("");
if trimmed.is_empty() {
return CorsLayer::new()
.allow_methods([Method::GET, Method::OPTIONS])
.allow_headers([
header::ACCEPT,
header::CONTENT_TYPE,
HeaderName::from_static("last-event-id"),
]);
}
if trimmed == "*" {
tracing::warn!(
"ALL_SMI_API_CORS_ALLOWED_ORIGINS=* selected; every origin may read /metrics, /snapshot, and /events. This exposes GPU telemetry, process command lines, and usernames cross-origin. Prefer an explicit origin list."
);
return CorsLayer::new()
.allow_origin(AllowOrigin::any())
.allow_methods([Method::GET, Method::OPTIONS])
.allow_headers([
header::ACCEPT,
header::CONTENT_TYPE,
HeaderName::from_static("last-event-id"),
]);
}
let origins: Vec<HeaderValue> = trimmed
.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.filter_map(|o| match HeaderValue::from_str(o) {
Ok(v) => Some(v),
Err(e) => {
tracing::warn!(origin = o, error = %e, "ignoring invalid CORS origin");
None
}
})
.collect();
if origins.is_empty() {
tracing::warn!(
"ALL_SMI_API_CORS_ALLOWED_ORIGINS was set but contained no valid origins; falling back to no-CORS default"
);
return CorsLayer::new()
.allow_methods([Method::GET, Method::OPTIONS])
.allow_headers([
header::ACCEPT,
header::CONTENT_TYPE,
HeaderName::from_static("last-event-id"),
]);
}
tracing::info!(
allowed_origins = origins.len(),
"CORS: allowlist configured from ALL_SMI_API_CORS_ALLOWED_ORIGINS"
);
CorsLayer::new()
.allow_origin(AllowOrigin::list(origins))
.allow_methods([Method::GET, Method::OPTIONS])
.allow_headers([
header::ACCEPT,
header::CONTENT_TYPE,
HeaderName::from_static("last-event-id"),
])
}
#[cfg(test)]
mod cors_tests {
use super::*;
struct EnvGuard {
key: &'static str,
original: Option<String>,
}
impl EnvGuard {
fn set(key: &'static str, value: &str) -> Self {
let original = std::env::var(key).ok();
unsafe { std::env::set_var(key, value) };
Self { key, original }
}
fn unset(key: &'static str) -> Self {
let original = std::env::var(key).ok();
unsafe { std::env::remove_var(key) };
Self { key, original }
}
}
impl Drop for EnvGuard {
fn drop(&mut self) {
match self.original.take() {
Some(v) => unsafe { std::env::set_var(self.key, v) },
None => unsafe { std::env::remove_var(self.key) },
}
}
}
#[test]
fn default_builds_without_panic() {
let _g = EnvGuard::unset("ALL_SMI_API_CORS_ALLOWED_ORIGINS");
let _layer = build_cors_layer();
}
#[test]
fn wildcard_allowed_when_explicitly_requested() {
let _g = EnvGuard::set("ALL_SMI_API_CORS_ALLOWED_ORIGINS", "*");
let _layer = build_cors_layer();
}
#[test]
fn invalid_origins_drop_to_default() {
let _g = EnvGuard::set("ALL_SMI_API_CORS_ALLOWED_ORIGINS", "\n\tnot a url");
let _layer = build_cors_layer();
}
}