1use borsh::{BorshDeserialize, BorshSerialize};
3
4use serde::{
5 Deserialize, Serialize,
6 de::{self, Deserializer, SeqAccess, Visitor},
7 ser::{self, SerializeTuple, Serializer},
8};
9use std::{convert::TryFrom, fmt, marker::PhantomData, vec::Vec};
10
11pub struct ShortU16(pub u16);
19
20impl Serialize for ShortU16 {
21 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
22 where
23 S: Serializer,
24 {
25 let mut seq = serializer.serialize_tuple(1)?;
28 let mut rem_val = self.0;
29 loop {
30 let mut elem = (rem_val & 0x7f) as u8;
31 rem_val >>= 7;
32 if rem_val == 0 {
33 seq.serialize_element(&elem)?;
34 break;
35 } else {
36 elem |= 0x80;
37 seq.serialize_element(&elem)?;
38 }
39 }
40 seq.end()
41 }
42}
43
44enum VisitStatus {
45 Done(u16),
46 More(u16),
47}
48
49#[derive(Debug)]
50enum VisitError {
51 TooLong(usize),
52 TooShort(usize),
53 Overflow(u32),
54 Alias,
55 ByteThreeContinues,
56}
57
58impl VisitError {
59 fn into_de_error<'de, A>(self) -> A::Error
60 where
61 A: SeqAccess<'de>,
62 {
63 match self {
64 VisitError::TooLong(len) => de::Error::invalid_length(len, &"three or fewer bytes"),
65 VisitError::TooShort(len) => de::Error::invalid_length(len, &"more bytes"),
66 VisitError::Overflow(val) => de::Error::invalid_value(
67 de::Unexpected::Unsigned(val as u64),
68 &"a value in the range [0, 65535]",
69 ),
70 VisitError::Alias => de::Error::invalid_value(
71 de::Unexpected::Other("alias encoding"),
72 &"strict form encoding",
73 ),
74 VisitError::ByteThreeContinues => de::Error::invalid_value(
75 de::Unexpected::Other("continue signal on byte-three"),
76 &"a terminal signal on or before byte-three",
77 ),
78 }
79 }
80}
81
82type VisitResult = Result<VisitStatus, VisitError>;
83
84const MAX_ENCODING_LENGTH: usize = 3;
85
86fn visit_byte(elem: u8, val: u16, nth_byte: usize) -> VisitResult {
87 if elem == 0 && nth_byte != 0 {
88 return Err(VisitError::Alias);
89 }
90 let val = u32::from(val);
91 let elem = u32::from(elem);
92 let elem_val = elem & 0x7f;
93 let elem_done = (elem & 0x80) == 0;
94 if nth_byte >= MAX_ENCODING_LENGTH {
95 return Err(VisitError::TooLong(nth_byte.saturating_add(1)));
96 } else if nth_byte == MAX_ENCODING_LENGTH.saturating_sub(1) && !elem_done {
97 return Err(VisitError::ByteThreeContinues);
98 }
99 let shift = u32::try_from(nth_byte)
100 .unwrap_or(u32::MAX)
101 .saturating_mul(7);
102 let elem_val = elem_val.checked_shl(shift).unwrap_or(u32::MAX);
103 let new_val = val | elem_val;
104 let val = u16::try_from(new_val).map_err(|_| VisitError::Overflow(new_val))?;
105 if elem_done {
106 Ok(VisitStatus::Done(val))
107 } else {
108 Ok(VisitStatus::More(val))
109 }
110}
111
112struct ShortU16Visitor;
113
114impl<'de> Visitor<'de> for ShortU16Visitor {
115 type Value = ShortU16;
116 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
117 formatter.write_str("a ShortU16")
118 }
119 fn visit_seq<A>(self, mut seq: A) -> Result<ShortU16, A::Error>
120 where
121 A: SeqAccess<'de>,
122 {
123 let mut val: u16 = 0;
128 for nth_byte in 0..MAX_ENCODING_LENGTH {
129 let elem: u8 = seq.next_element()?.ok_or_else(|| {
130 VisitError::TooShort(nth_byte.saturating_add(1)).into_de_error::<A>()
131 })?;
132 match visit_byte(elem, val, nth_byte).map_err(|e| e.into_de_error::<A>())? {
133 VisitStatus::Done(new_val) => return Ok(ShortU16(new_val)),
134 VisitStatus::More(new_val) => val = new_val,
135 }
136 }
137 Err(VisitError::ByteThreeContinues.into_de_error::<A>())
138 }
139}
140
141impl<'de> Deserialize<'de> for ShortU16 {
142 fn deserialize<D>(deserializer: D) -> Result<ShortU16, D::Error>
143 where
144 D: Deserializer<'de>,
145 {
146 deserializer.deserialize_tuple(3, ShortU16Visitor)
147 }
148}
149
150pub fn serialize<S: Serializer, T: Serialize>(
155 elements: &[T],
156 serializer: S,
157) -> Result<S::Ok, S::Error> {
158 let mut seq = serializer.serialize_tuple(1)?;
161 let len = elements.len();
162 if len > u16::MAX as usize {
163 return Err(ser::Error::custom("length larger than u16"));
164 }
165 let short_len = ShortU16(len as u16);
166 seq.serialize_element(&short_len)?;
167 for element in elements {
168 seq.serialize_element(element)?;
169 }
170 seq.end()
171}
172
173struct ShortVecVisitor<T> {
174 _t: PhantomData<T>,
175}
176
177impl<'de, T> Visitor<'de> for ShortVecVisitor<T>
178where
179 T: Deserialize<'de>,
180{
181 type Value = Vec<T>;
182 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
183 formatter.write_str("a Vec with a multi-byte length")
184 }
185 fn visit_seq<A>(self, mut seq: A) -> Result<Vec<T>, A::Error>
186 where
187 A: SeqAccess<'de>,
188 {
189 let short_len: ShortU16 = seq
190 .next_element()?
191 .ok_or_else(|| de::Error::invalid_length(0, &self))?;
192 let len = short_len.0 as usize;
193 let mut result = Vec::with_capacity(len.min(1024));
195 for i in 0..len {
196 let elem = seq
197 .next_element()?
198 .ok_or_else(|| de::Error::invalid_length(i, &self))?;
199 result.push(elem);
200 }
201 Ok(result)
202 }
203}
204
205pub fn encode_length_to_compact_u16_bytes(len: usize) -> Result<Vec<u8>, String> {
208 if len > u16::MAX as usize {
209 return Err(format!(
210 "Length {len} exceeds u16::MAX, cannot encode as Compact-U16"
211 ));
212 }
213 let mut bytes = Vec::new();
214 let mut rem_val = len as u16; loop {
216 let mut elem = (rem_val & 0x7f) as u8;
217 rem_val >>= 7;
218 if rem_val == 0 {
219 bytes.push(elem);
220 break;
221 } else {
222 elem |= 0x80; bytes.push(elem);
224 }
225 }
226 Ok(bytes)
227}
228
229pub fn decode_compact_u16_len(bytes: &[u8]) -> Result<(usize, usize), &'static str> {
232 if bytes.is_empty() {
233 return Err("Cannot decode length from empty slice");
234 }
235 let mut len: usize = 0;
236 let mut size_of_len_encoding: usize = 0;
237 loop {
238 if size_of_len_encoding >= bytes.len() {
239 return Err("Byte slice too short for compact u16 length (within loop)");
240 }
241 let current_byte = bytes[size_of_len_encoding];
242 len |= (current_byte as usize & 0x7F) << (size_of_len_encoding * 7);
243 size_of_len_encoding += 1;
244 if (current_byte & 0x80) == 0 {
245 break;
247 }
248 if size_of_len_encoding >= 3 && (current_byte & 0x80) != 0 {
253 return Err("Compact u16 length encoding too long (max 3 bytes for u16 values)");
258 }
259 }
260 if len > u16::MAX as usize {
262 return Err("Decoded length exceeds u16::MAX for compact-u16 encoding");
263 }
264 Ok((len, size_of_len_encoding))
265}
266
267pub fn deserialize<'de, D, T>(deserializer: D) -> Result<Vec<T>, D::Error>
272where
273 D: Deserializer<'de>,
274 T: Deserialize<'de>,
275{
276 deserializer.deserialize_seq(ShortVecVisitor { _t: PhantomData })
277}
278
279impl<T> Serialize for ShortVec<T>
281where
282 T: Serialize,
283{
284 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
285 where
286 S: Serializer,
287 {
288 self::serialize(&self.inner, serializer)
290 }
291}
292
293impl<'de, T> Deserialize<'de> for ShortVec<T>
294where
295 T: Deserialize<'de>,
296{
297 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
298 where
299 D: Deserializer<'de>,
300 {
301 Ok(ShortVec {
303 inner: self::deserialize(deserializer)?,
304 })
305 }
306}
307
308#[derive(BorshSerialize, BorshDeserialize)] pub struct ShortVec<T> {
310 pub inner: Vec<T>,
311}
312
313impl<T: Clone> Clone for ShortVec<T> {
316 fn clone(&self) -> Self {
317 ShortVec {
318 inner: self.inner.clone(),
319 }
320 }
321}
322
323impl<T: std::fmt::Debug> std::fmt::Debug for ShortVec<T> {
324 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
325 f.debug_tuple("ShortVec").field(&self.inner).finish()
326 }
327}
328
329impl<T: PartialEq> PartialEq for ShortVec<T> {
330 fn eq(&self, other: &Self) -> bool {
331 self.inner == other.inner
332 }
333}
334
335impl<T> ShortVec<T> {
337 pub fn new(inner: Vec<T>) -> Self {
338 ShortVec { inner }
339 }
340
341 pub fn into_inner(self) -> Vec<T> {
342 self.inner
343 }
344
345 pub fn as_inner(&self) -> &Vec<T> {
347 &self.inner
348 }
349
350 pub fn as_mut_inner(&mut self) -> &mut Vec<T> {
351 &mut self.inner
352 }
353}
354
355#[cfg(test)]
361mod tests {
362 use super::*;
363
364 #[test]
365 fn decode_compact_u16_len_rejects_values_above_u16_max() {
366 let bytes = [0xFF, 0xFF, 0x7F];
368 let result = decode_compact_u16_len(&bytes);
369 assert!(
370 result.is_err(),
371 "expected decode_compact_u16_len to reject a length exceeding u16::MAX, got {result:?}"
372 );
373 }
374
375 #[test]
376 fn decode_compact_u16_len_accepts_u16_max() {
377 let bytes = [0xFF, 0xFF, 0x03];
378 let (len, consumed) = decode_compact_u16_len(&bytes).unwrap();
379 assert_eq!(len, u16::MAX as usize);
380 assert_eq!(consumed, 3);
381 }
382}