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