use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex, RwLock};
use serde::Serialize;
use serde::de::DeserializeOwned;
use serde_json::Value;
use crate::{AppState, AutumnError, AutumnResult};
pub trait Event: Serialize + DeserializeOwned + Send + Sync + 'static {
const NAME: &'static str;
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum DispatchMode {
Sync,
Durable,
}
pub type ListenerHandler =
fn(AppState, Value) -> Pin<Box<dyn Future<Output = AutumnResult<()>> + Send + 'static>>;
#[derive(Clone)]
pub struct ListenerInfo {
pub event_name: &'static str,
pub listener_name: String,
pub mode: DispatchMode,
pub job_name: Option<String>,
pub max_attempts: u32,
pub initial_backoff_ms: u64,
pub handler: ListenerHandler,
}
#[derive(Clone, Default)]
pub struct EventRegistry {
by_event: Arc<HashMap<&'static str, Vec<ListenerInfo>>>,
}
impl EventRegistry {
#[must_use]
pub fn from_listeners(listeners: Vec<ListenerInfo>) -> Self {
let mut by_event: HashMap<&'static str, Vec<ListenerInfo>> = HashMap::new();
for listener in listeners {
by_event
.entry(listener.event_name)
.or_default()
.push(listener);
}
Self {
by_event: Arc::new(by_event),
}
}
#[must_use]
pub fn listeners_for(&self, event_name: &str) -> &[ListenerInfo] {
self.by_event.get(event_name).map_or(&[][..], Vec::as_slice)
}
#[must_use]
pub fn durable_job_infos(&self) -> Vec<crate::job::JobInfo> {
self.by_event
.values()
.flatten()
.filter(|listener| listener.mode == DispatchMode::Durable)
.map(|listener| crate::job::JobInfo {
name: listener
.job_name
.clone()
.expect("durable listener must carry a job_name"),
max_attempts: listener.max_attempts,
initial_backoff_ms: listener.initial_backoff_ms,
queue: "default".to_string(),
uniqueness: None,
concurrency: None,
version: 1,
handler: listener.handler,
})
.collect()
}
}
#[derive(Clone, Debug)]
pub struct RecordedEvent {
pub event_name: &'static str,
pub payload: Value,
}
#[derive(Default)]
pub struct EventRecorder {
events: Mutex<Vec<RecordedEvent>>,
}
impl EventRecorder {
fn record(&self, event_name: &'static str, payload: Value) {
self.events
.lock()
.expect("event recorder lock poisoned")
.push(RecordedEvent {
event_name,
payload,
});
}
#[must_use]
pub fn published<E: Event>(&self) -> Vec<E> {
self.events
.lock()
.expect("event recorder lock poisoned")
.iter()
.filter(|recorded| recorded.event_name == E::NAME)
.filter_map(|recorded| serde_json::from_value(recorded.payload.clone()).ok())
.collect()
}
#[must_use]
pub fn count<E: Event>(&self) -> usize {
self.events
.lock()
.expect("event recorder lock poisoned")
.iter()
.filter(|recorded| recorded.event_name == E::NAME)
.count()
}
#[must_use]
pub fn all(&self) -> Vec<RecordedEvent> {
self.events
.lock()
.expect("event recorder lock poisoned")
.clone()
}
}
#[derive(Clone)]
pub struct Events {
registry: Arc<EventRegistry>,
recorder: Option<Arc<EventRecorder>>,
state: AppState,
}
impl Events {
pub async fn publish<E: Event>(&self, event: E) -> AutumnResult<()> {
let payload = serialize_event(&event)?;
dispatch(
&self.registry,
self.recorder.as_deref(),
&self.state,
E::NAME,
payload,
)
.await
}
}
impl axum::extract::FromRequestParts<AppState> for Events {
type Rejection = AutumnError;
async fn from_request_parts(
_parts: &mut http::request::Parts,
state: &AppState,
) -> Result<Self, Self::Rejection> {
let registry = state
.extension::<EventRegistry>()
.unwrap_or_else(|| Arc::new(EventRegistry::default()));
let recorder = state.extension::<EventRecorder>();
Ok(Self {
registry,
recorder,
state: state.clone(),
})
}
}
fn serialize_event<E: Event>(event: &E) -> AutumnResult<Value> {
serde_json::to_value(event).map_err(|e| {
AutumnError::internal_server_error(std::io::Error::other(format!(
"event serialization failed: {e}"
)))
})
}
async fn dispatch(
registry: &EventRegistry,
recorder: Option<&EventRecorder>,
state: &AppState,
event_name: &'static str,
payload: Value,
) -> AutumnResult<()> {
if let Some(recorder) = recorder {
recorder.record(event_name, payload.clone());
}
let listeners = registry.listeners_for(event_name);
if listeners.is_empty() {
return Ok(());
}
let app_client = state.extension::<crate::job::JobClient>();
let mut durable_error = None;
for listener in listeners
.iter()
.filter(|listener| listener.mode == DispatchMode::Durable)
{
let job_name = listener
.job_name
.as_deref()
.expect("durable listener must carry a job_name");
let enqueued = if let Some(client) = &app_client {
client.enqueue_after_commit(job_name, payload.clone()).await
} else {
crate::job::enqueue_after_commit(job_name, payload.clone()).await
};
if let Err(error) = enqueued
&& durable_error.is_none()
{
durable_error = Some(error);
}
}
run_sync_listeners(state, listeners, &payload).await;
durable_error.map_or(Ok(()), Err)
}
async fn run_sync_listeners(state: &AppState, listeners: &[ListenerInfo], payload: &Value) {
use futures::FutureExt as _;
let runs = listeners
.iter()
.filter(|listener| listener.mode == DispatchMode::Sync)
.map(|listener| {
let state = state.clone();
let payload = payload.clone();
let run = listener.handler;
let name = listener.listener_name.clone();
async move {
match std::panic::AssertUnwindSafe(run(state, payload))
.catch_unwind()
.await
{
Ok(Ok(())) => {}
Ok(Err(error)) => {
tracing::error!(listener = %name, %error, "sync event listener failed");
}
Err(_panic) => {
tracing::error!(listener = %name, "sync event listener panicked");
}
}
}
});
futures::future::join_all(runs).await;
}
struct GlobalBus {
registry: Arc<EventRegistry>,
recorder: Option<Arc<EventRecorder>>,
state: AppState,
}
static GLOBAL_EVENT_BUS: RwLock<Option<Arc<GlobalBus>>> = RwLock::new(None);
fn global_bus() -> Option<Arc<GlobalBus>> {
GLOBAL_EVENT_BUS.read().ok().and_then(|guard| guard.clone())
}
pub(crate) fn init_global_event_bus(
registry: &EventRegistry,
state: &AppState,
recorder: Option<Arc<EventRecorder>>,
) {
let bus = Arc::new(GlobalBus {
registry: Arc::new(registry.clone()),
recorder,
state: state.clone(),
});
if let Ok(mut guard) = GLOBAL_EVENT_BUS.write() {
*guard = Some(bus);
}
}
pub fn clear_global_event_bus() {
if let Ok(mut guard) = GLOBAL_EVENT_BUS.write() {
*guard = None;
}
}
tokio::task_local! {
static CURRENT_EVENT_APP: AppState;
}
pub(crate) fn scope_event_app<F>(
state: AppState,
future: F,
) -> tokio::task::futures::TaskLocalFuture<AppState, F>
where
F: Future,
{
CURRENT_EVENT_APP.scope(state, future)
}
fn current_event_app() -> Option<AppState> {
CURRENT_EVENT_APP.try_with(AppState::clone).ok()
}
pub async fn publish<E: Event>(event: E) -> AutumnResult<()> {
let payload = serialize_event(&event)?;
if let Some(state) = current_event_app() {
let registry = state.extension::<EventRegistry>();
let empty = EventRegistry::default();
let registry_ref = registry.as_deref().unwrap_or(&empty);
let recorder = state.extension::<EventRecorder>();
return dispatch(registry_ref, recorder.as_deref(), &state, E::NAME, payload).await;
}
let Some(bus) = global_bus() else {
return Ok(());
};
dispatch(
&bus.registry,
bus.recorder.as_deref(),
&bus.state,
E::NAME,
payload,
)
.await
}
#[cfg(test)]
mod tests {
use super::*;
use serde::Deserialize;
use std::sync::atomic::{AtomicU32, Ordering};
#[derive(Serialize, Deserialize, Clone, Debug)]
struct Ping {
n: i64,
}
impl Event for Ping {
const NAME: &'static str = "Ping";
}
fn ok_handler() -> ListenerHandler {
|_state, _payload| Box::pin(async { Ok(()) })
}
fn sync_listener(name: &str, handler: ListenerHandler) -> ListenerInfo {
ListenerInfo {
event_name: Ping::NAME,
listener_name: name.to_string(),
mode: DispatchMode::Sync,
job_name: None,
max_attempts: 0,
initial_backoff_ms: 0,
handler,
}
}
fn durable_listener(name: &str) -> ListenerInfo {
ListenerInfo {
event_name: Ping::NAME,
listener_name: name.to_string(),
mode: DispatchMode::Durable,
job_name: Some(format!("__event_listener::{name}")),
max_attempts: 4,
initial_backoff_ms: 250,
handler: ok_handler(),
}
}
#[test]
fn registry_groups_by_event_name() {
let registry = EventRegistry::from_listeners(vec![
sync_listener("a", ok_handler()),
sync_listener("b", ok_handler()),
]);
assert_eq!(registry.listeners_for("Ping").len(), 2);
assert!(registry.listeners_for("Other").is_empty());
}
#[test]
fn durable_listeners_become_job_infos() {
let registry = EventRegistry::from_listeners(vec![
sync_listener("a", ok_handler()),
durable_listener("seed_workspace"),
]);
let jobs = registry.durable_job_infos();
assert_eq!(jobs.len(), 1, "only durable listeners become jobs");
assert_eq!(jobs[0].name, "__event_listener::seed_workspace");
assert_eq!(jobs[0].max_attempts, 4);
assert_eq!(jobs[0].initial_backoff_ms, 250);
}
#[tokio::test]
async fn sync_listeners_are_isolated_from_panics_and_errors() {
static RAN: AtomicU32 = AtomicU32::new(0);
RAN.store(0, Ordering::SeqCst);
let panicking: ListenerHandler = |_state, _payload| Box::pin(async { panic!("boom") });
let erroring: ListenerHandler = |_state, _payload| {
Box::pin(async {
Err(AutumnError::internal_server_error(std::io::Error::other(
"nope",
)))
})
};
let counting: ListenerHandler = |_state, _payload| {
Box::pin(async {
RAN.fetch_add(1, Ordering::SeqCst);
Ok(())
})
};
let registry = EventRegistry::from_listeners(vec![
sync_listener("panics", panicking),
sync_listener("errors", erroring),
sync_listener("counts", counting),
]);
let state = AppState::for_test();
let result = dispatch(
®istry,
None,
&state,
Ping::NAME,
serde_json::json!({"n": 1}),
)
.await;
assert!(result.is_ok(), "publish stays Ok despite listener failures");
assert_eq!(RAN.load(Ordering::SeqCst), 1, "surviving listener ran");
}
#[tokio::test]
async fn sync_listeners_run_even_when_a_durable_enqueue_fails() {
static RAN: AtomicU32 = AtomicU32::new(0);
crate::job::clear_global_job_client();
RAN.store(0, Ordering::SeqCst);
let counting: ListenerHandler = |_state, _payload| {
Box::pin(async {
RAN.fetch_add(1, Ordering::SeqCst);
Ok(())
})
};
let registry = EventRegistry::from_listeners(vec![
durable_listener("seed_workspace"),
sync_listener("counts", counting),
]);
let state = AppState::for_test();
let _ = dispatch(
®istry,
None,
&state,
Ping::NAME,
serde_json::json!({"n": 1}),
)
.await;
assert_eq!(RAN.load(Ordering::SeqCst), 1, "sync listener ran anyway");
}
#[tokio::test]
async fn missing_listener_is_a_noop() {
let registry = EventRegistry::default();
let state = AppState::for_test();
let result = dispatch(
®istry,
None,
&state,
Ping::NAME,
serde_json::json!({"n": 1}),
)
.await;
assert!(result.is_ok());
}
#[tokio::test]
async fn recorder_captures_published_events() {
let registry = EventRegistry::default();
let recorder = EventRecorder::default();
let state = AppState::for_test();
let payload = serialize_event(&Ping { n: 7 }).unwrap();
dispatch(®istry, Some(&recorder), &state, Ping::NAME, payload)
.await
.unwrap();
assert_eq!(recorder.count::<Ping>(), 1);
let published = recorder.published::<Ping>();
assert_eq!(published.len(), 1);
assert_eq!(published[0].n, 7);
}
}