use std::sync::Arc;
use std::time::Instant;
use axum::Extension;
use axum::extract::ConnectInfo;
use axum::extract::Query;
use axum::extract::ws::{CloseFrame, Message, Utf8Bytes, WebSocket, WebSocketUpgrade};
use axum::http::StatusCode;
use axum::response::IntoResponse;
use bytes::Bytes;
use futures_util::stream::StreamExt as _;
use futures_util::{SinkExt as _, stream::SplitSink};
use proto_blue_lex_cbor::encode as cbor_encode;
use proto_blue_lex_data::LexValue;
use serde::Deserialize;
use sqlx::{Pool, Sqlite};
use std::collections::BTreeMap;
use std::net::SocketAddr;
use tokio::sync::broadcast;
use tokio::time::{MissedTickBehavior, interval};
use crate::error::{Error, Result};
use crate::label::Label;
use crate::signing::label_to_lex_value_with_sig;
use crate::writer::{LabelEvent, WriterHandle};
use super::SubscribeConfig;
use super::limits::Limiter;
const CLOSE_POLICY_VIOLATION: u16 = 1008;
const CLOSE_NORMAL: u16 = 1000;
const CLOSE_GOING_AWAY: u16 = 1001;
#[derive(Debug, Deserialize)]
pub(super) struct CursorParams {
cursor: Option<String>,
}
#[derive(Debug, PartialEq, Eq)]
pub(super) enum CursorDecision {
LiveOnly,
Replay {
emit_outdated: bool,
from_exclusive: i64,
},
FutureCursor,
Malformed,
}
pub(super) fn decide_cursor(
raw: Option<&str>,
oldest_retained: Option<i64>,
head: i64,
) -> CursorDecision {
let Some(raw) = raw else {
return CursorDecision::LiveOnly;
};
let Ok(c) = raw.parse::<i64>() else {
return CursorDecision::Malformed;
};
if c > head {
return CursorDecision::FutureCursor;
}
let oldest = oldest_retained.unwrap_or(1);
if c <= 0 || c < oldest {
CursorDecision::Replay {
emit_outdated: true,
from_exclusive: (oldest - 1).max(0),
}
} else {
CursorDecision::Replay {
emit_outdated: false,
from_exclusive: c,
}
}
}
fn frame_bytes(header: LexValue, body: LexValue) -> Result<Bytes> {
let mut out = cbor_encode(&header)?;
out.extend_from_slice(&cbor_encode(&body)?);
Ok(Bytes::from(out))
}
fn header(op: i64, t: Option<&str>) -> LexValue {
let mut m = BTreeMap::new();
m.insert("op".to_string(), LexValue::Integer(op));
if let Some(t) = t {
m.insert("t".to_string(), LexValue::String(t.to_string()));
}
LexValue::Map(m)
}
pub(super) fn labels_frame(seq: i64, label: &Label) -> Result<Bytes> {
let h = header(1, Some("#labels"));
let mut body = BTreeMap::new();
body.insert("seq".to_string(), LexValue::Integer(seq));
body.insert(
"labels".to_string(),
LexValue::Array(vec![label_to_lex_value_with_sig(label)?]),
);
frame_bytes(h, LexValue::Map(body))
}
pub(super) fn info_frame(name: &str, message: Option<&str>) -> Result<Bytes> {
let h = header(1, Some("#info"));
let mut body = BTreeMap::new();
body.insert("name".to_string(), LexValue::String(name.to_string()));
if let Some(m) = message {
body.insert("message".to_string(), LexValue::String(m.to_string()));
}
frame_bytes(h, LexValue::Map(body))
}
pub(super) fn error_frame(error: &str, message: Option<&str>) -> Result<Bytes> {
let h = header(-1, None);
let mut body = BTreeMap::new();
body.insert("error".to_string(), LexValue::String(error.to_string()));
if let Some(m) = message {
body.insert("message".to_string(), LexValue::String(m.to_string()));
}
frame_bytes(h, LexValue::Map(body))
}
pub(super) async fn handler(
ws: WebSocketUpgrade,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
Query(params): Query<CursorParams>,
Extension(state): Extension<AppState>,
) -> axum::response::Response {
let Some(permit) = state.limiter.try_acquire(addr.ip()) else {
return (StatusCode::SERVICE_UNAVAILABLE, "subscriber cap reached").into_response();
};
let bcast = state.writer.subscribe();
let shutdown_rx = state.writer.shutdown_signal();
ws.on_upgrade(move |socket| async move {
let permit = permit; if let Err(err) = run_subscription(socket, state, params.cursor, bcast, shutdown_rx).await {
tracing::warn!(%addr, error = %err, "subscribeLabels connection ended with error");
}
drop(permit);
})
}
#[derive(Clone)]
pub(super) struct AppState {
pub pool: Pool<Sqlite>,
pub writer: WriterHandle,
pub limiter: Arc<Limiter>,
pub config: Arc<SubscribeConfig>,
}
async fn run_subscription(
socket: WebSocket,
state: AppState,
cursor_raw: Option<String>,
bcast: broadcast::Receiver<LabelEvent>,
shutdown_rx: tokio::sync::watch::Receiver<bool>,
) -> Result<()> {
let (mut sink, mut stream) = socket.split();
let oldest_retained = query_oldest_retained(&state.pool, &state.config).await?;
let head_at_join = query_current_head(&state.pool).await?;
let decision = decide_cursor(cursor_raw.as_deref(), oldest_retained, head_at_join);
match decision {
CursorDecision::Malformed => {
let _ = sink
.send(Message::Close(Some(CloseFrame {
code: CLOSE_POLICY_VIOLATION,
reason: Utf8Bytes::from_static("invalid cursor"),
})))
.await;
Ok(())
}
CursorDecision::FutureCursor => {
let frame = error_frame("FutureCursor", Some(&format!("head={head_at_join}")))?;
let _ = sink.send(Message::Binary(frame)).await;
let _ = sink
.send(Message::Close(Some(CloseFrame {
code: CLOSE_NORMAL,
reason: Utf8Bytes::from_static("FutureCursor"),
})))
.await;
Ok(())
}
CursorDecision::LiveOnly => {
live_tail(&mut sink, &mut stream, bcast, shutdown_rx, 0, &state).await
}
CursorDecision::Replay {
emit_outdated,
from_exclusive,
} => {
if emit_outdated {
let oldest_str = oldest_retained
.map(|n| n.to_string())
.unwrap_or_else(|| "0".into());
let frame = info_frame(
"OutdatedCursor",
Some(&format!("oldest_retained_seq={oldest_str}")),
)?;
sink.send(Message::Binary(frame)).await.ok();
}
replay_range(
&mut sink,
&state.pool,
from_exclusive,
head_at_join,
&state.config,
)
.await?;
live_tail(
&mut sink,
&mut stream,
bcast,
shutdown_rx,
head_at_join,
&state,
)
.await
}
}
}
type WsSink = SplitSink<WebSocket, Message>;
async fn live_tail(
sink: &mut WsSink,
stream: &mut futures_util::stream::SplitStream<WebSocket>,
mut bcast: broadcast::Receiver<LabelEvent>,
mut shutdown_rx: tokio::sync::watch::Receiver<bool>,
skip_seq_at_or_below: i64,
state: &AppState,
) -> Result<()> {
let mut ping_timer = interval(state.config.ping_interval);
ping_timer.set_missed_tick_behavior(MissedTickBehavior::Delay);
ping_timer.tick().await;
let mut last_pong = Instant::now();
loop {
tokio::select! {
biased;
res = bcast.recv() => match res {
Ok(event) => {
if event.seq <= skip_seq_at_or_below {
continue;
}
let frame = labels_frame(event.seq, &event.label)?;
if sink.send(Message::Binary(frame)).await.is_err() {
return Ok(()); }
}
Err(broadcast::error::RecvError::Lagged(_)) => {
let _ = sink.send(Message::Close(Some(CloseFrame {
code: CLOSE_GOING_AWAY,
reason: Utf8Bytes::from_static("subscriber lagged past buffer"),
}))).await;
return Ok(());
}
Err(broadcast::error::RecvError::Closed) => {
let _ = sink.send(Message::Close(Some(CloseFrame {
code: CLOSE_GOING_AWAY,
reason: Utf8Bytes::from_static("server shutting down"),
}))).await;
return Ok(());
}
},
msg = stream.next() => match msg {
None | Some(Err(_)) => return Ok(()), Some(Ok(Message::Close(_))) => return Ok(()),
Some(Ok(Message::Pong(_))) => {
last_pong = Instant::now();
}
Some(Ok(Message::Ping(_))) => {
last_pong = Instant::now();
}
Some(Ok(Message::Text(_) | Message::Binary(_))) => {
let _ = sink.send(Message::Close(Some(CloseFrame {
code: CLOSE_POLICY_VIOLATION,
reason: Utf8Bytes::from_static("client must not send application data"),
}))).await;
return Ok(());
}
},
res = shutdown_rx.changed() => {
if res.is_err() || *shutdown_rx.borrow() {
let _ = sink.send(Message::Close(Some(CloseFrame {
code: CLOSE_GOING_AWAY,
reason: Utf8Bytes::from_static("server shutting down"),
}))).await;
return Ok(());
}
},
_ = ping_timer.tick() => {
if last_pong.elapsed() > state.config.pong_timeout {
let _ = sink.send(Message::Close(Some(CloseFrame {
code: CLOSE_GOING_AWAY,
reason: Utf8Bytes::from_static("pong timeout"),
}))).await;
return Ok(());
}
if sink.send(Message::Ping(Bytes::new())).await.is_err() {
return Ok(());
}
}
}
}
}
async fn query_current_head(pool: &Pool<Sqlite>) -> Result<i64> {
let v: Option<i64> = sqlx::query_scalar!(r#"SELECT MAX(seq) AS "max_seq?: i64" FROM labels"#)
.fetch_one(pool)
.await?;
Ok(v.unwrap_or(0))
}
pub async fn current_retention_floor(
pool: &Pool<Sqlite>,
retention_days: Option<u32>,
) -> Result<Option<i64>> {
let cfg = SubscribeConfig {
retention_days,
..SubscribeConfig::default()
};
query_oldest_retained(pool, &cfg).await
}
async fn query_oldest_retained(
pool: &Pool<Sqlite>,
config: &SubscribeConfig,
) -> Result<Option<i64>> {
match config.retention_days {
None => {
let v: Option<i64> =
sqlx::query_scalar!(r#"SELECT MIN(seq) AS "min_seq?: i64" FROM labels"#)
.fetch_one(pool)
.await?;
Ok(v)
}
Some(days) => {
let cutoff_ms = crate::writer::epoch_ms_now() - (days as i64) * 86_400_000;
let v: Option<i64> = sqlx::query_scalar!(
r#"SELECT MIN(seq) AS "min_seq?: i64" FROM labels WHERE created_at >= ?1"#,
cutoff_ms
)
.fetch_one(pool)
.await?;
Ok(v)
}
}
}
async fn replay_range(
sink: &mut WsSink,
pool: &Pool<Sqlite>,
from_exclusive: i64,
head: i64,
config: &SubscribeConfig,
) -> Result<()> {
let mut cursor = from_exclusive;
loop {
let batch = sqlx::query!(
r#"SELECT
seq AS "seq!: i64",
ver AS "ver!: i64",
src, uri, cid, val,
neg AS "neg!: i64",
cts, exp, sig
FROM labels l1
WHERE l1.seq > ?1
AND l1.seq <= ?2
AND NOT (
l1.neg = 0
AND EXISTS (
SELECT 1 FROM labels l2
WHERE l2.src = l1.src
AND l2.uri = l1.uri
AND l2.val = l1.val
AND l2.neg = 1
AND l2.seq > l1.seq
)
)
ORDER BY l1.seq ASC
LIMIT ?3"#,
cursor,
head,
config.batch_size,
)
.fetch_all(pool)
.await?;
if batch.is_empty() {
return Ok(());
}
for row in &batch {
let sig: [u8; 64] = row.sig.as_slice().try_into().map_err(|_| {
Error::Signing(format!(
"corrupt sig in labels.seq={}: expected 64 bytes, got {}",
row.seq,
row.sig.len()
))
})?;
let label = Label {
ver: row.ver,
src: row.src.clone(),
uri: row.uri.clone(),
cid: row.cid.clone(),
val: row.val.clone(),
neg: row.neg != 0,
cts: row.cts.clone(),
exp: row.exp.clone(),
sig: Some(sig),
};
let frame = labels_frame(row.seq, &label)?;
if sink.send(Message::Binary(frame)).await.is_err() {
return Ok(());
}
}
cursor = batch.last().expect("non-empty checked above").seq;
if cursor >= head {
return Ok(());
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cursor_absent_is_live_only() {
let d = decide_cursor(None, Some(1), 10);
assert_eq!(d, CursorDecision::LiveOnly);
}
#[test]
fn cursor_zero_triggers_outdated_and_replay_from_oldest() {
let d = decide_cursor(Some("0"), Some(5), 10);
assert_eq!(
d,
CursorDecision::Replay {
emit_outdated: true,
from_exclusive: 4,
}
);
}
#[test]
fn cursor_below_oldest_triggers_outdated() {
let d = decide_cursor(Some("3"), Some(5), 10);
assert_eq!(
d,
CursorDecision::Replay {
emit_outdated: true,
from_exclusive: 4,
}
);
}
#[test]
fn cursor_at_or_above_oldest_within_head_is_plain_replay() {
assert_eq!(
decide_cursor(Some("5"), Some(5), 10),
CursorDecision::Replay {
emit_outdated: false,
from_exclusive: 5,
}
);
assert_eq!(
decide_cursor(Some("7"), Some(5), 10),
CursorDecision::Replay {
emit_outdated: false,
from_exclusive: 7,
}
);
}
#[test]
fn cursor_exceeds_head_is_future() {
assert_eq!(
decide_cursor(Some("11"), Some(5), 10),
CursorDecision::FutureCursor
);
}
#[test]
fn cursor_non_integer_is_malformed() {
assert_eq!(
decide_cursor(Some("abc"), Some(5), 10),
CursorDecision::Malformed
);
assert_eq!(
decide_cursor(Some(""), Some(5), 10),
CursorDecision::Malformed
);
}
#[test]
fn empty_table_cursor_zero_replays_nothing_but_emits_outdated() {
let d = decide_cursor(Some("0"), None, 0);
assert_eq!(
d,
CursorDecision::Replay {
emit_outdated: true,
from_exclusive: 0,
}
);
}
#[test]
fn labels_frame_starts_with_op_and_t_header() {
let label = Label {
ver: 1,
src: "did:plc:x".into(),
uri: "at://did:plc:x/a/b".into(),
cid: None,
val: "spam".into(),
neg: false,
cts: "2026-04-22T12:00:00.000Z".into(),
exp: None,
sig: Some([0xAA; 64]),
};
let bytes = labels_frame(1, &label).expect("frame");
assert_eq!(bytes[0], 0xA2, "first byte of header must be 2-key map");
assert!(
bytes.windows(7).any(|w| w == b"#labels"),
"labels frame must contain #labels tag"
);
}
#[test]
fn info_frame_carries_outdated_cursor_name() {
let bytes = info_frame("OutdatedCursor", Some("oldest=5")).expect("frame");
assert!(
bytes
.windows("OutdatedCursor".len())
.any(|w| w == b"OutdatedCursor"),
"info frame must contain OutdatedCursor"
);
}
#[test]
fn error_frame_op_minus_one() {
let bytes = error_frame("FutureCursor", None).expect("frame");
assert_eq!(bytes[0], 0xA1);
assert!(
bytes
.windows("FutureCursor".len())
.any(|w| w == b"FutureCursor"),
"error frame must contain FutureCursor"
);
}
}