1use 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#[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 pub baseline_reward: Option<f64>,
27 pub optimal_reward: Option<f64>,
29 pub drift: bool,
32}
33
34#[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#[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#[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 pub deterministic_replay_digest: String,
103}
104
105#[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 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 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}