#![cfg(not(miri))]
use mnesis::*;
use proptest::prelude::*;
use std::fmt;
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
struct PId(String);
impl PId {
fn new(v: u64) -> Self {
Self(format!("p-{v}"))
}
}
impl fmt::Display for PId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.0)
}
}
impl AsRef<[u8]> for PId {
fn as_ref(&self) -> &[u8] {
self.0.as_bytes()
}
}
#[derive(Debug, Clone, PartialEq)]
enum CountEvent {
Incremented,
Decremented,
Set(u64),
}
impl Message for CountEvent {}
impl DomainEvent for CountEvent {
fn name(&self) -> &'static str {
match self {
Self::Incremented => "Incremented",
Self::Decremented => "Decremented",
Self::Set(_) => "Set",
}
}
}
#[derive(Default, Debug, PartialEq, Clone)]
struct CountState {
value: i64,
}
impl AggregateState for CountState {
type Event = CountEvent;
fn initial() -> Self {
Self::default()
}
fn apply(mut self, event: &CountEvent) -> Self {
match event {
CountEvent::Incremented => self.value += 1,
CountEvent::Decremented => self.value -= 1,
CountEvent::Set(v) => self.value = (*v).cast_signed(),
}
self
}
}
#[derive(Debug)]
struct CountAgg;
#[derive(Debug, thiserror::Error)]
#[error("count error")]
struct CountErr;
impl Aggregate for CountAgg {
type State = CountState;
type Error = CountErr;
type Id = PId;
}
fn arb_event() -> impl Strategy<Value = CountEvent> {
prop_oneof![
Just(CountEvent::Incremented),
Just(CountEvent::Decremented),
(0..1000u64).prop_map(CountEvent::Set),
]
}
fn replay_events(events: &[CountEvent]) -> AggregateRoot<CountAgg> {
let mut agg = AggregateRoot::<CountAgg>::new(PId::new(1));
for (i, e) in events.iter().enumerate() {
let v = Version::new((i + 1) as u64).unwrap();
agg.replay(v, e).unwrap();
}
agg
}
fn apply_events_to(events: &[CountEvent]) -> AggregateRoot<CountAgg> {
let mut agg = AggregateRoot::<CountAgg>::new(PId::new(1));
for (i, e) in events.iter().enumerate() {
let v = Version::new((i + 1) as u64).unwrap();
let batch: Events<_, 0> = Events::new(e.clone());
agg.commit_persisted(v, &batch);
}
agg
}
proptest! {
#[test]
fn prop_replay_is_deterministic(raw_events in proptest::collection::vec(arb_event(), 0..50)) {
let agg1 = replay_events(&raw_events);
let agg2 = replay_events(&raw_events);
prop_assert_eq!(agg1.state(), agg2.state());
prop_assert_eq!(agg1.version(), agg2.version());
}
#[test]
fn prop_replay_version_equals_event_count(events in proptest::collection::vec(arb_event(), 0..50)) {
let agg = replay_events(&events);
let n = events.len() as u64;
if n == 0 {
prop_assert_eq!(agg.version(), None);
} else {
prop_assert_eq!(agg.version(), Version::new(n));
}
}
#[test]
fn prop_replay_equals_apply_events(raw_events in proptest::collection::vec(arb_event(), 0..50)) {
let replayed = replay_events(&raw_events);
let applied = apply_events_to(&raw_events);
prop_assert_eq!(replayed.state(), applied.state());
}
#[test]
fn prop_rejects_any_version_gap(
events in proptest::collection::vec(arb_event(), 2..50),
corrupt_idx in 1..50usize,
) {
let corrupt_idx = corrupt_idx % (events.len() - 1) + 1;
let mut agg = AggregateRoot::<CountAgg>::new(PId::new(1));
let mut found_error = false;
for (i, event) in events.iter().enumerate() {
let v = if i >= corrupt_idx { i + 2 } else { i + 1 };
let version = Version::new(v as u64).unwrap();
if agg.replay(version, event).is_err() {
found_error = true;
break;
}
}
prop_assert!(found_error, "Should reject sequence with version gap");
}
#[test]
fn prop_state_is_pure_function_of_events(
events in proptest::collection::vec(arb_event(), 1..50),
split_pct in 0..100u32,
) {
let full = replay_events(&events);
let mid = (events.len() * split_pct as usize) / 100;
let mut split = replay_events(&events[..mid]);
for (i, e) in events[mid..].iter().enumerate() {
let v = (mid + i + 1) as u64;
split.replay(Version::new(v).unwrap(), e).unwrap();
}
prop_assert_eq!(full.state(), split.state(), "full vs split diverged");
prop_assert_eq!(full.version(), split.version(), "full vs split version diverged");
}
#[test]
fn prop_replay_rejects_duplicate_version(events in proptest::collection::vec(arb_event(), 1..50)) {
let mut agg = replay_events(&events);
let last_version = agg.version().unwrap();
let result = agg.replay(last_version, &CountEvent::Incremented);
prop_assert!(result.is_err(), "Should reject duplicate version");
}
#[test]
fn prop_version_new_roundtrip(v in any::<u64>()) {
let version = Version::new(v);
if v == 0 {
prop_assert!(version.is_none(), "Version::new(0) must return None");
} else {
let version = version.unwrap();
prop_assert_eq!(version.as_u64(), v, "Version::new/as_u64 roundtrip failed");
}
}
#[test]
fn prop_version_next_is_monotonic(v in 1..u64::MAX) {
let version = Version::new(v).unwrap();
let next = version.next().unwrap();
prop_assert!(next > version, "next() must be strictly greater");
prop_assert_eq!(next.as_u64(), v + 1, "next() must increment by exactly 1");
}
}