use crate::error::{Error, Result};
use crate::model::{
BuildConfig, ConfidenceMode, MissingTimestampPolicy, Observation, Outcome, PriorAction,
PriorBook,
};
use crate::score::{
confidence, effective_sample_size, normalize, ratio, raw_score, shrink_toward,
time_decay_multiplier, wilson_lower_bound,
};
use serde::Serialize;
use std::cmp::Ordering;
use std::collections::HashMap;
#[derive(Debug, Default, Clone)]
struct ActionStats {
count: u64,
weighted_count: f64,
weighted_successes: f64,
weighted_trials: f64,
weighted_trial_weight_sq_sum: f64,
weighted_score_sum: f64,
weighted_score_count: f64,
}
fn wilson_confidence(stat: &ActionStats, z: f64) -> Option<f64> {
let n_eff = effective_sample_size(stat.weighted_trials, stat.weighted_trial_weight_sq_sum);
let p_hat = ratio(stat.weighted_successes, stat.weighted_trials)?;
wilson_lower_bound(p_hat * n_eff, n_eff, z)
}
fn action_confidence(stat: &ActionStats, config: &BuildConfig) -> f64 {
let heuristic_val = confidence(stat.weighted_count, config.confidence_k);
match config.confidence_mode {
ConfidenceMode::Heuristic => heuristic_val,
ConfidenceMode::WilsonLowerBound => {
wilson_confidence(stat, config.confidence_z).unwrap_or(heuristic_val)
}
ConfidenceMode::Hybrid => wilson_confidence(stat, config.confidence_z)
.map(|w| heuristic_val * w)
.unwrap_or(heuristic_val),
}
}
fn validate_config(config: &BuildConfig) -> Result<()> {
if let Some(half_life) = config.time_decay_half_life_days {
if !(half_life.is_finite() && half_life > 0.0) {
return Err(Error::InvalidConfig {
message: format!("time_decay_half_life_days must be > 0, got {half_life}"),
});
}
if config.time_decay_reference_unix_seconds.is_none() {
return Err(Error::InvalidConfig {
message: "time_decay_reference_unix_seconds is required when \
time_decay_half_life_days is set"
.to_string(),
});
}
}
for (name, weight) in &config.source_weights {
if !weight.is_finite() || *weight < 0.0 {
return Err(Error::InvalidConfig {
message: format!("source_weights[{name:?}] must be finite and >= 0, got {weight}"),
});
}
}
if !config.default_source_weight.is_finite() || config.default_source_weight < 0.0 {
return Err(Error::InvalidConfig {
message: format!(
"default_source_weight must be finite and >= 0, got {}",
config.default_source_weight
),
});
}
Ok(())
}
fn effective_weight(obs: &Observation, config: &BuildConfig) -> Option<f64> {
let mut weight = obs.weight;
if let Some(half_life_days) = config.time_decay_half_life_days {
let reference = config
.time_decay_reference_unix_seconds
.expect("validate_config requires this whenever time_decay_half_life_days is set");
match obs.observed_at_unix_seconds {
Some(observed_at) => {
let age_days = ((reference - observed_at) as f64 / 86_400.0).max(0.0);
weight *= time_decay_multiplier(age_days, half_life_days);
}
None => match config.missing_timestamp_policy {
MissingTimestampPolicy::KeepBaseWeight => {}
MissingTimestampPolicy::Drop => return None,
},
}
}
let source_multiplier = obs
.source
.as_deref()
.and_then(|name| config.source_weights.get(name))
.copied()
.unwrap_or(config.default_source_weight);
weight *= source_multiplier;
Some(weight)
}
#[derive(Debug, Clone, Copy, Default, Serialize)]
pub struct BuildStats {
pub observations_kept: u64,
pub observations_dropped_by_step_or_tag_filter: u64,
pub observations_dropped_by_missing_timestamp: u64,
pub candidates_before_filtering: u64,
pub candidates_dropped_by_min_count: u64,
pub candidates_dropped_by_min_weighted_count: u64,
pub candidates_dropped_by_min_confidence: u64,
pub candidates_dropped_by_max_actions_per_state: u64,
pub candidates_kept: u64,
}
pub(crate) struct PriorAccumulator<'a> {
stats: HashMap<String, HashMap<String, ActionStats>>,
global_weighted_successes: f64,
global_weighted_trials: f64,
global_weighted_score_sum: f64,
global_weighted_score_count: f64,
observations_kept: u64,
observations_dropped_by_step_or_tag_filter: u64,
observations_dropped_by_missing_timestamp: u64,
config: &'a BuildConfig,
}
impl<'a> PriorAccumulator<'a> {
pub(crate) fn new(config: &'a BuildConfig) -> Result<Self> {
validate_config(config)?;
Ok(Self {
stats: HashMap::new(),
global_weighted_successes: 0.0,
global_weighted_trials: 0.0,
global_weighted_score_sum: 0.0,
global_weighted_score_count: 0.0,
observations_kept: 0,
observations_dropped_by_step_or_tag_filter: 0,
observations_dropped_by_missing_timestamp: 0,
config,
})
}
pub(crate) fn observe(&mut self, obs: &Observation) {
if let Some(max_step) = self.config.max_step
&& obs.step > max_step
{
self.observations_dropped_by_step_or_tag_filter += 1;
return;
}
if let Some(required_tags) = &self.config.tag_filter
&& !required_tags.iter().any(|tag| obs.tags.contains(tag))
{
self.observations_dropped_by_step_or_tag_filter += 1;
return;
}
let Some(effective_weight) = effective_weight(obs, self.config) else {
self.observations_dropped_by_missing_timestamp += 1;
return;
};
self.observations_kept += 1;
let entry = self
.stats
.entry(obs.state.clone())
.or_default()
.entry(obs.action.clone())
.or_default();
entry.count += 1;
entry.weighted_count += effective_weight;
if obs.outcome != Outcome::Unknown {
entry.weighted_trials += effective_weight;
entry.weighted_trial_weight_sq_sum += effective_weight * effective_weight;
self.global_weighted_trials += effective_weight;
let success_credit = match obs.outcome {
Outcome::Success => 1.0,
Outcome::Draw => self.config.draw_value,
Outcome::Failure | Outcome::Unknown => 0.0,
};
entry.weighted_successes += success_credit * effective_weight;
self.global_weighted_successes += success_credit * effective_weight;
}
if let Some(score) = obs.score {
entry.weighted_score_sum += effective_weight * score;
entry.weighted_score_count += effective_weight;
self.global_weighted_score_sum += effective_weight * score;
self.global_weighted_score_count += effective_weight;
}
}
pub(crate) fn finish(self) -> PriorBook {
self.finish_with_stats().0
}
pub(crate) fn finish_with_stats(self) -> (PriorBook, BuildStats) {
let global_success_rate =
ratio(self.global_weighted_successes, self.global_weighted_trials);
let global_mean_score = ratio(
self.global_weighted_score_sum,
self.global_weighted_score_count,
);
let mut entries: HashMap<String, Vec<PriorAction>> = HashMap::new();
let mut candidates_before_filtering: u64 = 0;
let mut candidates_dropped_by_min_count: u64 = 0;
let mut candidates_dropped_by_min_weighted_count: u64 = 0;
let mut candidates_dropped_by_min_confidence: u64 = 0;
let mut candidates_dropped_by_max_actions_per_state: u64 = 0;
for (state, actions) in self.stats {
candidates_before_filtering += actions.len() as u64;
let mut kept: Vec<(String, ActionStats)> = Vec::new();
for (action, stat) in actions {
if stat.count < self.config.min_count {
candidates_dropped_by_min_count += 1;
continue;
}
if stat.weighted_count < self.config.min_weighted_count {
candidates_dropped_by_min_weighted_count += 1;
continue;
}
if action_confidence(&stat, self.config) < self.config.min_confidence {
candidates_dropped_by_min_confidence += 1;
continue;
}
kept.push((action, stat));
}
if kept.is_empty() {
continue;
}
let raw_scores: Vec<f64> = kept
.iter()
.map(|(_, stat)| {
let smoothed_success = global_success_rate.map(|global| {
shrink_toward(
stat.weighted_successes,
stat.weighted_trials,
self.config.smoothing_alpha,
global,
)
});
let smoothed_score = global_mean_score.map(|global| {
shrink_toward(
stat.weighted_score_sum,
stat.weighted_score_count,
self.config.smoothing_alpha,
global,
)
});
raw_score(
stat.weighted_count,
smoothed_success,
smoothed_score,
self.config,
)
})
.collect();
let priors = normalize(&raw_scores);
let mut actions_out: Vec<PriorAction> = kept
.into_iter()
.zip(priors)
.map(|((action, stat), prior)| PriorAction {
action,
count: stat.count,
weighted_count: stat.weighted_count,
success_rate: ratio(stat.weighted_successes, stat.weighted_trials),
mean_score: ratio(stat.weighted_score_sum, stat.weighted_score_count),
prior,
confidence: action_confidence(&stat, self.config),
})
.collect();
if let Some(max_actions) = self.config.max_actions_per_state {
actions_out.sort_by(|a, b| {
b.prior
.partial_cmp(&a.prior)
.unwrap_or(Ordering::Equal)
.then_with(|| a.action.cmp(&b.action))
});
if actions_out.len() > max_actions {
candidates_dropped_by_max_actions_per_state +=
(actions_out.len() - max_actions) as u64;
}
actions_out.truncate(max_actions);
}
entries.insert(state, actions_out);
}
let candidates_kept: u64 = entries.values().map(|v| v.len() as u64).sum();
let stats = BuildStats {
observations_kept: self.observations_kept,
observations_dropped_by_step_or_tag_filter: self
.observations_dropped_by_step_or_tag_filter,
observations_dropped_by_missing_timestamp: self
.observations_dropped_by_missing_timestamp,
candidates_before_filtering,
candidates_dropped_by_min_count,
candidates_dropped_by_min_weighted_count,
candidates_dropped_by_min_confidence,
candidates_dropped_by_max_actions_per_state,
candidates_kept,
};
(PriorBook { entries }, stats)
}
}
pub fn build_prior_book(observations: &[Observation], config: &BuildConfig) -> Result<PriorBook> {
if observations.is_empty() {
return Err(Error::NoObservations);
}
let mut acc = PriorAccumulator::new(config)?;
for obs in observations {
acc.observe(obs);
}
let book = acc.finish();
if book.entries.is_empty() {
return Err(Error::NoObservations);
}
Ok(book)
}
#[cfg(test)]
mod tests {
use super::*;
#[allow(clippy::too_many_arguments)]
fn obs(
sequence_id: &str,
step: u32,
state: &str,
action: &str,
outcome: Outcome,
score: Option<f64>,
weight: f64,
tags: Vec<&str>,
) -> Observation {
Observation {
sequence_id: sequence_id.to_string(),
step,
state: state.to_string(),
action: action.to_string(),
outcome,
score,
weight,
tags: tags.into_iter().map(str::to_string).collect(),
observed_at_unix_seconds: None,
source: None,
}
}
fn action<'a>(book: &'a PriorBook, state: &str, action: &str) -> &'a PriorAction {
book.entries[state]
.iter()
.find(|a| a.action == action)
.unwrap()
}
#[test]
fn aggregates_counts_and_weighted_counts() {
let observations = vec![
obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![]),
obs("c2", 0, "s", "a", Outcome::Unknown, None, 2.5, vec![]),
];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
let a = action(&book, "s", "a");
assert_eq!(a.count, 2);
assert_eq!(a.weighted_count, 3.5);
}
#[test]
fn computes_success_rate_and_mean_score() {
let observations = vec![
obs("c1", 0, "s", "a", Outcome::Success, Some(1.0), 1.0, vec![]),
obs("c2", 0, "s", "a", Outcome::Failure, Some(0.0), 1.0, vec![]),
];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
let a = action(&book, "s", "a");
assert_eq!(a.success_rate, Some(0.5));
assert_eq!(a.mean_score, Some(0.5));
}
#[test]
fn applies_min_count_filter() {
let observations = vec![
obs("c1", 0, "s", "rare", Outcome::Unknown, None, 1.0, vec![]),
obs("c2", 0, "s", "common", Outcome::Unknown, None, 1.0, vec![]),
obs("c3", 0, "s", "common", Outcome::Unknown, None, 1.0, vec![]),
];
let config = BuildConfig {
min_count: 2,
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
assert!(book.entries["s"].iter().all(|a| a.action != "rare"));
assert!(book.entries["s"].iter().any(|a| a.action == "common"));
}
#[test]
fn applies_min_weighted_count_filter_independent_of_raw_count() {
let observations = vec![
obs(
"c1",
0,
"s",
"many_tiny",
Outcome::Unknown,
None,
0.01,
vec![],
),
obs(
"c2",
0,
"s",
"many_tiny",
Outcome::Unknown,
None,
0.01,
vec![],
),
obs(
"c3",
0,
"s",
"few_heavy",
Outcome::Unknown,
None,
10.0,
vec![],
),
];
let config = BuildConfig {
min_weighted_count: 1.0,
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
assert!(book.entries["s"].iter().all(|a| a.action != "many_tiny"));
assert!(book.entries["s"].iter().any(|a| a.action == "few_heavy"));
}
#[test]
fn applies_min_confidence_filter() {
let observations = vec![
obs(
"c1",
0,
"s",
"unproven",
Outcome::Unknown,
None,
1.0,
vec![],
),
obs(
"c2",
0,
"s",
"proven",
Outcome::Unknown,
None,
100.0,
vec![],
),
];
let config = BuildConfig {
min_confidence: 0.5,
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
assert!(book.entries["s"].iter().all(|a| a.action != "unproven"));
assert!(book.entries["s"].iter().any(|a| a.action == "proven"));
}
#[test]
fn confidence_mode_heuristic_matches_pre_existing_formula() {
let observations = vec![
obs("c1", 0, "s", "a", Outcome::Success, None, 1.0, vec![]),
obs("c2", 0, "s", "a", Outcome::Failure, None, 1.0, vec![]),
];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
let got = action(&book, "s", "a").confidence;
assert_eq!(got, confidence(2.0, 20.0)); }
#[test]
fn confidence_mode_wilson_lower_bound_ranks_by_success_rate_at_equal_count() {
let mut observations = Vec::new();
for i in 0..18 {
observations.push(obs(
&format!("good_s_{i}"),
0,
"s",
"good",
Outcome::Success,
None,
1.0,
vec![],
));
}
for i in 0..2 {
observations.push(obs(
&format!("good_f_{i}"),
0,
"s",
"good",
Outcome::Failure,
None,
1.0,
vec![],
));
}
for i in 0..2 {
observations.push(obs(
&format!("bad_s_{i}"),
0,
"s",
"bad",
Outcome::Success,
None,
1.0,
vec![],
));
}
for i in 0..18 {
observations.push(obs(
&format!("bad_f_{i}"),
0,
"s",
"bad",
Outcome::Failure,
None,
1.0,
vec![],
));
}
let heuristic_book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
assert_eq!(
action(&heuristic_book, "s", "good").confidence,
action(&heuristic_book, "s", "bad").confidence
);
let wilson_config = BuildConfig {
confidence_mode: ConfidenceMode::WilsonLowerBound,
..Default::default()
};
let wilson_book = build_prior_book(&observations, &wilson_config).unwrap();
let good = action(&wilson_book, "s", "good").confidence;
let bad = action(&wilson_book, "s", "bad").confidence;
assert!(
good > bad,
"good={good} should outrank bad={bad} under WilsonLowerBound"
);
}
#[test]
fn confidence_mode_hybrid_multiplies_heuristic_and_wilson() {
let mut observations = Vec::new();
for i in 0..18 {
observations.push(obs(
&format!("s{i}"),
0,
"s",
"a",
Outcome::Success,
None,
1.0,
vec![],
));
}
for i in 0..2 {
observations.push(obs(
&format!("f{i}"),
0,
"s",
"a",
Outcome::Failure,
None,
1.0,
vec![],
));
}
let config = BuildConfig {
confidence_mode: ConfidenceMode::Hybrid,
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
let got = action(&book, "s", "a").confidence;
let heuristic_val = confidence(20.0, config.confidence_k);
let wilson_val = wilson_lower_bound(18.0, 20.0, config.confidence_z).unwrap();
assert!((got - heuristic_val * wilson_val).abs() < 1e-9);
}
#[test]
fn min_confidence_filter_behavior_depends_on_confidence_mode() {
let mut observations = Vec::new();
for i in 0..2 {
observations.push(obs(
&format!("rs{i}"),
0,
"s",
"risky",
Outcome::Success,
None,
1.0,
vec![],
));
}
for i in 0..18 {
observations.push(obs(
&format!("rf{i}"),
0,
"s",
"risky",
Outcome::Failure,
None,
1.0,
vec![],
));
}
for i in 0..18 {
observations.push(obs(
&format!("ss{i}"),
0,
"s",
"safe",
Outcome::Success,
None,
1.0,
vec![],
));
}
for i in 0..2 {
observations.push(obs(
&format!("sf{i}"),
0,
"s",
"safe",
Outcome::Failure,
None,
1.0,
vec![],
));
}
let heuristic_config = BuildConfig {
min_confidence: 0.5,
..Default::default()
};
let heuristic_book = build_prior_book(&observations, &heuristic_config).unwrap();
assert!(
heuristic_book.entries["s"]
.iter()
.any(|a| a.action == "risky"),
"heuristic min_confidence is blind to outcome, so risky should survive"
);
let wilson_config = BuildConfig {
min_confidence: 0.5,
confidence_mode: ConfidenceMode::WilsonLowerBound,
..Default::default()
};
let wilson_book = build_prior_book(&observations, &wilson_config).unwrap();
assert!(
wilson_book.entries["s"].iter().all(|a| a.action != "risky"),
"wilson-lower-bound min_confidence should drop a mostly-failing action at the same threshold"
);
assert!(wilson_book.entries["s"].iter().any(|a| a.action == "safe"));
}
#[test]
fn draw_value_contributes_fractional_success_under_wilson_mode() {
let mut observations = Vec::new();
for i in 0..18 {
observations.push(obs(
&format!("d{i}"),
0,
"s",
"a",
Outcome::Draw,
None,
1.0,
vec![],
));
}
for i in 0..2 {
observations.push(obs(
&format!("f{i}"),
0,
"s",
"a",
Outcome::Failure,
None,
1.0,
vec![],
));
}
let config = BuildConfig {
confidence_mode: ConfidenceMode::WilsonLowerBound,
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
let got = action(&book, "s", "a").confidence;
let expected = wilson_lower_bound(9.0, 20.0, config.confidence_z).unwrap();
assert!((got - expected).abs() < 1e-9);
}
#[test]
fn action_confidence_does_not_produce_nan_when_draw_value_exceeds_one() {
let observations = vec![
obs("d1", 0, "s", "a", Outcome::Draw, None, 1.0, vec![]),
obs("d2", 0, "s", "a", Outcome::Draw, None, 1.0, vec![]),
];
let config = BuildConfig {
draw_value: 1.5, confidence_mode: ConfidenceMode::WilsonLowerBound,
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
let got = action(&book, "s", "a").confidence;
assert!(got.is_finite());
assert!((0.0..=1.0).contains(&got));
}
#[test]
fn draw_earns_partial_success_credit_via_draw_value() {
let observations = vec![
obs("c1", 0, "s", "a", Outcome::Success, None, 1.0, vec![]),
obs("c2", 0, "s", "a", Outcome::Draw, None, 1.0, vec![]),
];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
assert_eq!(action(&book, "s", "a").success_rate, Some(0.75));
}
#[test]
fn draw_value_zero_reproduces_draw_as_failure_behavior() {
let observations = vec![
obs("c1", 0, "s", "a", Outcome::Success, None, 1.0, vec![]),
obs("c2", 0, "s", "a", Outcome::Draw, None, 1.0, vec![]),
];
let config = BuildConfig {
draw_value: 0.0,
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
assert_eq!(action(&book, "s", "a").success_rate, Some(0.5));
}
#[test]
fn draw_credit_also_shifts_the_shared_global_smoothing_rate() {
let observations = vec![
obs("c1", 0, "s", "steady", Outcome::Success, None, 1.0, vec![]),
obs("c2", 0, "s", "steady", Outcome::Failure, None, 1.0, vec![]),
obs("c3", 0, "s", "drawer", Outcome::Draw, None, 1.0, vec![]),
];
let base = BuildConfig {
count_weight: 0.0,
score_weight: 0.0,
success_weight: 1.0,
smoothing_alpha: 5.0,
..Default::default()
};
let credited_as_failure = build_prior_book(
&observations,
&BuildConfig {
draw_value: 0.0,
..base.clone()
},
)
.unwrap();
let credited_as_partial = build_prior_book(
&observations,
&BuildConfig {
draw_value: 0.5,
..base
},
)
.unwrap();
let steady_prior_failure = action(&credited_as_failure, "s", "steady").prior;
let steady_prior_partial = action(&credited_as_partial, "s", "steady").prior;
assert!((steady_prior_failure - 48.0 / 83.0).abs() < 1e-6);
assert!((steady_prior_partial - 0.5).abs() < 1e-6);
assert!((steady_prior_failure - steady_prior_partial).abs() > 0.05);
}
#[test]
fn smoothing_pulls_a_lone_success_toward_the_global_rate() {
let observations = vec![
obs("c1", 0, "s", "lucky", Outcome::Success, None, 1.0, vec![]),
obs("c2", 0, "s", "steady", Outcome::Success, None, 1.0, vec![]),
obs("c3", 0, "s", "steady", Outcome::Failure, None, 1.0, vec![]),
obs("c4", 0, "s", "steady", Outcome::Success, None, 1.0, vec![]),
obs("c5", 0, "s", "steady", Outcome::Failure, None, 1.0, vec![]),
];
let base = BuildConfig {
count_weight: 0.0,
score_weight: 0.0,
success_weight: 1.0,
min_count: 1,
..Default::default()
};
let unsmoothed = build_prior_book(
&observations,
&BuildConfig {
smoothing_alpha: 0.0,
..base.clone()
},
)
.unwrap();
let heavily_smoothed = build_prior_book(
&observations,
&BuildConfig {
smoothing_alpha: 50.0,
..base
},
)
.unwrap();
let lucky_unsmoothed = action(&unsmoothed, "s", "lucky").prior;
let lucky_smoothed = action(&heavily_smoothed, "s", "lucky").prior;
assert!(
lucky_smoothed < lucky_unsmoothed,
"smoothing should temper a lone perfect record: {lucky_smoothed} should be < {lucky_unsmoothed}"
);
}
#[test]
fn normalizes_priors_to_sum_to_one_per_state() {
let observations = vec![
obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![]),
obs("c2", 0, "s", "b", Outcome::Unknown, None, 3.0, vec![]),
];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
let sum: f64 = book.entries["s"].iter().map(|a| a.prior).sum();
assert!((sum - 1.0).abs() < 1e-9);
}
#[test]
fn deterministic_output_ordering_is_stable_across_runs() {
let observations = vec![
obs("c1", 0, "s2", "a", Outcome::Unknown, None, 1.0, vec![]),
obs("c2", 0, "s1", "b", Outcome::Unknown, None, 5.0, vec![]),
obs("c3", 0, "s1", "a", Outcome::Unknown, None, 1.0, vec![]),
];
let config = BuildConfig::default();
let first = build_prior_book(&observations, &config).unwrap();
let second = build_prior_book(&observations, &config).unwrap();
let first_json = serde_json::to_string(&first.entries_sorted()).unwrap();
let second_json = serde_json::to_string(&second.entries_sorted()).unwrap();
assert_eq!(first_json, second_json);
let sorted = first.entries_sorted();
assert_eq!(sorted[0].state, "s1");
assert_eq!(sorted[1].state, "s2");
assert_eq!(sorted[0].actions[0].action, "b");
}
#[test]
fn query_unseen_state_returns_no_candidates() {
let observations = vec![obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
assert!(book.query("nonexistent", None).is_empty());
}
#[test]
fn query_known_state_returns_candidates() {
let observations = vec![obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
assert_eq!(book.query("s", None).len(), 1);
}
#[test]
fn empty_input_is_an_error() {
let err = build_prior_book(&[], &BuildConfig::default()).unwrap_err();
assert!(matches!(err, Error::NoObservations));
}
#[test]
fn all_unknown_outcomes_drops_success_rate_entirely() {
let observations = vec![
obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![]),
obs("c2", 0, "s", "b", Outcome::Unknown, None, 1.0, vec![]),
];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
assert!(book.entries["s"].iter().all(|a| a.success_rate.is_none()));
}
#[test]
fn all_failures_reports_zero_success_rate_not_none() {
let observations = vec![obs("c1", 0, "s", "a", Outcome::Failure, None, 1.0, vec![])];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
assert_eq!(action(&book, "s", "a").success_rate, Some(0.0));
}
#[test]
fn all_successes_reports_full_success_rate() {
let observations = vec![obs("c1", 0, "s", "a", Outcome::Success, None, 1.0, vec![])];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
assert_eq!(action(&book, "s", "a").success_rate, Some(1.0));
}
#[test]
fn one_observation_builds_successfully() {
let observations = vec![obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
assert_eq!(action(&book, "s", "a").count, 1);
}
#[test]
fn extremely_large_counts_do_not_overflow_or_panic() {
let observations: Vec<Observation> = (0..50_000)
.map(|i| {
obs(
&format!("c{i}"),
0,
"s",
"a",
Outcome::Success,
None,
1.0,
vec![],
)
})
.collect();
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
let a = action(&book, "s", "a");
assert_eq!(a.count, 50_000);
assert!(a.confidence > 0.999);
}
#[test]
fn duplicate_sequence_ids_are_all_counted() {
let observations = vec![
obs("dup", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![]),
obs("dup", 1, "s", "a", Outcome::Unknown, None, 1.0, vec![]),
];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
assert_eq!(action(&book, "s", "a").count, 2);
}
#[test]
fn multiple_actions_per_state_all_appear() {
let observations = vec![
obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![]),
obs("c2", 0, "s", "b", Outcome::Unknown, None, 1.0, vec![]),
obs("c3", 0, "s", "c", Outcome::Unknown, None, 1.0, vec![]),
];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
assert_eq!(book.entries["s"].len(), 3);
}
#[test]
fn max_step_filters_out_later_steps() {
let observations = vec![
obs("c1", 0, "s", "early", Outcome::Unknown, None, 1.0, vec![]),
obs("c2", 99, "s", "late", Outcome::Unknown, None, 1.0, vec![]),
];
let config = BuildConfig {
max_step: Some(10),
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
assert!(book.entries["s"].iter().all(|a| a.action != "late"));
}
#[test]
fn tag_filter_keeps_only_matching_observations() {
let observations = vec![
obs(
"c1",
0,
"s",
"trusted",
Outcome::Unknown,
None,
1.0,
vec!["trusted"],
),
obs(
"c2",
0,
"s",
"untrusted",
Outcome::Unknown,
None,
1.0,
vec![],
),
];
let config = BuildConfig {
tag_filter: Some(vec!["trusted".to_string()]),
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
assert!(book.entries["s"].iter().all(|a| a.action != "untrusted"));
}
#[test]
fn no_observations_survive_filtering_is_an_error() {
let observations = vec![obs("c1", 99, "s", "a", Outcome::Unknown, None, 1.0, vec![])];
let config = BuildConfig {
max_step: Some(1),
..Default::default()
};
let err = build_prior_book(&observations, &config).unwrap_err();
assert!(matches!(err, Error::NoObservations));
}
#[test]
fn build_stats_invariant_holds_and_every_drop_bucket_is_exact() {
let mut observations = vec![
obs(
"late",
99,
"other_state",
"whatever",
Outcome::Unknown,
None,
1.0,
vec![],
),
obs(
"r1",
0,
"s",
"too_rare",
Outcome::Unknown,
None,
1.0,
vec![],
),
];
for i in 0..3 {
observations.push(obs(
&format!("l{i}"),
0,
"s",
"too_light",
Outcome::Unknown,
None,
0.5,
vec![],
));
}
for i in 0..3 {
observations.push(obs(
&format!("u{i}"),
0,
"s",
"unproven",
Outcome::Unknown,
None,
2.0,
vec![],
));
}
for name in ["surv_a", "surv_b", "surv_c"] {
for i in 0..10 {
observations.push(obs(
&format!("{name}_{i}"),
0,
"s",
name,
Outcome::Unknown,
None,
10.0,
vec![],
));
}
}
let config = BuildConfig {
min_count: 2,
min_weighted_count: 5.0,
min_confidence: 0.3,
max_step: Some(10),
max_actions_per_state: Some(2),
..Default::default()
};
let mut acc = PriorAccumulator::new(&config).unwrap();
for o in &observations {
acc.observe(o);
}
let (book, stats) = acc.finish_with_stats();
assert_eq!(stats.observations_dropped_by_step_or_tag_filter, 1);
assert_eq!(stats.observations_kept, observations.len() as u64 - 1);
assert_eq!(stats.candidates_before_filtering, 6);
assert_eq!(stats.candidates_dropped_by_min_count, 1);
assert_eq!(stats.candidates_dropped_by_min_weighted_count, 1);
assert_eq!(stats.candidates_dropped_by_min_confidence, 1);
assert_eq!(stats.candidates_dropped_by_max_actions_per_state, 1);
assert_eq!(stats.candidates_kept, 2);
assert_eq!(book.entries["s"].len(), 2);
assert_eq!(
stats.candidates_before_filtering,
stats.candidates_kept
+ stats.candidates_dropped_by_min_count
+ stats.candidates_dropped_by_min_weighted_count
+ stats.candidates_dropped_by_min_confidence
+ stats.candidates_dropped_by_max_actions_per_state
);
}
fn decay_config(half_life_days: f64, reference_unix_seconds: i64) -> BuildConfig {
BuildConfig {
time_decay_half_life_days: Some(half_life_days),
time_decay_reference_unix_seconds: Some(reference_unix_seconds),
..Default::default()
}
}
#[test]
fn time_decay_disabled_by_default_matches_pre_existing_weight() {
let observations = vec![obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])];
let book = build_prior_book(&observations, &BuildConfig::default()).unwrap();
assert_eq!(action(&book, "s", "a").weighted_count, 1.0);
}
#[test]
fn observation_one_half_life_old_has_half_weight() {
let reference = 1_000_000;
let observed_at = reference - 10 * 86_400; let observations = vec![Observation {
observed_at_unix_seconds: Some(observed_at),
..obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])
}];
let config = decay_config(10.0, reference);
let book = build_prior_book(&observations, &config).unwrap();
assert!((action(&book, "s", "a").weighted_count - 0.5).abs() < 1e-9);
}
#[test]
fn observation_two_half_lives_old_has_quarter_weight() {
let reference = 1_000_000;
let observed_at = reference - 20 * 86_400; let observations = vec![Observation {
observed_at_unix_seconds: Some(observed_at),
..obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])
}];
let config = decay_config(10.0, reference);
let book = build_prior_book(&observations, &config).unwrap();
assert!((action(&book, "s", "a").weighted_count - 0.25).abs() < 1e-9);
}
#[test]
fn future_timestamp_clamps_to_full_weight() {
let reference = 1_000_000;
let observed_at = reference + 10 * 86_400; let observations = vec![Observation {
observed_at_unix_seconds: Some(observed_at),
..obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])
}];
let config = decay_config(10.0, reference);
let book = build_prior_book(&observations, &config).unwrap();
assert_eq!(action(&book, "s", "a").weighted_count, 1.0);
}
#[test]
fn missing_timestamp_keep_base_weight_keeps_observation_at_full_weight() {
let observations = vec![obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])];
let config = decay_config(10.0, 1_000_000);
let mut acc = PriorAccumulator::new(&config).unwrap();
for o in &observations {
acc.observe(o);
}
let (book, stats) = acc.finish_with_stats();
assert_eq!(stats.observations_kept, 1);
assert_eq!(stats.observations_dropped_by_missing_timestamp, 0);
assert_eq!(action(&book, "s", "a").weighted_count, 1.0);
}
#[test]
fn missing_timestamp_drop_excludes_observation() {
let observations = vec![
obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![]), Observation {
observed_at_unix_seconds: Some(1_000_000),
..obs("c2", 0, "s", "b", Outcome::Unknown, None, 1.0, vec![])
},
];
let config = BuildConfig {
missing_timestamp_policy: MissingTimestampPolicy::Drop,
..decay_config(10.0, 1_000_000)
};
let mut acc = PriorAccumulator::new(&config).unwrap();
for o in &observations {
acc.observe(o);
}
let (book, stats) = acc.finish_with_stats();
assert_eq!(stats.observations_dropped_by_missing_timestamp, 1);
assert_eq!(stats.observations_kept, 1);
assert!(book.entries["s"].iter().all(|a| a.action != "a"));
assert!(book.entries["s"].iter().any(|a| a.action == "b"));
}
#[test]
fn source_weight_applies_only_to_matching_source() {
let observations = vec![
Observation {
source: Some("human".to_string()),
..obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])
},
Observation {
source: Some("engine".to_string()),
..obs("c2", 0, "s", "b", Outcome::Unknown, None, 1.0, vec![])
},
];
let config = BuildConfig {
source_weights: std::collections::BTreeMap::from([("human".to_string(), 0.5)]),
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
assert_eq!(action(&book, "s", "a").weighted_count, 0.5);
assert_eq!(action(&book, "s", "b").weighted_count, 1.0);
}
#[test]
fn unknown_source_uses_default_source_weight() {
let observations = vec![Observation {
source: Some("never_configured".to_string()),
..obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])
}];
let config = BuildConfig {
source_weights: std::collections::BTreeMap::from([("human".to_string(), 0.5)]),
default_source_weight: 0.25,
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
assert_eq!(action(&book, "s", "a").weighted_count, 0.25);
}
#[test]
fn source_weight_zero_makes_action_invisible_under_min_weighted_count() {
let observations = vec![
Observation {
source: Some("bad".to_string()),
..obs(
"c1",
0,
"s",
"bad_action",
Outcome::Unknown,
None,
1.0,
vec![],
)
},
obs("c2", 0, "s", "keep", Outcome::Unknown, None, 1.0, vec![]),
];
let config = BuildConfig {
source_weights: std::collections::BTreeMap::from([("bad".to_string(), 0.0)]),
min_weighted_count: 0.5,
..Default::default()
};
let book = build_prior_book(&observations, &config).unwrap();
assert!(book.entries["s"].iter().all(|a| a.action != "bad_action"));
assert!(book.entries["s"].iter().any(|a| a.action == "keep"));
}
fn fresh_and_stale_observations(reference: i64, half_life_days: f64) -> Vec<Observation> {
let stale_age_days = 10.0 * half_life_days;
let mut observations = Vec::new();
for i in 0..20 {
observations.push(Observation {
observed_at_unix_seconds: Some(reference),
..obs(
&format!("fresh_{i}"),
0,
"s",
"fresh",
Outcome::Success,
None,
1.0,
vec![],
)
});
}
for i in 0..20 {
observations.push(Observation {
observed_at_unix_seconds: Some(reference - (stale_age_days * 86_400.0) as i64),
..obs(
&format!("stale_{i}"),
0,
"s",
"stale",
Outcome::Success,
None,
1.0,
vec![],
)
});
}
observations
}
#[test]
fn wilson_confidence_is_invariant_to_uniform_time_decay_but_ranking_still_reflects_it() {
let reference = 1_000_000;
let half_life_days = 10.0;
let observations = fresh_and_stale_observations(reference, half_life_days);
let config = BuildConfig {
confidence_mode: ConfidenceMode::WilsonLowerBound,
..decay_config(half_life_days, reference)
};
let book = build_prior_book(&observations, &config).unwrap();
let fresh = action(&book, "s", "fresh");
let stale = action(&book, "s", "stale");
assert_eq!(fresh.confidence, stale.confidence);
assert!(
fresh.prior > stale.prior,
"fresh.prior={} should still outrank stale.prior={}: decay must shrink weighted_count \
even when it leaves this Wilson confidence number unchanged",
fresh.prior,
stale.prior
);
}
#[test]
fn hybrid_confidence_drops_with_time_decay() {
let reference = 1_000_000;
let half_life_days = 10.0;
let observations = fresh_and_stale_observations(reference, half_life_days);
let config = BuildConfig {
confidence_mode: ConfidenceMode::Hybrid,
..decay_config(half_life_days, reference)
};
let book = build_prior_book(&observations, &config).unwrap();
let fresh = action(&book, "s", "fresh").confidence;
let stale = action(&book, "s", "stale").confidence;
assert!(
fresh > stale,
"fresh={fresh} should outrank stale={stale}: the heuristic factor in Hybrid \
carries weighted_count, so decay must lower confidence here"
);
}
#[test]
fn validate_config_rejects_non_positive_half_life() {
let observations = vec![obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])];
let config = decay_config(0.0, 1_000_000);
let err = build_prior_book(&observations, &config).unwrap_err();
assert!(matches!(err, Error::InvalidConfig { .. }));
}
#[test]
fn validate_config_rejects_half_life_without_reference() {
let observations = vec![obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])];
let config = BuildConfig {
time_decay_half_life_days: Some(10.0),
time_decay_reference_unix_seconds: None,
..Default::default()
};
let err = build_prior_book(&observations, &config).unwrap_err();
assert!(matches!(err, Error::InvalidConfig { .. }));
}
#[test]
fn validate_config_rejects_negative_source_weight() {
let observations = vec![obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])];
let config = BuildConfig {
source_weights: std::collections::BTreeMap::from([("x".to_string(), -1.0)]),
..Default::default()
};
let err = build_prior_book(&observations, &config).unwrap_err();
assert!(matches!(err, Error::InvalidConfig { .. }));
}
#[test]
fn validate_config_rejects_negative_default_source_weight() {
let observations = vec![obs("c1", 0, "s", "a", Outcome::Unknown, None, 1.0, vec![])];
let config = BuildConfig {
default_source_weight: -1.0,
..Default::default()
};
let err = build_prior_book(&observations, &config).unwrap_err();
assert!(matches!(err, Error::InvalidConfig { .. }));
}
#[test]
fn build_config_fingerprint_changes_with_decay_and_source_config() {
let default_fp = crate::query::build_config_fingerprint(&BuildConfig::default());
let decay_fp = crate::query::build_config_fingerprint(&decay_config(10.0, 1_000_000));
let source_fp = crate::query::build_config_fingerprint(&BuildConfig {
source_weights: std::collections::BTreeMap::from([("human".to_string(), 0.5)]),
..Default::default()
});
assert_ne!(default_fp, decay_fp);
assert_ne!(default_fp, source_fp);
assert_ne!(decay_fp, source_fp);
}
}