Skip to main content

agsol_common/
max_len_vec.rs

1use super::{MaxLenResult, MaxSerializedLen, CONTENTS_FULL};
2use borsh::{BorshDeserialize, BorshSerialize};
3use std::convert::{From, TryFrom};
4
5// NOTE anyhow doesn't compile under bpf it seems
6
7#[repr(C)]
8#[derive(BorshDeserialize, BorshSerialize, Clone, Debug)]
9pub struct MaxLenVec<T, const N: usize> {
10    contents: Vec<T>,
11}
12
13impl<T, const N: usize> MaxSerializedLen for MaxLenVec<T, N>
14where
15    T: MaxSerializedLen,
16{
17    const MAX_SERIALIZED_LEN: usize = 4 + N * T::MAX_SERIALIZED_LEN;
18}
19
20impl<T, const N: usize> MaxLenVec<T, N> {
21    pub fn new() -> Self {
22        MaxLenVec {
23            contents: Vec::with_capacity(N),
24        }
25    }
26
27    pub fn is_full(&self) -> bool {
28        self.contents.len() == N
29    }
30
31    pub fn is_empty(&self) -> bool {
32        self.contents.is_empty()
33    }
34
35    pub fn len(&self) -> usize {
36        self.contents.len()
37    }
38
39    pub fn contents(&self) -> &[T] {
40        self.contents.as_slice()
41    }
42
43    pub fn contents_mut(&mut self) -> &mut [T] {
44        self.contents.as_mut_slice()
45    }
46
47    pub fn push(&mut self, elem: T) -> MaxLenResult {
48        if self.is_full() {
49            Err(CONTENTS_FULL)
50        } else {
51            self.contents.push(elem);
52            Ok(())
53        }
54    }
55
56    pub fn pop(&mut self) -> Option<T> {
57        self.contents.pop()
58    }
59
60    pub fn cyclic_push(&mut self, elem: T) {
61        if self.is_full() {
62            self.contents.remove(0);
63        }
64        self.contents.push(elem);
65    }
66
67    pub fn insert(&mut self, index: usize, value: T) -> MaxLenResult {
68        if self.is_full() {
69            Err(CONTENTS_FULL)
70        } else {
71            self.contents.insert(index, value);
72            Ok(())
73        }
74    }
75
76    pub fn remove(&mut self, index: usize) {
77        self.contents.remove(index);
78    }
79
80    pub fn get_last_element(&self) -> Option<&T> {
81        if self.is_empty() {
82            None
83        } else {
84            Some(&self.contents[self.contents.len() - 1])
85        }
86    }
87}
88
89impl<T, const N: usize> TryFrom<Vec<T>> for MaxLenVec<T, N> {
90    type Error = &'static str;
91
92    fn try_from(vec: Vec<T>) -> Result<Self, Self::Error> {
93        if vec.len() > N {
94            Err(CONTENTS_FULL)
95        } else {
96            Ok(Self { contents: vec })
97        }
98    }
99}
100
101impl<T, const N: usize> Default for MaxLenVec<T, N> {
102    fn default() -> Self {
103        Self::new()
104    }
105}
106
107impl<T, const N: usize> From<MaxLenVec<T, N>> for Vec<T> {
108    fn from(rhs: MaxLenVec<T, N>) -> Self {
109        rhs.contents
110    }
111}
112
113#[cfg(test)]
114mod test_max_len_vec {
115    use super::*;
116
117    const CAPACITY: usize = 5;
118    type TestVec = MaxLenVec<u8, CAPACITY>;
119
120    #[test]
121    fn initialization() {
122        assert_eq!(TestVec::new().contents.capacity(), CAPACITY);
123        let vec: Vec<u8> = vec![1, 2, 3, 4, 5];
124        assert!(TestVec::try_from(vec).is_ok());
125        let long_vec: Vec<u8> = vec![1, 2, 3, 4, 5, 6];
126        assert!(TestVec::try_from(long_vec).is_err());
127    }
128
129    #[test]
130    fn dynamic_updates() {
131        let mut vec = TestVec::new();
132        assert_eq!(vec.get_last_element(), None);
133        for i in 0..CAPACITY {
134            assert!(vec.push(i as u8).is_ok());
135        }
136        assert_eq!(vec.len(), CAPACITY);
137        assert!(vec.push(32).is_err());
138        vec.cyclic_push(32);
139        assert_eq!(vec.contents(), &[1, 2, 3, 4, 32]);
140        assert_eq!(vec.get_last_element(), Some(&32));
141        vec.pop();
142        vec.pop();
143        assert_eq!(vec.contents(), &[1, 2, 3]);
144        vec.cyclic_push(53);
145        assert_eq!(vec.contents(), &[1, 2, 3, 53]);
146        vec.cyclic_push(23);
147        vec.cyclic_push(33);
148        vec.cyclic_push(73);
149        assert_eq!(vec.contents(), &[3, 53, 23, 33, 73]);
150        assert_eq!(vec.get_last_element(), Some(&73));
151        assert!(vec.insert(3, 12).is_err());
152        vec.pop();
153        assert!(vec.insert(3, 12).is_ok());
154        assert_eq!(vec.contents(), &[3, 53, 23, 12, 33]);
155        vec.remove(2);
156        assert_eq!(vec.contents(), &[3, 53, 12, 33]);
157        vec.remove(1);
158        assert_eq!(vec.contents(), &[3, 12, 33]);
159        vec.remove(0);
160        assert_eq!(vec.contents(), &[12, 33]);
161        vec.remove(1);
162        assert_eq!(vec.contents(), &[12]);
163        vec.pop();
164        assert!(vec.is_empty());
165    }
166
167    #[test]
168    fn static_updates() {
169        let mut vec = TestVec::try_from(vec![3, 5, 2, 1, 4]).unwrap();
170        vec.contents_mut().sort_unstable();
171        assert_eq!(vec.contents(), &[1, 2, 3, 4, 5]);
172        vec.contents_mut()[2] = 10;
173        assert_eq!(vec.contents(), &[1, 2, 10, 4, 5]);
174
175        let std_vec = Vec::<u8>::from(vec);
176        assert_eq!(std_vec, vec![1, 2, 10, 4, 5]);
177    }
178
179    #[test]
180    fn max_len_vec_serialized_len() {
181        let mut test_vec = TestVec::new();
182        assert!(test_vec.try_to_vec().unwrap().len() <= TestVec::MAX_SERIALIZED_LEN);
183
184        for i in 0..4 {
185            assert!(test_vec.push(i).is_ok());
186        }
187        assert!(test_vec.try_to_vec().unwrap().len() <= TestVec::MAX_SERIALIZED_LEN);
188
189        assert!(test_vec.push(4).is_ok());
190        assert_eq!(
191            test_vec.try_to_vec().unwrap().len(),
192            TestVec::MAX_SERIALIZED_LEN
193        );
194    }
195}