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 =
95 u32::try_from(self.sets.len()).map_err(|_| io::Error::other("too many refsets"))?;
96 out.write_all(&count.to_le_bytes())?;
97 for (refset, set) in &self.sets {
98 out.write_all(&refset.to_le_bytes())?;
99 let size = u32::try_from(set.serialized_size())
100 .map_err(|_| io::Error::other("set too large"))?;
101 out.write_all(&size.to_le_bytes())?;
102 set.serialize_into(&mut *out)?;
103 }
104 Ok(())
105 }
106
107 pub fn read_from(input: &mut impl Read) -> Result<Self, MembersError> {
113 let mut magic = [0_u8; 8];
114 input.read_exact(&mut magic)?;
115 if &magic != MAGIC {
116 return Err(MembersError::Magic);
117 }
118 let version = read_u32(input)?;
119 if version != VERSION {
120 return Err(MembersError::Version {
121 found: version,
122 expected: VERSION,
123 });
124 }
125 let count = read_u32(input)?;
126 let mut sets = BTreeMap::new();
127 for _ in 0..count {
128 let mut long = [0_u8; 8];
129 input.read_exact(&mut long)?;
130 let size = read_u32(input)?;
131 let mut bytes = vec![0_u8; to_usize(size)];
132 input.read_exact(&mut bytes)?;
133 sets.insert(
134 u64::from_le_bytes(long),
135 RoaringBitmap::deserialize_from(bytes.as_slice())?,
136 );
137 }
138 Ok(Self { sets })
139 }
140}
141
142fn read_u32(input: &mut impl Read) -> io::Result<u32> {
143 let mut buffer = [0_u8; 4];
144 input.read_exact(&mut buffer)?;
145 Ok(u32::from_le_bytes(buffer))
146}
147
148#[cfg(test)]
149mod tests {
150 use roaring::RoaringBitmap;
151
152 use super::{MembersError, Memberships};
153 use crate::ordinal::Ordinal;
154
155 #[test]
156 fn memberships_round_trip_and_reject_foreign_bytes() {
157 let mut members = Memberships::new();
158 members.insert(31_000_147_101, Ordinal::new(2));
159 members.insert(31_000_147_101, Ordinal::new(3));
160 members.insert(900_000_000_000_497_000, Ordinal::new(3));
161 assert_eq!(members.len(), 2);
162 assert_eq!(members.total(), 3);
163 assert_eq!(
164 members.members(31_000_147_101).map(RoaringBitmap::len),
165 Some(2)
166 );
167 assert!(members.members(1).is_none());
168 assert_eq!(members.refsets().next(), Some(31_000_147_101));
169 let mut bytes = Vec::new();
170 members.write_to(&mut bytes).expect("writes");
171 assert_eq!(
172 Memberships::read_from(&mut bytes.as_slice()).expect("reads"),
173 members
174 );
175 assert!(matches!(
176 Memberships::read_from(&mut b"XXXXXXXX\0\0\0\0".as_slice()),
177 Err(MembersError::Magic)
178 ));
179 }
180}