Skip to main content

maxminddb/
decoder.rs

1//! Binary format decoder for MaxMind DB files.
2//!
3//! This module implements deserialization of the MaxMind DB binary format
4//! into Rust types via serde. The decoder handles all MaxMind DB data types
5//! including pointers, maps, arrays, and primitive types.
6//!
7//! Most users should not need to interact with this module directly.
8//! Use [`Reader::lookup()`](crate::Reader::lookup) for normal lookups.
9
10use 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
20// MaxMind DB type constants
21const 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
38/// Maximum recursion depth for nested data structures.
39/// This matches the value used in libmaxminddb and the Go reader.
40const MAXIMUM_DATA_STRUCTURE_DEPTH: u16 = 512;
41
42/// Lower limit for values skipped through unknown fields or IgnoredAny.
43/// Skipping is recursive and can be reached by corrupt data that callers did
44/// not explicitly request, so keep the limit below small default thread stacks.
45const 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/// Decoder for MaxMind DB binary format.
111///
112/// Implements serde's `Deserializer` trait. Handles pointer resolution,
113/// type coercion, and nested data structures.
114#[derive(Debug)]
115pub(crate) struct Decoder<'de> {
116    buf: &'de [u8],
117    limit: usize,
118    current_ptr: usize,
119    depth: u16,
120}
121
122/// Tracks data values visited by a single database verification pass.
123#[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    /// Check and increment depth, returning error if limit exceeded.
145    #[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    /// Decrement depth when exiting a nested structure.
157    #[inline]
158    fn exit_nested(&mut self) {
159        self.depth = self.depth.saturating_sub(1);
160    }
161
162    /// Create an InvalidDatabase error with current offset context.
163    #[inline]
164    fn invalid_db_error(&self, msg: &str) -> MaxMindDbError {
165        MaxMindDbError::invalid_database_at(msg, self.current_ptr)
166    }
167
168    /// Create a Decoding error with current offset context.
169    #[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        // SAFETY: Decoder constructors ensure `limit <= buf.len()`, and all
199        // callers reach this helper only after checking `end <= limit`.
200        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        // SAFETY: The check above proves `current_ptr < limit`, and decoder
220        // construction guarantees `limit <= buf.len()`.
221        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        // Extended type - size field is used differently
230        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        // Extended type: type 0 means read next byte for actual type
256        if type_num == TYPE_EXTENDED {
257            type_num = self.eat_byte()? + TYPE_MAP; // Extended types start at 7
258        }
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                // Pointer dereference - track depth
274                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            // Maps and arrays enter_nested in decode_any_value; exit when done
294            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                // Clamp to the end of the buffer so the next decode step fails
485                // with a normal bounds error instead of panicking here.
486                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        // SAFETY:
511        // A corrupt maxminddb will cause undefined behaviour.
512        // If the caller has verified the integrity of their database and trusts their upstream
513        // provider, they can opt-into the performance gains provided by this unsafe function via
514        // the `unsafe-str-decode` feature flag.
515        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            // ASCII is valid UTF-8, so this avoids the full validator fast path.
532            // SAFETY: `is_ascii()` guarantees UTF-8 validity.
533            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    // ========== Navigation methods for path decoding and verification ==========
543
544    /// Peeks at the type and size without consuming it.
545    /// Returns (size, type_num) and restores the position.
546    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    /// Consumes a map or array header in one pass, following a pointer if needed.
554    pub(crate) fn consume_container_header(&mut self) -> DecodeResult<(usize, u8)> {
555        self.size_and_type_following_pointers()
556    }
557
558    /// Gets size and type, following any pointers.
559    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    /// Reads a string's bytes directly, following pointers if needed.
626    /// Does NOT validate UTF-8.
627    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    /// Fast-path identifier decoding:
651    /// - Returns `Ok(Some(bytes))` and consumes the identifier when it is a string.
652    /// - Returns `Ok(None)` and restores `current_ptr` when the next value is not a string.
653    /// - Returns `Err` for malformed pointer chains or invalid string lengths.
654    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                // decode_pointer(size) temporarily dereferences by moving current_ptr
672                // to new_ptr; after size_and_type/read_string_bytes on the pointed
673                // value, restoring current_ptr = after_pointer resumes parsing right
674                // after the original pointer bytes. When result is Ok(None), also
675                // reset current_ptr = saved_ptr so the non-string identifier can be
676                // parsed normally by the caller without consuming the pointer token.
677                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    /// Skips the current value, following pointers.
691    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    /// Skips the current value and validates any referenced pointer targets.
697    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        // Headers and scalar payloads validate every cursor advance. A
749        // successful recursive skip therefore already guarantees that the
750        // cursor remains within the decoder limit.
751        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                // String or Bytes - skip size bytes
765                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                // Double - must be exactly 8 bytes
774                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                // Float - must be exactly 4 bytes
781                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                // Numeric types - skip size bytes
788                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                // Boolean - size field IS the value, no data bytes to skip
810                self.decode_bool(size).map(|_| ())
811            }
812            TYPE_MAP => {
813                // Map - skip size key-value pairs
814                let child_depth = self.check_skip_depth(skip_depth)?;
815                for _ in 0..size {
816                    // key
817                    self.skip_value_with_depth(child_depth)?;
818                    // value
819                    self.skip_value_with_depth(child_depth)?;
820                }
821                Ok(())
822            }
823            TYPE_ARRAY => {
824                // Array - skip size elements
825                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
924/// Deserializes any MaxMind DB value while exposing strings as raw bytes.
925///
926/// This helper is intended for format adapters that validate strings while
927/// converting them to another runtime's native string type. MMDB string values
928/// are delivered to [`Visitor::visit_newtype_struct`], which the adapter's
929/// visitor must implement. Its nested deserializer answers every
930/// `deserialize_*` call with [`Visitor::visit_borrowed_bytes`]; calling
931/// [`Deserializer::deserialize_bytes`] is the conventional choice. Genuine
932/// MMDB byte values continue to be delivered directly to
933/// [`Visitor::visit_borrowed_bytes`], so callers can distinguish the two
934/// types.
935///
936/// Callers decoding nested maps or arrays should invoke this helper again from
937/// the [`DeserializeSeed`] used for each nested value. Map keys can be read as
938/// unvalidated bytes with [`Deserializer::deserialize_identifier`]. Raw-string
939/// mode applies only to the value for which this helper is invoked. Nested
940/// values decoded without re-invoking it silently use normal string decoding,
941/// including the `unsafe-str-decode` behavior when that feature is enabled.
942/// Pointers are followed transparently and preserve the selected mode.
943///
944/// This function has its special effect only with this crate's deserializer;
945/// other Serde deserializers may treat the request as an ordinary newtype
946/// struct. The adapter is responsible for ensuring strict UTF-8 validation if
947/// malformed database strings must remain errors, as some runtimes replace
948/// invalid sequences instead. Unlike the `unsafe-str-decode` feature, this
949/// function itself never constructs an unvalidated Rust `str`.
950pub 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
1175// `SeqAccess` is provided to the `Visitor` to give it the ability to iterate
1176// through elements of the sequence.
1177impl<'de> SeqAccess<'de> for ArrayAccess<'_, 'de> {
1178    type Error = MaxMindDbError;
1179
1180    #[inline(always)]
1181    fn size_hint(&self) -> Option<usize> {
1182        // Never let a corrupt declared count drive an allocation larger than
1183        // the remaining encoded data can possibly fill.
1184        // Cursor advances are checked, so ordinary subtraction is sufficient.
1185        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        // Check if there are no more elements.
1194        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        // Deserialize an array element.
1205        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
1214// `MapAccess` is provided to the `Visitor` to give it the ability to iterate
1215// through entries of the map.
1216impl<'de> MapAccess<'de> for MapAccessor<'_, 'de> {
1217    type Error = MaxMindDbError;
1218
1219    #[inline(always)]
1220    fn size_hint(&self) -> Option<usize> {
1221        // Each map entry needs at least one control byte for both key and value.
1222        // Cursor advances are checked, so ordinary subtraction is sufficient.
1223        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        // Check if there are no more entries.
1232        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        // Deserialize a map key.
1243        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        // Check if there are no more entries.
1251        if self.count == 0 {
1252            return Err(self.de.decode_error("no more entries"));
1253        }
1254        self.count -= 1;
1255
1256        // Deserialize a map value.
1257        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        // Deserialize the variant identifier (string)
1274        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, // map with two entries
1493            0x44, b't', b'e', b'x', b't', // "text"
1494            0x41, 0xff, // invalid UTF-8 string value
1495            0x44, b'b', b'l', b'o', b'b', // "blob"
1496            0x81, 0xff, // byte value
1497        ];
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, // array with two elements
1515            0x41, 0xff, // invalid UTF-8 string value
1516            0x81, 0xff, // byte value
1517        ];
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, // pointer to offset two
1527            0x41, 0xff, // invalid UTF-8 string value
1528        ];
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, // map with two entries
1538            0x41, b'a', // "a"
1539            0x20, 0x0a, // pointer to the string at offset ten
1540            0x41, b'b', // "b"
1541            0x41, b'y', // "y"
1542            0x41, b'x', // pointed-to string "x"
1543        ];
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]; // map with eight entries
1560
1561        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, // map with one entry
1640            0x41, b'k', // "k"
1641            0x45, b'h', b'e', b'l', b'l', b'o', // "hello"
1642        ];
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        // Extended array with 284 declared elements and no element payload.
1823        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        // Extended bool type with an invalid size value of two.
1833        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        // A false boolean leaf followed by arrays containing two pointers to
1873        // the preceding value. Without caching, verification work doubles at
1874        // every level even though the encoded graph grows only linearly.
1875        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        // Two single-element arrays whose values point to each other.
1898        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}