1use std::collections::VecDeque;
4use std::marker::PhantomData;
5
6use bytes::Bytes;
7use serde::Serialize;
8use serde_json::Value;
9
10use super::op::{Header, Op};
11use crate::{Error, Result};
12
13pub(super) const MAX_GROUP_FRAMES: usize = 1024;
18
19pub(super) const MAX_INDEX: u64 = (1 << 53) - 1;
21
22#[derive(Debug, Clone)]
27#[non_exhaustive]
28pub struct ProducerConfig {
29 pub op_ratio: u32,
42
43 pub compression: bool,
49
50 pub checkpoint_records: Option<usize>,
56}
57
58impl Default for ProducerConfig {
59 fn default() -> Self {
60 Self {
61 op_ratio: 8,
62 compression: false,
63 checkpoint_records: None,
64 }
65 }
66}
67
68impl ProducerConfig {
69 pub fn with_op_ratio(mut self, op_ratio: u32) -> Self {
71 self.op_ratio = op_ratio;
72 self
73 }
74
75 pub fn with_compression(mut self, compression: bool) -> Self {
77 self.compression = compression;
78 self
79 }
80
81 pub fn with_checkpoint_records(mut self, checkpoint_records: usize) -> Self {
83 assert!(checkpoint_records > 0, "checkpoint_records must be positive");
84 self.checkpoint_records = Some(checkpoint_records);
85 self
86 }
87}
88
89#[derive(Clone, Debug)]
91#[non_exhaustive]
92pub struct Encoded {
93 pub payload: Bytes,
95
96 pub keyframe: bool,
101}
102
103#[must_use = "write the frame, then commit it"]
108pub struct Pending<'a, T> {
109 encoder: &'a mut Encoder<T>,
110 encoded: Encoded,
111 edit: Option<Edit>,
112}
113
114enum Edit {
116 Push(Value),
117 Pop(u64),
118}
119
120impl<T> std::ops::Deref for Pending<'_, T> {
121 type Target = Encoded;
122
123 fn deref(&self) -> &Encoded {
124 &self.encoded
125 }
126}
127
128impl<T> Pending<'_, T> {
129 pub fn commit(mut self) {
133 let edit = self.edit.take().expect("pending edit");
134 self.encoder.commit(edit);
135 }
136}
137
138impl<T> Drop for Pending<'_, T> {
139 fn drop(&mut self) {
140 if self.edit.is_some() {
141 self.encoder.resync();
142 }
143 }
144}
145
146pub struct Encoder<T> {
156 config: ProducerConfig,
157
158 window: VecDeque<Value>,
160
161 offset: u64,
163
164 start: u64,
166
167 flate: Option<moq_flate::Encoder>,
169
170 op_bytes: u64,
172
173 header_len: u64,
175
176 group_frames: usize,
178
179 resync: bool,
182
183 _marker: PhantomData<fn(T)>,
184}
185
186impl<T> Encoder<T> {
187 pub fn new(config: ProducerConfig) -> Self {
189 assert!(
190 config.checkpoint_records != Some(0),
191 "checkpoint_records must be positive"
192 );
193 Self {
194 config,
195 window: VecDeque::new(),
196 offset: 0,
197 start: 0,
198 flate: None,
199 op_bytes: 0,
200 header_len: 0,
201 group_frames: 0,
202 resync: true,
203 _marker: PhantomData,
204 }
205 }
206
207 pub fn window(&self) -> Vec<Value> {
211 self.window.iter().cloned().collect()
212 }
213
214 pub fn range(&self) -> std::ops::Range<u64> {
216 self.offset..self.start + self.window.len() as u64
217 }
218
219 fn resync(&mut self) {
221 self.flate = None;
222 self.op_bytes = 0;
223 self.header_len = 0;
224 self.group_frames = 0;
225 self.resync = true;
226 }
227
228 fn commit(&mut self, edit: Edit) {
230 match edit {
231 Edit::Push(record) => {
232 self.window.push_back(record);
233 if let Some(limit) = self.config.checkpoint_records {
234 while self.window.len() > limit {
235 self.window.pop_front();
236 self.start += 1;
237 }
238 }
239 }
240 Edit::Pop(count) => {
241 let offset = self.offset + count;
242 let stored = offset.saturating_sub(self.start).min(self.window.len() as u64);
243 self.window.drain(..stored as usize);
244 self.start += stored;
245 self.offset = offset;
246 }
247 }
248 }
249
250 fn op_allowed(&self) -> bool {
252 let ratio = u64::from(self.config.op_ratio);
253 ratio != 0
254 && self.group_frames > 0
255 && self.group_frames < MAX_GROUP_FRAMES
256 && self.op_bytes <= ratio * self.header_len
257 }
258
259 fn validate_plaintext(len: usize, kind: &str) -> Result<()> {
261 if u64::try_from(len).unwrap_or(u64::MAX) > moq_flate::DEFAULT_MAX_FRAME_SIZE {
262 return Err(Error::Json(format!(
263 "window {kind} exceeds the decoder's decompressed size limit"
264 )));
265 }
266 Ok(())
267 }
268
269 fn frame(&mut self, bytes: Vec<u8>) -> Result<Encoded> {
271 Self::validate_plaintext(bytes.len(), "frame")?;
272 let payload = match self.flate.as_mut() {
273 Some(flate) => flate.frame(&bytes),
274 None => Bytes::from(bytes),
275 };
276
277 self.op_bytes += payload.len() as u64;
278 self.group_frames += 1;
279
280 Ok(Encoded {
281 payload,
282 keyframe: false,
283 })
284 }
285
286 fn emit_op(&mut self, bytes: Vec<u8>) -> Result<Option<Encoded>> {
288 let encoded = self.frame(bytes)?;
289 let group_bytes = self.header_len.saturating_add(self.op_bytes);
290 if group_bytes > moq_net::group::MAX_CACHE_BYTES {
291 self.resync();
292 Ok(None)
293 } else {
294 Ok(Some(encoded))
295 }
296 }
297
298 fn header<'a>(
300 config: &ProducerConfig,
301 offset: u64,
302 start: u64,
303 len: usize,
304 records: impl Iterator<Item = &'a Value>,
305 ) -> Result<Vec<u8>> {
306 let skip = config
307 .checkpoint_records
308 .map(|limit| len.saturating_sub(limit))
309 .unwrap_or_default();
310 let start = start
311 .checked_add(skip as u64)
312 .ok_or_else(|| Error::Json("window checkpoint exceeds u64".into()))?;
313 let header = Header {
314 offset,
315 start: (start != offset).then_some(start),
316 records: records.skip(skip).collect(),
317 };
318 Ok(serde_json::to_vec(&header)?)
319 }
320
321 fn emit_header(&mut self, bytes: Vec<u8>) -> Result<Encoded> {
323 Self::validate_plaintext(bytes.len(), "header")?;
324
325 let (payload, flate) = match self.config.compression {
328 true => {
329 let mut flate = moq_flate::Encoder::new();
330 let payload = flate.frame(&bytes);
331 (payload, Some(flate))
332 }
333 false => (Bytes::from(bytes), None),
334 };
335 if payload.len() as u64 > moq_net::group::MAX_CACHE_BYTES {
336 return Err(Error::Json("window header exceeds the group cache limit".into()));
337 }
338
339 self.header_len = payload.len() as u64;
340 self.op_bytes = 0;
341 self.group_frames = 1;
342 self.flate = flate;
343 self.resync = false;
344
345 Ok(Encoded {
346 payload,
347 keyframe: true,
348 })
349 }
350
351 pub fn pop(&mut self, count: u64) -> Result<Option<Pending<'_, T>>> {
356 let count = count.min(self.range().end - self.offset);
357 if count == 0 {
358 return Ok(None);
359 }
360
361 let offset = self.offset + count;
362 let stored = offset.saturating_sub(self.start).min(self.window.len() as u64) as usize;
363 let start = self.start + stored as u64;
364 let encoded = match self.resync || !self.op_allowed() {
365 true => {
366 let bytes = Self::header(
367 &self.config,
368 offset,
369 start,
370 self.window.len() - stored,
371 self.window.iter().skip(stored),
372 )?;
373 self.emit_header(bytes)?
374 }
375 false => {
376 let bytes = serde_json::to_vec(&Op::<&Value>::Pop(count))?;
377 match self.emit_op(bytes)? {
378 Some(encoded) => encoded,
379 None => {
380 let bytes = Self::header(
381 &self.config,
382 offset,
383 start,
384 self.window.len() - stored,
385 self.window.iter().skip(stored),
386 )?;
387 self.emit_header(bytes)?
388 }
389 }
390 }
391 };
392
393 Ok(Some(self.pending(encoded, Edit::Pop(count))))
394 }
395
396 fn pending(&mut self, encoded: Encoded, edit: Edit) -> Pending<'_, T> {
398 Pending {
399 encoder: self,
400 encoded,
401 edit: Some(edit),
402 }
403 }
404}
405
406impl<T: Serialize> Encoder<T> {
407 pub fn push(&mut self, value: &T) -> Result<Pending<'_, T>> {
412 let bytes = serde_json::to_vec(value)?;
416 let record: Value = serde_json::from_slice(&bytes)?;
417 if self.range().end >= MAX_INDEX {
418 return Err(crate::Error::Json("window index exceeds the safe integer range".into()));
419 }
420
421 let encoded = match self.resync || !self.op_allowed() {
422 true => {
423 let bytes = Self::header(
424 &self.config,
425 self.offset,
426 self.start,
427 self.window.len() + 1,
428 self.window.iter().chain(std::iter::once(&record)),
429 )?;
430 self.emit_header(bytes)?
431 }
432 false => {
433 let bytes = serde_json::to_vec(&Op::Push(&record))?;
434 match self.emit_op(bytes)? {
435 Some(encoded) => encoded,
436 None => {
437 let bytes = Self::header(
438 &self.config,
439 self.offset,
440 self.start,
441 self.window.len() + 1,
442 self.window.iter().chain(std::iter::once(&record)),
443 )?;
444 self.emit_header(bytes)?
445 }
446 }
447 }
448 };
449
450 Ok(self.pending(encoded, Edit::Push(record)))
451 }
452}
453
454#[cfg(test)]
455mod test {
456 use super::*;
457
458 #[test]
459 fn an_op_that_would_evict_the_header_rolls_first() {
460 let mut encoder = Encoder::<String>::new(ProducerConfig::default().with_op_ratio(u32::MAX));
461 let first = "a".repeat(16 * 1024 * 1024);
462 let next = "b".repeat(15 * 1024 * 1024);
463
464 let frame = encoder.push(&first).unwrap();
465 assert!(frame.keyframe);
466 frame.commit();
467
468 let frame = encoder.push(&next).unwrap();
469 assert!(!frame.keyframe);
470 frame.commit();
471
472 let frame = encoder.pop(1).unwrap().unwrap();
473 assert!(!frame.keyframe);
474 frame.commit();
475
476 let frame = encoder.push(&next).unwrap();
477 assert!(frame.keyframe);
478 assert!(frame.payload.len() < moq_net::group::MAX_CACHE_BYTES as usize);
479 frame.commit();
480 }
481
482 #[test]
483 fn an_uncommitted_edit_leaves_the_window_unchanged() {
484 let mut encoder = Encoder::<u64>::new(ProducerConfig::default());
485
486 drop(encoder.push(&1).unwrap());
487 assert!(encoder.window().is_empty());
488
489 let frame = encoder.push(&2).unwrap();
490 assert!(frame.keyframe);
491 frame.commit();
492 assert_eq!(encoder.window(), vec![Value::from(2)]);
493
494 drop(encoder.pop(1).unwrap().unwrap());
495 assert_eq!(encoder.window(), vec![Value::from(2)]);
496 }
497
498 #[test]
499 fn a_header_larger_than_the_group_cache_is_rejected() {
500 let mut encoder = Encoder::<String>::new(ProducerConfig::default());
501 let record = "x".repeat(moq_net::group::MAX_CACHE_BYTES as usize);
502
503 let err = encoder.push(&record).err().expect("oversized header should fail");
504 assert!(err.to_string().contains("group cache limit"));
505 assert!(encoder.window().is_empty());
506
507 let frame = encoder.push(&"ok".to_string()).unwrap();
508 assert!(frame.keyframe);
509 }
510
511 #[test]
512 fn plaintext_is_bounded_by_the_decoder_limit() {
513 let len = usize::try_from(moq_flate::DEFAULT_MAX_FRAME_SIZE + 1).unwrap();
514 assert!(Encoder::<()>::validate_plaintext(len, "frame").is_err());
515 }
516}