use diapause::{Coroutine, CoroutineState, Fingerprinted};
use serde::{Deserialize, Serialize};
#[diapause::coroutine(yield = u32, resume = u32)]
#[derive(Serialize, Deserialize)]
fn running_sum(n: u32) -> u32 {
let mut sum: u32 = 0;
for i in 0u32..n {
let w = yield_!(i);
sum += i * w;
}
sum
}
#[test]
fn suspended_for_loop_round_trips_through_json() {
let mut c = running_sum(4);
assert_eq!(c.start(), CoroutineState::Yielded(0));
assert_eq!(c.resume(1), CoroutineState::Yielded(1));
let json = serde_json::to_string(&c).unwrap();
let mut restored: running_sum::State = serde_json::from_str(&json).unwrap();
for c in [&mut c, &mut restored] {
assert_eq!(c.resume(10), CoroutineState::Yielded(2));
assert_eq!(c.resume(10), CoroutineState::Yielded(3));
assert_eq!(c.resume(10), CoroutineState::Complete(60));
}
}
#[diapause::coroutine(yield = u32, resume = u32, fingerprint)]
#[derive(Serialize, Deserialize)]
fn tally_v1(n: u32) -> u32 {
let mut sum: u32 = 0;
for i in 0u32..n {
let w = yield_!(i);
sum += i * w;
}
sum
}
#[diapause::coroutine(yield = u32, resume = u32, fingerprint)]
#[derive(Serialize, Deserialize)]
fn tally_v2(n: u32) -> u32 {
let mut sum: u32 = 0;
for i in 0u32..n {
let w = yield_!(i);
sum += i + w;
}
sum
}
#[test]
fn fingerprinted_state_round_trips_through_json() {
let mut c = tally_v1(4);
assert_eq!(c.start(), CoroutineState::Yielded(0));
assert_eq!(c.resume(1), CoroutineState::Yielded(1));
let json = serde_json::to_string(&c).unwrap();
let mut restored: tally_v1::State = serde_json::from_str(&json).unwrap();
restored.check_fingerprint().unwrap();
for c in [&mut c, &mut restored] {
assert_eq!(c.resume(10), CoroutineState::Yielded(2));
assert_eq!(c.resume(10), CoroutineState::Yielded(3));
assert_eq!(c.resume(10), CoroutineState::Complete(60));
}
}
#[test]
fn fingerprints_differ_between_edited_sources() {
assert_ne!(tally_v1::State::FINGERPRINT, tally_v2::State::FINGERPRINT);
}
#[test]
fn check_fingerprint_rejects_a_state_from_edited_source() {
let mut c = tally_v1(4);
let _ = c.start();
let json = serde_json::to_string(&c).unwrap();
let restored: tally_v2::State = serde_json::from_str(&json).unwrap();
let err = restored.check_fingerprint().unwrap_err();
assert_eq!(err.expected, tally_v2::State::FINGERPRINT);
assert_eq!(err.found, tally_v1::State::FINGERPRINT);
let err: &dyn std::error::Error = &err;
assert!(err.to_string().contains("fingerprint mismatch"));
}
#[test]
#[should_panic(expected = "this state was created by a different version of `tally_v2`")]
fn resume_panics_on_a_state_from_edited_source() {
let mut c = tally_v1(4);
let _ = c.start();
let json = serde_json::to_string(&c).unwrap();
let mut restored: tally_v2::State = serde_json::from_str(&json).unwrap();
let _ = restored.resume(1);
}
#[test]
#[should_panic(expected = "this state was created by a different version of `tally_v2`")]
fn start_panics_on_a_state_from_edited_source() {
let c = tally_v1(4);
let json = serde_json::to_string(&c).unwrap();
let mut restored: tally_v2::State = serde_json::from_str(&json).unwrap();
let _ = restored.start();
}
#[test]
fn enabling_fingerprint_invalidates_old_persisted_states() {
let mut c = running_sum(4);
let _ = c.start();
let json = serde_json::to_string(&c).unwrap();
assert!(serde_json::from_str::<tally_v1::State>(&json).is_err());
}
#[diapause::coroutine(yield = u32, resume = u32, fingerprint = "tally-pin")]
#[derive(Serialize, Deserialize)]
fn pinned_v1(n: u32) -> u32 {
let mut sum: u32 = 0;
for i in 0u32..n {
let w = yield_!(i);
sum += i * w;
}
sum
}
#[diapause::coroutine(yield = u32, resume = u32, fingerprint = "tally-pin")]
#[derive(Serialize, Deserialize)]
fn pinned_v2(n: u32) -> u32 {
let mut sum: u32 = 0;
for i in 0u32..n {
let w = yield_!(i);
sum += i + w;
}
sum
}
#[test]
fn manual_fingerprint_pins_compatibility_across_edits() {
assert_eq!(pinned_v1::State::FINGERPRINT, pinned_v2::State::FINGERPRINT);
let mut c = pinned_v1(4);
let _ = c.start();
let json = serde_json::to_string(&c).unwrap();
let mut restored: pinned_v2::State = serde_json::from_str(&json).unwrap();
restored.check_fingerprint().unwrap();
assert_eq!(restored.resume(10), CoroutineState::Yielded(1));
}
#[test]
fn fingerprinted_is_usable_generically() {
fn validate<S: Fingerprinted>(s: &S) -> Result<(), diapause::FingerprintMismatch> {
s.check_fingerprint()
}
let c = tally_v1(4);
validate(&c).unwrap();
assert_eq!(
<tally_v1::State as Fingerprinted>::FINGERPRINT,
tally_v1::State::FINGERPRINT
);
}
#[test]
fn fingerprint_const_is_generated_without_the_flag() {
let fp: u64 = running_sum::State::FINGERPRINT;
assert_eq!(fp, running_sum::State::FINGERPRINT);
}
#[diapause::coroutine(yield = u32, resume = u32)]
#[derive(Serialize, Deserialize)]
fn sub_sum(n: u32) -> u32 {
let mut sum: u32 = 0;
for i in 0u32..n {
let w = yield_!(i);
sum += w;
}
sum
}
#[diapause::coroutine(yield = u32, resume = u32)]
#[derive(Serialize, Deserialize)]
fn delegating(n: u32) -> u32 {
let g: sub_sum::State = sub_sum(n);
let total: u32 = yield_all!(g);
total * 2
}
#[test]
fn suspended_delegation_round_trips_through_json() {
let mut c = delegating(3);
assert_eq!(c.start(), CoroutineState::Yielded(0));
assert_eq!(c.resume(10), CoroutineState::Yielded(1));
let value: serde_json::Value = serde_json::to_value(&c).unwrap();
let inner = &value["S2"]["__dg0"];
assert!(inner.is_object(), "inner state missing: {value}");
let json = serde_json::to_string(&c).unwrap();
let mut restored: delegating::State = serde_json::from_str(&json).unwrap();
for c in [&mut c, &mut restored] {
assert_eq!(c.resume(20), CoroutineState::Yielded(2));
assert_eq!(c.resume(30), CoroutineState::Complete(120));
}
}
#[diapause::coroutine(yield = u32)]
#[derive(Serialize, Deserialize)]
fn inclusive_sum(n: u32) -> u32 {
let mut sum: u32 = 0;
for i in 0u32..=n {
yield_!(i);
sum += i;
}
sum
}
#[test]
fn inclusive_range_round_trips_at_every_suspension() {
let mut c = inclusive_sum(2);
let mut step = c.start();
let mut yields = Vec::new();
while let CoroutineState::Yielded(v) = step {
yields.push(v);
let json = serde_json::to_string(&c).unwrap();
c = serde_json::from_str(&json).unwrap();
step = c.resume(());
}
assert_eq!(yields, [0, 1, 2]);
assert_eq!(step, CoroutineState::Complete(3));
}
#[test]
fn serialized_inclusive_iterator_exposes_exhaustion() {
let mut c = inclusive_sum(1);
let _ = c.start();
assert_eq!(c.resume(()), CoroutineState::Yielded(1));
let value: serde_json::Value = serde_json::to_value(&c).unwrap();
let it = &value["S1"]["__iter0"];
assert_eq!(it["start"], 1);
assert_eq!(it["end"], 1);
assert_eq!(it["done"], true);
}
#[test]
fn serialized_state_exposes_the_iterator_cursor() {
let mut c = running_sum(3);
let _ = c.start();
let value: serde_json::Value = serde_json::to_value(&c).unwrap();
let s1 = &value["S1"];
assert_eq!(s1["__iter0"]["start"], 1);
assert_eq!(s1["__iter0"]["end"], 3);
}