1use serde::de::{
11 self, value::BorrowedBytesDeserializer, DeserializeSeed, Deserializer, MapAccess, SeqAccess,
12 Visitor,
13};
14use serde::forward_to_deserialize_any;
15use std::collections::HashSet;
16use std::convert::TryInto;
17
18use crate::error::MaxMindDbError;
19
20const TYPE_EXTENDED: usize = 0;
22pub(crate) const TYPE_POINTER: usize = 1;
23const TYPE_STRING: usize = 2;
24const TYPE_DOUBLE: usize = 3;
25const TYPE_BYTES: usize = 4;
26const TYPE_UINT16: usize = 5;
27const TYPE_UINT32: usize = 6;
28pub(crate) const TYPE_MAP: usize = 7;
29const TYPE_INT32: usize = 8;
30const TYPE_UINT64: usize = 9;
31const TYPE_UINT128: usize = 10;
32pub(crate) const TYPE_ARRAY: usize = 11;
33const TYPE_BOOL: usize = 14;
34const TYPE_FLOAT: usize = 15;
35
36const RAW_STRINGS_NEWTYPE: &str = "$maxminddb::raw_strings";
37
38const MAXIMUM_DATA_STRUCTURE_DEPTH: u16 = 512;
41
42const MAXIMUM_SKIPPED_DATA_STRUCTURE_DEPTH: u16 = 128;
46
47#[inline(always)]
48fn to_usize(base: u8, bytes: &[u8]) -> usize {
49 bytes
50 .iter()
51 .fold(base as usize, |acc, &b| (acc << 8) | b as usize)
52}
53
54macro_rules! decode_int_like {
55 ($name:ident, $ty:ty, $max_size:expr, $label:literal, $zero:expr) => {
56 fn $name(&mut self, size: usize) -> DecodeResult<$ty> {
57 match size {
58 s if s <= $max_size => {
59 let new_offset = self
60 .current_ptr
61 .checked_add(size)
62 .filter(|&offset| offset <= self.limit)
63 .ok_or_else(|| {
64 self.invalid_db_error(&format!("{} of size {}", $label, size))
65 })?;
66 let value = self
67 .slice(self.current_ptr, new_offset)
68 .iter()
69 .fold($zero, |acc, &b| (acc << 8) | <$ty>::from(b));
70 self.current_ptr = new_offset;
71 Ok(value)
72 }
73 s => Err(self.invalid_db_error(&format!("{} of size {}", $label, s))),
74 }
75 }
76 };
77}
78
79macro_rules! deserialize_direct_scalar {
80 ($name:ident, $expected_type:expr, $label:literal, $visit:ident, $decode:ident) => {
81 fn $name<V>(self, visitor: V) -> DecodeResult<V::Value>
82 where
83 V: Visitor<'de>,
84 {
85 let (size, type_num) = self.size_and_type()?;
86 self.decode_direct(size, type_num, $expected_type, $label, |de, size| {
87 visitor.$visit(de.$decode(size)?)
88 })
89 }
90 };
91}
92
93enum Value<'a, 'de> {
94 Any { prev_ptr: usize },
95 Bytes(&'de [u8]),
96 String(&'de str),
97 RawString(&'de [u8]),
98 Bool(bool),
99 I32(i32),
100 U16(u16),
101 U32(u32),
102 U64(u64),
103 U128(u128),
104 F64(f64),
105 F32(f32),
106 Map(MapAccessor<'a, 'de>),
107 Array(ArrayAccess<'a, 'de>),
108}
109
110#[derive(Debug)]
115pub(crate) struct Decoder<'de> {
116 buf: &'de [u8],
117 limit: usize,
118 current_ptr: usize,
119 depth: u16,
120}
121
122#[derive(Debug, Default)]
124pub(crate) struct VerificationState {
125 validated: HashSet<usize>,
126 active: HashSet<usize>,
127}
128
129impl<'de> Decoder<'de> {
130 pub(crate) fn new(buf: &'de [u8], start_ptr: usize) -> Decoder<'de> {
131 Decoder::new_with_limit(buf, start_ptr, buf.len())
132 }
133
134 pub(crate) fn new_with_limit(buf: &'de [u8], start_ptr: usize, limit: usize) -> Decoder<'de> {
135 debug_assert!(limit <= buf.len());
136 Decoder {
137 buf,
138 limit,
139 current_ptr: start_ptr,
140 depth: 0,
141 }
142 }
143
144 #[inline]
146 fn enter_nested(&mut self) -> DecodeResult<()> {
147 if self.depth >= MAXIMUM_DATA_STRUCTURE_DEPTH {
148 return Err(self.invalid_db_error(
149 "exceeded maximum data structure depth; database is likely corrupt",
150 ));
151 }
152 self.depth += 1;
153 Ok(())
154 }
155
156 #[inline]
158 fn exit_nested(&mut self) {
159 self.depth = self.depth.saturating_sub(1);
160 }
161
162 #[inline]
164 fn invalid_db_error(&self, msg: &str) -> MaxMindDbError {
165 MaxMindDbError::invalid_database_at(msg, self.current_ptr)
166 }
167
168 #[inline]
170 fn decode_error(&self, msg: &str) -> MaxMindDbError {
171 MaxMindDbError::decoding_at(msg, self.current_ptr)
172 }
173
174 #[inline(always)]
175 fn type_mismatch(&self, label: &str, type_num: usize) -> MaxMindDbError {
176 if type_num > usize::from(u8::MAX) {
177 self.invalid_db_error(&format!("unknown data type: {type_num}"))
178 } else {
179 self.decode_error(&format!("expected {label}, got type {type_num}"))
180 }
181 }
182
183 #[inline]
184 pub(crate) fn offset(&self) -> usize {
185 self.current_ptr
186 }
187
188 #[inline(always)]
189 fn checked_offset(&self, size: usize, label: &str) -> DecodeResult<usize> {
190 let new_offset = self.current_ptr.wrapping_add(size);
191 if new_offset < self.current_ptr || new_offset > self.limit {
192 return Err(self.invalid_db_error(&format!("{label} of size {size}")));
193 }
194 Ok(new_offset)
195 }
196
197 #[inline(always)]
198 fn slice(&self, start: usize, end: usize) -> &'de [u8] {
199 debug_assert!(start <= end);
200 debug_assert!(end <= self.limit);
201 debug_assert!(self.limit <= self.buf.len());
202 unsafe { self.buf.get_unchecked(start..end) }
205 }
206
207 #[inline(always)]
208 fn skip_bytes(&mut self, size: usize, label: &str) -> DecodeResult<()> {
209 debug_assert!(self.current_ptr <= self.limit);
210 if size > self.limit - self.current_ptr {
211 return Err(self.invalid_db_error(&format!("{label} of size {size}")));
212 }
213 self.current_ptr += size;
214 Ok(())
215 }
216
217 #[inline(always)]
218 fn eat_byte(&mut self) -> DecodeResult<u8> {
219 if self.current_ptr >= self.limit {
220 return Err(self.invalid_db_error("unexpected end of buffer"));
221 }
222 debug_assert!(self.limit <= self.buf.len());
223 let b = unsafe { *self.buf.get_unchecked(self.current_ptr) };
226 self.current_ptr += 1;
227 Ok(b)
228 }
229
230 #[inline(always)]
231 fn size_from_ctrl_byte(&mut self, ctrl_byte: u8, type_num: usize) -> DecodeResult<usize> {
232 let size = (ctrl_byte & 0x1f) as usize;
233 if type_num == TYPE_EXTENDED {
235 return Ok(size);
236 }
237
238 match size {
239 s if s < 29 => Ok(s),
240 29 => Ok(29_usize + self.eat_byte()? as usize),
241 30 => {
242 let b0 = self.eat_byte()? as usize;
243 let b1 = self.eat_byte()? as usize;
244 Ok(285_usize + (b0 << 8) + b1)
245 }
246 _ => {
247 let b0 = self.eat_byte()? as usize;
248 let b1 = self.eat_byte()? as usize;
249 let b2 = self.eat_byte()? as usize;
250 Ok(65_821_usize + (b0 << 16) + (b1 << 8) + b2)
251 }
252 }
253 }
254
255 #[inline(always)]
256 fn size_and_type(&mut self) -> DecodeResult<(usize, usize)> {
257 let ctrl_byte = self.eat_byte()?;
258 let mut type_num = usize::from(ctrl_byte >> 5);
259 if type_num == TYPE_EXTENDED {
261 type_num = usize::from(self.eat_byte()?) + TYPE_MAP;
263 }
264 let size = self.size_from_ctrl_byte(ctrl_byte, type_num)?;
265 Ok((size, type_num))
266 }
267
268 fn decode_any<V: Visitor<'de>>(&mut self, visitor: V) -> DecodeResult<V::Value> {
269 self.decode_any_impl::<false, V>(visitor)
270 }
271
272 fn decode_any_impl<const RAW_STRINGS: bool, V: Visitor<'de>>(
273 &mut self,
274 visitor: V,
275 ) -> DecodeResult<V::Value> {
276 match self.decode_any_value::<RAW_STRINGS>()? {
277 Value::Any { prev_ptr } => {
278 self.enter_nested()?;
280 let res = self.decode_any_impl::<RAW_STRINGS, V>(visitor);
281 self.exit_nested();
282 self.current_ptr = prev_ptr;
283 res
284 }
285 Value::Bool(x) => visitor.visit_bool(x),
286 Value::Bytes(x) => visitor.visit_borrowed_bytes(x),
287 Value::String(x) => visitor.visit_borrowed_str(x),
288 Value::RawString(x) => {
289 visitor.visit_newtype_struct(BorrowedBytesDeserializer::<MaxMindDbError>::new(x))
290 }
291 Value::I32(x) => visitor.visit_i32(x),
292 Value::U16(x) => visitor.visit_u16(x),
293 Value::U32(x) => visitor.visit_u32(x),
294 Value::U64(x) => visitor.visit_u64(x),
295 Value::U128(x) => visitor.visit_u128(x),
296 Value::F64(x) => visitor.visit_f64(x),
297 Value::F32(x) => visitor.visit_f32(x),
298 Value::Map(x) => {
300 let res = visitor.visit_map(x);
301 self.exit_nested();
302 res
303 }
304 Value::Array(x) => {
305 let res = visitor.visit_seq(x);
306 self.exit_nested();
307 res
308 }
309 }
310 }
311
312 fn deserialize_fixed_size_array<V>(&mut self, len: usize, visitor: V) -> DecodeResult<V::Value>
313 where
314 V: Visitor<'de>,
315 {
316 let (size, type_num) = self.size_and_type()?;
317 self.decode_direct(size, type_num, TYPE_ARRAY, "array", |de, size| {
318 if size != len {
319 return Err(de.decode_error(&format!(
320 "expected tuple of length {len}, got array of length {size}"
321 )));
322 }
323
324 de.enter_nested()?;
325 let res = visitor.visit_seq(ArrayAccess { de, count: size });
326 de.exit_nested();
327 res
328 })
329 }
330
331 #[inline(always)]
332 fn decode_any_value<const RAW_STRINGS: bool>(&mut self) -> DecodeResult<Value<'_, 'de>> {
333 let (size, type_num) = self.size_and_type()?;
334
335 Ok(match type_num {
336 TYPE_POINTER => {
337 let new_ptr = self.decode_pointer(size);
338 let prev_ptr = self.current_ptr;
339 self.current_ptr = new_ptr;
340
341 Value::Any { prev_ptr }
342 }
343 TYPE_STRING if RAW_STRINGS => Value::RawString(self.read_string_bytes(size)?),
344 TYPE_STRING => Value::String(self.decode_string(size)?),
345 TYPE_DOUBLE => Value::F64(self.decode_double(size)?),
346 TYPE_BYTES => Value::Bytes(self.decode_bytes(size)?),
347 TYPE_UINT16 => Value::U16(self.decode_uint16(size)?),
348 TYPE_UINT32 => Value::U32(self.decode_uint32(size)?),
349 TYPE_MAP => {
350 self.enter_nested()?;
351 self.decode_map(size)
352 }
353 TYPE_INT32 => Value::I32(self.decode_int(size)?),
354 TYPE_UINT64 => Value::U64(self.decode_uint64(size)?),
355 TYPE_UINT128 => Value::U128(self.decode_uint128(size)?),
356 TYPE_ARRAY => {
357 self.enter_nested()?;
358 self.decode_array(size)
359 }
360 TYPE_BOOL => Value::Bool(self.decode_bool(size)?),
361 TYPE_FLOAT => Value::F32(self.decode_float(size)?),
362 u => return Err(self.invalid_db_error(&format!("unknown data type: {u}"))),
363 })
364 }
365
366 fn decode_array(&mut self, size: usize) -> Value<'_, 'de> {
367 Value::Array(ArrayAccess {
368 de: self,
369 count: size,
370 })
371 }
372
373 fn decode_bool(&mut self, size: usize) -> DecodeResult<bool> {
374 match size {
375 0 | 1 => Ok(size != 0),
376 s => Err(self.invalid_db_error(&format!("bool of size {s}"))),
377 }
378 }
379
380 fn decode_bytes(&mut self, size: usize) -> DecodeResult<&'de [u8]> {
381 let new_offset = self.checked_offset(size, "bytes")?;
382 let u8_slice = self.slice(self.current_ptr, new_offset);
383 self.current_ptr = new_offset;
384
385 Ok(u8_slice)
386 }
387
388 fn decode_float(&mut self, size: usize) -> DecodeResult<f32> {
389 let new_offset = self.checked_offset(size, "float")?;
390 let value: [u8; 4] = self
391 .slice(self.current_ptr, new_offset)
392 .try_into()
393 .map_err(|_| self.invalid_db_error(&format!("float of size {size}")))?;
394 self.current_ptr = new_offset;
395 let float_value = f32::from_be_bytes(value);
396 Ok(float_value)
397 }
398
399 fn decode_double(&mut self, size: usize) -> DecodeResult<f64> {
400 let new_offset = self.checked_offset(size, "double")?;
401 let value: [u8; 8] = self
402 .slice(self.current_ptr, new_offset)
403 .try_into()
404 .map_err(|_| self.invalid_db_error(&format!("double of size {size}")))?;
405 self.current_ptr = new_offset;
406 let float_value = f64::from_be_bytes(value);
407 Ok(float_value)
408 }
409
410 decode_int_like!(decode_uint64, u64, 8, "u64", 0_u64);
411 decode_int_like!(decode_uint128, u128, 16, "u128", 0_u128);
412
413 #[inline(always)]
414 fn read_u32_be(&mut self, size: usize, label: &str) -> DecodeResult<u32> {
415 if size > 4 {
416 return Err(self.invalid_db_error(&format!("{label} of size {size}")));
417 }
418 let new_offset = self
419 .current_ptr
420 .checked_add(size)
421 .filter(|&offset| offset <= self.limit)
422 .ok_or_else(|| self.invalid_db_error(&format!("{label} of size {}", size)))?;
423 let p = self.current_ptr;
424 let value = match size {
425 0 => 0,
426 1 => self.buf[p] as u32,
427 2 => ((self.buf[p] as u32) << 8) | self.buf[p + 1] as u32,
428 3 => {
429 ((self.buf[p] as u32) << 16)
430 | ((self.buf[p + 1] as u32) << 8)
431 | self.buf[p + 2] as u32
432 }
433 _ => {
434 ((self.buf[p] as u32) << 24)
435 | ((self.buf[p + 1] as u32) << 16)
436 | ((self.buf[p + 2] as u32) << 8)
437 | self.buf[p + 3] as u32
438 }
439 };
440 self.current_ptr = new_offset;
441 Ok(value)
442 }
443
444 #[inline(always)]
445 fn decode_uint32(&mut self, size: usize) -> DecodeResult<u32> {
446 self.read_u32_be(size, "u32")
447 }
448
449 #[inline(always)]
450 fn decode_uint16(&mut self, size: usize) -> DecodeResult<u16> {
451 if size > 2 {
452 return Err(self.invalid_db_error(&format!("u16 of size {size}")));
453 }
454 let new_offset = self
455 .current_ptr
456 .checked_add(size)
457 .filter(|&offset| offset <= self.limit)
458 .ok_or_else(|| self.invalid_db_error(&format!("u16 of size {}", size)))?;
459 let p = self.current_ptr;
460 let value = match size {
461 0 => 0,
462 1 => self.buf[p] as u16,
463 _ => ((self.buf[p] as u16) << 8) | self.buf[p + 1] as u16,
464 };
465 self.current_ptr = new_offset;
466 Ok(value)
467 }
468
469 fn decode_int(&mut self, size: usize) -> DecodeResult<i32> {
470 self.read_u32_be(size, "i32").map(|value| value as i32)
471 }
472
473 fn decode_map(&mut self, size: usize) -> Value<'_, 'de> {
474 Value::Map(MapAccessor {
475 de: self,
476 count: size * 2,
477 })
478 }
479
480 #[inline(always)]
481 fn decode_pointer(&mut self, size: usize) -> usize {
482 let pointer_value_offset = [0, 0, 2048, 526_336, 0];
483 let pointer_size = ((size >> 3) & 0x3) + 1;
484 let p = self.current_ptr;
485 let limit = self.limit;
486 let new_offset = match p.checked_add(pointer_size) {
487 Some(offset) if offset <= limit => offset,
488 _ => {
489 self.current_ptr = limit;
492 return limit;
493 }
494 };
495 let pointer_bytes = self.slice(p, new_offset);
496 self.current_ptr = new_offset;
497
498 let base = if pointer_size == 4 {
499 0
500 } else {
501 (size & 0x7) as u8
502 };
503 let unpacked = to_usize(base, pointer_bytes);
504
505 unpacked + pointer_value_offset[pointer_size]
506 }
507
508 #[cfg(feature = "unsafe-str-decode")]
509 fn decode_string(&mut self, size: usize) -> DecodeResult<&'de str> {
510 use std::str::from_utf8_unchecked;
511
512 let new_offset = self.checked_offset(size, "string")?;
513 let bytes = self.slice(self.current_ptr, new_offset);
514 self.current_ptr = new_offset;
515 let v = unsafe { from_utf8_unchecked(bytes) };
521 Ok(v)
522 }
523
524 #[cfg(not(feature = "unsafe-str-decode"))]
525 fn decode_string(&mut self, size: usize) -> DecodeResult<&'de str> {
526 #[cfg(feature = "simdutf8")]
527 use simdutf8::basic::from_utf8;
528 #[cfg(not(feature = "simdutf8"))]
529 use std::str::from_utf8;
530 use std::str::from_utf8_unchecked;
531
532 let new_offset = self.checked_offset(size, "string")?;
533 let bytes = self.slice(self.current_ptr, new_offset);
534 self.current_ptr = new_offset;
535 if bytes.is_ascii() {
536 let v = unsafe { from_utf8_unchecked(bytes) };
539 return Ok(v);
540 }
541 match from_utf8(bytes) {
542 Ok(v) => Ok(v),
543 Err(_) => Err(self.invalid_db_error("invalid UTF-8 in string")),
544 }
545 }
546
547 pub(crate) fn peek_type(&mut self) -> DecodeResult<(usize, usize)> {
552 let saved_ptr = self.current_ptr;
553 let result = self.size_and_type_following_pointers()?;
554 self.current_ptr = saved_ptr;
555 Ok(result)
556 }
557
558 pub(crate) fn consume_container_header(&mut self) -> DecodeResult<(usize, usize)> {
560 self.size_and_type_following_pointers()
561 }
562
563 fn size_and_type_following_pointers(&mut self) -> DecodeResult<(usize, usize)> {
565 let (size, type_num) = self.size_and_type()?;
566 if type_num != TYPE_POINTER {
567 return Ok((size, type_num));
568 }
569
570 self.current_ptr = self.decode_pointer(size);
571 let (size, type_num) = self.size_and_type()?;
572 if type_num == TYPE_POINTER {
573 return Err(self.invalid_db_error("pointer points to another pointer"));
574 }
575
576 Ok((size, type_num))
577 }
578
579 #[inline(always)]
580 fn decode_direct<T, F>(
581 &mut self,
582 size: usize,
583 type_num: usize,
584 expected_type: usize,
585 label: &str,
586 decode: F,
587 ) -> DecodeResult<T>
588 where
589 F: FnOnce(&mut Self, usize) -> DecodeResult<T>,
590 {
591 match type_num {
592 TYPE_POINTER => {
593 let new_ptr = self.decode_pointer(size);
594 let saved_ptr = self.current_ptr;
595 self.current_ptr = new_ptr;
596 self.enter_nested()?;
597 let result = (|| {
598 let (size, type_num) = self.size_and_type()?;
599 if type_num == TYPE_POINTER {
600 return Err(self.invalid_db_error("pointer points to another pointer"));
601 }
602 if type_num != expected_type {
603 return Err(self.type_mismatch(label, type_num));
604 }
605 decode(self, size)
606 })();
607 self.exit_nested();
608 self.current_ptr = saved_ptr;
609 result
610 }
611 t if t == expected_type => decode(self, size),
612 _ => Err(self.type_mismatch(label, type_num)),
613 }
614 }
615
616 #[inline(always)]
617 fn read_string_bytes(&mut self, size: usize) -> DecodeResult<&'de [u8]> {
618 let new_offset = self
619 .current_ptr
620 .checked_add(size)
621 .ok_or_else(|| self.invalid_db_error("string length exceeds buffer"))?;
622 if new_offset > self.limit {
623 return Err(self.invalid_db_error("string length exceeds buffer"));
624 }
625 let bytes = self.slice(self.current_ptr, new_offset);
626 self.current_ptr = new_offset;
627 Ok(bytes)
628 }
629
630 pub(crate) fn read_str_as_bytes(&mut self) -> DecodeResult<&'de [u8]> {
633 let (size, type_num) = self.size_and_type()?;
634 match type_num {
635 TYPE_POINTER => {
636 let new_ptr = self.decode_pointer(size);
637 let saved_ptr = self.current_ptr;
638 self.current_ptr = new_ptr;
639 let (size, type_num) = self.size_and_type()?;
640 let result = if type_num == TYPE_POINTER {
641 Err(self.invalid_db_error("pointer points to another pointer"))
642 } else if type_num == TYPE_STRING {
643 self.read_string_bytes(size)
644 } else {
645 Err(self.invalid_db_error(&format!("expected string, got type {type_num}")))
646 };
647 self.current_ptr = saved_ptr;
648 result
649 }
650 TYPE_STRING => self.read_string_bytes(size),
651 _ => Err(self.invalid_db_error(&format!("expected string, got type {type_num}"))),
652 }
653 }
654
655 fn try_read_identifier_bytes(&mut self) -> DecodeResult<Option<&'de [u8]>> {
660 let saved_ptr = self.current_ptr;
661 let (size, type_num) = self.size_and_type()?;
662 match type_num {
663 TYPE_STRING => self.read_string_bytes(size).map(Some),
664 TYPE_POINTER => {
665 let new_ptr = self.decode_pointer(size);
666 let after_pointer = self.current_ptr;
667 self.current_ptr = new_ptr;
668 let (inner_size, inner_type) = self.size_and_type()?;
669 let result = if inner_type == TYPE_POINTER {
670 Err(self.invalid_db_error("pointer points to another pointer"))
671 } else if inner_type == TYPE_STRING {
672 self.read_string_bytes(inner_size).map(Some)
673 } else {
674 Ok(None)
675 };
676 self.current_ptr = after_pointer;
683 if matches!(result, Ok(None)) {
684 self.current_ptr = saved_ptr;
685 }
686 result
687 }
688 _ => {
689 self.current_ptr = saved_ptr;
690 Ok(None)
691 }
692 }
693 }
694
695 pub(crate) fn skip_value(&mut self) -> DecodeResult<()> {
697 let (size, type_num) = self.size_and_type()?;
698 self.skip_value_inner(size, type_num, 0)
699 }
700
701 pub(crate) fn skip_value_for_verification(
703 &mut self,
704 state: &mut VerificationState,
705 ) -> DecodeResult<()> {
706 let offset = self.current_ptr;
707 if state.validated.contains(&offset) {
708 return Ok(());
709 }
710 if !state.active.insert(offset) {
711 return Err(
712 self.invalid_db_error(&format!("cyclic data pointer references offset {offset}"))
713 );
714 }
715
716 let result = (|| {
717 let (size, type_num) = self.size_and_type()?;
718 self.skip_value_inner_for_verification(size, type_num, 0, state)?;
719 self.validate_skip_end()
720 })();
721
722 state.active.remove(&offset);
723 if result.is_ok() {
724 state.validated.insert(offset);
725 }
726 result
727 }
728
729 #[inline(always)]
730 pub(crate) fn validate_skip_end(&mut self) -> DecodeResult<()> {
731 if self.current_ptr > self.limit {
732 return Err(self.invalid_db_error("skipped value extends beyond buffer"));
733 }
734 Ok(())
735 }
736
737 #[inline(always)]
738 fn check_skip_depth(&self, skip_depth: u16) -> DecodeResult<u16> {
739 if skip_depth == MAXIMUM_SKIPPED_DATA_STRUCTURE_DEPTH {
740 return self.skip_depth_error();
741 }
742 Ok(skip_depth + 1)
743 }
744
745 #[cold]
746 fn skip_depth_error(&self) -> DecodeResult<u16> {
747 Err(self
748 .invalid_db_error("exceeded maximum data structure depth; database is likely corrupt"))
749 }
750
751 #[inline(always)]
752 fn skip_value_inner(
753 &mut self,
754 size: usize,
755 type_num: usize,
756 skip_depth: u16,
757 ) -> DecodeResult<()> {
758 match type_num {
762 TYPE_POINTER => {
763 let new_ptr = self.decode_pointer(size);
764 let saved_ptr = self.current_ptr;
765 self.current_ptr = new_ptr;
766 let result = match self.check_skip_depth(skip_depth) {
767 Ok(child_depth) => self.skip_value_with_depth(child_depth),
768 Err(err) => Err(err),
769 };
770 self.current_ptr = saved_ptr;
771 result
772 }
773 TYPE_STRING | TYPE_BYTES => {
774 let label = if type_num == TYPE_STRING {
776 "string"
777 } else {
778 "bytes"
779 };
780 self.skip_bytes(size, label)
781 }
782 TYPE_DOUBLE => {
783 if size != 8 {
785 return Err(self.invalid_db_error(&format!("double of size {size}")));
786 }
787 self.skip_bytes(size, "double")
788 }
789 TYPE_FLOAT => {
790 if size != 4 {
792 return Err(self.invalid_db_error(&format!("float of size {size}")));
793 }
794 self.skip_bytes(size, "float")
795 }
796 TYPE_UINT16 | TYPE_UINT32 | TYPE_INT32 | TYPE_UINT64 | TYPE_UINT128 => {
797 let label = match type_num {
799 TYPE_UINT16 => "u16",
800 TYPE_UINT32 => "u32",
801 TYPE_INT32 => "i32",
802 TYPE_UINT64 => "u64",
803 TYPE_UINT128 => "u128",
804 _ => unreachable!(),
805 };
806 let max_size = match type_num {
807 TYPE_UINT16 => 2,
808 TYPE_UINT32 | TYPE_INT32 => 4,
809 TYPE_UINT64 => 8,
810 TYPE_UINT128 => 16,
811 _ => unreachable!(),
812 };
813 if size > max_size {
814 return Err(self.invalid_db_error(&format!("{label} of size {size}")));
815 }
816 self.skip_bytes(size, label)
817 }
818 TYPE_BOOL => {
819 self.decode_bool(size).map(|_| ())
821 }
822 TYPE_MAP => {
823 let child_depth = self.check_skip_depth(skip_depth)?;
825 for _ in 0..size {
826 self.skip_value_with_depth(child_depth)?;
828 self.skip_value_with_depth(child_depth)?;
830 }
831 Ok(())
832 }
833 TYPE_ARRAY => {
834 let child_depth = self.check_skip_depth(skip_depth)?;
836 for _ in 0..size {
837 self.skip_value_with_depth(child_depth)?;
838 }
839 Ok(())
840 }
841 u => Err(self.invalid_db_error(&format!("unknown data type: {u}"))),
842 }
843 }
844
845 #[inline(always)]
846 fn skip_value_with_depth(&mut self, skip_depth: u16) -> DecodeResult<()> {
847 let (size, type_num) = self.size_and_type()?;
848 self.skip_value_inner(size, type_num, skip_depth)
849 }
850
851 fn skip_value_inner_for_verification(
852 &mut self,
853 size: usize,
854 type_num: usize,
855 skip_depth: u16,
856 state: &mut VerificationState,
857 ) -> DecodeResult<()> {
858 match type_num {
859 TYPE_STRING => {
860 let end = self.checked_offset(size, "string")?;
861 let bytes = self.slice(self.current_ptr, end);
862 self.current_ptr = end;
863 std::str::from_utf8(bytes)
864 .map(|_| ())
865 .map_err(|_| self.invalid_db_error("invalid UTF-8 in string"))
866 }
867 TYPE_POINTER => {
868 let target = self.decode_pointer(size);
869 let child_depth = self.check_skip_depth(skip_depth)?;
870 self.verify_pointer_target(target, child_depth, state)
871 }
872 TYPE_MAP => {
873 let child_depth = self.check_skip_depth(skip_depth)?;
874 for _ in 0..size {
875 self.skip_value_with_verification(child_depth, state)?;
876 self.skip_value_with_verification(child_depth, state)?;
877 }
878 self.validate_skip_end()
879 }
880 TYPE_ARRAY => {
881 let child_depth = self.check_skip_depth(skip_depth)?;
882 for _ in 0..size {
883 self.skip_value_with_verification(child_depth, state)?;
884 }
885 self.validate_skip_end()
886 }
887 _ => self.skip_value_inner(size, type_num, skip_depth),
888 }
889 }
890
891 fn skip_value_with_verification(
892 &mut self,
893 skip_depth: u16,
894 state: &mut VerificationState,
895 ) -> DecodeResult<()> {
896 let (size, type_num) = self.size_and_type()?;
897 self.skip_value_inner_for_verification(size, type_num, skip_depth, state)
898 }
899
900 fn verify_pointer_target(
901 &mut self,
902 target: usize,
903 skip_depth: u16,
904 state: &mut VerificationState,
905 ) -> DecodeResult<()> {
906 if state.validated.contains(&target) {
907 return Ok(());
908 }
909 if !state.active.insert(target) {
910 return Err(
911 self.invalid_db_error(&format!("cyclic data pointer references offset {target}"))
912 );
913 }
914
915 let continuation = self.current_ptr;
916 self.current_ptr = target;
917 let result = (|| {
918 let (size, type_num) = self.size_and_type()?;
919 self.skip_value_inner_for_verification(size, type_num, skip_depth, state)?;
920 self.validate_skip_end()
921 })();
922 self.current_ptr = continuation;
923
924 state.active.remove(&target);
925 if result.is_ok() {
926 state.validated.insert(target);
927 }
928 result
929 }
930}
931
932pub type DecodeResult<T> = Result<T, MaxMindDbError>;
933
934pub fn deserialize_any_with_raw_strings<'de, D, V>(
961 deserializer: D,
962 visitor: V,
963) -> Result<V::Value, D::Error>
964where
965 D: Deserializer<'de>,
966 V: Visitor<'de>,
967{
968 deserializer.deserialize_newtype_struct(RAW_STRINGS_NEWTYPE, visitor)
969}
970
971impl<'de: 'a, 'a> de::Deserializer<'de> for &'a mut Decoder<'de> {
972 type Error = MaxMindDbError;
973
974 fn deserialize_any<V>(self, visitor: V) -> DecodeResult<V::Value>
975 where
976 V: Visitor<'de>,
977 {
978 self.decode_any(visitor)
979 }
980
981 fn deserialize_option<V>(self, visitor: V) -> DecodeResult<V::Value>
982 where
983 V: Visitor<'de>,
984 {
985 visitor.visit_some(self)
986 }
987
988 deserialize_direct_scalar!(deserialize_bool, TYPE_BOOL, "bool", visit_bool, decode_bool);
989
990 deserialize_direct_scalar!(
991 deserialize_u16,
992 TYPE_UINT16,
993 "u16",
994 visit_u16,
995 decode_uint16
996 );
997
998 deserialize_direct_scalar!(
999 deserialize_u32,
1000 TYPE_UINT32,
1001 "u32",
1002 visit_u32,
1003 decode_uint32
1004 );
1005
1006 deserialize_direct_scalar!(
1007 deserialize_u64,
1008 TYPE_UINT64,
1009 "u64",
1010 visit_u64,
1011 decode_uint64
1012 );
1013
1014 deserialize_direct_scalar!(
1015 deserialize_u128,
1016 TYPE_UINT128,
1017 "u128",
1018 visit_u128,
1019 decode_uint128
1020 );
1021
1022 deserialize_direct_scalar!(deserialize_i32, TYPE_INT32, "i32", visit_i32, decode_int);
1023
1024 deserialize_direct_scalar!(
1025 deserialize_f32,
1026 TYPE_FLOAT,
1027 "float",
1028 visit_f32,
1029 decode_float
1030 );
1031
1032 deserialize_direct_scalar!(
1033 deserialize_f64,
1034 TYPE_DOUBLE,
1035 "double",
1036 visit_f64,
1037 decode_double
1038 );
1039
1040 deserialize_direct_scalar!(
1041 deserialize_str,
1042 TYPE_STRING,
1043 "string",
1044 visit_borrowed_str,
1045 decode_string
1046 );
1047
1048 fn deserialize_string<V>(self, visitor: V) -> DecodeResult<V::Value>
1049 where
1050 V: Visitor<'de>,
1051 {
1052 self.deserialize_str(visitor)
1053 }
1054
1055 deserialize_direct_scalar!(
1056 deserialize_bytes,
1057 TYPE_BYTES,
1058 "bytes",
1059 visit_borrowed_bytes,
1060 decode_bytes
1061 );
1062
1063 fn deserialize_byte_buf<V>(self, visitor: V) -> DecodeResult<V::Value>
1064 where
1065 V: Visitor<'de>,
1066 {
1067 self.deserialize_bytes(visitor)
1068 }
1069
1070 fn deserialize_seq<V>(self, visitor: V) -> DecodeResult<V::Value>
1071 where
1072 V: Visitor<'de>,
1073 {
1074 let (size, type_num) = self.size_and_type()?;
1075 self.decode_direct(size, type_num, TYPE_ARRAY, "array", |de, size| {
1076 de.enter_nested()?;
1077 let res = visitor.visit_seq(ArrayAccess { de, count: size });
1078 de.exit_nested();
1079 res
1080 })
1081 }
1082
1083 fn deserialize_tuple<V>(self, len: usize, visitor: V) -> DecodeResult<V::Value>
1084 where
1085 V: Visitor<'de>,
1086 {
1087 self.deserialize_fixed_size_array(len, visitor)
1088 }
1089
1090 fn deserialize_tuple_struct<V>(
1091 self,
1092 _name: &'static str,
1093 len: usize,
1094 visitor: V,
1095 ) -> DecodeResult<V::Value>
1096 where
1097 V: Visitor<'de>,
1098 {
1099 self.deserialize_fixed_size_array(len, visitor)
1100 }
1101
1102 fn deserialize_map<V>(self, visitor: V) -> DecodeResult<V::Value>
1103 where
1104 V: Visitor<'de>,
1105 {
1106 let (size, type_num) = self.size_and_type()?;
1107 self.decode_direct(size, type_num, TYPE_MAP, "map", |de, size| {
1108 de.enter_nested()?;
1109 let res = visitor.visit_map(MapAccessor {
1110 de,
1111 count: size * 2,
1112 });
1113 de.exit_nested();
1114 res
1115 })
1116 }
1117
1118 fn deserialize_struct<V>(
1119 self,
1120 _name: &'static str,
1121 _fields: &'static [&'static str],
1122 visitor: V,
1123 ) -> DecodeResult<V::Value>
1124 where
1125 V: Visitor<'de>,
1126 {
1127 self.deserialize_map(visitor)
1128 }
1129
1130 fn is_human_readable(&self) -> bool {
1131 false
1132 }
1133
1134 fn deserialize_ignored_any<V>(self, visitor: V) -> DecodeResult<V::Value>
1135 where
1136 V: Visitor<'de>,
1137 {
1138 self.skip_value()?;
1139 visitor.visit_unit()
1140 }
1141
1142 fn deserialize_enum<V>(
1143 self,
1144 _name: &'static str,
1145 _variants: &'static [&'static str],
1146 visitor: V,
1147 ) -> DecodeResult<V::Value>
1148 where
1149 V: Visitor<'de>,
1150 {
1151 visitor.visit_enum(EnumAccessor { de: self })
1152 }
1153
1154 fn deserialize_identifier<V>(self, visitor: V) -> DecodeResult<V::Value>
1155 where
1156 V: Visitor<'de>,
1157 {
1158 match self.try_read_identifier_bytes()? {
1159 Some(bytes) => visitor.visit_borrowed_bytes(bytes),
1160 None => self.decode_any(visitor),
1161 }
1162 }
1163
1164 fn deserialize_newtype_struct<V>(self, name: &'static str, visitor: V) -> DecodeResult<V::Value>
1165 where
1166 V: Visitor<'de>,
1167 {
1168 if name == RAW_STRINGS_NEWTYPE {
1169 self.decode_any_impl::<true, V>(visitor)
1170 } else {
1171 self.decode_any(visitor)
1172 }
1173 }
1174
1175 forward_to_deserialize_any! {
1176 i8 i16 i64 i128 u8 char unit unit_struct
1177 }
1178}
1179
1180struct ArrayAccess<'a, 'de: 'a> {
1181 de: &'a mut Decoder<'de>,
1182 count: usize,
1183}
1184
1185impl<'de> SeqAccess<'de> for ArrayAccess<'_, 'de> {
1188 type Error = MaxMindDbError;
1189
1190 #[inline(always)]
1191 fn size_hint(&self) -> Option<usize> {
1192 debug_assert!(self.de.current_ptr <= self.de.limit);
1196 Some(self.count.min(self.de.limit - self.de.current_ptr))
1197 }
1198
1199 fn next_element_seed<T>(&mut self, seed: T) -> DecodeResult<Option<T::Value>>
1200 where
1201 T: DeserializeSeed<'de>,
1202 {
1203 if self.count == 0 {
1205 if self.de.current_ptr > self.de.limit {
1206 return Err(self
1207 .de
1208 .invalid_db_error("skipped value extends beyond buffer"));
1209 }
1210 return Ok(None);
1211 }
1212 self.count -= 1;
1213
1214 seed.deserialize(&mut *self.de).map(Some)
1216 }
1217}
1218
1219struct MapAccessor<'a, 'de: 'a> {
1220 de: &'a mut Decoder<'de>,
1221 count: usize,
1222}
1223
1224impl<'de> MapAccess<'de> for MapAccessor<'_, 'de> {
1227 type Error = MaxMindDbError;
1228
1229 #[inline(always)]
1230 fn size_hint(&self) -> Option<usize> {
1231 debug_assert!(self.de.current_ptr <= self.de.limit);
1234 Some((self.count / 2).min((self.de.limit - self.de.current_ptr) / 2))
1235 }
1236
1237 fn next_key_seed<K>(&mut self, seed: K) -> DecodeResult<Option<K::Value>>
1238 where
1239 K: DeserializeSeed<'de>,
1240 {
1241 if self.count == 0 {
1243 if self.de.current_ptr > self.de.limit {
1244 return Err(self
1245 .de
1246 .invalid_db_error("skipped value extends beyond buffer"));
1247 }
1248 return Ok(None);
1249 }
1250 self.count -= 1;
1251
1252 seed.deserialize(&mut *self.de).map(Some)
1254 }
1255
1256 fn next_value_seed<V>(&mut self, seed: V) -> DecodeResult<V::Value>
1257 where
1258 V: DeserializeSeed<'de>,
1259 {
1260 if self.count == 0 {
1262 return Err(self.de.decode_error("no more entries"));
1263 }
1264 self.count -= 1;
1265
1266 seed.deserialize(&mut *self.de)
1268 }
1269}
1270
1271struct EnumAccessor<'a, 'de: 'a> {
1272 de: &'a mut Decoder<'de>,
1273}
1274
1275impl<'de> de::EnumAccess<'de> for EnumAccessor<'_, 'de> {
1276 type Error = MaxMindDbError;
1277 type Variant = Self;
1278
1279 fn variant_seed<V>(self, seed: V) -> DecodeResult<(V::Value, Self::Variant)>
1280 where
1281 V: DeserializeSeed<'de>,
1282 {
1283 let variant = seed.deserialize(&mut *self.de)?;
1285 Ok((variant, self))
1286 }
1287}
1288
1289impl<'de> de::VariantAccess<'de> for EnumAccessor<'_, 'de> {
1290 type Error = MaxMindDbError;
1291
1292 fn unit_variant(self) -> DecodeResult<()> {
1293 Ok(())
1294 }
1295
1296 fn newtype_variant_seed<T>(self, seed: T) -> DecodeResult<T::Value>
1297 where
1298 T: DeserializeSeed<'de>,
1299 {
1300 seed.deserialize(&mut *self.de)
1301 }
1302
1303 fn tuple_variant<V>(self, len: usize, visitor: V) -> DecodeResult<V::Value>
1304 where
1305 V: Visitor<'de>,
1306 {
1307 self.de.deserialize_fixed_size_array(len, visitor)
1308 }
1309
1310 fn struct_variant<V>(
1311 self,
1312 _fields: &'static [&'static str],
1313 visitor: V,
1314 ) -> DecodeResult<V::Value>
1315 where
1316 V: Visitor<'de>,
1317 {
1318 de::Deserializer::deserialize_map(&mut *self.de, visitor)
1319 }
1320}
1321
1322#[cfg(test)]
1323mod tests {
1324 use std::fmt;
1325
1326 use serde::de::{DeserializeSeed, Deserializer, MapAccess, SeqAccess, Visitor};
1327 use serde::Deserialize;
1328
1329 use crate::{deserialize_any_with_raw_strings, MaxMindDbError, Reader};
1330
1331 use super::{Decoder, VerificationState};
1332
1333 #[derive(Debug, PartialEq)]
1334 enum RawValue<'de> {
1335 String(&'de [u8]),
1336 Bytes(&'de [u8]),
1337 Bool(bool),
1338 I32(i32),
1339 U16(u16),
1340 U32(u32),
1341 U64(u64),
1342 U128(u128),
1343 F32(f32),
1344 F64(f64),
1345 Array(Vec<RawValue<'de>>),
1346 Map(Vec<(Vec<u8>, RawValue<'de>)>),
1347 }
1348
1349 impl<'de> Deserialize<'de> for RawValue<'de> {
1350 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
1351 where
1352 D: Deserializer<'de>,
1353 {
1354 RawValueSeed.deserialize(deserializer)
1355 }
1356 }
1357
1358 struct RawValueSeed;
1359
1360 impl<'de> DeserializeSeed<'de> for RawValueSeed {
1361 type Value = RawValue<'de>;
1362
1363 fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
1364 where
1365 D: Deserializer<'de>,
1366 {
1367 deserialize_any_with_raw_strings(deserializer, RawValueVisitor)
1368 }
1369 }
1370
1371 struct RawValueVisitor;
1372
1373 impl<'de> Visitor<'de> for RawValueVisitor {
1374 type Value = RawValue<'de>;
1375
1376 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
1377 formatter.write_str("an MMDB value")
1378 }
1379
1380 fn visit_newtype_struct<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
1381 where
1382 D: Deserializer<'de>,
1383 {
1384 deserializer.deserialize_bytes(RawStringVisitor)
1385 }
1386
1387 fn visit_borrowed_bytes<E>(self, bytes: &'de [u8]) -> Result<Self::Value, E> {
1388 Ok(RawValue::Bytes(bytes))
1389 }
1390
1391 fn visit_bool<E>(self, value: bool) -> Result<Self::Value, E> {
1392 Ok(RawValue::Bool(value))
1393 }
1394
1395 fn visit_i32<E>(self, value: i32) -> Result<Self::Value, E> {
1396 Ok(RawValue::I32(value))
1397 }
1398
1399 fn visit_u16<E>(self, value: u16) -> Result<Self::Value, E> {
1400 Ok(RawValue::U16(value))
1401 }
1402
1403 fn visit_u32<E>(self, value: u32) -> Result<Self::Value, E> {
1404 Ok(RawValue::U32(value))
1405 }
1406
1407 fn visit_u64<E>(self, value: u64) -> Result<Self::Value, E> {
1408 Ok(RawValue::U64(value))
1409 }
1410
1411 fn visit_u128<E>(self, value: u128) -> Result<Self::Value, E> {
1412 Ok(RawValue::U128(value))
1413 }
1414
1415 fn visit_f32<E>(self, value: f32) -> Result<Self::Value, E> {
1416 Ok(RawValue::F32(value))
1417 }
1418
1419 fn visit_f64<E>(self, value: f64) -> Result<Self::Value, E> {
1420 Ok(RawValue::F64(value))
1421 }
1422
1423 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
1424 where
1425 A: MapAccess<'de>,
1426 {
1427 let mut entries = Vec::with_capacity(map.size_hint().unwrap_or(0));
1428 while let Some(key) = map.next_key_seed(RawIdentifierSeed)? {
1429 let value = map.next_value_seed(RawValueSeed)?;
1430 entries.push((key, value));
1431 }
1432 Ok(RawValue::Map(entries))
1433 }
1434
1435 fn visit_seq<A>(self, mut sequence: A) -> Result<Self::Value, A::Error>
1436 where
1437 A: SeqAccess<'de>,
1438 {
1439 let mut values = Vec::with_capacity(sequence.size_hint().unwrap_or(0));
1440 while let Some(value) = sequence.next_element_seed(RawValueSeed)? {
1441 values.push(value);
1442 }
1443 Ok(RawValue::Array(values))
1444 }
1445 }
1446
1447 struct RawStringVisitor;
1448
1449 impl<'de> Visitor<'de> for RawStringVisitor {
1450 type Value = RawValue<'de>;
1451
1452 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
1453 formatter.write_str("borrowed MMDB string bytes")
1454 }
1455
1456 fn visit_borrowed_bytes<E>(self, bytes: &'de [u8]) -> Result<Self::Value, E> {
1457 Ok(RawValue::String(bytes))
1458 }
1459 }
1460
1461 struct RawIdentifierSeed;
1462
1463 impl<'de> DeserializeSeed<'de> for RawIdentifierSeed {
1464 type Value = Vec<u8>;
1465
1466 fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
1467 where
1468 D: Deserializer<'de>,
1469 {
1470 deserializer.deserialize_identifier(RawIdentifierVisitor)
1471 }
1472 }
1473
1474 struct RawIdentifierVisitor;
1475
1476 impl<'de> Visitor<'de> for RawIdentifierVisitor {
1477 type Value = Vec<u8>;
1478
1479 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
1480 formatter.write_str("borrowed MMDB map-key bytes")
1481 }
1482
1483 fn visit_borrowed_bytes<E>(self, bytes: &'de [u8]) -> Result<Self::Value, E> {
1484 Ok(bytes.to_vec())
1485 }
1486 }
1487
1488 #[test]
1489 fn raw_string_mode_distinguishes_strings_from_bytes() {
1490 let mut string_decoder = Decoder::new(&[0x42, 0xff, 0xfe], 0);
1491 let string = RawValueSeed.deserialize(&mut string_decoder).unwrap();
1492 assert_eq!(string, RawValue::String(&[0xff, 0xfe]));
1493
1494 let mut bytes_decoder = Decoder::new(&[0x82, 0xff, 0xfe], 0);
1495 let bytes = RawValueSeed.deserialize(&mut bytes_decoder).unwrap();
1496 assert_eq!(bytes, RawValue::Bytes(&[0xff, 0xfe]));
1497 }
1498
1499 #[test]
1500 fn raw_string_mode_recurses_through_maps() {
1501 let encoded = [
1502 0x02, 0x00, 0x44, b't', b'e', b'x', b't', 0x41, 0xff, 0x44, b'b', b'l', b'o', b'b', 0x81, 0xff, ];
1508 let mut decoder = Decoder::new(&encoded, 0);
1509
1510 let value = RawValueSeed.deserialize(&mut decoder).unwrap();
1511
1512 assert_eq!(
1513 value,
1514 RawValue::Map(vec![
1515 (b"text".to_vec(), RawValue::String(&[0xff])),
1516 (b"blob".to_vec(), RawValue::Bytes(&[0xff])),
1517 ])
1518 );
1519 }
1520
1521 #[test]
1522 fn raw_string_mode_recurses_through_arrays_and_pointers() {
1523 let encoded_array = [
1524 0x02, 0x04, 0x41, 0xff, 0x81, 0xff, ];
1528 let mut array_decoder = Decoder::new(&encoded_array, 0);
1529 let array = RawValueSeed.deserialize(&mut array_decoder).unwrap();
1530 assert_eq!(
1531 array,
1532 RawValue::Array(vec![RawValue::String(&[0xff]), RawValue::Bytes(&[0xff]),])
1533 );
1534
1535 let encoded_pointer = [
1536 0x20, 0x02, 0x41, 0xff, ];
1539 let mut pointer_decoder = Decoder::new(&encoded_pointer, 0);
1540 let pointer = RawValueSeed.deserialize(&mut pointer_decoder).unwrap();
1541 assert_eq!(pointer, RawValue::String(&[0xff]));
1542 }
1543
1544 #[test]
1545 fn raw_string_mode_restores_pointer_continuation_in_maps() {
1546 let encoded = [
1547 0x02, 0x00, 0x41, b'a', 0x20, 0x0a, 0x41, b'b', 0x41, b'y', 0x41, b'x', ];
1554 let mut decoder = Decoder::new(&encoded, 0);
1555
1556 let value = RawValueSeed.deserialize(&mut decoder).unwrap();
1557
1558 assert_eq!(
1559 value,
1560 RawValue::Map(vec![
1561 (b"a".to_vec(), RawValue::String(b"x")),
1562 (b"b".to_vec(), RawValue::String(b"y")),
1563 ])
1564 );
1565 }
1566
1567 #[test]
1568 fn raw_string_mode_decodes_all_scalar_types() {
1569 let mut encoded = vec![0x08, 0x00]; encoded.extend_from_slice(&[0x41, b'd', 0x68]);
1572 encoded.extend_from_slice(&1.5_f64.to_be_bytes());
1573 encoded.extend_from_slice(&[0x41, b's', 0xa2, 0x01, 0x02]);
1574 encoded.extend_from_slice(&[0x41, b'i', 0xc4, 0x01, 0x02, 0x03, 0x04]);
1575 encoded.extend_from_slice(&[0x41, b'n', 0x04, 0x01]);
1576 encoded.extend_from_slice(&(-2_i32).to_be_bytes());
1577 encoded.extend_from_slice(&[0x41, b'l', 0x08, 0x02]);
1578 encoded.extend_from_slice(&0x0102_0304_0506_0708_u64.to_be_bytes());
1579 encoded.extend_from_slice(&[0x41, b'x', 0x10, 0x03]);
1580 encoded.extend_from_slice(&0x0102_0304_0506_0708_1112_1314_1516_1718_u128.to_be_bytes());
1581 encoded.extend_from_slice(&[0x41, b'b', 0x01, 0x07]);
1582 encoded.extend_from_slice(&[0x41, b'f', 0x04, 0x08]);
1583 encoded.extend_from_slice(&2.5_f32.to_be_bytes());
1584
1585 let mut decoder = Decoder::new(&encoded, 0);
1586 let value = RawValueSeed.deserialize(&mut decoder).unwrap();
1587
1588 assert_eq!(
1589 value,
1590 RawValue::Map(vec![
1591 (b"d".to_vec(), RawValue::F64(1.5)),
1592 (b"s".to_vec(), RawValue::U16(0x0102)),
1593 (b"i".to_vec(), RawValue::U32(0x0102_0304)),
1594 (b"n".to_vec(), RawValue::I32(-2)),
1595 (b"l".to_vec(), RawValue::U64(0x0102_0304_0506_0708)),
1596 (
1597 b"x".to_vec(),
1598 RawValue::U128(0x0102_0304_0506_0708_1112_1314_1516_1718)
1599 ),
1600 (b"b".to_vec(), RawValue::Bool(true)),
1601 (b"f".to_vec(), RawValue::F32(2.5)),
1602 ])
1603 );
1604 }
1605
1606 #[test]
1607 fn raw_string_mode_rejects_excessive_pointer_depth_and_unknown_types() {
1608 std::thread::Builder::new()
1609 .stack_size(8 * 1024 * 1024)
1610 .spawn(|| {
1611 let mut cyclic_decoder = Decoder::new(&[0x20, 0x00], 0);
1612 let depth_err = RawValueSeed.deserialize(&mut cyclic_decoder).unwrap_err();
1613 assert!(depth_err
1614 .to_string()
1615 .contains("exceeded maximum data structure depth"));
1616 })
1617 .unwrap()
1618 .join()
1619 .unwrap();
1620
1621 let mut unknown_decoder = Decoder::new(&[0x00, 0x06], 0);
1622 let type_err = RawValueSeed.deserialize(&mut unknown_decoder).unwrap_err();
1623 assert!(type_err.to_string().contains("unknown data type: 13"));
1624 }
1625
1626 #[test]
1627 fn malformed_extended_types_return_errors_instead_of_overflowing() {
1628 for extended_type in 249..=u8::MAX {
1629 let encoded = [0x00, extended_type];
1630 let mut decoder = Decoder::new(&encoded, 0);
1631 let error = RawValueSeed.deserialize(&mut decoder).unwrap_err();
1632
1633 assert!(matches!(error, MaxMindDbError::InvalidDatabase { .. }));
1634 assert!(error.to_string().contains(&format!(
1635 "unknown data type: {}",
1636 u16::from(extended_type) + 7
1637 )));
1638
1639 let mut typed_decoder = Decoder::new(&encoded, 0);
1640 let typed_error =
1641 <u32 as serde::Deserialize>::deserialize(&mut typed_decoder).unwrap_err();
1642 assert!(matches!(
1643 typed_error,
1644 MaxMindDbError::InvalidDatabase { .. }
1645 ));
1646 }
1647 }
1648
1649 #[test]
1650 fn nested_values_without_raw_opt_in_use_normal_string_decoding() {
1651 struct NestedNormalVisitor;
1652
1653 impl<'de> Visitor<'de> for NestedNormalVisitor {
1654 type Value = &'de str;
1655
1656 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
1657 formatter.write_str("an MMDB map containing a string")
1658 }
1659
1660 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
1661 where
1662 A: MapAccess<'de>,
1663 {
1664 let Some(_key) = map.next_key_seed(RawIdentifierSeed)? else {
1665 return Err(serde::de::Error::custom("expected one map entry"));
1666 };
1667 map.next_value::<&'de str>()
1668 }
1669 }
1670
1671 let encoded = [
1672 0x01, 0x00, 0x41, b'k', 0x45, b'h', b'e', b'l', b'l', b'o', ];
1676 let mut decoder = Decoder::new(&encoded, 0);
1677 let value = deserialize_any_with_raw_strings(&mut decoder, NestedNormalVisitor).unwrap();
1678
1679 assert_eq!(value, "hello");
1680 }
1681
1682 fn raw_map_value<'value, 'de>(
1683 value: &'value RawValue<'de>,
1684 key: &[u8],
1685 ) -> &'value RawValue<'de> {
1686 let RawValue::Map(entries) = value else {
1687 panic!("expected map, got {value:?}");
1688 };
1689 entries
1690 .iter()
1691 .find_map(|(entry_key, value)| (entry_key == key).then_some(value))
1692 .unwrap_or_else(|| panic!("missing map key {:?}", String::from_utf8_lossy(key)))
1693 }
1694
1695 #[test]
1696 fn raw_string_mode_decodes_reader_lookup_results() {
1697 let reader = Reader::open_readfile("test-data/test-data/GeoIP2-City-Test.mmdb").unwrap();
1698 let lookup = reader.lookup("89.160.20.128".parse().unwrap()).unwrap();
1699 let value = lookup.decode::<RawValue<'_>>().unwrap().unwrap();
1700
1701 let city = raw_map_value(&value, b"city");
1702 let city_names = raw_map_value(city, b"names");
1703 assert_eq!(
1704 raw_map_value(city_names, b"en"),
1705 &RawValue::String("Linköping".as_bytes())
1706 );
1707
1708 let country = raw_map_value(&value, b"country");
1709 assert_eq!(
1710 raw_map_value(country, b"is_in_european_union"),
1711 &RawValue::Bool(true)
1712 );
1713
1714 let location = raw_map_value(&value, b"location");
1715 assert_eq!(
1716 raw_map_value(location, b"accuracy_radius"),
1717 &RawValue::U16(76)
1718 );
1719 assert_eq!(
1720 raw_map_value(location, b"latitude"),
1721 &RawValue::F64(58.4167)
1722 );
1723
1724 let subdivisions = raw_map_value(&value, b"subdivisions");
1725 let RawValue::Array(subdivisions) = subdivisions else {
1726 panic!("expected subdivisions array, got {subdivisions:?}");
1727 };
1728 assert!(!subdivisions.is_empty());
1729 }
1730
1731 #[test]
1732 fn ordinary_string_decoding_remains_validated() {
1733 struct OrdinaryNewtypeSeed;
1734
1735 impl<'de> DeserializeSeed<'de> for OrdinaryNewtypeSeed {
1736 type Value = &'de str;
1737
1738 fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
1739 where
1740 D: Deserializer<'de>,
1741 {
1742 deserializer.deserialize_newtype_struct("ordinary", StringVisitor)
1743 }
1744 }
1745
1746 struct StringVisitor;
1747
1748 impl<'de> Visitor<'de> for StringVisitor {
1749 type Value = &'de str;
1750
1751 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
1752 formatter.write_str("a borrowed string")
1753 }
1754
1755 fn visit_borrowed_str<E>(self, value: &'de str) -> Result<Self::Value, E> {
1756 Ok(value)
1757 }
1758 }
1759
1760 let mut valid_decoder = Decoder::new(&[0x45, b'h', b'e', b'l', b'l', b'o'], 0);
1761 assert_eq!(String::deserialize(&mut valid_decoder).unwrap(), "hello");
1762
1763 let mut newtype_decoder = Decoder::new(&[0x45, b'h', b'e', b'l', b'l', b'o'], 0);
1764 assert_eq!(
1765 OrdinaryNewtypeSeed
1766 .deserialize(&mut newtype_decoder)
1767 .unwrap(),
1768 "hello"
1769 );
1770
1771 #[cfg(not(feature = "unsafe-str-decode"))]
1772 {
1773 let mut invalid_decoder = Decoder::new(&[0x41, 0xff], 0);
1774 let err = String::deserialize(&mut invalid_decoder).unwrap_err();
1775 assert!(err.to_string().contains("invalid UTF-8"));
1776 }
1777 }
1778
1779 #[test]
1780 fn test_decoder_accepts_tuple_with_matching_length() {
1781 #[allow(dead_code)]
1782 #[derive(Debug, serde::Deserialize)]
1783 struct TupleRecord {
1784 array: (u32, u32, u32),
1785 }
1786
1787 #[allow(dead_code)]
1788 #[derive(Debug, serde::Deserialize)]
1789 struct TupleStructRecord {
1790 array: TupleStruct,
1791 }
1792
1793 #[allow(dead_code)]
1794 #[derive(Debug, serde::Deserialize)]
1795 struct TupleStruct(u32, u32, u32);
1796
1797 let reader =
1798 Reader::open_readfile("test-data/test-data/MaxMind-DB-test-decoder.mmdb").unwrap();
1799 let lookup = reader.lookup("1.1.1.0".parse().unwrap()).unwrap();
1800
1801 let tuple = lookup.decode::<TupleRecord>().unwrap().unwrap();
1802 assert_eq!(tuple.array, (1, 2, 3));
1803
1804 let tuple_struct = lookup.decode::<TupleStructRecord>().unwrap().unwrap();
1805 assert_eq!(tuple_struct.array.0, 1);
1806 assert_eq!(tuple_struct.array.1, 2);
1807 assert_eq!(tuple_struct.array.2, 3);
1808 }
1809
1810 #[test]
1811 fn test_decoder_rejects_tuple_length_mismatch() {
1812 #[allow(dead_code)]
1813 #[derive(Debug, serde::Deserialize)]
1814 struct TupleRecord {
1815 array: (u32, u32),
1816 }
1817
1818 #[allow(dead_code)]
1819 #[derive(Debug, serde::Deserialize)]
1820 struct TupleStructRecord {
1821 array: TupleStruct,
1822 }
1823
1824 #[allow(dead_code)]
1825 #[derive(Debug, serde::Deserialize)]
1826 struct TupleStruct(u32, u32);
1827
1828 let reader =
1829 Reader::open_readfile("test-data/test-data/MaxMind-DB-test-decoder.mmdb").unwrap();
1830 let lookup = reader.lookup("1.1.1.0".parse().unwrap()).unwrap();
1831
1832 let tuple_err = lookup.decode::<TupleRecord>().unwrap_err();
1833 assert!(tuple_err
1834 .to_string()
1835 .contains("expected tuple of length 2, got array of length 3"));
1836
1837 let tuple_struct_err = lookup.decode::<TupleStructRecord>().unwrap_err();
1838 assert!(tuple_struct_err
1839 .to_string()
1840 .contains("expected tuple of length 2, got array of length 3"));
1841 }
1842
1843 #[test]
1844 fn test_skip_value_for_verification_rejects_truncated_pointer_payload() {
1845 let mut decoder = Decoder::new(&[0x28], 0);
1846 let err = decoder
1847 .skip_value_for_verification(&mut VerificationState::default())
1848 .unwrap_err();
1849
1850 assert!(matches!(err, MaxMindDbError::InvalidDatabase { .. }));
1851 }
1852
1853 #[test]
1854 fn test_decoder_caps_impossible_container_size_hint() {
1855 let mut decoder = Decoder::new(&[0x1d, 0x04, 0xff], 0);
1857 let err = Vec::<serde::de::IgnoredAny>::deserialize(&mut decoder).unwrap_err();
1858
1859 assert!(matches!(err, MaxMindDbError::InvalidDatabase { .. }));
1860 assert!(err.to_string().contains("unexpected end of buffer"));
1861 }
1862
1863 #[test]
1864 fn test_verification_rejects_invalid_bool_size() {
1865 let mut decoder = Decoder::new(&[0x02, 0x07], 0);
1867 let err = decoder
1868 .skip_value_for_verification(&mut VerificationState::default())
1869 .unwrap_err();
1870
1871 assert!(matches!(err, MaxMindDbError::InvalidDatabase { .. }));
1872 }
1873
1874 #[test]
1875 fn test_verification_rejects_and_does_not_cache_invalid_utf8() {
1876 let buf = [0x41, 0xff];
1877 let mut state = VerificationState::default();
1878
1879 for _ in 0..2 {
1880 let mut decoder = Decoder::new(&buf, 0);
1881 let err = decoder.skip_value_for_verification(&mut state).unwrap_err();
1882
1883 assert!(matches!(err, MaxMindDbError::InvalidDatabase { .. }));
1884 assert!(err.to_string().contains("invalid UTF-8"));
1885 assert!(state.validated.is_empty());
1886 assert!(state.active.is_empty());
1887 }
1888
1889 #[cfg(not(feature = "unsafe-str-decode"))]
1890 {
1891 let mut decoder = Decoder::new(&buf, 0);
1892 let err = String::deserialize(&mut decoder).unwrap_err();
1893 assert!(err.to_string().contains("invalid UTF-8"));
1894 }
1895 }
1896
1897 fn append_pointer(buf: &mut Vec<u8>, target: usize) {
1898 assert!(target < 2048);
1899 buf.push(0x20 | ((target >> 8) as u8));
1900 buf.push(target as u8);
1901 }
1902
1903 #[test]
1904 fn test_verification_caches_shared_pointer_targets() {
1905 let mut buf = vec![0x00, 0x07];
1909 let mut target = 0;
1910 const LEVELS: usize = 20;
1911
1912 for _ in 0..LEVELS {
1913 let array = buf.len();
1914 buf.extend_from_slice(&[0x02, 0x04]);
1915 append_pointer(&mut buf, target);
1916 append_pointer(&mut buf, target);
1917 target = array;
1918 }
1919
1920 let mut decoder = Decoder::new(&buf, target);
1921 let mut state = VerificationState::default();
1922 decoder.skip_value_for_verification(&mut state).unwrap();
1923
1924 assert_eq!(state.validated.len(), LEVELS + 1);
1925 assert!(state.active.is_empty());
1926 }
1927
1928 #[test]
1929 fn test_verification_rejects_data_pointer_cycles() {
1930 let mut buf = vec![0x01, 0x04];
1932 append_pointer(&mut buf, 4);
1933 buf.extend_from_slice(&[0x01, 0x04]);
1934 append_pointer(&mut buf, 0);
1935
1936 let mut decoder = Decoder::new(&buf, 0);
1937 let err = decoder
1938 .skip_value_for_verification(&mut VerificationState::default())
1939 .unwrap_err();
1940
1941 assert!(matches!(err, MaxMindDbError::InvalidDatabase { .. }));
1942 assert!(err.to_string().contains("cyclic data pointer"));
1943 }
1944}