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