Skip to main content

sim_lib_discrete_comb/
subset.rs

1//! Subsets of `{0, ..., n-1}` in bitmask order, with rank/unrank.
2//!
3//! A subset is represented as its sorted member indices. Ordinal `r` is the
4//! bitmask whose bit `i` indicates membership of element `i`.
5
6use crate::error::CombError;
7use num_bigint::BigUint;
8use std::collections::BTreeSet;
9
10const MAX_N: usize = 127;
11
12/// Iterator over all `2^n` subsets of `{0, ..., n-1}` in bitmask order.
13#[derive(Debug, Clone)]
14pub struct SubsetIter {
15    n: usize,
16    next: u128,
17    total: u128,
18}
19
20impl SubsetIter {
21    /// Total number of subsets in this iterator's finite domain.
22    pub fn total_ordinals(&self) -> u128 {
23        self.total
24    }
25
26    /// Number of subsets not yet emitted.
27    pub fn remaining_ordinals(&self) -> u128 {
28        self.total.saturating_sub(self.next)
29    }
30}
31
32/// Construct a subset iterator, rejecting `n` too large for the `u128` cursor.
33pub fn subsets(n: usize) -> Result<SubsetIter, CombError> {
34    if n > MAX_N {
35        return Err(CombError::LimitExceeded(format!(
36            "subset cardinality {n} exceeds {MAX_N}"
37        )));
38    }
39    let total = 1u128 << n;
40    Ok(SubsetIter { n, next: 0, total })
41}
42
43impl Iterator for SubsetIter {
44    type Item = Vec<usize>;
45
46    fn next(&mut self) -> Option<Self::Item> {
47        if self.next >= self.total {
48            return None;
49        }
50        let r = self.next;
51        self.next += 1;
52        Some((0..self.n).filter(|&i| (r >> i) & 1 == 1).collect())
53    }
54}
55
56/// The bitmask ordinal of `subset` (a list of distinct indices `< n`).
57pub fn subset_rank(subset: &[usize], n: usize) -> Result<BigUint, CombError> {
58    let mut rank = BigUint::from(0u32);
59    let mut seen = BTreeSet::new();
60    for &i in subset {
61        if i >= n {
62            return Err(CombError::OutOfRange {
63                value: i.to_string(),
64                bound: n.to_string(),
65            });
66        }
67        if !seen.insert(i) {
68            return Err(CombError::InvalidParameters(format!(
69                "subset member {i} appears more than once"
70            )));
71        }
72        rank.set_bit(i as u64, true);
73    }
74    Ok(rank)
75}
76
77/// The subset (sorted member indices) for bitmask ordinal `rank` over `n`.
78pub fn subset_unrank(rank: &BigUint, n: usize) -> Result<Vec<usize>, CombError> {
79    let bound = BigUint::from(1u32) << n;
80    if rank >= &bound {
81        return Err(CombError::OutOfRange {
82            value: rank.to_string(),
83            bound: bound.to_string(),
84        });
85    }
86    Ok((0..n).filter(|&i| rank.bit(i as u64)).collect())
87}
88
89#[cfg(test)]
90mod tests {
91    use super::*;
92
93    #[test]
94    fn powerset_count_and_order() {
95        let all: Vec<_> = subsets(3).unwrap().collect();
96        assert_eq!(all.len(), 8);
97        assert_eq!(all[0], Vec::<usize>::new());
98        assert_eq!(all[1], vec![0]);
99        assert_eq!(all[7], vec![0, 1, 2]);
100    }
101
102    #[test]
103    fn rank_unrank_round_trip() {
104        for (i, s) in subsets(4).unwrap().enumerate() {
105            let r = subset_rank(&s, 4).unwrap();
106            assert_eq!(r, BigUint::from(i as u32));
107            assert_eq!(subset_unrank(&r, 4).unwrap(), s);
108        }
109    }
110
111    #[test]
112    fn rank_rejects_out_of_range() {
113        assert!(matches!(
114            subset_rank(&[5], 4),
115            Err(CombError::OutOfRange { .. })
116        ));
117    }
118
119    #[test]
120    fn rank_rejects_duplicates() {
121        assert!(matches!(
122            subset_rank(&[1, 1], 4),
123            Err(CombError::InvalidParameters(_))
124        ));
125    }
126
127    #[test]
128    fn unrank_rejects_cardinality() {
129        assert!(matches!(
130            subset_unrank(&BigUint::from(8u32), 3),
131            Err(CombError::OutOfRange { .. })
132        ));
133    }
134
135    #[test]
136    fn n_127_total_is_exact_domain_size() {
137        let iter = subsets(127).unwrap();
138        assert_eq!(iter.total_ordinals(), 1u128 << 127);
139        assert_eq!(iter.remaining_ordinals(), 1u128 << 127);
140    }
141}