use std::collections::{HashMap, VecDeque};
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
use serde::{Deserialize, Serialize};
use crate::callback::{Callback, CallbackError};
use crate::schema::Schema;
#[derive(Debug, Clone)]
pub enum CallbackOutcome {
Ok(serde_json::Value),
Err(String),
Suspend(String),
}
impl CallbackOutcome {
fn from_invoke(result: &Result<serde_json::Value, CallbackError>) -> Self {
match result {
Ok(value) => Self::Ok(value.clone()),
Err(CallbackError::Suspend(reason)) => Self::Suspend(reason.clone()),
Err(other) => Self::Err(other.to_string()),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CallbackJournalEntry {
pub index: u32,
pub name: String,
pub args_hash: u64,
pub args_json: String,
pub result: Result<serde_json::Value, String>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct CallbackJournal {
pub code: String,
pub entries: Vec<CallbackJournalEntry>,
}
impl CallbackJournal {
#[must_use]
pub fn new(code: impl Into<String>) -> Self {
Self {
code: code.into(),
entries: Vec::new(),
}
}
#[must_use]
pub fn len(&self) -> usize {
self.entries.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SuspendedCallback {
pub name: String,
pub args_json: String,
pub reason: String,
}
#[derive(Debug)]
pub struct ReplayState {
cached: HashMap<(String, String), VecDeque<Result<serde_json::Value, String>>>,
live_mode: bool,
next_seq: usize,
entries: Vec<Option<CallbackJournalEntry>>,
suspended: Option<SuspendedCallback>,
suspend_seq: Option<usize>,
replayed_count: u32,
}
enum Decision {
Hit(Result<serde_json::Value, String>),
Miss { seq: usize },
AlreadySuspended,
}
impl ReplayState {
#[must_use]
pub fn new(previous: CallbackJournal) -> Self {
let mut cached: HashMap<(String, String), VecDeque<Result<serde_json::Value, String>>> =
HashMap::new();
for entry in previous.entries {
cached
.entry((entry.name, entry.args_json))
.or_default()
.push_back(entry.result);
}
Self {
cached,
live_mode: false,
next_seq: 0,
entries: Vec::new(),
suspended: None,
suspend_seq: None,
replayed_count: 0,
}
}
fn decide(&mut self, name: &str, args_hash: u64, args_json: &str) -> Decision {
if self.suspended.is_some() {
return Decision::AlreadySuspended;
}
let seq = self.next_seq;
self.next_seq += 1;
self.ensure_slot(seq);
if self.live_mode {
return Decision::Miss { seq };
}
let key = (name.to_string(), args_json.to_string());
if let Some(queue) = self.cached.get_mut(&key)
&& let Some(result) = queue.pop_front()
{
self.replayed_count += 1;
self.entries[seq] = Some(CallbackJournalEntry {
index: u32::try_from(seq).unwrap_or(u32::MAX),
name: name.to_string(),
args_hash,
args_json: args_json.to_string(),
result: result.clone(),
});
return Decision::Hit(result);
}
self.live_mode = true;
Decision::Miss { seq }
}
fn ensure_slot(&mut self, seq: usize) {
if self.entries.len() <= seq {
self.entries.resize(seq + 1, None);
}
}
fn record_live(&mut self, entry: CallbackJournalEntry) {
let seq = entry.index as usize;
self.ensure_slot(seq);
self.entries[seq] = Some(entry);
}
fn record_suspend(&mut self, seq: usize, suspended: SuspendedCallback) {
if self.suspended.is_none() {
self.suspended = Some(suspended);
self.suspend_seq = Some(seq);
}
}
#[must_use]
pub fn suspended(&self) -> Option<&SuspendedCallback> {
self.suspended.as_ref()
}
#[must_use]
pub fn replayed_count(&self) -> u32 {
self.replayed_count
}
#[must_use]
pub fn build_journal(&self, code: impl Into<String>) -> CallbackJournal {
let entries = self
.entries
.iter()
.flatten()
.filter(|entry| match self.suspend_seq {
Some(seq) => (entry.index as usize) < seq,
None => true,
})
.cloned()
.collect();
CallbackJournal {
code: code.into(),
entries,
}
}
}
pub struct ReplayCallback {
inner: Arc<dyn Callback>,
state: Arc<Mutex<ReplayState>>,
}
impl std::fmt::Debug for ReplayCallback {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ReplayCallback")
.field("name", &self.inner.name())
.finish_non_exhaustive()
}
}
impl ReplayCallback {
#[must_use]
pub fn new(inner: Arc<dyn Callback>, state: Arc<Mutex<ReplayState>>) -> Self {
Self { inner, state }
}
}
impl Callback for ReplayCallback {
fn name(&self) -> &str {
self.inner.name()
}
fn description(&self) -> &str {
self.inner.description()
}
fn parameters_schema(&self) -> Schema {
self.inner.parameters_schema()
}
fn invoke(
&self,
args: serde_json::Value,
) -> Pin<Box<dyn Future<Output = Result<serde_json::Value, CallbackError>> + Send + '_>> {
let name = self.inner.name().to_string();
let args_json = canonical_json(&args);
let args_hash = fnv1a_64(args_json.as_bytes());
let decision = lock_state(&self.state).decide(&name, args_hash, &args_json);
match decision {
Decision::AlreadySuspended => Box::pin(async {
Err(CallbackError::Suspend(
"execution already suspended by a previous callback".into(),
))
}),
Decision::Hit(result) => {
let invoke_result = match result {
Ok(value) => Ok(value),
Err(message) => Err(CallbackError::Replayed(message)),
};
Box::pin(async move { invoke_result })
}
Decision::Miss { seq } => {
let inner = Arc::clone(&self.inner);
let state = Arc::clone(&self.state);
Box::pin(async move {
let result = inner.invoke(args).await;
let index = u32::try_from(seq).unwrap_or(u32::MAX);
match CallbackOutcome::from_invoke(&result) {
CallbackOutcome::Ok(value) => {
lock_state(&state).record_live(CallbackJournalEntry {
index,
name,
args_hash,
args_json,
result: Ok(value),
});
}
CallbackOutcome::Err(message) => {
lock_state(&state).record_live(CallbackJournalEntry {
index,
name,
args_hash,
args_json,
result: Err(message),
});
}
CallbackOutcome::Suspend(reason) => {
lock_state(&state).record_suspend(
seq,
SuspendedCallback {
name,
args_json,
reason,
},
);
}
}
result
})
}
}
}
}
pub(crate) fn wrap_callbacks(
callbacks: &HashMap<String, Arc<dyn Callback>>,
state: &Arc<Mutex<ReplayState>>,
) -> HashMap<String, Arc<dyn Callback>> {
callbacks
.iter()
.map(|(name, callback)| {
let wrapped: Arc<dyn Callback> =
Arc::new(ReplayCallback::new(Arc::clone(callback), Arc::clone(state)));
(name.clone(), wrapped)
})
.collect()
}
fn lock_state(state: &Mutex<ReplayState>) -> MutexGuard<'_, ReplayState> {
state.lock().unwrap_or_else(PoisonError::into_inner)
}
fn canonical_json(value: &serde_json::Value) -> String {
canonicalize(value).to_string()
}
fn canonicalize(value: &serde_json::Value) -> serde_json::Value {
match value {
serde_json::Value::Object(map) => {
let mut keys: Vec<&String> = map.keys().collect();
keys.sort();
let mut sorted = serde_json::Map::new();
for key in keys {
if let Some(v) = map.get(key) {
sorted.insert(key.clone(), canonicalize(v));
}
}
serde_json::Value::Object(sorted)
}
serde_json::Value::Array(items) => {
serde_json::Value::Array(items.iter().map(canonicalize).collect())
}
other => other.clone(),
}
}
fn fnv1a_64(bytes: &[u8]) -> u64 {
const OFFSET_BASIS: u64 = 0xcbf2_9ce4_8422_2325;
const PRIME: u64 = 0x0000_0100_0000_01b3;
let mut hash = OFFSET_BASIS;
for &byte in bytes {
hash ^= u64::from(byte);
hash = hash.wrapping_mul(PRIME);
}
hash
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use serde_json::json;
use std::sync::atomic::{AtomicU32, Ordering};
struct ProgrammableCallback {
name: String,
calls: Arc<AtomicU32>,
outcome: Outcome,
}
#[derive(Clone)]
enum Outcome {
Ok(serde_json::Value),
Err(String),
Suspend(String),
}
impl Callback for ProgrammableCallback {
fn name(&self) -> &str {
&self.name
}
fn description(&self) -> &str {
"test"
}
fn parameters_schema(&self) -> Schema {
Schema::empty()
}
fn invoke(
&self,
_args: serde_json::Value,
) -> Pin<Box<dyn Future<Output = Result<serde_json::Value, CallbackError>> + Send + '_>>
{
self.calls.fetch_add(1, Ordering::SeqCst);
let outcome = self.outcome.clone();
Box::pin(async move {
match outcome {
Outcome::Ok(v) => Ok(v),
Outcome::Err(m) => Err(CallbackError::ExecutionFailed(m)),
Outcome::Suspend(r) => Err(CallbackError::Suspend(r)),
}
})
}
}
fn programmable(name: &str, outcome: Outcome) -> (Arc<ProgrammableCallback>, Arc<AtomicU32>) {
let calls = Arc::new(AtomicU32::new(0));
let cb = Arc::new(ProgrammableCallback {
name: name.to_string(),
calls: Arc::clone(&calls),
outcome,
});
(cb, calls)
}
fn entry_ok(
index: u32,
name: &str,
args: &serde_json::Value,
value: serde_json::Value,
) -> CallbackJournalEntry {
let args_json = canonical_json(args);
CallbackJournalEntry {
index,
name: name.to_string(),
args_hash: fnv1a_64(args_json.as_bytes()),
args_json,
result: Ok(value),
}
}
#[test]
fn outcome_classifies_ok_err() {
assert!(matches!(
CallbackOutcome::from_invoke(&Ok(json!(1))),
CallbackOutcome::Ok(_)
));
assert!(matches!(
CallbackOutcome::from_invoke(&Err(CallbackError::ExecutionFailed("x".into()))),
CallbackOutcome::Err(_)
));
}
#[test]
fn canonicalization_is_key_order_independent() {
let a = json!({"a": 1, "b": {"c": 2, "d": 3}});
let b = json!({"b": {"d": 3, "c": 2}, "a": 1});
assert_eq!(canonical_json(&a), canonical_json(&b));
assert_eq!(
fnv1a_64(canonical_json(&a).as_bytes()),
fnv1a_64(canonical_json(&b).as_bytes())
);
}
#[tokio::test]
async fn full_match_replays_without_invoking() {
let args = json!({"q": "hello"});
let previous = CallbackJournal {
code: "code".into(),
entries: vec![entry_ok(0, "fetch", &args, json!({"v": 1}))],
};
let state = Arc::new(Mutex::new(ReplayState::new(previous)));
let (inner, calls) = programmable("fetch", Outcome::Ok(json!({"v": 999})));
let wrapper = ReplayCallback::new(inner, Arc::clone(&state));
let result = wrapper.invoke(args.clone()).await.unwrap();
assert_eq!(result, json!({"v": 1}), "should return cached value");
assert_eq!(calls.load(Ordering::SeqCst), 0, "real callback not invoked");
let st = lock_state(&state);
assert_eq!(st.replayed_count(), 1);
let journal = st.build_journal("code");
assert_eq!(journal.entries.len(), 1);
}
#[tokio::test]
async fn empty_journal_goes_live_and_records() {
let args = json!({"q": "hi"});
let state = Arc::new(Mutex::new(ReplayState::new(CallbackJournal::new("code"))));
let (inner, calls) = programmable("fetch", Outcome::Ok(json!({"v": 7})));
let wrapper = ReplayCallback::new(inner, Arc::clone(&state));
let result = wrapper.invoke(args.clone()).await.unwrap();
assert_eq!(result, json!({"v": 7}));
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"real callback invoked once"
);
let st = lock_state(&state);
assert_eq!(st.replayed_count(), 0);
let journal = st.build_journal("code");
assert_eq!(journal.entries.len(), 1);
assert_eq!(journal.entries[0].result, Ok(json!({"v": 7})));
}
#[tokio::test]
async fn first_miss_switches_to_live_mode() {
let a_args = json!({"id": "A"});
let b_args = json!({"id": "B"});
let previous = CallbackJournal {
code: "code".into(),
entries: vec![
entry_ok(0, "a", &a_args, json!("cached-a")),
entry_ok(1, "b", &b_args, json!("cached-b")),
],
};
let state = Arc::new(Mutex::new(ReplayState::new(previous)));
let (c, c_calls) = programmable("c", Outcome::Ok(json!("live-c")));
let c_wrapper = ReplayCallback::new(c, Arc::clone(&state));
let rc = c_wrapper.invoke(json!({"id": "C"})).await.unwrap();
assert_eq!(rc, json!("live-c"));
assert_eq!(c_calls.load(Ordering::SeqCst), 1);
let (a, a_calls) = programmable("a", Outcome::Ok(json!("live-a")));
let a_wrapper = ReplayCallback::new(a, Arc::clone(&state));
let ra = a_wrapper.invoke(a_args.clone()).await.unwrap();
assert_eq!(
ra,
json!("live-a"),
"a runs live after divergence, not replayed"
);
assert_eq!(a_calls.load(Ordering::SeqCst), 1, "a actually invoked");
let st = lock_state(&state);
assert_eq!(
st.replayed_count(),
0,
"nothing replayed after the first miss"
);
}
#[tokio::test]
async fn write_miss_before_read_prevents_stale_replay() {
let read_args = json!({});
let previous = CallbackJournal {
code: "code".into(),
entries: vec![entry_ok(0, "read_counter", &read_args, json!(5))],
};
let state = Arc::new(Mutex::new(ReplayState::new(previous)));
let (write, write_calls) = programmable("set_counter", Outcome::Ok(json!("ok")));
let write_wrapper = ReplayCallback::new(write, Arc::clone(&state));
let _ = write_wrapper.invoke(json!({"value": 10})).await.unwrap();
assert_eq!(write_calls.load(Ordering::SeqCst), 1);
let (read, read_calls) = programmable("read_counter", Outcome::Ok(json!(10)));
let read_wrapper = ReplayCallback::new(read, Arc::clone(&state));
let r = read_wrapper.invoke(read_args.clone()).await.unwrap();
assert_eq!(
r,
json!(10),
"read runs live, returning the true post-write value"
);
assert_eq!(
read_calls.load(Ordering::SeqCst),
1,
"read not replayed from stale cache"
);
assert_eq!(lock_state(&state).replayed_count(), 0);
}
#[tokio::test]
async fn matching_is_order_independent() {
let a_args = json!({"id": "A"});
let b_args = json!({"id": "B"});
let previous = CallbackJournal {
code: "code".into(),
entries: vec![
entry_ok(0, "a", &a_args, json!("cached-a")),
entry_ok(1, "b", &b_args, json!("cached-b")),
],
};
let state = Arc::new(Mutex::new(ReplayState::new(previous)));
let (a, a_calls) = programmable("a", Outcome::Ok(json!("live-a")));
let (b, b_calls) = programmable("b", Outcome::Ok(json!("live-b")));
let a_wrapper = ReplayCallback::new(a, Arc::clone(&state));
let b_wrapper = ReplayCallback::new(b, Arc::clone(&state));
let rb = b_wrapper.invoke(b_args.clone()).await.unwrap();
let ra = a_wrapper.invoke(a_args.clone()).await.unwrap();
assert_eq!(rb, json!("cached-b"), "B replayed out of order");
assert_eq!(ra, json!("cached-a"), "A replayed out of order");
assert_eq!(b_calls.load(Ordering::SeqCst), 0, "B not invoked live");
assert_eq!(a_calls.load(Ordering::SeqCst), 0, "A not invoked live");
let st = lock_state(&state);
assert_eq!(st.replayed_count(), 2);
}
#[tokio::test]
async fn repeated_identical_calls_replay_fifo() {
let args = json!({"q": "x"});
let previous = CallbackJournal {
code: "code".into(),
entries: vec![
entry_ok(0, "fetch", &args, json!("first")),
entry_ok(1, "fetch", &args, json!("second")),
],
};
let state = Arc::new(Mutex::new(ReplayState::new(previous)));
let (inner, calls) = programmable("fetch", Outcome::Ok(json!("live")));
let w1 = ReplayCallback::new(Arc::clone(&inner) as Arc<dyn Callback>, Arc::clone(&state));
let r1 = w1.invoke(args.clone()).await.unwrap();
let w2 = ReplayCallback::new(inner as Arc<dyn Callback>, Arc::clone(&state));
let r2 = w2.invoke(args.clone()).await.unwrap();
assert_eq!(r1, json!("first"));
assert_eq!(r2, json!("second"));
assert_eq!(calls.load(Ordering::SeqCst), 0, "neither invoked live");
assert_eq!(lock_state(&state).replayed_count(), 2);
}
#[tokio::test]
async fn args_mismatch_is_a_miss() {
let cached_args = json!({"q": "a"});
let previous = CallbackJournal {
code: "code".into(),
entries: vec![entry_ok(0, "fetch", &cached_args, json!("cached"))],
};
let state = Arc::new(Mutex::new(ReplayState::new(previous)));
let (inner, calls) = programmable("fetch", Outcome::Ok(json!("live")));
let wrapper = ReplayCallback::new(inner, Arc::clone(&state));
let result = wrapper.invoke(json!({"q": "DIFFERENT"})).await.unwrap();
assert_eq!(result, json!("live"));
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn replayed_error_is_returned_transparently() {
let args = json!({});
let previous = CallbackJournal {
code: "code".into(),
entries: vec![CallbackJournalEntry {
index: 0,
name: "fail".into(),
args_hash: fnv1a_64(canonical_json(&args).as_bytes()),
args_json: canonical_json(&args),
result: Err("execution failed: boom".into()),
}],
};
let state = Arc::new(Mutex::new(ReplayState::new(previous)));
let (inner, calls) = programmable("fail", Outcome::Ok(json!("unused")));
let wrapper = ReplayCallback::new(inner, Arc::clone(&state));
let err = wrapper.invoke(args).await.unwrap_err();
assert_eq!(err.to_string(), "execution failed: boom");
assert_eq!(calls.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn live_error_is_journaled() {
let state = Arc::new(Mutex::new(ReplayState::new(CallbackJournal::new("code"))));
let (inner, _) = programmable("fail", Outcome::Err("boom".into()));
let wrapper = ReplayCallback::new(inner, Arc::clone(&state));
let err = wrapper.invoke(json!({})).await.unwrap_err();
assert!(matches!(err, CallbackError::ExecutionFailed(_)));
let st = lock_state(&state);
let journal = st.build_journal("code");
assert_eq!(journal.entries.len(), 1);
assert_eq!(
journal.entries[0].result,
Err("execution failed: boom".to_string())
);
}
#[test]
fn outcome_classifies_suspend() {
assert!(matches!(
CallbackOutcome::from_invoke(&Err(CallbackError::Suspend("wait".into()))),
CallbackOutcome::Suspend(_)
));
}
#[tokio::test]
async fn suspend_records_metadata_without_journaling() {
let state = Arc::new(Mutex::new(ReplayState::new(CallbackJournal::new("code"))));
let (suspender, suspend_calls) = programmable("approve", Outcome::Suspend("wait".into()));
let suspend_wrapper = ReplayCallback::new(suspender, Arc::clone(&state));
let err = suspend_wrapper.invoke(json!({"id": 5})).await.unwrap_err();
assert!(matches!(err, CallbackError::Suspend(_)));
assert_eq!(suspend_calls.load(Ordering::SeqCst), 1);
let st = lock_state(&state);
let suspended = st.suspended().expect("suspension should be recorded");
assert_eq!(suspended.name, "approve");
assert_eq!(suspended.reason, "wait");
assert!(
st.build_journal("code").is_empty(),
"suspended callback must not be journaled"
);
}
#[tokio::test]
async fn callbacks_dispatched_after_suspension_are_rejected() {
let state = Arc::new(Mutex::new(ReplayState::new(CallbackJournal::new("code"))));
let (suspender, _) = programmable("approve", Outcome::Suspend("wait".into()));
let suspend_wrapper = ReplayCallback::new(suspender, Arc::clone(&state));
let _ = suspend_wrapper.invoke(json!({})).await;
let (fetch, fetch_calls) = programmable("fetch", Outcome::Ok(json!("live")));
let fetch_wrapper = ReplayCallback::new(fetch, Arc::clone(&state));
let err = fetch_wrapper.invoke(json!({})).await.unwrap_err();
assert!(matches!(err, CallbackError::Suspend(_)));
assert_eq!(
fetch_calls.load(Ordering::SeqCst),
0,
"callback dispatched after suspension must not run"
);
let st = lock_state(&state);
assert!(st.build_journal("code").is_empty());
}
#[test]
fn build_journal_truncates_entries_after_suspend_point() {
let mut st = ReplayState::new(CallbackJournal::new("code"));
let args = json!({});
let h = fnv1a_64(canonical_json(&args).as_bytes());
assert!(matches!(st.decide("a", h, "{}"), Decision::Miss { seq: 0 }));
assert!(matches!(st.decide("b", h, "{}"), Decision::Miss { seq: 1 }));
assert!(matches!(st.decide("c", h, "{}"), Decision::Miss { seq: 2 }));
st.record_live(entry_ok(0, "a", &args, json!("a-done")));
st.record_live(entry_ok(2, "c", &args, json!("c-done")));
st.record_suspend(
1,
SuspendedCallback {
name: "b".into(),
args_json: "{}".into(),
reason: "wait".into(),
},
);
let journal = st.build_journal("code");
assert_eq!(
journal.entries.len(),
1,
"only the pre-suspend prefix is kept"
);
assert_eq!(journal.entries[0].index, 0);
assert_eq!(journal.entries[0].name, "a");
}
#[tokio::test]
async fn resume_replays_prefix_and_reruns_suspended_call_live() {
let fetch_args = json!({"q": "x"});
let previous = CallbackJournal {
code: "code".into(),
entries: vec![entry_ok(0, "fetch", &fetch_args, json!("cached"))],
};
let state = Arc::new(Mutex::new(ReplayState::new(previous)));
let (fetch, fetch_calls) = programmable("fetch", Outcome::Ok(json!("LIVE")));
let fetch_wrapper = ReplayCallback::new(fetch, Arc::clone(&state));
let rf = fetch_wrapper.invoke(fetch_args.clone()).await.unwrap();
assert_eq!(rf, json!("cached"), "fetch replayed from cache");
assert_eq!(
fetch_calls.load(Ordering::SeqCst),
0,
"fetch not re-invoked"
);
let (approve, approve_calls) = programmable("approve", Outcome::Ok(json!("approved")));
let approve_wrapper = ReplayCallback::new(approve, Arc::clone(&state));
let ra = approve_wrapper.invoke(json!({})).await.unwrap();
assert_eq!(ra, json!("approved"));
assert_eq!(approve_calls.load(Ordering::SeqCst), 1, "approve runs live");
let st = lock_state(&state);
assert_eq!(st.replayed_count(), 1, "only fetch replayed");
assert!(st.suspended().is_none(), "no suspension on the resume run");
}
}