1use std::collections::BTreeMap;
11use std::io::{self, Read, Write};
12
13use roaring::RoaringBitmap;
14
15use crate::ordinal::{Ordinal, to_usize};
16
17const MAGIC: &[u8; 8] = b"FTRSETS\0";
18const VERSION: u32 = 1;
19
20#[derive(Debug, thiserror::Error)]
22pub enum MembersError {
23 #[error("memberships I/O failed")]
25 Io(#[from] io::Error),
26 #[error("not a memberships artifact")]
28 Magic,
29 #[error("memberships layout version {found}, expected {expected}")]
31 Version {
32 found: u32,
34 expected: u32,
36 },
37}
38
39#[derive(Debug, Clone, PartialEq, Eq, Default)]
41pub struct Memberships {
42 sets: BTreeMap<u64, RoaringBitmap>,
43}
44
45impl Memberships {
46 #[must_use]
48 pub fn new() -> Self {
49 Self::default()
50 }
51
52 pub fn insert(&mut self, refset: u64, concept: Ordinal) {
54 self.sets.entry(refset).or_default().insert(concept.index());
55 }
56
57 #[must_use]
59 pub fn members(&self, refset: u64) -> Option<&RoaringBitmap> {
60 self.sets.get(&refset)
61 }
62
63 pub fn refsets(&self) -> impl Iterator<Item = u64> + '_ {
65 self.sets.keys().copied()
66 }
67
68 #[must_use]
70 pub fn len(&self) -> usize {
71 self.sets.len()
72 }
73
74 #[must_use]
76 pub fn is_empty(&self) -> bool {
77 self.sets.is_empty()
78 }
79
80 #[must_use]
82 pub fn total(&self) -> u64 {
83 self.sets.values().map(RoaringBitmap::len).sum()
84 }
85
86 pub fn write_to(&self, out: &mut impl Write) -> Result<(), MembersError> {
92 out.write_all(MAGIC)?;
93 out.write_all(&VERSION.to_le_bytes())?;
94 let count = crate::persist::u32_len(self.sets.len(), "the refset count")?;
95 out.write_all(&count.to_le_bytes())?;
96 for (refset, set) in &self.sets {
97 out.write_all(&refset.to_le_bytes())?;
98 let size = crate::persist::u32_len(set.serialized_size(), "a set size")?;
99 out.write_all(&size.to_le_bytes())?;
100 set.serialize_into(&mut *out)?;
101 }
102 Ok(())
103 }
104
105 pub fn read_from(input: &mut impl Read) -> Result<Self, MembersError> {
111 let mut magic = [0_u8; 8];
112 input.read_exact(&mut magic)?;
113 if &magic != MAGIC {
114 return Err(MembersError::Magic);
115 }
116 let version = read_u32(input)?;
117 if version != VERSION {
118 return Err(MembersError::Version {
119 found: version,
120 expected: VERSION,
121 });
122 }
123 let count = read_u32(input)?;
124 let mut sets = BTreeMap::new();
125 for _ in 0..count {
126 let mut long = [0_u8; 8];
127 input.read_exact(&mut long)?;
128 let size = read_u32(input)?;
129 let mut bytes = vec![0_u8; to_usize(size)];
130 input.read_exact(&mut bytes)?;
131 sets.insert(
132 u64::from_le_bytes(long),
133 RoaringBitmap::deserialize_from(bytes.as_slice())?,
134 );
135 }
136 Ok(Self { sets })
137 }
138}
139
140fn read_u32(input: &mut impl Read) -> io::Result<u32> {
141 let mut buffer = [0_u8; 4];
142 input.read_exact(&mut buffer)?;
143 Ok(u32::from_le_bytes(buffer))
144}
145
146#[cfg(test)]
147mod tests {
148 use roaring::RoaringBitmap;
149
150 use super::{MembersError, Memberships};
151 use crate::ordinal::Ordinal;
152
153 #[test]
154 fn memberships_round_trip_and_reject_foreign_bytes() {
155 let mut members = Memberships::new();
156 members.insert(31_000_147_101, Ordinal::new(2));
157 members.insert(31_000_147_101, Ordinal::new(3));
158 members.insert(900_000_000_000_497_000, Ordinal::new(3));
159 assert_eq!(members.len(), 2);
160 assert_eq!(members.total(), 3);
161 assert_eq!(
162 members.members(31_000_147_101).map(RoaringBitmap::len),
163 Some(2)
164 );
165 assert!(members.members(1).is_none());
166 assert_eq!(members.refsets().next(), Some(31_000_147_101));
167 let mut bytes = Vec::new();
168 members.write_to(&mut bytes).expect("writes");
169 assert_eq!(
170 Memberships::read_from(&mut bytes.as_slice()).expect("reads"),
171 members
172 );
173 assert!(matches!(
174 Memberships::read_from(&mut b"XXXXXXXX\0\0\0\0".as_slice()),
175 Err(MembersError::Magic)
176 ));
177 }
178}