use async_trait::async_trait;
use aion_core::{ActivityEvent, ActivityId, WorkflowId};
use crate::StoreError;
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub struct ActivityStreamKey {
pub workflow_id: WorkflowId,
pub activity_id: ActivityId,
pub attempt: u32,
}
impl ActivityStreamKey {
#[must_use]
pub const fn new(workflow_id: WorkflowId, activity_id: ActivityId, attempt: u32) -> Self {
Self {
workflow_id,
activity_id,
attempt,
}
}
#[must_use]
pub fn of(event: &ActivityEvent) -> Self {
Self {
workflow_id: event.workflow_id.clone(),
activity_id: event.activity_id.clone(),
attempt: event.attempt,
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct ActivityRecord {
pub store_seq: u64,
pub event: ActivityEvent,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ActivityStreamSummary {
pub key: ActivityStreamKey,
pub head: u64,
}
#[async_trait]
pub trait ObservabilityStore: Send + Sync + 'static {
async fn append_activity_event(
&self,
expected_seq: u64,
event: &ActivityEvent,
) -> Result<u64, StoreError>;
async fn activity_head(&self, key: &ActivityStreamKey) -> Result<u64, StoreError>;
async fn read_activity_events_from(
&self,
key: &ActivityStreamKey,
from_seq: u64,
) -> Result<Vec<ActivityRecord>, StoreError>;
async fn list_activity_streams(
&self,
workflow_id: &WorkflowId,
) -> Result<Vec<ActivityStreamSummary>, StoreError>;
}
#[derive(Debug, Default)]
pub struct InMemoryObservabilityStore {
streams:
std::sync::Mutex<std::collections::HashMap<ActivityStreamKeyBytes, Vec<ActivityRecord>>>,
}
type ActivityStreamKeyBytes = (uuid::Uuid, u64, u32);
fn key_bytes(key: &ActivityStreamKey) -> ActivityStreamKeyBytes {
(
key.workflow_id.as_uuid(),
key.activity_id.sequence_position(),
key.attempt,
)
}
fn stream_head(stream: &[ActivityRecord]) -> u64 {
u64::try_from(stream.len()).unwrap_or(u64::MAX)
}
#[async_trait]
impl ObservabilityStore for InMemoryObservabilityStore {
async fn append_activity_event(
&self,
expected_seq: u64,
event: &ActivityEvent,
) -> Result<u64, StoreError> {
let key = ActivityStreamKey::of(event);
let mut streams = self.streams.lock().map_err(|error| {
StoreError::Backend(format!("observability mutex poisoned: {error}"))
})?;
let stream = streams.entry(key_bytes(&key)).or_default();
let head = stream_head(stream);
if head != expected_seq {
return Err(StoreError::SequenceConflict {
expected: expected_seq,
found: head,
});
}
let mut event = event.clone();
event.store_seq = Some(head);
stream.push(ActivityRecord {
store_seq: head,
event,
});
Ok(head)
}
async fn activity_head(&self, key: &ActivityStreamKey) -> Result<u64, StoreError> {
let streams = self.streams.lock().map_err(|error| {
StoreError::Backend(format!("observability mutex poisoned: {error}"))
})?;
Ok(streams
.get(&key_bytes(key))
.map_or(0, |stream| stream_head(stream)))
}
async fn read_activity_events_from(
&self,
key: &ActivityStreamKey,
from_seq: u64,
) -> Result<Vec<ActivityRecord>, StoreError> {
let streams = self.streams.lock().map_err(|error| {
StoreError::Backend(format!("observability mutex poisoned: {error}"))
})?;
Ok(streams
.get(&key_bytes(key))
.map_or_else(Vec::new, |stream| {
stream
.iter()
.filter(|record| record.store_seq >= from_seq)
.cloned()
.collect()
}))
}
async fn list_activity_streams(
&self,
workflow_id: &WorkflowId,
) -> Result<Vec<ActivityStreamSummary>, StoreError> {
let streams = self.streams.lock().map_err(|error| {
StoreError::Backend(format!("observability mutex poisoned: {error}"))
})?;
let mut summaries: Vec<ActivityStreamSummary> = streams
.iter()
.filter(|((workflow, _activity, _attempt), _records)| {
*workflow == workflow_id.as_uuid()
})
.map(
|(&(workflow, activity_seq, attempt), records)| ActivityStreamSummary {
key: ActivityStreamKey::new(
WorkflowId::new(workflow),
ActivityId::from_sequence_position(activity_seq),
attempt,
),
head: stream_head(records),
},
)
.collect();
summaries.sort_by_key(|summary| {
(
summary.key.activity_id.sequence_position(),
summary.key.attempt,
)
});
Ok(summaries)
}
}
#[cfg(test)]
mod tests {
use super::*;
use aion_core::{ActivityEventKind, MessageRole};
use chrono::Utc;
use uuid::Uuid;
fn event(attempt: u32, worker_seq: u64, text: &str) -> ActivityEvent {
ActivityEvent {
workflow_id: WorkflowId::new(Uuid::from_u128(1)),
activity_id: ActivityId::from_sequence_position(3),
attempt,
agent_id: Uuid::from_u128(9),
agent_role: "orchestrator".to_owned(),
emitted_at: Utc::now(),
worker_seq,
store_seq: None,
ephemeral: false,
kind: ActivityEventKind::Message {
role: MessageRole::Assistant,
text: text.to_owned(),
},
}
}
#[tokio::test]
async fn append_assigns_contiguous_store_seq_from_zero() -> Result<(), StoreError> {
let store = InMemoryObservabilityStore::default();
let key = ActivityStreamKey::new(
WorkflowId::new(Uuid::from_u128(1)),
ActivityId::from_sequence_position(3),
0,
);
assert_eq!(store.activity_head(&key).await?, 0);
assert_eq!(store.append_activity_event(0, &event(0, 1, "a")).await?, 0);
assert_eq!(store.append_activity_event(1, &event(0, 2, "b")).await?, 1);
assert_eq!(store.activity_head(&key).await?, 2);
let records = store.read_activity_events_from(&key, 0).await?;
assert_eq!(records.len(), 2);
assert_eq!(records[0].store_seq, 0);
assert_eq!(records[0].event.store_seq, Some(0));
assert_eq!(records[1].store_seq, 1);
Ok(())
}
#[tokio::test]
async fn stale_expected_seq_conflicts_and_writes_nothing() -> Result<(), StoreError> {
let store = InMemoryObservabilityStore::default();
store.append_activity_event(0, &event(0, 1, "a")).await?;
let conflict = store.append_activity_event(0, &event(0, 2, "dup")).await;
assert_eq!(
conflict,
Err(StoreError::SequenceConflict {
expected: 0,
found: 1
})
);
let key = ActivityStreamKey::of(&event(0, 0, ""));
assert_eq!(store.read_activity_events_from(&key, 0).await?.len(), 1);
Ok(())
}
#[tokio::test]
async fn attempts_are_disjoint_streams() -> Result<(), StoreError> {
let store = InMemoryObservabilityStore::default();
store
.append_activity_event(0, &event(0, 1, "attempt-0"))
.await?;
store
.append_activity_event(0, &event(1, 1, "attempt-1"))
.await?;
let key0 = ActivityStreamKey::new(
WorkflowId::new(Uuid::from_u128(1)),
ActivityId::from_sequence_position(3),
0,
);
let key1 = ActivityStreamKey::new(
WorkflowId::new(Uuid::from_u128(1)),
ActivityId::from_sequence_position(3),
1,
);
assert_eq!(store.activity_head(&key0).await?, 1);
assert_eq!(store.activity_head(&key1).await?, 1);
Ok(())
}
#[tokio::test]
async fn list_activity_streams_orders_by_activity_then_attempt() -> Result<(), StoreError> {
let store = InMemoryObservabilityStore::default();
let event_for = |activity_seq: u64, attempt: u32, workflow: u128| {
let mut event = event(attempt, 1, "x");
event.workflow_id = WorkflowId::new(Uuid::from_u128(workflow));
event.activity_id = ActivityId::from_sequence_position(activity_seq);
event
};
store.append_activity_event(0, &event_for(5, 0, 1)).await?;
store.append_activity_event(0, &event_for(3, 1, 1)).await?;
store.append_activity_event(0, &event_for(3, 0, 1)).await?;
store.append_activity_event(1, &event_for(3, 0, 1)).await?;
store.append_activity_event(0, &event_for(3, 0, 2)).await?;
let summaries = store
.list_activity_streams(&WorkflowId::new(Uuid::from_u128(1)))
.await?;
let listed: Vec<(u64, u32, u64)> = summaries
.iter()
.map(|summary| {
(
summary.key.activity_id.sequence_position(),
summary.key.attempt,
summary.head,
)
})
.collect();
assert_eq!(listed, vec![(3, 0, 2), (3, 1, 1), (5, 0, 1)]);
Ok(())
}
#[tokio::test]
async fn list_activity_streams_is_empty_for_unknown_workflow() -> Result<(), StoreError> {
let store = InMemoryObservabilityStore::default();
store.append_activity_event(0, &event(0, 1, "a")).await?;
let summaries = store
.list_activity_streams(&WorkflowId::new(Uuid::from_u128(99)))
.await?;
assert!(summaries.is_empty(), "an unwritten workflow lists empty");
Ok(())
}
#[tokio::test]
async fn read_from_resumes_by_store_seq() -> Result<(), StoreError> {
let store = InMemoryObservabilityStore::default();
for seq in 0..5u64 {
store
.append_activity_event(seq, &event(0, seq, "x"))
.await?;
}
let key = ActivityStreamKey::of(&event(0, 0, ""));
let tail = store.read_activity_events_from(&key, 3).await?;
assert_eq!(tail.len(), 2);
assert_eq!(tail[0].store_seq, 3);
assert_eq!(tail[1].store_seq, 4);
Ok(())
}
}