use std::cell::RefCell;
use std::marker::PhantomData;
use serde::de::DeserializeOwned;
use serde_json::Value;
use super::consumer::Config;
use crate::{Error, Result};
pub struct Decoder<T> {
compression: bool,
flate: Option<moq_flate::Decoder>,
plain: Vec<u8>,
check: RefCell<crate::merge::CheckScratch>,
current: Option<Value>,
_marker: PhantomData<fn() -> T>,
}
impl<T> Decoder<T> {
pub fn new(config: Config) -> Self {
Self {
compression: config.compression.is_deflate(),
flate: None,
plain: Vec::new(),
check: RefCell::new(crate::merge::CheckScratch::default()),
current: None,
_marker: PhantomData,
}
}
pub fn snapshot(&mut self, payload: &[u8]) -> Result<()> {
self.flate = self.compression.then(moq_flate::Decoder::new);
self.current = Some(match self.flate.as_mut() {
Some(flate) => serde_json::from_slice(&flate.frame(payload)?)?,
None => serde_json::from_slice(payload)?,
});
Ok(())
}
pub fn delta(&mut self, payload: &[u8]) -> Result<()> {
if self.current.is_none() {
return Err(Error::MissingSnapshot);
}
let plain = match self.flate.as_mut() {
Some(flate) => {
flate.frame_into(payload, &mut self.plain)?;
self.plain.as_slice()
}
None => payload,
};
crate::merge::apply_bytes(
self.current.as_mut().expect("a snapshot precedes any delta"),
plain,
&self.check,
)?;
Ok(())
}
pub fn value(&self) -> Option<&Value> {
self.current.as_ref()
}
}
impl<T: DeserializeOwned> Decoder<T> {
pub fn decode(&self) -> Result<Option<T>> {
let Some(current) = self.current.as_ref() else {
return Ok(None);
};
let value = serde_path_to_error::deserialize(current).map_err(|err| {
let path = err.path().to_string();
match path.as_str() {
"." => Error::Json(err.into_inner().to_string()),
_ => Error::Json(format!("{}: {}", path, err.into_inner())),
}
})?;
Ok(Some(value))
}
}
#[cfg(test)]
mod test {
use super::super::consumer::Config as ConsumerConfig;
use super::super::{Config, Encoder};
use super::*;
use crate::Compression;
use serde_json::{Value, json};
fn consume(compression: Compression) -> ConsumerConfig {
ConsumerConfig { compression }
}
fn deflate() -> Config {
Config {
compression: Compression::Deflate,
..Default::default()
}
}
fn roundtrip(config: Config, values: &[Value]) -> Vec<Value> {
let compression = config.compression;
let mut encoder = Encoder::<Value>::new(config);
let mut decoder = Decoder::<Value>::new(consume(compression));
let mut out = Vec::new();
for value in values {
let Some(frame) = encoder.update(value).unwrap() else {
continue;
};
match frame.keyframe {
true => decoder.snapshot(&frame.payload).unwrap(),
false => decoder.delta(&frame.payload).unwrap(),
}
frame.commit();
out.push(decoder.decode().unwrap().unwrap());
}
out
}
#[test]
fn plaintext_roundtrip() {
let values = vec![
json!({ "a": 1, "b": 1 }),
json!({ "a": 1, "b": 2 }),
json!({ "a": 5, "b": 2 }),
];
assert_eq!(roundtrip(Config::default(), &values), values);
}
#[test]
fn compressed_roundtrip() {
let values = vec![
json!({ "a": 1, "b": 1 }),
json!({ "a": 1, "b": 2 }),
json!({ "a": 5, "b": 2 }),
];
assert_eq!(roundtrip(deflate(), &values), values);
}
#[test]
fn compressed_roundtrip_across_a_group_boundary() {
let values: Vec<Value> = (0..=40).map(|n| json!({ "n": n })).collect();
let config = deflate().with_delta_ratio(2);
assert_eq!(roundtrip(config, &values).last().unwrap(), &json!({ "n": 40 }));
}
#[test]
fn no_value_before_the_first_snapshot() {
let decoder = Decoder::<Value>::new(ConsumerConfig::default());
assert_eq!(decoder.value(), None);
assert_eq!(decoder.decode().unwrap(), None);
}
#[test]
fn a_delta_before_a_snapshot_is_an_error() {
let mut decoder = Decoder::<Value>::new(ConsumerConfig::default());
assert!(matches!(decoder.delta(br#"{"a":1}"#), Err(Error::MissingSnapshot)));
}
#[test]
fn frames_apply_without_materializing() {
let mut encoder = Encoder::<Value>::new(Config::default().with_delta_ratio(100));
let mut decoder = Decoder::<Value>::new(ConsumerConfig::default());
for n in 0..=20 {
let frame = encoder.update(&json!({ "n": n })).unwrap().unwrap();
match frame.keyframe {
true => decoder.snapshot(&frame.payload).unwrap(),
false => decoder.delta(&frame.payload).unwrap(),
}
frame.commit();
}
assert_eq!(decoder.decode().unwrap(), Some(json!({ "n": 20 })));
}
#[test]
fn a_rejected_field_names_its_path() {
#[derive(serde::Deserialize, Debug)]
#[allow(dead_code)]
struct Inner {
count: u8,
}
#[derive(serde::Deserialize, Debug)]
#[allow(dead_code)]
struct Outer {
inner: Inner,
}
let mut decoder = Decoder::<Outer>::new(ConsumerConfig::default());
decoder.snapshot(br#"{"inner":{"count":300}}"#).unwrap();
let err = decoder.decode().unwrap_err();
assert!(err.to_string().starts_with("json: inner.count: "), "{err}");
}
}