use std::collections::VecDeque;
use std::marker::PhantomData;
use bytes::Bytes;
use serde::Serialize;
use serde_json::Value;
use super::op::{Header, Op};
use crate::{Error, Result};
pub(super) const MAX_GROUP_FRAMES: usize = 256;
pub(super) const MAX_INDEX: u64 = (1 << 53) - 1;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct ProducerConfig {
pub op_ratio: u32,
pub compression: bool,
}
impl Default for ProducerConfig {
fn default() -> Self {
Self {
op_ratio: 8,
compression: false,
}
}
}
impl ProducerConfig {
pub fn with_op_ratio(mut self, op_ratio: u32) -> Self {
self.op_ratio = op_ratio;
self
}
pub fn with_compression(mut self, compression: bool) -> Self {
self.compression = compression;
self
}
}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct Encoded {
pub payload: Bytes,
pub keyframe: bool,
}
#[must_use = "write the frame, then commit it"]
pub struct Pending<'a, T> {
encoder: &'a mut Encoder<T>,
encoded: Encoded,
edit: Option<Edit>,
}
enum Edit {
Push(Value),
Pop(u64),
}
impl<T> std::ops::Deref for Pending<'_, T> {
type Target = Encoded;
fn deref(&self) -> &Encoded {
&self.encoded
}
}
impl<T> Pending<'_, T> {
pub fn commit(mut self) {
let edit = self.edit.take().expect("pending edit");
self.encoder.commit(edit);
}
}
impl<T> Drop for Pending<'_, T> {
fn drop(&mut self) {
if self.edit.is_some() {
self.encoder.resync();
}
}
}
pub struct Encoder<T> {
config: ProducerConfig,
window: VecDeque<Value>,
offset: u64,
flate: Option<moq_flate::Encoder>,
op_bytes: u64,
header_len: u64,
group_frames: usize,
resync: bool,
_marker: PhantomData<fn(T)>,
}
impl<T> Encoder<T> {
pub fn new(config: ProducerConfig) -> Self {
Self {
config,
window: VecDeque::new(),
offset: 0,
flate: None,
op_bytes: 0,
header_len: 0,
group_frames: 0,
resync: true,
_marker: PhantomData,
}
}
pub fn window(&self) -> Vec<Value> {
self.window.iter().cloned().collect()
}
pub fn range(&self) -> std::ops::Range<u64> {
self.offset..self.offset + self.window.len() as u64
}
fn resync(&mut self) {
self.flate = None;
self.op_bytes = 0;
self.header_len = 0;
self.group_frames = 0;
self.resync = true;
}
fn commit(&mut self, edit: Edit) {
match edit {
Edit::Push(record) => self.window.push_back(record),
Edit::Pop(count) => {
self.window.drain(..count as usize);
self.offset += count;
}
}
}
fn op_allowed(&self) -> bool {
let ratio = u64::from(self.config.op_ratio);
ratio != 0
&& self.group_frames > 0
&& self.group_frames < MAX_GROUP_FRAMES
&& self.op_bytes <= ratio * self.header_len
}
fn validate_plaintext(len: usize, kind: &str) -> Result<()> {
if u64::try_from(len).unwrap_or(u64::MAX) > moq_flate::DEFAULT_MAX_FRAME_SIZE {
return Err(Error::Json(format!(
"window {kind} exceeds the decoder's decompressed size limit"
)));
}
Ok(())
}
fn frame(&mut self, bytes: Vec<u8>) -> Result<Encoded> {
Self::validate_plaintext(bytes.len(), "frame")?;
let payload = match self.flate.as_mut() {
Some(flate) => flate.frame(&bytes),
None => Bytes::from(bytes),
};
self.op_bytes += payload.len() as u64;
self.group_frames += 1;
Ok(Encoded {
payload,
keyframe: false,
})
}
fn emit_op(&mut self, bytes: Vec<u8>) -> Result<Option<Encoded>> {
let encoded = self.frame(bytes)?;
let group_bytes = self.header_len.saturating_add(self.op_bytes);
if group_bytes > moq_net::group::MAX_CACHE_BYTES {
self.resync();
Ok(None)
} else {
Ok(Some(encoded))
}
}
fn header(offset: u64, records: Vec<&Value>) -> Result<Vec<u8>> {
Ok(serde_json::to_vec(&Header { offset, records })?)
}
fn emit_header(&mut self, bytes: Vec<u8>) -> Result<Encoded> {
Self::validate_plaintext(bytes.len(), "header")?;
let (payload, flate) = match self.config.compression {
true => {
let mut flate = moq_flate::Encoder::new();
let payload = flate.frame(&bytes);
(payload, Some(flate))
}
false => (Bytes::from(bytes), None),
};
if payload.len() as u64 > moq_net::group::MAX_CACHE_BYTES {
return Err(Error::Json("window header exceeds the group cache limit".into()));
}
self.header_len = payload.len() as u64;
self.op_bytes = 0;
self.group_frames = 1;
self.flate = flate;
self.resync = false;
Ok(Encoded {
payload,
keyframe: true,
})
}
pub fn pop(&mut self, count: u64) -> Result<Option<Pending<'_, T>>> {
let count = count.min(self.window.len() as u64);
if count == 0 {
return Ok(None);
}
let offset = self.offset + count;
let encoded = match self.resync || !self.op_allowed() {
true => {
let bytes = Self::header(offset, self.window.iter().skip(count as usize).collect())?;
self.emit_header(bytes)?
}
false => {
let bytes = serde_json::to_vec(&Op::<&Value>::Pop(count))?;
match self.emit_op(bytes)? {
Some(encoded) => encoded,
None => {
let bytes = Self::header(offset, self.window.iter().skip(count as usize).collect())?;
self.emit_header(bytes)?
}
}
}
};
Ok(Some(self.pending(encoded, Edit::Pop(count))))
}
fn pending(&mut self, encoded: Encoded, edit: Edit) -> Pending<'_, T> {
Pending {
encoder: self,
encoded,
edit: Some(edit),
}
}
}
impl<T: Serialize> Encoder<T> {
pub fn push(&mut self, value: &T) -> Result<Pending<'_, T>> {
let bytes = serde_json::to_vec(value)?;
let record: Value = serde_json::from_slice(&bytes)?;
if self.range().end >= MAX_INDEX {
return Err(crate::Error::Json("window index exceeds the safe integer range".into()));
}
let encoded = match self.resync || !self.op_allowed() {
true => {
let bytes = Self::header(
self.offset,
self.window.iter().chain(std::iter::once(&record)).collect(),
)?;
self.emit_header(bytes)?
}
false => {
let bytes = serde_json::to_vec(&Op::Push(&record))?;
match self.emit_op(bytes)? {
Some(encoded) => encoded,
None => {
let bytes = Self::header(
self.offset,
self.window.iter().chain(std::iter::once(&record)).collect(),
)?;
self.emit_header(bytes)?
}
}
}
};
Ok(self.pending(encoded, Edit::Push(record)))
}
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn an_op_that_would_evict_the_header_rolls_first() {
let mut encoder = Encoder::<String>::new(ProducerConfig::default().with_op_ratio(u32::MAX));
let first = "a".repeat(16 * 1024 * 1024);
let next = "b".repeat(15 * 1024 * 1024);
let frame = encoder.push(&first).unwrap();
assert!(frame.keyframe);
frame.commit();
let frame = encoder.push(&next).unwrap();
assert!(!frame.keyframe);
frame.commit();
let frame = encoder.pop(1).unwrap().unwrap();
assert!(!frame.keyframe);
frame.commit();
let frame = encoder.push(&next).unwrap();
assert!(frame.keyframe);
assert!(frame.payload.len() < moq_net::group::MAX_CACHE_BYTES as usize);
frame.commit();
}
#[test]
fn an_uncommitted_edit_leaves_the_window_unchanged() {
let mut encoder = Encoder::<u64>::new(ProducerConfig::default());
drop(encoder.push(&1).unwrap());
assert!(encoder.window().is_empty());
let frame = encoder.push(&2).unwrap();
assert!(frame.keyframe);
frame.commit();
assert_eq!(encoder.window(), vec![Value::from(2)]);
drop(encoder.pop(1).unwrap().unwrap());
assert_eq!(encoder.window(), vec![Value::from(2)]);
}
#[test]
fn a_header_larger_than_the_group_cache_is_rejected() {
let mut encoder = Encoder::<String>::new(ProducerConfig::default());
let record = "x".repeat(moq_net::group::MAX_CACHE_BYTES as usize);
let err = encoder.push(&record).err().expect("oversized header should fail");
assert!(err.to_string().contains("group cache limit"));
assert!(encoder.window().is_empty());
let frame = encoder.push(&"ok".to_string()).unwrap();
assert!(frame.keyframe);
}
#[test]
fn plaintext_is_bounded_by_the_decoder_limit() {
let len = usize::try_from(moq_flate::DEFAULT_MAX_FRAME_SIZE + 1).unwrap();
assert!(Encoder::<()>::validate_plaintext(len, "frame").is_err());
}
}