use std::collections::BTreeMap;
use typesayer_types::field::FieldValue;
use crate::prediction::Prediction;
#[derive(Debug, Clone)]
pub struct TraceEntry {
pub predictor_name: String,
pub inputs: BTreeMap<String, FieldValue>,
pub prediction: Prediction,
}
#[derive(Debug, Clone, Default)]
pub struct ExecutionTrace {
entries: Vec<TraceEntry>,
}
impl ExecutionTrace {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn push(&mut self, entry: TraceEntry) {
self.entries.push(entry);
}
#[must_use]
pub fn entries(&self) -> &[TraceEntry] {
&self.entries
}
#[must_use]
pub fn into_entries(self) -> Vec<TraceEntry> {
self.entries
}
#[must_use]
pub fn len(&self) -> usize {
self.entries.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn extend(&mut self, other: Self) {
self.entries.extend(other.entries);
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_entry(name: &str) -> TraceEntry {
TraceEntry {
predictor_name: name.to_owned(),
inputs: BTreeMap::from([("q".into(), FieldValue::Str("test".into()))]),
prediction: Prediction::new(
BTreeMap::from([("a".into(), FieldValue::Str("answer".into()))]),
None,
),
}
}
#[test]
fn empty_trace() {
let trace = ExecutionTrace::new();
assert!(trace.is_empty());
assert_eq!(trace.len(), 0);
assert!(trace.entries().is_empty());
}
#[test]
fn push_and_entries() {
let mut trace = ExecutionTrace::new();
trace.push(sample_entry("first"));
trace.push(sample_entry("second"));
assert_eq!(trace.len(), 2);
assert!(!trace.is_empty());
assert_eq!(trace.entries()[0].predictor_name, "first");
assert_eq!(trace.entries()[1].predictor_name, "second");
}
#[test]
fn into_entries_consumes() {
let mut trace = ExecutionTrace::new();
trace.push(sample_entry("qa"));
let entries = trace.into_entries();
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].predictor_name, "qa");
}
#[test]
fn extend_merges_traces() {
let mut trace1 = ExecutionTrace::new();
trace1.push(sample_entry("first"));
let mut trace2 = ExecutionTrace::new();
trace2.push(sample_entry("second"));
trace2.push(sample_entry("third"));
trace1.extend(trace2);
assert_eq!(trace1.len(), 3);
assert_eq!(trace1.entries()[0].predictor_name, "first");
assert_eq!(trace1.entries()[1].predictor_name, "second");
assert_eq!(trace1.entries()[2].predictor_name, "third");
}
#[test]
fn trace_entry_fields_accessible() {
let entry = sample_entry("qa");
assert_eq!(entry.predictor_name, "qa");
assert_eq!(entry.inputs["q"], FieldValue::Str("test".into()));
assert_eq!(entry.prediction.get::<String>("a").unwrap(), "answer");
}
}