diff --git a/src/history.rs b/src/history.rs index cd96865..b361780 100644 --- a/src/history.rs +++ b/src/history.rs @@ -507,6 +507,32 @@ impl HistoryBuilder Self { Self::default() } + + /// Set the drift rate, in skill units per unit of time. + /// + /// Shorthand for `.drift(ConstantDrift::new(gamma))`. Drift is the + /// most-tuned parameter after `sigma` and [`GAMMA`](crate::GAMMA) is a + /// public constant, but reaching it otherwise means first discovering + /// [`ConstantDrift`] — a type a caller has no other reason to name. + /// + /// Only on the `ConstantDrift` builder: `gamma` is that model's parameter, + /// not a property every [`Drift`] has. Use + /// [`drift`](HistoryBuilder::drift) for anything else. + /// + /// # Panics + /// + /// Panics unless `gamma` is finite and non-negative. Drift enters as + /// `gamma^2` per elapsed tick, so a negative value would behave as its + /// absolute value — the same sign-absorption already rejected for `sigma` + /// and `beta`. + pub fn gamma(self, gamma: f64) -> Self { + assert!( + gamma.is_finite() && gamma >= 0.0, + "gamma must be finite and non-negative (got {gamma}); it is only ever \ + squared, so a negative value would silently behave as its absolute value" + ); + self.drift(ConstantDrift::new(gamma)) + } } impl, O: Observer, K: Eq + Hash + Clone> History { @@ -531,7 +557,7 @@ impl, O: Observer, K: Eq + Hash + Clone> History(&self, key: &Q) -> Option where K: Borrow, - Q: Hash + Eq + ToOwned + ?Sized, + Q: Hash + Eq + ?Sized, { self.keys.get(key) } @@ -1128,9 +1154,10 @@ impl, O: Observer, K: Eq + Hash + Clone> History Result>, InferenceError> + fn member_skills(&self, teams: &[&[&Q]]) -> Result>, InferenceError> where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { if teams.len() < 2 { return Err(InferenceError::NotEnoughTeams { got: teams.len() }); @@ -1244,9 +1271,13 @@ impl, O: Observer, K: Eq + Hash + Clone> History Result<(Vec, Vec), InferenceError> + fn performances( + &self, + teams: &[&[&Q]], + ) -> Result<(Vec, Vec), InferenceError> where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { let skills = self.member_skills(teams)?; @@ -1310,9 +1341,10 @@ impl, O: Observer, K: Eq + Hash + Clone> History Result + pub fn predict_quality(&self, teams: &[&[&Q]]) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { let groups = self.member_skills(teams)?; let group_refs: Vec<&[Gaussian]> = groups.iter().map(Vec::as_slice).collect(); @@ -1448,14 +1480,15 @@ impl, O: Observer, K: Eq + Hash + Clone> History( &self, - terms: &[(&K, f64)], + terms: &[(&Q, f64)], width: usize, row_for: impl Fn(Index) -> Option<(usize, usize)>, ) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { let mut contrast = vec![0.0; width]; let mut unseen: BTreeMap = BTreeMap::new(); @@ -1557,9 +1590,10 @@ impl, O: Observer, K: Eq + Hash + Clone> History Result + pub fn posterior_of(&self, terms: &[(&Q, f64)]) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { self.joint()?.posterior_of(terms) } @@ -1579,9 +1613,14 @@ impl, O: Observer, K: Eq + Hash + Clone> History Result + pub fn posterior_of_at( + &self, + time: T, + terms: &[(&Q, f64)], + ) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { self.joint()?.posterior_of_at(time, terms) } @@ -1625,13 +1664,14 @@ impl, O: Observer, K: Eq + Hash + Clone> History( &self, - teams: &[&[&K]], - target: &[(&K, f64)], + teams: &[&[&Q]], + target: &[(&Q, f64)], ) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { self.joint()?.expected_variance_reduction(teams, target) } @@ -1747,9 +1787,10 @@ impl, O: Observer, K: Eq + Hash + Clone> History Result + pub fn predict_margin(&self, teams: &[&[&Q]]) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { if teams.len() != 2 { return Err(InferenceError::MismatchedShape { @@ -1759,7 +1800,7 @@ impl, O: Observer, K: Eq + Hash + Clone> History = Vec::new(); + let mut terms: Vec<(&Q, f64)> = Vec::new(); let mut performance_noise = 0.0; for (team_idx, team) in teams.iter().enumerate() { @@ -1824,9 +1865,10 @@ impl, O: Observer, K: Eq + Hash + Clone> History Result + pub fn expected_information_gain(&self, teams: &[&[&Q]]) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { let skills = self.member_skills(teams)?; @@ -1884,9 +1926,10 @@ impl, O: Observer, K: Eq + Hash + Clone> History Result, InferenceError> + pub fn predict_win_probabilities(&self, teams: &[&[&Q]]) -> Result, InferenceError> where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { let (performances, sizes) = self.performances(teams)?; Ok(crate::predict::win_probabilities( @@ -1938,9 +1981,10 @@ impl, O: Observer, K: Eq + Hash + Clone> History Result + pub fn predict_outcome(&self, teams: &[&[&Q]]) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { if teams.len() > crate::MAX_PREDICTED_TEAMS { return Err(InferenceError::TooManyTeams { @@ -1987,9 +2031,10 @@ impl, O: Observer, K: Eq + Hash + Clone> History Result + pub fn predict_ranking(&self, teams: &[&[&Q]], ranks: &[u32]) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { if ranks.len() != teams.len() { return Err(InferenceError::MismatchedShape { @@ -2877,9 +2922,10 @@ impl, O: Observer, K: Eq + Hash + Clone> Joint<'_, T, D, /// # Errors /// /// `UnknownKey` for a competitor the history has never seen. - pub fn posterior_of(&self, terms: &[(&K, f64)]) -> Result + pub fn posterior_of(&self, terms: &[(&Q, f64)]) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { let resolved = self .history @@ -2895,9 +2941,14 @@ impl, O: Observer, K: Eq + Hash + Clone> Joint<'_, T, D, /// # Errors /// /// `UnknownKey` for a competitor with no appearance at or before `time`. - pub fn posterior_of_at(&self, time: T, terms: &[(&K, f64)]) -> Result + pub fn posterior_of_at( + &self, + time: T, + terms: &[(&Q, f64)], + ) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { let as_of = self.rows_as_of(time); let resolved = self @@ -2932,13 +2983,14 @@ impl, O: Observer, K: Eq + Hash + Clone> Joint<'_, T, D, /// /// `MismatchedShape` unless exactly two teams are supplied, `EmptyTeam` for /// an empty one, and `UnknownKey` for an unseen competitor. - pub fn expected_variance_reduction( + pub fn expected_variance_reduction( &self, - teams: &[&[&K]], - target: &[(&K, f64)], + teams: &[&[&Q]], + target: &[(&Q, f64)], ) -> Result where - K: std::fmt::Debug, + K: Borrow, + Q: Hash + Eq + ?Sized + std::fmt::Debug, { if teams.len() != 2 { return Err(InferenceError::MismatchedShape { @@ -2950,7 +3002,7 @@ impl, O: Observer, K: Eq + Hash + Clone> Joint<'_, T, D, // The candidate matchup, expressed as the same kind of linear // functional as the target. - let mut matchup: Vec<(&K, f64)> = Vec::new(); + let mut matchup: Vec<(&Q, f64)> = Vec::new(); let mut noise = self.history.score_sigma * self.history.score_sigma; for (team_idx, team) in teams.iter().enumerate() { if team.is_empty() { diff --git a/tests/key_ergonomics.rs b/tests/key_ergonomics.rs new file mode 100644 index 0000000..1078bd6 --- /dev/null +++ b/tests/key_ergonomics.rs @@ -0,0 +1,119 @@ +//! The realistic program: keys arrive owned, queries are written with literals. +//! +//! Every prediction and joint query used to take `&[&[&K]]`, which at +//! `K = String` made a string literal *impossible* — the shape required three +//! levels of temporaries that all had to outlive the call. They are generic +//! over the borrowed key now, so one spelling works at both key types. +//! +//! Both key types are exercised in every test, because the point is that the +//! spelling is the same. + +use trueskill_tt::{ConstantDrift, History, NullObserver}; + +type Owned = History; +type Borrowed = History; + +fn owned() -> Owned { + let mut h: Owned = History::builder().key_type::().build(); + for t in 1..=4 { + h.record_winner(&"alice".to_string(), &"bob".to_string(), t) + .expect("ingests"); + } + h.converge().expect("converges"); + h +} + +fn borrowed() -> Borrowed { + let mut h = History::default(); + for t in 1..=4 { + h.record_winner(&"alice", &"bob", t).expect("ingests"); + } + h.converge().expect("converges"); + h +} + +#[test] +fn predictions_take_literals_at_either_key_type() { + let teams: &[&[&str]] = &[&["alice"], &["bob"]]; + + let a = owned() + .predict_win_probabilities(teams) + .expect("K = String"); + let b = borrowed() + .predict_win_probabilities(teams) + .expect("K = &'static str"); + + assert_eq!(a, b, "the same fit through the same spelling"); + assert!(a[0] > a[1], "alice won every game"); +} + +#[test] +fn every_team_shaped_query_accepts_the_same_slice() { + let h = owned(); + let teams: &[&[&str]] = &[&["alice"], &["bob"]]; + + h.predict_quality(teams).expect("quality"); + let _ = h.predict_outcome(teams).expect("outcome"); + h.predict_ranking(teams, &[0, 1]).expect("ranking"); + h.expected_information_gain(teams) + .expect("information gain"); +} + +#[test] +fn linear_combinations_take_bare_keys() { + // `&[(&K, f64)]` at `K = String` meant `&[(&String, f64)]` — no literals. + let h = owned(); + let terms: &[(&str, f64)] = &[("alice", 1.0), ("bob", -1.0)]; + + // Ranked history, so the joint is unavailable — but the *call* compiles, + // which is what this pins. The error proves it reached the joint check + // rather than failing to resolve a key. + let err = h + .posterior_of(terms) + .expect_err("ranked history has no joint"); + assert!( + format!("{err}").contains("ranked"), + "expected the joint-unavailable path, got {err}" + ); +} + +#[test] +fn lookup_accepts_a_borrowed_key_like_its_neighbours() { + // `lookup` carried `ToOwned`, copy-pasted from `intern`, which + // genuinely needs it to create the entry. `lookup` never creates. + let h = owned(); + assert!(h.lookup("alice").is_some()); + assert!(h.lookup("nobody").is_none()); + + // Control: its neighbours already accepted this and must still. + assert!(h.current_skill("alice").is_some()); + assert!(h.rating("alice").is_some()); +} + +#[test] +fn gamma_sets_drift_without_naming_constant_drift() { + let mut a: Borrowed = History::builder().gamma(0.5).build(); + let mut b: Borrowed = History::builder().drift(ConstantDrift::new(0.5)).build(); + + for h in [&mut a, &mut b] { + h.record_winner(&"x", &"y", 1).unwrap(); + h.record_winner(&"y", &"x", 100).unwrap(); + h.converge().unwrap(); + } + + let (ga, gb) = (a.current_skill("x").unwrap(), b.current_skill("x").unwrap()); + assert_eq!((ga.mu(), ga.sigma()), (gb.mu(), gb.sigma())); + + // Control: the shorthand is not a no-op — a different gamma differs. + let mut c: Borrowed = History::builder().gamma(0.0).build(); + c.record_winner(&"x", &"y", 1).unwrap(); + c.record_winner(&"y", &"x", 100).unwrap(); + c.converge().unwrap(); + assert_ne!(c.current_skill("x").unwrap().sigma(), ga.sigma()); +} + +#[test] +#[should_panic(expected = "gamma must be finite and non-negative")] +fn a_negative_gamma_is_rejected_rather_than_squared_away() { + let _: Borrowed = History::builder().gamma(-0.5).build(); +} diff --git a/tests/prediction.rs b/tests/prediction.rs index 5a0f792..c18e5d1 100644 --- a/tests/prediction.rs +++ b/tests/prediction.rs @@ -60,8 +60,12 @@ fn degenerate_team_shapes_are_errors_rather_than_panics() { h.predict_outcome(&[&[&"a"]]).unwrap_err(), InferenceError::NotEnoughTeams { got: 1, .. } ),); + // An empty team list cannot infer the key type — nothing in `&[]` names it. + // The annotation is the cost of `predict_*` being generic over the borrowed + // key, and it only bites on the degenerate call. + let none: &[&[&str]] = &[]; assert!(matches!( - h.predict_outcome(&[]).unwrap_err(), + h.predict_outcome(none).unwrap_err(), InferenceError::NotEnoughTeams { got: 0, .. } ),); assert!(matches!(