use crate::error::KernelError;
use crate::event::DomainEvent;
use crate::events::Events;
use crate::id::Id;
use crate::version::Version;
use core::error::Error;
use core::fmt;
use core::fmt::Debug;
use core::mem;
use core::num::NonZeroUsize;
pub trait AggregateState: Send + Sync + Debug + 'static {
type Event: DomainEvent;
fn initial() -> Self;
#[must_use]
fn apply(self, event: &Self::Event) -> Self;
}
pub trait Aggregate: Sized {
type State: AggregateState;
type Error: Error + Send + Sync + Debug + 'static;
type Id: Id;
const MAX_REHYDRATION_EVENTS: NonZeroUsize = DEFAULT_MAX_REHYDRATION_EVENTS;
}
pub trait Handle<C, const N: usize = 0>: Aggregate {
fn handle(state: &Self::State, cmd: C) -> Result<Events<EventOf<Self>, N>, Self::Error>;
}
pub type EventOf<A> = <<A as Aggregate>::State as AggregateState>::Event;
#[allow(clippy::unwrap_used, reason = "1_000_000 is non-zero by inspection")]
pub const DEFAULT_MAX_REHYDRATION_EVENTS: NonZeroUsize = NonZeroUsize::new(1_000_000).unwrap();
pub struct AggregateRoot<A: Aggregate> {
id: A::Id,
state: A::State,
version: Option<Version>,
}
impl<A: Aggregate> fmt::Debug for AggregateRoot<A> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("AggregateRoot")
.field("id", &self.id)
.field("version", &self.version)
.finish_non_exhaustive()
}
}
impl<A: Aggregate> AggregateRoot<A> {
pub fn new(id: A::Id) -> Self {
Self {
id,
state: A::State::initial(),
version: None,
}
}
#[must_use]
pub const fn restore(id: A::Id, state: A::State, version: Version) -> Self {
Self {
id,
state,
version: Some(version),
}
}
#[must_use]
pub const fn id(&self) -> &A::Id {
&self.id
}
#[must_use]
pub const fn state(&self) -> &A::State {
&self.state
}
#[must_use]
pub const fn version(&self) -> Option<Version> {
self.version
}
pub fn handle<C, const N: usize>(&self, cmd: C) -> Result<Events<EventOf<A>, N>, A::Error>
where
A: Handle<C, N>,
{
A::handle(self.state(), cmd)
}
#[allow(
clippy::expect_used,
reason = "u64::try_from(usize) cannot fail on supported platforms (max 64-bit)"
)]
pub fn replay(&mut self, version: Version, event: &EventOf<A>) -> Result<(), KernelError> {
let expected = match self.version {
None => Version::INITIAL,
Some(v) => v.next().ok_or(KernelError::VersionOverflow)?,
};
if version != expected {
return Err(KernelError::VersionMismatch {
expected,
actual: version,
});
}
if version.as_u64()
> u64::try_from(A::MAX_REHYDRATION_EVENTS.get())
.expect("MAX_REHYDRATION_EVENTS exceeds u64 on this platform")
{
return Err(KernelError::RehydrationLimitExceeded {
max: A::MAX_REHYDRATION_EVENTS.get(),
});
}
let taken = mem::replace(&mut self.state, A::State::initial());
self.state = taken.apply(event);
self.version = Some(version);
Ok(())
}
pub fn commit_persisted<const N: usize>(
&mut self,
version: Version,
events: &Events<EventOf<A>, N>,
) {
self.advance_version(version);
self.apply_events(events);
}
const fn advance_version(&mut self, new_version: Version) {
self.version = Some(new_version);
}
pub(crate) fn apply_events<const N: usize>(&mut self, events: &Events<EventOf<A>, N>) {
for event in events {
self.apply_event(event);
}
}
fn apply_event(&mut self, event: &EventOf<A>) {
let taken = mem::replace(&mut self.state, A::State::initial());
self.state = taken.apply(event);
}
}
#[cfg(test)]
#[allow(
clippy::expect_used,
clippy::panic,
reason = "test code: panic-safety test deliberately panics in apply()"
)]
mod purist_dispatch_tests {
use super::{Aggregate, AggregateRoot, AggregateState, Handle};
use crate::event::DomainEvent;
use crate::events;
use crate::events::Events;
use crate::message::Message;
use crate::version::Version;
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
struct CtrId([u8; 8]);
impl CtrId {
fn new(n: u64) -> Self {
Self(n.to_le_bytes())
}
}
impl std::fmt::Display for CtrId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", u64::from_le_bytes(self.0))
}
}
impl AsRef<[u8]> for CtrId {
fn as_ref(&self) -> &[u8] {
&self.0
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum CtrEvent {
Added(u64),
}
impl Message for CtrEvent {}
impl DomainEvent for CtrEvent {
fn name(&self) -> &'static str {
match self {
Self::Added(_) => "Added",
}
}
}
#[derive(Debug)]
struct CtrState {
total: u64,
}
impl AggregateState for CtrState {
type Event = CtrEvent;
fn initial() -> Self {
Self { total: 0 }
}
fn apply(mut self, event: &CtrEvent) -> Self {
match event {
CtrEvent::Added(n) => self.total = self.total.wrapping_add(*n),
}
self
}
}
#[derive(Debug, thiserror::Error, PartialEq)]
#[error("counter error")]
struct CtrError;
struct Counter;
impl Aggregate for Counter {
type State = CtrState;
type Error = CtrError;
type Id = CtrId;
}
struct Add(u64);
impl Handle<Add> for Counter {
fn handle(state: &CtrState, cmd: Add) -> Result<Events<CtrEvent>, CtrError> {
if cmd.0 == 0 {
return Err(CtrError);
}
let _ = state.total;
Ok(events![CtrEvent::Added(cmd.0)])
}
}
#[test]
fn dispatches_to_handle_on_the_marker() {
let root = AggregateRoot::<Counter>::new(CtrId::new(1));
let decided = root.handle(Add(5)).expect("ok");
assert_eq!(
decided.into_iter().collect::<Vec<_>>(),
vec![CtrEvent::Added(5)]
);
}
#[test]
fn surfaces_domain_error_from_handle() {
assert_eq!(
AggregateRoot::<Counter>::new(CtrId::new(1)).handle(Add(0)),
Err(CtrError)
);
}
#[test]
fn commit_persisted_advances_version_and_folds_state_atomically() {
let v2 = Version::new(2).expect("nonzero");
let persisted: Events<CtrEvent, 1> = events![CtrEvent::Added(10), CtrEvent::Added(5)];
let mut committed = AggregateRoot::<Counter>::new(CtrId::new(42));
committed.commit_persisted(v2, &persisted);
assert_eq!(committed.version(), Some(v2));
assert_eq!(committed.state().total, 15);
let mut replayed = AggregateRoot::<Counter>::new(CtrId::new(42));
replayed
.replay(Version::INITIAL, &CtrEvent::Added(10))
.expect("replay v1");
replayed.replay(v2, &CtrEvent::Added(5)).expect("replay v2");
assert_eq!(committed.version(), replayed.version());
assert_eq!(committed.state().total, replayed.state().total);
}
#[test]
fn advance_version_sets_version_without_applying_state() {
let mut agg = AggregateRoot::<Counter>::new(CtrId::new(1));
assert_eq!(agg.version(), None);
agg.advance_version(Version::INITIAL);
assert_eq!(agg.version(), Version::new(1));
assert_eq!(agg.state().total, 0);
agg.advance_version(Version::INITIAL);
assert_eq!(agg.version(), Version::new(1));
}
#[test]
fn apply_events_folds_state_without_advancing_version() {
let mut agg = AggregateRoot::<Counter>::new(CtrId::new(1));
let decided: Events<CtrEvent, 1> = events![CtrEvent::Added(2), CtrEvent::Added(3)];
agg.apply_events(&decided);
assert_eq!(agg.state().total, 5);
assert_eq!(agg.version(), None);
}
#[test]
fn apply_event_accumulates_state_without_advancing_version() {
let mut agg = AggregateRoot::<Counter>::new(CtrId::new(1));
agg.apply_event(&CtrEvent::Added(1));
assert_eq!(agg.state().total, 1);
agg.apply_event(&CtrEvent::Added(9));
assert_eq!(agg.state().total, 10);
assert_eq!(agg.version(), None);
}
#[test]
fn apply_events_mid_batch_panic_leaves_initial_state() {
use std::panic;
#[derive(Debug, Clone)]
enum BoomEvent {
Inc,
Boom,
}
impl Message for BoomEvent {}
impl DomainEvent for BoomEvent {
fn name(&self) -> &'static str {
match self {
Self::Inc => "Inc",
Self::Boom => "Boom",
}
}
}
#[derive(Default, Debug)]
struct BoomState {
count: u64,
}
impl AggregateState for BoomState {
type Event = BoomEvent;
fn initial() -> Self {
Self::default()
}
fn apply(mut self, event: &BoomEvent) -> Self {
match event {
BoomEvent::Inc => self.count += 1,
BoomEvent::Boom => panic!("boom in apply_events"),
}
self
}
}
struct BoomAgg;
impl Aggregate for BoomAgg {
type State = BoomState;
type Error = CtrError;
type Id = CtrId;
}
let mut agg = AggregateRoot::<BoomAgg>::new(CtrId::new(1));
let events: Events<BoomEvent, 1> = events![BoomEvent::Inc, BoomEvent::Boom];
let result = panic::catch_unwind(panic::AssertUnwindSafe(|| {
agg.apply_events(&events);
}));
assert!(result.is_err(), "apply_events should have panicked");
assert_eq!(
agg.state().count,
0,
"state must be left at initial() after a mid-batch panic"
);
assert_eq!(
agg.version(),
None,
"version must remain None — apply_events does not set version"
);
}
#[test]
fn replay_folds_state_without_clone() {
let mut root = AggregateRoot::<Counter>::new(CtrId::new(7));
root.replay(Version::INITIAL, &CtrEvent::Added(10))
.expect("replay v1");
root.replay(Version::new(2).expect("nonzero"), &CtrEvent::Added(5))
.expect("replay v2");
assert_eq!(root.state().total, 15);
assert_eq!(root.version(), Version::new(2));
}
}