use crate::error::ApiError;
use crate::rest::mission_paths;
use crate::ServerState;
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::extract::{Path as UrlPath, Query, State};
use axum::http::{header, HeaderMap, StatusCode};
use axum::response::{IntoResponse, Response};
use axum::Extension;
use kranz_engine::event_log::EventLog;
use kranz_engine::events::{Event, EventKind};
use kranz_engine::paths::MissionPaths;
use kranz_engine::reducer;
use kranz_engine::types::MissionState;
use serde_json::{json, Value};
use std::collections::HashMap;
use std::path::Path;
use std::sync::Arc;
use std::time::Duration;
const MAX_REPLAY_GAP: u64 = 5000;
const POLL_INTERVAL: Duration = Duration::from_millis(250);
pub(crate) async fn ws_handler(
Extension(read_work): Extension<crate::read_work::ReadWork>,
State(server): State<Arc<ServerState>>,
UrlPath(id): UrlPath<String>,
Query(params): Query<HashMap<String, String>>,
headers: HeaderMap,
ws: WebSocketUpgrade,
) -> Response {
let origin = headers
.get(header::ORIGIN)
.and_then(|value| value.to_str().ok());
if !crate::ws_origin_allowed(origin, server.bind_addr, server.bind_is_loopback) {
return StatusCode::FORBIDDEN.into_response();
}
let since = match params.get("since") {
None => None,
Some(raw) => match raw.parse::<u64>() {
Ok(n) => Some(n),
Err(_) => {
return ApiError::bad_request(format!("invalid 'since' value: '{raw}'"))
.into_response();
}
},
};
let paths = match read_work
.run(move || {
let paths = mission_paths(&server, &id)?;
if !paths.events_file().is_file() {
return Err(ApiError::not_found(format!("unknown mission '{id}'")));
}
Ok(paths)
})
.await
{
Ok(paths) => paths,
Err(error) => return error.into_response(),
};
ws.on_upgrade(move |socket| session(socket, paths, since, read_work))
}
async fn session(
mut socket: WebSocket,
paths: MissionPaths,
since: Option<u64>,
read_work: crate::read_work::ReadWork,
) {
let events_path = paths.events_file();
let initial_path = events_path.clone();
let Ok((events, mut state)) = read_work
.run(move || {
let events = EventLog::read_events(&initial_path)?;
let state = reducer::fold(&events)?;
Ok((events, state))
})
.await
else {
close(&mut socket).await;
return;
};
let head = state.last_seq;
match since {
Some(s) if s <= head && head - s <= MAX_REPLAY_GAP => {
for event in events.iter().filter(|e| e.seq > s) {
if send_json(&mut socket, &event_frame(event)).await.is_err() {
return;
}
}
}
_ => {
let frame = json!({ "type": "snapshot", "seq": head, "state": &state });
if send_json(&mut socket, &frame).await.is_err() {
return;
}
}
}
let mut last_seq = head;
let mut ticker = tokio::time::interval(POLL_INTERVAL);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
_ = ticker.tick() => {
let read_path = events_path.clone();
let new_events = match read_work.run(move || {
Ok(EventLog::read_events_after(&read_path, last_seq)?)
}).await {
Ok(events) => events,
Err(error) if error.status == StatusCode::SERVICE_UNAVAILABLE => continue,
Err(_) => {
close(&mut socket).await;
return;
}
};
for event in &new_events {
if send_json(&mut socket, &event_frame(event)).await.is_err() {
return;
}
if reducer::apply(&mut state, event).is_err() {
let refold_path = events_path.clone();
let seq = event.seq;
match read_work.run(move || {
Ok(refold_at(&refold_path, seq)?)
}).await {
Ok(rebuilt) => state = rebuilt,
Err(_) => {
close(&mut socket).await;
return;
}
}
}
last_seq = event.seq;
if !matches!(event.kind, EventKind::WorkerMessage { .. }) {
let frame =
json!({ "type": "state", "seq": state.last_seq, "state": &state });
if send_json(&mut socket, &frame).await.is_err() {
return;
}
}
}
}
incoming = socket.recv() => {
match incoming {
Some(Ok(Message::Text(text))) => {
if is_ping(&text)
&& send_json(&mut socket, &json!({ "type": "pong" })).await.is_err()
{
return;
}
}
Some(Ok(Message::Close(_))) | Some(Err(_)) | None => return,
Some(Ok(_)) => {} }
}
}
}
}
fn event_frame(event: &Event) -> Value {
json!({ "type": "event", "seq": event.seq, "event": event })
}
fn is_ping(text: &str) -> bool {
serde_json::from_str::<Value>(text)
.ok()
.and_then(|v| v.get("type").map(|t| t == "ping"))
.unwrap_or(false)
}
fn refold_at(events_path: &Path, seq: u64) -> kranz_engine::error::Result<MissionState> {
let events = EventLog::read_events(events_path)?;
let upto = usize::try_from(seq).unwrap_or(usize::MAX).min(events.len());
reducer::fold(&events[..upto])
}
async fn send_json(socket: &mut WebSocket, value: &Value) -> Result<(), axum::Error> {
socket.send(Message::Text(value.to_string().into())).await
}
async fn close(socket: &mut WebSocket) {
let _ = socket.send(Message::Close(None)).await;
}