use std::collections::VecDeque;
use serde::de::DeserializeOwned;
use super::encoder::MAX_INDEX;
use super::op::{Header, Op};
use crate::{Error, Result};
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct ConsumerConfig {
pub compression: bool,
}
impl ConsumerConfig {
pub fn with_compression(mut self, compression: bool) -> Self {
self.compression = compression;
self
}
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub enum Event<T> {
Push {
index: u64,
value: T,
},
Pop(std::ops::Range<u64>),
Skip(std::ops::Range<u64>),
}
enum Queued<T> {
Event(Event<T>),
Push { index: u64, records: std::vec::IntoIter<T> },
}
pub struct Group<'a, T> {
decoder: &'a mut Decoder<T>,
codec: Codec,
}
pub(super) struct Codec {
flate: Option<moq_flate::Decoder>,
positioned: bool,
}
impl Codec {
pub(super) fn new() -> Self {
Self {
flate: None,
positioned: false,
}
}
}
pub struct Decoder<T> {
config: ConsumerConfig,
front: u64,
len: u64,
delivered: Option<u64>,
events: VecDeque<Queued<T>>,
}
impl<T> Decoder<T> {
pub fn new(config: ConsumerConfig) -> Self {
Self {
config,
front: 0,
len: 0,
delivered: None,
events: VecDeque::new(),
}
}
pub fn group(&mut self) -> Group<'_, T> {
Group {
decoder: self,
codec: Codec::new(),
}
}
pub fn next_event(&mut self) -> Option<Event<T>> {
match self.events.pop_front()? {
Queued::Event(event) => Some(event),
Queued::Push { index, mut records } => {
let value = records.next().expect("queued push batch is not empty");
if !records.as_slice().is_empty() {
self.events.push_front(Queued::Push {
index: index + 1,
records,
});
}
Some(Event::Push { index, value })
}
}
}
pub fn range(&self) -> std::ops::Range<u64> {
self.front..self.front + self.len
}
}
impl<T: DeserializeOwned> Decoder<T> {
pub(super) fn decode(&mut self, group: &mut Codec, payload: &[u8]) -> Result<()> {
let inflated = match self.config.compression {
true => Some(group.flate.get_or_insert_with(moq_flate::Decoder::new).frame(payload)?),
false => None,
};
let bytes = inflated.as_deref().unwrap_or(payload);
if !group.positioned {
let header: Header<T> = serde_path_to_error::deserialize(&mut serde_json::Deserializer::from_slice(bytes))
.map_err(|err| Error::Json(err.to_string()))?;
self.apply_header(header.offset, header.records)?;
group.positioned = true;
return Ok(());
}
match serde_path_to_error::deserialize(&mut serde_json::Deserializer::from_slice(bytes))
.map_err(|err| Error::Json(err.to_string()))?
{
Op::Push(record) => self.apply_push(record),
Op::Pop(count) => self.apply_pop(count),
}
}
fn apply_header(&mut self, offset: u64, records: Vec<T>) -> Result<()> {
if offset > MAX_INDEX {
return Err(Error::Json("window offset exceeds the safe integer range".into()));
}
let len = u64::try_from(records.len()).map_err(|_| Error::Json("window length exceeds u64".into()))?;
let end = offset
.checked_add(len)
.filter(|end| *end <= MAX_INDEX)
.ok_or_else(|| Error::Json("window range exceeds the safe integer range".into()))?;
let delivered = match self.delivered {
None => offset,
Some(delivered) => {
if offset < self.front || end < delivered {
return Err(Error::Json("window header moved backwards".into()));
}
let popped = self.front..delivered.min(offset);
if !popped.is_empty() {
self.events.push_back(Queued::Event(Event::Pop(popped)));
}
let skipped = delivered..offset;
if !skipped.is_empty() {
self.events.push_back(Queued::Event(Event::Skip(skipped)));
}
delivered
}
};
let skip = usize::try_from(delivered.saturating_sub(offset))
.map_err(|_| Error::Json("window length exceeds usize".into()))?;
let mut records = records.into_iter();
if skip > 0 {
records.nth(skip - 1);
}
if !records.as_slice().is_empty() {
self.events.push_back(Queued::Push {
index: offset + skip as u64,
records,
});
}
self.front = offset;
self.len = end - offset;
self.delivered = Some(delivered.max(end));
Ok(())
}
fn apply_push(&mut self, record: T) -> Result<()> {
let delivered = self.delivered.expect("group header positioned the decoder");
let index = self
.front
.checked_add(self.len)
.ok_or_else(|| Error::Json("window range exceeds u64".into()))?;
let end = index
.checked_add(1)
.filter(|end| *end <= MAX_INDEX)
.ok_or_else(|| Error::Json("window range exceeds the safe integer range".into()))?;
self.len = end - self.front;
if index >= delivered {
self.events
.push_back(Queued::Event(Event::Push { index, value: record }));
self.delivered = Some(end);
}
Ok(())
}
fn apply_pop(&mut self, count: u64) -> Result<()> {
let delivered = self.delivered.expect("group header positioned the decoder");
if count > self.len {
return Err(Error::Json(format!(
"pop of {count} exceeds the {} record(s) in the window",
self.len
)));
}
let end = self
.front
.checked_add(count)
.ok_or_else(|| Error::Json("window range exceeds u64".into()))?;
let popped = self.front..delivered.min(end);
if !popped.is_empty() {
self.events.push_back(Queued::Event(Event::Pop(popped)));
}
let skipped = delivered.max(self.front)..end;
if !skipped.is_empty() {
self.events.push_back(Queued::Event(Event::Skip(skipped)));
}
self.front = end;
self.len -= count;
self.delivered = Some(delivered.max(self.front));
Ok(())
}
}
impl<T> Group<'_, T> {
pub fn next_event(&mut self) -> Option<Event<T>> {
self.decoder.next_event()
}
pub fn range(&self) -> std::ops::Range<u64> {
self.decoder.range()
}
}
impl<T: DeserializeOwned> Group<'_, T> {
pub fn decode(&mut self, payload: &[u8]) -> Result<()> {
self.decoder.decode(&mut self.codec, payload)
}
}