use borsh::{BorshDeserialize, BorshSerialize};
use serde::{
Deserialize, Serialize,
de::{self, Deserializer, SeqAccess, Visitor},
ser::{self, SerializeTuple, Serializer},
};
use std::{convert::TryFrom, fmt, marker::PhantomData, vec::Vec};
pub struct ShortU16(pub u16);
impl Serialize for ShortU16 {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut seq = serializer.serialize_tuple(1)?;
let mut rem_val = self.0;
loop {
let mut elem = (rem_val & 0x7f) as u8;
rem_val >>= 7;
if rem_val == 0 {
seq.serialize_element(&elem)?;
break;
} else {
elem |= 0x80;
seq.serialize_element(&elem)?;
}
}
seq.end()
}
}
enum VisitStatus {
Done(u16),
More(u16),
}
#[derive(Debug)]
enum VisitError {
TooLong(usize),
TooShort(usize),
Overflow(u32),
Alias,
ByteThreeContinues,
}
impl VisitError {
fn into_de_error<'de, A>(self) -> A::Error
where
A: SeqAccess<'de>,
{
match self {
VisitError::TooLong(len) => de::Error::invalid_length(len, &"three or fewer bytes"),
VisitError::TooShort(len) => de::Error::invalid_length(len, &"more bytes"),
VisitError::Overflow(val) => de::Error::invalid_value(
de::Unexpected::Unsigned(val as u64),
&"a value in the range [0, 65535]",
),
VisitError::Alias => de::Error::invalid_value(
de::Unexpected::Other("alias encoding"),
&"strict form encoding",
),
VisitError::ByteThreeContinues => de::Error::invalid_value(
de::Unexpected::Other("continue signal on byte-three"),
&"a terminal signal on or before byte-three",
),
}
}
}
type VisitResult = Result<VisitStatus, VisitError>;
const MAX_ENCODING_LENGTH: usize = 3;
fn visit_byte(elem: u8, val: u16, nth_byte: usize) -> VisitResult {
if elem == 0 && nth_byte != 0 {
return Err(VisitError::Alias);
}
let val = u32::from(val);
let elem = u32::from(elem);
let elem_val = elem & 0x7f;
let elem_done = (elem & 0x80) == 0;
if nth_byte >= MAX_ENCODING_LENGTH {
return Err(VisitError::TooLong(nth_byte.saturating_add(1)));
} else if nth_byte == MAX_ENCODING_LENGTH.saturating_sub(1) && !elem_done {
return Err(VisitError::ByteThreeContinues);
}
let shift = u32::try_from(nth_byte)
.unwrap_or(u32::MAX)
.saturating_mul(7);
let elem_val = elem_val.checked_shl(shift).unwrap_or(u32::MAX);
let new_val = val | elem_val;
let val = u16::try_from(new_val).map_err(|_| VisitError::Overflow(new_val))?;
if elem_done {
Ok(VisitStatus::Done(val))
} else {
Ok(VisitStatus::More(val))
}
}
struct ShortU16Visitor;
impl<'de> Visitor<'de> for ShortU16Visitor {
type Value = ShortU16;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("a ShortU16")
}
fn visit_seq<A>(self, mut seq: A) -> Result<ShortU16, A::Error>
where
A: SeqAccess<'de>,
{
let mut val: u16 = 0;
for nth_byte in 0..MAX_ENCODING_LENGTH {
let elem: u8 = seq.next_element()?.ok_or_else(|| {
VisitError::TooShort(nth_byte.saturating_add(1)).into_de_error::<A>()
})?;
match visit_byte(elem, val, nth_byte).map_err(|e| e.into_de_error::<A>())? {
VisitStatus::Done(new_val) => return Ok(ShortU16(new_val)),
VisitStatus::More(new_val) => val = new_val,
}
}
Err(VisitError::ByteThreeContinues.into_de_error::<A>())
}
}
impl<'de> Deserialize<'de> for ShortU16 {
fn deserialize<D>(deserializer: D) -> Result<ShortU16, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_tuple(3, ShortU16Visitor)
}
}
pub fn serialize<S: Serializer, T: Serialize>(
elements: &[T],
serializer: S,
) -> Result<S::Ok, S::Error> {
let mut seq = serializer.serialize_tuple(1)?;
let len = elements.len();
if len > u16::MAX as usize {
return Err(ser::Error::custom("length larger than u16"));
}
let short_len = ShortU16(len as u16);
seq.serialize_element(&short_len)?;
for element in elements {
seq.serialize_element(element)?;
}
seq.end()
}
struct ShortVecVisitor<T> {
_t: PhantomData<T>,
}
impl<'de, T> Visitor<'de> for ShortVecVisitor<T>
where
T: Deserialize<'de>,
{
type Value = Vec<T>;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("a Vec with a multi-byte length")
}
fn visit_seq<A>(self, mut seq: A) -> Result<Vec<T>, A::Error>
where
A: SeqAccess<'de>,
{
let short_len: ShortU16 = seq
.next_element()?
.ok_or_else(|| de::Error::invalid_length(0, &self))?;
let len = short_len.0 as usize;
let mut result = Vec::with_capacity(len.min(1024));
for i in 0..len {
let elem = seq
.next_element()?
.ok_or_else(|| de::Error::invalid_length(i, &self))?;
result.push(elem);
}
Ok(result)
}
}
pub fn encode_length_to_compact_u16_bytes(len: usize) -> Result<Vec<u8>, String> {
if len > u16::MAX as usize {
return Err(format!(
"Length {len} exceeds u16::MAX, cannot encode as Compact-U16"
));
}
let mut bytes = Vec::new();
let mut rem_val = len as u16; loop {
let mut elem = (rem_val & 0x7f) as u8;
rem_val >>= 7;
if rem_val == 0 {
bytes.push(elem);
break;
} else {
elem |= 0x80; bytes.push(elem);
}
}
Ok(bytes)
}
pub fn decode_compact_u16_len(bytes: &[u8]) -> Result<(usize, usize), &'static str> {
if bytes.is_empty() {
return Err("Cannot decode length from empty slice");
}
let mut len: usize = 0;
let mut size_of_len_encoding: usize = 0;
loop {
if size_of_len_encoding >= bytes.len() {
return Err("Byte slice too short for compact u16 length (within loop)");
}
let current_byte = bytes[size_of_len_encoding];
len |= (current_byte as usize & 0x7F) << (size_of_len_encoding * 7);
size_of_len_encoding += 1;
if (current_byte & 0x80) == 0 {
break;
}
if size_of_len_encoding >= 3 && (current_byte & 0x80) != 0 {
return Err("Compact u16 length encoding too long (max 3 bytes for u16 values)");
}
}
if len > u16::MAX as usize {
return Err("Decoded length exceeds u16::MAX for compact-u16 encoding");
}
Ok((len, size_of_len_encoding))
}
pub fn deserialize<'de, D, T>(deserializer: D) -> Result<Vec<T>, D::Error>
where
D: Deserializer<'de>,
T: Deserialize<'de>,
{
deserializer.deserialize_seq(ShortVecVisitor { _t: PhantomData })
}
impl<T> Serialize for ShortVec<T>
where
T: Serialize,
{
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
self::serialize(&self.inner, serializer)
}
}
impl<'de, T> Deserialize<'de> for ShortVec<T>
where
T: Deserialize<'de>,
{
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
Ok(ShortVec {
inner: self::deserialize(deserializer)?,
})
}
}
#[derive(BorshSerialize, BorshDeserialize)] pub struct ShortVec<T> {
pub inner: Vec<T>,
}
impl<T: Clone> Clone for ShortVec<T> {
fn clone(&self) -> Self {
ShortVec {
inner: self.inner.clone(),
}
}
}
impl<T: std::fmt::Debug> std::fmt::Debug for ShortVec<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_tuple("ShortVec").field(&self.inner).finish()
}
}
impl<T: PartialEq> PartialEq for ShortVec<T> {
fn eq(&self, other: &Self) -> bool {
self.inner == other.inner
}
}
impl<T> ShortVec<T> {
pub fn new(inner: Vec<T>) -> Self {
ShortVec { inner }
}
pub fn into_inner(self) -> Vec<T> {
self.inner
}
pub fn as_inner(&self) -> &Vec<T> {
&self.inner
}
pub fn as_mut_inner(&mut self) -> &mut Vec<T> {
&mut self.inner
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn decode_compact_u16_len_rejects_values_above_u16_max() {
let bytes = [0xFF, 0xFF, 0x7F];
let result = decode_compact_u16_len(&bytes);
assert!(
result.is_err(),
"expected decode_compact_u16_len to reject a length exceeding u16::MAX, got {result:?}"
);
}
#[test]
fn decode_compact_u16_len_accepts_u16_max() {
let bytes = [0xFF, 0xFF, 0x03];
let (len, consumed) = decode_compact_u16_len(&bytes).unwrap();
assert_eq!(len, u16::MAX as usize);
assert_eq!(consumed, 3);
}
}