use std::future::Future;
use std::marker::PhantomData;
use std::pin::Pin;
use crate::entity::Entity;
use crate::outbox::OutboxPublisherConfig;
use crate::queued_repo::{GetAllWithOpts, GetWithOpts, ReadOpts, UnlockableRepository};
use crate::repository::{
CommitBatch, GetStream, RepositoryError, SnapshotWrite, StreamIdentity, StreamWrite,
TransactionalCommit,
};
use crate::snapshot::SnapshotRecord;
use super::{hydrate, Aggregate};
fn stream_identity_for<A: Aggregate>(
aggregate_id: &str,
) -> Result<StreamIdentity, RepositoryError> {
StreamIdentity::new(A::aggregate_type(), aggregate_id)
}
pub trait AggregateBuilder: Sized {
fn aggregate<A: Aggregate>(self) -> AggregateRepository<Self, A> {
AggregateRepository::new(self)
}
}
impl<T> AggregateBuilder for T {}
pub(crate) struct SnapshotPolicy<R, A> {
frequency: u64,
record: fn(&A, u64) -> Result<Option<SnapshotRecord>, RepositoryError>,
hydrate: HydrateFn<R, A>,
}
type HydrateFn<R, A> =
for<'a> fn(
&'a R,
&'a StreamIdentity,
Entity,
) -> Pin<Box<dyn Future<Output = Result<A, RepositoryError>> + Send + 'a>>;
impl<R, A> SnapshotPolicy<R, A> {
pub(crate) fn new(
frequency: u64,
record: fn(&A, u64) -> Result<Option<SnapshotRecord>, RepositoryError>,
hydrate: HydrateFn<R, A>,
) -> Self {
Self {
frequency,
record,
hydrate,
}
}
}
pub struct AggregateRepository<R, A> {
repo: R,
snapshot: Option<SnapshotPolicy<R, A>>,
outbox_publisher: Option<OutboxPublisherConfig>,
_marker: PhantomData<A>,
}
impl<R, A> AggregateRepository<R, A> {
pub fn new(repo: R) -> Self {
Self {
repo,
snapshot: None,
outbox_publisher: None,
_marker: PhantomData,
}
}
pub fn repo(&self) -> &R {
&self.repo
}
pub fn repo_mut(&mut self) -> &mut R {
&mut self.repo
}
pub(crate) fn set_snapshot_policy(&mut self, policy: SnapshotPolicy<R, A>) {
self.snapshot = Some(policy);
}
pub(crate) fn set_outbox_publisher(&mut self, config: OutboxPublisherConfig) {
self.outbox_publisher = Some(config);
}
pub(crate) fn outbox_publisher(&self) -> Option<&OutboxPublisherConfig> {
self.outbox_publisher.as_ref()
}
}
impl<R, A> AggregateRepository<R, A>
where
A: Aggregate + Send,
{
async fn hydrate_entity(
&self,
identity: &StreamIdentity,
entity: Entity,
) -> Result<A, RepositoryError> {
match &self.snapshot {
Some(policy) => (policy.hydrate)(&self.repo, identity, entity).await,
None => hydrate::<A>(entity),
}
}
fn snapshot_writes(
&self,
aggregate: &A,
) -> Result<(Vec<SnapshotWrite>, Option<u64>), RepositoryError> {
let Some(policy) = &self.snapshot else {
return Ok((Vec::new(), None));
};
let Some(record) = (policy.record)(aggregate, policy.frequency)? else {
return Ok((Vec::new(), None));
};
let version = record.version;
let identity = stream_identity_for::<A>(aggregate.entity().id())?;
Ok((
vec![SnapshotWrite::Save { identity, record }],
Some(version),
))
}
pub(crate) fn snapshot_writes_for(
&self,
aggregate: &A,
) -> Result<(Vec<SnapshotWrite>, Option<u64>), RepositoryError> {
self.snapshot_writes(aggregate)
}
}
impl<R, A> AggregateRepository<R, A>
where
R: GetStream,
A: Aggregate + Send,
{
pub async fn get(&self, id: &str) -> Result<Option<A>, RepositoryError> {
let identity = stream_identity_for::<A>(id)?;
let entity = self.repo.get_stream(&identity).await?;
let Some(entity) = entity else {
return Ok(None);
};
Ok(Some(self.hydrate_entity(&identity, entity).await?))
}
pub async fn get_all(&self, ids: &[&str]) -> Result<Vec<A>, RepositoryError> {
let identities = ids
.iter()
.map(|id| stream_identity_for::<A>(id))
.collect::<Result<Vec<_>, _>>()?;
let entities = self.repo.get_streams(&identities).await?;
self.hydrate_entities(entities).await
}
}
impl<R, A> AggregateRepository<R, A>
where
A: Aggregate + Send,
{
async fn hydrate_entities(&self, entities: Vec<Entity>) -> Result<Vec<A>, RepositoryError> {
let mut aggregates = Vec::with_capacity(entities.len());
for entity in entities {
let identity = stream_identity_for::<A>(entity.id())?;
aggregates.push(self.hydrate_entity(&identity, entity).await?);
}
Ok(aggregates)
}
}
impl<R, A> AggregateRepository<R, A>
where
R: TransactionalCommit,
A: Aggregate + Send,
{
pub async fn commit(&self, aggregate: &mut A) -> Result<(), RepositoryError> {
let (snapshots, snapshot_version) = self.snapshot_writes(aggregate)?;
let identity = stream_identity_for::<A>(aggregate.entity().id())?;
let stream = StreamWrite::new(identity, aggregate.entity_mut());
let mut batch = CommitBatch::new(vec![stream]);
batch.snapshots = snapshots;
self.repo.commit_batch(batch).await?;
if let Some(version) = snapshot_version {
aggregate.entity_mut().set_snapshot_version(version);
}
Ok(())
}
pub async fn commit_all(&self, aggregates: &mut [&mut A]) -> Result<(), RepositoryError> {
let mut snapshots = Vec::new();
let mut snapshot_versions = Vec::with_capacity(aggregates.len());
for aggregate in aggregates.iter() {
let (mut writes, version) = self.snapshot_writes(aggregate)?;
snapshots.append(&mut writes);
snapshot_versions.push(version);
}
let mut streams = Vec::with_capacity(aggregates.len());
for aggregate in aggregates.iter_mut() {
let identity = stream_identity_for::<A>((*aggregate).entity().id())?;
streams.push(StreamWrite::new(identity, (*aggregate).entity_mut()));
}
let mut batch = CommitBatch::new(streams);
batch.snapshots = snapshots;
self.repo.commit_batch(batch).await?;
for (aggregate, version) in aggregates.iter_mut().zip(snapshot_versions) {
if let Some(version) = version {
(*aggregate).entity_mut().set_snapshot_version(version);
}
}
Ok(())
}
pub async fn commit_entities(
&self,
streams: Vec<(StreamIdentity, &mut Entity)>,
) -> Result<(), RepositoryError> {
let streams = streams
.into_iter()
.map(|(identity, entity)| StreamWrite::new(identity, entity))
.collect();
self.repo.commit_batch(CommitBatch::new(streams)).await
}
}
impl<R, A> AggregateRepository<R, A>
where
R: GetWithOpts,
A: Aggregate + Send,
{
pub async fn get_with(&self, id: &str, opts: ReadOpts) -> Result<Option<A>, RepositoryError> {
let identity = stream_identity_for::<A>(id)?;
let Some(entity) = self.repo.get_stream_with(&identity, opts).await? else {
return Ok(None);
};
Ok(Some(self.hydrate_entity(&identity, entity).await?))
}
pub async fn peek(&self, id: &str) -> Result<Option<A>, RepositoryError> {
self.get_with(id, ReadOpts::no_lock()).await
}
}
impl<R, A> AggregateRepository<R, A>
where
R: GetAllWithOpts,
A: Aggregate + Send,
{
pub async fn get_all_with(
&self,
ids: &[&str],
opts: ReadOpts,
) -> Result<Vec<A>, RepositoryError> {
let identities = ids
.iter()
.map(|id| stream_identity_for::<A>(id))
.collect::<Result<Vec<_>, _>>()?;
let entities = self.repo.get_streams_with(&identities, opts).await?;
self.hydrate_entities(entities).await
}
pub async fn peek_all(&self, ids: &[&str]) -> Result<Vec<A>, RepositoryError> {
self.get_all_with(ids, ReadOpts::no_lock()).await
}
}
impl<R, A> AggregateRepository<R, A>
where
R: UnlockableRepository,
A: Aggregate,
{
pub async fn abort(&self, aggregate: &A) -> Result<(), RepositoryError> {
let identity = stream_identity_for::<A>(aggregate.entity().id())?;
self.repo.abort(&identity).await
}
pub async fn unlock(&self, id: &str) -> Result<(), RepositoryError> {
let identity = stream_identity_for::<A>(id)?;
self.repo.unlock(&identity).await
}
}