1use serde::de::{
11 self, value::BorrowedBytesDeserializer, DeserializeSeed, Deserializer, MapAccess, SeqAccess,
12 Visitor,
13};
14use serde::forward_to_deserialize_any;
15
16use crate::error::MaxMindDbError;
17
18mod key;
19mod verification;
20use key::{CachedKey, KeyDeserializer};
21pub(crate) use verification::VerificationState;
22
23const TYPE_EXTENDED: usize = 0;
25pub(crate) const TYPE_POINTER: usize = 1;
26const TYPE_STRING: usize = 2;
27const TYPE_DOUBLE: usize = 3;
28const TYPE_BYTES: usize = 4;
29const TYPE_UINT16: usize = 5;
30const TYPE_UINT32: usize = 6;
31pub(crate) const TYPE_MAP: usize = 7;
32const TYPE_INT32: usize = 8;
33const TYPE_UINT64: usize = 9;
34const TYPE_UINT128: usize = 10;
35pub(crate) const TYPE_ARRAY: usize = 11;
36const TYPE_BOOL: usize = 14;
37const TYPE_FLOAT: usize = 15;
38
39const RAW_STRINGS_NEWTYPE: &str = "$maxminddb::raw_strings";
40
41const MAXIMUM_DATA_STRUCTURE_DEPTH: u16 = 512;
44
45const MAXIMUM_DATA_STRUCTURE_VALUES: u32 = 1 << 16;
50
51const MAXIMUM_DATA_STRUCTURE_BYTES: usize = 2 << 20;
58
59const MAXIMUM_UNCHARGED_IDENTIFIER_BYTES: usize =
64 MAXIMUM_DATA_STRUCTURE_BYTES / MAXIMUM_DATA_STRUCTURE_VALUES as usize;
65
66const DEPTH_MASK: u32 = (1 << 10) - 1;
70const BUDGET_VALUES_SHIFT: u32 = 10;
71const BUDGET_VALUES_MASK: u32 = ((1 << 17) - 1) << BUDGET_VALUES_SHIFT;
72const BUDGET_ACTIVE_MASK: u32 = 1 << 27;
73
74const MAXIMUM_SKIPPED_DATA_STRUCTURE_DEPTH: u16 = 128;
78
79#[cfg(not(feature = "unsafe-str-decode"))]
80#[inline]
81fn is_ascii(bytes: &[u8]) -> bool {
82 match bytes.len() {
84 4..=7 => {
85 let first = u32::from_ne_bytes(bytes[..4].try_into().unwrap());
86 let last = u32::from_ne_bytes(bytes[bytes.len() - 4..].try_into().unwrap());
87 (first | last) & 0x8080_8080 == 0
88 }
89 8..=16 => {
90 let first = u64::from_ne_bytes(bytes[..8].try_into().unwrap());
91 let last = u64::from_ne_bytes(bytes[bytes.len() - 8..].try_into().unwrap());
92 (first | last) & 0x8080_8080_8080_8080 == 0
93 }
94 _ => bytes.is_ascii(),
95 }
96}
97
98macro_rules! decode_int_like {
99 ($name:ident, $ty:ty, $max_size:expr, $label:literal, $zero:expr) => {
100 fn $name(&mut self, size: usize) -> DecodeResult<$ty> {
101 match size {
102 s if s <= $max_size => {
103 let new_offset = self
104 .current_ptr
105 .checked_add(size)
106 .filter(|&offset| offset <= self.limit)
107 .ok_or_else(|| {
108 self.invalid_db_error(&format!("{} of size {}", $label, size))
109 })?;
110 let value = self
111 .slice(self.current_ptr, new_offset)
112 .iter()
113 .fold($zero, |acc, &b| (acc << 8) | <$ty>::from(b));
114 self.current_ptr = new_offset;
115 Ok(value)
116 }
117 s => Err(self.invalid_db_error(&format!("{} of size {}", $label, s))),
118 }
119 }
120 };
121}
122
123macro_rules! deserialize_direct_scalar {
124 ($name:ident, $expected_type:expr, $label:literal, $visit:ident, $decode:ident) => {
125 fn $name<V>(self, visitor: V) -> DecodeResult<V::Value>
126 where
127 V: Visitor<'de>,
128 {
129 let (size, type_num) = self.size_and_type()?;
130 self.decode_direct(size, type_num, $expected_type, $label, |de, size| {
131 visitor.$visit(de.$decode(size)?)
132 })
133 }
134 };
135}
136
137macro_rules! deserialize_direct_payload {
138 ($name:ident, $expected_type:expr, $label:literal, $visit:ident, $decode:ident) => {
139 #[cfg_attr(feature = "unsafe-str-decode", inline(always))]
140 fn $name<V>(self, visitor: V) -> DecodeResult<V::Value>
141 where
142 V: Visitor<'de>,
143 {
144 let (size, type_num) = self.size_and_type()?;
145 self.decode_direct(size, type_num, $expected_type, $label, |de, size| {
146 de.count_payload(size)?;
147 visitor.$visit(de.$decode(size)?)
148 })
149 }
150 };
151}
152
153enum Value<'a, 'de> {
154 Any { prev_ptr: usize },
155 Bytes(&'de [u8]),
156 String(&'de str),
157 RawString(&'de [u8]),
158 Bool(bool),
159 I32(i32),
160 U16(u16),
161 U32(u32),
162 U64(u64),
163 U128(u128),
164 F64(f64),
165 F32(f32),
166 Map(MapAccessor<'a, 'de, true>),
167 Array(ArrayAccess<'a, 'de>),
168}
169
170#[derive(Debug)]
175pub(crate) struct Decoder<'de> {
176 buf: &'de [u8],
177 limit: usize,
178 current_ptr: usize,
179 state: u32,
180 payload_remaining: u32,
181}
182
183impl<'de> Decoder<'de> {
184 pub(crate) fn new(buf: &'de [u8], start_ptr: usize) -> Decoder<'de> {
185 Decoder::new_with_limit(buf, start_ptr, buf.len())
186 }
187
188 pub(crate) fn new_with_limit(buf: &'de [u8], start_ptr: usize, limit: usize) -> Decoder<'de> {
189 debug_assert!(limit <= buf.len());
190 Decoder {
191 buf,
192 limit,
193 current_ptr: start_ptr,
194 state: 0,
195 payload_remaining: MAXIMUM_DATA_STRUCTURE_BYTES as u32,
196 }
197 }
198
199 #[inline(always)]
200 fn activate_budget(&mut self) {
201 if self.state & BUDGET_ACTIVE_MASK == 0 {
202 self.state |= BUDGET_ACTIVE_MASK | (1 << BUDGET_VALUES_SHIFT);
204 }
205 }
206
207 #[inline]
209 fn enter_nested(&mut self) -> DecodeResult<()> {
210 if self.state & DEPTH_MASK >= u32::from(MAXIMUM_DATA_STRUCTURE_DEPTH) {
211 return Err(self.invalid_db_error(
212 "exceeded maximum data structure depth; database is likely corrupt",
213 ));
214 }
215 self.state += 1;
216 Ok(())
217 }
218
219 #[inline]
221 fn exit_nested(&mut self) {
222 if self.state & DEPTH_MASK != 0 {
223 self.state -= 1;
224 }
225 }
226
227 #[inline(always)]
231 fn reserve_values(&mut self, count: usize) -> DecodeResult<()> {
232 if self.state & BUDGET_ACTIVE_MASK == 0 {
233 return Ok(());
234 }
235 let values_used = ((self.state & BUDGET_VALUES_MASK) >> BUDGET_VALUES_SHIFT) as usize;
236 if count > MAXIMUM_DATA_STRUCTURE_VALUES as usize - values_used {
237 return Err(
238 self.resource_limit_error("exceeded maximum number of data structure values")
239 );
240 }
241 self.state += (count as u32) << BUDGET_VALUES_SHIFT;
242 Ok(())
243 }
244
245 #[inline(always)]
246 fn reserve_container_values(&mut self, size: usize, type_num: usize) -> DecodeResult<()> {
247 self.activate_budget();
251 let count = if type_num == TYPE_MAP {
252 size.saturating_mul(2)
253 } else {
254 debug_assert_eq!(type_num, TYPE_ARRAY);
255 size
256 };
257 self.reserve_values(count)
258 }
259
260 #[inline(always)]
268 fn count_payload(&mut self, size: usize) -> DecodeResult<()> {
269 if self.state & BUDGET_ACTIVE_MASK == 0 {
270 return Ok(());
271 }
272 if size > self.payload_remaining as usize {
273 return Err(self.resource_limit_error(
274 "exceeded maximum size of data structure string and bytes values",
275 ));
276 }
277 self.payload_remaining -= size as u32;
278 Ok(())
279 }
280
281 #[inline(always)]
285 fn count_identifier_payload(&mut self, size: usize) -> DecodeResult<()> {
286 if size <= MAXIMUM_UNCHARGED_IDENTIFIER_BYTES {
287 return Ok(());
288 }
289 self.count_long_identifier_payload(size)
290 }
291
292 #[cold]
293 #[inline(never)]
294 fn count_long_identifier_payload(&mut self, size: usize) -> DecodeResult<()> {
295 self.count_payload(size - MAXIMUM_UNCHARGED_IDENTIFIER_BYTES)
296 }
297
298 #[inline]
303 fn count_payload_at_current(&mut self) -> DecodeResult<(bool, Option<CachedKey<'de>>)> {
304 let saved_ptr = self.current_ptr;
305 let result = self.count_payload_at_current_inner();
306 self.current_ptr = saved_ptr;
307 result
308 }
309
310 #[inline(always)]
314 fn count_payload_at_current_inner(&mut self) -> DecodeResult<(bool, Option<CachedKey<'de>>)> {
315 let (mut size, mut type_num) = self.size_and_type()?;
316 let mut continuation = None;
317 if type_num == TYPE_POINTER {
318 let target = self.decode_pointer(size);
319 continuation = Some(self.current_ptr);
320 self.current_ptr = target;
321 (size, type_num) = self.size_and_type()?;
322 if type_num == TYPE_POINTER {
323 return Err(self.invalid_db_error("pointer points to another pointer"));
324 }
325 }
326 if type_num == TYPE_STRING || type_num == TYPE_BYTES {
327 self.count_payload(size)?;
328 let cached = if type_num == TYPE_STRING {
332 self.current_ptr
333 .checked_add(size)
334 .filter(|&end| end <= self.limit)
335 .map(|end| CachedKey {
336 bytes: self.slice(self.current_ptr, end),
337 continuation: continuation.unwrap_or(end),
338 })
339 } else {
340 None
341 };
342 return Ok((true, cached));
343 }
344 Ok((false, None))
345 }
346
347 #[inline]
349 fn invalid_db_error(&self, msg: &str) -> DecoderError {
350 MaxMindDbError::invalid_database_at(msg, self.current_ptr).into()
351 }
352
353 #[inline]
355 fn decode_error(&self, msg: &str) -> DecoderError {
356 MaxMindDbError::decoding_at(msg, self.current_ptr).into()
357 }
358
359 #[cold]
361 #[inline(never)]
362 fn resource_limit_error(&self, msg: &str) -> DecoderError {
363 MaxMindDbError::resource_limit_at(msg, self.current_ptr).into()
364 }
365
366 #[inline(always)]
367 fn type_mismatch(&self, label: &str, type_num: usize) -> DecoderError {
368 if type_num > usize::from(u8::MAX) {
369 self.invalid_db_error(&format!("unknown data type: {type_num}"))
370 } else {
371 self.decode_error(&format!("expected {label}, got type {type_num}"))
372 }
373 }
374
375 #[inline]
376 pub(crate) fn offset(&self) -> usize {
377 self.current_ptr
378 }
379
380 #[inline(always)]
381 fn checked_offset(&self, size: usize, label: &str) -> DecodeResult<usize> {
382 let new_offset = self.current_ptr.wrapping_add(size);
383 if new_offset < self.current_ptr || new_offset > self.limit {
384 return Err(self.invalid_db_error(&format!("{label} of size {size}")));
385 }
386 Ok(new_offset)
387 }
388
389 #[inline(always)]
390 fn slice(&self, start: usize, end: usize) -> &'de [u8] {
391 debug_assert!(start <= end);
392 debug_assert!(end <= self.limit);
393 debug_assert!(self.limit <= self.buf.len());
394 unsafe { self.buf.get_unchecked(start..end) }
397 }
398
399 #[inline(always)]
400 fn skip_bytes(&mut self, size: usize, label: &str) -> DecodeResult<()> {
401 debug_assert!(self.current_ptr <= self.limit);
402 if size > self.limit - self.current_ptr {
403 return Err(self.invalid_db_error(&format!("{label} of size {size}")));
404 }
405 self.current_ptr += size;
406 Ok(())
407 }
408
409 #[inline(always)]
410 fn eat_byte(&mut self) -> DecodeResult<u8> {
411 if self.current_ptr >= self.limit {
412 return Err(self.invalid_db_error("unexpected end of buffer"));
413 }
414 debug_assert!(self.limit <= self.buf.len());
415 let b = unsafe { *self.buf.get_unchecked(self.current_ptr) };
418 self.current_ptr += 1;
419 Ok(b)
420 }
421
422 #[inline(always)]
423 fn size_from_ctrl_byte(&mut self, ctrl_byte: u8, type_num: usize) -> DecodeResult<usize> {
424 let size = (ctrl_byte & 0x1f) as usize;
425 if type_num == TYPE_EXTENDED {
427 return Ok(size);
428 }
429
430 match size {
431 s if s < 29 => Ok(s),
432 29 => Ok(29_usize + self.eat_byte()? as usize),
433 30 => {
434 let b0 = self.eat_byte()? as usize;
435 let b1 = self.eat_byte()? as usize;
436 Ok(285_usize + (b0 << 8) + b1)
437 }
438 _ => {
439 let b0 = self.eat_byte()? as usize;
440 let b1 = self.eat_byte()? as usize;
441 let b2 = self.eat_byte()? as usize;
442 Ok(65_821_usize + (b0 << 16) + (b1 << 8) + b2)
443 }
444 }
445 }
446
447 #[inline(always)]
448 fn size_and_type(&mut self) -> DecodeResult<(usize, usize)> {
449 let ctrl_byte = self.eat_byte()?;
450 let mut type_num = usize::from(ctrl_byte >> 5);
451 if type_num == TYPE_EXTENDED {
453 type_num = usize::from(self.eat_byte()?) + TYPE_MAP;
455 }
456 self.size_from_ctrl_byte(ctrl_byte, type_num)
457 .map(|size| (size, type_num))
458 }
459
460 fn decode_any<V: Visitor<'de>>(&mut self, visitor: V) -> DecodeResult<V::Value> {
461 self.activate_budget();
462 self.decode_any_impl::<false, V>(visitor)
463 }
464
465 fn decode_any_impl<const RAW_STRINGS: bool, V: Visitor<'de>>(
466 &mut self,
467 visitor: V,
468 ) -> DecodeResult<V::Value> {
469 match self.decode_any_value::<RAW_STRINGS>()? {
470 Value::Any { prev_ptr } => {
471 self.enter_nested()?;
473 let res = self.decode_any_impl::<RAW_STRINGS, V>(visitor);
474 self.exit_nested();
475 self.current_ptr = prev_ptr;
476 res
477 }
478 Value::Bool(x) => visitor.visit_bool(x),
479 Value::Bytes(x) => visitor.visit_borrowed_bytes(x),
480 Value::String(x) => visitor.visit_borrowed_str(x),
481 Value::RawString(x) => {
482 visitor.visit_newtype_struct(BorrowedBytesDeserializer::<DecoderError>::new(x))
483 }
484 Value::I32(x) => visitor.visit_i32(x),
485 Value::U16(x) => visitor.visit_u16(x),
486 Value::U32(x) => visitor.visit_u32(x),
487 Value::U64(x) => visitor.visit_u64(x),
488 Value::U128(x) => visitor.visit_u128(x),
489 Value::F64(x) => visitor.visit_f64(x),
490 Value::F32(x) => visitor.visit_f32(x),
491 Value::Map(x) => {
493 let res = visitor.visit_map(x);
494 self.exit_nested();
495 res
496 }
497 Value::Array(x) => {
498 let res = visitor.visit_seq(x);
499 self.exit_nested();
500 res
501 }
502 }
503 }
504
505 fn deserialize_fixed_size_array<V>(&mut self, len: usize, visitor: V) -> DecodeResult<V::Value>
506 where
507 V: Visitor<'de>,
508 {
509 let (size, type_num) = self.size_and_type()?;
510 self.decode_direct(size, type_num, TYPE_ARRAY, "array", |de, size| {
511 if size != len {
512 return Err(de.decode_error(&format!(
513 "expected tuple of length {len}, got array of length {size}"
514 )));
515 }
516
517 de.reserve_container_values(size, TYPE_ARRAY)?;
518 de.enter_nested()?;
519 let res = visitor.visit_seq(ArrayAccess { de, count: size });
520 de.exit_nested();
521 res
522 })
523 }
524
525 #[inline(always)]
526 fn decode_any_value<const RAW_STRINGS: bool>(&mut self) -> DecodeResult<Value<'_, 'de>> {
527 let (size, type_num) = self.size_and_type()?;
528
529 Ok(match type_num {
530 TYPE_POINTER => {
531 let new_ptr = self.decode_pointer(size);
532 let prev_ptr = self.current_ptr;
533 self.current_ptr = new_ptr;
534
535 Value::Any { prev_ptr }
536 }
537 TYPE_STRING if RAW_STRINGS => {
538 self.count_payload(size)?;
539 Value::RawString(self.read_string_bytes(size)?)
540 }
541 TYPE_STRING => {
542 self.count_payload(size)?;
543 Value::String(self.decode_string(size)?)
544 }
545 TYPE_DOUBLE => Value::F64(self.decode_double(size)?),
546 TYPE_BYTES => {
547 self.count_payload(size)?;
548 Value::Bytes(self.decode_bytes(size)?)
549 }
550 TYPE_UINT16 => Value::U16(self.decode_uint16(size)?),
551 TYPE_UINT32 => Value::U32(self.decode_uint32(size)?),
552 TYPE_MAP => {
553 self.reserve_container_values(size, TYPE_MAP)?;
554 self.enter_nested()?;
555 self.decode_map(size)
556 }
557 TYPE_INT32 => Value::I32(self.decode_int(size)?),
558 TYPE_UINT64 => Value::U64(self.decode_uint64(size)?),
559 TYPE_UINT128 => Value::U128(self.decode_uint128(size)?),
560 TYPE_ARRAY => {
561 self.reserve_container_values(size, TYPE_ARRAY)?;
562 self.enter_nested()?;
563 self.decode_array(size)
564 }
565 TYPE_BOOL => Value::Bool(self.decode_bool(size)?),
566 TYPE_FLOAT => Value::F32(self.decode_float(size)?),
567 u => return Err(self.invalid_db_error(&format!("unknown data type: {u}"))),
568 })
569 }
570
571 fn decode_array(&mut self, size: usize) -> Value<'_, 'de> {
572 Value::Array(ArrayAccess {
573 de: self,
574 count: size,
575 })
576 }
577
578 fn decode_bool(&mut self, size: usize) -> DecodeResult<bool> {
579 match size {
580 0 | 1 => Ok(size != 0),
581 s => Err(self.invalid_db_error(&format!("bool of size {s}"))),
582 }
583 }
584
585 fn decode_bytes(&mut self, size: usize) -> DecodeResult<&'de [u8]> {
586 let new_offset = self.checked_offset(size, "bytes")?;
587 let u8_slice = self.slice(self.current_ptr, new_offset);
588 self.current_ptr = new_offset;
589
590 Ok(u8_slice)
591 }
592
593 fn decode_float(&mut self, size: usize) -> DecodeResult<f32> {
594 let new_offset = self.checked_offset(size, "float")?;
595 let value: [u8; 4] = self
596 .slice(self.current_ptr, new_offset)
597 .try_into()
598 .map_err(|_| self.invalid_db_error(&format!("float of size {size}")))?;
599 self.current_ptr = new_offset;
600 let float_value = f32::from_be_bytes(value);
601 Ok(float_value)
602 }
603
604 fn decode_double(&mut self, size: usize) -> DecodeResult<f64> {
605 let new_offset = self.checked_offset(size, "double")?;
606 let value: [u8; 8] = self
607 .slice(self.current_ptr, new_offset)
608 .try_into()
609 .map_err(|_| self.invalid_db_error(&format!("double of size {size}")))?;
610 self.current_ptr = new_offset;
611 let float_value = f64::from_be_bytes(value);
612 Ok(float_value)
613 }
614
615 decode_int_like!(decode_uint64, u64, 8, "u64", 0_u64);
616 decode_int_like!(decode_uint128, u128, 16, "u128", 0_u128);
617
618 #[inline(always)]
619 fn read_u32_be(&mut self, size: usize, label: &str) -> DecodeResult<u32> {
620 if size > 4 {
621 return Err(self.invalid_db_error(&format!("{label} of size {size}")));
622 }
623 let new_offset = self
624 .current_ptr
625 .checked_add(size)
626 .filter(|&offset| offset <= self.limit)
627 .ok_or_else(|| self.invalid_db_error(&format!("{label} of size {}", size)))?;
628 let p = self.current_ptr;
629 let value = match size {
630 0 => 0,
631 1 => self.buf[p] as u32,
632 2 => ((self.buf[p] as u32) << 8) | self.buf[p + 1] as u32,
633 3 => {
634 ((self.buf[p] as u32) << 16)
635 | ((self.buf[p + 1] as u32) << 8)
636 | self.buf[p + 2] as u32
637 }
638 _ => {
639 ((self.buf[p] as u32) << 24)
640 | ((self.buf[p + 1] as u32) << 16)
641 | ((self.buf[p + 2] as u32) << 8)
642 | self.buf[p + 3] as u32
643 }
644 };
645 self.current_ptr = new_offset;
646 Ok(value)
647 }
648
649 #[inline(always)]
650 fn decode_uint32(&mut self, size: usize) -> DecodeResult<u32> {
651 self.read_u32_be(size, "u32")
652 }
653
654 #[inline(always)]
655 fn decode_uint16(&mut self, size: usize) -> DecodeResult<u16> {
656 if size > 2 {
657 return Err(self.invalid_db_error(&format!("u16 of size {size}")));
658 }
659 let new_offset = self
660 .current_ptr
661 .checked_add(size)
662 .filter(|&offset| offset <= self.limit)
663 .ok_or_else(|| self.invalid_db_error(&format!("u16 of size {}", size)))?;
664 let p = self.current_ptr;
665 let value = match size {
666 0 => 0,
667 1 => self.buf[p] as u16,
668 _ => ((self.buf[p] as u16) << 8) | self.buf[p + 1] as u16,
669 };
670 self.current_ptr = new_offset;
671 Ok(value)
672 }
673
674 fn decode_int(&mut self, size: usize) -> DecodeResult<i32> {
675 self.read_u32_be(size, "i32").map(|value| value as i32)
676 }
677
678 fn decode_map(&mut self, size: usize) -> Value<'_, 'de> {
679 Value::Map(MapAccessor {
680 de: self,
681 count: size * 2,
682 })
683 }
684
685 #[inline(always)]
686 fn decode_pointer(&mut self, size: usize) -> usize {
687 let pointer_size = ((size >> 3) & 0x3) + 1;
688 let p = self.current_ptr;
689 let limit = self.limit;
690 let new_offset = match p.checked_add(pointer_size) {
691 Some(offset) if offset <= limit => offset,
692 _ => {
693 self.current_ptr = limit;
696 return limit;
697 }
698 };
699 let pointer_bytes = self.slice(p, new_offset);
700 self.current_ptr = new_offset;
701
702 match pointer_size {
703 1 => ((size & 0x7) << 8) | usize::from(pointer_bytes[0]),
704 2 => {
705 (((size & 0x7) << 16)
706 | (usize::from(pointer_bytes[0]) << 8)
707 | usize::from(pointer_bytes[1]))
708 + 2048
709 }
710 3 => {
711 (((size & 0x7) << 24)
712 | (usize::from(pointer_bytes[0]) << 16)
713 | (usize::from(pointer_bytes[1]) << 8)
714 | usize::from(pointer_bytes[2]))
715 + 526_336
716 }
717 _ => {
718 (usize::from(pointer_bytes[0]) << 24)
719 | (usize::from(pointer_bytes[1]) << 16)
720 | (usize::from(pointer_bytes[2]) << 8)
721 | usize::from(pointer_bytes[3])
722 }
723 }
724 }
725
726 #[cfg(feature = "unsafe-str-decode")]
727 #[inline(always)]
728 fn decode_string(&mut self, size: usize) -> DecodeResult<&'de str> {
729 use std::str::from_utf8_unchecked;
730
731 let new_offset = self.checked_offset(size, "string")?;
732 let bytes = self.slice(self.current_ptr, new_offset);
733 self.current_ptr = new_offset;
734 let v = unsafe { from_utf8_unchecked(bytes) };
740 Ok(v)
741 }
742
743 #[cfg(not(feature = "unsafe-str-decode"))]
744 #[inline(always)]
745 fn decode_string(&mut self, size: usize) -> DecodeResult<&'de str> {
746 use std::str::from_utf8;
747 use std::str::from_utf8_unchecked;
748
749 let new_offset = self.checked_offset(size, "string")?;
750 let bytes = self.slice(self.current_ptr, new_offset);
751 self.current_ptr = new_offset;
752 if is_ascii(bytes) {
753 let v = unsafe { from_utf8_unchecked(bytes) };
756 return Ok(v);
757 }
758 match from_utf8(bytes) {
759 Ok(v) => Ok(v),
760 Err(_) => Err(self.invalid_db_error("invalid UTF-8 in string")),
761 }
762 }
763
764 pub(crate) fn peek_type(&mut self) -> DecodeResult<(usize, usize)> {
769 let saved_ptr = self.current_ptr;
770 let result = self.size_and_type_following_pointers()?;
771 self.current_ptr = saved_ptr;
772 Ok(result)
773 }
774
775 pub(crate) fn consume_container_header(&mut self) -> DecodeResult<(usize, usize)> {
777 let (size, type_num) = self.size_and_type_following_pointers()?;
778 if type_num == TYPE_MAP || type_num == TYPE_ARRAY {
779 self.reserve_container_values(size, type_num)?;
780 }
781 Ok((size, type_num))
782 }
783
784 fn size_and_type_following_pointers(&mut self) -> DecodeResult<(usize, usize)> {
786 let (size, type_num) = self.size_and_type()?;
787 if type_num != TYPE_POINTER {
788 return Ok((size, type_num));
789 }
790
791 self.current_ptr = self.decode_pointer(size);
792 let (size, type_num) = self.size_and_type()?;
793 if type_num == TYPE_POINTER {
794 return Err(self.invalid_db_error("pointer points to another pointer"));
795 }
796
797 Ok((size, type_num))
798 }
799
800 #[inline(always)]
801 fn decode_direct<T, F>(
802 &mut self,
803 size: usize,
804 type_num: usize,
805 expected_type: usize,
806 label: &str,
807 decode: F,
808 ) -> DecodeResult<T>
809 where
810 F: FnOnce(&mut Self, usize) -> DecodeResult<T>,
811 {
812 let (size, continuation) = match type_num {
813 TYPE_POINTER => {
814 let new_ptr = self.decode_pointer(size);
815 let saved_ptr = self.current_ptr;
816 self.current_ptr = new_ptr;
817 if let Err(error) = self.enter_nested() {
818 self.current_ptr = saved_ptr;
819 return Err(error);
820 }
821 let header = self.size_and_type().and_then(|(size, type_num)| {
822 if type_num == TYPE_POINTER {
823 Err(self.invalid_db_error("pointer points to another pointer"))
824 } else if type_num != expected_type {
825 Err(self.type_mismatch(label, type_num))
826 } else {
827 Ok(size)
828 }
829 });
830 match header {
831 Ok(size) => (size, Some(saved_ptr)),
832 Err(error) => {
833 self.exit_nested();
834 self.current_ptr = saved_ptr;
835 return Err(error);
836 }
837 }
838 }
839 t if t == expected_type => (size, None),
840 _ => return Err(self.type_mismatch(label, type_num)),
841 };
842 let result = decode(self, size);
845 if let Some(saved_ptr) = continuation {
846 self.exit_nested();
847 self.current_ptr = saved_ptr;
848 }
849 result
850 }
851
852 #[inline(always)]
853 fn read_string_bytes(&mut self, size: usize) -> DecodeResult<&'de [u8]> {
854 let new_offset = self
855 .current_ptr
856 .checked_add(size)
857 .ok_or_else(|| self.invalid_db_error("string length exceeds buffer"))?;
858 if new_offset > self.limit {
859 return Err(self.invalid_db_error("string length exceeds buffer"));
860 }
861 let bytes = self.slice(self.current_ptr, new_offset);
862 self.current_ptr = new_offset;
863 Ok(bytes)
864 }
865
866 #[inline]
869 pub(crate) fn read_str_as_bytes(&mut self) -> DecodeResult<&'de [u8]> {
870 let offset = self.current_ptr;
874 let ctrl = self.eat_byte()?;
875 match usize::from(ctrl >> 5) {
876 TYPE_STRING => {
877 let size = self.size_from_ctrl_byte(ctrl, TYPE_STRING)?;
878 self.count_payload(size)?;
879 return self.read_string_bytes(size);
880 }
881 TYPE_POINTER => {
882 let size = self.size_from_ctrl_byte(ctrl, TYPE_POINTER)?;
883 let target = self.decode_pointer(size);
884 let continuation = self.current_ptr;
885 self.current_ptr = target;
886 let ctrl = self.eat_byte()?;
887 if usize::from(ctrl >> 5) == TYPE_STRING {
888 let size = self.size_from_ctrl_byte(ctrl, TYPE_STRING)?;
889 let result = self
890 .count_payload(size)
891 .and_then(|()| self.read_string_bytes(size));
892 self.current_ptr = continuation;
893 return result;
894 }
895 }
896 _ => {}
897 }
898 self.current_ptr = offset;
899 self.read_str_as_bytes_slow()
900 }
901
902 #[cold]
905 fn read_str_as_bytes_slow(&mut self) -> DecodeResult<&'de [u8]> {
906 let (size, type_num) = self.size_and_type()?;
907 match type_num {
908 TYPE_POINTER => {
909 let new_ptr = self.decode_pointer(size);
910 let saved_ptr = self.current_ptr;
911 self.current_ptr = new_ptr;
912 let (size, type_num) = self.size_and_type()?;
913 let result = if type_num == TYPE_POINTER {
914 Err(self.invalid_db_error("pointer points to another pointer"))
915 } else if type_num == TYPE_STRING {
916 self.count_payload(size)
917 .and_then(|()| self.read_string_bytes(size))
918 } else {
919 Err(self.invalid_db_error(&format!("expected string, got type {type_num}")))
920 };
921 self.current_ptr = saved_ptr;
922 result
923 }
924 TYPE_STRING => {
925 self.count_payload(size)?;
926 self.read_string_bytes(size)
927 }
928 _ => Err(self.invalid_db_error(&format!("expected string, got type {type_num}"))),
929 }
930 }
931
932 #[inline(always)]
937 fn try_read_identifier_bytes(&mut self) -> DecodeResult<Option<&'de [u8]>> {
938 let saved_ptr = self.current_ptr;
939 let (size, type_num) = self.size_and_type()?;
940 match type_num {
941 TYPE_STRING => {
942 self.count_identifier_payload(size)?;
943 self.read_string_bytes(size).map(Some)
944 }
945 TYPE_POINTER => {
946 let new_ptr = self.decode_pointer(size);
947 let after_pointer = self.current_ptr;
948 self.current_ptr = new_ptr;
949 let (inner_size, inner_type) = self.size_and_type()?;
950 let result = if inner_type == TYPE_POINTER {
951 Err(self.invalid_db_error("pointer points to another pointer"))
952 } else if inner_type == TYPE_STRING {
953 let payload_result = self.count_identifier_payload(inner_size);
954 match payload_result {
955 Ok(()) => self.read_string_bytes(inner_size).map(Some),
956 Err(error) => Err(error),
957 }
958 } else {
959 Ok(None)
960 };
961 self.current_ptr = after_pointer;
968 if matches!(result, Ok(None)) {
969 self.current_ptr = saved_ptr;
970 }
971 result
972 }
973 _ => {
974 self.current_ptr = saved_ptr;
975 Ok(None)
976 }
977 }
978 }
979
980 pub(crate) fn skip_value(&mut self) -> DecodeResult<()> {
982 let (size, type_num) = self.size_and_type()?;
983 self.skip_value_inner(size, type_num, 0)
984 }
985
986 pub(crate) fn skip_value_for_verification(
988 &mut self,
989 state: &mut VerificationState,
990 ) -> DecodeResult<()> {
991 let offset = self.current_ptr;
992 if state.validated.contains(&offset) {
993 return Ok(());
994 }
995 if !state.active.insert(offset) {
996 return Err(
997 self.invalid_db_error(&format!("cyclic data pointer references offset {offset}"))
998 );
999 }
1000
1001 let result = (|| {
1002 let (size, type_num) = self.size_and_type()?;
1003 self.skip_value_inner_for_verification(size, type_num, 0, state)?;
1004 self.validate_skip_end()
1005 })();
1006
1007 state.active.remove(&offset);
1008 if result.is_ok() {
1009 state.validated.insert(offset);
1010 }
1011 result
1012 }
1013
1014 #[inline(always)]
1015 pub(crate) fn validate_skip_end(&mut self) -> DecodeResult<()> {
1016 if self.current_ptr > self.limit {
1017 return Err(self.invalid_db_error("skipped value extends beyond buffer"));
1018 }
1019 Ok(())
1020 }
1021
1022 #[inline(always)]
1023 fn check_skip_depth(&self, skip_depth: u16) -> DecodeResult<u16> {
1024 if skip_depth == MAXIMUM_SKIPPED_DATA_STRUCTURE_DEPTH {
1025 return self.skip_depth_error();
1026 }
1027 Ok(skip_depth + 1)
1028 }
1029
1030 #[cold]
1031 fn skip_depth_error(&self) -> DecodeResult<u16> {
1032 Err(self
1033 .invalid_db_error("exceeded maximum data structure depth; database is likely corrupt"))
1034 }
1035
1036 #[inline(always)]
1037 fn skip_value_inner(
1038 &mut self,
1039 size: usize,
1040 type_num: usize,
1041 skip_depth: u16,
1042 ) -> DecodeResult<()> {
1043 match type_num {
1047 TYPE_POINTER => {
1048 let pointer_size = ((size >> 3) & 0x3) + 1;
1053 self.checked_offset(pointer_size, "pointer")?;
1054 self.decode_pointer(size);
1055 Ok(())
1056 }
1057 TYPE_STRING | TYPE_BYTES => {
1058 let label = if type_num == TYPE_STRING {
1060 "string"
1061 } else {
1062 "bytes"
1063 };
1064 self.skip_bytes(size, label)
1065 }
1066 TYPE_DOUBLE => {
1067 if size != 8 {
1069 return Err(self.invalid_db_error(&format!("double of size {size}")));
1070 }
1071 self.skip_bytes(size, "double")
1072 }
1073 TYPE_FLOAT => {
1074 if size != 4 {
1076 return Err(self.invalid_db_error(&format!("float of size {size}")));
1077 }
1078 self.skip_bytes(size, "float")
1079 }
1080 TYPE_UINT16 | TYPE_UINT32 | TYPE_INT32 | TYPE_UINT64 | TYPE_UINT128 => {
1081 let label = match type_num {
1083 TYPE_UINT16 => "u16",
1084 TYPE_UINT32 => "u32",
1085 TYPE_INT32 => "i32",
1086 TYPE_UINT64 => "u64",
1087 TYPE_UINT128 => "u128",
1088 _ => unreachable!(),
1089 };
1090 let max_size = match type_num {
1091 TYPE_UINT16 => 2,
1092 TYPE_UINT32 | TYPE_INT32 => 4,
1093 TYPE_UINT64 => 8,
1094 TYPE_UINT128 => 16,
1095 _ => unreachable!(),
1096 };
1097 if size > max_size {
1098 return Err(self.invalid_db_error(&format!("{label} of size {size}")));
1099 }
1100 self.skip_bytes(size, label)
1101 }
1102 TYPE_BOOL => {
1103 self.decode_bool(size).map(|_| ())
1105 }
1106 TYPE_MAP => {
1107 self.reserve_container_values(size, TYPE_MAP)?;
1109 let child_depth = self.check_skip_depth(skip_depth)?;
1110 for _ in 0..size {
1111 self.skip_value_with_depth(child_depth)?;
1113 self.skip_value_with_depth(child_depth)?;
1115 }
1116 Ok(())
1117 }
1118 TYPE_ARRAY => {
1119 self.reserve_container_values(size, TYPE_ARRAY)?;
1121 let child_depth = self.check_skip_depth(skip_depth)?;
1122 for _ in 0..size {
1123 self.skip_value_with_depth(child_depth)?;
1124 }
1125 Ok(())
1126 }
1127 u => Err(self.invalid_db_error(&format!("unknown data type: {u}"))),
1128 }
1129 }
1130
1131 #[inline(always)]
1132 fn skip_value_with_depth(&mut self, skip_depth: u16) -> DecodeResult<()> {
1133 let (size, type_num) = self.size_and_type()?;
1134 self.skip_value_inner(size, type_num, skip_depth)
1135 }
1136
1137 fn skip_value_inner_for_verification(
1138 &mut self,
1139 size: usize,
1140 type_num: usize,
1141 skip_depth: u16,
1142 state: &mut VerificationState,
1143 ) -> DecodeResult<()> {
1144 state.charge(1, self.current_ptr)?;
1147 match type_num {
1148 TYPE_STRING => {
1149 let end = self.checked_offset(size, "string")?;
1150 state.charge(size, self.current_ptr)?;
1151 let bytes = self.slice(self.current_ptr, end);
1152 self.current_ptr = end;
1153 std::str::from_utf8(bytes)
1154 .map(|_| ())
1155 .map_err(|_| self.invalid_db_error("invalid UTF-8 in string"))
1156 }
1157 TYPE_POINTER => {
1158 let target = self.decode_pointer(size);
1159 let child_depth = self.check_skip_depth(skip_depth)?;
1160 self.verify_pointer_target(target, child_depth, state)
1161 }
1162 TYPE_MAP => {
1163 let child_depth = self.check_skip_depth(skip_depth)?;
1164 for _ in 0..size {
1165 self.skip_value_with_verification(child_depth, state)?;
1166 self.skip_value_with_verification(child_depth, state)?;
1167 }
1168 self.validate_skip_end()
1169 }
1170 TYPE_ARRAY => {
1171 let child_depth = self.check_skip_depth(skip_depth)?;
1172 for _ in 0..size {
1173 self.skip_value_with_verification(child_depth, state)?;
1174 }
1175 self.validate_skip_end()
1176 }
1177 _ => self.skip_value_inner(size, type_num, skip_depth),
1178 }
1179 }
1180
1181 fn skip_value_with_verification(
1182 &mut self,
1183 skip_depth: u16,
1184 state: &mut VerificationState,
1185 ) -> DecodeResult<()> {
1186 let (size, type_num) = self.size_and_type()?;
1187 self.skip_value_inner_for_verification(size, type_num, skip_depth, state)
1188 }
1189
1190 fn verify_pointer_target(
1191 &mut self,
1192 target: usize,
1193 skip_depth: u16,
1194 state: &mut VerificationState,
1195 ) -> DecodeResult<()> {
1196 if state.validated.contains(&target) {
1197 return Ok(());
1198 }
1199 if !state.active.insert(target) {
1200 return Err(
1201 self.invalid_db_error(&format!("cyclic data pointer references offset {target}"))
1202 );
1203 }
1204
1205 let continuation = self.current_ptr;
1206 self.current_ptr = target;
1207 let result = (|| {
1208 let (size, type_num) = self.size_and_type()?;
1209 self.skip_value_inner_for_verification(size, type_num, skip_depth, state)?;
1210 self.validate_skip_end()
1211 })();
1212 self.current_ptr = continuation;
1213
1214 state.active.remove(&target);
1215 if result.is_ok() {
1216 state.validated.insert(target);
1217 }
1218 result
1219 }
1220}
1221
1222pub(crate) type DecoderError = Box<MaxMindDbError>;
1225
1226pub type DecodeResult<T> = Result<T, DecoderError>;
1227
1228pub fn deserialize_any_with_raw_strings<'de, D, V>(
1256 deserializer: D,
1257 visitor: V,
1258) -> Result<V::Value, D::Error>
1259where
1260 D: Deserializer<'de>,
1261 V: Visitor<'de>,
1262{
1263 deserializer.deserialize_newtype_struct(RAW_STRINGS_NEWTYPE, visitor)
1264}
1265
1266impl<'de: 'a, 'a> de::Deserializer<'de> for &'a mut Decoder<'de> {
1267 type Error = DecoderError;
1268
1269 fn deserialize_any<V>(self, visitor: V) -> DecodeResult<V::Value>
1270 where
1271 V: Visitor<'de>,
1272 {
1273 self.decode_any(visitor)
1274 }
1275
1276 fn deserialize_option<V>(self, visitor: V) -> DecodeResult<V::Value>
1277 where
1278 V: Visitor<'de>,
1279 {
1280 visitor.visit_some(self)
1281 }
1282
1283 deserialize_direct_scalar!(deserialize_bool, TYPE_BOOL, "bool", visit_bool, decode_bool);
1284
1285 deserialize_direct_scalar!(
1286 deserialize_u16,
1287 TYPE_UINT16,
1288 "u16",
1289 visit_u16,
1290 decode_uint16
1291 );
1292
1293 deserialize_direct_scalar!(
1294 deserialize_u32,
1295 TYPE_UINT32,
1296 "u32",
1297 visit_u32,
1298 decode_uint32
1299 );
1300
1301 deserialize_direct_scalar!(
1302 deserialize_u64,
1303 TYPE_UINT64,
1304 "u64",
1305 visit_u64,
1306 decode_uint64
1307 );
1308
1309 deserialize_direct_scalar!(
1310 deserialize_u128,
1311 TYPE_UINT128,
1312 "u128",
1313 visit_u128,
1314 decode_uint128
1315 );
1316
1317 deserialize_direct_scalar!(deserialize_i32, TYPE_INT32, "i32", visit_i32, decode_int);
1318
1319 deserialize_direct_scalar!(
1320 deserialize_f32,
1321 TYPE_FLOAT,
1322 "float",
1323 visit_f32,
1324 decode_float
1325 );
1326
1327 deserialize_direct_scalar!(
1328 deserialize_f64,
1329 TYPE_DOUBLE,
1330 "double",
1331 visit_f64,
1332 decode_double
1333 );
1334
1335 deserialize_direct_payload!(
1336 deserialize_str,
1337 TYPE_STRING,
1338 "string",
1339 visit_borrowed_str,
1340 decode_string
1341 );
1342
1343 fn deserialize_string<V>(self, visitor: V) -> DecodeResult<V::Value>
1344 where
1345 V: Visitor<'de>,
1346 {
1347 self.deserialize_str(visitor)
1348 }
1349
1350 deserialize_direct_payload!(
1351 deserialize_bytes,
1352 TYPE_BYTES,
1353 "bytes",
1354 visit_borrowed_bytes,
1355 decode_bytes
1356 );
1357
1358 fn deserialize_byte_buf<V>(self, visitor: V) -> DecodeResult<V::Value>
1359 where
1360 V: Visitor<'de>,
1361 {
1362 self.deserialize_bytes(visitor)
1363 }
1364
1365 fn deserialize_seq<V>(self, visitor: V) -> DecodeResult<V::Value>
1366 where
1367 V: Visitor<'de>,
1368 {
1369 let (size, type_num) = self.size_and_type()?;
1370 self.decode_direct(size, type_num, TYPE_ARRAY, "array", |de, size| {
1371 de.reserve_container_values(size, TYPE_ARRAY)?;
1372 de.enter_nested()?;
1373 let res = visitor.visit_seq(ArrayAccess { de, count: size });
1374 de.exit_nested();
1375 res
1376 })
1377 }
1378
1379 fn deserialize_tuple<V>(self, len: usize, visitor: V) -> DecodeResult<V::Value>
1380 where
1381 V: Visitor<'de>,
1382 {
1383 self.deserialize_fixed_size_array(len, visitor)
1384 }
1385
1386 fn deserialize_tuple_struct<V>(
1387 self,
1388 _name: &'static str,
1389 len: usize,
1390 visitor: V,
1391 ) -> DecodeResult<V::Value>
1392 where
1393 V: Visitor<'de>,
1394 {
1395 self.deserialize_fixed_size_array(len, visitor)
1396 }
1397
1398 fn deserialize_map<V>(self, visitor: V) -> DecodeResult<V::Value>
1399 where
1400 V: Visitor<'de>,
1401 {
1402 self.deserialize_map_impl(visitor)
1403 }
1404
1405 fn deserialize_struct<V>(
1406 self,
1407 _name: &'static str,
1408 _fields: &'static [&'static str],
1409 visitor: V,
1410 ) -> DecodeResult<V::Value>
1411 where
1412 V: Visitor<'de>,
1413 {
1414 self.deserialize_map_impl(visitor)
1415 }
1416
1417 fn is_human_readable(&self) -> bool {
1418 false
1419 }
1420
1421 fn deserialize_ignored_any<V>(self, visitor: V) -> DecodeResult<V::Value>
1422 where
1423 V: Visitor<'de>,
1424 {
1425 self.skip_value()?;
1426 visitor.visit_unit()
1427 }
1428
1429 fn deserialize_enum<V>(
1430 self,
1431 _name: &'static str,
1432 _variants: &'static [&'static str],
1433 visitor: V,
1434 ) -> DecodeResult<V::Value>
1435 where
1436 V: Visitor<'de>,
1437 {
1438 self.activate_budget();
1439 visitor.visit_enum(EnumAccessor { de: self })
1440 }
1441
1442 fn deserialize_identifier<V>(self, visitor: V) -> DecodeResult<V::Value>
1443 where
1444 V: Visitor<'de>,
1445 {
1446 match self.try_read_identifier_bytes()? {
1447 Some(bytes) => visitor.visit_borrowed_bytes(bytes),
1448 None => self.decode_any(visitor),
1449 }
1450 }
1451
1452 fn deserialize_newtype_struct<V>(self, name: &'static str, visitor: V) -> DecodeResult<V::Value>
1453 where
1454 V: Visitor<'de>,
1455 {
1456 if name == RAW_STRINGS_NEWTYPE {
1457 self.activate_budget();
1458 self.decode_any_impl::<true, V>(visitor)
1459 } else {
1460 self.decode_any(visitor)
1461 }
1462 }
1463
1464 forward_to_deserialize_any! {
1465 i8 i16 i64 i128 u8 char unit unit_struct
1466 }
1467}
1468
1469impl<'de> Decoder<'de> {
1470 fn deserialize_map_impl<V>(&mut self, visitor: V) -> DecodeResult<V::Value>
1471 where
1472 V: Visitor<'de>,
1473 {
1474 let (size, type_num) = self.size_and_type()?;
1475 self.decode_direct(size, type_num, TYPE_MAP, "map", |de, size| {
1476 de.reserve_container_values(size, TYPE_MAP)?;
1477 de.enter_nested()?;
1478 let res = visitor.visit_map(MapAccessor::<false> {
1479 de,
1480 count: size * 2,
1481 });
1482 de.exit_nested();
1483 res
1484 })
1485 }
1486}
1487
1488struct ArrayAccess<'a, 'de: 'a> {
1489 de: &'a mut Decoder<'de>,
1490 count: usize,
1491}
1492
1493impl<'de> SeqAccess<'de> for ArrayAccess<'_, 'de> {
1496 type Error = DecoderError;
1497
1498 #[inline(always)]
1499 fn size_hint(&self) -> Option<usize> {
1500 debug_assert!(self.de.current_ptr <= self.de.limit);
1504 Some(self.count.min(self.de.limit - self.de.current_ptr))
1505 }
1506
1507 fn next_element_seed<T>(&mut self, seed: T) -> DecodeResult<Option<T::Value>>
1508 where
1509 T: DeserializeSeed<'de>,
1510 {
1511 if self.count == 0 {
1513 if self.de.current_ptr > self.de.limit {
1514 return Err(self
1515 .de
1516 .invalid_db_error("skipped value extends beyond buffer"));
1517 }
1518 return Ok(None);
1519 }
1520 self.count -= 1;
1521
1522 seed.deserialize(&mut *self.de).map(Some)
1524 }
1525}
1526
1527struct MapAccessor<'a, 'de: 'a, const BUDGETED: bool = false> {
1528 de: &'a mut Decoder<'de>,
1529 count: usize,
1530}
1531
1532impl<'de, const BUDGETED: bool> MapAccess<'de> for MapAccessor<'_, 'de, BUDGETED> {
1535 type Error = DecoderError;
1536
1537 #[inline(always)]
1538 fn size_hint(&self) -> Option<usize> {
1539 debug_assert!(self.de.current_ptr <= self.de.limit);
1542 Some((self.count / 2).min((self.de.limit - self.de.current_ptr) / 2))
1543 }
1544
1545 fn next_key_seed<K>(&mut self, seed: K) -> DecodeResult<Option<K::Value>>
1546 where
1547 K: DeserializeSeed<'de>,
1548 {
1549 if self.count == 0 {
1551 if self.de.current_ptr > self.de.limit {
1552 return Err(self
1553 .de
1554 .invalid_db_error("skipped value extends beyond buffer"));
1555 }
1556 return Ok(None);
1557 }
1558 self.count -= 1;
1559
1560 if BUDGETED {
1561 let payload_remaining_before = self.de.payload_remaining;
1562 let (payload_precharged, cached) = self.de.count_payload_at_current()?;
1563 let payload_remaining_after = self.de.payload_remaining;
1564
1565 self.de.payload_remaining = payload_remaining_before;
1570 let result = seed
1571 .deserialize(KeyDeserializer {
1572 decoder: self.de,
1573 cached,
1574 })
1575 .map(Some);
1576 if payload_precharged {
1577 self.de.payload_remaining = self.de.payload_remaining.min(payload_remaining_after);
1578 }
1579 return result;
1580 }
1581
1582 seed.deserialize(&mut *self.de).map(Some)
1584 }
1585
1586 fn next_value_seed<V>(&mut self, seed: V) -> DecodeResult<V::Value>
1587 where
1588 V: DeserializeSeed<'de>,
1589 {
1590 if self.count == 0 {
1592 return Err(self.de.decode_error("no more entries"));
1593 }
1594 self.count -= 1;
1595
1596 seed.deserialize(&mut *self.de)
1598 }
1599}
1600
1601struct EnumAccessor<'a, 'de: 'a> {
1602 de: &'a mut Decoder<'de>,
1603}
1604
1605impl<'de> de::EnumAccess<'de> for EnumAccessor<'_, 'de> {
1606 type Error = DecoderError;
1607 type Variant = Self;
1608
1609 fn variant_seed<V>(self, seed: V) -> DecodeResult<(V::Value, Self::Variant)>
1610 where
1611 V: DeserializeSeed<'de>,
1612 {
1613 let variant = seed.deserialize(&mut *self.de)?;
1615 Ok((variant, self))
1616 }
1617}
1618
1619impl<'de> de::VariantAccess<'de> for EnumAccessor<'_, 'de> {
1620 type Error = DecoderError;
1621
1622 fn unit_variant(self) -> DecodeResult<()> {
1623 Ok(())
1624 }
1625
1626 fn newtype_variant_seed<T>(self, seed: T) -> DecodeResult<T::Value>
1627 where
1628 T: DeserializeSeed<'de>,
1629 {
1630 self.de.reserve_values(1)?;
1631 self.de.enter_nested()?;
1632 let result = seed.deserialize(&mut *self.de);
1633 self.de.exit_nested();
1634 result
1635 }
1636
1637 fn tuple_variant<V>(self, len: usize, visitor: V) -> DecodeResult<V::Value>
1638 where
1639 V: Visitor<'de>,
1640 {
1641 self.de.deserialize_fixed_size_array(len, visitor)
1642 }
1643
1644 fn struct_variant<V>(
1645 self,
1646 _fields: &'static [&'static str],
1647 visitor: V,
1648 ) -> DecodeResult<V::Value>
1649 where
1650 V: Visitor<'de>,
1651 {
1652 de::Deserializer::deserialize_map(&mut *self.de, visitor)
1653 }
1654}
1655
1656#[cfg(test)]
1657mod tests {
1658 use std::fmt;
1659
1660 use serde::de::{DeserializeSeed, Deserializer, MapAccess, SeqAccess, Visitor};
1661 use serde::Deserialize;
1662
1663 use crate::{deserialize_any_with_raw_strings, MaxMindDbError, Reader};
1664
1665 use super::{Decoder, VerificationState};
1666
1667 #[derive(Debug, PartialEq)]
1668 enum RawValue<'de> {
1669 String(&'de [u8]),
1670 Bytes(&'de [u8]),
1671 Bool(bool),
1672 I32(i32),
1673 U16(u16),
1674 U32(u32),
1675 U64(u64),
1676 U128(u128),
1677 F32(f32),
1678 F64(f64),
1679 Array(Vec<RawValue<'de>>),
1680 Map(Vec<(Vec<u8>, RawValue<'de>)>),
1681 }
1682
1683 impl<'de> Deserialize<'de> for RawValue<'de> {
1684 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
1685 where
1686 D: Deserializer<'de>,
1687 {
1688 RawValueSeed.deserialize(deserializer)
1689 }
1690 }
1691
1692 struct RawValueSeed;
1693
1694 impl<'de> DeserializeSeed<'de> for RawValueSeed {
1695 type Value = RawValue<'de>;
1696
1697 fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
1698 where
1699 D: Deserializer<'de>,
1700 {
1701 deserialize_any_with_raw_strings(deserializer, RawValueVisitor)
1702 }
1703 }
1704
1705 struct RawValueVisitor;
1706
1707 impl<'de> Visitor<'de> for RawValueVisitor {
1708 type Value = RawValue<'de>;
1709
1710 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
1711 formatter.write_str("an MMDB value")
1712 }
1713
1714 fn visit_newtype_struct<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
1715 where
1716 D: Deserializer<'de>,
1717 {
1718 deserializer.deserialize_bytes(RawStringVisitor)
1719 }
1720
1721 fn visit_borrowed_bytes<E>(self, bytes: &'de [u8]) -> Result<Self::Value, E> {
1722 Ok(RawValue::Bytes(bytes))
1723 }
1724
1725 fn visit_bool<E>(self, value: bool) -> Result<Self::Value, E> {
1726 Ok(RawValue::Bool(value))
1727 }
1728
1729 fn visit_i32<E>(self, value: i32) -> Result<Self::Value, E> {
1730 Ok(RawValue::I32(value))
1731 }
1732
1733 fn visit_u16<E>(self, value: u16) -> Result<Self::Value, E> {
1734 Ok(RawValue::U16(value))
1735 }
1736
1737 fn visit_u32<E>(self, value: u32) -> Result<Self::Value, E> {
1738 Ok(RawValue::U32(value))
1739 }
1740
1741 fn visit_u64<E>(self, value: u64) -> Result<Self::Value, E> {
1742 Ok(RawValue::U64(value))
1743 }
1744
1745 fn visit_u128<E>(self, value: u128) -> Result<Self::Value, E> {
1746 Ok(RawValue::U128(value))
1747 }
1748
1749 fn visit_f32<E>(self, value: f32) -> Result<Self::Value, E> {
1750 Ok(RawValue::F32(value))
1751 }
1752
1753 fn visit_f64<E>(self, value: f64) -> Result<Self::Value, E> {
1754 Ok(RawValue::F64(value))
1755 }
1756
1757 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
1758 where
1759 A: MapAccess<'de>,
1760 {
1761 let mut entries = Vec::with_capacity(map.size_hint().unwrap_or(0));
1762 while let Some(key) = map.next_key_seed(RawIdentifierSeed)? {
1763 let value = map.next_value_seed(RawValueSeed)?;
1764 entries.push((key, value));
1765 }
1766 Ok(RawValue::Map(entries))
1767 }
1768
1769 fn visit_seq<A>(self, mut sequence: A) -> Result<Self::Value, A::Error>
1770 where
1771 A: SeqAccess<'de>,
1772 {
1773 let mut values = Vec::with_capacity(sequence.size_hint().unwrap_or(0));
1774 while let Some(value) = sequence.next_element_seed(RawValueSeed)? {
1775 values.push(value);
1776 }
1777 Ok(RawValue::Array(values))
1778 }
1779 }
1780
1781 struct RawStringVisitor;
1782
1783 impl<'de> Visitor<'de> for RawStringVisitor {
1784 type Value = RawValue<'de>;
1785
1786 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
1787 formatter.write_str("borrowed MMDB string bytes")
1788 }
1789
1790 fn visit_borrowed_bytes<E>(self, bytes: &'de [u8]) -> Result<Self::Value, E> {
1791 Ok(RawValue::String(bytes))
1792 }
1793 }
1794
1795 #[derive(Clone, Copy)]
1796 struct RawIdentifierSeed;
1797
1798 impl<'de> DeserializeSeed<'de> for RawIdentifierSeed {
1799 type Value = Vec<u8>;
1800
1801 fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
1802 where
1803 D: Deserializer<'de>,
1804 {
1805 deserializer.deserialize_identifier(RawIdentifierVisitor)
1806 }
1807 }
1808
1809 struct RawIdentifierVisitor;
1810
1811 impl<'de> Visitor<'de> for RawIdentifierVisitor {
1812 type Value = Vec<u8>;
1813
1814 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
1815 formatter.write_str("borrowed MMDB map-key bytes")
1816 }
1817
1818 fn visit_borrowed_bytes<E>(self, bytes: &'de [u8]) -> Result<Self::Value, E> {
1819 Ok(bytes.to_vec())
1820 }
1821 }
1822
1823 struct AnyIdentifierSeed;
1824
1825 impl<'de> DeserializeSeed<'de> for AnyIdentifierSeed {
1826 type Value = usize;
1827
1828 fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
1829 where
1830 D: Deserializer<'de>,
1831 {
1832 deserializer.deserialize_any(AnyIdentifierVisitor)
1833 }
1834 }
1835
1836 struct AnyIdentifierVisitor;
1837
1838 impl<'de> Visitor<'de> for AnyIdentifierVisitor {
1839 type Value = usize;
1840
1841 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
1842 formatter.write_str("an MMDB map key")
1843 }
1844
1845 fn visit_borrowed_str<E>(self, value: &'de str) -> Result<Self::Value, E> {
1846 Ok(value.len())
1847 }
1848
1849 fn visit_borrowed_bytes<E>(self, value: &'de [u8]) -> Result<Self::Value, E> {
1850 Ok(value.len())
1851 }
1852 }
1853
1854 #[derive(Debug, PartialEq)]
1855 struct AnyKeyMap(usize);
1856
1857 struct AnyKeyMapVisitor;
1858
1859 impl<'de> Visitor<'de> for AnyKeyMapVisitor {
1860 type Value = AnyKeyMap;
1861
1862 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
1863 formatter.write_str("an MMDB map")
1864 }
1865
1866 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
1867 where
1868 A: MapAccess<'de>,
1869 {
1870 let mut count = 0;
1871 while let Some(key_size) = map.next_key_seed(AnyIdentifierSeed)? {
1872 assert_eq!(key_size, 4096);
1873 map.next_value::<serde::de::IgnoredAny>()?;
1874 count += 1;
1875 }
1876 Ok(AnyKeyMap(count))
1877 }
1878 }
1879
1880 #[allow(dead_code)]
1881 #[derive(Debug)]
1882 struct OwnedBytes(Vec<u8>);
1883
1884 impl<'de> Deserialize<'de> for OwnedBytes {
1885 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
1886 where
1887 D: Deserializer<'de>,
1888 {
1889 deserializer.deserialize_byte_buf(OwnedBytesVisitor)
1890 }
1891 }
1892
1893 struct OwnedBytesVisitor;
1894
1895 impl<'de> Visitor<'de> for OwnedBytesVisitor {
1896 type Value = OwnedBytes;
1897
1898 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
1899 formatter.write_str("an MMDB byte value")
1900 }
1901
1902 fn visit_borrowed_bytes<E>(self, value: &'de [u8]) -> Result<Self::Value, E> {
1903 Ok(OwnedBytes(value.to_vec()))
1904 }
1905
1906 fn visit_byte_buf<E>(self, value: Vec<u8>) -> Result<Self::Value, E> {
1907 Ok(OwnedBytes(value))
1908 }
1909 }
1910
1911 #[test]
1912 fn raw_string_mode_distinguishes_strings_from_bytes() {
1913 let mut string_decoder = Decoder::new(&[0x42, 0xff, 0xfe], 0);
1914 let string = RawValueSeed.deserialize(&mut string_decoder).unwrap();
1915 assert_eq!(string, RawValue::String(&[0xff, 0xfe]));
1916
1917 let mut bytes_decoder = Decoder::new(&[0x82, 0xff, 0xfe], 0);
1918 let bytes = RawValueSeed.deserialize(&mut bytes_decoder).unwrap();
1919 assert_eq!(bytes, RawValue::Bytes(&[0xff, 0xfe]));
1920 }
1921
1922 #[test]
1923 fn raw_string_mode_recurses_through_maps() {
1924 let encoded = [
1925 0x02, 0x00, 0x44, b't', b'e', b'x', b't', 0x41, 0xff, 0x44, b'b', b'l', b'o', b'b', 0x81, 0xff, ];
1931 let mut decoder = Decoder::new(&encoded, 0);
1932
1933 let value = RawValueSeed.deserialize(&mut decoder).unwrap();
1934
1935 assert_eq!(
1936 value,
1937 RawValue::Map(vec![
1938 (b"text".to_vec(), RawValue::String(&[0xff])),
1939 (b"blob".to_vec(), RawValue::Bytes(&[0xff])),
1940 ])
1941 );
1942 }
1943
1944 #[test]
1945 fn raw_string_mode_recurses_through_arrays_and_pointers() {
1946 let encoded_array = [
1947 0x02, 0x04, 0x41, 0xff, 0x81, 0xff, ];
1951 let mut array_decoder = Decoder::new(&encoded_array, 0);
1952 let array = RawValueSeed.deserialize(&mut array_decoder).unwrap();
1953 assert_eq!(
1954 array,
1955 RawValue::Array(vec![RawValue::String(&[0xff]), RawValue::Bytes(&[0xff]),])
1956 );
1957
1958 let encoded_pointer = [
1959 0x20, 0x02, 0x41, 0xff, ];
1962 let mut pointer_decoder = Decoder::new(&encoded_pointer, 0);
1963 let pointer = RawValueSeed.deserialize(&mut pointer_decoder).unwrap();
1964 assert_eq!(pointer, RawValue::String(&[0xff]));
1965 }
1966
1967 #[test]
1968 fn raw_string_mode_restores_pointer_continuation_in_maps() {
1969 let encoded = [
1970 0x02, 0x00, 0x41, b'a', 0x20, 0x0a, 0x41, b'b', 0x41, b'y', 0x41, b'x', ];
1977 let mut decoder = Decoder::new(&encoded, 0);
1978
1979 let value = RawValueSeed.deserialize(&mut decoder).unwrap();
1980
1981 assert_eq!(
1982 value,
1983 RawValue::Map(vec![
1984 (b"a".to_vec(), RawValue::String(b"x")),
1985 (b"b".to_vec(), RawValue::String(b"y")),
1986 ])
1987 );
1988 }
1989
1990 #[test]
1991 fn raw_string_mode_decodes_all_scalar_types() {
1992 let mut encoded = vec![0x08, 0x00]; encoded.extend_from_slice(&[0x41, b'd', 0x68]);
1995 encoded.extend_from_slice(&1.5_f64.to_be_bytes());
1996 encoded.extend_from_slice(&[0x41, b's', 0xa2, 0x01, 0x02]);
1997 encoded.extend_from_slice(&[0x41, b'i', 0xc4, 0x01, 0x02, 0x03, 0x04]);
1998 encoded.extend_from_slice(&[0x41, b'n', 0x04, 0x01]);
1999 encoded.extend_from_slice(&(-2_i32).to_be_bytes());
2000 encoded.extend_from_slice(&[0x41, b'l', 0x08, 0x02]);
2001 encoded.extend_from_slice(&0x0102_0304_0506_0708_u64.to_be_bytes());
2002 encoded.extend_from_slice(&[0x41, b'x', 0x10, 0x03]);
2003 encoded.extend_from_slice(&0x0102_0304_0506_0708_1112_1314_1516_1718_u128.to_be_bytes());
2004 encoded.extend_from_slice(&[0x41, b'b', 0x01, 0x07]);
2005 encoded.extend_from_slice(&[0x41, b'f', 0x04, 0x08]);
2006 encoded.extend_from_slice(&2.5_f32.to_be_bytes());
2007
2008 let mut decoder = Decoder::new(&encoded, 0);
2009 let value = RawValueSeed.deserialize(&mut decoder).unwrap();
2010
2011 assert_eq!(
2012 value,
2013 RawValue::Map(vec![
2014 (b"d".to_vec(), RawValue::F64(1.5)),
2015 (b"s".to_vec(), RawValue::U16(0x0102)),
2016 (b"i".to_vec(), RawValue::U32(0x0102_0304)),
2017 (b"n".to_vec(), RawValue::I32(-2)),
2018 (b"l".to_vec(), RawValue::U64(0x0102_0304_0506_0708)),
2019 (
2020 b"x".to_vec(),
2021 RawValue::U128(0x0102_0304_0506_0708_1112_1314_1516_1718)
2022 ),
2023 (b"b".to_vec(), RawValue::Bool(true)),
2024 (b"f".to_vec(), RawValue::F32(2.5)),
2025 ])
2026 );
2027 }
2028
2029 #[test]
2030 fn raw_string_mode_rejects_excessive_pointer_depth_and_unknown_types() {
2031 std::thread::Builder::new()
2032 .stack_size(8 * 1024 * 1024)
2033 .spawn(|| {
2034 let mut cyclic_decoder = Decoder::new(&[0x20, 0x00], 0);
2035 let depth_err = RawValueSeed.deserialize(&mut cyclic_decoder).unwrap_err();
2036 assert!(depth_err
2037 .to_string()
2038 .contains("exceeded maximum data structure depth"));
2039 })
2040 .unwrap()
2041 .join()
2042 .unwrap();
2043
2044 let mut unknown_decoder = Decoder::new(&[0x00, 0x06], 0);
2045 let type_err = RawValueSeed.deserialize(&mut unknown_decoder).unwrap_err();
2046 assert!(type_err.to_string().contains("unknown data type: 13"));
2047 }
2048
2049 #[test]
2050 fn malformed_extended_types_return_errors_instead_of_overflowing() {
2051 for extended_type in 249..=u8::MAX {
2052 let encoded = [0x00, extended_type];
2053 let mut decoder = Decoder::new(&encoded, 0);
2054 let error = RawValueSeed.deserialize(&mut decoder).unwrap_err();
2055
2056 assert!(matches!(*error, MaxMindDbError::InvalidDatabase { .. }));
2057 assert!(error.to_string().contains(&format!(
2058 "unknown data type: {}",
2059 u16::from(extended_type) + 7
2060 )));
2061
2062 let mut typed_decoder = Decoder::new(&encoded, 0);
2063 let typed_error =
2064 <u32 as serde::Deserialize>::deserialize(&mut typed_decoder).unwrap_err();
2065 assert!(matches!(
2066 *typed_error,
2067 MaxMindDbError::InvalidDatabase { .. }
2068 ));
2069 }
2070 }
2071
2072 #[test]
2073 fn nested_values_without_raw_opt_in_use_normal_string_decoding() {
2074 struct NestedNormalVisitor;
2075
2076 impl<'de> Visitor<'de> for NestedNormalVisitor {
2077 type Value = &'de str;
2078
2079 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
2080 formatter.write_str("an MMDB map containing a string")
2081 }
2082
2083 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
2084 where
2085 A: MapAccess<'de>,
2086 {
2087 let Some(_key) = map.next_key_seed(RawIdentifierSeed)? else {
2088 return Err(serde::de::Error::custom("expected one map entry"));
2089 };
2090 map.next_value::<&'de str>()
2091 }
2092 }
2093
2094 let encoded = [
2095 0x01, 0x00, 0x41, b'k', 0x45, b'h', b'e', b'l', b'l', b'o', ];
2099 let mut decoder = Decoder::new(&encoded, 0);
2100 let value = deserialize_any_with_raw_strings(&mut decoder, NestedNormalVisitor).unwrap();
2101
2102 assert_eq!(value, "hello");
2103 }
2104
2105 fn raw_map_value<'value, 'de>(
2106 value: &'value RawValue<'de>,
2107 key: &[u8],
2108 ) -> &'value RawValue<'de> {
2109 let RawValue::Map(entries) = value else {
2110 panic!("expected map, got {value:?}");
2111 };
2112 entries
2113 .iter()
2114 .find_map(|(entry_key, value)| (entry_key == key).then_some(value))
2115 .unwrap_or_else(|| panic!("missing map key {:?}", String::from_utf8_lossy(key)))
2116 }
2117
2118 #[test]
2119 fn raw_string_mode_decodes_reader_lookup_results() {
2120 let reader = Reader::open_readfile("test-data/test-data/GeoIP2-City-Test.mmdb").unwrap();
2121 let lookup = reader.lookup("89.160.20.128".parse().unwrap()).unwrap();
2122 let value = lookup.decode::<RawValue<'_>>().unwrap().unwrap();
2123
2124 let city = raw_map_value(&value, b"city");
2125 let city_names = raw_map_value(city, b"names");
2126 assert_eq!(
2127 raw_map_value(city_names, b"en"),
2128 &RawValue::String("Linköping".as_bytes())
2129 );
2130
2131 let country = raw_map_value(&value, b"country");
2132 assert_eq!(
2133 raw_map_value(country, b"is_in_european_union"),
2134 &RawValue::Bool(true)
2135 );
2136
2137 let location = raw_map_value(&value, b"location");
2138 assert_eq!(
2139 raw_map_value(location, b"accuracy_radius"),
2140 &RawValue::U16(76)
2141 );
2142 assert_eq!(
2143 raw_map_value(location, b"latitude"),
2144 &RawValue::F64(58.4167)
2145 );
2146
2147 let subdivisions = raw_map_value(&value, b"subdivisions");
2148 let RawValue::Array(subdivisions) = subdivisions else {
2149 panic!("expected subdivisions array, got {subdivisions:?}");
2150 };
2151 assert!(!subdivisions.is_empty());
2152 }
2153
2154 #[test]
2155 fn ordinary_string_decoding_remains_validated() {
2156 struct OrdinaryNewtypeSeed;
2157
2158 impl<'de> DeserializeSeed<'de> for OrdinaryNewtypeSeed {
2159 type Value = &'de str;
2160
2161 fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
2162 where
2163 D: Deserializer<'de>,
2164 {
2165 deserializer.deserialize_newtype_struct("ordinary", StringVisitor)
2166 }
2167 }
2168
2169 struct StringVisitor;
2170
2171 impl<'de> Visitor<'de> for StringVisitor {
2172 type Value = &'de str;
2173
2174 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
2175 formatter.write_str("a borrowed string")
2176 }
2177
2178 fn visit_borrowed_str<E>(self, value: &'de str) -> Result<Self::Value, E> {
2179 Ok(value)
2180 }
2181 }
2182
2183 let mut valid_decoder = Decoder::new(&[0x45, b'h', b'e', b'l', b'l', b'o'], 0);
2184 assert_eq!(String::deserialize(&mut valid_decoder).unwrap(), "hello");
2185
2186 let mut newtype_decoder = Decoder::new(&[0x45, b'h', b'e', b'l', b'l', b'o'], 0);
2187 assert_eq!(
2188 OrdinaryNewtypeSeed
2189 .deserialize(&mut newtype_decoder)
2190 .unwrap(),
2191 "hello"
2192 );
2193
2194 #[cfg(not(feature = "unsafe-str-decode"))]
2195 {
2196 let mut invalid_decoder = Decoder::new(&[0x41, 0xff], 0);
2197 let err = String::deserialize(&mut invalid_decoder).unwrap_err();
2198 assert!(err.to_string().contains("invalid UTF-8"));
2199 }
2200 }
2201
2202 #[test]
2203 fn test_decoder_accepts_tuple_with_matching_length() {
2204 #[allow(dead_code)]
2205 #[derive(Debug, serde::Deserialize)]
2206 struct TupleRecord {
2207 array: (u32, u32, u32),
2208 }
2209
2210 #[allow(dead_code)]
2211 #[derive(Debug, serde::Deserialize)]
2212 struct TupleStructRecord {
2213 array: TupleStruct,
2214 }
2215
2216 #[allow(dead_code)]
2217 #[derive(Debug, serde::Deserialize)]
2218 struct TupleStruct(u32, u32, u32);
2219
2220 let reader =
2221 Reader::open_readfile("test-data/test-data/MaxMind-DB-test-decoder.mmdb").unwrap();
2222 let lookup = reader.lookup("1.1.1.0".parse().unwrap()).unwrap();
2223
2224 let tuple = lookup.decode::<TupleRecord>().unwrap().unwrap();
2225 assert_eq!(tuple.array, (1, 2, 3));
2226
2227 let tuple_struct = lookup.decode::<TupleStructRecord>().unwrap().unwrap();
2228 assert_eq!(tuple_struct.array.0, 1);
2229 assert_eq!(tuple_struct.array.1, 2);
2230 assert_eq!(tuple_struct.array.2, 3);
2231 }
2232
2233 #[test]
2234 fn test_decoder_rejects_tuple_length_mismatch() {
2235 #[allow(dead_code)]
2236 #[derive(Debug, serde::Deserialize)]
2237 struct TupleRecord {
2238 array: (u32, u32),
2239 }
2240
2241 #[allow(dead_code)]
2242 #[derive(Debug, serde::Deserialize)]
2243 struct TupleStructRecord {
2244 array: TupleStruct,
2245 }
2246
2247 #[allow(dead_code)]
2248 #[derive(Debug, serde::Deserialize)]
2249 struct TupleStruct(u32, u32);
2250
2251 let reader =
2252 Reader::open_readfile("test-data/test-data/MaxMind-DB-test-decoder.mmdb").unwrap();
2253 let lookup = reader.lookup("1.1.1.0".parse().unwrap()).unwrap();
2254
2255 let tuple_err = lookup.decode::<TupleRecord>().unwrap_err();
2256 assert!(tuple_err
2257 .to_string()
2258 .contains("expected tuple of length 2, got array of length 3"));
2259
2260 let tuple_struct_err = lookup.decode::<TupleStructRecord>().unwrap_err();
2261 assert!(tuple_struct_err
2262 .to_string()
2263 .contains("expected tuple of length 2, got array of length 3"));
2264 }
2265
2266 #[test]
2267 fn test_skip_value_for_verification_rejects_truncated_pointer_payload() {
2268 let mut decoder = Decoder::new(&[0x28], 0);
2269 let err = decoder
2270 .skip_value_for_verification(&mut VerificationState::new(decoder.limit))
2271 .unwrap_err();
2272
2273 assert!(matches!(*err, MaxMindDbError::InvalidDatabase { .. }));
2274 }
2275
2276 #[test]
2277 fn pointer_widths_preserve_offsets_and_decoder_limits() {
2278 let cases: &[(usize, &[u8], usize)] = &[
2279 (0x00, &[0x00], 0),
2280 (0x07, &[0xff], 2_047),
2281 (0x08, &[0x00, 0x00], 2_048),
2282 (0x0f, &[0xff, 0xff], 526_335),
2283 (0x10, &[0x00, 0x00, 0x00], 526_336),
2284 (0x10, &[0x01, 0x02, 0x03], 592_387),
2285 (0x17, &[0xff, 0xff, 0xff], 134_744_063),
2286 (0x18, &[0x00, 0x00, 0x00, 0x00], 0),
2287 (0x1f, &[0xff, 0xff, 0xff, 0xff], 4_294_967_295),
2288 ];
2289 for &(size, payload, target) in cases {
2290 let mut buf = vec![0];
2291 buf.extend_from_slice(payload);
2292 let continuation = buf.len();
2293 buf.extend([0xa1, 42]);
2294 let mut decoder = Decoder::new(&buf, 1);
2295 assert_eq!(decoder.decode_pointer(size), target);
2296 assert_eq!(decoder.offset(), continuation);
2297 assert_eq!(u16::deserialize(&mut decoder).unwrap(), 42);
2298
2299 for limit in 1..continuation {
2300 let mut decoder = Decoder::new_with_limit(&buf, 1, limit);
2301 assert_eq!(decoder.decode_pointer(size), limit);
2302 assert_eq!(decoder.offset(), limit);
2303 assert!(u16::deserialize(&mut decoder).is_err());
2304 }
2305 }
2306 }
2307
2308 fn compare_key_decoders(buf: &[u8], start: usize, limit: usize, remaining: u32) {
2309 for budgeted in [false, true] {
2310 let mut fast = Decoder::new_with_limit(buf, start, limit);
2311 let mut general = Decoder::new_with_limit(buf, start, limit);
2312 if budgeted {
2313 fast.activate_budget();
2314 general.activate_budget();
2315 }
2316 fast.payload_remaining = remaining;
2317 general.payload_remaining = remaining;
2318 let actual = fast.read_str_as_bytes().map_err(|e| format!("{e:?}"));
2319 let expected = general
2320 .read_str_as_bytes_slow()
2321 .map_err(|e| format!("{e:?}"));
2322 assert_eq!(
2323 actual, expected,
2324 "start={start}, limit={limit}, remaining={remaining}, budgeted={budgeted}"
2325 );
2326 assert_eq!(fast.current_ptr, general.current_ptr);
2327 assert_eq!(fast.state, general.state);
2328 assert_eq!(fast.payload_remaining, general.payload_remaining);
2329 }
2330 }
2331
2332 #[test]
2333 fn key_decoding_preserves_values_errors_cursors_and_budgets() {
2334 let mut buf = [0x61; 80];
2335 buf[..2].copy_from_slice(&[0x20, 10]);
2336 for control in 0..=255 {
2337 buf[10] = control;
2338 for limit in 0..=buf.len() {
2339 for remaining in [0, 1, 2, 28, 29, 4096] {
2340 compare_key_decoders(&buf, 0, limit, remaining);
2341 compare_key_decoders(&buf, 10, limit, remaining);
2342 }
2343 }
2344 }
2345
2346 let mut state = 0x4D59_5DF4_D0F3_3173_u64;
2347 for _ in 0..8192 {
2348 for byte in &mut buf {
2349 state = state
2350 .wrapping_mul(6_364_136_223_846_793_005)
2351 .wrapping_add(1_442_695_040_888_963_407);
2352 *byte = (state >> 32) as u8;
2353 }
2354 for start in [0, 1, 10, 79, 80, usize::MAX] {
2355 compare_key_decoders(&buf, start, buf.len(), 0);
2356 compare_key_decoders(&buf, start, buf.len(), 4096);
2357 }
2358 }
2359
2360 for pointer in 0..2048 {
2361 let mut buf = [0x61; 2112];
2362 buf[2080..2082].copy_from_slice(&[0x20 | ((pointer >> 8) as u8), pointer as u8]);
2363 buf[pointer..pointer + 3].copy_from_slice(&[0x42, 0xFF, 0xFE]);
2366 compare_key_decoders(&buf, 2080, buf.len(), 1);
2367 compare_key_decoders(&buf, 2080, buf.len(), 2);
2368 }
2369 }
2370
2371 #[test]
2372 fn key_decoders_agree_on_pointer_widths_and_extended_strings() {
2373 let pointers: &[(&[u8], usize)] = &[
2374 (&[0x20, 8], 8),
2375 (&[0x28, 0, 0], 2048),
2376 (&[0x30, 0, 0, 0], 526_336),
2377 (&[0x38, 0, 0, 0, 8], 8),
2378 ];
2379 for &(pointer, target) in pointers {
2380 for size in [0, 1, 28, 29, 285] {
2381 let mut buf = pointer.to_vec();
2382 buf.resize(target, 0);
2383 match size {
2384 0..=28 => buf.push(0x40 | size as u8),
2385 29 => buf.extend_from_slice(&[0x5D, 0]),
2386 285 => buf.extend_from_slice(&[0x5E, 0, 0]),
2387 _ => unreachable!(),
2388 }
2389 buf.resize(buf.len() + size, b'k');
2390 for limit in [pointer.len() - 1, target, buf.len() - 1, buf.len()] {
2391 for remaining in [0, 28, 4096] {
2392 compare_key_decoders(&buf, 0, limit, remaining);
2393 }
2394 }
2395 }
2396 }
2397 }
2398
2399 #[test]
2400 fn typed_pointer_errors_restore_continuation_and_depth() {
2401 let targets: &[&[u8]] = &[
2402 &[], &[0x5d], &[0xc0], &[0x20, 0], &[0x41], &[0x41, b'x'], ];
2409 for &target in targets {
2410 let mut buf = vec![0x20, 4, 0xa1, 42];
2411 buf.extend_from_slice(target);
2412 let mut decoder = Decoder::new(&buf, 0);
2413 let result = <&str>::deserialize(&mut decoder);
2414 if target == [0x41, b'x'] {
2415 assert_eq!(result.unwrap(), "x");
2416 } else {
2417 assert!(result.is_err());
2418 }
2419 assert_eq!(decoder.offset(), 2);
2420 assert_eq!(decoder.state & super::DEPTH_MASK, 0);
2421 assert_eq!(u16::deserialize(&mut decoder).unwrap(), 42);
2422 }
2423 }
2424
2425 #[test]
2426 fn typed_pointer_depth_limit_restores_continuation() {
2427 let buf = [0x20, 4, 0xa1, 42, 0x41, b'x'];
2428 let mut decoder = Decoder::new(&buf, 0);
2429 for _ in 0..super::MAXIMUM_DATA_STRUCTURE_DEPTH {
2430 decoder.enter_nested().unwrap();
2431 }
2432
2433 let error = <&str>::deserialize(&mut decoder).unwrap_err();
2434 assert!(matches!(
2435 *error,
2436 MaxMindDbError::InvalidDatabase {
2437 offset: Some(4),
2438 ..
2439 }
2440 ));
2441 assert!(error
2442 .to_string()
2443 .contains("exceeded maximum data structure depth"));
2444 assert_eq!(decoder.offset(), 2);
2445 assert_eq!(
2446 decoder.state & super::DEPTH_MASK,
2447 u32::from(super::MAXIMUM_DATA_STRUCTURE_DEPTH)
2448 );
2449 assert_eq!(u16::deserialize(&mut decoder).unwrap(), 42);
2450 }
2451
2452 #[cfg(not(feature = "unsafe-str-decode"))]
2453 #[test]
2454 fn ascii_check_covers_every_byte_at_word_boundaries() {
2455 for len in 0..=80 {
2456 let mut storage = vec![0x7f; len + 7];
2457 for offset in 0..8 {
2458 let bytes = &mut storage[offset..offset + len];
2459 assert!(super::is_ascii(bytes));
2460 for index in 0..len {
2461 for byte in 0x80..=0xff {
2462 bytes[index] = byte;
2463 assert!(
2464 !super::is_ascii(bytes),
2465 "accepted non-ASCII byte {byte} at index {index}, length {len}, offset {offset}"
2466 );
2467 }
2468 bytes[index] = 0x7f;
2469 }
2470 }
2471 }
2472 }
2473
2474 #[test]
2475 fn ignored_any_skips_pointer_targets_but_verification_follows_them() {
2476 for (pointer_size, control) in [(1, 0x20), (2, 0x28), (3, 0x30), (4, 0x38)] {
2477 let mut encoded = vec![control];
2480 encoded.resize(pointer_size + 1, 0xff);
2481 let continuation = encoded.len();
2482 encoded.extend([0xa1, 42]); let mut decoder = Decoder::new(&encoded, 0);
2484 serde::de::IgnoredAny::deserialize(&mut decoder).unwrap();
2485 assert_eq!(decoder.offset(), continuation);
2486 assert_eq!(u16::deserialize(&mut decoder).unwrap(), 42);
2487 assert_eq!(decoder.offset(), encoded.len());
2488
2489 let mut decoder = Decoder::new(&encoded, 0);
2490 let err = decoder
2491 .skip_value_for_verification(&mut VerificationState::new(decoder.limit))
2492 .unwrap_err();
2493 assert!(matches!(*err, MaxMindDbError::InvalidDatabase { .. }));
2494
2495 for limit in 1..continuation {
2498 for buf in [&encoded[..limit], &encoded[..]] {
2499 let mut decoder = Decoder::new_with_limit(buf, 0, limit);
2500 let err = serde::de::IgnoredAny::deserialize(&mut decoder).unwrap_err();
2501 assert!(matches!(
2502 *err,
2503 MaxMindDbError::InvalidDatabase { message, offset: Some(1) }
2504 if message == format!("pointer of size {pointer_size}")
2505 ));
2506 assert_eq!(decoder.offset(), 1);
2507 }
2508 }
2509 }
2510 }
2511
2512 #[test]
2513 fn test_decoder_caps_impossible_container_size_hint() {
2514 let mut decoder = Decoder::new(&[0x1d, 0x04, 0xff], 0);
2516 let err = Vec::<serde::de::IgnoredAny>::deserialize(&mut decoder).unwrap_err();
2517
2518 assert!(matches!(*err, MaxMindDbError::InvalidDatabase { .. }));
2519 assert!(err.to_string().contains("unexpected end of buffer"));
2520 }
2521
2522 #[test]
2523 fn ignored_any_rejects_excessive_inline_container_work() {
2524 let mut decoder = Decoder::new(&[0x1e, 0x04, 0xfe, 0xe3], 0);
2528 let err = serde::de::IgnoredAny::deserialize(&mut decoder).unwrap_err();
2529
2530 assert!(err
2531 .to_string()
2532 .contains("maximum number of data structure values"));
2533 }
2534
2535 #[test]
2536 fn ignored_any_rejects_excessive_inline_map_work() {
2537 let mut decoder = Decoder::new(&[0xfe, 0x7e, 0xe3], 0);
2540 let err = serde::de::IgnoredAny::deserialize(&mut decoder).unwrap_err();
2541
2542 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
2543 assert!(err
2544 .to_string()
2545 .contains("maximum number of data structure values"));
2546 }
2547
2548 #[test]
2549 fn navigation_rejects_excessive_declared_container_work() {
2550 let mut decoder = Decoder::new(&[0x1e, 0x04, 0xfe, 0xe3], 0);
2551 let err = decoder.consume_container_header().unwrap_err();
2552
2553 assert!(err
2554 .to_string()
2555 .contains("maximum number of data structure values"));
2556 }
2557
2558 #[test]
2559 fn oversized_containers_fail_before_visitor_entry() {
2560 use std::cell::Cell;
2561
2562 struct EntryVisitor<'a>(&'a Cell<bool>);
2563 impl<'de> Visitor<'de> for EntryVisitor<'_> {
2564 type Value = ();
2565
2566 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
2567 formatter.write_str("a container")
2568 }
2569
2570 fn visit_seq<A: SeqAccess<'de>>(self, _: A) -> Result<(), A::Error> {
2571 self.0.set(true);
2572 Ok(())
2573 }
2574
2575 fn visit_map<A: MapAccess<'de>>(self, _: A) -> Result<(), A::Error> {
2576 self.0.set(true);
2577 Ok(())
2578 }
2579 }
2580
2581 let mut array = vec![0x1e, 0x04, 0xfe, 0xe3]; array.extend([0x00, 0x07].repeat(65_536)); let mut map = vec![0xfe, 0x7e, 0xe3]; map.extend([0x40, 0x00, 0x07].repeat(32_768)); for (is_map, bytes) in [(false, array), (true, map)] {
2589 for dynamic in [false, true] {
2590 let entered = Cell::new(false);
2591 let mut decoder = Decoder::new(&bytes, 0);
2592 let visitor = EntryVisitor(&entered);
2593 let result = if dynamic {
2594 decoder.deserialize_any(visitor)
2595 } else if is_map {
2596 decoder.deserialize_map(visitor)
2597 } else {
2598 decoder.deserialize_seq(visitor)
2599 };
2600 assert!(matches!(
2601 result.map_err(|error| *error),
2602 Err(MaxMindDbError::ResourceLimit { .. })
2603 ));
2604 assert!(
2605 !entered.get(),
2606 "visitor entered: map={is_map}, dynamic={dynamic}"
2607 );
2608 }
2609 }
2610 }
2611
2612 #[test]
2613 fn city_subdivisions_reject_declared_size_before_allocating() {
2614 let mut decoder = Decoder::new(&[0x1e, 0x04, 0xfe, 0xe3], 0);
2618 let err =
2619 Vec::<crate::geoip2::city::Subdivision<'_>>::deserialize(&mut decoder).unwrap_err();
2620
2621 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
2622 assert!(err
2623 .to_string()
2624 .contains("maximum number of data structure values"));
2625 }
2626
2627 #[test]
2628 fn concrete_map_rejects_declared_size_before_allocating() {
2629 let mut decoder = Decoder::new(&[0xfe, 0x7e, 0xe3], 0);
2634 let err =
2635 std::collections::HashMap::<String, serde::de::IgnoredAny>::deserialize(&mut decoder)
2636 .unwrap_err();
2637
2638 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
2639 assert!(err
2640 .to_string()
2641 .contains("maximum number of data structure values"));
2642 }
2643
2644 #[test]
2645 fn concrete_sequence_charges_repeated_shared_struct_maps() {
2646 #[derive(Debug, Deserialize)]
2647 struct EmptyRecord {}
2648
2649 let mut buf = vec![0xfd, 171]; for _ in 0..200 {
2656 buf.extend_from_slice(&[0x40, 0xe0]); }
2658 let array_offset = buf.len();
2659 buf.extend_from_slice(&[0x1e, 0x04, 0x00, 0x0f]); for _ in 0..300 {
2661 append_pointer(&mut buf, 0);
2662 }
2663
2664 let mut decoder = Decoder::new(&buf, array_offset);
2665 let err = Vec::<EmptyRecord>::deserialize(&mut decoder).unwrap_err();
2666
2667 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
2668 assert!(err
2669 .to_string()
2670 .contains("maximum number of data structure values"));
2671 }
2672
2673 #[test]
2674 fn concrete_sequence_charges_repeated_string_payloads() {
2675 let (buf, array_offset) = repeated_4k_payload_array(0x5e, 600);
2679
2680 let mut decoder = Decoder::new(&buf, array_offset);
2681 let err = Vec::<&str>::deserialize(&mut decoder).unwrap_err();
2682
2683 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
2684 assert!(err
2685 .to_string()
2686 .contains("maximum size of data structure string and bytes"));
2687 }
2688
2689 #[test]
2690 fn direct_byte_buf_decoding_charges_repeated_payloads() {
2691 let (buf, array_offset) = repeated_4k_payload_array(0x9e, 600);
2692 let mut decoder = Decoder::new(&buf, array_offset);
2693 let err = Vec::<OwnedBytes>::deserialize(&mut decoder).unwrap_err();
2694
2695 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
2696 assert!(err
2697 .to_string()
2698 .contains("maximum size of data structure string and bytes"));
2699 }
2700
2701 #[test]
2702 fn raw_string_decoding_charges_repeated_payloads() {
2703 let (buf, array_offset) = repeated_4k_payload_array(0x5e, 600);
2704 let mut decoder = Decoder::new(&buf, array_offset);
2705 let err = RawValueSeed.deserialize(&mut decoder).unwrap_err();
2706
2707 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
2708 assert!(err
2709 .to_string()
2710 .contains("maximum size of data structure string and bytes"));
2711 }
2712
2713 #[allow(dead_code)]
2714 #[derive(Debug, Deserialize)]
2715 enum StandaloneScalarEnum {
2716 Known,
2717 }
2718
2719 fn large_string(size: usize) -> Vec<u8> {
2720 assert!((65_821..=65_821 + 0xff_ff_ff).contains(&size));
2721 let encoded_size = (size - 65_821) as u32;
2722 let mut buf = vec![0x5f]; buf.extend_from_slice(&encoded_size.to_be_bytes()[1..]);
2724 buf.resize(buf.len() + size, b'a');
2725 buf
2726 }
2727
2728 #[test]
2729 fn partial_struct_skips_unknown_inline_payload_over_budget() {
2730 #[derive(Deserialize)]
2731 struct PartialRecord {
2732 known: bool,
2733 }
2734
2735 let payload_size = super::MAXIMUM_DATA_STRUCTURE_BYTES + 1;
2736 for control in [0x5f, 0x9f] {
2737 let mut payload = large_string(payload_size);
2738 payload[0] = control; let mut buf = vec![0xe2, 0x47]; buf.extend_from_slice(b"unknown");
2742 buf.extend_from_slice(&payload);
2743 buf.extend_from_slice(&[0x45]); buf.extend_from_slice(b"known");
2745 buf.extend_from_slice(&[0x01, 0x07]); let mut decoder = Decoder::new(&buf, 0);
2748 let decoded = PartialRecord::deserialize(&mut decoder).unwrap();
2749 assert!(decoded.known);
2750 }
2751 }
2752
2753 #[test]
2754 fn budgeted_standalone_scalar_entry_points_enforce_payload_limit() {
2755 let size =
2758 super::MAXIMUM_DATA_STRUCTURE_BYTES + super::MAXIMUM_UNCHARGED_IDENTIFIER_BYTES + 1;
2759 let buf = large_string(size);
2760
2761 let mut decoder = Decoder::new(&buf, 0);
2762 let err = serde_json::Value::deserialize(&mut decoder).unwrap_err();
2763 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
2764
2765 let mut decoder = Decoder::new(&buf, 0);
2766 let err = RawValueSeed.deserialize(&mut decoder).unwrap_err();
2767 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
2768
2769 let mut decoder = Decoder::new(&buf, 0);
2770 let err = StandaloneScalarEnum::deserialize(&mut decoder).unwrap_err();
2771 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
2772
2773 let mut decoder = Decoder::new(&buf, 0);
2777 let decoded = <&str>::deserialize(&mut decoder).unwrap();
2778 assert_eq!(decoded.len(), size);
2779 }
2780
2781 #[test]
2782 fn standalone_scalar_retains_unbudgeted_fast_path() {
2783 let size = super::MAXIMUM_DATA_STRUCTURE_BYTES + 1;
2784 let buf = large_string(size);
2785
2786 let mut decoder = Decoder::new(&buf, 0);
2787 let decoded = <&str>::deserialize(&mut decoder).unwrap();
2788 assert_eq!(decoded.len(), size);
2789 }
2790
2791 fn repeated_4k_payload_array(control: u8, element_count: usize) -> (Vec<u8>, usize) {
2792 assert!(matches!(control, 0x5e | 0x9e));
2795 let mut buf = vec![control, 0x0e, 0xe3]; buf.resize(buf.len() + 4096, b'a');
2797 let array_offset = buf.len();
2798
2799 assert!((285..=u16::MAX as usize + 285).contains(&element_count));
2800 buf.extend_from_slice(&[0x1e, 0x04]); buf.extend_from_slice(&((element_count - 285) as u16).to_be_bytes());
2802 for _ in 0..element_count {
2803 append_pointer(&mut buf, 0);
2804 }
2805
2806 (buf, array_offset)
2807 }
2808
2809 fn pointer_key_map(entry_count: usize) -> (Vec<u8>, usize) {
2810 let mut buf = vec![0x5e, 0x0e, 0xe3]; buf.resize(buf.len() + 4096, b'k');
2813 let map_offset = buf.len();
2814
2815 assert!((285..=u16::MAX as usize + 285).contains(&entry_count));
2816 buf.push(0xfe); buf.extend_from_slice(&((entry_count - 285) as u16).to_be_bytes());
2818 for _ in 0..entry_count {
2819 append_pointer(&mut buf, 0);
2820 buf.extend_from_slice(&[0x00, 0x07]); }
2822
2823 (buf, map_offset)
2824 }
2825
2826 fn repeated_inline_key_map_array(element_count: usize) -> (Vec<u8>, usize) {
2827 let mut buf = vec![0xe1, 0x5e, 0x0e, 0xe3]; buf.resize(buf.len() + 4096, b'k');
2831 buf.extend_from_slice(&[0x00, 0x07]); let array_offset = buf.len();
2833
2834 assert!((285..=u16::MAX as usize + 285).contains(&element_count));
2835 buf.extend_from_slice(&[0x1e, 0x04]); buf.extend_from_slice(&((element_count - 285) as u16).to_be_bytes());
2837 for _ in 0..element_count {
2838 append_pointer(&mut buf, 0);
2839 }
2840
2841 (buf, array_offset)
2842 }
2843
2844 #[test]
2845 fn dynamic_map_keys_are_charged_once() {
2846 let (buf, map_offset) = pointer_key_map(512);
2850 let mut decoder = Decoder::new(&buf, map_offset);
2851 let value = RawValueSeed.deserialize(&mut decoder).unwrap();
2852
2853 let RawValue::Map(entries) = value else {
2854 panic!("expected map");
2855 };
2856 assert_eq!(entries.len(), 512);
2857
2858 let (buf, map_offset) = pointer_key_map(513);
2861 let mut decoder = Decoder::new(&buf, map_offset);
2862 let err = RawValueSeed.deserialize(&mut decoder).unwrap_err();
2863 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
2864 assert!(err
2865 .to_string()
2866 .contains("maximum size of data structure string and bytes"));
2867 }
2868
2869 fn compare_cached_key<'de, K>(
2870 buf: &'de [u8],
2871 start: usize,
2872 limit: usize,
2873 remaining: u32,
2874 seed: K,
2875 ) where
2876 K: DeserializeSeed<'de> + Copy,
2877 K::Value: fmt::Debug + PartialEq,
2878 {
2879 let mut cached = Decoder::new_with_limit(buf, start, limit);
2880 let mut original = Decoder::new_with_limit(buf, start, limit);
2881 for decoder in [&mut cached, &mut original] {
2882 decoder.activate_budget();
2883 decoder.payload_remaining = remaining;
2884 }
2885
2886 let actual = super::MapAccessor::<true> {
2887 de: &mut cached,
2888 count: 2,
2889 }
2890 .next_key_seed(seed)
2891 .map_err(|error| format!("{error:?}"));
2892 let expected = (|| {
2895 let (charged, _) = original.count_payload_at_current()?;
2896 let after = original.payload_remaining;
2897 original.payload_remaining = remaining;
2898 let result = seed.deserialize(&mut original).map(Some);
2899 if charged {
2900 original.payload_remaining = original.payload_remaining.min(after);
2901 }
2902 result
2903 })()
2904 .map_err(|error: super::DecoderError| format!("{error:?}"));
2905
2906 assert_eq!(
2907 actual, expected,
2908 "start={start}, limit={limit}, remaining={remaining}"
2909 );
2910 assert_eq!(cached.current_ptr, original.current_ptr);
2911 assert_eq!(cached.state, original.state);
2912 assert_eq!(cached.payload_remaining, original.payload_remaining);
2913 }
2914
2915 #[test]
2916 fn cached_map_keys_preserve_values_errors_cursors_and_budgets() {
2917 let mut buf = [0x61; 80];
2918 buf[..2].copy_from_slice(&[0x20, 10]);
2919 for control in 0..=255 {
2920 buf[10] = control;
2921 for limit in 0..=buf.len() {
2922 for remaining in [0, 1, 28, 32, 4096] {
2923 compare_cached_key(&buf, 0, limit, remaining, RawIdentifierSeed);
2924 compare_cached_key(&buf, 10, limit, remaining, RawIdentifierSeed);
2925 }
2926 }
2927 }
2928 for &(pointer, target) in &[
2929 (&[0x20, 8][..], 8),
2930 (&[0x28, 0, 0][..], 2048),
2931 (&[0x30, 0, 0, 0][..], 526_336),
2932 (&[0x38, 0, 0, 0, 8][..], 8),
2933 ] {
2934 let mut buf = pointer.to_vec();
2935 buf.resize(target, 0);
2936 buf.extend([0x5e, 0, 0]); buf.resize(buf.len() + 285, 0xff); for limit in [pointer.len() - 1, target, buf.len() - 1, buf.len()] {
2939 for remaining in [0, 284, 285, 4096] {
2940 compare_cached_key(&buf, 0, limit, remaining, RawIdentifierSeed);
2941 }
2942 }
2943 }
2944 }
2945
2946 #[test]
2947 fn cached_map_keys_preserve_other_serde_entry_points() {
2948 use std::marker::PhantomData;
2949
2950 let cases: &[&[u8]] = &[
2951 &[0x41, b'x'],
2952 &[0x20, 4, 0xa1, 42, 0x41, b'x'],
2953 &[0x20, 4, 0xa1, 42, 0x41],
2954 &[0x81, 0xff], &[0xa1, 42],
2956 &[0xe0],
2957 ];
2958 for &buf in cases {
2959 for remaining in [0, 1, 4096] {
2960 compare_cached_key(buf, 0, buf.len(), remaining, PhantomData::<String>);
2961 compare_cached_key(
2962 buf,
2963 0,
2964 buf.len(),
2965 remaining,
2966 PhantomData::<serde_json::Value>,
2967 );
2968 compare_cached_key(
2969 buf,
2970 0,
2971 buf.len(),
2972 remaining,
2973 PhantomData::<serde::de::IgnoredAny>,
2974 );
2975 }
2976 }
2977 }
2978
2979 #[test]
2980 #[cfg(not(feature = "unsafe-str-decode"))]
2981 fn cached_map_keys_do_not_bypass_utf8_validation_for_strings() {
2982 use std::marker::PhantomData;
2983
2984 for buf in [&[0x41, 0xff][..], &[0x20, 2, 0x41, 0xff][..]] {
2985 compare_cached_key(buf, 0, buf.len(), 4096, PhantomData::<String>);
2986 compare_cached_key(buf, 0, buf.len(), 4096, PhantomData::<serde_json::Value>);
2987 compare_cached_key(buf, 0, buf.len(), 4096, RawIdentifierSeed);
2988 }
2989 }
2990
2991 #[test]
2992 fn dynamic_any_map_keys_are_charged_once() {
2993 let (buf, map_offset) = pointer_key_map(512);
2997 let mut decoder = Decoder::new(&buf, map_offset);
2998 assert_eq!(
2999 decoder.deserialize_any(AnyKeyMapVisitor).unwrap(),
3000 AnyKeyMap(512)
3001 );
3002
3003 let (buf, map_offset) = pointer_key_map(513);
3005 let mut decoder = Decoder::new(&buf, map_offset);
3006 let err = decoder.deserialize_any(AnyKeyMapVisitor).unwrap_err();
3007 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
3008 assert!(err
3009 .to_string()
3010 .contains("maximum size of data structure string and bytes"));
3011 }
3012
3013 #[test]
3014 fn navigation_charges_repeated_pointer_keys() {
3015 let (buf, map_offset) = pointer_key_map(513);
3016 let mut decoder = Decoder::new(&buf, map_offset);
3017 let (size, type_num) = decoder.consume_container_header().unwrap();
3018 assert_eq!((size, type_num), (513, super::TYPE_MAP));
3019
3020 for _ in 0..512 {
3021 assert_eq!(decoder.read_str_as_bytes().unwrap().len(), 4096);
3022 decoder.skip_value().unwrap();
3023 }
3024 let err = decoder.read_str_as_bytes().unwrap_err();
3025
3026 assert!(err
3027 .to_string()
3028 .contains("maximum size of data structure string and bytes"));
3029 }
3030
3031 #[test]
3032 fn flattened_pointer_keys_cannot_bypass_payload_budget() {
3033 #[derive(Debug, Deserialize)]
3034 struct Flattened {
3035 #[serde(flatten)]
3036 fields: std::collections::HashMap<String, serde::de::IgnoredAny>,
3037 }
3038
3039 let (buf, map_offset) = pointer_key_map(516);
3042 let mut decoder = Decoder::new(&buf, map_offset);
3043 let decoded = Flattened::deserialize(&mut decoder).unwrap();
3044 assert_eq!(decoded.fields.len(), 1);
3045
3046 let (buf, map_offset) = pointer_key_map(517);
3050 let mut decoder = Decoder::new(&buf, map_offset);
3051 let err = Flattened::deserialize(&mut decoder).unwrap_err();
3052
3053 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
3054 assert!(err
3055 .to_string()
3056 .contains("maximum size of data structure string and bytes"));
3057 }
3058
3059 #[test]
3060 fn repeated_maps_with_inline_keys_cannot_bypass_payload_budget() {
3061 #[derive(Debug, Deserialize)]
3062 struct Flattened {
3063 #[serde(flatten)]
3064 fields: std::collections::HashMap<String, serde::de::IgnoredAny>,
3065 }
3066
3067 let (buf, array_offset) = repeated_inline_key_map_array(516);
3069 let mut decoder = Decoder::new(&buf, array_offset);
3070 let decoded = Vec::<Flattened>::deserialize(&mut decoder).unwrap();
3071 assert_eq!(decoded.len(), 516);
3072 assert!(decoded.iter().all(|value| value.fields.len() == 1));
3073
3074 let (buf, array_offset) = repeated_inline_key_map_array(517);
3076 let mut decoder = Decoder::new(&buf, array_offset);
3077 let err = Vec::<Flattened>::deserialize(&mut decoder).unwrap_err();
3078
3079 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
3080 assert!(err
3081 .to_string()
3082 .contains("maximum size of data structure string and bytes"));
3083 }
3084
3085 #[allow(dead_code)]
3086 #[derive(Debug, Deserialize)]
3087 enum RecursiveEnum {
3088 Next(Box<RecursiveEnum>),
3089 End,
3090 }
3091
3092 fn recursive_enum_chain(in_array: bool) -> (Vec<u8>, usize) {
3093 let mut buf = vec![0x44, b'N', b'e', b'x', b't'];
3094 let start = buf.len();
3095 if in_array {
3096 buf.extend_from_slice(&[0x01, 0x04]); }
3098 for _ in 0..=super::MAXIMUM_DATA_STRUCTURE_DEPTH {
3099 append_pointer(&mut buf, 0);
3100 }
3101 buf.extend_from_slice(&[0x43, b'E', b'n', b'd']);
3102 (buf, start)
3103 }
3104
3105 #[test]
3106 fn recursive_newtype_enum_is_depth_bounded() {
3107 let (buf, start) = recursive_enum_chain(false);
3108 let mut decoder = Decoder::new(&buf, start);
3109 let err = RecursiveEnum::deserialize(&mut decoder).unwrap_err();
3110
3111 assert!(err
3112 .to_string()
3113 .contains("exceeded maximum data structure depth"));
3114 }
3115
3116 #[test]
3117 fn recursive_newtype_enum_inside_array_is_depth_bounded() {
3118 let (buf, start) = recursive_enum_chain(true);
3119 let mut decoder = Decoder::new(&buf, start);
3120 let err = Vec::<RecursiveEnum>::deserialize(&mut decoder).unwrap_err();
3121
3122 assert!(err
3123 .to_string()
3124 .contains("exceeded maximum data structure depth"));
3125 }
3126
3127 #[test]
3128 fn verification_bounds_overlapping_string_scans() {
3129 let size = 65_821 + 65_536;
3132 let mut buf = [0x5f, 0x01, 0x00, 0x00].repeat(64);
3133 buf.resize(buf.len() + size, b'a');
3134 let mut state = VerificationState::new(buf.len());
3135 let allowance = state.work_remaining;
3136 for index in 0..8 {
3137 Decoder::new(&buf, index * 4)
3138 .skip_value_for_verification(&mut state)
3139 .unwrap();
3140 }
3141 assert_eq!(state.work_remaining, allowance - 8 * (size + 1));
3142
3143 let mut decoder = Decoder::new(&buf, 8 * 4);
3144 let err = decoder.skip_value_for_verification(&mut state).unwrap_err();
3145 assert!(matches!(
3146 *err,
3147 MaxMindDbError::ResourceLimit {
3148 offset: Some(36),
3149 ..
3150 }
3151 ));
3152 assert_eq!(state.work_remaining, allowance - 8 * (size + 1) - 1);
3155 assert_eq!(decoder.offset(), 36);
3156 assert_eq!(state.validated.len(), 8);
3157 assert!(state.active.is_empty());
3158 }
3159
3160 #[test]
3161 fn verification_bounds_repeated_inline_traversal_across_roots() {
3162 let mut buf = [0x01, 0x04].repeat(64); buf.extend_from_slice(&[0x00, 0x07]); let mut state = VerificationState::new(buf.len());
3165 let err = (0..64)
3168 .try_for_each(|index| {
3169 Decoder::new(&buf, index * 2).skip_value_for_verification(&mut state)
3170 })
3171 .unwrap_err();
3172 assert!(matches!(*err, MaxMindDbError::ResourceLimit { .. }));
3173 assert_eq!(state.work_remaining, 0);
3174 assert!(state.active.is_empty());
3175 }
3176
3177 #[test]
3178 fn verification_allows_large_nonoverlapping_payloads() {
3179 let size = super::MAXIMUM_DATA_STRUCTURE_BYTES + 1;
3180 let mut buf = vec![0x5f];
3181 buf.extend_from_slice(&((size - 65_821) as u32).to_be_bytes()[1..]);
3182 buf.resize(buf.len() + size, b'a');
3183 let mut state = VerificationState::new(buf.len());
3184 Decoder::new(&buf, 0)
3185 .skip_value_for_verification(&mut state)
3186 .unwrap();
3187 }
3188
3189 #[test]
3190 fn verification_work_allowance_does_not_overflow() {
3191 let mut state = VerificationState::new(usize::MAX);
3192 assert_eq!(state.work_remaining, usize::MAX);
3193 state.charge(usize::MAX, 0).unwrap();
3194 assert!(matches!(
3195 state.charge(1, 0).map_err(|error| *error),
3196 Err(MaxMindDbError::ResourceLimit { .. })
3197 ));
3198 assert_eq!(state.work_remaining, 0);
3199 }
3200
3201 #[test]
3202 fn test_verification_rejects_invalid_bool_size() {
3203 let mut decoder = Decoder::new(&[0x02, 0x07], 0);
3205 let err = decoder
3206 .skip_value_for_verification(&mut VerificationState::new(decoder.limit))
3207 .unwrap_err();
3208
3209 assert!(matches!(*err, MaxMindDbError::InvalidDatabase { .. }));
3210 }
3211
3212 #[test]
3213 fn test_verification_rejects_and_does_not_cache_invalid_utf8() {
3214 let buf = [0x41, 0xff];
3215 let mut state = VerificationState::new(buf.len());
3216
3217 for _ in 0..2 {
3218 let mut decoder = Decoder::new(&buf, 0);
3219 let err = decoder.skip_value_for_verification(&mut state).unwrap_err();
3220
3221 assert!(matches!(*err, MaxMindDbError::InvalidDatabase { .. }));
3222 assert!(err.to_string().contains("invalid UTF-8"));
3223 assert!(state.validated.is_empty());
3224 assert!(state.active.is_empty());
3225 }
3226
3227 #[cfg(not(feature = "unsafe-str-decode"))]
3228 {
3229 let mut decoder = Decoder::new(&buf, 0);
3230 let err = String::deserialize(&mut decoder).unwrap_err();
3231 assert!(err.to_string().contains("invalid UTF-8"));
3232 }
3233 }
3234
3235 fn append_pointer(buf: &mut Vec<u8>, target: usize) {
3236 assert!(target < 2048);
3237 buf.push(0x20 | ((target >> 8) as u8));
3238 buf.push(target as u8);
3239 }
3240
3241 #[test]
3242 fn test_verification_caches_shared_pointer_targets() {
3243 let mut buf = vec![0x00, 0x07];
3247 let mut target = 0;
3248 const LEVELS: usize = 20;
3249
3250 for _ in 0..LEVELS {
3251 let array = buf.len();
3252 buf.extend_from_slice(&[0x02, 0x04]);
3253 append_pointer(&mut buf, target);
3254 append_pointer(&mut buf, target);
3255 target = array;
3256 }
3257
3258 let mut decoder = Decoder::new(&buf, target);
3259 let mut state = VerificationState::new(buf.len());
3260 decoder.skip_value_for_verification(&mut state).unwrap();
3261
3262 assert_eq!(state.validated.len(), LEVELS + 1);
3263 assert!(state.active.is_empty());
3264 }
3265
3266 #[test]
3267 fn test_verification_rejects_data_pointer_cycles() {
3268 let mut buf = vec![0x01, 0x04];
3270 append_pointer(&mut buf, 4);
3271 buf.extend_from_slice(&[0x01, 0x04]);
3272 append_pointer(&mut buf, 0);
3273
3274 let mut decoder = Decoder::new(&buf, 0);
3275 let err = decoder
3276 .skip_value_for_verification(&mut VerificationState::new(decoder.limit))
3277 .unwrap_err();
3278
3279 assert!(matches!(*err, MaxMindDbError::InvalidDatabase { .. }));
3280 assert!(err.to_string().contains("cyclic data pointer"));
3281 }
3282}