use std::convert::Infallible;
use std::sync::Arc;
use std::time::{Duration, Instant};
use axum::extract::{Query, State};
use axum::http::{HeaderMap, HeaderValue, StatusCode, header};
use axum::response::IntoResponse;
use axum::response::Response;
use axum::response::sse::{Event, KeepAlive, Sse};
use futures_util::Stream;
use futures_util::stream::unfold;
use serde::Deserialize;
use tokio::sync::broadcast::Receiver;
use tokio::sync::broadcast::error::RecvError;
use crate::api::frame_bus::FrameBus;
use crate::snapshot::Snapshot;
use super::snapshot::{SectionFilter, filter_snapshot_value, parse_include};
pub const DEFAULT_HEARTBEAT_SECS: u64 = 30;
pub const MAX_INTERVAL_SECS: u64 = 86_400;
pub const DEFAULT_MAX_SSE_SUBSCRIBERS: usize = 256;
fn configured_max_subscribers() -> usize {
match std::env::var("ALL_SMI_API_MAX_SSE_SUBSCRIBERS") {
Ok(v) => match v.trim().parse::<usize>() {
Ok(n) => n,
Err(e) => {
tracing::warn!(
value = %v,
error = %e,
"ALL_SMI_API_MAX_SSE_SUBSCRIBERS is not a valid usize; falling back to default"
);
DEFAULT_MAX_SSE_SUBSCRIBERS
}
},
Err(_) => DEFAULT_MAX_SSE_SUBSCRIBERS,
}
}
#[derive(Debug, Default, Deserialize)]
pub struct EventsQuery {
pub include: Option<String>,
pub throttle: Option<u64>,
pub heartbeat: Option<u64>,
}
pub async fn events_handler(
State(bus): State<FrameBus>,
Query(params): Query<EventsQuery>,
headers: HeaderMap,
) -> Response {
let filter = parse_include(params.include.as_deref());
let throttle = resolve_throttle(params.throttle, bus.collection_interval());
let heartbeat = resolve_heartbeat(params.heartbeat);
let cap = configured_max_subscribers();
if cap > 0 && bus.subscriber_count() >= cap {
tracing::warn!(
current_subscribers = bus.subscriber_count(),
cap,
"rejecting SSE subscription: subscriber cap reached. Tune ALL_SMI_API_MAX_SSE_SUBSCRIBERS or 0 to disable the cap."
);
let body = serde_json::json!({
"error": "subscriber_cap_exceeded",
"message": format!(
"SSE subscriber cap {cap} reached; retry later or tune ALL_SMI_API_MAX_SSE_SUBSCRIBERS"
),
})
.to_string();
let mut resp = (StatusCode::SERVICE_UNAVAILABLE, body).into_response();
resp.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/json"),
);
resp.headers_mut()
.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
resp.headers_mut()
.insert(header::RETRY_AFTER, HeaderValue::from_static("5"));
return resp;
}
if let Some(id) = headers.get("last-event-id").and_then(|v| v.to_str().ok()) {
const MAX_LOGGED_ID: usize = 256;
let mut boundary = id.len().min(MAX_LOGGED_ID);
while boundary > 0 && !id.is_char_boundary(boundary) {
boundary -= 1;
}
tracing::debug!(
client_last_event_id = %&id[..boundary],
truncated = id.len() > boundary,
"SSE client reconnected; history replay not supported, resuming with next live frame"
);
}
let stream = build_sse_stream(bus.subscribe(), filter, throttle);
let sse = Sse::new(stream).keep_alive(KeepAlive::new().interval(heartbeat).text("keep-alive"));
let mut extra = HeaderMap::new();
extra.insert("X-Accel-Buffering", HeaderValue::from_static("no"));
extra.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
(extra, sse).into_response()
}
fn resolve_throttle(user: Option<u64>, collection_interval: Duration) -> Duration {
let floor = collection_interval.as_secs();
let effective_floor = floor.min(MAX_INTERVAL_SECS);
let secs = user.unwrap_or(0).clamp(effective_floor, MAX_INTERVAL_SECS);
if secs == 0 {
collection_interval
} else {
Duration::from_secs(secs)
}
}
fn resolve_heartbeat(user: Option<u64>) -> Duration {
let secs = user.unwrap_or(0);
if secs == 0 {
Duration::from_secs(DEFAULT_HEARTBEAT_SECS)
} else {
Duration::from_secs(secs.clamp(1, MAX_INTERVAL_SECS))
}
}
struct StreamState {
rx: Receiver<Arc<Snapshot>>,
filter: SectionFilter,
throttle: Duration,
last_emit: Option<Instant>,
}
pub fn build_sse_stream(
rx: Receiver<Arc<Snapshot>>,
filter: SectionFilter,
throttle: Duration,
) -> impl Stream<Item = Result<Event, Infallible>> {
unfold(
StreamState {
rx,
filter,
throttle,
last_emit: None,
},
|mut state| async move {
loop {
match state.rx.recv().await {
Ok(frame) => {
if let Some(prev) = state.last_emit
&& prev.elapsed() < state.throttle
{
continue;
}
let event = build_snapshot_event(&frame, &state.filter);
state.last_emit = Some(Instant::now());
return Some((Ok(event), state));
}
Err(RecvError::Lagged(n)) => {
let event = build_lag_event(n);
return Some((Ok(event), state));
}
Err(RecvError::Closed) => {
return None;
}
}
}
},
)
}
fn build_snapshot_event(snapshot: &Arc<Snapshot>, filter: &SectionFilter) -> Event {
let value = filter_snapshot_value(snapshot, filter);
let event = Event::default()
.event("snapshot")
.id(event_id_for(snapshot));
match event.clone().json_data(&value) {
Ok(e) => e,
Err(err) => error_event(&err.to_string()),
}
}
fn build_lag_event(dropped: u64) -> Event {
let payload = serde_json::json!({ "dropped": dropped });
Event::default()
.event("lag")
.json_data(&payload)
.unwrap_or_else(|e| error_event(&e.to_string()))
}
fn error_event(message: &str) -> Event {
let payload = serde_json::json!({ "error": message });
Event::default()
.event("error")
.json_data(&payload)
.unwrap_or_else(|_| Event::default().comment("serialization failure"))
}
fn event_id_for(snapshot: &Arc<Snapshot>) -> String {
snapshot.timestamp.clone()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::api::frame_bus::FrameBus;
use crate::snapshot::Snapshot;
use futures_util::StreamExt;
fn minimal_snapshot() -> Snapshot {
Snapshot {
schema: 1,
timestamp: "2026-04-20T00:00:01Z".to_string(),
hostname: "h".to_string(),
gpus: Some(Vec::new()),
cpus: Some(Vec::new()),
memory: Some(Vec::new()),
chassis: Some(Vec::new()),
processes: None,
storage: None,
errors: Vec::new(),
}
}
#[test]
fn resolve_throttle_clamps_below_collection_interval() {
let d = resolve_throttle(Some(1), Duration::from_secs(5));
assert_eq!(d, Duration::from_secs(5));
}
#[test]
fn resolve_throttle_defaults_to_collection_interval() {
let d = resolve_throttle(None, Duration::from_secs(3));
assert_eq!(d, Duration::from_secs(3));
}
#[test]
fn resolve_heartbeat_defaults_to_thirty() {
let d = resolve_heartbeat(None);
assert_eq!(d, Duration::from_secs(DEFAULT_HEARTBEAT_SECS));
}
#[test]
fn resolve_heartbeat_accepts_custom_value() {
let d = resolve_heartbeat(Some(10));
assert_eq!(d, Duration::from_secs(10));
}
#[test]
fn resolve_throttle_does_not_panic_when_collection_interval_exceeds_max() {
let huge = Duration::from_secs(MAX_INTERVAL_SECS + 3600);
let d = resolve_throttle(Some(60), huge);
assert!(d <= Duration::from_secs(MAX_INTERVAL_SECS));
}
#[test]
fn resolve_throttle_handles_none_with_oversize_interval() {
let huge = Duration::from_secs(MAX_INTERVAL_SECS * 2);
let d = resolve_throttle(None, huge);
assert!(d <= Duration::from_secs(MAX_INTERVAL_SECS));
}
#[tokio::test]
async fn stream_emits_published_frame() {
let bus = FrameBus::new(Duration::from_millis(10));
let filter = SectionFilter::default_http();
let rx = bus.subscribe();
bus.publish(minimal_snapshot()).await;
let stream = build_sse_stream(rx, filter, Duration::from_millis(10));
futures_util::pin_mut!(stream);
let next = stream.next().await.expect("stream yields at least once");
assert!(next.is_ok());
}
#[tokio::test]
async fn lag_event_emitted_when_receiver_falls_behind() {
let bus = FrameBus::new(Duration::from_millis(10));
let filter = SectionFilter::default_http();
let rx = bus.subscribe();
for _ in 0..(crate::api::frame_bus::FRAME_BUFFER + 4) {
bus.publish(minimal_snapshot()).await;
}
let stream = build_sse_stream(rx, filter, Duration::from_millis(10));
futures_util::pin_mut!(stream);
let first = stream.next().await.expect("stream yields at least once");
assert!(first.is_ok());
}
}