1use crate::diagnostics::prediction_interval::{ResidualInterval, ResidualIntervalConfig};
13use crate::diagnostics::training_summary::{TrainingSummary, TrainingSummaryConfig};
14use crate::diagnostics::warmup::{WarmupConfig, WarmupState, WarmupTracker};
15use crate::error::{RillError, ensure_finite};
16
17#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
21#[non_exhaustive]
22pub enum Confidence {
23 None,
25 Low,
27 Medium,
29 High,
31}
32
33impl Confidence {
34 pub const fn as_str(&self) -> &'static str {
38 match self {
39 Confidence::None => "none",
40 Confidence::Low => "low",
41 Confidence::Medium => "medium",
42 Confidence::High => "high",
43 }
44 }
45}
46
47#[derive(Debug, Clone)]
49#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
50pub struct PredictionReport {
51 prediction: f64,
52 lower_bound: Option<f64>,
53 upper_bound: Option<f64>,
54 confidence: Confidence,
55 samples_seen: u64,
56 recent_error: Option<f64>,
57 baseline_error: Option<f64>,
58 warmup_state: WarmupState,
59 beats_baseline: Option<bool>,
60}
61
62impl PredictionReport {
63 pub const fn prediction(&self) -> f64 {
65 self.prediction
66 }
67
68 pub const fn lower_bound(&self) -> Option<f64> {
70 self.lower_bound
71 }
72
73 pub const fn upper_bound(&self) -> Option<f64> {
75 self.upper_bound
76 }
77
78 pub const fn confidence(&self) -> Confidence {
80 self.confidence
81 }
82
83 pub const fn samples_seen(&self) -> u64 {
85 self.samples_seen
86 }
87
88 pub const fn recent_error(&self) -> Option<f64> {
90 self.recent_error
91 }
92
93 pub const fn baseline_error(&self) -> Option<f64> {
95 self.baseline_error
96 }
97
98 pub const fn warmup_state(&self) -> WarmupState {
100 self.warmup_state
101 }
102
103 pub const fn beats_baseline(&self) -> Option<bool> {
107 self.beats_baseline
108 }
109}
110
111#[derive(Debug, Clone)]
134#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
135pub struct PredictionReporter {
136 interval: ResidualInterval,
137 warmup: WarmupTracker,
138 summary: TrainingSummary,
139}
140
141impl PredictionReporter {
142 pub fn new(
147 interval_config: ResidualIntervalConfig,
148 warmup_config: WarmupConfig,
149 summary_config: TrainingSummaryConfig,
150 ) -> Result<Self, RillError> {
151 Ok(Self {
152 interval: ResidualInterval::new(interval_config)?,
153 warmup: WarmupTracker::new(warmup_config)?,
154 summary: TrainingSummary::new(summary_config)?,
155 })
156 }
157
158 pub fn observe(&mut self, prediction: f64, truth: f64) -> Result<(), RillError> {
163 self.interval.observe(prediction, truth)?;
164 let error = (truth - prediction).abs();
165 self.warmup.observe_sample(Some(error))?;
166 self.summary.record_error(error)?;
167 self.summary.record_sample()?;
168 Ok(())
169 }
170
171 pub fn set_baseline(&mut self, baseline: f64) -> Result<(), RillError> {
175 self.warmup.set_baseline(baseline)?;
176 self.summary.set_baseline_error(baseline)?;
177 Ok(())
178 }
179
180 pub fn report(&self, prediction: f64) -> Result<PredictionReport, RillError> {
186 ensure_finite("prediction", prediction)?;
187
188 let (lower_bound, upper_bound) = match self.interval.interval(prediction) {
189 Ok(iv) => (Some(iv.lower()), Some(iv.upper())),
190 Err(RillError::InsufficientData) => (None, None),
191 Err(e) => return Err(e),
192 };
193
194 let warmup_state = self.warmup.state();
195 let beats_baseline = self.summary.beats_baseline();
196 let samples_seen = self.summary.total_samples();
197 let recent_error = self.summary.recent_error();
198 let baseline_error = self.summary.baseline_error();
199
200 let confidence = match warmup_state {
201 WarmupState::NoData => Confidence::None,
202 WarmupState::WarmingUp | WarmupState::Degraded => Confidence::Low,
203 WarmupState::Usable => Confidence::Medium,
204 WarmupState::Stable => {
205 if matches!(beats_baseline, Some(true)) {
206 Confidence::High
207 } else {
208 Confidence::Medium
209 }
210 }
211 };
212
213 Ok(PredictionReport {
214 prediction,
215 lower_bound,
216 upper_bound,
217 confidence,
218 samples_seen,
219 recent_error,
220 baseline_error,
221 warmup_state,
222 beats_baseline,
223 })
224 }
225
226 pub fn summary(&self) -> &TrainingSummary {
228 &self.summary
229 }
230
231 pub fn warmup_state(&self) -> WarmupState {
233 self.warmup.state()
234 }
235
236 pub fn recent_error(&self) -> Option<f64> {
238 self.summary.recent_error()
239 }
240
241 pub fn reset(&mut self) {
243 self.interval.reset();
244 self.warmup.reset();
245 self.summary.reset();
246 }
247}
248
249impl Default for PredictionReporter {
250 fn default() -> Self {
251 Self::new(
252 ResidualIntervalConfig::default(),
253 WarmupConfig::default(),
254 TrainingSummaryConfig::default(),
255 )
256 .expect("default configs are valid")
257 }
258}
259
260#[cfg(test)]
261mod tests {
262 use super::*;
263
264 #[test]
265 fn confidence_as_str() {
266 assert_eq!(Confidence::None.as_str(), "none");
267 assert_eq!(Confidence::Low.as_str(), "low");
268 assert_eq!(Confidence::Medium.as_str(), "medium");
269 assert_eq!(Confidence::High.as_str(), "high");
270 }
271
272 #[test]
273 fn default_reporter_no_data() {
274 let reporter = PredictionReporter::default();
275 let r = reporter.report(0.0).unwrap();
276 assert_eq!(r.prediction(), 0.0);
277 assert_eq!(r.lower_bound(), None);
278 assert_eq!(r.upper_bound(), None);
279 assert_eq!(r.confidence(), Confidence::None);
280 assert_eq!(r.warmup_state(), WarmupState::NoData);
281 assert_eq!(r.samples_seen(), 0);
282 assert_eq!(r.recent_error(), None);
283 assert_eq!(r.baseline_error(), None);
284 assert_eq!(r.beats_baseline(), None);
285 }
286
287 #[test]
288 fn observe_then_report() {
289 let mut reporter = PredictionReporter::default();
290 reporter.observe(10.0, 11.0).unwrap(); reporter.observe(10.0, 9.0).unwrap(); let r = reporter.report(10.0).unwrap();
293 assert_eq!(r.prediction(), 10.0);
294 assert!(r.lower_bound().is_some());
295 assert!(r.upper_bound().is_some());
296 assert!(r.lower_bound().unwrap() < 10.0);
297 assert!(r.upper_bound().unwrap() > 10.0);
298 assert_eq!(r.samples_seen(), 2);
299 assert!(r.recent_error().is_some());
300 }
301
302 #[test]
303 fn set_baseline_enables_comparison() {
304 let mut reporter = PredictionReporter::default();
305 reporter.observe(0.0, 1.0).unwrap();
306 let r = reporter.report(0.0).unwrap();
307 assert_eq!(r.beats_baseline(), None);
308 assert_eq!(r.baseline_error(), None);
309 reporter.set_baseline(2.0).unwrap();
310 let r = reporter.report(0.0).unwrap();
311 assert_eq!(r.baseline_error(), Some(2.0));
312 assert_eq!(r.beats_baseline(), Some(true)); }
314
315 #[test]
316 fn confidence_progression() {
317 let warmup_config = WarmupConfig {
318 warming_up_threshold: 2,
319 usable_threshold: 5,
320 stable_threshold: 10,
321 degraded_error_ratio: 2.0,
322 };
323 let summary_config = TrainingSummaryConfig { error_alpha: 1.0 };
324 let mut reporter = PredictionReporter::new(
325 ResidualIntervalConfig::default(),
326 warmup_config,
327 summary_config,
328 )
329 .unwrap();
330
331 let r = reporter.report(0.0).unwrap();
333 assert_eq!(r.warmup_state(), WarmupState::NoData);
334 assert_eq!(r.confidence(), Confidence::None);
335
336 reporter.observe(0.0, 0.5).unwrap();
338 let r = reporter.report(0.0).unwrap();
339 assert_eq!(r.warmup_state(), WarmupState::WarmingUp);
340 assert_eq!(r.confidence(), Confidence::Low);
341
342 for _ in 0..4 {
344 reporter.observe(0.0, 0.5).unwrap();
345 }
346 let r = reporter.report(0.0).unwrap();
347 assert_eq!(r.warmup_state(), WarmupState::Usable);
348 assert_eq!(r.confidence(), Confidence::Medium);
349
350 reporter.set_baseline(1.0).unwrap();
352 let r = reporter.report(0.0).unwrap();
353 assert_eq!(r.warmup_state(), WarmupState::Usable);
354 assert_eq!(r.confidence(), Confidence::Medium);
355
356 for _ in 0..5 {
358 reporter.observe(0.0, 0.5).unwrap();
359 }
360 let r = reporter.report(0.0).unwrap();
361 assert_eq!(r.warmup_state(), WarmupState::Stable);
362 assert_eq!(r.confidence(), Confidence::High);
363 }
364
365 #[test]
366 fn degraded_state() {
367 let mut reporter = PredictionReporter::default();
368 reporter.set_baseline(0.4).unwrap();
369 for _ in 0..5 {
371 reporter.observe(0.0, 1.0).unwrap();
372 }
373 let r = reporter.report(0.0).unwrap();
374 assert_eq!(r.warmup_state(), WarmupState::Degraded);
375 assert_eq!(r.confidence(), Confidence::Low);
376 }
377
378 #[test]
379 fn reset_clears_all() {
380 let mut reporter = PredictionReporter::default();
381 reporter.observe(10.0, 12.0).unwrap();
382 reporter.set_baseline(2.0).unwrap();
383 reporter.reset();
384 let r = reporter.report(0.0).unwrap();
385 assert_eq!(r.lower_bound(), None);
386 assert_eq!(r.upper_bound(), None);
387 assert_eq!(r.confidence(), Confidence::None);
388 assert_eq!(r.warmup_state(), WarmupState::NoData);
389 assert_eq!(r.samples_seen(), 0);
390 assert_eq!(r.recent_error(), None);
391 assert_eq!(r.baseline_error(), None);
392 assert_eq!(r.beats_baseline(), None);
393 }
394
395 #[test]
396 fn report_with_non_finite_prediction_errors() {
397 let mut reporter = PredictionReporter::default();
398 reporter.observe(0.0, 1.0).unwrap();
399 assert!(reporter.report(f64::NAN).is_err());
400 assert!(reporter.report(f64::INFINITY).is_err());
401 assert!(reporter.report(f64::NEG_INFINITY).is_err());
402 }
403
404 #[test]
405 fn observe_with_non_finite_rejected() {
406 let mut reporter = PredictionReporter::default();
407 assert!(reporter.observe(0.0, f64::NAN).is_err());
408 assert!(reporter.observe(0.0, f64::INFINITY).is_err());
409 assert!(reporter.observe(f64::NAN, 0.0).is_err());
410 let r = reporter.report(0.0).unwrap();
412 assert_eq!(r.samples_seen(), 0);
413 assert_eq!(r.warmup_state(), WarmupState::NoData);
414 }
415
416 #[test]
417 fn samples_seen_tracked() {
418 let mut reporter = PredictionReporter::default();
419 for i in 0..10 {
420 reporter.observe(0.0, i as f64).unwrap();
421 }
422 let r = reporter.report(0.0).unwrap();
423 assert_eq!(r.samples_seen(), 10);
424 }
425
426 #[cfg(feature = "serde")]
427 #[test]
428 fn serde_roundtrip() {
429 let mut reporter = PredictionReporter::default();
430 reporter.observe(10.0, 12.0).unwrap();
431 reporter.observe(10.0, 9.0).unwrap();
432 reporter.set_baseline(3.0).unwrap();
433
434 let json = serde_json::to_string(&reporter).unwrap();
435 let restored: PredictionReporter = serde_json::from_str(&json).unwrap();
436
437 let r = restored.report(10.0).unwrap();
438 assert_eq!(r.samples_seen(), 2);
439 assert_eq!(r.baseline_error(), Some(3.0));
440 assert!(r.beats_baseline().is_some());
441 }
442}