Skip to main content

moq_archive/
segment.rs

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/// One complete group in a segment object, in sequence order.
10#[derive(Debug, Clone, PartialEq, Eq)]
11pub struct Group {
12	/// Group sequence number.
13	pub sequence: u64,
14	/// Frames in their original order within the group.
15	pub frames: Vec<Frame>,
16}
17
18/// One frame's timestamp and original payload.
19#[derive(Debug, Clone, PartialEq, Eq)]
20pub struct Frame {
21	/// Absolute timestamp in the track's timescale.
22	pub timestamp: u64,
23	/// Original frame payload, including any group-scoped compression.
24	pub payload: Bytes,
25}
26
27/// A versioned group/frame table followed by concatenated payload bytes.
28#[derive(Debug, Clone, PartialEq, Eq)]
29pub struct Object {
30	/// Groups in strictly ascending sequence order.
31	pub groups: Vec<Group>,
32}
33
34impl Object {
35	/// Inclusive group sequences from first to last.
36	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	/// Encode the binary envelope. The first sequence is absolute; later ones are `current - previous - 1`.
42	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	/// Decode a complete table before slicing any payload, and refuse an unknown version.
79	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	/// Decode and require the table's sequences to match `range`.
164	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	/// Require this object's sequences to match a range-named key.
172	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); // version
282		assert_eq!(read_varint(&mut buf).unwrap(), 3); // group count
283		assert_eq!(read_varint(&mut buf).unwrap(), 5); // absolute
284		assert_eq!(read_varint(&mut buf).unwrap(), 1); // one frame
285		read_varint(&mut buf).unwrap(); // timestamp
286		read_varint(&mut buf).unwrap(); // offset
287		read_varint(&mut buf).unwrap(); // length
288		assert_eq!(read_varint(&mut buf).unwrap(), 0); // consecutive
289		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); // 10 - 6 - 1
294	}
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(); // no frames
344		write_varint(&mut bytes, 0).unwrap(); // consecutive => ID_MAX + 1
345		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		// group_count equals remaining bytes, so a remaining-bytes check would
393		// accept it, but each group needs at least two varints.
394		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(); // length 4, no payload
411		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}