very fast at protocol indexer with flexible filtering, xrpc queries, cursor-backed event stream, and more, built on fjall
rust fjall at-protocol atproto indexer
Something went wrong. Try again.
Rust
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382use 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<String>, #[serde(default, rename = "wantedDids")] wanted_dids: Vec<String>, #[serde(default, rename = "maxMessageSizeBytes", alias = "maxSize")] max_message_size_bytes: Option<i64>, cursor: Option<i64>, #[serde(default)] compress: bool, #[serde(default, rename = "requireHello")] require_hello: bool, #[serde(default, rename = "wantedEventTypes")] wanted_event_types: Vec<String>,}
pub async fn handle_subscribe( State(hydrant): State<Hydrant>, MultiQuery(query): MultiQuery<JetstreamQuery>, 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<i64>, 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<bool, String> { 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<zstd::bulk::Compressor<'static>> = 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<Option<Message>, 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, String> { JetstreamSubscriberOptions::parse( wanted_collections, wanted_dids, max_message_size_bytes, wanted_event_types, )}
fn parse_max_message_size(value: Option<i64>) -> 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<String>, #[serde(default, rename = "wantedDids")] wanted_dids: Vec<String>, #[serde(default, rename = "maxMessageSizeBytes", alias = "maxSize")] max_message_size_bytes: i64, #[serde(default, rename = "wantedEventTypes")] wanted_event_types: Vec<String>,}
#[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); }}