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}