use std::fmt;
use std::num::NonZeroUsize;
use mnesis::Aggregate;
use mnesis::AggregateRoot;
use mnesis::AggregateState;
use mnesis::DomainEvent;
use mnesis::Events;
use mnesis::Handle;
use mnesis::KernelError;
use mnesis::Message;
use mnesis::Version;
use mnesis::events;
use mnesis::testing::AggregateFixture;
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
struct TestId(String);
impl fmt::Display for TestId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.0)
}
}
impl AsRef<[u8]> for TestId {
fn as_ref(&self) -> &[u8] {
self.0.as_bytes()
}
}
#[derive(Debug, Clone, PartialEq)]
enum CounterEvent {
Incremented,
Decremented,
IncrementedBy(u64),
}
impl Message for CounterEvent {}
impl DomainEvent for CounterEvent {
fn name(&self) -> &'static str {
match self {
Self::Incremented => "Incremented",
Self::Decremented => "Decremented",
Self::IncrementedBy(_) => "IncrementedBy",
}
}
}
#[derive(Default, Debug, Clone)]
struct CounterState {
value: i64,
}
impl AggregateState for CounterState {
type Event = CounterEvent;
fn initial() -> Self {
Self::default()
}
fn apply(mut self, event: &CounterEvent) -> Self {
match event {
CounterEvent::Incremented => self.value += 1,
CounterEvent::Decremented => self.value -= 1,
CounterEvent::IncrementedBy(n) => self.value += i64::try_from(*n).unwrap_or(i64::MAX),
}
self
}
}
#[derive(Debug, thiserror::Error)]
enum CounterError {
#[error("counter would go negative")]
WouldGoNegative,
#[error("increment amount must be positive")]
ZeroIncrement,
}
struct Counter;
impl Aggregate for Counter {
type State = CounterState;
type Error = CounterError;
type Id = TestId;
}
struct Increment;
struct IncrementBy {
amount: u64,
}
struct Decrement;
impl Handle<Increment> for Counter {
fn handle(
_state: &CounterState,
_cmd: Increment,
) -> Result<Events<CounterEvent>, CounterError> {
Ok(events![CounterEvent::Incremented])
}
}
impl Handle<IncrementBy> for Counter {
fn handle(
_state: &CounterState,
cmd: IncrementBy,
) -> Result<Events<CounterEvent>, CounterError> {
if cmd.amount == 0 {
return Err(CounterError::ZeroIncrement);
}
Ok(events![CounterEvent::IncrementedBy(cmd.amount)])
}
}
impl Handle<Decrement> for Counter {
fn handle(state: &CounterState, _cmd: Decrement) -> Result<Events<CounterEvent>, CounterError> {
if state.value <= 0 {
return Err(CounterError::WouldGoNegative);
}
Ok(events![CounterEvent::Decremented])
}
}
#[test]
fn new_aggregate_has_none_version() {
let agg = AggregateRoot::<Counter>::new(TestId("1".into()));
assert_eq!(agg.version(), None);
}
#[test]
fn new_aggregate_has_initial_state() {
let agg = AggregateRoot::<Counter>::new(TestId("1".into()));
assert_eq!(agg.state().value, 0);
}
#[test]
fn aggregate_id_is_accessible() {
let agg = AggregateRoot::<Counter>::new(TestId("abc".into()));
assert_eq!(agg.id(), &TestId("abc".into()));
}
#[test]
fn replay_single_event_advances_version_and_state() {
let mut agg = AggregateRoot::<Counter>::new(TestId("1".into()));
agg.replay(Version::INITIAL, &CounterEvent::Incremented)
.unwrap();
assert_eq!(agg.version(), Some(Version::INITIAL));
assert_eq!(agg.state().value, 1);
}
#[test]
fn replay_multiple_events_sequentially() {
let mut agg = AggregateRoot::<Counter>::new(TestId("1".into()));
agg.replay(Version::new(1).unwrap(), &CounterEvent::Incremented)
.unwrap();
agg.replay(Version::new(2).unwrap(), &CounterEvent::Incremented)
.unwrap();
agg.replay(Version::new(3).unwrap(), &CounterEvent::Decremented)
.unwrap();
assert_eq!(agg.version(), Version::new(3));
assert_eq!(agg.state().value, 1);
}
#[test]
fn replay_rejects_version_gap() {
let mut agg = AggregateRoot::<Counter>::new(TestId("1".into()));
agg.replay(Version::new(1).unwrap(), &CounterEvent::Incremented)
.unwrap();
let err = agg
.replay(Version::new(3).unwrap(), &CounterEvent::Incremented)
.unwrap_err();
match err {
KernelError::VersionMismatch { expected, actual } => {
assert_eq!(expected, Version::new(2).unwrap());
assert_eq!(actual, Version::new(3).unwrap());
}
other => panic!("expected VersionMismatch, got {other:?}"),
}
}
#[test]
fn replay_rejects_gap_on_fresh_aggregate() {
let mut agg = AggregateRoot::<Counter>::new(TestId("1".into()));
let err = agg
.replay(Version::new(5).unwrap(), &CounterEvent::Incremented)
.unwrap_err();
match err {
KernelError::VersionMismatch { expected, actual } => {
assert_eq!(expected, Version::INITIAL);
assert_eq!(actual, Version::new(5).unwrap());
}
other => panic!("expected VersionMismatch, got {other:?}"),
}
}
#[test]
fn replay_rejects_duplicate_version() {
let mut agg = AggregateRoot::<Counter>::new(TestId("1".into()));
agg.replay(Version::new(1).unwrap(), &CounterEvent::Incremented)
.unwrap();
let err = agg
.replay(Version::new(1).unwrap(), &CounterEvent::Incremented)
.unwrap_err();
match err {
KernelError::VersionMismatch { expected, actual } => {
assert_eq!(expected, Version::new(2).unwrap());
assert_eq!(actual, Version::new(1).unwrap());
}
other => panic!("expected VersionMismatch, got {other:?}"),
}
}
#[test]
fn replay_does_not_mutate_state_on_version_gap() {
let mut agg = AggregateRoot::<Counter>::new(TestId("1".into()));
agg.replay(Version::new(1).unwrap(), &CounterEvent::Incremented)
.unwrap();
let _ = agg.replay(Version::new(3).unwrap(), &CounterEvent::Incremented);
assert_eq!(agg.version(), Version::new(1));
assert_eq!(agg.state().value, 1);
}
#[test]
fn handle_increment_returns_event() {
AggregateFixture::<Counter>::with_id(TestId("1".into()))
.given([])
.when(Increment)
.then_expect_events([CounterEvent::Incremented]);
}
#[test]
fn handle_increment_by_returns_event_with_amount() {
AggregateFixture::<Counter>::with_id(TestId("1".into()))
.given([])
.when(IncrementBy { amount: 42 })
.then_expect_events([CounterEvent::IncrementedBy(42)]);
}
#[test]
fn handle_rejects_invalid_command() {
AggregateFixture::<Counter>::with_id(TestId("1".into()))
.given([])
.when(Decrement)
.then_expect_error_matching(|e| matches!(e, CounterError::WouldGoNegative));
}
#[test]
fn handle_rejects_zero_increment() {
AggregateFixture::<Counter>::with_id(TestId("1".into()))
.given([])
.when(IncrementBy { amount: 0 })
.then_expect_error_matching(|e| matches!(e, CounterError::ZeroIncrement));
}
#[test]
fn handle_uses_current_state_for_decision() {
AggregateFixture::<Counter>::with_id(TestId("1".into()))
.given([CounterEvent::Incremented])
.when(Decrement)
.then_expect_events([CounterEvent::Decremented]);
}
#[test]
fn commit_persisted_advances_version_and_folds_state() {
let mut agg = AggregateRoot::<Counter>::new(TestId("1".into()));
let decided: Events<_, 1> = events![CounterEvent::Incremented, CounterEvent::IncrementedBy(9)];
let new_version = Version::new(2).unwrap();
agg.commit_persisted(new_version, &decided);
assert_eq!(agg.version(), Version::new(2));
assert_eq!(agg.state().value, 10);
}
#[test]
fn commit_persisted_then_replay_continues_from_committed_version() {
let mut agg = AggregateRoot::<Counter>::new(TestId("1".into()));
agg.replay(Version::new(1).unwrap(), &CounterEvent::Incremented)
.unwrap();
let decided: Events<_, 1> = events![CounterEvent::Incremented, CounterEvent::Incremented];
agg.commit_persisted(Version::new(3).unwrap(), &decided);
assert_eq!(agg.version(), Version::new(3));
assert_eq!(agg.state().value, 3);
agg.replay(Version::new(4).unwrap(), &CounterEvent::Decremented)
.unwrap();
assert_eq!(agg.version(), Version::new(4));
assert_eq!(agg.state().value, 2);
}
struct TinyLimitAggregate;
#[allow(clippy::unwrap_used, reason = "3 is non-zero by inspection")]
impl Aggregate for TinyLimitAggregate {
type State = CounterState;
type Error = CounterError;
type Id = TestId;
const MAX_REHYDRATION_EVENTS: NonZeroUsize = NonZeroUsize::new(3).unwrap();
}
#[test]
fn replay_respects_rehydration_limit() {
let mut agg = AggregateRoot::<TinyLimitAggregate>::new(TestId("1".into()));
agg.replay(Version::new(1).unwrap(), &CounterEvent::Incremented)
.unwrap();
agg.replay(Version::new(2).unwrap(), &CounterEvent::Incremented)
.unwrap();
agg.replay(Version::new(3).unwrap(), &CounterEvent::Incremented)
.unwrap();
let err = agg
.replay(Version::new(4).unwrap(), &CounterEvent::Incremented)
.unwrap_err();
assert!(matches!(
err,
KernelError::RehydrationLimitExceeded { max: 3 }
));
}
#[test]
fn version_new_rejects_zero() {
assert!(Version::new(0).is_none());
}
#[test]
fn version_new_accepts_nonzero() {
let v = Version::new(42).unwrap();
assert_eq!(v.as_u64(), 42);
}
#[test]
fn version_initial_is_one() {
assert_eq!(Version::INITIAL.as_u64(), 1);
}
#[test]
fn restore_creates_root_at_given_state_and_version() {
let id = TestId("1".into());
let state = CounterState { value: 42 };
let version = Version::new(10).unwrap();
let root = AggregateRoot::<Counter>::restore(id.clone(), state, version);
assert_eq!(root.id(), &id);
assert_eq!(root.state().value, 42);
assert_eq!(root.version(), Some(version));
}
#[test]
fn restore_then_replay_continues_from_snapshot_version() {
let id = TestId("1".into());
let state = CounterState { value: 42 };
let version = Version::new(10).unwrap();
let mut root = AggregateRoot::<Counter>::restore(id, state, version);
let result = root.replay(Version::new(11).unwrap(), &CounterEvent::Incremented);
assert!(result.is_ok());
assert_eq!(root.state().value, 43);
assert_eq!(root.version(), Some(Version::new(11).unwrap()));
}
#[test]
fn restore_then_replay_rejects_wrong_version() {
let id = TestId("1".into());
let state = CounterState { value: 42 };
let version = Version::new(10).unwrap();
let mut root = AggregateRoot::<Counter>::restore(id, state, version);
let result = root.replay(Version::new(10).unwrap(), &CounterEvent::Incremented);
assert!(result.is_err());
}