use std::sync::Arc;
use aion_core::{ActivityEvent, ActivityEventKind, ProgressDetail};
use aion_store::{ActivityRecord, ActivityStreamKey, ObservabilityStore, StoreError};
use futures::stream::{self, BoxStream};
use tokio::sync::broadcast;
use crate::activity_bounds::{TranscriptBounds, bound_event};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct TranscriptStreamLagged {
pub skipped: u64,
}
impl std::fmt::Display for TranscriptStreamLagged {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
formatter,
"transcript stream lagged: {} events dropped",
self.skipped
)
}
}
impl std::error::Error for TranscriptStreamLagged {}
#[derive(Clone)]
pub struct ActivityEventPublisher {
store: Arc<dyn ObservabilityStore>,
live: broadcast::Sender<ActivityEvent>,
bounds: TranscriptBounds,
}
impl std::fmt::Debug for ActivityEventPublisher {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ActivityEventPublisher")
.field("live_receivers", &self.live.receiver_count())
.finish_non_exhaustive()
}
}
const MAX_SEQUENCE_CONFLICT_RETRIES: usize = 16;
impl ActivityEventPublisher {
#[must_use]
pub fn new(store: Arc<dyn ObservabilityStore>, capacity: std::num::NonZeroUsize) -> Self {
let (live, _receiver) = broadcast::channel(capacity.get());
Self {
store,
live,
bounds: TranscriptBounds::default(),
}
}
#[must_use]
pub(crate) fn with_bounds(mut self, bounds: TranscriptBounds) -> Self {
self.bounds = bounds;
self
}
pub async fn publish(&self, event: &ActivityEvent) -> Result<Option<u64>, StoreError> {
if event.ephemeral {
let mut ephemeral = event.clone();
ephemeral.store_seq = None;
let send_result = self.live.send(ephemeral);
drop(send_result);
return Ok(None);
}
let event = bound_event(event, self.bounds.max_event_bytes)?;
let key = ActivityStreamKey::of(&event);
let mut expected_seq = self.store.activity_head(&key).await?;
for _attempt in 0..MAX_SEQUENCE_CONFLICT_RETRIES {
if expected_seq > self.bounds.max_stream_events {
self.fan_out_live_only(&event);
return Ok(None);
}
if expected_seq == self.bounds.max_stream_events {
match self.append_cap_marker(&event, expected_seq).await {
Ok(()) => {
self.fan_out_live_only(&event);
return Ok(None);
}
Err(StoreError::SequenceConflict { found, .. }) => {
expected_seq = found;
continue;
}
Err(error) => return Err(error),
}
}
match self.store.append_activity_event(expected_seq, &event).await {
Ok(store_seq) => {
let mut persisted = event.clone();
persisted.store_seq = Some(store_seq);
let send_result = self.live.send(persisted);
drop(send_result);
return Ok(Some(store_seq));
}
Err(StoreError::SequenceConflict { found, .. }) => {
expected_seq = found;
}
Err(error) => return Err(error),
}
}
Err(StoreError::Backend(format!(
"observability append exceeded {MAX_SEQUENCE_CONFLICT_RETRIES} sequence-conflict retries for {key:?}"
)))
}
fn fan_out_live_only(&self, event: &ActivityEvent) {
let mut live_only = event.clone();
live_only.store_seq = None;
let send_result = self.live.send(live_only);
drop(send_result);
}
async fn append_cap_marker(
&self,
event: &ActivityEvent,
cap_seq: u64,
) -> Result<(), StoreError> {
let cap = self.bounds.max_stream_events;
let mut marker = event.clone();
marker.kind = ActivityEventKind::Progress {
detail: ProgressDetail::Note {
text: format!(
"transcript retention cap reached ({cap} events); further events are live-only and not persisted"
),
},
};
let store_seq = self.store.append_activity_event(cap_seq, &marker).await?;
marker.store_seq = Some(store_seq);
let send_result = self.live.send(marker);
drop(send_result);
Ok(())
}
pub async fn replay_from(
&self,
key: &ActivityStreamKey,
from_seq: u64,
) -> Result<Vec<ActivityRecord>, StoreError> {
self.store.read_activity_events_from(key, from_seq).await
}
pub async fn list_streams(
&self,
workflow_id: &aion_core::WorkflowId,
run_id: &aion_core::RunId,
) -> Result<Vec<aion_store::ActivityStreamSummary>, StoreError> {
self.store.list_activity_streams(workflow_id, run_id).await
}
#[must_use]
pub fn subscribe(
&self,
key: ActivityStreamKey,
after_seq: Option<u64>,
) -> BoxStream<'static, Result<ActivityEvent, TranscriptStreamLagged>> {
let receiver = self.live.subscribe();
Box::pin(stream::unfold(
(receiver, key, after_seq),
|(mut receiver, key, after_seq)| async move {
loop {
match receiver.recv().await {
Ok(event) => {
if ActivityStreamKey::of(&event) != key {
continue;
}
match (event.store_seq, after_seq) {
(Some(seq), Some(cursor)) if seq <= cursor => {}
_ => return Some((Ok(event), (receiver, key, after_seq))),
}
}
Err(broadcast::error::RecvError::Lagged(skipped)) => {
return Some((
Err(TranscriptStreamLagged { skipped }),
(receiver, key, after_seq),
));
}
Err(broadcast::error::RecvError::Closed) => return None,
}
}
},
))
}
}
#[cfg(test)]
#[path = "activity_publisher_tests.rs"]
mod tests;