use std::sync::Arc;
use std::time::Duration;
use aion_core::{ActivityEvent, ActivityEventKind, ProgressDetail};
use aion_store::{ActivityRecord, ActivityStreamKey, ObservabilityStore, StoreError};
use futures::stream::{self, BoxStream};
use tokio::sync::{broadcast, mpsc};
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,
batch: TranscriptBatchPolicy,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct TranscriptBatchPolicy {
pub max_batch_events: std::num::NonZeroUsize,
pub max_hold: Duration,
}
impl TranscriptBatchPolicy {
pub const UNBATCHED: Self = Self {
max_batch_events: std::num::NonZeroUsize::MIN,
max_hold: Duration::ZERO,
};
}
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,
batch: TranscriptBatchPolicy,
) -> Self {
let (live, _receiver) = broadcast::channel(capacity.get());
Self {
store,
live,
bounds: TranscriptBounds::default(),
batch,
}
}
#[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> {
let assigned = self.publish_all(std::slice::from_ref(event)).await?;
Ok(assigned.into_iter().next().flatten())
}
pub async fn publish_all(
&self,
events: &[ActivityEvent],
) -> Result<Vec<Option<u64>>, StoreError> {
let (outcomes, first_error) = self.publish_all_outcomes(events).await;
match first_error {
Some(error) => Err(error),
None => Ok(outcomes
.into_iter()
.map(|outcome| match outcome {
EventOutcome::Persisted(store_seq) => Some(store_seq),
EventOutcome::NotPersisted | EventOutcome::Refused => None,
})
.collect()),
}
}
async fn publish_all_outcomes(
&self,
events: &[ActivityEvent],
) -> (Vec<EventOutcome>, Option<StoreError>) {
let mut outcomes: Vec<EventOutcome> = vec![EventOutcome::NotPersisted; events.len()];
let mut first_error: Option<StoreError> = None;
let mut groups: Vec<(ActivityStreamKey, Vec<(usize, ActivityEvent)>)> = Vec::new();
for (index, event) in events.iter().enumerate() {
if event.ephemeral {
let mut ephemeral = event.clone();
ephemeral.store_seq = None;
let send_result = self.live.send(ephemeral);
drop(send_result);
continue;
}
let bounded = match bound_event(event, self.bounds.max_event_bytes) {
Ok(bounded) => bounded,
Err(error) => {
outcomes[index] = EventOutcome::Refused;
if first_error.is_none() {
first_error = Some(error);
}
continue;
}
};
let key = ActivityStreamKey::of(&bounded);
match groups.iter_mut().find(|(existing, _)| *existing == key) {
Some((_, items)) => items.push((index, bounded)),
None => groups.push((key, vec![(index, bounded)])),
}
}
for (key, items) in groups {
if let Err(error) = self.persist_group(&key, &items, &mut outcomes).await {
if first_error.is_none() {
first_error = Some(error);
}
}
}
(outcomes, first_error)
}
async fn persist_group(
&self,
key: &ActivityStreamKey,
items: &[(usize, ActivityEvent)],
outcomes: &mut [EventOutcome],
) -> Result<(), StoreError> {
let mut expected_seq = match self.store.activity_head(key).await {
Ok(head) => head,
Err(error) => {
mark_refused(items, outcomes, 0);
return Err(error);
}
};
let mut committed = 0usize;
let mut conflicts = 0usize;
while committed < items.len() {
if expected_seq > self.bounds.max_stream_events {
for (_, event) in &items[committed..] {
self.fan_out_live_only(event);
}
return Ok(());
}
if expected_seq == self.bounds.max_stream_events {
match self
.append_cap_marker(&items[committed].1, expected_seq)
.await
{
Ok(()) => {
for (_, event) in &items[committed..] {
self.fan_out_live_only(event);
}
return Ok(());
}
Err(StoreError::SequenceConflict { found, .. }) => {
expected_seq = found;
conflicts += 1;
}
Err(error) => {
mark_refused(items, outcomes, committed);
return Err(error);
}
}
if conflicts >= MAX_SEQUENCE_CONFLICT_RETRIES {
break;
}
continue;
}
let room = usize::try_from(self.bounds.max_stream_events - expected_seq)
.unwrap_or(usize::MAX)
.min(items.len() - committed);
let batch: Vec<ActivityEvent> = items[committed..committed + room]
.iter()
.map(|(_, event)| event.clone())
.collect();
match self
.store
.append_activity_events(expected_seq, &batch)
.await
{
Ok(_new_head) => {
for (offset, (index, event)) in
items[committed..committed + room].iter().enumerate()
{
let store_seq =
expected_seq.saturating_add(u64::try_from(offset).unwrap_or(u64::MAX));
outcomes[*index] = EventOutcome::Persisted(store_seq);
let mut persisted = event.clone();
persisted.store_seq = Some(store_seq);
let send_result = self.live.send(persisted);
drop(send_result);
}
committed += room;
expected_seq =
expected_seq.saturating_add(u64::try_from(room).unwrap_or(u64::MAX));
}
Err(StoreError::SequenceConflict { found, .. }) => {
expected_seq = found;
conflicts += 1;
if conflicts >= MAX_SEQUENCE_CONFLICT_RETRIES {
break;
}
}
Err(error) => {
mark_refused(items, outcomes, committed);
return Err(error);
}
}
}
if committed < items.len() {
mark_refused(items, outcomes, committed);
return Err(StoreError::Backend(format!(
"observability append exceeded {MAX_SEQUENCE_CONFLICT_RETRIES} sequence-conflict retries for {key:?}"
)));
}
Ok(())
}
pub(crate) async fn drain<R: TranscriptEventReceiver>(
&self,
receiver: &mut R,
operation: &'static str,
) -> u64 {
let limit = self.batch.max_batch_events.get();
let mut buffer: Vec<ActivityEvent> = Vec::with_capacity(limit);
let mut dropped: u64 = 0;
loop {
let mut closed = receiver.recv_many(&mut buffer, limit).await == 0;
if !closed && buffer.len() < limit && !self.batch.max_hold.is_zero() {
let deadline = tokio::time::Instant::now() + self.batch.max_hold;
while buffer.len() < limit {
let remaining = limit - buffer.len();
match tokio::time::timeout_at(
deadline,
receiver.recv_many(&mut buffer, remaining),
)
.await
{
Ok(0) => {
closed = true;
break;
}
Ok(_) => {}
Err(_elapsed) => break,
}
}
}
if !buffer.is_empty() {
let (outcomes, error) = self.publish_all_outcomes(&buffer).await;
if let Some(error) = error {
let refused = outcomes
.iter()
.filter(|outcome| matches!(outcome, EventOutcome::Refused))
.count();
dropped = dropped.saturating_add(u64::try_from(refused).unwrap_or(u64::MAX));
let first = buffer.first();
tracing::warn!(
%error,
operation,
batch_events = buffer.len(),
refused_events = refused,
workflow_id = ?first.map(|event| event.workflow_id.to_string()),
activity_id = ?first.map(|event| event.activity_id.to_string()),
attempt = ?first.map(|event| event.attempt),
"transcript drain: the sequencer refused part of a batch; the producer \
is unaffected and the refused events are not retained"
);
}
buffer.clear();
}
if closed {
return dropped;
}
}
}
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,
}
}
},
))
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum EventOutcome {
NotPersisted,
Persisted(u64),
Refused,
}
fn mark_refused(items: &[(usize, ActivityEvent)], outcomes: &mut [EventOutcome], from: usize) {
for (index, _) in &items[from..] {
outcomes[*index] = EventOutcome::Refused;
}
}
#[async_trait::async_trait]
pub(crate) trait TranscriptEventReceiver: Send {
async fn recv_many(&mut self, buffer: &mut Vec<ActivityEvent>, limit: usize) -> usize;
}
#[async_trait::async_trait]
impl TranscriptEventReceiver for mpsc::Receiver<ActivityEvent> {
async fn recv_many(&mut self, buffer: &mut Vec<ActivityEvent>, limit: usize) -> usize {
Self::recv_many(self, buffer, limit).await
}
}
#[async_trait::async_trait]
impl TranscriptEventReceiver for mpsc::UnboundedReceiver<ActivityEvent> {
async fn recv_many(&mut self, buffer: &mut Vec<ActivityEvent>, limit: usize) -> usize {
Self::recv_many(self, buffer, limit).await
}
}
#[cfg(test)]
#[path = "activity_publisher_tests.rs"]
mod tests;
#[cfg(test)]
#[path = "activity_publisher_batching_tests.rs"]
mod batching_tests;