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