1use std::ops::RangeInclusive;
2
3use bytes::{Buf, BufMut, Bytes, BytesMut};
4use moq_net::VarInt;
5
6use crate::path::{check_id, check_range};
7use crate::{Error, Result, VERSION};
8
9#[derive(Debug, Clone, PartialEq, Eq)]
11pub struct Group {
12 pub sequence: u64,
14 pub frames: Vec<Frame>,
16}
17
18#[derive(Debug, Clone, PartialEq, Eq)]
20pub struct Frame {
21 pub timestamp: u64,
23 pub payload: Bytes,
25}
26
27#[derive(Debug, Clone, PartialEq, Eq)]
29pub struct Object {
30 pub groups: Vec<Group>,
32}
33
34impl Object {
35 pub fn bounds(&self) -> Result<RangeInclusive<u64>> {
37 validate(&self.groups)?;
38 Ok(self.groups[0].sequence..=self.groups.last().unwrap().sequence)
39 }
40
41 pub fn encode(&self) -> Result<Bytes> {
43 validate(&self.groups)?;
44
45 let mut table = BytesMut::new();
46 write_varint(&mut table, VERSION)?;
47 write_varint(&mut table, self.groups.len() as u64)?;
48
49 let mut payload = BytesMut::new();
50 let mut prev = None;
51 for group in &self.groups {
52 let delta = match prev {
53 None => group.sequence,
54 Some(previous) => group
55 .sequence
56 .checked_sub(previous)
57 .ok_or(Error::Overflow)?
58 .checked_sub(1)
59 .ok_or(Error::Overflow)?,
60 };
61 prev = Some(group.sequence);
62 write_varint(&mut table, delta)?;
63 write_varint(&mut table, group.frames.len() as u64)?;
64 for frame in &group.frames {
65 let offset = u64::try_from(payload.len()).map_err(|_| Error::Overflow)?;
66 let length = u64::try_from(frame.payload.len()).map_err(|_| Error::Overflow)?;
67 write_varint(&mut table, frame.timestamp)?;
68 write_varint(&mut table, offset)?;
69 write_varint(&mut table, length)?;
70 payload.extend_from_slice(&frame.payload);
71 }
72 }
73
74 table.extend_from_slice(&payload);
75 Ok(table.freeze())
76 }
77
78 pub fn decode(mut buf: impl Buf) -> Result<Self> {
80 let version = read_varint(&mut buf)?;
81 if version != VERSION {
82 return Err(Error::Version(version));
83 }
84
85 let group_count = read_count(&mut buf, 2)?;
86 if group_count == 0 {
87 return Err(Error::Empty);
88 }
89
90 struct Entry {
91 sequence: u64,
92 frames: Vec<(u64, u64, u64)>,
93 }
94
95 let mut entries = Vec::new();
96 let mut prev: Option<u64> = None;
97 for i in 0..group_count {
98 let delta = read_varint(&mut buf)?;
99 let sequence = if i == 0 {
100 check_id(delta)?
101 } else {
102 let previous = prev.unwrap();
103 let sequence = previous
104 .checked_add(1)
105 .ok_or(Error::Overflow)?
106 .checked_add(delta)
107 .ok_or(Error::Overflow)?;
108 check_id(sequence)?
109 };
110 if let Some(previous) = prev
111 && sequence <= previous
112 {
113 return Err(Error::Sequence);
114 }
115 prev = Some(sequence);
116
117 let frame_count = read_count(&mut buf, 3)?;
118 let mut frames = Vec::new();
119 for _ in 0..frame_count {
120 let timestamp = check_id(read_varint(&mut buf)?)?;
121 let offset = read_varint(&mut buf)?;
122 let length = read_varint(&mut buf)?;
123 frames.push((timestamp, offset, length));
124 }
125 entries.push(Entry { sequence, frames });
126 }
127
128 let payload = buf.copy_to_bytes(buf.remaining());
129 let payload_len = u64::try_from(payload.len()).map_err(|_| Error::Overflow)?;
130
131 let mut expected = 0u64;
132 let mut groups = Vec::with_capacity(entries.len());
133 for entry in entries {
134 let mut frames = Vec::with_capacity(entry.frames.len());
135 for (timestamp, offset, length) in entry.frames {
136 if offset != expected {
137 return Err(Error::Table);
138 }
139 let end = offset.checked_add(length).ok_or(Error::Overflow)?;
140 if end > payload_len {
141 return Err(Error::Table);
142 }
143 let start = usize::try_from(offset).map_err(|_| Error::Overflow)?;
144 let stop = usize::try_from(end).map_err(|_| Error::Overflow)?;
145 frames.push(Frame {
146 timestamp,
147 payload: payload.slice(start..stop),
148 });
149 expected = end;
150 }
151 groups.push(Group {
152 sequence: entry.sequence,
153 frames,
154 });
155 }
156 if expected != payload_len {
157 return Err(Error::Table);
158 }
159
160 Ok(Self { groups })
161 }
162
163 pub fn decode_groups(buf: impl Buf, range: RangeInclusive<u64>) -> Result<Self> {
165 check_range(&range)?;
166 let object = Self::decode(buf)?;
167 object.check_bounds(range)?;
168 Ok(object)
169 }
170
171 pub fn check_bounds(&self, range: RangeInclusive<u64>) -> Result<()> {
173 check_range(&range)?;
174 let got = self.bounds()?;
175 if got != range {
176 return Err(Error::Bounds {
177 smallest: *got.start(),
178 largest: *got.end(),
179 });
180 }
181 Ok(())
182 }
183}
184
185fn validate(groups: &[Group]) -> Result<()> {
186 if groups.is_empty() {
187 return Err(Error::Empty);
188 }
189 let mut prev = None;
190 for group in groups {
191 check_id(group.sequence)?;
192 if let Some(previous) = prev
193 && group.sequence <= previous
194 {
195 return Err(Error::Sequence);
196 }
197 prev = Some(group.sequence);
198 for frame in &group.frames {
199 check_id(frame.timestamp)?;
200 }
201 }
202 Ok(())
203}
204
205fn write_varint(buf: &mut impl BufMut, value: u64) -> Result<()> {
206 let value = VarInt::try_from(value).map_err(|_| Error::Overflow)?;
207 value.encode_quic(buf).map_err(|_| Error::Overflow)
208}
209
210fn read_varint(buf: &mut impl Buf) -> Result<u64> {
211 Ok(VarInt::decode_quic(buf).map_err(|_| Error::Table)?.into_inner())
212}
213
214fn read_count(buf: &mut impl Buf, min_entry: usize) -> Result<usize> {
215 let n = read_varint(buf)?;
216 let n = usize::try_from(n).map_err(|_| Error::Overflow)?;
217 if min_entry == 0 || n > buf.remaining() / min_entry {
218 return Err(Error::Table);
219 }
220 Ok(n)
221}
222
223#[cfg(test)]
224mod tests {
225 use super::*;
226 use crate::ID_MAX;
227
228 fn frame(timestamp: u64, payload: &'static [u8]) -> Frame {
229 Frame {
230 timestamp,
231 payload: Bytes::from_static(payload),
232 }
233 }
234
235 fn object(groups: Vec<Group>) -> Object {
236 Object { groups }
237 }
238
239 #[test]
240 fn roundtrip_consecutive_and_sparse() {
241 let original = object(vec![
242 Group {
243 sequence: 0,
244 frames: vec![frame(0, b"a"), frame(1, b"bb")],
245 },
246 Group {
247 sequence: 1,
248 frames: vec![frame(2, b"ccc")],
249 },
250 Group {
251 sequence: 4,
252 frames: vec![frame(ID_MAX, b"d")],
253 },
254 ]);
255 let bytes = original.encode().unwrap();
256 assert_eq!(Object::decode(&bytes[..]).unwrap(), original);
257 assert_eq!(original.bounds().unwrap(), 0..=4);
258 original.check_bounds(0..=4).unwrap();
259 }
260
261 #[test]
262 fn first_id_is_absolute_then_minus_one_deltas() {
263 let bytes = object(vec![
264 Group {
265 sequence: 5,
266 frames: vec![frame(0, b"x")],
267 },
268 Group {
269 sequence: 6,
270 frames: vec![frame(1, b"y")],
271 },
272 Group {
273 sequence: 10,
274 frames: vec![frame(2, b"z")],
275 },
276 ])
277 .encode()
278 .unwrap();
279
280 let mut buf = &bytes[..];
281 assert_eq!(read_varint(&mut buf).unwrap(), 1); assert_eq!(read_varint(&mut buf).unwrap(), 3); assert_eq!(read_varint(&mut buf).unwrap(), 5); assert_eq!(read_varint(&mut buf).unwrap(), 1); read_varint(&mut buf).unwrap(); read_varint(&mut buf).unwrap(); read_varint(&mut buf).unwrap(); assert_eq!(read_varint(&mut buf).unwrap(), 0); assert_eq!(read_varint(&mut buf).unwrap(), 1);
290 read_varint(&mut buf).unwrap();
291 read_varint(&mut buf).unwrap();
292 read_varint(&mut buf).unwrap();
293 assert_eq!(read_varint(&mut buf).unwrap(), 3); }
295
296 #[test]
297 fn empty_object_is_rejected() {
298 assert!(matches!(object(vec![]).encode(), Err(Error::Empty)));
299 let mut bytes = BytesMut::new();
300 write_varint(&mut bytes, 1).unwrap();
301 write_varint(&mut bytes, 0).unwrap();
302 assert!(matches!(Object::decode(bytes.freeze()), Err(Error::Empty)));
303 }
304
305 #[test]
306 fn sequences_must_ascend() {
307 assert!(matches!(
308 object(vec![
309 Group {
310 sequence: 2,
311 frames: vec![frame(0, b"a")],
312 },
313 Group {
314 sequence: 2,
315 frames: vec![frame(1, b"b")],
316 },
317 ])
318 .encode(),
319 Err(Error::Sequence)
320 ));
321 assert!(matches!(
322 object(vec![
323 Group {
324 sequence: 2,
325 frames: vec![frame(0, b"a")],
326 },
327 Group {
328 sequence: 1,
329 frames: vec![frame(1, b"b")],
330 },
331 ])
332 .encode(),
333 Err(Error::Sequence)
334 ));
335 }
336
337 #[test]
338 fn reconstructed_ids_must_stay_in_range() {
339 let mut bytes = BytesMut::new();
340 write_varint(&mut bytes, 1).unwrap();
341 write_varint(&mut bytes, 2).unwrap();
342 write_varint(&mut bytes, ID_MAX).unwrap();
343 write_varint(&mut bytes, 0).unwrap(); write_varint(&mut bytes, 0).unwrap(); write_varint(&mut bytes, 0).unwrap();
346 assert!(matches!(Object::decode(bytes.freeze()), Err(Error::Id(_))));
347 }
348
349 #[test]
350 fn timestamp_bounds() {
351 assert!(
352 object(vec![Group {
353 sequence: 0,
354 frames: vec![frame(ID_MAX, b"a")],
355 }])
356 .encode()
357 .is_ok()
358 );
359 assert!(matches!(
360 object(vec![Group {
361 sequence: 0,
362 frames: vec![frame(ID_MAX + 1, b"a")],
363 }])
364 .encode(),
365 Err(Error::Id(_))
366 ));
367 }
368
369 #[test]
370 fn unknown_version_is_refused() {
371 let mut bytes = BytesMut::new();
372 write_varint(&mut bytes, 2).unwrap();
373 write_varint(&mut bytes, 1).unwrap();
374 assert!(matches!(Object::decode(bytes.freeze()), Err(Error::Version(2))));
375 }
376
377 #[test]
378 fn payload_must_be_contiguous() {
379 let mut good = object(vec![Group {
380 sequence: 0,
381 frames: vec![frame(0, b"ab")],
382 }])
383 .encode()
384 .unwrap()
385 .to_vec();
386 good.push(b'x');
387 assert!(matches!(Object::decode(Bytes::from(good)), Err(Error::Table)));
388 }
389
390 #[test]
391 fn advertised_counts_must_fit_minimum_entry_size() {
392 let mut bytes = BytesMut::new();
395 write_varint(&mut bytes, 1).unwrap();
396 write_varint(&mut bytes, 4).unwrap();
397 bytes.extend_from_slice(&[0, 0, 0, 0]);
398 assert!(matches!(Object::decode(bytes.freeze()), Err(Error::Table)));
399 }
400
401 #[test]
402 fn offset_overflow_and_truncated_table() {
403 let mut bytes = BytesMut::new();
404 write_varint(&mut bytes, 1).unwrap();
405 write_varint(&mut bytes, 1).unwrap();
406 write_varint(&mut bytes, 0).unwrap();
407 write_varint(&mut bytes, 1).unwrap();
408 write_varint(&mut bytes, 0).unwrap();
409 write_varint(&mut bytes, 0).unwrap();
410 write_varint(&mut bytes, 4).unwrap(); assert!(matches!(Object::decode(bytes.freeze()), Err(Error::Table)));
412
413 assert!(matches!(Object::decode(&b"\x01"[..]), Err(Error::Table)));
414 }
415
416 #[test]
417 fn filename_bounds_must_match_the_table() {
418 let object = object(vec![
419 Group {
420 sequence: 5,
421 frames: vec![frame(0, b"a")],
422 },
423 Group {
424 sequence: 7,
425 frames: vec![frame(1, b"b")],
426 },
427 ]);
428 let bytes = object.encode().unwrap();
429 assert!(Object::decode_groups(&bytes[..], 5..=7).is_ok());
430 assert!(matches!(
431 Object::decode_groups(&bytes[..], 6..=7),
432 Err(Error::Bounds {
433 smallest: 5,
434 largest: 7
435 })
436 ));
437 let (smallest, largest) = (7, 5);
438 assert!(matches!(
439 Object::decode_groups(&bytes[..], smallest..=largest),
440 Err(Error::Bounds {
441 smallest: 7,
442 largest: 5
443 })
444 ));
445 }
446
447 #[test]
448 fn endpoints_roundtrip() {
449 let original = object(vec![Group {
450 sequence: ID_MAX,
451 frames: vec![frame(0, b"")],
452 }]);
453 assert_eq!(Object::decode(original.encode().unwrap()).unwrap(), original);
454 }
455}