1use crate::error::{Error, Result};
8use hashbrown::HashMap;
9
10#[derive(Debug, Clone, Copy)]
16pub struct OrderedFloat(pub f64);
17
18impl OrderedFloat {
19 fn key(self) -> u64 {
20 (self.0 + 0.0).to_bits()
22 }
23}
24
25impl PartialEq for OrderedFloat {
26 fn eq(&self, other: &Self) -> bool {
27 self.key() == other.key()
28 }
29}
30
31impl Eq for OrderedFloat {}
32
33impl std::hash::Hash for OrderedFloat {
34 fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
35 self.key().hash(state);
36 }
37}
38
39impl From<f64> for OrderedFloat {
40 fn from(v: f64) -> Self {
41 Self(v)
42 }
43}
44
45impl std::fmt::Display for OrderedFloat {
46 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
47 write!(f, "{}", self.0)
48 }
49}
50
51pub type Canonical = HashMap<OrderedFloat, f64>;
55
56#[derive(Debug, Clone, PartialEq)]
63pub enum Distribution {
64 Fixed(f64),
66
67 Weighted(HashMap<OrderedFloat, f64>),
70}
71
72impl Distribution {
73 pub fn fixed(read_fraction: f64) -> Result<Self> {
79 validate_fraction(read_fraction)?;
80 Ok(Self::Fixed(read_fraction))
81 }
82
83 pub fn weighted(weights: &[(f64, f64)]) -> Result<Self> {
94 if weights.is_empty() {
95 return Err(Error::InvalidDistribution(
96 "distribution cannot be empty".into(),
97 ));
98 }
99 let mut mapped: HashMap<OrderedFloat, f64> = HashMap::new();
100 for &(frac, weight) in weights {
101 *mapped.entry(OrderedFloat(frac)).or_default() += weight;
102 }
103 let d = Self::Weighted(mapped);
104 d.canonicalize()?;
105 Ok(d)
106 }
107
108 #[must_use]
110 pub fn fractions(&self) -> Vec<f64> {
111 let mut v: Vec<f64> = match self {
112 Self::Fixed(f) => vec![*f],
113 Self::Weighted(map) => map.keys().map(|k| k.0).collect(),
114 };
115 v.sort_by(f64::total_cmp);
116 v
117 }
118
119 pub fn canonicalize(&self) -> Result<Canonical> {
128 match self {
129 Self::Fixed(f) => {
130 validate_fraction(*f)?;
131 let mut m = HashMap::with_capacity(1);
132 m.insert(OrderedFloat(*f), 1.0);
133 Ok(m)
134 }
135 Self::Weighted(weights) => {
136 if weights.is_empty() {
137 return Err(Error::InvalidDistribution(
138 "distribution cannot be empty".into(),
139 ));
140 }
141 for (frac, &weight) in weights {
142 validate_fraction(frac.0)?;
143 if !weight.is_finite() || weight < 0.0 {
144 return Err(Error::InvalidDistribution(format!(
145 "weight must be finite and non-negative, got \
146 {weight} for fraction {frac}"
147 )));
148 }
149 }
150 let total: f64 = weights.values().sum();
151 if !total.is_finite() || total <= 0.0 {
152 return Err(Error::InvalidDistribution(
153 "total weight must be finite and positive".into(),
154 ));
155 }
156 let m: Canonical = weights
157 .iter()
158 .filter(|(_, &w)| w > 0.0)
159 .map(|(k, &w)| (*k, w / total))
160 .collect();
161 Ok(m)
162 }
163 }
164 }
165}
166
167impl TryFrom<f64> for Distribution {
169 type Error = Error;
170
171 fn try_from(value: f64) -> Result<Self> {
172 Self::fixed(value)
173 }
174}
175
176impl TryFrom<i32> for Distribution {
177 type Error = Error;
178
179 fn try_from(value: i32) -> Result<Self> {
180 Self::fixed(f64::from(value))
181 }
182}
183
184pub fn canonicalize_rw(
195 read_fraction: Option<&Distribution>,
196 write_fraction: Option<&Distribution>,
197) -> Result<Canonical> {
198 match (read_fraction, write_fraction) {
199 (None, None) => Err(Error::InvalidDistribution(
200 "either read_fraction or write_fraction \
201 must be provided"
202 .into(),
203 )),
204 (Some(_), Some(_)) => Err(Error::InvalidDistribution(
205 "only one of read_fraction or \
206 write_fraction can be provided"
207 .into(),
208 )),
209 (Some(d), None) => d.canonicalize(),
210 (None, Some(d)) => {
211 let canon = d.canonicalize()?;
212 let flipped: Canonical = canon
213 .into_iter()
214 .map(|(k, p)| (OrderedFloat(1.0 - k.0), p))
215 .collect();
216 Ok(flipped)
217 }
218 }
219}
220
221fn validate_fraction(f: f64) -> Result<()> {
222 if !(0.0..=1.0).contains(&f) {
224 return Err(Error::InvalidDistribution(format!(
225 "fraction must be in [0, 1], got {f}"
226 )));
227 }
228 Ok(())
229}
230
231#[cfg(test)]
232#[expect(clippy::expect_used)]
233mod tests {
234 use super::*;
235
236 #[test]
239 fn ordered_float_eq_and_hash() {
240 use hashbrown::HashSet;
241 let a = OrderedFloat(0.5);
242 let b = OrderedFloat(0.5);
243 assert_eq!(a, b);
244
245 let mut set = HashSet::new();
246 set.insert(a);
247 assert!(set.contains(&b));
248 }
249
250 #[test]
251 fn ordered_float_signed_zero() {
252 use std::hash::BuildHasher;
253 let h = hashbrown::DefaultHashBuilder::default();
254 assert_eq!(OrderedFloat(0.0), OrderedFloat(-0.0));
255 assert_eq!(
256 h.hash_one(OrderedFloat(0.0)),
257 h.hash_one(OrderedFloat(-0.0))
258 );
259 }
260
261 #[test]
262 fn invalid_values_rejected() {
263 assert!(Distribution::fixed(f64::NAN).is_err());
264 assert!(Distribution::weighted(&[(0.5, f64::NAN)]).is_err());
265 assert!(Distribution::weighted(&[(0.5, f64::INFINITY)]).is_err());
266 assert!(Distribution::weighted(&[(f64::NAN, 1.0)]).is_err());
267 assert!(Distribution::Fixed(2.0).canonicalize().is_err());
269 assert!(Distribution::Weighted(HashMap::new()).canonicalize().is_err());
270 let mut m = HashMap::new();
271 m.insert(OrderedFloat(2.0), 1.0);
272 assert!(Distribution::Weighted(m).canonicalize().is_err());
273 let mut m = HashMap::new();
274 m.insert(OrderedFloat(0.5), -1.0);
275 assert!(Distribution::Weighted(m).canonicalize().is_err());
276 }
277
278 #[test]
279 fn weighted_merges_duplicates_and_sorts_fractions() {
280 let d = Distribution::weighted(&[(0.8, 1.0), (0.2, 1.0), (0.8, 2.0)])
281 .expect("valid");
282 assert_eq!(d.fractions(), vec![0.2, 0.8]);
283 let c = d.canonicalize().expect("valid");
284 assert!((c[&OrderedFloat(0.8)] - 0.75).abs() < 1e-12);
285 }
286
287 #[test]
288 fn ordered_float_display() {
289 assert_eq!(format!("{}", OrderedFloat(0.25)), "0.25");
290 }
291
292 #[test]
293 fn ordered_float_from_f64() {
294 let of: OrderedFloat = 0.75.into();
295 assert_eq!(of.0, 0.75);
296 }
297
298 #[test]
301 fn fixed_valid() {
302 let d = Distribution::fixed(0.0);
303 assert!(d.is_ok());
304 let d = Distribution::fixed(0.5);
305 assert!(d.is_ok());
306 let d = Distribution::fixed(1.0);
307 assert!(d.is_ok());
308 }
309
310 #[test]
311 fn fixed_out_of_range() {
312 assert!(Distribution::fixed(-0.1).is_err());
313 assert!(Distribution::fixed(1.1).is_err());
314 }
315
316 #[test]
317 fn fixed_fractions() {
318 let d = Distribution::fixed(0.3).expect("valid");
319 assert_eq!(d.fractions(), vec![0.3]);
320 }
321
322 #[test]
323 fn fixed_canonicalize() {
324 let d = Distribution::fixed(0.8).expect("valid");
325 let c = d.canonicalize().expect("valid");
326 assert_eq!(c.len(), 1);
327 assert!((c[&OrderedFloat(0.8)] - 1.0).abs() < f64::EPSILON);
328 }
329
330 #[test]
333 fn weighted_valid() {
334 let d = Distribution::weighted(&[(0.25, 1.0), (0.8, 2.0)]);
335 assert!(d.is_ok());
336 }
337
338 #[test]
339 fn weighted_empty() {
340 assert!(Distribution::weighted(&[]).is_err());
341 }
342
343 #[test]
344 fn weighted_negative_weight() {
345 assert!(Distribution::weighted(&[(0.5, -1.0)]).is_err());
346 }
347
348 #[test]
349 fn weighted_zero_total_weight() {
350 assert!(Distribution::weighted(&[(0.5, 0.0)]).is_err());
351 }
352
353 #[test]
354 fn weighted_fraction_out_of_range() {
355 assert!(Distribution::weighted(&[(1.5, 1.0)]).is_err());
356 }
357
358 #[test]
359 fn weighted_canonicalize_normalizes() {
360 let d =
361 Distribution::weighted(&[(0.25, 1.0), (0.8, 2.0)]).expect("valid");
362 let c = d.canonicalize().expect("valid");
363
364 assert_eq!(c.len(), 2);
365 let p_25 = c[&OrderedFloat(0.25)];
366 let p_80 = c[&OrderedFloat(0.8)];
367 assert!((p_25 - 1.0 / 3.0).abs() < 1e-10);
368 assert!((p_80 - 2.0 / 3.0).abs() < 1e-10);
369 assert!((p_25 + p_80 - 1.0).abs() < 1e-10);
370 }
371
372 #[test]
373 fn weighted_canonicalize_excludes_zero_weight() {
374 let d =
375 Distribution::weighted(&[(0.1, 0.0), (0.9, 3.0)]).expect("valid");
376 let c = d.canonicalize().expect("valid");
377 assert_eq!(c.len(), 1);
378 assert!((c[&OrderedFloat(0.9)] - 1.0).abs() < f64::EPSILON);
379 }
380
381 #[test]
384 fn try_from_f64() {
385 let d: Distribution = (0.5_f64).try_into().expect("valid");
386 assert_eq!(d, Distribution::Fixed(0.5));
387 }
388
389 #[test]
390 fn try_from_f64_invalid() {
391 let d: std::result::Result<Distribution, _> = (2.0_f64).try_into();
392 assert!(d.is_err());
393 }
394
395 #[test]
396 fn try_from_i32() {
397 let d: Distribution = 1_i32.try_into().expect("valid");
398 assert_eq!(d, Distribution::Fixed(1.0));
399 }
400
401 #[test]
402 fn try_from_i32_invalid() {
403 let d: std::result::Result<Distribution, _> = (-1_i32).try_into();
404 assert!(d.is_err());
405 }
406
407 #[test]
410 fn canonicalize_rw_read_fraction() {
411 let d = Distribution::fixed(0.6).expect("valid");
412 let c = canonicalize_rw(Some(&d), None).expect("valid");
413 assert_eq!(c.len(), 1);
414 assert!((c[&OrderedFloat(0.6)] - 1.0).abs() < f64::EPSILON);
415 }
416
417 #[test]
418 fn canonicalize_rw_write_fraction() {
419 let d = Distribution::fixed(0.3).expect("valid");
420 let c = canonicalize_rw(None, Some(&d)).expect("valid");
421 assert_eq!(c.len(), 1);
422 assert!((c[&OrderedFloat(0.7)] - 1.0).abs() < f64::EPSILON);
424 }
425
426 #[test]
427 fn canonicalize_rw_write_fraction_weighted() {
428 let d =
429 Distribution::weighted(&[(0.2, 1.0), (0.5, 1.0)]).expect("valid");
430 let c = canonicalize_rw(None, Some(&d)).expect("valid");
431 assert_eq!(c.len(), 2);
432 assert!((c[&OrderedFloat(0.8)] - 0.5).abs() < 1e-10);
434 assert!((c[&OrderedFloat(0.5)] - 0.5).abs() < 1e-10);
435 }
436
437 #[test]
438 fn canonicalize_rw_both_none() {
439 assert!(canonicalize_rw(None, None).is_err());
440 }
441
442 #[test]
443 fn canonicalize_rw_both_some() {
444 let d = Distribution::fixed(0.5).expect("valid");
445 assert!(canonicalize_rw(Some(&d), Some(&d)).is_err());
446 }
447}