use std::time::Duration; use axum::extract::State; use axum::http::{HeaderMap, StatusCode}; use axum::response::{IntoResponse, Response}; use axum_tws::{Message, WebSocket, WebSocketUpgrade}; use futures::{SinkExt, StreamExt}; use serde::Deserialize; use tokio::time::{MissedTickBehavior, interval, timeout}; use tracing::{debug, warn}; use crate::api::MultiQuery; use crate::control::{Hydrant, JetstreamFilter, JetstreamSubscriberOptions}; const PING_INTERVAL: Duration = Duration::from_secs(30); const CLOSE_TIMEOUT: Duration = Duration::from_secs(1); const MAX_SUBSCRIBER_MESSAGE_BYTES: usize = 10_000_000; #[derive(Deserialize)] pub struct JetstreamQuery { #[serde(default, rename = "wantedCollections")] wanted_collections: Vec, #[serde(default, rename = "wantedDids")] wanted_dids: Vec, #[serde(default, rename = "maxMessageSizeBytes", alias = "maxSize")] max_message_size_bytes: Option, cursor: Option, #[serde(default)] compress: bool, #[serde(default, rename = "requireHello")] require_hello: bool, #[serde(default, rename = "wantedEventTypes")] wanted_event_types: Vec, } pub async fn handle_subscribe( State(hydrant): State, MultiQuery(query): MultiQuery, headers: HeaderMap, ws: WebSocketUpgrade, ) -> Response { let options = match options_from_parts( &query.wanted_collections, &query.wanted_dids, parse_max_message_size(query.max_message_size_bytes), &query.wanted_event_types, ) { Ok(options) => JetstreamFilter::new(options), Err(err) => return (StatusCode::BAD_REQUEST, err).into_response(), }; let socket_encoding = headers .get("Socket-Encoding") .and_then(|v| v.to_str().ok()) .unwrap_or_default(); let compress = query.compress || socket_encoding.contains("zstd"); let cursor = query .cursor .filter(|cursor| *cursor <= chrono::Utc::now().timestamp_micros()); ws.on_upgrade(move |socket| { handle_socket( socket, hydrant, cursor, options, compress, query.require_hello, ) }) .into_response() } async fn handle_socket( socket: WebSocket, hydrant: Hydrant, cursor: Option, options: JetstreamFilter, compress: bool, require_hello: bool, ) { let send_timeout = hydrant.stream_send_timeout(); let (mut sink, mut ws_recv) = socket.split(); if require_hello { loop { match ws_recv.next().await { Some(Ok(m)) if m.is_ping() => { if sink.send(Message::pong(m.into_payload())).await.is_err() { return; } } Some(Ok(m)) if m.is_close() => return, Some(Ok(m)) if m.is_text() => match handle_options_message(m, &options) { Ok(true) => break, Ok(false) => return, Err(_) => return, }, Some(Ok(_)) => {} Some(Err(err)) => { warn!(err = %err, "Jetstream ws recv error while waiting for hello"); return; } None => return, } } } let mut events = hydrant.subscribe_jetstream(cursor, options.clone()); let mut ping_timer = interval(PING_INTERVAL); ping_timer.set_missed_tick_behavior(MissedTickBehavior::Delay); ping_timer.tick().await; loop { tokio::select! { inbound = ws_recv.next() => match inbound { Some(Ok(m)) if m.is_close() => break, Some(Ok(m)) if m.is_ping() => { if let Err(err) = sink.send(Message::pong(m.into_payload())).await { warn!(err = %err, "Jetstream ws pong send error"); break; } } Some(Ok(m)) if m.is_text() => { if let Err(err) = handle_options_message(m, &options) { debug!(err = %err, "invalid Jetstream options update"); break; } } Some(Ok(m)) if m.is_binary() => { debug!("Jetstream client sent binary frame, closing"); break; } Some(Ok(_)) => {} Some(Err(err)) => { warn!(err = %err, "Jetstream ws recv error"); break; } None => break, }, evt = events.next() => match evt { Some(Ok(json)) => { let msg = match frame_for_event(&json, compress, &options) { Ok(Some(msg)) => msg, Ok(None) => continue, Err(err) => { warn!(err = %err, "failed to encode Jetstream frame"); break; } }; match timeout(send_timeout, sink.send(msg)).await { Ok(Ok(())) => {} Ok(Err(err)) => { warn!(err = %err, "Jetstream ws send error"); break; } Err(_) => { let _ = timeout( CLOSE_TIMEOUT, sink.send(error_message("ConsumerTooSlow", &format!( "Jetstream socket send blocked for at least {} seconds", send_timeout.as_secs() ))), ) .await; break; } } } Some(Err(err)) => { let _ = timeout(send_timeout, sink.send(error_message(err.code(), &err.to_string()))).await; break; } None => break, }, _ = ping_timer.tick() => { if let Err(err) = sink.send(Message::ping(bytes::Bytes::new())).await { warn!(err = %err, "Jetstream ws ping send error"); break; } } } } let _ = timeout(CLOSE_TIMEOUT, sink.close()).await; } fn handle_options_message(msg: Message, options: &JetstreamFilter) -> Result { let payload: bytes::Bytes = msg.into_payload().into(); if payload.len() > MAX_SUBSCRIBER_MESSAGE_BYTES { return Err("subscriber message too large".into()); } let msg: SubscriberSourcedMessage = serde_json::from_slice(&payload).map_err(|e| e.to_string())?; if msg.kind != "options_update" { return Ok(false); } let payload: SubscriberOptionsUpdatePayload = serde_json::from_value(msg.payload).map_err(|e| e.to_string())?; let next = options_from_parts( &payload.wanted_collections, &payload.wanted_dids, parse_max_message_size(Some(payload.max_message_size_bytes)), &payload.wanted_event_types, )?; options.update(next); Ok(true) } thread_local! { static ZSTD_COMPRESSOR: std::cell::RefCell> = std::cell::RefCell::new( zstd::bulk::Compressor::new(3).expect("failed to initialize zstd compressor") ); } fn frame_for_event( json: &[u8], compress: bool, options: &JetstreamFilter, ) -> Result, String> { if exceeds_max(json.len(), options) { return Ok(None); } if compress { let compressed = ZSTD_COMPRESSOR .with(|c| c.borrow_mut().compress(json)) .map_err(|e| e.to_string())?; if exceeds_max(compressed.len(), options) { return Ok(None); } return Ok(Some(Message::binary(compressed))); } let text = std::str::from_utf8(json).map_err(|e| e.to_string())?; Ok(Some(Message::text(text.to_string()))) } fn exceeds_max(size: usize, options: &JetstreamFilter) -> bool { let max = options.max_message_size_bytes(); max > 0 && size > max as usize } fn options_from_parts( wanted_collections: &[String], wanted_dids: &[String], max_message_size_bytes: u32, wanted_event_types: &[String], ) -> Result { JetstreamSubscriberOptions::parse( wanted_collections, wanted_dids, max_message_size_bytes, wanted_event_types, ) } fn parse_max_message_size(value: Option) -> u32 { value .filter(|v| *v > 0) .and_then(|v| u32::try_from(v).ok()) .unwrap_or(0) } fn error_message(code: &str, message: &str) -> Message { Message::text( serde_json::json!({ "type": "error", "error": code, "message": message, }) .to_string(), ) } #[derive(Deserialize)] struct SubscriberSourcedMessage { #[serde(rename = "type")] kind: String, payload: serde_json::Value, } #[derive(Deserialize)] struct SubscriberOptionsUpdatePayload { #[serde(default, rename = "wantedCollections")] wanted_collections: Vec, #[serde(default, rename = "wantedDids")] wanted_dids: Vec, #[serde(default, rename = "maxMessageSizeBytes", alias = "maxSize")] max_message_size_bytes: i64, #[serde(default, rename = "wantedEventTypes")] wanted_event_types: Vec, } #[cfg(test)] mod tests { use super::*; #[test] fn test_frame_for_event_size_limits() { let json = b"{\"hello\": \"world\"}"; // options with max size 5 let options = JetstreamFilter::new(JetstreamSubscriberOptions::parse(&[], &[], 5, &[]).unwrap()); // Uncompressed exceeds limit let res = frame_for_event(json, false, &options).unwrap(); assert!(res.is_none()); // Compressed with uncompressed exceeding limit let res = frame_for_event(json, true, &options).unwrap(); assert!(res.is_none()); // Fits within limit let options_large = JetstreamFilter::new(JetstreamSubscriberOptions::parse(&[], &[], 100, &[]).unwrap()); let res_uncompressed = frame_for_event(json, false, &options_large) .unwrap() .unwrap(); assert!(res_uncompressed.is_text()); let res_compressed = frame_for_event(json, true, &options_large) .unwrap() .unwrap(); assert!(res_compressed.is_binary()); } #[test] fn jetstream_query_accepts_repeated_values() { let query: JetstreamQuery = serde_html_form::from_str( "wantedCollections=app.bsky.feed.post&wantedCollections=app.bsky.feed.like\ &wantedDids=did%3Aplc%3Atest&wantedEventTypes=live", ) .unwrap(); assert_eq!( query.wanted_collections, ["app.bsky.feed.post", "app.bsky.feed.like"] ); assert_eq!(query.wanted_dids, ["did:plc:test"]); assert_eq!(query.wanted_event_types, ["live"]); let defaults: JetstreamQuery = serde_html_form::from_str("").unwrap(); assert!(defaults.wanted_collections.is_empty()); assert!(defaults.wanted_dids.is_empty()); assert!(defaults.wanted_event_types.is_empty()); } #[test] fn jetstream_options_message_json_contract() { let options = JetstreamFilter::new(JetstreamSubscriberOptions::parse(&[], &[], 0, &[]).unwrap()); let message = Message::text( serde_json::json!({ "type": "options_update", "payload": { "wantedCollections": ["app.bsky.feed.post"], "wantedDids": ["did:web:example.com"], "maxSize": 4096, "wantedEventTypes": ["live"], }, }) .to_string(), ); assert_eq!(handle_options_message(message, &options), Ok(true)); assert_eq!(options.max_message_size_bytes(), 4096); let ignored = Message::text(r#"{"type":"stinkpot","payload":{}}"#); assert_eq!(handle_options_message(ignored, &options), Ok(false)); assert!(handle_options_message(Message::text("not json"), &options).is_err()); } #[test] fn jetstream_message_size_conversion_rejects_invalid_ranges() { assert_eq!(parse_max_message_size(None), 0); assert_eq!(parse_max_message_size(Some(-1)), 0); assert_eq!(parse_max_message_size(Some(0)), 0); assert_eq!(parse_max_message_size(Some(1)), 1); assert_eq!(parse_max_message_size(Some(u32::MAX.into())), u32::MAX); assert_eq!(parse_max_message_size(Some(i64::from(u32::MAX) + 1)), 0); } }