Skip to main content

commonware_utils/
range.rs

1//! Non-empty [`Range`] type that guarantees at least one element.
2
3use bytes::{Buf, BufMut};
4use commonware_codec::{BufsMut, EncodeSize, Error as CodecError, Read, Write};
5use core::{fmt, ops::Range};
6
7/// Error returned when attempting to create a non-empty range from an empty range.
8#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
9#[error("range is empty")]
10pub struct EmptyRange;
11
12/// A non-empty [`Range`] (`start..end`) where `start < end` is guaranteed.
13#[derive(Clone, PartialEq, Eq, Hash)]
14pub struct NonEmptyRange<Idx>(Range<Idx>);
15
16impl<Idx: fmt::Debug> fmt::Debug for NonEmptyRange<Idx> {
17    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
18        self.0.fmt(f)
19    }
20}
21
22impl<Idx: PartialOrd> NonEmptyRange<Idx> {
23    /// Creates a `NonEmptyRange` if `start < end`.
24    pub fn new(range: Range<Idx>) -> Result<Self, EmptyRange> {
25        (range.start < range.end)
26            .then_some(Self(range))
27            .ok_or(EmptyRange)
28    }
29}
30
31impl<Idx: Copy> NonEmptyRange<Idx> {
32    /// Returns the start of the range.
33    pub const fn start(&self) -> Idx {
34        self.0.start
35    }
36
37    /// Returns the end of the range (exclusive).
38    pub const fn end(&self) -> Idx {
39        self.0.end
40    }
41}
42
43impl<Idx: PartialOrd> TryFrom<Range<Idx>> for NonEmptyRange<Idx> {
44    type Error = EmptyRange;
45
46    fn try_from(range: Range<Idx>) -> Result<Self, Self::Error> {
47        Self::new(range)
48    }
49}
50
51impl<Idx> From<NonEmptyRange<Idx>> for Range<Idx> {
52    fn from(r: NonEmptyRange<Idx>) -> Self {
53        r.0
54    }
55}
56
57impl<Idx> IntoIterator for NonEmptyRange<Idx>
58where
59    Range<Idx>: Iterator,
60{
61    type Item = <Range<Idx> as Iterator>::Item;
62    type IntoIter = Range<Idx>;
63
64    fn into_iter(self) -> Self::IntoIter {
65        self.0
66    }
67}
68
69impl<Idx: Write> Write for NonEmptyRange<Idx> {
70    #[inline]
71    fn write(&self, buf: &mut impl BufMut) {
72        self.0.start.write(buf);
73        self.0.end.write(buf);
74    }
75
76    #[inline]
77    fn write_bufs(&self, buf: &mut impl BufsMut) {
78        self.0.start.write_bufs(buf);
79        self.0.end.write_bufs(buf);
80    }
81}
82
83impl<Idx: EncodeSize> EncodeSize for NonEmptyRange<Idx> {
84    #[inline]
85    fn encode_size(&self) -> usize {
86        self.0.start.encode_size() + self.0.end.encode_size()
87    }
88
89    #[inline]
90    fn encode_inline_size(&self) -> usize {
91        self.0.start.encode_inline_size() + self.0.end.encode_inline_size()
92    }
93}
94
95impl<Idx: Read + PartialOrd> Read for NonEmptyRange<Idx> {
96    type Cfg = Idx::Cfg;
97
98    #[inline]
99    fn read_cfg(buf: &mut impl Buf, cfg: &Self::Cfg) -> Result<Self, CodecError> {
100        let start = Idx::read_cfg(buf, cfg)?;
101        let end = Idx::read_cfg(buf, cfg)?;
102        if !start.partial_cmp(&end).is_some_and(|o| o.is_lt()) {
103            return Err(CodecError::Invalid("NonEmptyRange", "start must be < end"));
104        }
105        Ok(Self(start..end))
106    }
107}
108
109#[cfg(feature = "arbitrary")]
110impl<'a, Idx: arbitrary::Arbitrary<'a> + Ord> arbitrary::Arbitrary<'a> for NonEmptyRange<Idx> {
111    fn arbitrary(u: &mut arbitrary::Unstructured<'a>) -> arbitrary::Result<Self> {
112        let a = Idx::arbitrary(u)?;
113        let b = Idx::arbitrary(u)?;
114        let (start, end) = if a < b {
115            (a, b)
116        } else if b < a {
117            (b, a)
118        } else {
119            return Err(arbitrary::Error::IncorrectFormat);
120        };
121        Ok(Self(start..end))
122    }
123}
124
125/// A macro to create a [`NonEmptyRange`] from a range expression, panicking if the range is empty.
126#[macro_export]
127macro_rules! non_empty_range {
128    ($start:expr, $end:expr) => {
129        $crate::range::NonEmptyRange::new($start..$end).expect("range must be non-empty")
130    };
131}
132
133#[cfg(test)]
134mod tests {
135    use super::*;
136    use commonware_codec::{DecodeExt, Encode};
137
138    #[test]
139    fn test_non_empty_range_valid() {
140        let r = NonEmptyRange::new(0u32..5).unwrap();
141        assert_eq!(r.start(), 0);
142        assert_eq!(r.end(), 5);
143        assert_eq!(Range::from(r), 0..5);
144    }
145
146    #[test]
147    fn test_non_empty_range_single_element() {
148        let r = NonEmptyRange::new(3u32..4).unwrap();
149        assert_eq!(r.start(), 3);
150        assert_eq!(r.end(), 4);
151    }
152
153    #[test]
154    fn test_non_empty_range_empty() {
155        assert_eq!(NonEmptyRange::new(5u32..5), Err(EmptyRange));
156        #[allow(clippy::reversed_empty_ranges)]
157        let reversed = NonEmptyRange::new(5u32..3);
158        assert_eq!(reversed, Err(EmptyRange));
159    }
160
161    #[test]
162    fn test_non_empty_range_into() {
163        let r = NonEmptyRange::new(1u32..10).unwrap();
164        let range: Range<u32> = r.into();
165        assert_eq!(range, 1..10);
166    }
167
168    #[test]
169    fn test_non_empty_range_debug() {
170        let r = NonEmptyRange::new(1u32..5).unwrap();
171        assert_eq!(format!("{r:?}"), "1..5");
172    }
173
174    #[test]
175    fn test_non_empty_range_iter() {
176        let r = NonEmptyRange::new(0u32..4).unwrap();
177        let items: Vec<_> = r.into_iter().collect();
178        assert_eq!(items, vec![0, 1, 2, 3]);
179    }
180
181    #[test]
182    fn test_non_empty_range_encode_decode() {
183        let r = NonEmptyRange::new(10u32..20).unwrap();
184        let encoded = r.encode();
185        let decoded = NonEmptyRange::<u32>::decode(encoded).unwrap();
186        assert_eq!(r, decoded);
187    }
188
189    #[test]
190    fn test_non_empty_range_decode_invalid() {
191        for (start, end) in [(20u32, 10u32), (5, 5)] {
192            let mut buf = Vec::new();
193            buf.extend_from_slice(&start.to_be_bytes());
194            buf.extend_from_slice(&end.to_be_bytes());
195            assert!(matches!(
196                NonEmptyRange::<u32>::decode(bytes::Bytes::from(buf)),
197                Err(CodecError::Invalid("NonEmptyRange", "start must be < end"))
198            ));
199        }
200    }
201
202    #[cfg(feature = "arbitrary")]
203    mod conformance {
204        use super::*;
205        use commonware_codec::conformance::CodecConformance;
206
207        commonware_conformance::conformance_tests! {
208            CodecConformance<NonEmptyRange<u32>>,
209            CodecConformance<NonEmptyRange<u64>>,
210        }
211    }
212}