use core::num::NonZeroU32;
use mnesis::{Aggregate, AggregateRoot, DomainEvent, EventOf, Events, KernelError, Version};
use crate::repository::{ReplayFrom, Repository};
use crate::state;
pub struct Snapshotting<R, SS, T> {
inner: R,
snapshot_store: SS,
trigger: T,
schema_version: NonZeroU32,
snapshot_on_read: bool,
}
impl<R, SS, T> Snapshotting<R, SS, T> {
pub const fn new(
inner: R,
snapshot_store: SS,
trigger: T,
schema_version: NonZeroU32,
snapshot_on_read: bool,
) -> Self {
Self {
inner,
snapshot_store,
trigger,
schema_version,
snapshot_on_read,
}
}
}
impl<A, R, SS, T> Repository<A> for Snapshotting<R, SS, T>
where
A: Aggregate,
R: Repository<A> + ReplayFrom<A, Error = <R as Repository<A>>::Error>,
<R as Repository<A>>::Error: From<KernelError>,
SS: state::SnapshotStore<A::State, Version>,
T: state::PersistTrigger,
EventOf<A>: DomainEvent,
{
type Error = <R as Repository<A>>::Error;
type Position = <R as Repository<A>>::Position;
async fn load(&self, id: A::Id) -> Result<AggregateRoot<A>, Self::Error> {
if let Some((root, from)) = self.try_load_from_snapshot::<A>(&id).await {
return self.inner.replay_from(root, from).await;
}
let root = self.inner.load(id).await?;
if let (true, Some(version)) = (self.snapshot_on_read, root.version()) {
self.try_save_snapshot::<A>(&root, version).await;
}
Ok(root)
}
async fn save<const N: usize>(
&self,
aggregate: &mut AggregateRoot<A>,
events: &Events<EventOf<A>, N>,
) -> Result<Self::Position, Self::Error> {
let old_version = aggregate.version();
let position = self.inner.save(aggregate, events).await?;
let Some(new_version) = aggregate.version() else {
return Ok(position);
};
if self.trigger.should_persist(
old_version,
new_version,
events.iter().map(DomainEvent::name),
) {
self.try_save_snapshot::<A>(aggregate, new_version).await;
}
Ok(position)
}
}
impl<R, SS, T> Snapshotting<R, SS, T>
where
R: Send + Sync,
SS: Send + Sync,
T: Send + Sync,
{
#[cfg_attr(
feature = "tracing",
tracing::instrument(
name = "mnesis.snapshot.hydrate",
level = "debug",
skip_all,
fields(stream = %id, hit = tracing::field::Empty)
)
)]
async fn try_load_from_snapshot<A>(&self, id: &A::Id) -> Option<(AggregateRoot<A>, Version)>
where
A: Aggregate,
SS: state::SnapshotStore<A::State, Version>,
{
let hydrated = match self.snapshot_store.hydrate(id, self.schema_version).await {
Ok(hydrated) => hydrated,
Err(_snapshot_read_failed) => {
#[cfg(feature = "tracing")]
tracing::Span::current().record("hit", "error");
return None;
}
};
let (version, typed_state) = match hydrated {
state::Hydrated::Found { position, state } => {
#[cfg(feature = "tracing")]
tracing::Span::current().record("hit", "found");
(position, state)
}
state::Hydrated::Stale { .. } => {
#[cfg(feature = "tracing")]
tracing::Span::current().record("hit", "stale");
return None;
}
state::Hydrated::Absent => {
#[cfg(feature = "tracing")]
tracing::Span::current().record("hit", "absent");
return None;
}
};
let root = AggregateRoot::<A>::restore(id.clone(), typed_state, version);
let next = version.next()?;
Some((root, next))
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(
name = "mnesis.snapshot.commit",
level = "debug",
skip_all,
fields(stream = %aggregate.id(), version = %version)
)
)]
async fn try_save_snapshot<A>(&self, aggregate: &AggregateRoot<A>, version: Version)
where
A: Aggregate,
SS: state::SnapshotStore<A::State, Version>,
{
let _ = self
.snapshot_store
.commit(
aggregate.id(),
self.schema_version,
version,
aggregate.state(),
)
.await;
}
}