reth_db_api/models/
integer_list.rs1use crate::table::{Compress, Decompress};
4use bytes::BufMut;
5use core::{fmt, ops::RangeBounds};
6use derive_more::Deref;
7use reth_codecs::DecompressError;
8use roaring::RoaringTreemap;
9
10#[derive(Clone, PartialEq, Eq, Default, Deref)]
21pub struct IntegerList(pub RoaringTreemap);
22
23impl fmt::Debug for IntegerList {
24 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
25 f.write_str("IntegerList")?;
26 f.debug_list().entries(self.0.iter()).finish()
27 }
28}
29
30impl IntegerList {
31 pub fn empty() -> Self {
33 Self(RoaringTreemap::new())
34 }
35
36 pub fn new(list: impl IntoIterator<Item = u64>) -> Result<Self, IntegerListError> {
40 RoaringTreemap::from_sorted_iter(list)
41 .map(Self)
42 .map_err(|_| IntegerListError::UnsortedInput)
43 }
44
45 #[inline]
51 #[track_caller]
52 pub fn new_pre_sorted(list: impl IntoIterator<Item = u64>) -> Self {
53 Self::new(list).expect("IntegerList must be pre-sorted and non-empty")
54 }
55
56 pub fn append(&mut self, list: impl IntoIterator<Item = u64>) -> Result<u64, IntegerListError> {
58 self.0.append(list).map_err(|_| IntegerListError::UnsortedInput)
59 }
60
61 pub fn push(&mut self, value: u64) -> Result<(), IntegerListError> {
63 self.0.try_push(value).map_err(|_| IntegerListError::UnsortedInput)
64 }
65
66 pub fn clear(&mut self) {
68 self.0.clear();
69 }
70
71 pub fn remove_range<R: RangeBounds<u64>>(&mut self, range: R) -> u64 {
73 self.0.remove_range(range)
74 }
75
76 pub fn to_bytes(&self) -> Vec<u8> {
78 let mut vec = Vec::with_capacity(self.0.serialized_size());
79 self.0.serialize_into(&mut vec).expect("not able to encode IntegerList");
80 vec
81 }
82
83 pub fn to_mut_bytes<B: bytes::BufMut>(&self, buf: &mut B) {
85 self.0.serialize_into(buf.writer()).unwrap();
86 }
87
88 pub fn from_bytes(data: &[u8]) -> Result<Self, IntegerListError> {
90 RoaringTreemap::deserialize_from(data)
91 .map(Self)
92 .map_err(|_| IntegerListError::FailedToDeserialize)
93 }
94}
95
96impl serde::Serialize for IntegerList {
97 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
98 where
99 S: serde::Serializer,
100 {
101 use serde::ser::SerializeSeq;
102
103 let mut seq = serializer.serialize_seq(Some(self.len() as usize))?;
104 for e in &self.0 {
105 seq.serialize_element(&e)?;
106 }
107 seq.end()
108 }
109}
110
111struct IntegerListVisitor;
112
113impl<'de> serde::de::Visitor<'de> for IntegerListVisitor {
114 type Value = IntegerList;
115
116 fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
117 f.write_str("a usize array")
118 }
119
120 fn visit_seq<E>(self, mut seq: E) -> Result<Self::Value, E::Error>
121 where
122 E: serde::de::SeqAccess<'de>,
123 {
124 let mut list = IntegerList::empty();
125 while let Some(item) = seq.next_element()? {
126 list.push(item).map_err(serde::de::Error::custom)?;
127 }
128 Ok(list)
129 }
130}
131
132impl<'de> serde::Deserialize<'de> for IntegerList {
133 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
134 where
135 D: serde::Deserializer<'de>,
136 {
137 deserializer.deserialize_byte_buf(IntegerListVisitor)
138 }
139}
140
141#[cfg(any(test, feature = "arbitrary"))]
142use arbitrary::{Arbitrary, Unstructured};
143
144#[cfg(any(test, feature = "arbitrary"))]
145impl<'a> Arbitrary<'a> for IntegerList {
146 fn arbitrary(u: &mut Unstructured<'a>) -> Result<Self, arbitrary::Error> {
147 let mut nums: Vec<u64> = Vec::arbitrary(u)?;
148 nums.sort_unstable();
149 Self::new(nums).map_err(|_| arbitrary::Error::IncorrectFormat)
150 }
151}
152
153#[derive(Debug, derive_more::Display, derive_more::Error)]
155pub enum IntegerListError {
156 #[display("the provided input is unsorted")]
158 UnsortedInput,
159 #[display("failed to deserialize data into type")]
161 FailedToDeserialize,
162}
163
164impl Compress for IntegerList {
165 type Compressed = Vec<u8>;
166
167 fn compress(self) -> Self::Compressed {
168 self.to_bytes()
169 }
170
171 fn compress_to_buf<B: bytes::BufMut + AsMut<[u8]>>(&self, buf: &mut B) {
172 self.to_mut_bytes(buf)
173 }
174}
175
176impl Decompress for IntegerList {
177 fn decompress(value: &[u8]) -> Result<Self, DecompressError> {
178 Self::from_bytes(value).map_err(DecompressError::new)
179 }
180}
181
182#[cfg(test)]
183mod tests {
184 use super::*;
185
186 #[test]
187 fn empty_list() {
188 assert_eq!(IntegerList::empty().len(), 0);
189 assert_eq!(IntegerList::new_pre_sorted(std::iter::empty()).len(), 0);
190 }
191
192 #[test]
193 fn test_integer_list() {
194 let original_list = [1, 2, 3];
195 let ef_list = IntegerList::new(original_list).unwrap();
196 assert_eq!(ef_list.iter().collect::<Vec<_>>(), original_list);
197 }
198
199 #[test]
200 fn test_integer_list_serialization() {
201 let original_list = [1, 2, 3];
202 let ef_list = IntegerList::new(original_list).unwrap();
203
204 let blist = ef_list.to_bytes();
205 assert_eq!(IntegerList::from_bytes(&blist).unwrap(), ef_list)
206 }
207
208 #[test]
209 fn remove_range_matches_filtering() {
210 let values = [1u64, 2, 100, 65_535, 65_536, 70_000, 200_000];
212
213 for to_block in [0u64, 1, 99, 100, 65_535, 65_536, 199_999, 200_000, 200_001] {
214 let mut list = IntegerList::new(values).unwrap();
215 let removed = list.remove_range(0..=to_block);
216
217 let expected = values.into_iter().filter(|value| *value > to_block).collect::<Vec<_>>();
218 assert_eq!(list.iter().collect::<Vec<_>>(), expected, "to_block {to_block}");
219 assert_eq!(removed, (values.len() - expected.len()) as u64, "to_block {to_block}");
220 assert_eq!(list.is_empty(), expected.is_empty(), "to_block {to_block}");
221 }
222 }
223
224 #[test]
225 fn remove_range_on_empty_list_removes_nothing() {
226 let mut list = IntegerList::empty();
227 assert_eq!(list.remove_range(0..=100), 0);
228 assert!(list.is_empty());
229 }
230}