moq_json/snapshot/
decoder.rs1use std::cell::RefCell;
4use std::marker::PhantomData;
5
6use serde::de::DeserializeOwned;
7use serde_json::Value;
8
9use super::consumer::Config;
10use crate::{Error, Result};
11
12pub struct Decoder<T> {
32 compression: bool,
34 max_size: Option<usize>,
35
36 flate: Option<moq_flate::Decoder>,
38
39 plain: Vec<u8>,
41
42 check: RefCell<crate::merge::CheckScratch>,
44
45 current: Option<Value>,
47
48 _marker: PhantomData<fn() -> T>,
49}
50
51impl<T> Decoder<T> {
52 pub fn new(config: Config) -> Self {
54 Self {
55 compression: config.compression.is_deflate(),
56 max_size: config.max_size,
57 flate: None,
58 plain: Vec::new(),
59 check: RefCell::new(crate::merge::CheckScratch::default()),
60 current: None,
61 _marker: PhantomData,
62 }
63 }
64
65 pub fn snapshot(&mut self, payload: &[u8]) -> Result<()> {
70 self.current = None;
71 self.flate = self.compression.then(|| {
73 moq_flate::Decoder::with_max_frame_size(self.max_size.map_or(moq_flate::DEFAULT_MAX_FRAME_SIZE, |size| {
74 (size as u64).min(moq_flate::DEFAULT_MAX_FRAME_SIZE)
75 }))
76 });
77 let inflated = self.flate.as_mut().map(|flate| flate.frame(payload)).transpose()?;
78 let plain = inflated.as_deref().unwrap_or(payload);
79 if let Some(limit) = self.max_size
80 && plain.len() > limit
81 {
82 return Err(Error::TooLarge(limit));
83 }
84 self.current = Some(serde_json::from_slice(plain)?);
85 self.check_size()
86 }
87
88 pub fn delta(&mut self, payload: &[u8]) -> Result<()> {
91 if self.current.is_none() {
92 return Err(Error::MissingSnapshot);
93 }
94 let result = (|| {
95 let plain = match self.flate.as_mut() {
96 Some(flate) => {
97 flate.frame_into(payload, &mut self.plain)?;
98 self.plain.as_slice()
99 }
100 None => payload,
101 };
102 if let Some(limit) = self.max_size
103 && plain.len() > limit
104 {
105 return Err(Error::TooLarge(limit));
106 }
107 crate::merge::apply_bytes(
108 self.current.as_mut().expect("a snapshot precedes any delta"),
109 plain,
110 &self.check,
111 )?;
112 self.check_size()
113 })();
114 if result.is_err() {
115 self.current = None;
116 }
117 result
118 }
119
120 fn check_size(&mut self) -> Result<()> {
121 let Some(limit) = self.max_size else { return Ok(()) };
122 if serde_json::to_writer(Budget(limit), self.current.as_ref().unwrap()).is_err() {
126 self.current = None;
127 return Err(Error::TooLarge(limit));
128 }
129 Ok(())
130 }
131
132 pub fn value(&self) -> Option<&Value> {
134 self.current.as_ref()
135 }
136}
137
138struct Budget(usize);
139
140impl std::io::Write for Budget {
141 fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
142 if bytes.len() > self.0 {
143 return Err(std::io::Error::other("JSON size budget exceeded"));
144 }
145 self.0 -= bytes.len();
146 Ok(bytes.len())
147 }
148
149 fn flush(&mut self) -> std::io::Result<()> {
150 Ok(())
151 }
152}
153
154impl<T: DeserializeOwned> Decoder<T> {
155 pub fn decode(&self) -> Result<Option<T>> {
162 let Some(current) = self.current.as_ref() else {
163 return Ok(None);
164 };
165
166 if let Ok(value) = T::deserialize(current) {
169 return Ok(Some(value));
170 }
171
172 let value = serde_path_to_error::deserialize(current).map_err(|err| {
173 let path = err.path().to_string();
174 match path.as_str() {
175 "." => Error::Json(err.into_inner().to_string()),
177 _ => Error::Json(format!("{}: {}", path, err.into_inner())),
178 }
179 })?;
180
181 Ok(Some(value))
182 }
183}
184
185#[cfg(test)]
186mod test {
187 use super::super::consumer::Config as ConsumerConfig;
188 use super::super::{Config, Encoder};
189 use super::*;
190 use crate::Compression;
191 use serde_json::{Value, json};
192
193 fn consume(compression: Compression) -> ConsumerConfig {
194 ConsumerConfig {
195 compression,
196 ..Default::default()
197 }
198 }
199
200 fn deflate() -> Config {
201 Config {
202 compression: Compression::Deflate,
203 ..Default::default()
204 }
205 }
206
207 fn roundtrip(config: Config, values: &[Value]) -> Vec<Value> {
210 let compression = config.compression;
211 let mut encoder = Encoder::<Value>::new(config);
212 let mut decoder = Decoder::<Value>::new(consume(compression));
213
214 let mut out = Vec::new();
215 for value in values {
216 let Some(frame) = encoder.update(value).unwrap() else {
217 continue;
218 };
219 match frame.keyframe {
220 true => decoder.snapshot(&frame.payload).unwrap(),
221 false => decoder.delta(&frame.payload).unwrap(),
222 }
223 frame.commit();
224 out.push(decoder.decode().unwrap().unwrap());
225 }
226 out
227 }
228
229 #[test]
230 fn inflated_snapshots_and_patches_obey_the_budget() {
231 for snapshot in [true, false] {
232 let mut config = consume(Compression::Deflate);
233 config.max_size = Some(1024);
234 let mut decoder = Decoder::<Value>::new(config);
235 let mut encoder = moq_flate::Encoder::new();
236 if !snapshot {
237 decoder.snapshot(&encoder.frame(b"{}")).unwrap();
238 }
239 let payload = encoder.frame(
240 serde_json::to_string(&json!({"large": "x".repeat(4096)}))
241 .unwrap()
242 .as_bytes(),
243 );
244 assert!(payload.len() < 1024);
245 let result = if snapshot {
246 decoder.snapshot(&payload)
247 } else {
248 decoder.delta(&payload)
249 };
250 assert!(matches!(result, Err(Error::Flate(moq_flate::Error::TooLarge(1024)))));
251 assert!(decoder.value().is_none());
252 }
253 }
254
255 #[test]
256 fn patches_cannot_accumulate_past_the_budget() {
257 for compression in [Compression::None, Compression::Deflate] {
258 let mut config = consume(compression);
259 config.max_size = Some(13);
260 let mut decoder = Decoder::<Value>::new(config);
261 let mut encoder = moq_flate::Encoder::new();
262 let mut payload = |bytes: &[u8]| {
263 if compression == Compression::Deflate {
264 encoder.frame(bytes).to_vec()
265 } else {
266 bytes.to_vec()
267 }
268 };
269 decoder.snapshot(&payload(br#"{"a":1}"#)).unwrap();
270 decoder.delta(&payload(br#"{"b":2}"#)).unwrap(); assert_eq!(decoder.value(), Some(&json!({"a":1,"b":2})));
272 decoder.delta(&payload(br#"{"a":null}"#)).unwrap(); decoder.delta(&payload(br#"{"c":3}"#)).unwrap();
274 assert!(matches!(
275 decoder.delta(&payload(br#"{"d":4}"#)),
276 Err(Error::TooLarge(13))
277 ));
278 assert!(decoder.value().is_none());
279 assert!(matches!(decoder.delta(b"{}"), Err(Error::MissingSnapshot)));
280 let bytes = if compression == Compression::Deflate {
282 moq_flate::Encoder::new().frame(b"{}").to_vec()
283 } else {
284 b"{}".to_vec()
285 };
286 decoder.snapshot(&bytes).unwrap();
287 }
288 }
289
290 #[test]
291 fn plain_frames_obey_the_budget_before_parsing() {
292 let config = ConsumerConfig {
293 max_size: Some(2),
294 ..Default::default()
295 };
296 let mut decoder = Decoder::<Value>::new(config);
297 assert!(matches!(decoder.snapshot(b"not json"), Err(Error::TooLarge(2))));
298 decoder.snapshot(b"{}").unwrap();
299 assert!(matches!(decoder.delta(b"not json"), Err(Error::TooLarge(2))));
300 assert!(decoder.value().is_none());
301 }
302
303 #[test]
304 fn plaintext_roundtrip() {
305 let values = vec![
306 json!({ "a": 1, "b": 1 }),
307 json!({ "a": 1, "b": 2 }),
308 json!({ "a": 5, "b": 2 }),
309 ];
310 assert_eq!(roundtrip(Config::default(), &values), values);
311 }
312
313 #[test]
314 fn compressed_roundtrip() {
315 let values = vec![
316 json!({ "a": 1, "b": 1 }),
317 json!({ "a": 1, "b": 2 }),
318 json!({ "a": 5, "b": 2 }),
319 ];
320 assert_eq!(roundtrip(deflate(), &values), values);
321 }
322
323 #[test]
326 fn compressed_roundtrip_across_a_group_boundary() {
327 let values: Vec<Value> = (0..=40).map(|n| json!({ "n": n })).collect();
329 let config = deflate().with_delta_ratio(2);
330 assert_eq!(roundtrip(config, &values).last().unwrap(), &json!({ "n": 40 }));
331 }
332
333 #[test]
334 fn no_value_before_the_first_snapshot() {
335 let decoder = Decoder::<Value>::new(ConsumerConfig::default());
336 assert_eq!(decoder.value(), None);
337 assert_eq!(decoder.decode().unwrap(), None);
338 }
339
340 #[test]
341 fn a_delta_before_a_snapshot_is_an_error() {
342 let mut decoder = Decoder::<Value>::new(ConsumerConfig::default());
343 assert!(matches!(decoder.delta(br#"{"a":1}"#), Err(Error::MissingSnapshot)));
344 }
345
346 #[test]
349 fn frames_apply_without_materializing() {
350 let mut encoder = Encoder::<Value>::new(Config::default().with_delta_ratio(100));
351 let mut decoder = Decoder::<Value>::new(ConsumerConfig::default());
352
353 for n in 0..=20 {
354 let frame = encoder.update(&json!({ "n": n })).unwrap().unwrap();
355 match frame.keyframe {
356 true => decoder.snapshot(&frame.payload).unwrap(),
357 false => decoder.delta(&frame.payload).unwrap(),
358 }
359 frame.commit();
360 }
361
362 assert_eq!(decoder.decode().unwrap(), Some(json!({ "n": 20 })));
363 }
364
365 #[test]
366 fn a_rejected_field_names_its_path() {
367 #[derive(serde::Deserialize, Debug)]
368 #[allow(dead_code)]
369 struct Inner {
370 count: u8,
371 }
372 #[derive(serde::Deserialize, Debug)]
373 #[allow(dead_code)]
374 struct Outer {
375 inner: Inner,
376 }
377
378 let mut decoder = Decoder::<Outer>::new(ConsumerConfig::default());
379 decoder.snapshot(br#"{"inner":{"count":300}}"#).unwrap();
380
381 let err = decoder.decode().unwrap_err();
382 assert!(err.to_string().starts_with("json: inner.count: "), "{err}");
383 }
384}