Skip to main content

solana_primitives/
short_vec.rs

1// Compact serde-encoding of vectors with small length.
2use 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
11/// Represents a ShortU16.
12///
13/// Same as u16, but serialized with 1 to 3 bytes. If the value is above
14/// 0x7f, the top bit is set and the remaining value is stored in the next
15/// bytes. Each byte follows the same pattern until the 3rd byte. The 3rd
16/// byte, if needed, uses all 8 bits to store the last byte of the original
17/// value.
18pub 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        // Pass a non-zero value to serialize_tuple() so that serde_json will
26        // generate an open bracket.
27        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        // Decodes an unsigned 16 bit integer one-to-one encoded as follows:
124        // 1 byte : 0xxxxxxx => 00000000 0xxxxxxx : 0 - 127
125        // 2 bytes : 1xxxxxxx 0yyyyyyy => 00yyyyyy yxxxxxxx : 128 - 16,383
126        // 3 bytes : 1xxxxxxx 1yyyyyyy 000000zz => zzyyyyyy yxxxxxxx : 16,384 - 65,535
127        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
150/// If you don't want to use the ShortVec newtype, you can do ShortVec
151/// serialization on an ordinary vector with the following field annotation:
152///
153/// #[serde(with = "short_vec")]
154pub fn serialize<S: Serializer, T: Serialize>(
155    elements: &[T],
156    serializer: S,
157) -> Result<S::Ok, S::Error> {
158    // Pass a non-zero value to serialize_tuple() so that serde_json will
159    // generate an open bracket.
160    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        // Cap the reservation so a malformed payload claiming a large count can't allocate big.
194        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
205// Helper function to encode a usize length into Compact-U16 format bytes.
206// Returns a Vec<u8> with the encoded length or an Err if length is too large for u16.
207pub 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; // Safe to cast now
215    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; // More bytes to follow, set MSB
223            bytes.push(elem);
224        }
225    }
226    Ok(bytes)
227}
228
229// Helper function to decode Compact-U16 length
230// Returns Ok((length, bytes_consumed)) or Err(message)
231pub 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            // MSB is 0, this is the last byte for the length
246            break;
247        }
248        // According to Solana's short_vec.rs, max 3 bytes for u16 values (up to 65535)
249        // 1 byte for 0-127
250        // 2 bytes for 128 - 16383
251        // 3 bytes for 16384 - 65535
252        if size_of_len_encoding >= 3 && (current_byte & 0x80) != 0 {
253            // If we've read 3 bytes and the 3rd byte still has MSB set, it's an invalid encoding for u16.
254            // Or if we are about to read a 4th byte for a u16 value.
255            // This check is to prevent overruns for u16. If len can be > u16::MAX, this check changes.
256            // For typical Solana message elements, lengths are expected to fit u16.
257            return Err("Compact u16 length encoding too long (max 3 bytes for u16 values)");
258        }
259    }
260    // A 3rd byte can still contribute up to 2,097,151; every caller expects u16-bounded.
261    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
267/// If you don't want to use the ShortVec newtype, you can do ShortVec
268/// deserialization on an ordinary vector with the following field annotation:
269///
270/// #[serde(with = "short_vec::deserialize")]
271pub 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
279/// A newtype to provide Compact-U16 (AKA short_vec) serialization for `Vec<T>`
280impl<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        // Calls the module-level serialize function
289        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        // Calls the module-level deserialize function
302        Ok(ShortVec {
303            inner: self::deserialize(deserializer)?,
304        })
305    }
306}
307
308#[derive(BorshSerialize, BorshDeserialize)] // Derives for Borsh
309pub struct ShortVec<T> {
310    pub inner: Vec<T>,
311}
312
313// Manual impls for common traits, forwarding to Vec<T>
314// We need to be careful with bounds if T itself is complex.
315impl<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
335// Add a constructor and a way to get the inner Vec
336impl<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    // Optional: provide a way to borrow the inner vec
346    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// We still need to ensure T itself is bound correctly where ShortVec<T> is used.
356// For Borsh: T must be BorshSerialize + BorshDeserialize.
357// For Serde (via our custom impls): T must be Serialize + Deserialize<'de>.
358// The derive for BorshSerialize/Deserialize on ShortVec<T> will require T to also implement them for Vec<T>.
359
360#[cfg(test)]
361mod tests {
362    use super::*;
363
364    #[test]
365    fn decode_compact_u16_len_rejects_values_above_u16_max() {
366        // Decodes to 2,097,151 without the u16 bound.
367        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}