Skip to main content

rill_ml/
replay.rs

1//! Deterministic delayed-decision replay harness.
2//!
3//! The harness is a development/verification tool. It does not execute
4//! product actions and does not infer reward semantics.
5
6use crate::bandit::{ContextualBandit, LinUcb};
7use crate::decision::{
8    DecisionId, DecisionLedger, DecisionLedgerConfig, DecisionLedgerError, DecisionOutcome,
9    PendingDecision, RegistrationStatus, apply_contextual_outcome,
10};
11use crate::{RillError, ValidateState};
12
13/// One recorded decision plus its optional delayed outcome.
14#[derive(Debug, Clone, PartialEq)]
15#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
16pub struct DecisionReplayRecord {
17    pub timestamp: u64,
18    pub decision_id: DecisionId,
19    pub context: Vec<f64>,
20    pub selected_arm: usize,
21    pub outcome_time: Option<u64>,
22    pub reward: Option<f64>,
23    pub generation: u64,
24    pub feature_schema_hash: String,
25    /// Optional counterfactual baseline reward supplied by the caller.
26    pub baseline_reward: Option<f64>,
27    /// Optional best-known reward used for an explainable regret estimate.
28    pub optimal_reward: Option<f64>,
29    /// Caller-recorded drift boundary. The harness resets the replay model
30    /// before selecting this decision and records the event.
31    pub drift: bool,
32}
33
34/// Bounded replay configuration.
35#[derive(Debug, Clone, PartialEq, Eq)]
36#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
37#[non_exhaustive]
38pub struct DecisionReplayConfig {
39    pub max_records: usize,
40    pub max_feedback_delay: u64,
41    pub model_generation: u64,
42    pub feature_schema_hash: String,
43    pub max_pending: usize,
44    pub max_completed: usize,
45}
46
47impl DecisionReplayConfig {
48    pub fn new(
49        max_records: usize,
50        max_feedback_delay: u64,
51        model_generation: u64,
52        feature_schema_hash: String,
53    ) -> Result<Self, RillError> {
54        if max_records == 0 || feature_schema_hash.is_empty() || feature_schema_hash.len() > 128 {
55            return Err(RillError::InvalidState(
56                "invalid decision replay configuration".into(),
57            ));
58        }
59        Ok(Self {
60            max_records,
61            max_feedback_delay,
62            model_generation,
63            feature_schema_hash,
64            max_pending: max_records,
65            max_completed: max_records,
66        })
67    }
68}
69
70/// Bounded feedback-latency summary.
71#[derive(Debug, Clone, Copy, PartialEq)]
72#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
73pub struct FeedbackLatencySummary {
74    pub count: usize,
75    pub min: u64,
76    pub max: u64,
77    pub mean: f64,
78    pub p95: u64,
79}
80
81/// Replay result. Regret is only accumulated for records that provide an
82/// `optimal_reward`; baseline comparison only uses records that provide a
83/// `baseline_reward`.
84#[derive(Debug, Clone, PartialEq)]
85#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
86pub struct DecisionReplayReport {
87    pub cumulative_reward: f64,
88    pub approximate_regret: f64,
89    pub baseline_cumulative_reward: f64,
90    pub reward_over_baseline: f64,
91    pub per_arm_pulls: Vec<u64>,
92    pub pending: usize,
93    pub completed: usize,
94    pub expired: usize,
95    pub missing_feedback: usize,
96    pub duplicate_feedback: usize,
97    pub rejected: usize,
98    pub generation_transitions: usize,
99    pub drift_events: usize,
100    pub feedback_latency: Option<FeedbackLatencySummary>,
101    /// Stable, non-cryptographic FNV-1a digest of accepted replay facts.
102    pub deterministic_replay_digest: String,
103}
104
105/// Stateful harness that can be serialized/restored between replay chunks.
106#[derive(Debug, Clone)]
107#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
108pub struct DecisionReplayHarness {
109    config: DecisionReplayConfig,
110    model: LinUcb,
111    ledger: DecisionLedger<Vec<f64>, usize>,
112    current_generation: u64,
113    cumulative_reward: f64,
114    approximate_regret: f64,
115    baseline_cumulative_reward: f64,
116    per_arm_pulls: Vec<u64>,
117    latencies: Vec<u64>,
118    expired: usize,
119    duplicate_feedback: usize,
120    rejected: usize,
121    generation_transitions: usize,
122    drift_events: usize,
123    missing_feedback: usize,
124    processed_records: usize,
125    digest: u64,
126}
127
128impl DecisionReplayHarness {
129    pub fn new(model: LinUcb, config: DecisionReplayConfig) -> Result<Self, RillError> {
130        if model.arm_count() == 0 || model.feature_count() == 0 {
131            return Err(RillError::InvalidState("invalid replay model".into()));
132        }
133        let ledger_config = DecisionLedgerConfig::new(config.max_pending, config.max_completed)
134            .map_err(ledger_error)?;
135        let arm_count = model.arm_count();
136        let current_generation = config.model_generation;
137        Ok(Self {
138            config,
139            model,
140            ledger: DecisionLedger::new(ledger_config).map_err(ledger_error)?,
141            current_generation,
142            cumulative_reward: 0.0,
143            approximate_regret: 0.0,
144            baseline_cumulative_reward: 0.0,
145            per_arm_pulls: vec![0; arm_count],
146            latencies: Vec::new(),
147            expired: 0,
148            duplicate_feedback: 0,
149            rejected: 0,
150            generation_transitions: 0,
151            drift_events: 0,
152            missing_feedback: 0,
153            processed_records: 0,
154            digest: 0xcbf29ce484222325,
155        })
156    }
157
158    pub fn model(&self) -> &LinUcb {
159        &self.model
160    }
161
162    pub fn ledger(&self) -> &DecisionLedger<Vec<f64>, usize> {
163        &self.ledger
164    }
165
166    /// Replay a bounded chunk. Decisions sort before outcomes at the same
167    /// timestamp; equal-type ties preserve input order.
168    pub fn replay(
169        &mut self,
170        records: &[DecisionReplayRecord],
171    ) -> Result<DecisionReplayReport, RillError> {
172        let processed_records = self
173            .processed_records
174            .checked_add(records.len())
175            .ok_or_else(|| RillError::InvalidState("replay record count overflow".into()))?;
176        if processed_records > self.config.max_records {
177            return Err(RillError::InvalidState(
178                "replay record capacity exceeded".into(),
179            ));
180        }
181        for record in records {
182            validate_record(record, self.model.feature_count(), self.model.arm_count())?;
183        }
184
185        // A replay chunk is transactional: any unexpected model/ledger or
186        // arithmetic failure leaves the caller's harness unchanged.
187        let mut next = self.clone();
188        next.processed_records = processed_records;
189        next.replay_in_place(records)?;
190        next.validate_state()?;
191        let report = next.report();
192        *self = next;
193        Ok(report)
194    }
195
196    fn replay_in_place(&mut self, records: &[DecisionReplayRecord]) -> Result<(), RillError> {
197        #[derive(Clone, Copy)]
198        enum EventKind {
199            Decision,
200            Outcome,
201        }
202        let mut events = Vec::with_capacity(records.len().saturating_mul(2));
203        for (index, record) in records.iter().enumerate() {
204            events.push((record.timestamp, 0_u8, index, EventKind::Decision));
205            if let Some(outcome_time) = record.outcome_time {
206                events.push((outcome_time, 1_u8, index, EventKind::Outcome));
207            }
208        }
209        events.sort_by_key(|event| (event.0, event.1, event.2));
210
211        for (_, _, index, kind) in events {
212            let record = &records[index];
213            match kind {
214                EventKind::Decision => self.replay_decision(record)?,
215                EventKind::Outcome => self.replay_outcome(record)?,
216            }
217        }
218        Ok(())
219    }
220
221    fn replay_decision(&mut self, record: &DecisionReplayRecord) -> Result<(), RillError> {
222        if record.feature_schema_hash != self.config.feature_schema_hash {
223            checked_increment_usize(&mut self.rejected, "rejected")?;
224            return Ok(());
225        }
226        if record.generation < self.current_generation {
227            checked_increment_usize(&mut self.rejected, "rejected")?;
228            return Ok(());
229        }
230        if record.generation > self.current_generation {
231            self.model.reset();
232            self.current_generation = record.generation;
233            checked_increment_usize(&mut self.generation_transitions, "generation transitions")?;
234        }
235        if record.drift {
236            self.model.reset();
237            checked_increment_usize(&mut self.drift_events, "drift events")?;
238        }
239
240        let samples_before = self.model.samples_seen();
241        let selected = self.model.select_deterministic(&record.context)?;
242        if self.model.samples_seen() != samples_before {
243            return Err(RillError::InvalidState(
244                "LinUCB select modified learning state during replay".into(),
245            ));
246        }
247        if selected != record.selected_arm {
248            checked_increment_usize(&mut self.rejected, "rejected")?;
249            return Ok(());
250        }
251        let expires_at = record
252            .timestamp
253            .checked_add(self.config.max_feedback_delay)
254            .ok_or_else(|| RillError::InvalidState("decision expiry overflow".into()))?;
255        let decision = PendingDecision::new(
256            record.decision_id,
257            record.context.clone(),
258            record.selected_arm,
259            record.timestamp,
260            expires_at,
261            record.generation,
262        );
263        let registered = match self.ledger.register(decision).map_err(ledger_error)?.status {
264            RegistrationStatus::Registered => {
265                self.digest_record(record, b'd');
266                true
267            }
268            RegistrationStatus::PendingReplay | RegistrationStatus::CompletedReplay => false,
269        };
270        if registered && record.outcome_time.is_none() {
271            self.missing_feedback = self.missing_feedback.checked_add(1).ok_or_else(|| {
272                RillError::InvalidState("missing feedback counter overflow".into())
273            })?;
274        }
275        Ok(())
276    }
277
278    fn replay_outcome(&mut self, record: &DecisionReplayRecord) -> Result<(), RillError> {
279        let Some(outcome_time) = record.outcome_time else {
280            return Ok(());
281        };
282        let Some(reward) = record.reward else {
283            checked_increment_usize(&mut self.rejected, "rejected")?;
284            return Ok(());
285        };
286        if record.generation != self.current_generation {
287            checked_increment_usize(&mut self.rejected, "rejected")?;
288            return Ok(());
289        }
290        if self.ledger.completed(record.decision_id).is_some() {
291            self.duplicate_feedback = self.duplicate_feedback.checked_add(1).ok_or_else(|| {
292                RillError::InvalidState("duplicate feedback counter overflow".into())
293            })?;
294            return Ok(());
295        }
296        let Some(pending) = self.ledger.pending(record.decision_id) else {
297            self.rejected = self
298                .rejected
299                .checked_add(1)
300                .ok_or_else(|| RillError::InvalidState("rejected counter overflow".into()))?;
301            return Ok(());
302        };
303        if outcome_time > pending.expires_at {
304            let removed = self.ledger.clear_expired(outcome_time).len();
305            self.expired = self
306                .expired
307                .checked_add(removed)
308                .ok_or_else(|| RillError::InvalidState("expired counter overflow".into()))?;
309            return Ok(());
310        }
311        if outcome_time < pending.created_at
312            || record.generation != pending.model_generation
313            || record.selected_arm != pending.action
314        {
315            self.rejected = self
316                .rejected
317                .checked_add(1)
318                .ok_or_else(|| RillError::InvalidState("rejected counter overflow".into()))?;
319            return Ok(());
320        }
321        let outcome = DecisionOutcome {
322            decision_id: record.decision_id,
323            action: record.selected_arm,
324            reward,
325            observed_at: outcome_time,
326            model_generation: record.generation,
327        };
328        let next_reward = finite_add(self.cumulative_reward, reward, "cumulative reward")?;
329        let next_baseline = if let Some(baseline) = record.baseline_reward {
330            finite_add(self.baseline_cumulative_reward, baseline, "baseline reward")?
331        } else {
332            self.baseline_cumulative_reward
333        };
334        let next_regret = if let Some(optimal) = record.optimal_reward {
335            let regret = (optimal - reward).max(0.0);
336            if !regret.is_finite() {
337                return Err(RillError::InvalidState("regret overflow".into()));
338            }
339            finite_add(self.approximate_regret, regret, "approximate regret")?
340        } else {
341            self.approximate_regret
342        };
343        let next_pulls = self.per_arm_pulls[record.selected_arm]
344            .checked_add(1)
345            .ok_or_else(|| RillError::InvalidState("per-arm pull overflow".into()))?;
346        if self.latencies.len() >= self.config.max_records {
347            return Err(RillError::InvalidState(
348                "feedback latency capacity exceeded".into(),
349            ));
350        }
351
352        apply_contextual_outcome(&mut self.ledger, &mut self.model, outcome)?;
353        self.cumulative_reward = next_reward;
354        self.baseline_cumulative_reward = next_baseline;
355        self.approximate_regret = next_regret;
356        self.per_arm_pulls[record.selected_arm] = next_pulls;
357        self.latencies.push(outcome_time - record.timestamp);
358        self.digest_record(record, b'o');
359        Ok(())
360    }
361
362    fn report(&self) -> DecisionReplayReport {
363        let reward_over_baseline = self.cumulative_reward - self.baseline_cumulative_reward;
364        DecisionReplayReport {
365            cumulative_reward: self.cumulative_reward,
366            approximate_regret: self.approximate_regret,
367            baseline_cumulative_reward: self.baseline_cumulative_reward,
368            reward_over_baseline,
369            per_arm_pulls: self.per_arm_pulls.clone(),
370            pending: self.ledger.pending_len(),
371            completed: self.ledger.completed_len(),
372            expired: self.expired,
373            missing_feedback: self.missing_feedback,
374            duplicate_feedback: self.duplicate_feedback,
375            rejected: self.rejected,
376            generation_transitions: self.generation_transitions,
377            drift_events: self.drift_events,
378            feedback_latency: latency_summary(&self.latencies),
379            deterministic_replay_digest: format!("{:016x}", self.digest),
380        }
381    }
382
383    fn digest_record(&mut self, record: &DecisionReplayRecord, marker: u8) {
384        fn mix(digest: &mut u64, bytes: &[u8]) {
385            for byte in bytes {
386                *digest ^= u64::from(*byte);
387                *digest = digest.wrapping_mul(0x100000001b3);
388            }
389        }
390        mix(&mut self.digest, &[marker]);
391        mix(&mut self.digest, &record.timestamp.to_le_bytes());
392        mix(&mut self.digest, &record.decision_id.0.to_le_bytes());
393        mix(
394            &mut self.digest,
395            &(record.selected_arm as u64).to_le_bytes(),
396        );
397        mix(&mut self.digest, &record.generation.to_le_bytes());
398        for value in &record.context {
399            mix(&mut self.digest, &value.to_bits().to_le_bytes());
400        }
401        if let Some(reward) = record.reward {
402            mix(&mut self.digest, &reward.to_bits().to_le_bytes());
403        }
404    }
405
406    #[cfg(feature = "serde")]
407    pub fn checkpoint_json(&self) -> Result<String, RillError> {
408        serde_json::to_string(self).map_err(|error| RillError::InvalidState(error.to_string()))
409    }
410
411    #[cfg(feature = "serde")]
412    pub fn restore_json(json: &str) -> Result<Self, RillError> {
413        let restored: Self = serde_json::from_str(json)
414            .map_err(|error| RillError::InvalidState(error.to_string()))?;
415        restored.validate_state()?;
416        Ok(restored)
417    }
418}
419
420impl ValidateState for DecisionReplayHarness {
421    fn validate_state(&self) -> Result<(), RillError> {
422        if self.config.max_records == 0
423            || self.config.max_pending == 0
424            || self.config.max_completed == 0
425            || self.config.feature_schema_hash.is_empty()
426            || self.config.feature_schema_hash.len() > 128
427            || self.current_generation < self.config.model_generation
428            || self.processed_records > self.config.max_records
429            || self.per_arm_pulls.len() != self.model.arm_count()
430            || self.latencies.len() > self.config.max_records
431            || !self.cumulative_reward.is_finite()
432            || !self.approximate_regret.is_finite()
433            || !self.baseline_cumulative_reward.is_finite()
434            || !(self.cumulative_reward - self.baseline_cumulative_reward).is_finite()
435        {
436            return Err(RillError::InvalidState("invalid replay state".into()));
437        }
438        let total_pulls = self.per_arm_pulls.iter().try_fold(0_u64, |sum, pulls| {
439            sum.checked_add(*pulls)
440                .ok_or_else(|| RillError::InvalidState("per-arm pull sum overflow".into()))
441        })?;
442        if total_pulls < self.model.samples_seen()
443            || total_pulls as usize != self.ledger.completed_len()
444        {
445            return Err(RillError::InvalidState(
446                "replay pulls, model samples, and completed ledger disagree".into(),
447            ));
448        }
449        self.model.validate()?;
450        self.ledger.validate()?;
451        Ok(())
452    }
453}
454
455fn validate_record(
456    record: &DecisionReplayRecord,
457    feature_count: usize,
458    arm_count: usize,
459) -> Result<(), RillError> {
460    if record.context.len() != feature_count {
461        return Err(RillError::DimensionMismatch {
462            expected: feature_count,
463            actual: record.context.len(),
464        });
465    }
466    if record.context.iter().any(|value| !value.is_finite()) {
467        return Err(RillError::InvalidState("non-finite replay context".into()));
468    }
469    if record.selected_arm >= arm_count {
470        return Err(RillError::InvalidArm {
471            expected: arm_count,
472            actual: record.selected_arm,
473        });
474    }
475    if record.outcome_time.is_some() != record.reward.is_some() {
476        return Err(RillError::InvalidState(
477            "outcome time and reward must both be present or absent".into(),
478        ));
479    }
480    if let Some(outcome_time) = record.outcome_time
481        && outcome_time < record.timestamp
482    {
483        return Err(RillError::InvalidState("outcome precedes decision".into()));
484    }
485    for value in [record.reward, record.baseline_reward, record.optimal_reward]
486        .into_iter()
487        .flatten()
488    {
489        if !value.is_finite() {
490            return Err(RillError::InvalidState("non-finite replay reward".into()));
491        }
492    }
493    Ok(())
494}
495
496fn latency_summary(latencies: &[u64]) -> Option<FeedbackLatencySummary> {
497    if latencies.is_empty() {
498        return None;
499    }
500    let mut sorted = latencies.to_vec();
501    sorted.sort_unstable();
502    let sum = sorted
503        .iter()
504        .fold(0_u128, |sum, value| sum + u128::from(*value));
505    let index = ((sorted.len() - 1) * 95).div_ceil(100);
506    Some(FeedbackLatencySummary {
507        count: sorted.len(),
508        min: sorted[0],
509        max: *sorted.last().unwrap(),
510        mean: sum as f64 / sorted.len() as f64,
511        p95: sorted[index],
512    })
513}
514
515fn ledger_error(error: DecisionLedgerError) -> RillError {
516    RillError::InvalidState(error.to_string())
517}
518
519fn finite_add(current: f64, value: f64, field: &'static str) -> Result<f64, RillError> {
520    let next = current + value;
521    if next.is_finite() {
522        Ok(next)
523    } else {
524        Err(RillError::InvalidState(format!("{field} overflow")))
525    }
526}
527
528fn checked_increment_usize(value: &mut usize, field: &'static str) -> Result<(), RillError> {
529    *value = value
530        .checked_add(1)
531        .ok_or_else(|| RillError::InvalidState(format!("{field} counter overflow")))?;
532    Ok(())
533}
534
535#[cfg(test)]
536mod tests {
537    use super::*;
538    use crate::bandit::LinUcbConfig;
539
540    fn harness() -> DecisionReplayHarness {
541        DecisionReplayHarness::new(
542            LinUcb::new(LinUcbConfig::default()).unwrap(),
543            DecisionReplayConfig::new(16, 10, 1, "schema-a".into()).unwrap(),
544        )
545        .unwrap()
546    }
547
548    fn record(id: u128, timestamp: u64, outcome: Option<(u64, f64)>) -> DecisionReplayRecord {
549        let model = LinUcb::new(LinUcbConfig::default()).unwrap();
550        let arm = model.select_deterministic(&[1.0]).unwrap();
551        DecisionReplayRecord {
552            timestamp,
553            decision_id: DecisionId(id),
554            context: vec![1.0],
555            selected_arm: arm,
556            outcome_time: outcome.map(|value| value.0),
557            reward: outcome.map(|value| value.1),
558            generation: 1,
559            feature_schema_hash: "schema-a".into(),
560            baseline_reward: Some(0.25),
561            optimal_reward: Some(1.0),
562            drift: false,
563        }
564    }
565
566    #[test]
567    fn update_waits_for_outcome_and_reports_metrics() {
568        let mut harness = harness();
569        let report = harness
570            .replay(&[record(1, 0, Some((5, 0.75))), record(2, 1, None)])
571            .unwrap();
572        assert_eq!(harness.model().samples_seen(), 1);
573        assert_eq!(report.cumulative_reward, 0.75);
574        assert_eq!(report.approximate_regret, 0.25);
575        assert_eq!(report.pending, 1);
576        assert_eq!(report.completed, 1);
577        assert_eq!(report.missing_feedback, 1);
578        assert_eq!(report.feedback_latency.unwrap().mean, 5.0);
579    }
580
581    #[test]
582    fn duplicate_expired_and_schema_mismatch_are_counted() {
583        let mut harness = harness();
584        let duplicate = record(1, 0, Some((5, 0.5)));
585        let expired = record(2, 1, Some((20, 0.5)));
586        let mut mismatch = record(3, 2, Some((3, 0.5)));
587        mismatch.feature_schema_hash = "other".into();
588        let report = harness
589            .replay(&[duplicate.clone(), duplicate, expired, mismatch])
590            .unwrap();
591        assert_eq!(report.duplicate_feedback, 1);
592        assert_eq!(report.expired, 1);
593        assert!(report.rejected >= 1);
594    }
595
596    #[test]
597    fn replay_digest_is_deterministic_and_restore_continues() {
598        let records = [record(1, 0, Some((2, 0.5)))];
599        let mut a = harness();
600        let mut b = harness();
601        let report_a = a.replay(&records).unwrap();
602        let report_b = b.replay(&records).unwrap();
603        assert_eq!(
604            report_a.deterministic_replay_digest,
605            report_b.deterministic_replay_digest
606        );
607
608        #[cfg(feature = "serde")]
609        {
610            let checkpoint = a.checkpoint_json().unwrap();
611            let restored = DecisionReplayHarness::restore_json(&checkpoint).unwrap();
612            assert_eq!(restored.model().samples_seen(), a.model().samples_seen());
613            assert_eq!(restored.ledger().completed_len(), 1);
614        }
615    }
616}