1use crate::entropy::model::{MAX_SCALE_BITS, MIN_SCALE_BITS};
11use crate::error::{Error, Result};
12use crate::limits::Limits;
13
14pub const CODER_ORDER0_BYTE_RANS: u8 = 1;
16pub const CODER_VERSION_1: u16 = 1;
18
19pub const WIRE_HEADER_LEN: usize = 33;
21
22#[derive(Debug, Clone, PartialEq, Eq)]
24pub struct EntropyChannelDescriptor {
25 pub coder: u8,
27 pub coder_version: u16,
29 pub scale_bits: u8,
31 pub lane_count: u8,
33 pub model_id: u32,
35 pub symbol_count: u64,
37 pub decoded_length: u64,
39 pub initial_state: u32,
41 pub payload: Vec<u8>,
43}
44
45impl EntropyChannelDescriptor {
46 pub fn encode(&self) -> Result<Vec<u8>> {
51 let mut out = self.header_bytes()?;
52 out.extend_from_slice(&self.payload);
53 Ok(out)
54 }
55
56 pub fn header_bytes(&self) -> Result<Vec<u8>> {
64 let payload_len = u32::try_from(self.payload.len())
65 .map_err(|_| Error::resource_limit("entropy channel payload exceeds 4 GiB"))?;
66 let mut out = Vec::with_capacity(WIRE_HEADER_LEN);
67 out.push(self.coder);
68 out.extend_from_slice(&self.coder_version.to_le_bytes());
69 out.push(self.scale_bits);
70 out.push(self.lane_count);
71 out.extend_from_slice(&self.model_id.to_le_bytes());
72 out.extend_from_slice(&self.symbol_count.to_le_bytes());
73 out.extend_from_slice(&self.decoded_length.to_le_bytes());
74 out.extend_from_slice(&self.initial_state.to_le_bytes());
75 out.extend_from_slice(&payload_len.to_le_bytes());
76 Ok(out)
77 }
78
79 pub fn decode(bytes: &[u8], limits: Limits) -> Result<EntropyChannelDescriptor> {
82 if bytes.len() < WIRE_HEADER_LEN {
83 return Err(Error::invalid_container(
84 "entropy channel descriptor shorter than header",
85 ));
86 }
87 let coder = bytes[0];
88 if coder != CODER_ORDER0_BYTE_RANS {
89 return Err(Error::unsupported_feature(format!(
90 "entropy coder {coder} is not supported"
91 )));
92 }
93 let coder_version = u16::from_le_bytes([bytes[1], bytes[2]]);
94 if coder_version != CODER_VERSION_1 {
95 return Err(Error::unsupported_version(format!(
96 "entropy coder version {coder_version}, expected {CODER_VERSION_1}"
97 )));
98 }
99 let scale_bits = bytes[3];
100 if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&scale_bits) {
101 let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
102 return Err(Error::invalid_model(format!(
103 "scale_bits {scale_bits} outside {min}..={max}"
104 )));
105 }
106 let lane_count = bytes[4];
107 if lane_count != 1 {
108 return Err(Error::unsupported_feature(format!(
109 "lane_count {lane_count}, only single-lane channels are supported"
110 )));
111 }
112 let model_id = u32::from_le_bytes([bytes[5], bytes[6], bytes[7], bytes[8]]);
113 let symbol_count = u64::from_le_bytes([
114 bytes[9], bytes[10], bytes[11], bytes[12], bytes[13], bytes[14], bytes[15], bytes[16],
115 ]);
116 let decoded_length = u64::from_le_bytes([
117 bytes[17], bytes[18], bytes[19], bytes[20], bytes[21], bytes[22], bytes[23], bytes[24],
118 ]);
119 let initial_state = u32::from_le_bytes([bytes[25], bytes[26], bytes[27], bytes[28]]);
120 let payload_len = u32::from_le_bytes([bytes[29], bytes[30], bytes[31], bytes[32]]);
121
122 if symbol_count > limits.max_channel_symbols {
123 return Err(Error::resource_limit(format!(
124 "symbol_count {symbol_count} exceeds limit {}",
125 limits.max_channel_symbols
126 )));
127 }
128 if decoded_length > limits.max_output_bytes {
129 return Err(Error::resource_limit(format!(
130 "decoded_length {decoded_length} exceeds limit {}",
131 limits.max_output_bytes
132 )));
133 }
134 if payload_len > limits.max_record_len {
135 return Err(Error::resource_limit(format!(
136 "payload length {payload_len} exceeds limit {}",
137 limits.max_record_len
138 )));
139 }
140 let expected = WIRE_HEADER_LEN
141 .checked_add(payload_len as usize)
142 .ok_or_else(|| Error::resource_limit("entropy channel length overflow"))?;
143 if bytes.len() != expected {
144 return Err(Error::invalid_container(
145 "entropy channel payload length mismatch (trailing or truncated)",
146 ));
147 }
148
149 Ok(EntropyChannelDescriptor {
150 coder,
151 coder_version,
152 scale_bits,
153 lane_count,
154 model_id,
155 symbol_count,
156 decoded_length,
157 initial_state,
158 payload: bytes[WIRE_HEADER_LEN..].to_vec(),
159 })
160 }
161}
162
163#[cfg(test)]
164mod tests {
165 use super::*;
166 use crate::error::ErrorClass;
167
168 fn descriptor() -> EntropyChannelDescriptor {
169 EntropyChannelDescriptor {
170 coder: CODER_ORDER0_BYTE_RANS,
171 coder_version: CODER_VERSION_1,
172 scale_bits: 12,
173 lane_count: 1,
174 model_id: 7,
175 symbol_count: 1234,
176 decoded_length: 1234,
177 initial_state: 0xABCD_1234,
178 payload: vec![1, 2, 3, 4, 5, 6, 7, 8],
179 }
180 }
181
182 #[test]
183 fn wire_length_is_header_plus_payload() {
184 let d = descriptor();
185 let bytes = d.encode().unwrap();
186 assert_eq!(bytes.len(), 33 + d.payload.len());
187 }
188
189 #[test]
190 fn header_bytes_is_the_encode_prefix() {
191 let d = descriptor();
192 let header = d.header_bytes().unwrap();
193 assert_eq!(header.len(), WIRE_HEADER_LEN);
194 let full = d.encode().unwrap();
195 assert_eq!(&full[..WIRE_HEADER_LEN], &header[..]);
196 }
197
198 #[test]
199 fn header_bytes_depends_on_payload_length_not_bytes() {
200 let mut a = descriptor();
202 let mut b = descriptor();
203 b.payload[0] ^= 0xff;
204 assert_ne!(a.payload, b.payload);
205 assert_eq!(a.header_bytes().unwrap(), b.header_bytes().unwrap());
206 a.payload.push(0);
208 assert_ne!(a.header_bytes().unwrap(), b.header_bytes().unwrap());
209 }
210
211 #[test]
212 fn roundtrip() {
213 let d = descriptor();
214 let bytes = d.encode().unwrap();
215 let back = EntropyChannelDescriptor::decode(&bytes, Limits::DEFAULT).unwrap();
216 assert_eq!(back, d);
217 }
218
219 #[test]
220 fn empty_payload_roundtrips() {
221 let mut d = descriptor();
222 d.payload.clear();
223 let bytes = d.encode().unwrap();
224 assert_eq!(bytes.len(), 33);
225 let back = EntropyChannelDescriptor::decode(&bytes, Limits::DEFAULT).unwrap();
226 assert_eq!(back, d);
227 }
228
229 #[test]
230 fn rejects_unknown_coder() {
231 let mut d = descriptor();
232 d.coder = 2;
233 let bytes = d.encode().unwrap();
234 let e = EntropyChannelDescriptor::decode(&bytes, Limits::DEFAULT).unwrap_err();
235 assert_eq!(e.class(), ErrorClass::UnsupportedFeature);
236 }
237
238 #[test]
239 fn rejects_unknown_version() {
240 let mut d = descriptor();
241 d.coder_version = 2;
242 let bytes = d.encode().unwrap();
243 let e = EntropyChannelDescriptor::decode(&bytes, Limits::DEFAULT).unwrap_err();
244 assert_eq!(e.class(), ErrorClass::UnsupportedVersion);
245 }
246
247 #[test]
248 fn rejects_bad_scale_bits() {
249 for bits in [0u8, 16] {
250 let mut d = descriptor();
251 d.scale_bits = bits;
252 let bytes = d.encode().unwrap();
253 let e = EntropyChannelDescriptor::decode(&bytes, Limits::DEFAULT).unwrap_err();
254 assert_eq!(e.class(), ErrorClass::InvalidModel);
255 }
256 }
257
258 #[test]
259 fn rejects_multilane() {
260 let mut d = descriptor();
261 d.lane_count = 2;
262 let bytes = d.encode().unwrap();
263 let e = EntropyChannelDescriptor::decode(&bytes, Limits::DEFAULT).unwrap_err();
264 assert_eq!(e.class(), ErrorClass::UnsupportedFeature);
265 }
266
267 #[test]
268 fn rejects_truncated_header() {
269 let e = EntropyChannelDescriptor::decode(&[0u8; 32], Limits::DEFAULT).unwrap_err();
270 assert_eq!(e.class(), ErrorClass::InvalidContainer);
271 }
272
273 #[test]
274 fn rejects_trailing_bytes() {
275 let mut bytes = descriptor().encode().unwrap();
276 bytes.push(0);
277 let e = EntropyChannelDescriptor::decode(&bytes, Limits::DEFAULT).unwrap_err();
278 assert_eq!(e.class(), ErrorClass::InvalidContainer);
279 }
280
281 #[test]
282 fn rejects_truncated_payload() {
283 let mut bytes = descriptor().encode().unwrap();
284 bytes.pop();
285 let e = EntropyChannelDescriptor::decode(&bytes, Limits::DEFAULT).unwrap_err();
286 assert_eq!(e.class(), ErrorClass::InvalidContainer);
287 }
288
289 #[test]
290 fn enforces_payload_limit() {
291 let mut d = descriptor();
292 d.payload = vec![0u8; 100];
293 let bytes = d.encode().unwrap();
294 let limits = Limits {
295 max_record_len: 10,
296 ..Limits::DEFAULT
297 };
298 let e = EntropyChannelDescriptor::decode(&bytes, limits).unwrap_err();
299 assert_eq!(e.class(), ErrorClass::ResourceLimit);
300 }
301
302 #[test]
303 fn enforces_decoded_length_limit() {
304 let mut d = descriptor();
305 d.decoded_length = 1 << 30;
306 let bytes = d.encode().unwrap();
307 let limits = Limits {
308 max_output_bytes: 1 << 20,
309 ..Limits::DEFAULT
310 };
311 let e = EntropyChannelDescriptor::decode(&bytes, limits).unwrap_err();
312 assert_eq!(e.class(), ErrorClass::ResourceLimit);
313 }
314
315 #[test]
316 fn enforces_symbol_count_limit() {
317 let mut d = descriptor();
318 d.symbol_count = 1 << 30;
319 let bytes = d.encode().unwrap();
320 let limits = Limits {
321 max_channel_symbols: 1 << 20,
322 ..Limits::DEFAULT
323 };
324 let e = EntropyChannelDescriptor::decode(&bytes, limits).unwrap_err();
325 assert_eq!(e.class(), ErrorClass::ResourceLimit);
326 }
327}