use std::collections::HashSet;
use std::convert::Infallible;
use std::time::Duration;
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::extract::{Query, State};
use axum::response::sse::{Event, KeepAlive, Sse};
use axum::response::{IntoResponse, Response};
use futures_util::stream::{self, Stream, StreamExt};
use serde::Deserialize;
use tokio::sync::broadcast::error::RecvError;
use tracing::debug;
use super::dto::EventEnvelope;
use super::projection::{EventsSince, Projection};
use super::RemoteControlState;
const KEEP_ALIVE: Duration = Duration::from_secs(15);
#[derive(Debug, Clone, Default, Deserialize)]
pub struct StreamParams {
#[serde(default)]
pub after_sequence: Option<u64>,
#[serde(default)]
pub instance_id: Option<String>,
}
fn initial_events(projection: &Projection, params: &StreamParams) -> Vec<EventEnvelope> {
let Some(after) = params.after_sequence else {
return Vec::new();
};
if params
.instance_id
.as_deref()
.is_some_and(|id| id != projection.instance_id())
{
return vec![projection.gap_envelope(after)];
}
match projection.events_after(after) {
EventsSince::Replay(events) => events,
EventsSince::Gap => vec![projection.gap_envelope(after)],
}
}
fn event_stream(
projection: std::sync::Arc<Projection>,
params: StreamParams,
) -> impl Stream<Item = EventEnvelope> {
let rx = projection.subscribe();
let replayed = initial_events(&projection, ¶ms);
let already_sent: HashSet<u64> = replayed.iter().map(|e| e.event_sequence).collect();
let live = stream::unfold(
(rx, projection, already_sent),
|(mut rx, projection, already_sent)| async move {
loop {
match rx.recv().await {
Ok(event) => {
if already_sent.contains(&event.event_sequence) {
continue;
}
return Some((event, (rx, projection, already_sent)));
}
Err(RecvError::Lagged(skipped)) => {
debug!("v2 event subscriber lagged by {skipped} events");
let gap = projection.gap_envelope(0);
return Some((gap, (rx, projection, already_sent)));
}
Err(RecvError::Closed) => return None,
}
}
},
);
stream::iter(replayed).chain(live)
}
#[utoipa::path(
get,
path = "/api/v2/events",
tag = "remote-control",
params(
("after_sequence" = Option<u64>, Query, description = "Resume strictly after this event sequence"),
("instance_id" = Option<String>, Query, description = "Incarnation the cursor belongs to")
),
responses(
(
status = 200,
description = "Ordered SSE stream read with fetch() response streaming. Each `data:` \
frame is one EventEnvelope; a `gap` category means the requested cursor \
is no longer replayable and the client must re-read GET /api/v2/state.",
body = super::dto::EventEnvelope,
content_type = "text/event-stream"
),
(status = 401, description = "Missing or invalid bearer credentials", body = super::dto::ApiError)
)
)]
pub async fn events(
State(state): State<RemoteControlState>,
Query(params): Query<StreamParams>,
) -> Sse<impl Stream<Item = Result<Event, Infallible>>> {
let stream = event_stream(state.projection.clone(), params).map(|envelope| {
let event = Event::default()
.id(envelope.event_sequence.to_string())
.event(envelope.event_type.clone());
Ok(match serde_json::to_string(&envelope) {
Ok(json) => event.data(json),
Err(_) => event.data("{}"),
})
});
Sse::new(stream).keep_alive(KeepAlive::new().interval(KEEP_ALIVE))
}
#[utoipa::path(
get,
path = "/api/v2/ws",
tag = "remote-control",
params(
("after_sequence" = Option<u64>, Query, description = "Resume strictly after this event sequence"),
("instance_id" = Option<String>, Query, description = "Incarnation the cursor belongs to")
),
responses(
(
status = 101,
description = "Switching protocols. Every subsequent text frame is one EventEnvelope, \
in the same order and with the same `gap` semantics as the SSE stream."
),
(status = 401, description = "Missing bearer header, or credentials in query/subprotocol", body = super::dto::ApiError)
)
)]
pub async fn ws(
ws: WebSocketUpgrade,
State(state): State<RemoteControlState>,
Query(params): Query<StreamParams>,
) -> Response {
ws.on_upgrade(move |socket| drive_socket(socket, state, params))
.into_response()
}
async fn drive_socket(mut socket: WebSocket, state: RemoteControlState, params: StreamParams) {
let mut stream = Box::pin(event_stream(state.projection.clone(), params));
loop {
tokio::select! {
next = stream.next() => {
let Some(envelope) = next else { break };
let Ok(json) = serde_json::to_string(&envelope) else { continue };
if socket.send(Message::Text(json.into())).await.is_err() {
break;
}
}
incoming = socket.recv() => {
match incoming {
Some(Ok(Message::Ping(data))) => {
if socket.send(Message::Pong(data)).await.is_err() {
break;
}
}
Some(Ok(Message::Close(_))) | None => break,
Some(Err(_)) => break,
Some(Ok(_)) => {}
}
}
}
}
}