Skip to main content

radixdb_plugin/
codec.rs

1use crate::{PluginError, PluginResult};
2
3#[derive(Debug)]
4pub struct CodecWriter {
5    bytes: Vec<u8>,
6    max_bytes: usize,
7}
8
9impl CodecWriter {
10    pub fn new(max_bytes: usize) -> Self {
11        Self {
12            bytes: Vec::new(),
13            max_bytes,
14        }
15    }
16
17    pub fn write(&mut self, bytes: &[u8]) -> PluginResult<()> {
18        let next = self
19            .bytes
20            .len()
21            .checked_add(bytes.len())
22            .ok_or_else(|| PluginError::limit_exceeded("codec output size overflow"))?;
23        if next > self.max_bytes {
24            return Err(PluginError::limit_exceeded(
25                "codec output exceeds declared max_bytes",
26            ));
27        }
28        self.bytes.extend_from_slice(bytes);
29        Ok(())
30    }
31
32    pub fn len(&self) -> usize {
33        self.bytes.len()
34    }
35
36    pub fn is_empty(&self) -> bool {
37        self.bytes.is_empty()
38    }
39
40    pub fn into_bytes(self) -> Vec<u8> {
41        self.bytes
42    }
43}
44
45#[derive(Debug, Clone, Copy)]
46pub struct CodecReader<'a> {
47    bytes: &'a [u8],
48    offset: usize,
49}
50
51impl<'a> CodecReader<'a> {
52    pub fn new(bytes: &'a [u8]) -> Self {
53        Self { bytes, offset: 0 }
54    }
55
56    pub fn read(&mut self, len: usize) -> PluginResult<&'a [u8]> {
57        let end = self
58            .offset
59            .checked_add(len)
60            .ok_or_else(|| PluginError::invalid_input("codec input offset overflow"))?;
61        let result = self
62            .bytes
63            .get(self.offset..end)
64            .ok_or_else(|| PluginError::invalid_input("truncated canonical value"))?;
65        self.offset = end;
66        Ok(result)
67    }
68
69    pub fn finish(self) -> PluginResult<()> {
70        if self.offset != self.bytes.len() {
71            return Err(PluginError::invalid_input(
72                "canonical value has trailing bytes",
73            ));
74        }
75        Ok(())
76    }
77}
78
79pub trait ManualCodec<T>: Send + Sync + 'static {
80    fn encode(value: &T, output: &mut CodecWriter) -> PluginResult<()>;
81    fn decode(input: &mut CodecReader<'_>) -> PluginResult<T>;
82    fn corpus() -> Vec<T>;
83}
84
85#[doc(hidden)]
86pub trait CanonicalField: Sized + Clone {
87    fn encode_field(&self, output: &mut CodecWriter) -> PluginResult<()>;
88    fn decode_field(input: &mut CodecReader<'_>) -> PluginResult<Self>;
89    fn edge_values() -> Vec<Self>;
90}
91
92macro_rules! integer_field {
93    ($type:ty) => {
94        impl CanonicalField for $type {
95            fn encode_field(&self, output: &mut CodecWriter) -> PluginResult<()> {
96                output.write(&self.to_le_bytes())
97            }
98
99            fn decode_field(input: &mut CodecReader<'_>) -> PluginResult<Self> {
100                let bytes: [u8; std::mem::size_of::<Self>()] = input
101                    .read(std::mem::size_of::<Self>())?
102                    .try_into()
103                    .expect("fixed-width slice length was checked");
104                Ok(Self::from_le_bytes(bytes))
105            }
106
107            fn edge_values() -> Vec<Self> {
108                vec![Self::MIN, 0, Self::MAX]
109            }
110        }
111    };
112}
113
114integer_field!(i8);
115integer_field!(i16);
116integer_field!(i32);
117integer_field!(i64);
118integer_field!(u8);
119integer_field!(u16);
120integer_field!(u32);
121integer_field!(u64);
122
123impl CanonicalField for bool {
124    fn encode_field(&self, output: &mut CodecWriter) -> PluginResult<()> {
125        output.write(&[u8::from(*self)])
126    }
127
128    fn decode_field(input: &mut CodecReader<'_>) -> PluginResult<Self> {
129        match input.read(1)?[0] {
130            0 => Ok(false),
131            1 => Ok(true),
132            _ => Err(PluginError::invalid_input(
133                "boolean field must be encoded as 0 or 1",
134            )),
135        }
136    }
137
138    fn edge_values() -> Vec<Self> {
139        vec![false, true]
140    }
141}
142
143impl CanonicalField for f32 {
144    fn encode_field(&self, output: &mut CodecWriter) -> PluginResult<()> {
145        output.write(&self.to_bits().to_le_bytes())
146    }
147
148    fn decode_field(input: &mut CodecReader<'_>) -> PluginResult<Self> {
149        let bytes = input.read(4)?.try_into().expect("fixed width");
150        Ok(Self::from_bits(u32::from_le_bytes(bytes)))
151    }
152
153    fn edge_values() -> Vec<Self> {
154        vec![
155            Self::NEG_INFINITY,
156            -0.0,
157            0.0,
158            Self::INFINITY,
159            Self::from_bits(0x7fc0_0001),
160        ]
161    }
162}
163
164impl CanonicalField for f64 {
165    fn encode_field(&self, output: &mut CodecWriter) -> PluginResult<()> {
166        output.write(&self.to_bits().to_le_bytes())
167    }
168
169    fn decode_field(input: &mut CodecReader<'_>) -> PluginResult<Self> {
170        let bytes = input.read(8)?.try_into().expect("fixed width");
171        Ok(Self::from_bits(u64::from_le_bytes(bytes)))
172    }
173
174    fn edge_values() -> Vec<Self> {
175        vec![
176            Self::NEG_INFINITY,
177            -0.0,
178            0.0,
179            Self::INFINITY,
180            Self::from_bits(0x7ff8_0000_0000_0001),
181        ]
182    }
183}
184
185impl<T: CanonicalField, const N: usize> CanonicalField for [T; N] {
186    fn encode_field(&self, output: &mut CodecWriter) -> PluginResult<()> {
187        for value in self {
188            value.encode_field(output)?;
189        }
190        Ok(())
191    }
192
193    fn decode_field(input: &mut CodecReader<'_>) -> PluginResult<Self> {
194        let mut values = Vec::with_capacity(N);
195        for _ in 0..N {
196            values.push(T::decode_field(input)?);
197        }
198        values
199            .try_into()
200            .map_err(|_| PluginError::internal("fixed array decoder length mismatch"))
201    }
202
203    fn edge_values() -> Vec<Self> {
204        T::edge_values()
205            .into_iter()
206            .map(|value| std::array::from_fn(|_| value.clone()))
207            .collect()
208    }
209}
210
211#[doc(hidden)]
212pub fn encode_sequence<T: CanonicalField>(
213    values: &[T],
214    max_items: usize,
215    max_bytes: usize,
216    output: &mut CodecWriter,
217) -> PluginResult<()> {
218    if values.len() > max_items || values.len() > u32::MAX as usize {
219        return Err(PluginError::limit_exceeded(
220            "sequence exceeds declared max_items",
221        ));
222    }
223    let before = output.len();
224    output.write(&(values.len() as u32).to_le_bytes())?;
225    if output.len() - before > max_bytes {
226        return Err(PluginError::limit_exceeded(
227            "sequence exceeds declared max_bytes",
228        ));
229    }
230    for value in values {
231        value.encode_field(output)?;
232        if output.len() - before > max_bytes {
233            return Err(PluginError::limit_exceeded(
234                "sequence exceeds declared max_bytes",
235            ));
236        }
237    }
238    Ok(())
239}
240
241#[doc(hidden)]
242pub fn decode_sequence<T: CanonicalField>(
243    input: &mut CodecReader<'_>,
244    max_items: usize,
245    max_bytes: usize,
246) -> PluginResult<Vec<T>> {
247    let count = u32::from_le_bytes(input.read(4)?.try_into().expect("fixed width")) as usize;
248    if count > max_items {
249        return Err(PluginError::limit_exceeded(
250            "sequence exceeds declared max_items",
251        ));
252    }
253    let mut output = CodecWriter::new(max_bytes);
254    output.write(&(count as u32).to_le_bytes())?;
255    let mut values = Vec::with_capacity(count);
256    for _ in 0..count {
257        let value = T::decode_field(input)?;
258        value.encode_field(&mut output)?;
259        values.push(value);
260    }
261    Ok(values)
262}