1use serde::de as serde_de;
2use serde::de::{
3 self, DeserializeSeed, Deserializer, EnumAccess, MapAccess, SeqAccess, VariantAccess, Visitor,
4};
5use std::fmt;
6use std::marker::PhantomData;
7
8mod cost {
11 pub const UNIT: usize = 1;
12}
13
14struct Meter {
16 #[cfg(test)]
17 limit: usize,
18 remaining: usize,
19 exceeded: bool,
20}
21
22impl Meter {
23 pub fn new(limit: usize) -> Self {
25 Self {
26 #[cfg(test)]
27 limit,
28 remaining: limit,
29 exceeded: false,
30 }
31 }
32
33 pub fn wrap<'de, D: Deserializer<'de>>(
35 &mut self,
36 deserializer: D,
37 ) -> MeteredDeserializer<'_, D> {
38 MeteredDeserializer::new(self, deserializer)
39 }
40
41 #[cfg(test)]
42 fn spent(&self) -> usize {
43 self.limit - self.remaining
44 }
45
46 pub fn exceeded(&self) -> bool {
48 self.exceeded
49 }
50
51 pub fn spend<E: serde_de::Error>(&mut self, amount: usize) -> Result<(), E> {
54 match self.remaining.checked_sub(amount) {
55 Some(remaining) => {
56 self.remaining = remaining;
57 Ok(())
58 }
59 None => {
60 self.remaining = 0;
61 self.exceeded = true;
62 Err(serde_de::Error::custom(LimitExceeded {}))
63 }
64 }
65 }
66}
67
68struct LimitExceeded();
70
71impl fmt::Display for LimitExceeded {
72 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
73 write!(f, "deserialization exceeds the operation budget")
74 }
75}
76
77#[derive(Debug)]
79pub enum Error<E> {
80 LimitExceeded(usize),
82 Serde(E),
84}
85
86impl<E> Error<E> {
87 pub fn is_limit_exceeded(&self) -> bool {
89 matches!(self, Self::LimitExceeded(_))
90 }
91
92 pub fn into_serde(self) -> Option<E> {
94 match self {
95 Self::LimitExceeded(_) => None,
96 Self::Serde(error) => Some(error),
97 }
98 }
99}
100
101impl<E: fmt::Display> fmt::Display for Error<E> {
102 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
103 match self {
104 Self::LimitExceeded(limit) => {
105 write!(f, "value exceeds the {limit} operation limit")
106 }
107 Self::Serde(error) => error.fmt(f),
108 }
109 }
110}
111
112impl<E: std::error::Error + 'static> std::error::Error for Error<E> {
113 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
114 match self {
115 Self::LimitExceeded(_) => None,
116 Self::Serde(error) => Some(error),
117 }
118 }
119}
120
121pub fn deserialize<'de, T, D>(deserializer: D, max_ops: usize) -> Result<T, Error<D::Error>>
125where
126 T: serde_de::Deserialize<'de>,
127 D: Deserializer<'de>,
128{
129 deserialize_seed(PhantomData::<T>, deserializer, max_ops)
130}
131
132pub fn deserialize_seed<'de, S, D>(
137 seed: S,
138 deserializer: D,
139 max_ops: usize,
140) -> Result<S::Value, Error<D::Error>>
141where
142 S: DeserializeSeed<'de>,
143 D: Deserializer<'de>,
144{
145 let mut meter = Meter::new(max_ops);
146
147 match seed.deserialize(meter.wrap(deserializer)) {
148 Ok(value) => Ok(value),
149 Err(_) if meter.exceeded() => Err(Error::LimitExceeded(max_ops)),
152 Err(error) => Err(Error::Serde(error)),
153 }
154}
155
156pub struct MeteredDeserializer<'m, D> {
159 meter: &'m mut Meter,
160 inner: D,
161}
162
163impl<'m, 'de, D: Deserializer<'de>> MeteredDeserializer<'m, D> {
164 fn new(meter: &'m mut Meter, inner: D) -> Self {
165 Self { meter, inner }
166 }
167}
168
169macro_rules! forward {
171 ($($method:ident($($arg:ident: $ty:ty),*)),* $(,)?) => {
172 $(
173 fn $method<V: Visitor<'de>>(
174 self,
175 $($arg: $ty,)*
176 visitor: V,
177 ) -> Result<V::Value, Self::Error> {
178 let visitor = MeteredVisitor::new(self.meter, visitor);
179 self.inner.$method($($arg,)* visitor)
180 }
181 )*
182 };
183}
184
185impl<'de, D: Deserializer<'de>> Deserializer<'de> for MeteredDeserializer<'_, D> {
186 type Error = D::Error;
187
188 forward! {
189 deserialize_any(),
190 deserialize_bool(),
191 deserialize_i8(),
192 deserialize_i16(),
193 deserialize_i32(),
194 deserialize_i64(),
195 deserialize_i128(),
196 deserialize_u8(),
197 deserialize_u16(),
198 deserialize_u32(),
199 deserialize_u64(),
200 deserialize_u128(),
201 deserialize_f32(),
202 deserialize_f64(),
203 deserialize_char(),
204 deserialize_str(),
205 deserialize_string(),
206 deserialize_bytes(),
207 deserialize_byte_buf(),
208 deserialize_option(),
209 deserialize_unit(),
210 deserialize_seq(),
211 deserialize_map(),
212 deserialize_identifier(),
213 deserialize_ignored_any(),
214 deserialize_unit_struct(name: &'static str),
215 deserialize_newtype_struct(name: &'static str),
216 deserialize_tuple(len: usize),
217 deserialize_tuple_struct(name: &'static str, len: usize),
218 deserialize_struct(name: &'static str, fields: &'static [&'static str]),
219 deserialize_enum(name: &'static str, variants: &'static [&'static str]),
220 }
221
222 fn is_human_readable(&self) -> bool {
223 self.inner.is_human_readable()
224 }
225}
226
227struct MeteredSeed<'m, S> {
229 meter: &'m mut Meter,
230 inner: S,
231}
232
233impl<'de, S: DeserializeSeed<'de>> DeserializeSeed<'de> for MeteredSeed<'_, S> {
234 type Value = S::Value;
235
236 fn deserialize<D: Deserializer<'de>>(self, deserializer: D) -> Result<Self::Value, D::Error> {
237 self.inner
238 .deserialize(MeteredDeserializer::new(self.meter, deserializer))
239 }
240}
241
242struct MeteredVisitor<'m, V> {
244 meter: &'m mut Meter,
245 inner: V,
246}
247
248impl<'m, V> MeteredVisitor<'m, V> {
249 fn new(meter: &'m mut Meter, inner: V) -> Self {
250 Self { meter, inner }
251 }
252}
253
254macro_rules! visit_scalar {
256 ($($method:ident($ty:ty)),* $(,)?) => {
257 $(
258 fn $method<E: de::Error>(self, v: $ty) -> Result<Self::Value, E> {
259 self.meter.spend(cost::UNIT)?;
260 self.inner.$method(v)
261 }
262 )*
263 };
264}
265
266impl<'de, V: Visitor<'de>> Visitor<'de> for MeteredVisitor<'_, V> {
267 type Value = V::Value;
268
269 fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
270 self.inner.expecting(f)
271 }
272
273 visit_scalar! {
274 visit_bool(bool),
275 visit_i8(i8),
276 visit_i16(i16),
277 visit_i32(i32),
278 visit_i64(i64),
279 visit_i128(i128),
280 visit_u8(u8),
281 visit_u16(u16),
282 visit_u32(u32),
283 visit_u64(u64),
284 visit_u128(u128),
285 visit_f32(f32),
286 visit_f64(f64),
287 visit_char(char),
288 }
289
290 fn visit_str<E: de::Error>(self, v: &str) -> Result<Self::Value, E> {
291 self.meter.spend(cost::UNIT)?;
292 self.inner.visit_str(v)
293 }
294
295 fn visit_borrowed_str<E: de::Error>(self, v: &'de str) -> Result<Self::Value, E> {
296 self.meter.spend(cost::UNIT)?;
297 self.inner.visit_borrowed_str(v)
298 }
299
300 fn visit_string<E: de::Error>(self, v: String) -> Result<Self::Value, E> {
301 self.meter.spend(cost::UNIT)?;
302 self.inner.visit_string(v)
303 }
304
305 fn visit_bytes<E: de::Error>(self, v: &[u8]) -> Result<Self::Value, E> {
306 self.meter.spend(cost::UNIT)?;
307 self.inner.visit_bytes(v)
308 }
309
310 fn visit_borrowed_bytes<E: de::Error>(self, v: &'de [u8]) -> Result<Self::Value, E> {
311 self.meter.spend(cost::UNIT)?;
312 self.inner.visit_borrowed_bytes(v)
313 }
314
315 fn visit_byte_buf<E: de::Error>(self, v: Vec<u8>) -> Result<Self::Value, E> {
316 self.meter.spend(cost::UNIT)?;
317 self.inner.visit_byte_buf(v)
318 }
319
320 fn visit_none<E: de::Error>(self) -> Result<Self::Value, E> {
321 self.meter.spend(cost::UNIT)?;
322 self.inner.visit_none()
323 }
324
325 fn visit_unit<E: de::Error>(self) -> Result<Self::Value, E> {
326 self.meter.spend(cost::UNIT)?;
327 self.inner.visit_unit()
328 }
329
330 fn visit_some<D: Deserializer<'de>>(self, d: D) -> Result<Self::Value, D::Error> {
331 self.inner
333 .visit_some(MeteredDeserializer::new(self.meter, d))
334 }
335
336 fn visit_newtype_struct<D: Deserializer<'de>>(self, d: D) -> Result<Self::Value, D::Error> {
337 self.inner
339 .visit_newtype_struct(MeteredDeserializer::new(self.meter, d))
340 }
341
342 fn visit_seq<A: SeqAccess<'de>>(self, seq: A) -> Result<Self::Value, A::Error> {
343 self.meter.spend(cost::UNIT)?;
344 self.inner.visit_seq(MeteredSeqAccess {
345 meter: self.meter,
346 inner: seq,
347 })
348 }
349
350 fn visit_map<A: MapAccess<'de>>(self, map: A) -> Result<Self::Value, A::Error> {
351 self.meter.spend(cost::UNIT)?;
352
353 self.inner.visit_map(MeteredMapAccess {
354 meter: self.meter,
355 inner: map,
356 })
357 }
358
359 fn visit_enum<A: EnumAccess<'de>>(self, data: A) -> Result<Self::Value, A::Error> {
360 self.inner.visit_enum(MeteredEnumAccess {
361 meter: self.meter,
362 inner: data,
363 })
364 }
365}
366
367struct MeteredSeqAccess<'m, A> {
369 meter: &'m mut Meter,
370 inner: A,
371}
372
373impl<'de, A: SeqAccess<'de>> SeqAccess<'de> for MeteredSeqAccess<'_, A> {
374 type Error = A::Error;
375
376 fn next_element_seed<T: DeserializeSeed<'de>>(
377 &mut self,
378 seed: T,
379 ) -> Result<Option<T::Value>, Self::Error> {
380 let element = self.inner.next_element_seed(MeteredSeed {
381 meter: self.meter,
382 inner: seed,
383 })?;
384
385 Ok(element)
386 }
387
388 fn size_hint(&self) -> Option<usize> {
389 self.inner.size_hint()
390 }
391}
392
393struct MeteredMapAccess<'m, A> {
395 meter: &'m mut Meter,
396 inner: A,
397}
398
399impl<'de, A: MapAccess<'de>> MapAccess<'de> for MeteredMapAccess<'_, A> {
400 type Error = A::Error;
401
402 fn next_key_seed<K: DeserializeSeed<'de>>(
403 &mut self,
404 seed: K,
405 ) -> Result<Option<K::Value>, Self::Error> {
406 let key = self.inner.next_key_seed(MeteredSeed {
407 meter: self.meter,
408 inner: seed,
409 })?;
410
411 Ok(key)
412 }
413
414 fn next_value_seed<Va: DeserializeSeed<'de>>(
415 &mut self,
416 seed: Va,
417 ) -> Result<Va::Value, Self::Error> {
418 self.inner.next_value_seed(MeteredSeed {
419 meter: self.meter,
420 inner: seed,
421 })
422 }
423
424 fn size_hint(&self) -> Option<usize> {
425 self.inner.size_hint()
426 }
427}
428
429struct MeteredEnumAccess<'m, A> {
431 meter: &'m mut Meter,
432 inner: A,
433}
434
435impl<'de, 'm, A: EnumAccess<'de>> EnumAccess<'de> for MeteredEnumAccess<'m, A> {
436 type Error = A::Error;
437 type Variant = MeteredVariantAccess<'m, A::Variant>;
438
439 fn variant_seed<S: DeserializeSeed<'de>>(
440 self,
441 seed: S,
442 ) -> Result<(S::Value, Self::Variant), Self::Error> {
443 let meter = self.meter;
444 let (value, variant) = self
445 .inner
446 .variant_seed(MeteredSeed { meter, inner: seed })?;
447
448 Ok((
449 value,
450 MeteredVariantAccess {
451 meter,
452 inner: variant,
453 },
454 ))
455 }
456}
457
458struct MeteredVariantAccess<'m, A> {
460 meter: &'m mut Meter,
461 inner: A,
462}
463
464impl<'de, A: VariantAccess<'de>> VariantAccess<'de> for MeteredVariantAccess<'_, A> {
465 type Error = A::Error;
466
467 fn unit_variant(self) -> Result<(), Self::Error> {
468 self.meter.spend(cost::UNIT)?;
469 self.inner.unit_variant()
470 }
471
472 fn newtype_variant_seed<S: DeserializeSeed<'de>>(
473 self,
474 seed: S,
475 ) -> Result<S::Value, Self::Error> {
476 self.inner.newtype_variant_seed(MeteredSeed {
477 meter: self.meter,
478 inner: seed,
479 })
480 }
481
482 fn tuple_variant<V: Visitor<'de>>(
483 self,
484 len: usize,
485 visitor: V,
486 ) -> Result<V::Value, Self::Error> {
487 self.inner
488 .tuple_variant(len, MeteredVisitor::new(self.meter, visitor))
489 }
490
491 fn struct_variant<V: Visitor<'de>>(
492 self,
493 fields: &'static [&'static str],
494 visitor: V,
495 ) -> Result<V::Value, Self::Error> {
496 self.inner
497 .struct_variant(fields, MeteredVisitor::new(self.meter, visitor))
498 }
499}
500
501#[cfg(test)]
502mod tests {
503 use std::collections::BTreeMap;
504
505 use serde::Deserialize;
506
507 use super::*;
508
509 fn deserialize_return_meter<'de, T, D>(deserializer: D, max_ops: usize) -> (T, Meter)
510 where
511 D: Deserializer<'de>,
512 T: Deserialize<'de>,
513 {
514 let mut meter = Meter::new(max_ops);
515
516 let metered_deserializer = meter.wrap(deserializer);
517 let t = T::deserialize(metered_deserializer).unwrap();
518
519 (t, meter)
520 }
521
522 fn json_deserializer(payload: &str) -> serde_json::Deserializer<serde_json::de::StrRead<'_>> {
523 serde_json::Deserializer::from_str(payload)
524 }
525
526 #[derive(Debug, Deserialize, PartialEq)]
527 struct Nested {
528 name: Vec<String>,
529 values: Vec<u64>,
530 inner: Option<Box<Nested>>,
531 }
532
533 #[test]
534 fn test_deserialize_exceeds_budget() {
535 let payload = format!(r#"[{}]"#, vec!["\"a\""; 4096].join(","));
536
537 let error =
538 deserialize::<Vec<String>, _>(&mut json_deserializer(&payload), 128).unwrap_err();
539 assert!(error.is_limit_exceeded());
540 }
541
542 #[test]
543 fn test_deserialize_exceeds_budget_when_nested() {
544 let payload = format!(
545 r#"{{"name": [{}], "values": [], "inner": null}}"#,
546 vec!["\"a\""; 4096].join(",")
547 );
548
549 let error = deserialize::<Nested, _>(&mut json_deserializer(&payload), 256).unwrap_err();
550 assert!(error.is_limit_exceeded());
551 }
552
553 #[test]
554 fn test_deserialize_empty_values_are_not_free() {
555 let payload = format!("[{}]", vec!["{}"; 10_000].join(","));
557
558 let error = deserialize::<Vec<serde_json::Map<String, serde_json::Value>>, _>(
559 &mut json_deserializer(&payload),
560 1 << 8,
561 )
562 .unwrap_err();
563 assert!(error.is_limit_exceeded());
564 }
565
566 #[test]
567 fn test_deserialize_invalid_payload() {
568 let error =
569 deserialize::<Nested, _>(&mut json_deserializer("{invalid"), 1 << 20).unwrap_err();
570 assert!(!error.is_limit_exceeded());
571 assert!(error.into_serde().is_some());
572 }
573
574 #[test]
575 fn test_deserialize_does_not_consume_trailing_input() {
576 let mut de = json_deserializer("[1] trailing");
579 let value: Vec<u64> = deserialize(&mut de, 1 << 20).unwrap();
580 assert_eq!(value, [1]);
581 assert!(de.end().is_err());
582 }
583
584 fn prim_array_cost(num_elements: usize) -> usize {
585 cost::UNIT + (cost::UNIT * num_elements)
586 }
587
588 fn str_cost() -> usize {
589 cost::UNIT
590 }
591
592 fn prim_cost() -> usize {
593 cost::UNIT
594 }
595
596 fn map_cost() -> usize {
597 cost::UNIT
598 }
599
600 fn scalar_struct_cost(num_fields: usize) -> usize {
601 map_cost() + num_fields * (str_cost() + prim_cost())
602 }
603
604 fn variant_cost(payload: usize) -> usize {
605 str_cost() + payload
606 }
607
608 #[derive(Debug, Deserialize, PartialEq)]
609 struct UnitStruct;
610
611 #[derive(Debug, Deserialize, PartialEq)]
612 struct NewtypeStruct(u32);
613
614 #[derive(Debug, Deserialize, PartialEq)]
615 struct TupleStruct(u8, String, bool);
616
617 #[derive(Debug, Deserialize, PartialEq)]
618 enum Variants {
619 Unit,
620 Newtype(i32),
621 Tuple(u8, String),
622 Struct { key: char, value: Option<f64> },
623 }
624
625 #[derive(Debug, Deserialize, PartialEq)]
626 struct Signed {
627 a: i8,
628 b: i16,
629 c: i32,
630 d: i64,
631 e: i128,
632 }
633
634 #[derive(Debug, Deserialize, PartialEq)]
635 struct Unsigned {
636 a: u8,
637 b: u16,
638 c: u32,
639 d: u64,
640 e: u128,
641 }
642
643 #[derive(Debug, Deserialize, PartialEq)]
645 struct Everything<'a> {
646 flag: bool,
647 signed: Signed,
648 unsigned: Unsigned,
649 f32_: f32,
650 f64_: f64,
651 letter: char,
652 owned: String,
653 #[serde(borrow)]
654 borrowed: &'a str,
655 nothing: (),
656 unit_struct: UnitStruct,
657 newtype: NewtypeStruct,
658 tuple: (u8, String, bool),
659 tuple_struct: TupleStruct,
660 some: Option<u32>,
661 none: Option<u32>,
662 seq: Vec<Vec<u8>>,
663 map: BTreeMap<String, Variants>,
664 variants: Vec<Variants>,
665 nested: Option<Box<Nested>>,
666 }
667
668 const EVERYTHING_PAYLOAD: &str = r#"{
670 "flag": true,
671 "signed": {"a": -8, "b": -16, "c": -32, "d": -64, "e": -128},
672 "unsigned": {"a": 8, "b": 16, "c": 32, "d": 64, "e": 128},
673 "f32_": 1.5,
674 "f64_": -2.25,
675 "letter": "x",
676 "owned": "owned string",
677 "borrowed": "borrowed string",
678 "nothing": null,
679 "unit_struct": null,
680 "newtype": 7,
681 "tuple": [1, "two", false],
682 "tuple_struct": [2, "three", true],
683 "some": 9,
684 "none": null,
685 "seq": [[1, 2], [], [3]],
686 "map": {"unit": "Unit", "newtype": {"Newtype": -5}},
687 "variants": [
688 "Unit",
689 {"Newtype": 5},
690 {"Tuple": [1, "one"]},
691 {"Struct": {"key": "k", "value": 0.5}},
692 {"Struct": {"key": "l", "value": null}}
693 ],
694 "nested": {"name": ["a", "b"], "values": [1], "inner": null},
695 "ignored": {"deeply": [{"nested": [1, 2, 3]}]}
696 }"#;
697
698 #[test]
699 fn test_deserialize_all_paths() {
700 let mut de = json_deserializer(EVERYTHING_PAYLOAD);
701 let value: Everything<'_> = deserialize(&mut de, 1 << 20).unwrap();
702
703 assert_eq!(
704 value,
705 Everything {
706 flag: true,
707 signed: Signed {
708 a: -8,
709 b: -16,
710 c: -32,
711 d: -64,
712 e: -128,
713 },
714 unsigned: Unsigned {
715 a: 8,
716 b: 16,
717 c: 32,
718 d: 64,
719 e: 128,
720 },
721 f32_: 1.5,
722 f64_: -2.25,
723 letter: 'x',
724 owned: "owned string".to_owned(),
725 borrowed: "borrowed string",
726 nothing: (),
727 unit_struct: UnitStruct,
728 newtype: NewtypeStruct(7),
729 tuple: (1, "two".to_owned(), false),
730 tuple_struct: TupleStruct(2, "three".to_owned(), true),
731 some: Some(9),
732 none: None,
733 seq: vec![vec![1, 2], vec![], vec![3]],
734 map: BTreeMap::from([
735 ("unit".to_owned(), Variants::Unit),
736 ("newtype".to_owned(), Variants::Newtype(-5)),
737 ]),
738 variants: vec![
739 Variants::Unit,
740 Variants::Newtype(5),
741 Variants::Tuple(1, "one".to_owned()),
742 Variants::Struct {
743 key: 'k',
744 value: Some(0.5),
745 },
746 Variants::Struct {
747 key: 'l',
748 value: None,
749 },
750 ],
751 nested: Some(Box::new(Nested {
752 name: vec!["a".to_owned(), "b".to_owned()],
753 values: vec![1],
754 inner: None,
755 })),
756 }
757 );
758 }
759
760 #[test]
761 fn test_expected_costs_everything() {
762 let mut de = json_deserializer(EVERYTHING_PAYLOAD);
763 let (_, meter): (Everything<'_>, _) = deserialize_return_meter(&mut de, 1 << 20);
764
765 let values = [
768 prim_cost(),
770 scalar_struct_cost(5),
772 scalar_struct_cost(5),
773 prim_cost(),
775 prim_cost(),
776 prim_cost(),
777 str_cost(),
779 str_cost(),
780 prim_cost(),
782 prim_cost(),
783 prim_cost(),
785 prim_array_cost(3),
787 prim_array_cost(3),
788 prim_cost(),
790 prim_cost(),
792 cost::UNIT + prim_array_cost(2) + prim_array_cost(0) + prim_array_cost(1),
794 map_cost()
796 + (str_cost() + variant_cost(prim_cost()))
797 + (str_cost() + variant_cost(prim_cost())),
798 cost::UNIT
800 + variant_cost(prim_cost())
801 + variant_cost(prim_cost())
802 + variant_cost(prim_array_cost(2))
803 + variant_cost(scalar_struct_cost(2))
804 + variant_cost(scalar_struct_cost(2)),
805 map_cost()
807 + (str_cost() + prim_array_cost(2))
808 + (str_cost() + prim_array_cost(1))
809 + (str_cost() + prim_cost()),
810 cost::UNIT,
812 ];
813
814 let expected = map_cost() + values.len() * str_cost() + values.iter().sum::<usize>();
815 assert_eq!(meter.spent(), expected);
816 }
817}