refactor!: close the remaining API gaps from #21

Three unrelated small defects, all requiring signature changes:

- `Game::one_v_one` hardcoded `GameOptions::default()`, so a 1v1 could
  never set `p_draw` or convergence options — and a drawn 1v1 was
  therefore unreachable through it, since the default `p_draw` is zero.
  It now takes `&GameOptions` like every other constructor.

- `Observer::on_batch_processed` was declared on the trait and never
  called from anywhere: implementors wired up a callback that could not
  fire. It is now called after each slice sweep, and renamed
  `on_slice_processed` to match the vocabulary the codebase adopted in
  T2 — the unit of work is a `TimeSlice`, not a batch. A slice is swept
  once travelling backward and once forward, so a multi-slice history
  fires it twice per slice per iteration; the doc comment says so.

- `pub mod factors` sat beside `pub(crate) mod factor`, two module paths
  differing by one character with only one of them importable. The
  public facade is now `graph`.

Tests cover each as a behaviour rather than a compile check: a drawn 1v1
succeeds only when p_draw is supplied, and the observer tests fail if
any callback stops firing.

BREAKING CHANGE: `Game::one_v_one` takes a fourth `&GameOptions`
argument; `Observer::on_batch_processed` is renamed
`on_slice_processed`; the `factors` module is renamed `graph`.

Closes #21

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_011hcFjNDmHXZF8URGLku5zZ
This commit is contained in:
2026-09-07 14:57:39 +02:00
co-authored by Claude Opus 5
parent bb2a845882
commit 507894dae7
8 changed files with 190 additions and 14 deletions
+11 -9
View File
@@ -532,18 +532,20 @@ impl<T: Time, D: Drift<T>> Game<'_, T, D> {
)) ))
} }
/// Convenience wrapper over [`Game::ranked`] for two single-player teams.
///
/// # Errors /// # Errors
/// ///
/// Delegates to [`Game::ranked`] with default options, so it returns the /// Delegates to [`Game::ranked`], so it returns the same errors — in
/// same errors — in practice `WrongOutcomeKind` for a non-ranked outcome, /// practice `WrongOutcomeKind` for a non-ranked outcome, or
/// or `TieWithoutDrawProbability` for a draw, since the default `p_draw` /// `TieWithoutDrawProbability` for a draw when `options.p_draw` is zero.
/// applies rather than one you chose.
pub fn one_v_one( pub fn one_v_one(
a: &Rating<T, D>, a: &Rating<T, D>,
b: &Rating<T, D>, b: &Rating<T, D>,
outcome: crate::Outcome, outcome: crate::Outcome,
options: &GameOptions,
) -> Result<(Gaussian, Gaussian), crate::InferenceError> { ) -> Result<(Gaussian, Gaussian), crate::InferenceError> {
let game = Self::ranked(&[&[*a], &[*b]], outcome, &GameOptions::default())?; let game = Self::ranked(&[&[*a], &[*b]], outcome, options)?;
let post = game.posteriors(); let post = game.posteriors();
Ok((post[0][0], post[1][0])) Ok((post[0][0], post[1][0]))
} }
@@ -563,11 +565,11 @@ impl<T: Time, D: Drift<T>> Game<'_, T, D> {
} }
#[doc(hidden)] #[doc(hidden)]
pub fn custom<S: crate::factors::Schedule>( pub fn custom<S: crate::graph::Schedule>(
factors: &mut [crate::factors::BuiltinFactor], factors: &mut [crate::graph::BuiltinFactor],
vars: &mut crate::factors::VarStore, vars: &mut crate::graph::VarStore,
schedule: &S, schedule: &S,
) -> crate::factors::ScheduleReport { ) -> crate::graph::ScheduleReport {
schedule.run(factors, vars) schedule.run(factors, vars)
} }
} }
+4
View File
@@ -1,5 +1,9 @@
//! Factor-graph public API. //! Factor-graph public API.
//! //!
//! Named `graph` rather than `factors` because the private implementation
//! module beside it is `factor`: two module paths differing by one character,
//! one public and one not, was a standing invitation to import the wrong one.
//!
//! The factor types, `VarStore` and the `Schedule` trait are public so custom //! The factor types, `VarStore` and the `Schedule` trait are public so custom
//! schedules can be written against them. //! schedules can be written against them.
//! //!
+15
View File
@@ -261,6 +261,11 @@ impl<T: Time, D: Drift<T>, O: Observer<T>, K: Eq + Hash + Clone> History<T, D, O
let old = self.time_slices[j].posteriors(); let old = self.time_slices[j].posteriors();
self.time_slices[j].new_backward_info(&self.agents); self.time_slices[j].new_backward_info(&self.agents);
self.observer.on_slice_processed(
&self.time_slices[j].time,
j,
self.time_slices[j].events.len(),
);
let new = self.time_slices[j].posteriors(); let new = self.time_slices[j].posteriors();
@@ -280,6 +285,11 @@ impl<T: Time, D: Drift<T>, O: Observer<T>, K: Eq + Hash + Clone> History<T, D, O
let old = self.time_slices[j].posteriors(); let old = self.time_slices[j].posteriors();
self.time_slices[j].new_forward_info(&self.agents); self.time_slices[j].new_forward_info(&self.agents);
self.observer.on_slice_processed(
&self.time_slices[j].time,
j,
self.time_slices[j].events.len(),
);
let new = self.time_slices[j].posteriors(); let new = self.time_slices[j].posteriors();
@@ -292,6 +302,11 @@ impl<T: Time, D: Drift<T>, O: Observer<T>, K: Eq + Hash + Clone> History<T, D, O
let old = self.time_slices[0].posteriors(); let old = self.time_slices[0].posteriors();
self.time_slices[0].iteration(0, &self.agents); self.time_slices[0].iteration(0, &self.agents);
self.observer.on_slice_processed(
&self.time_slices[0].time,
0,
self.time_slices[0].events.len(),
);
let new = self.time_slices[0].posteriors(); let new = self.time_slices[0].posteriors();
+1 -1
View File
@@ -118,9 +118,9 @@ mod error;
mod event; mod event;
mod event_builder; mod event_builder;
pub(crate) mod factor; pub(crate) mod factor;
pub mod factors;
mod game; mod game;
pub mod gaussian; pub mod gaussian;
pub mod graph;
mod history; mod history;
mod key_table; mod key_table;
mod matrix; mod matrix;
+8 -2
View File
@@ -14,8 +14,13 @@ pub trait Observer<T: Time>: Send + Sync {
/// Called after each convergence iteration across the whole history. /// Called after each convergence iteration across the whole history.
fn on_iteration_end(&self, _iter: usize, _max_step: (f64, f64)) {} fn on_iteration_end(&self, _iter: usize, _max_step: (f64, f64)) {}
/// Called after each time slice is processed within an iteration. /// Called after each time slice is swept within an iteration.
fn on_batch_processed(&self, _time: &T, _slice_idx: usize, _n_events: usize) {} ///
/// A convergence iteration sweeps every slice twice — once travelling
/// backward through the history and once forward — so a multi-slice
/// history fires this twice per slice per iteration. A single-slice
/// history is swept once and fires once.
fn on_slice_processed(&self, _time: &T, _slice_idx: usize, _n_events: usize) {}
/// Called once when convergence completes (or max iters is reached). /// Called once when convergence completes (or max iters is reached).
fn on_converged(&self, _iters: usize, _final_step: (f64, f64), _converged: bool) {} fn on_converged(&self, _iters: usize, _final_step: (f64, f64), _converged: bool) {}
@@ -35,6 +40,7 @@ mod tests {
fn null_observer_compiles_for_i64() { fn null_observer_compiles_for_i64() {
let o = NullObserver; let o = NullObserver;
<NullObserver as Observer<i64>>::on_iteration_end(&o, 1, (0.0, 0.0)); <NullObserver as Observer<i64>>::on_iteration_end(&o, 1, (0.0, 0.0));
<NullObserver as Observer<i64>>::on_slice_processed(&o, &7, 0, 3);
<NullObserver as Observer<i64>>::on_converged(&o, 5, (1e-6, 1e-6), true); <NullObserver as Observer<i64>>::on_converged(&o, 5, (1e-6, 1e-6), true);
} }
+2 -1
View File
@@ -19,7 +19,8 @@ fn ts_rating(mu: f64, sigma: f64, beta: f64, gamma: f64) -> R {
fn game_1v1_golden_matches_historical() { fn game_1v1_golden_matches_historical() {
let a = ts_rating(25.0, 25.0 / 3.0, 25.0 / 6.0, 25.0 / 300.0); let a = ts_rating(25.0, 25.0 / 3.0, 25.0 / 6.0, 25.0 / 300.0);
let b = ts_rating(25.0, 25.0 / 3.0, 25.0 / 6.0, 25.0 / 300.0); let b = ts_rating(25.0, 25.0 / 3.0, 25.0 / 6.0, 25.0 / 300.0);
let (a_post, b_post) = Game::<i64, _>::one_v_one(&a, &b, Outcome::winner(0, 2)).unwrap(); let (a_post, b_post) =
Game::<i64, _>::one_v_one(&a, &b, Outcome::winner(0, 2), &GameOptions::default()).unwrap();
// Historical golden from pre-T2 test_1vs1 (team 0 wins): // Historical golden from pre-T2 test_1vs1 (team 0 wins):
assert_ulps_eq!( assert_ulps_eq!(
a_post, a_post,
+44 -1
View File
@@ -32,7 +32,8 @@ fn game_ranked_1v1_golden() {
fn game_one_v_one_shortcut() { fn game_one_v_one_shortcut() {
let a = default_rating(); let a = default_rating();
let b = default_rating(); let b = default_rating();
let (a_post, b_post) = Game::<i64, _>::one_v_one(&a, &b, Outcome::winner(0, 2)).unwrap(); let (a_post, b_post) =
Game::<i64, _>::one_v_one(&a, &b, Outcome::winner(0, 2), &GameOptions::default()).unwrap();
assert!(a_post.mu() > 25.0); assert!(a_post.mu() > 25.0);
assert!(b_post.mu() < 25.0); assert!(b_post.mu() < 25.0);
} }
@@ -95,3 +96,45 @@ fn game_log_evidence_is_finite() {
assert!(g.log_evidence().is_finite()); assert!(g.log_evidence().is_finite());
assert!(g.log_evidence() < 0.0); assert!(g.log_evidence() < 0.0);
} }
/// `one_v_one` used to hardcode `GameOptions::default()`, so a 1v1 could
/// never set `p_draw` and a drawn 1v1 was unreachable through it.
#[test]
fn one_v_one_honours_the_draw_probability_it_is_given() {
let a = default_rating();
let b = default_rating();
// Default options still reject a draw, because the default p_draw is zero.
let err = Game::<i64, _>::one_v_one(&a, &b, Outcome::draw(2), &GameOptions::default())
.expect_err("a draw needs a positive p_draw");
assert!(matches!(
err,
InferenceError::TieWithoutDrawProbability { .. }
));
// With a draw probability supplied it succeeds — which was impossible
// before the signature took options.
let options = GameOptions {
p_draw: 0.25,
..GameOptions::default()
};
let (a_post, b_post) = Game::<i64, _>::one_v_one(&a, &b, Outcome::draw(2), &options)
.expect("a draw is representable once p_draw is positive");
// A symmetric draw leaves the means alone and sharpens both sides.
assert!((a_post.mu() - b_post.mu()).abs() < 1e-9);
assert!(a_post.sigma() < 25.0 / 3.0);
}
/// Convergence options reach the 1v1 path too, not just `p_draw`.
#[test]
fn one_v_one_honours_convergence_options() {
let a = default_rating();
let b = default_rating();
let options = GameOptions {
convergence: ConvergenceOptions::default(),
..GameOptions::default()
};
let (a_post, _) = Game::<i64, _>::one_v_one(&a, &b, Outcome::winner(0, 2), &options).unwrap();
assert!(a_post.mu() > 25.0);
}
+105
View File
@@ -0,0 +1,105 @@
//! `Observer` callbacks must actually fire.
//!
//! `on_slice_processed` (formerly `on_batch_processed`) was declared on the
//! trait and never called from anywhere, so implementors wired up a callback
//! that could not run. These tests exist so that cannot silently recur.
use std::sync::{Arc, Mutex};
use trueskill_tt::{History, Observer};
/// `History` takes its observer by value and never hands it back, so a test
/// that wants to read what was recorded shares the storage rather than the
/// observer: the handles are cloned, the buffers are not.
#[derive(Clone, Default)]
struct Recorder {
iterations: Arc<Mutex<Vec<usize>>>,
slices: Arc<Mutex<Vec<(i64, usize, usize)>>>,
converged: Arc<Mutex<Vec<(usize, bool)>>>,
}
impl Observer<i64> for Recorder {
fn on_iteration_end(&self, iter: usize, _max_step: (f64, f64)) {
self.iterations.lock().unwrap().push(iter);
}
fn on_slice_processed(&self, time: &i64, slice_idx: usize, n_events: usize) {
self.slices
.lock()
.unwrap()
.push((*time, slice_idx, n_events));
}
fn on_converged(&self, iters: usize, _final_step: (f64, f64), converged: bool) {
self.converged.lock().unwrap().push((iters, converged));
}
}
#[test]
fn every_observer_callback_fires() {
let recorder = Recorder::default();
let mut h = History::builder().observer(recorder.clone()).build();
h.record_winner(&"a", &"b", 1).unwrap();
h.record_winner(&"b", &"c", 2).unwrap();
h.record_winner(&"c", &"a", 3).unwrap();
h.converge().unwrap();
assert!(
!recorder.iterations.lock().unwrap().is_empty(),
"on_iteration_end never fired"
);
assert!(
!recorder.converged.lock().unwrap().is_empty(),
"on_converged never fired"
);
assert!(
!recorder.slices.lock().unwrap().is_empty(),
"on_slice_processed never fired — the defect this test exists for"
);
}
#[test]
fn slice_callbacks_report_the_slice_they_swept() {
let recorder = Recorder::default();
let mut h = History::builder().observer(recorder.clone()).build();
h.record_winner(&"a", &"b", 10).unwrap();
h.record_winner(&"a", &"b", 20).unwrap();
h.converge().unwrap();
let slices = recorder.slices.lock().unwrap();
// Only the times actually in the history, and each with its own events.
for &(time, idx, events) in slices.iter() {
assert!(time == 10 || time == 20, "unexpected slice time {time}");
assert!(idx < 2, "slice index {idx} out of range");
assert_eq!(events, 1, "each slice holds exactly one event");
}
// Both slices must be reported, not just one end of the sweep.
assert!(
slices.iter().any(|&(t, ..)| t == 10),
"slice 10 never reported"
);
assert!(
slices.iter().any(|&(t, ..)| t == 20),
"slice 20 never reported"
);
}
#[test]
fn a_single_slice_history_still_reports_its_sweep() {
let recorder = Recorder::default();
let mut h = History::builder().observer(recorder.clone()).build();
h.record_winner(&"a", &"b", 1).unwrap();
h.converge().unwrap();
let slices = recorder.slices.lock().unwrap();
assert!(
!slices.is_empty(),
"the single-slice path must report its sweep too"
);
assert!(slices.iter().all(|&(t, idx, _)| t == 1 && idx == 0));
}