use anyhow::Result;
use async_trait::async_trait;
use futures::{SinkExt, StreamExt};
use http::Uri;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use std::{collections::HashMap, str::FromStr};
use tokio::sync::{RwLock, broadcast};
use tokio::time::{Instant, sleep};
use tokio_util::sync::CancellationToken;
use tokio_websockets::{ClientBuilder, Message};
use tracing::Instrument;
const MAX_MESSAGE_SIZE: usize = 56000;
#[cfg_attr(debug_assertions, derive(Debug))]
#[derive(Clone)]
pub struct ConsumerTaskConfig {
pub user_agent: String,
pub compression: bool,
pub zstd_dictionary_location: String,
pub jetstream_hostname: String,
pub collections: Vec<String>,
pub dids: Vec<String>,
pub max_message_size_bytes: Option<u64>,
pub cursor: Option<i64>,
pub require_hello: bool,
}
#[cfg_attr(debug_assertions, derive(Debug))]
#[derive(Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum JetstreamEvent {
Commit {
did: String,
time_us: u64,
kind: String,
#[serde(rename = "commit")]
commit: JetstreamEventCommit,
},
Delete {
did: String,
time_us: u64,
kind: String,
#[serde(rename = "commit")]
commit: JetstreamEventDelete,
},
Identity {
did: String,
time_us: u64,
kind: String,
#[serde(rename = "identity")]
identity: serde_json::Value,
},
Account {
did: String,
time_us: u64,
kind: String,
#[serde(rename = "account")]
identity: serde_json::Value,
},
}
#[cfg_attr(debug_assertions, derive(Debug))]
#[derive(Clone, Serialize, Deserialize)]
pub struct JetstreamEventCommit {
pub rev: String,
pub operation: String,
pub collection: String,
pub rkey: String,
pub cid: String,
pub record: serde_json::Value,
}
#[cfg_attr(debug_assertions, derive(Debug))]
#[derive(Clone, Serialize, Deserialize)]
pub struct JetstreamEventDelete {
pub rev: String,
pub operation: String,
pub collection: String,
pub rkey: String,
}
#[async_trait]
pub trait EventHandler: Send + Sync {
async fn handle_event(&self, event: JetstreamEvent) -> Result<()>;
fn handler_id(&self) -> String;
}
#[derive(thiserror::Error, Debug)]
pub enum ConsumerError {
#[error("error-atproto-jetstream-consumer-1 WebSocket connection failed: {0}")]
ConnectionFailed(String),
#[error("error-atproto-jetstream-consumer-2 Message decompression failed: {0}")]
DecompressionFailed(String),
#[error("error-atproto-jetstream-consumer-3 Event deserialization failed: {0}")]
DeserializationFailed(String),
#[error("error-atproto-jetstream-consumer-4 Handler registration failed: {0}")]
HandlerRegistrationFailed(String),
#[error("error-atproto-jetstream-consumer-5 Event sender not initialized: {0}")]
EventSenderNotInitialized(String),
#[error("error-atproto-jetstream-consumer-6 Message conversion failed: {0}")]
MessageConversionFailed(String),
#[error("error-atproto-jetstream-consumer-7 Update serialization failed: {0}")]
UpdateSerializationFailed(String),
#[error("error-atproto-jetstream-consumer-8 Update send failed: {0}")]
UpdateSendFailed(String),
#[error("error-atproto-jetstream-consumer-9 Decompressor creation failed: {0}")]
DecompressorCreationFailed(String),
}
#[cfg_attr(debug_assertions, derive(Debug))]
#[derive(Clone, Serialize, Deserialize)]
#[serde(tag = "type", content = "payload")]
pub(crate) enum SubscriberSourcedMessage {
#[serde(rename = "options_update")]
Update {
#[serde(rename = "wantedCollections")]
wanted_collections: Vec<String>,
#[serde(rename = "wantedDids", skip_serializing_if = "Vec::is_empty", default)]
wanted_dids: Vec<String>,
#[serde(rename = "maxMessageSizeBytes")]
max_message_size_bytes: u64,
#[serde(skip_serializing_if = "Option::is_none")]
cursor: Option<i64>,
},
}
pub struct Consumer {
config: ConsumerTaskConfig,
handlers: Arc<RwLock<HashMap<String, Arc<dyn EventHandler>>>>,
event_sender: Arc<RwLock<Option<broadcast::Sender<JetstreamEvent>>>>,
}
impl Consumer {
pub fn new(config: ConsumerTaskConfig) -> Self {
Self {
config,
handlers: Arc::new(RwLock::new(HashMap::new())),
event_sender: Arc::new(RwLock::new(None)),
}
}
pub async fn register_handler(&self, handler: Arc<dyn EventHandler>) -> Result<()> {
let handler_id = handler.handler_id();
let mut handlers = self.handlers.write().await;
if handlers.contains_key(&handler_id) {
return Err(ConsumerError::HandlerRegistrationFailed(format!(
"Handler with ID '{}' already registered",
handler_id
))
.into());
}
handlers.insert(handler_id.clone(), handler);
Ok(())
}
pub async fn unregister_handler(&self, handler_id: &str) -> Result<()> {
let mut handlers = self.handlers.write().await;
handlers.remove(handler_id);
Ok(())
}
pub async fn get_event_receiver(&self) -> Result<broadcast::Receiver<JetstreamEvent>> {
let sender_guard = self.event_sender.read().await;
match sender_guard.as_ref() {
Some(sender) => Ok(sender.subscribe()),
None => Err(ConsumerError::EventSenderNotInitialized(
"consumer not running".to_string(),
)
.into()),
}
}
pub async fn run_background(&self, cancellation_token: CancellationToken) -> Result<()> {
tracing::info!("Starting Jetstream consumer");
let mut query_params = vec![];
query_params.push(format!("compress={}", self.config.compression));
query_params.push(format!("requireHello={}", self.config.require_hello));
if !self.config.collections.is_empty() {
let collections = self
.config
.collections
.iter()
.map(|c| urlencoding::encode(c))
.collect::<Vec<_>>()
.join(",");
query_params.push(format!("wantedCollections={}", collections));
}
if !self.config.dids.is_empty() {
let dids = self
.config
.dids
.iter()
.map(|d| urlencoding::encode(d))
.collect::<Vec<_>>()
.join(",");
query_params.push(format!("wantedDids={}", dids));
}
if let Some(max_size) = self.config.max_message_size_bytes {
query_params.push(format!("maxMessageSizeBytes={}", max_size));
}
if let Some(cursor) = self.config.cursor {
query_params.push(format!("cursor={}", cursor));
}
let query_string = query_params.join("&");
let ws_url = Uri::from_str(&format!(
"wss://{}/subscribe?{}",
self.config.jetstream_hostname, query_string
))?;
let (mut client, _) = ClientBuilder::from_uri(ws_url)
.add_header(
http::header::USER_AGENT,
http::HeaderValue::from_str(&self.config.user_agent)?,
)?
.connect()
.await?;
let update = SubscriberSourcedMessage::Update {
wanted_collections: self.config.collections.clone(),
wanted_dids: self.config.dids.clone(),
max_message_size_bytes: self
.config
.max_message_size_bytes
.unwrap_or(MAX_MESSAGE_SIZE as u64),
cursor: self.config.cursor,
};
let serialized_update = serde_json::to_string(&update)
.map_err(|err| ConsumerError::UpdateSerializationFailed(err.to_string()))?;
client
.send(Message::text(serialized_update))
.await
.map_err(|err| ConsumerError::UpdateSendFailed(err.to_string()))?;
let mut decompressor = if self.config.compression {
let data: Vec<u8> = std::fs::read(self.config.zstd_dictionary_location.clone())?;
zstd::bulk::Decompressor::with_dictionary(&data)
.map_err(|err| ConsumerError::DecompressorCreationFailed(err.to_string()))?
} else {
zstd::bulk::Decompressor::new()
.map_err(|err| ConsumerError::DecompressorCreationFailed(err.to_string()))?
};
let interval = std::time::Duration::from_secs(120);
let sleeper = sleep(interval);
tokio::pin!(sleeper);
loop {
tokio::select! {
() = cancellation_token.cancelled() => {
break;
},
() = &mut sleeper => {
sleeper.as_mut().reset(Instant::now() + interval);
},
item = client.next() => {
if item.is_none() {
tracing::warn!("jetstream connection closed");
break;
}
let item = item.unwrap();
if let Err(err) = item {
tracing::error!(error = ?err, "error processing jetstream message");
continue;
}
let item = item.unwrap();
let event = if self.config.compression {
if !item.is_binary() {
tracing::debug!("compression enabled but message from jetstream is not binary");
continue;
}
let payload = item.into_payload();
let decoded = decompressor.decompress(&payload, MAX_MESSAGE_SIZE * 3);
if let Err(err) = decoded {
tracing::debug!(err = ?err, "cannot decompress message");
continue;
}
let decoded = decoded.unwrap();
serde_json::from_slice::<JetstreamEvent>(&decoded)
.map_err(|err| ConsumerError::DeserializationFailed(err.to_string()))
} else {
if !item.is_text() {
tracing::debug!("compression disabled but message from jetstream is binary");
continue;
}
item.as_text()
.ok_or_else(|| ConsumerError::MessageConversionFailed("cannot convert message to text".to_string()))
.and_then(|value| {
serde_json::from_str::<JetstreamEvent>(value)
.map_err(|err| ConsumerError::DeserializationFailed(err.to_string()))
})
};
if let Err(err) = event {
tracing::error!(error = ?err, "error processing jetstream message");
continue;
}
let event = event.unwrap();
if let Err(err) = self.dispatch_to_handlers(event).await {
tracing::error!(error = ?err, "Failed to process message");
}
}
}
}
{
let mut sender_guard = self.event_sender.write().await;
*sender_guard = None;
}
Ok(())
}
async fn dispatch_to_handlers(&self, event: JetstreamEvent) -> Result<()> {
let handlers = self.handlers.read().await;
for (handler_id, handler) in handlers.iter() {
let handler_span = tracing::debug_span!("handler_dispatch", handler_id = %handler_id);
async {
if let Err(err) = handler.handle_event(event.clone()).await {
tracing::error!(
error = ?err,
handler_id = %handler_id,
"Handler failed to process event"
);
}
}
.instrument(handler_span)
.await;
}
Ok(())
}
}
pub struct LoggingHandler {
id: String,
}
impl LoggingHandler {
pub fn new(id: String) -> Self {
Self { id }
}
}
#[async_trait]
impl EventHandler for LoggingHandler {
async fn handle_event(&self, _event: JetstreamEvent) -> Result<()> {
Ok(())
}
fn handler_id(&self) -> String {
self.id.clone()
}
}