Skip to main content

quoracle/
distribution.rs

1//! Workload distribution types for modeling read/write ratios.
2//!
3//! A `Distribution` describes the probability distribution over
4//! read fractions in a workload. A read fraction `fr` means that
5//! `fr` of the workload is reads and `1 - fr` is writes.
6
7use crate::error::{Error, Result};
8use hashbrown::HashMap;
9
10/// An `f64` usable as a map key.
11///
12/// Equality and hashing use the bit pattern after normalizing `-0.0` to
13/// `0.0`, so `Eq` and `Hash` agree. Values stored by this crate are always
14/// finite and in `[0, 1]`.
15#[derive(Debug, Clone, Copy)]
16pub struct OrderedFloat(pub f64);
17
18impl OrderedFloat {
19    fn key(self) -> u64 {
20        // `0.0 + -0.0 == 0.0`, which folds -0.0 into 0.0.
21        (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
51/// A canonicalized distribution mapping read fractions to
52/// probabilities. All probabilities sum to 1.0 and all
53/// fractions are in [0, 1].
54pub type Canonical = HashMap<OrderedFloat, f64>;
55
56/// A distribution over read fractions.
57///
58/// Build one with [`Distribution::fixed`] or [`Distribution::weighted`],
59/// which validate their inputs. The variants are public for pattern
60/// matching; values built directly are validated again whenever they are
61/// used.
62#[derive(Debug, Clone, PartialEq)]
63pub enum Distribution {
64    /// A single fixed read fraction (e.g. 0.5 means 50% reads).
65    Fixed(f64),
66
67    /// A weighted distribution over multiple read fractions.
68    /// Maps `read_fraction -> weight` (not yet normalized).
69    Weighted(HashMap<OrderedFloat, f64>),
70}
71
72impl Distribution {
73    /// Create a fixed distribution with a single read fraction.
74    ///
75    /// # Errors
76    ///
77    /// Returns an error if `read_fraction` is not in [0.0, 1.0].
78    pub fn fixed(read_fraction: f64) -> Result<Self> {
79        validate_fraction(read_fraction)?;
80        Ok(Self::Fixed(read_fraction))
81    }
82
83    /// Create a weighted distribution from pairs of
84    /// `(read_fraction, weight)`. Validates but does not
85    /// normalize; call [`Distribution::canonicalize`] for normalization.
86    ///
87    /// # Errors
88    ///
89    /// Returns an error if any fraction is not in [0.0, 1.0], any weight
90    /// is negative or not finite, the slice is empty, or all weights are 0.
91    ///
92    /// Repeated fractions have their weights added together.
93    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    /// The distinct read fractions in this distribution, sorted.
109    #[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    /// Canonicalize this distribution into a map of
120    /// `read_fraction -> probability` where probabilities
121    /// sum to 1.0. Zero-weight entries are excluded.
122    ///
123    /// # Errors
124    ///
125    /// Returns an error if any fraction is outside [0, 1], any weight is
126    /// negative or not finite, or the total weight is not positive.
127    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
167/// Convenience conversions so callers can pass a plain `f64`.
168impl 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
184/// Resolve the `read_fraction` / `write_fraction` pair into a
185/// canonical distribution. Exactly one must be `Some`.
186///
187/// When `write_fraction` is provided, each fraction `fw` is
188/// converted to a read fraction as `1.0 - fw`.
189///
190/// # Errors
191///
192/// Returns an error if both parameters are `None`, both are `Some`,
193/// or if canonicalization of the provided distribution fails.
194pub 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    // `contains` is false for NaN, so NaN is rejected too.
223    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    // ---- OrderedFloat -----------------------------------------
237
238    #[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        // Directly-built variants are validated on use.
268        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    // ---- Distribution::fixed ----------------------------------
299
300    #[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    // ---- Distribution::weighted -------------------------------
331
332    #[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    // ---- TryFrom conversions ----------------------------------
382
383    #[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    // ---- canonicalize_rw --------------------------------------
408
409    #[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        // write_fraction 0.3 -> read_fraction 0.7
423        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        // write 0.2 -> read 0.8, write 0.5 -> read 0.5
433        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}