use serde::de::{
self, value::BorrowedBytesDeserializer, DeserializeSeed, Deserializer, MapAccess, SeqAccess,
Visitor,
};
use serde::forward_to_deserialize_any;
use std::convert::TryInto;
use crate::error::MaxMindDbError;
mod verification;
pub(crate) use verification::VerificationState;
const TYPE_EXTENDED: usize = 0;
pub(crate) const TYPE_POINTER: usize = 1;
const TYPE_STRING: usize = 2;
const TYPE_DOUBLE: usize = 3;
const TYPE_BYTES: usize = 4;
const TYPE_UINT16: usize = 5;
const TYPE_UINT32: usize = 6;
pub(crate) const TYPE_MAP: usize = 7;
const TYPE_INT32: usize = 8;
const TYPE_UINT64: usize = 9;
const TYPE_UINT128: usize = 10;
pub(crate) const TYPE_ARRAY: usize = 11;
const TYPE_BOOL: usize = 14;
const TYPE_FLOAT: usize = 15;
const RAW_STRINGS_NEWTYPE: &str = "$maxminddb::raw_strings";
const MAXIMUM_DATA_STRUCTURE_DEPTH: u16 = 512;
const MAXIMUM_DATA_STRUCTURE_VALUES: u32 = 1 << 16;
const MAXIMUM_DATA_STRUCTURE_BYTES: usize = 2 << 20;
const MAXIMUM_UNCHARGED_IDENTIFIER_BYTES: usize =
MAXIMUM_DATA_STRUCTURE_BYTES / MAXIMUM_DATA_STRUCTURE_VALUES as usize;
const DEPTH_MASK: u32 = (1 << 10) - 1;
const BUDGET_VALUES_SHIFT: u32 = 10;
const BUDGET_VALUES_MASK: u32 = ((1 << 17) - 1) << BUDGET_VALUES_SHIFT;
const BUDGET_ACTIVE_MASK: u32 = 1 << 27;
const MAXIMUM_SKIPPED_DATA_STRUCTURE_DEPTH: u16 = 128;
#[inline(always)]
fn to_usize(base: u8, bytes: &[u8]) -> usize {
bytes
.iter()
.fold(base as usize, |acc, &b| (acc << 8) | b as usize)
}
#[cfg(not(feature = "unsafe-str-decode"))]
#[inline]
fn is_ascii(bytes: &[u8]) -> bool {
match bytes.len() {
4..=7 => {
let first = u32::from_ne_bytes(bytes[..4].try_into().unwrap());
let last = u32::from_ne_bytes(bytes[bytes.len() - 4..].try_into().unwrap());
(first | last) & 0x8080_8080 == 0
}
8..=16 => {
let first = u64::from_ne_bytes(bytes[..8].try_into().unwrap());
let last = u64::from_ne_bytes(bytes[bytes.len() - 8..].try_into().unwrap());
(first | last) & 0x8080_8080_8080_8080 == 0
}
_ => bytes.is_ascii(),
}
}
macro_rules! decode_int_like {
($name:ident, $ty:ty, $max_size:expr, $label:literal, $zero:expr) => {
fn $name(&mut self, size: usize) -> DecodeResult<$ty> {
match size {
s if s <= $max_size => {
let new_offset = self
.current_ptr
.checked_add(size)
.filter(|&offset| offset <= self.limit)
.ok_or_else(|| {
self.invalid_db_error(&format!("{} of size {}", $label, size))
})?;
let value = self
.slice(self.current_ptr, new_offset)
.iter()
.fold($zero, |acc, &b| (acc << 8) | <$ty>::from(b));
self.current_ptr = new_offset;
Ok(value)
}
s => Err(self.invalid_db_error(&format!("{} of size {}", $label, s))),
}
}
};
}
macro_rules! deserialize_direct_scalar {
($name:ident, $expected_type:expr, $label:literal, $visit:ident, $decode:ident) => {
fn $name<V>(self, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
let (size, type_num) = self.size_and_type()?;
self.decode_direct(size, type_num, $expected_type, $label, |de, size| {
visitor.$visit(de.$decode(size)?)
})
}
};
}
macro_rules! deserialize_direct_payload {
($name:ident, $expected_type:expr, $label:literal, $visit:ident, $decode:ident) => {
fn $name<V>(self, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
let (size, type_num) = self.size_and_type()?;
self.decode_direct(size, type_num, $expected_type, $label, |de, size| {
de.count_payload(size)?;
visitor.$visit(de.$decode(size)?)
})
}
};
}
enum Value<'a, 'de> {
Any { prev_ptr: usize },
Bytes(&'de [u8]),
String(&'de str),
RawString(&'de [u8]),
Bool(bool),
I32(i32),
U16(u16),
U32(u32),
U64(u64),
U128(u128),
F64(f64),
F32(f32),
Map(MapAccessor<'a, 'de, true>),
Array(ArrayAccess<'a, 'de>),
}
#[derive(Debug)]
pub(crate) struct Decoder<'de> {
buf: &'de [u8],
limit: usize,
current_ptr: usize,
state: u32,
payload_remaining: u32,
}
impl<'de> Decoder<'de> {
pub(crate) fn new(buf: &'de [u8], start_ptr: usize) -> Decoder<'de> {
Decoder::new_with_limit(buf, start_ptr, buf.len())
}
pub(crate) fn new_with_limit(buf: &'de [u8], start_ptr: usize, limit: usize) -> Decoder<'de> {
debug_assert!(limit <= buf.len());
Decoder {
buf,
limit,
current_ptr: start_ptr,
state: 0,
payload_remaining: MAXIMUM_DATA_STRUCTURE_BYTES as u32,
}
}
#[inline(always)]
fn activate_budget(&mut self) {
if self.state & BUDGET_ACTIVE_MASK == 0 {
self.state |= BUDGET_ACTIVE_MASK | (1 << BUDGET_VALUES_SHIFT);
}
}
#[inline]
fn enter_nested(&mut self) -> DecodeResult<()> {
if self.state & DEPTH_MASK >= u32::from(MAXIMUM_DATA_STRUCTURE_DEPTH) {
return Err(self.invalid_db_error(
"exceeded maximum data structure depth; database is likely corrupt",
));
}
self.state += 1;
Ok(())
}
#[inline]
fn exit_nested(&mut self) {
if self.state & DEPTH_MASK != 0 {
self.state -= 1;
}
}
#[inline(always)]
fn reserve_values(&mut self, count: usize) -> DecodeResult<()> {
if self.state & BUDGET_ACTIVE_MASK == 0 {
return Ok(());
}
let values_used = ((self.state & BUDGET_VALUES_MASK) >> BUDGET_VALUES_SHIFT) as usize;
if count > MAXIMUM_DATA_STRUCTURE_VALUES as usize - values_used {
return Err(
self.resource_limit_error("exceeded maximum number of data structure values")
);
}
self.state += (count as u32) << BUDGET_VALUES_SHIFT;
Ok(())
}
#[inline(always)]
fn reserve_container_values(&mut self, size: usize, type_num: usize) -> DecodeResult<()> {
self.activate_budget();
let count = if type_num == TYPE_MAP {
size.saturating_mul(2)
} else {
debug_assert_eq!(type_num, TYPE_ARRAY);
size
};
self.reserve_values(count)
}
#[inline(always)]
fn count_payload(&mut self, size: usize) -> DecodeResult<()> {
if self.state & BUDGET_ACTIVE_MASK == 0 {
return Ok(());
}
if size > self.payload_remaining as usize {
return Err(self.resource_limit_error(
"exceeded maximum size of data structure string and bytes values",
));
}
self.payload_remaining -= size as u32;
Ok(())
}
#[inline(always)]
fn count_identifier_payload(&mut self, size: usize) -> DecodeResult<()> {
if size <= MAXIMUM_UNCHARGED_IDENTIFIER_BYTES {
return Ok(());
}
self.count_long_identifier_payload(size)
}
#[cold]
#[inline(never)]
fn count_long_identifier_payload(&mut self, size: usize) -> DecodeResult<()> {
self.count_payload(size - MAXIMUM_UNCHARGED_IDENTIFIER_BYTES)
}
fn count_payload_at_current(&mut self) -> DecodeResult<bool> {
let saved_ptr = self.current_ptr;
let result = (|| {
let (mut size, mut type_num) = self.size_and_type()?;
if type_num == TYPE_POINTER {
self.current_ptr = self.decode_pointer(size);
(size, type_num) = self.size_and_type()?;
if type_num == TYPE_POINTER {
return Err(self.invalid_db_error("pointer points to another pointer"));
}
}
if type_num == TYPE_STRING || type_num == TYPE_BYTES {
self.count_payload(size)?;
return Ok(true);
}
Ok(false)
})();
self.current_ptr = saved_ptr;
result
}
#[inline]
fn invalid_db_error(&self, msg: &str) -> MaxMindDbError {
MaxMindDbError::invalid_database_at(msg, self.current_ptr)
}
#[inline]
fn decode_error(&self, msg: &str) -> MaxMindDbError {
MaxMindDbError::decoding_at(msg, self.current_ptr)
}
#[cold]
#[inline(never)]
fn resource_limit_error(&self, msg: &str) -> MaxMindDbError {
MaxMindDbError::resource_limit_at(msg, self.current_ptr)
}
#[inline(always)]
fn type_mismatch(&self, label: &str, type_num: usize) -> MaxMindDbError {
if type_num > usize::from(u8::MAX) {
self.invalid_db_error(&format!("unknown data type: {type_num}"))
} else {
self.decode_error(&format!("expected {label}, got type {type_num}"))
}
}
#[inline]
pub(crate) fn offset(&self) -> usize {
self.current_ptr
}
#[inline(always)]
fn checked_offset(&self, size: usize, label: &str) -> DecodeResult<usize> {
let new_offset = self.current_ptr.wrapping_add(size);
if new_offset < self.current_ptr || new_offset > self.limit {
return Err(self.invalid_db_error(&format!("{label} of size {size}")));
}
Ok(new_offset)
}
#[inline(always)]
fn slice(&self, start: usize, end: usize) -> &'de [u8] {
debug_assert!(start <= end);
debug_assert!(end <= self.limit);
debug_assert!(self.limit <= self.buf.len());
unsafe { self.buf.get_unchecked(start..end) }
}
#[inline(always)]
fn skip_bytes(&mut self, size: usize, label: &str) -> DecodeResult<()> {
debug_assert!(self.current_ptr <= self.limit);
if size > self.limit - self.current_ptr {
return Err(self.invalid_db_error(&format!("{label} of size {size}")));
}
self.current_ptr += size;
Ok(())
}
#[inline(always)]
fn eat_byte(&mut self) -> DecodeResult<u8> {
if self.current_ptr >= self.limit {
return Err(self.invalid_db_error("unexpected end of buffer"));
}
debug_assert!(self.limit <= self.buf.len());
let b = unsafe { *self.buf.get_unchecked(self.current_ptr) };
self.current_ptr += 1;
Ok(b)
}
#[inline(always)]
fn size_from_ctrl_byte(&mut self, ctrl_byte: u8, type_num: usize) -> DecodeResult<usize> {
let size = (ctrl_byte & 0x1f) as usize;
if type_num == TYPE_EXTENDED {
return Ok(size);
}
match size {
s if s < 29 => Ok(s),
29 => Ok(29_usize + self.eat_byte()? as usize),
30 => {
let b0 = self.eat_byte()? as usize;
let b1 = self.eat_byte()? as usize;
Ok(285_usize + (b0 << 8) + b1)
}
_ => {
let b0 = self.eat_byte()? as usize;
let b1 = self.eat_byte()? as usize;
let b2 = self.eat_byte()? as usize;
Ok(65_821_usize + (b0 << 16) + (b1 << 8) + b2)
}
}
}
#[inline(always)]
fn size_and_type(&mut self) -> DecodeResult<(usize, usize)> {
let ctrl_byte = self.eat_byte()?;
let mut type_num = usize::from(ctrl_byte >> 5);
if type_num == TYPE_EXTENDED {
type_num = usize::from(self.eat_byte()?) + TYPE_MAP;
}
self.size_from_ctrl_byte(ctrl_byte, type_num)
.map(|size| (size, type_num))
}
fn decode_any<V: Visitor<'de>>(&mut self, visitor: V) -> DecodeResult<V::Value> {
self.activate_budget();
self.decode_any_impl::<false, V>(visitor)
}
fn decode_any_impl<const RAW_STRINGS: bool, V: Visitor<'de>>(
&mut self,
visitor: V,
) -> DecodeResult<V::Value> {
match self.decode_any_value::<RAW_STRINGS>()? {
Value::Any { prev_ptr } => {
self.enter_nested()?;
let res = self.decode_any_impl::<RAW_STRINGS, V>(visitor);
self.exit_nested();
self.current_ptr = prev_ptr;
res
}
Value::Bool(x) => visitor.visit_bool(x),
Value::Bytes(x) => visitor.visit_borrowed_bytes(x),
Value::String(x) => visitor.visit_borrowed_str(x),
Value::RawString(x) => {
visitor.visit_newtype_struct(BorrowedBytesDeserializer::<MaxMindDbError>::new(x))
}
Value::I32(x) => visitor.visit_i32(x),
Value::U16(x) => visitor.visit_u16(x),
Value::U32(x) => visitor.visit_u32(x),
Value::U64(x) => visitor.visit_u64(x),
Value::U128(x) => visitor.visit_u128(x),
Value::F64(x) => visitor.visit_f64(x),
Value::F32(x) => visitor.visit_f32(x),
Value::Map(x) => {
let res = visitor.visit_map(x);
self.exit_nested();
res
}
Value::Array(x) => {
let res = visitor.visit_seq(x);
self.exit_nested();
res
}
}
}
fn deserialize_fixed_size_array<V>(&mut self, len: usize, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
let (size, type_num) = self.size_and_type()?;
self.decode_direct(size, type_num, TYPE_ARRAY, "array", |de, size| {
if size != len {
return Err(de.decode_error(&format!(
"expected tuple of length {len}, got array of length {size}"
)));
}
de.reserve_container_values(size, TYPE_ARRAY)?;
de.enter_nested()?;
let res = visitor.visit_seq(ArrayAccess { de, count: size });
de.exit_nested();
res
})
}
#[inline(always)]
fn decode_any_value<const RAW_STRINGS: bool>(&mut self) -> DecodeResult<Value<'_, 'de>> {
let (size, type_num) = self.size_and_type()?;
Ok(match type_num {
TYPE_POINTER => {
let new_ptr = self.decode_pointer(size);
let prev_ptr = self.current_ptr;
self.current_ptr = new_ptr;
Value::Any { prev_ptr }
}
TYPE_STRING if RAW_STRINGS => {
self.count_payload(size)?;
Value::RawString(self.read_string_bytes(size)?)
}
TYPE_STRING => {
self.count_payload(size)?;
Value::String(self.decode_string(size)?)
}
TYPE_DOUBLE => Value::F64(self.decode_double(size)?),
TYPE_BYTES => {
self.count_payload(size)?;
Value::Bytes(self.decode_bytes(size)?)
}
TYPE_UINT16 => Value::U16(self.decode_uint16(size)?),
TYPE_UINT32 => Value::U32(self.decode_uint32(size)?),
TYPE_MAP => {
self.reserve_container_values(size, TYPE_MAP)?;
self.enter_nested()?;
self.decode_map(size)
}
TYPE_INT32 => Value::I32(self.decode_int(size)?),
TYPE_UINT64 => Value::U64(self.decode_uint64(size)?),
TYPE_UINT128 => Value::U128(self.decode_uint128(size)?),
TYPE_ARRAY => {
self.reserve_container_values(size, TYPE_ARRAY)?;
self.enter_nested()?;
self.decode_array(size)
}
TYPE_BOOL => Value::Bool(self.decode_bool(size)?),
TYPE_FLOAT => Value::F32(self.decode_float(size)?),
u => return Err(self.invalid_db_error(&format!("unknown data type: {u}"))),
})
}
fn decode_array(&mut self, size: usize) -> Value<'_, 'de> {
Value::Array(ArrayAccess {
de: self,
count: size,
})
}
fn decode_bool(&mut self, size: usize) -> DecodeResult<bool> {
match size {
0 | 1 => Ok(size != 0),
s => Err(self.invalid_db_error(&format!("bool of size {s}"))),
}
}
fn decode_bytes(&mut self, size: usize) -> DecodeResult<&'de [u8]> {
let new_offset = self.checked_offset(size, "bytes")?;
let u8_slice = self.slice(self.current_ptr, new_offset);
self.current_ptr = new_offset;
Ok(u8_slice)
}
fn decode_float(&mut self, size: usize) -> DecodeResult<f32> {
let new_offset = self.checked_offset(size, "float")?;
let value: [u8; 4] = self
.slice(self.current_ptr, new_offset)
.try_into()
.map_err(|_| self.invalid_db_error(&format!("float of size {size}")))?;
self.current_ptr = new_offset;
let float_value = f32::from_be_bytes(value);
Ok(float_value)
}
fn decode_double(&mut self, size: usize) -> DecodeResult<f64> {
let new_offset = self.checked_offset(size, "double")?;
let value: [u8; 8] = self
.slice(self.current_ptr, new_offset)
.try_into()
.map_err(|_| self.invalid_db_error(&format!("double of size {size}")))?;
self.current_ptr = new_offset;
let float_value = f64::from_be_bytes(value);
Ok(float_value)
}
decode_int_like!(decode_uint64, u64, 8, "u64", 0_u64);
decode_int_like!(decode_uint128, u128, 16, "u128", 0_u128);
#[inline(always)]
fn read_u32_be(&mut self, size: usize, label: &str) -> DecodeResult<u32> {
if size > 4 {
return Err(self.invalid_db_error(&format!("{label} of size {size}")));
}
let new_offset = self
.current_ptr
.checked_add(size)
.filter(|&offset| offset <= self.limit)
.ok_or_else(|| self.invalid_db_error(&format!("{label} of size {}", size)))?;
let p = self.current_ptr;
let value = match size {
0 => 0,
1 => self.buf[p] as u32,
2 => ((self.buf[p] as u32) << 8) | self.buf[p + 1] as u32,
3 => {
((self.buf[p] as u32) << 16)
| ((self.buf[p + 1] as u32) << 8)
| self.buf[p + 2] as u32
}
_ => {
((self.buf[p] as u32) << 24)
| ((self.buf[p + 1] as u32) << 16)
| ((self.buf[p + 2] as u32) << 8)
| self.buf[p + 3] as u32
}
};
self.current_ptr = new_offset;
Ok(value)
}
#[inline(always)]
fn decode_uint32(&mut self, size: usize) -> DecodeResult<u32> {
self.read_u32_be(size, "u32")
}
#[inline(always)]
fn decode_uint16(&mut self, size: usize) -> DecodeResult<u16> {
if size > 2 {
return Err(self.invalid_db_error(&format!("u16 of size {size}")));
}
let new_offset = self
.current_ptr
.checked_add(size)
.filter(|&offset| offset <= self.limit)
.ok_or_else(|| self.invalid_db_error(&format!("u16 of size {}", size)))?;
let p = self.current_ptr;
let value = match size {
0 => 0,
1 => self.buf[p] as u16,
_ => ((self.buf[p] as u16) << 8) | self.buf[p + 1] as u16,
};
self.current_ptr = new_offset;
Ok(value)
}
fn decode_int(&mut self, size: usize) -> DecodeResult<i32> {
self.read_u32_be(size, "i32").map(|value| value as i32)
}
fn decode_map(&mut self, size: usize) -> Value<'_, 'de> {
Value::Map(MapAccessor {
de: self,
count: size * 2,
})
}
#[inline(always)]
fn decode_pointer(&mut self, size: usize) -> usize {
let pointer_value_offset = [0, 0, 2048, 526_336, 0];
let pointer_size = ((size >> 3) & 0x3) + 1;
let p = self.current_ptr;
let limit = self.limit;
let new_offset = match p.checked_add(pointer_size) {
Some(offset) if offset <= limit => offset,
_ => {
self.current_ptr = limit;
return limit;
}
};
let pointer_bytes = self.slice(p, new_offset);
self.current_ptr = new_offset;
let base = if pointer_size == 4 {
0
} else {
(size & 0x7) as u8
};
let unpacked = to_usize(base, pointer_bytes);
unpacked + pointer_value_offset[pointer_size]
}
#[cfg(feature = "unsafe-str-decode")]
#[inline(always)]
fn decode_string(&mut self, size: usize) -> DecodeResult<&'de str> {
use std::str::from_utf8_unchecked;
let new_offset = self.checked_offset(size, "string")?;
let bytes = self.slice(self.current_ptr, new_offset);
self.current_ptr = new_offset;
let v = unsafe { from_utf8_unchecked(bytes) };
Ok(v)
}
#[cfg(not(feature = "unsafe-str-decode"))]
#[inline(always)]
fn decode_string(&mut self, size: usize) -> DecodeResult<&'de str> {
#[cfg(feature = "simdutf8")]
use simdutf8::basic::from_utf8;
#[cfg(not(feature = "simdutf8"))]
use std::str::from_utf8;
use std::str::from_utf8_unchecked;
let new_offset = self.checked_offset(size, "string")?;
let bytes = self.slice(self.current_ptr, new_offset);
self.current_ptr = new_offset;
if is_ascii(bytes) {
let v = unsafe { from_utf8_unchecked(bytes) };
return Ok(v);
}
match from_utf8(bytes) {
Ok(v) => Ok(v),
Err(_) => Err(self.invalid_db_error("invalid UTF-8 in string")),
}
}
pub(crate) fn peek_type(&mut self) -> DecodeResult<(usize, usize)> {
let saved_ptr = self.current_ptr;
let result = self.size_and_type_following_pointers()?;
self.current_ptr = saved_ptr;
Ok(result)
}
pub(crate) fn consume_container_header(&mut self) -> DecodeResult<(usize, usize)> {
let (size, type_num) = self.size_and_type_following_pointers()?;
if type_num == TYPE_MAP || type_num == TYPE_ARRAY {
self.reserve_container_values(size, type_num)?;
}
Ok((size, type_num))
}
fn size_and_type_following_pointers(&mut self) -> DecodeResult<(usize, usize)> {
let (size, type_num) = self.size_and_type()?;
if type_num != TYPE_POINTER {
return Ok((size, type_num));
}
self.current_ptr = self.decode_pointer(size);
let (size, type_num) = self.size_and_type()?;
if type_num == TYPE_POINTER {
return Err(self.invalid_db_error("pointer points to another pointer"));
}
Ok((size, type_num))
}
#[inline(always)]
fn decode_direct<T, F>(
&mut self,
size: usize,
type_num: usize,
expected_type: usize,
label: &str,
decode: F,
) -> DecodeResult<T>
where
F: FnOnce(&mut Self, usize) -> DecodeResult<T>,
{
match type_num {
TYPE_POINTER => {
let new_ptr = self.decode_pointer(size);
let saved_ptr = self.current_ptr;
self.current_ptr = new_ptr;
self.enter_nested()?;
let result = (|| {
let (size, type_num) = self.size_and_type()?;
if type_num == TYPE_POINTER {
return Err(self.invalid_db_error("pointer points to another pointer"));
}
if type_num != expected_type {
return Err(self.type_mismatch(label, type_num));
}
decode(self, size)
})();
self.exit_nested();
self.current_ptr = saved_ptr;
result
}
t if t == expected_type => decode(self, size),
_ => Err(self.type_mismatch(label, type_num)),
}
}
#[inline(always)]
fn read_string_bytes(&mut self, size: usize) -> DecodeResult<&'de [u8]> {
let new_offset = self
.current_ptr
.checked_add(size)
.ok_or_else(|| self.invalid_db_error("string length exceeds buffer"))?;
if new_offset > self.limit {
return Err(self.invalid_db_error("string length exceeds buffer"));
}
let bytes = self.slice(self.current_ptr, new_offset);
self.current_ptr = new_offset;
Ok(bytes)
}
pub(crate) fn read_str_as_bytes(&mut self) -> DecodeResult<&'de [u8]> {
let (size, type_num) = self.size_and_type()?;
match type_num {
TYPE_POINTER => {
let new_ptr = self.decode_pointer(size);
let saved_ptr = self.current_ptr;
self.current_ptr = new_ptr;
let (size, type_num) = self.size_and_type()?;
let result = if type_num == TYPE_POINTER {
Err(self.invalid_db_error("pointer points to another pointer"))
} else if type_num == TYPE_STRING {
self.count_payload(size)
.and_then(|()| self.read_string_bytes(size))
} else {
Err(self.invalid_db_error(&format!("expected string, got type {type_num}")))
};
self.current_ptr = saved_ptr;
result
}
TYPE_STRING => {
self.count_payload(size)?;
self.read_string_bytes(size)
}
_ => Err(self.invalid_db_error(&format!("expected string, got type {type_num}"))),
}
}
#[inline(always)]
fn try_read_identifier_bytes(&mut self) -> DecodeResult<Option<&'de [u8]>> {
let saved_ptr = self.current_ptr;
let (size, type_num) = self.size_and_type()?;
match type_num {
TYPE_STRING => {
self.count_identifier_payload(size)?;
self.read_string_bytes(size).map(Some)
}
TYPE_POINTER => {
let new_ptr = self.decode_pointer(size);
let after_pointer = self.current_ptr;
self.current_ptr = new_ptr;
let (inner_size, inner_type) = self.size_and_type()?;
let result = if inner_type == TYPE_POINTER {
Err(self.invalid_db_error("pointer points to another pointer"))
} else if inner_type == TYPE_STRING {
let payload_result = self.count_identifier_payload(inner_size);
match payload_result {
Ok(()) => self.read_string_bytes(inner_size).map(Some),
Err(error) => Err(error),
}
} else {
Ok(None)
};
self.current_ptr = after_pointer;
if matches!(result, Ok(None)) {
self.current_ptr = saved_ptr;
}
result
}
_ => {
self.current_ptr = saved_ptr;
Ok(None)
}
}
}
pub(crate) fn skip_value(&mut self) -> DecodeResult<()> {
let (size, type_num) = self.size_and_type()?;
self.skip_value_inner(size, type_num, 0)
}
pub(crate) fn skip_value_for_verification(
&mut self,
state: &mut VerificationState,
) -> DecodeResult<()> {
let offset = self.current_ptr;
if state.validated.contains(&offset) {
return Ok(());
}
if !state.active.insert(offset) {
return Err(
self.invalid_db_error(&format!("cyclic data pointer references offset {offset}"))
);
}
let result = (|| {
let (size, type_num) = self.size_and_type()?;
self.skip_value_inner_for_verification(size, type_num, 0, state)?;
self.validate_skip_end()
})();
state.active.remove(&offset);
if result.is_ok() {
state.validated.insert(offset);
}
result
}
#[inline(always)]
pub(crate) fn validate_skip_end(&mut self) -> DecodeResult<()> {
if self.current_ptr > self.limit {
return Err(self.invalid_db_error("skipped value extends beyond buffer"));
}
Ok(())
}
#[inline(always)]
fn check_skip_depth(&self, skip_depth: u16) -> DecodeResult<u16> {
if skip_depth == MAXIMUM_SKIPPED_DATA_STRUCTURE_DEPTH {
return self.skip_depth_error();
}
Ok(skip_depth + 1)
}
#[cold]
fn skip_depth_error(&self) -> DecodeResult<u16> {
Err(self
.invalid_db_error("exceeded maximum data structure depth; database is likely corrupt"))
}
#[inline(always)]
fn skip_value_inner(
&mut self,
size: usize,
type_num: usize,
skip_depth: u16,
) -> DecodeResult<()> {
match type_num {
TYPE_POINTER => {
let pointer_size = ((size >> 3) & 0x3) + 1;
self.checked_offset(pointer_size, "pointer")?;
self.decode_pointer(size);
Ok(())
}
TYPE_STRING | TYPE_BYTES => {
let label = if type_num == TYPE_STRING {
"string"
} else {
"bytes"
};
self.skip_bytes(size, label)
}
TYPE_DOUBLE => {
if size != 8 {
return Err(self.invalid_db_error(&format!("double of size {size}")));
}
self.skip_bytes(size, "double")
}
TYPE_FLOAT => {
if size != 4 {
return Err(self.invalid_db_error(&format!("float of size {size}")));
}
self.skip_bytes(size, "float")
}
TYPE_UINT16 | TYPE_UINT32 | TYPE_INT32 | TYPE_UINT64 | TYPE_UINT128 => {
let label = match type_num {
TYPE_UINT16 => "u16",
TYPE_UINT32 => "u32",
TYPE_INT32 => "i32",
TYPE_UINT64 => "u64",
TYPE_UINT128 => "u128",
_ => unreachable!(),
};
let max_size = match type_num {
TYPE_UINT16 => 2,
TYPE_UINT32 | TYPE_INT32 => 4,
TYPE_UINT64 => 8,
TYPE_UINT128 => 16,
_ => unreachable!(),
};
if size > max_size {
return Err(self.invalid_db_error(&format!("{label} of size {size}")));
}
self.skip_bytes(size, label)
}
TYPE_BOOL => {
self.decode_bool(size).map(|_| ())
}
TYPE_MAP => {
self.reserve_container_values(size, TYPE_MAP)?;
let child_depth = self.check_skip_depth(skip_depth)?;
for _ in 0..size {
self.skip_value_with_depth(child_depth)?;
self.skip_value_with_depth(child_depth)?;
}
Ok(())
}
TYPE_ARRAY => {
self.reserve_container_values(size, TYPE_ARRAY)?;
let child_depth = self.check_skip_depth(skip_depth)?;
for _ in 0..size {
self.skip_value_with_depth(child_depth)?;
}
Ok(())
}
u => Err(self.invalid_db_error(&format!("unknown data type: {u}"))),
}
}
#[inline(always)]
fn skip_value_with_depth(&mut self, skip_depth: u16) -> DecodeResult<()> {
let (size, type_num) = self.size_and_type()?;
self.skip_value_inner(size, type_num, skip_depth)
}
fn skip_value_inner_for_verification(
&mut self,
size: usize,
type_num: usize,
skip_depth: u16,
state: &mut VerificationState,
) -> DecodeResult<()> {
state.charge(1, self.current_ptr)?;
match type_num {
TYPE_STRING => {
let end = self.checked_offset(size, "string")?;
state.charge(size, self.current_ptr)?;
let bytes = self.slice(self.current_ptr, end);
self.current_ptr = end;
std::str::from_utf8(bytes)
.map(|_| ())
.map_err(|_| self.invalid_db_error("invalid UTF-8 in string"))
}
TYPE_POINTER => {
let target = self.decode_pointer(size);
let child_depth = self.check_skip_depth(skip_depth)?;
self.verify_pointer_target(target, child_depth, state)
}
TYPE_MAP => {
let child_depth = self.check_skip_depth(skip_depth)?;
for _ in 0..size {
self.skip_value_with_verification(child_depth, state)?;
self.skip_value_with_verification(child_depth, state)?;
}
self.validate_skip_end()
}
TYPE_ARRAY => {
let child_depth = self.check_skip_depth(skip_depth)?;
for _ in 0..size {
self.skip_value_with_verification(child_depth, state)?;
}
self.validate_skip_end()
}
_ => self.skip_value_inner(size, type_num, skip_depth),
}
}
fn skip_value_with_verification(
&mut self,
skip_depth: u16,
state: &mut VerificationState,
) -> DecodeResult<()> {
let (size, type_num) = self.size_and_type()?;
self.skip_value_inner_for_verification(size, type_num, skip_depth, state)
}
fn verify_pointer_target(
&mut self,
target: usize,
skip_depth: u16,
state: &mut VerificationState,
) -> DecodeResult<()> {
if state.validated.contains(&target) {
return Ok(());
}
if !state.active.insert(target) {
return Err(
self.invalid_db_error(&format!("cyclic data pointer references offset {target}"))
);
}
let continuation = self.current_ptr;
self.current_ptr = target;
let result = (|| {
let (size, type_num) = self.size_and_type()?;
self.skip_value_inner_for_verification(size, type_num, skip_depth, state)?;
self.validate_skip_end()
})();
self.current_ptr = continuation;
state.active.remove(&target);
if result.is_ok() {
state.validated.insert(target);
}
result
}
}
pub type DecodeResult<T> = Result<T, MaxMindDbError>;
pub fn deserialize_any_with_raw_strings<'de, D, V>(
deserializer: D,
visitor: V,
) -> Result<V::Value, D::Error>
where
D: Deserializer<'de>,
V: Visitor<'de>,
{
deserializer.deserialize_newtype_struct(RAW_STRINGS_NEWTYPE, visitor)
}
impl<'de: 'a, 'a> de::Deserializer<'de> for &'a mut Decoder<'de> {
type Error = MaxMindDbError;
fn deserialize_any<V>(self, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
self.decode_any(visitor)
}
fn deserialize_option<V>(self, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
visitor.visit_some(self)
}
deserialize_direct_scalar!(deserialize_bool, TYPE_BOOL, "bool", visit_bool, decode_bool);
deserialize_direct_scalar!(
deserialize_u16,
TYPE_UINT16,
"u16",
visit_u16,
decode_uint16
);
deserialize_direct_scalar!(
deserialize_u32,
TYPE_UINT32,
"u32",
visit_u32,
decode_uint32
);
deserialize_direct_scalar!(
deserialize_u64,
TYPE_UINT64,
"u64",
visit_u64,
decode_uint64
);
deserialize_direct_scalar!(
deserialize_u128,
TYPE_UINT128,
"u128",
visit_u128,
decode_uint128
);
deserialize_direct_scalar!(deserialize_i32, TYPE_INT32, "i32", visit_i32, decode_int);
deserialize_direct_scalar!(
deserialize_f32,
TYPE_FLOAT,
"float",
visit_f32,
decode_float
);
deserialize_direct_scalar!(
deserialize_f64,
TYPE_DOUBLE,
"double",
visit_f64,
decode_double
);
deserialize_direct_payload!(
deserialize_str,
TYPE_STRING,
"string",
visit_borrowed_str,
decode_string
);
fn deserialize_string<V>(self, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
self.deserialize_str(visitor)
}
deserialize_direct_payload!(
deserialize_bytes,
TYPE_BYTES,
"bytes",
visit_borrowed_bytes,
decode_bytes
);
fn deserialize_byte_buf<V>(self, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
self.deserialize_bytes(visitor)
}
fn deserialize_seq<V>(self, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
let (size, type_num) = self.size_and_type()?;
self.decode_direct(size, type_num, TYPE_ARRAY, "array", |de, size| {
de.reserve_container_values(size, TYPE_ARRAY)?;
de.enter_nested()?;
let res = visitor.visit_seq(ArrayAccess { de, count: size });
de.exit_nested();
res
})
}
fn deserialize_tuple<V>(self, len: usize, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
self.deserialize_fixed_size_array(len, visitor)
}
fn deserialize_tuple_struct<V>(
self,
_name: &'static str,
len: usize,
visitor: V,
) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
self.deserialize_fixed_size_array(len, visitor)
}
fn deserialize_map<V>(self, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
self.deserialize_map_impl(visitor)
}
fn deserialize_struct<V>(
self,
_name: &'static str,
_fields: &'static [&'static str],
visitor: V,
) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
self.deserialize_map_impl(visitor)
}
fn is_human_readable(&self) -> bool {
false
}
fn deserialize_ignored_any<V>(self, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
self.skip_value()?;
visitor.visit_unit()
}
fn deserialize_enum<V>(
self,
_name: &'static str,
_variants: &'static [&'static str],
visitor: V,
) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
self.activate_budget();
visitor.visit_enum(EnumAccessor { de: self })
}
fn deserialize_identifier<V>(self, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
match self.try_read_identifier_bytes()? {
Some(bytes) => visitor.visit_borrowed_bytes(bytes),
None => self.decode_any(visitor),
}
}
fn deserialize_newtype_struct<V>(self, name: &'static str, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
if name == RAW_STRINGS_NEWTYPE {
self.activate_budget();
self.decode_any_impl::<true, V>(visitor)
} else {
self.decode_any(visitor)
}
}
forward_to_deserialize_any! {
i8 i16 i64 i128 u8 char unit unit_struct
}
}
impl<'de> Decoder<'de> {
fn deserialize_map_impl<V>(&mut self, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
let (size, type_num) = self.size_and_type()?;
self.decode_direct(size, type_num, TYPE_MAP, "map", |de, size| {
de.reserve_container_values(size, TYPE_MAP)?;
de.enter_nested()?;
let res = visitor.visit_map(MapAccessor::<false> {
de,
count: size * 2,
});
de.exit_nested();
res
})
}
}
struct ArrayAccess<'a, 'de: 'a> {
de: &'a mut Decoder<'de>,
count: usize,
}
impl<'de> SeqAccess<'de> for ArrayAccess<'_, 'de> {
type Error = MaxMindDbError;
#[inline(always)]
fn size_hint(&self) -> Option<usize> {
debug_assert!(self.de.current_ptr <= self.de.limit);
Some(self.count.min(self.de.limit - self.de.current_ptr))
}
fn next_element_seed<T>(&mut self, seed: T) -> DecodeResult<Option<T::Value>>
where
T: DeserializeSeed<'de>,
{
if self.count == 0 {
if self.de.current_ptr > self.de.limit {
return Err(self
.de
.invalid_db_error("skipped value extends beyond buffer"));
}
return Ok(None);
}
self.count -= 1;
seed.deserialize(&mut *self.de).map(Some)
}
}
struct MapAccessor<'a, 'de: 'a, const BUDGETED: bool = false> {
de: &'a mut Decoder<'de>,
count: usize,
}
impl<'de, const BUDGETED: bool> MapAccess<'de> for MapAccessor<'_, 'de, BUDGETED> {
type Error = MaxMindDbError;
#[inline(always)]
fn size_hint(&self) -> Option<usize> {
debug_assert!(self.de.current_ptr <= self.de.limit);
Some((self.count / 2).min((self.de.limit - self.de.current_ptr) / 2))
}
fn next_key_seed<K>(&mut self, seed: K) -> DecodeResult<Option<K::Value>>
where
K: DeserializeSeed<'de>,
{
if self.count == 0 {
if self.de.current_ptr > self.de.limit {
return Err(self
.de
.invalid_db_error("skipped value extends beyond buffer"));
}
return Ok(None);
}
self.count -= 1;
if BUDGETED {
let payload_remaining_before = self.de.payload_remaining;
let payload_precharged = self.de.count_payload_at_current()?;
let payload_remaining_after = self.de.payload_remaining;
self.de.payload_remaining = payload_remaining_before;
let result = seed.deserialize(&mut *self.de).map(Some);
if payload_precharged {
self.de.payload_remaining = self.de.payload_remaining.min(payload_remaining_after);
}
return result;
}
seed.deserialize(&mut *self.de).map(Some)
}
fn next_value_seed<V>(&mut self, seed: V) -> DecodeResult<V::Value>
where
V: DeserializeSeed<'de>,
{
if self.count == 0 {
return Err(self.de.decode_error("no more entries"));
}
self.count -= 1;
seed.deserialize(&mut *self.de)
}
}
struct EnumAccessor<'a, 'de: 'a> {
de: &'a mut Decoder<'de>,
}
impl<'de> de::EnumAccess<'de> for EnumAccessor<'_, 'de> {
type Error = MaxMindDbError;
type Variant = Self;
fn variant_seed<V>(self, seed: V) -> DecodeResult<(V::Value, Self::Variant)>
where
V: DeserializeSeed<'de>,
{
let variant = seed.deserialize(&mut *self.de)?;
Ok((variant, self))
}
}
impl<'de> de::VariantAccess<'de> for EnumAccessor<'_, 'de> {
type Error = MaxMindDbError;
fn unit_variant(self) -> DecodeResult<()> {
Ok(())
}
fn newtype_variant_seed<T>(self, seed: T) -> DecodeResult<T::Value>
where
T: DeserializeSeed<'de>,
{
self.de.reserve_values(1)?;
self.de.enter_nested()?;
let result = seed.deserialize(&mut *self.de);
self.de.exit_nested();
result
}
fn tuple_variant<V>(self, len: usize, visitor: V) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
self.de.deserialize_fixed_size_array(len, visitor)
}
fn struct_variant<V>(
self,
_fields: &'static [&'static str],
visitor: V,
) -> DecodeResult<V::Value>
where
V: Visitor<'de>,
{
de::Deserializer::deserialize_map(&mut *self.de, visitor)
}
}
#[cfg(test)]
mod tests {
use std::fmt;
use serde::de::{DeserializeSeed, Deserializer, MapAccess, SeqAccess, Visitor};
use serde::Deserialize;
use crate::{deserialize_any_with_raw_strings, MaxMindDbError, Reader};
use super::{Decoder, VerificationState};
#[derive(Debug, PartialEq)]
enum RawValue<'de> {
String(&'de [u8]),
Bytes(&'de [u8]),
Bool(bool),
I32(i32),
U16(u16),
U32(u32),
U64(u64),
U128(u128),
F32(f32),
F64(f64),
Array(Vec<RawValue<'de>>),
Map(Vec<(Vec<u8>, RawValue<'de>)>),
}
impl<'de> Deserialize<'de> for RawValue<'de> {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
RawValueSeed.deserialize(deserializer)
}
}
struct RawValueSeed;
impl<'de> DeserializeSeed<'de> for RawValueSeed {
type Value = RawValue<'de>;
fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: Deserializer<'de>,
{
deserialize_any_with_raw_strings(deserializer, RawValueVisitor)
}
}
struct RawValueVisitor;
impl<'de> Visitor<'de> for RawValueVisitor {
type Value = RawValue<'de>;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("an MMDB value")
}
fn visit_newtype_struct<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_bytes(RawStringVisitor)
}
fn visit_borrowed_bytes<E>(self, bytes: &'de [u8]) -> Result<Self::Value, E> {
Ok(RawValue::Bytes(bytes))
}
fn visit_bool<E>(self, value: bool) -> Result<Self::Value, E> {
Ok(RawValue::Bool(value))
}
fn visit_i32<E>(self, value: i32) -> Result<Self::Value, E> {
Ok(RawValue::I32(value))
}
fn visit_u16<E>(self, value: u16) -> Result<Self::Value, E> {
Ok(RawValue::U16(value))
}
fn visit_u32<E>(self, value: u32) -> Result<Self::Value, E> {
Ok(RawValue::U32(value))
}
fn visit_u64<E>(self, value: u64) -> Result<Self::Value, E> {
Ok(RawValue::U64(value))
}
fn visit_u128<E>(self, value: u128) -> Result<Self::Value, E> {
Ok(RawValue::U128(value))
}
fn visit_f32<E>(self, value: f32) -> Result<Self::Value, E> {
Ok(RawValue::F32(value))
}
fn visit_f64<E>(self, value: f64) -> Result<Self::Value, E> {
Ok(RawValue::F64(value))
}
fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
where
A: MapAccess<'de>,
{
let mut entries = Vec::with_capacity(map.size_hint().unwrap_or(0));
while let Some(key) = map.next_key_seed(RawIdentifierSeed)? {
let value = map.next_value_seed(RawValueSeed)?;
entries.push((key, value));
}
Ok(RawValue::Map(entries))
}
fn visit_seq<A>(self, mut sequence: A) -> Result<Self::Value, A::Error>
where
A: SeqAccess<'de>,
{
let mut values = Vec::with_capacity(sequence.size_hint().unwrap_or(0));
while let Some(value) = sequence.next_element_seed(RawValueSeed)? {
values.push(value);
}
Ok(RawValue::Array(values))
}
}
struct RawStringVisitor;
impl<'de> Visitor<'de> for RawStringVisitor {
type Value = RawValue<'de>;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("borrowed MMDB string bytes")
}
fn visit_borrowed_bytes<E>(self, bytes: &'de [u8]) -> Result<Self::Value, E> {
Ok(RawValue::String(bytes))
}
}
struct RawIdentifierSeed;
impl<'de> DeserializeSeed<'de> for RawIdentifierSeed {
type Value = Vec<u8>;
fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_identifier(RawIdentifierVisitor)
}
}
struct RawIdentifierVisitor;
impl<'de> Visitor<'de> for RawIdentifierVisitor {
type Value = Vec<u8>;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("borrowed MMDB map-key bytes")
}
fn visit_borrowed_bytes<E>(self, bytes: &'de [u8]) -> Result<Self::Value, E> {
Ok(bytes.to_vec())
}
}
struct AnyIdentifierSeed;
impl<'de> DeserializeSeed<'de> for AnyIdentifierSeed {
type Value = usize;
fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_any(AnyIdentifierVisitor)
}
}
struct AnyIdentifierVisitor;
impl<'de> Visitor<'de> for AnyIdentifierVisitor {
type Value = usize;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("an MMDB map key")
}
fn visit_borrowed_str<E>(self, value: &'de str) -> Result<Self::Value, E> {
Ok(value.len())
}
fn visit_borrowed_bytes<E>(self, value: &'de [u8]) -> Result<Self::Value, E> {
Ok(value.len())
}
}
#[derive(Debug, PartialEq)]
struct AnyKeyMap(usize);
struct AnyKeyMapVisitor;
impl<'de> Visitor<'de> for AnyKeyMapVisitor {
type Value = AnyKeyMap;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("an MMDB map")
}
fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
where
A: MapAccess<'de>,
{
let mut count = 0;
while let Some(key_size) = map.next_key_seed(AnyIdentifierSeed)? {
assert_eq!(key_size, 4096);
map.next_value::<serde::de::IgnoredAny>()?;
count += 1;
}
Ok(AnyKeyMap(count))
}
}
#[allow(dead_code)]
#[derive(Debug)]
struct OwnedBytes(Vec<u8>);
impl<'de> Deserialize<'de> for OwnedBytes {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_byte_buf(OwnedBytesVisitor)
}
}
struct OwnedBytesVisitor;
impl<'de> Visitor<'de> for OwnedBytesVisitor {
type Value = OwnedBytes;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("an MMDB byte value")
}
fn visit_borrowed_bytes<E>(self, value: &'de [u8]) -> Result<Self::Value, E> {
Ok(OwnedBytes(value.to_vec()))
}
fn visit_byte_buf<E>(self, value: Vec<u8>) -> Result<Self::Value, E> {
Ok(OwnedBytes(value))
}
}
#[test]
fn raw_string_mode_distinguishes_strings_from_bytes() {
let mut string_decoder = Decoder::new(&[0x42, 0xff, 0xfe], 0);
let string = RawValueSeed.deserialize(&mut string_decoder).unwrap();
assert_eq!(string, RawValue::String(&[0xff, 0xfe]));
let mut bytes_decoder = Decoder::new(&[0x82, 0xff, 0xfe], 0);
let bytes = RawValueSeed.deserialize(&mut bytes_decoder).unwrap();
assert_eq!(bytes, RawValue::Bytes(&[0xff, 0xfe]));
}
#[test]
fn raw_string_mode_recurses_through_maps() {
let encoded = [
0x02, 0x00, 0x44, b't', b'e', b'x', b't', 0x41, 0xff, 0x44, b'b', b'l', b'o', b'b', 0x81, 0xff, ];
let mut decoder = Decoder::new(&encoded, 0);
let value = RawValueSeed.deserialize(&mut decoder).unwrap();
assert_eq!(
value,
RawValue::Map(vec![
(b"text".to_vec(), RawValue::String(&[0xff])),
(b"blob".to_vec(), RawValue::Bytes(&[0xff])),
])
);
}
#[test]
fn raw_string_mode_recurses_through_arrays_and_pointers() {
let encoded_array = [
0x02, 0x04, 0x41, 0xff, 0x81, 0xff, ];
let mut array_decoder = Decoder::new(&encoded_array, 0);
let array = RawValueSeed.deserialize(&mut array_decoder).unwrap();
assert_eq!(
array,
RawValue::Array(vec![RawValue::String(&[0xff]), RawValue::Bytes(&[0xff]),])
);
let encoded_pointer = [
0x20, 0x02, 0x41, 0xff, ];
let mut pointer_decoder = Decoder::new(&encoded_pointer, 0);
let pointer = RawValueSeed.deserialize(&mut pointer_decoder).unwrap();
assert_eq!(pointer, RawValue::String(&[0xff]));
}
#[test]
fn raw_string_mode_restores_pointer_continuation_in_maps() {
let encoded = [
0x02, 0x00, 0x41, b'a', 0x20, 0x0a, 0x41, b'b', 0x41, b'y', 0x41, b'x', ];
let mut decoder = Decoder::new(&encoded, 0);
let value = RawValueSeed.deserialize(&mut decoder).unwrap();
assert_eq!(
value,
RawValue::Map(vec![
(b"a".to_vec(), RawValue::String(b"x")),
(b"b".to_vec(), RawValue::String(b"y")),
])
);
}
#[test]
fn raw_string_mode_decodes_all_scalar_types() {
let mut encoded = vec![0x08, 0x00];
encoded.extend_from_slice(&[0x41, b'd', 0x68]);
encoded.extend_from_slice(&1.5_f64.to_be_bytes());
encoded.extend_from_slice(&[0x41, b's', 0xa2, 0x01, 0x02]);
encoded.extend_from_slice(&[0x41, b'i', 0xc4, 0x01, 0x02, 0x03, 0x04]);
encoded.extend_from_slice(&[0x41, b'n', 0x04, 0x01]);
encoded.extend_from_slice(&(-2_i32).to_be_bytes());
encoded.extend_from_slice(&[0x41, b'l', 0x08, 0x02]);
encoded.extend_from_slice(&0x0102_0304_0506_0708_u64.to_be_bytes());
encoded.extend_from_slice(&[0x41, b'x', 0x10, 0x03]);
encoded.extend_from_slice(&0x0102_0304_0506_0708_1112_1314_1516_1718_u128.to_be_bytes());
encoded.extend_from_slice(&[0x41, b'b', 0x01, 0x07]);
encoded.extend_from_slice(&[0x41, b'f', 0x04, 0x08]);
encoded.extend_from_slice(&2.5_f32.to_be_bytes());
let mut decoder = Decoder::new(&encoded, 0);
let value = RawValueSeed.deserialize(&mut decoder).unwrap();
assert_eq!(
value,
RawValue::Map(vec![
(b"d".to_vec(), RawValue::F64(1.5)),
(b"s".to_vec(), RawValue::U16(0x0102)),
(b"i".to_vec(), RawValue::U32(0x0102_0304)),
(b"n".to_vec(), RawValue::I32(-2)),
(b"l".to_vec(), RawValue::U64(0x0102_0304_0506_0708)),
(
b"x".to_vec(),
RawValue::U128(0x0102_0304_0506_0708_1112_1314_1516_1718)
),
(b"b".to_vec(), RawValue::Bool(true)),
(b"f".to_vec(), RawValue::F32(2.5)),
])
);
}
#[test]
fn raw_string_mode_rejects_excessive_pointer_depth_and_unknown_types() {
std::thread::Builder::new()
.stack_size(8 * 1024 * 1024)
.spawn(|| {
let mut cyclic_decoder = Decoder::new(&[0x20, 0x00], 0);
let depth_err = RawValueSeed.deserialize(&mut cyclic_decoder).unwrap_err();
assert!(depth_err
.to_string()
.contains("exceeded maximum data structure depth"));
})
.unwrap()
.join()
.unwrap();
let mut unknown_decoder = Decoder::new(&[0x00, 0x06], 0);
let type_err = RawValueSeed.deserialize(&mut unknown_decoder).unwrap_err();
assert!(type_err.to_string().contains("unknown data type: 13"));
}
#[test]
fn malformed_extended_types_return_errors_instead_of_overflowing() {
for extended_type in 249..=u8::MAX {
let encoded = [0x00, extended_type];
let mut decoder = Decoder::new(&encoded, 0);
let error = RawValueSeed.deserialize(&mut decoder).unwrap_err();
assert!(matches!(error, MaxMindDbError::InvalidDatabase { .. }));
assert!(error.to_string().contains(&format!(
"unknown data type: {}",
u16::from(extended_type) + 7
)));
let mut typed_decoder = Decoder::new(&encoded, 0);
let typed_error =
<u32 as serde::Deserialize>::deserialize(&mut typed_decoder).unwrap_err();
assert!(matches!(
typed_error,
MaxMindDbError::InvalidDatabase { .. }
));
}
}
#[test]
fn nested_values_without_raw_opt_in_use_normal_string_decoding() {
struct NestedNormalVisitor;
impl<'de> Visitor<'de> for NestedNormalVisitor {
type Value = &'de str;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("an MMDB map containing a string")
}
fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
where
A: MapAccess<'de>,
{
let Some(_key) = map.next_key_seed(RawIdentifierSeed)? else {
return Err(serde::de::Error::custom("expected one map entry"));
};
map.next_value::<&'de str>()
}
}
let encoded = [
0x01, 0x00, 0x41, b'k', 0x45, b'h', b'e', b'l', b'l', b'o', ];
let mut decoder = Decoder::new(&encoded, 0);
let value = deserialize_any_with_raw_strings(&mut decoder, NestedNormalVisitor).unwrap();
assert_eq!(value, "hello");
}
fn raw_map_value<'value, 'de>(
value: &'value RawValue<'de>,
key: &[u8],
) -> &'value RawValue<'de> {
let RawValue::Map(entries) = value else {
panic!("expected map, got {value:?}");
};
entries
.iter()
.find_map(|(entry_key, value)| (entry_key == key).then_some(value))
.unwrap_or_else(|| panic!("missing map key {:?}", String::from_utf8_lossy(key)))
}
#[test]
fn raw_string_mode_decodes_reader_lookup_results() {
let reader = Reader::open_readfile("test-data/test-data/GeoIP2-City-Test.mmdb").unwrap();
let lookup = reader.lookup("89.160.20.128".parse().unwrap()).unwrap();
let value = lookup.decode::<RawValue<'_>>().unwrap().unwrap();
let city = raw_map_value(&value, b"city");
let city_names = raw_map_value(city, b"names");
assert_eq!(
raw_map_value(city_names, b"en"),
&RawValue::String("Linköping".as_bytes())
);
let country = raw_map_value(&value, b"country");
assert_eq!(
raw_map_value(country, b"is_in_european_union"),
&RawValue::Bool(true)
);
let location = raw_map_value(&value, b"location");
assert_eq!(
raw_map_value(location, b"accuracy_radius"),
&RawValue::U16(76)
);
assert_eq!(
raw_map_value(location, b"latitude"),
&RawValue::F64(58.4167)
);
let subdivisions = raw_map_value(&value, b"subdivisions");
let RawValue::Array(subdivisions) = subdivisions else {
panic!("expected subdivisions array, got {subdivisions:?}");
};
assert!(!subdivisions.is_empty());
}
#[test]
fn ordinary_string_decoding_remains_validated() {
struct OrdinaryNewtypeSeed;
impl<'de> DeserializeSeed<'de> for OrdinaryNewtypeSeed {
type Value = &'de str;
fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_newtype_struct("ordinary", StringVisitor)
}
}
struct StringVisitor;
impl<'de> Visitor<'de> for StringVisitor {
type Value = &'de str;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("a borrowed string")
}
fn visit_borrowed_str<E>(self, value: &'de str) -> Result<Self::Value, E> {
Ok(value)
}
}
let mut valid_decoder = Decoder::new(&[0x45, b'h', b'e', b'l', b'l', b'o'], 0);
assert_eq!(String::deserialize(&mut valid_decoder).unwrap(), "hello");
let mut newtype_decoder = Decoder::new(&[0x45, b'h', b'e', b'l', b'l', b'o'], 0);
assert_eq!(
OrdinaryNewtypeSeed
.deserialize(&mut newtype_decoder)
.unwrap(),
"hello"
);
#[cfg(not(feature = "unsafe-str-decode"))]
{
let mut invalid_decoder = Decoder::new(&[0x41, 0xff], 0);
let err = String::deserialize(&mut invalid_decoder).unwrap_err();
assert!(err.to_string().contains("invalid UTF-8"));
}
}
#[test]
fn test_decoder_accepts_tuple_with_matching_length() {
#[allow(dead_code)]
#[derive(Debug, serde::Deserialize)]
struct TupleRecord {
array: (u32, u32, u32),
}
#[allow(dead_code)]
#[derive(Debug, serde::Deserialize)]
struct TupleStructRecord {
array: TupleStruct,
}
#[allow(dead_code)]
#[derive(Debug, serde::Deserialize)]
struct TupleStruct(u32, u32, u32);
let reader =
Reader::open_readfile("test-data/test-data/MaxMind-DB-test-decoder.mmdb").unwrap();
let lookup = reader.lookup("1.1.1.0".parse().unwrap()).unwrap();
let tuple = lookup.decode::<TupleRecord>().unwrap().unwrap();
assert_eq!(tuple.array, (1, 2, 3));
let tuple_struct = lookup.decode::<TupleStructRecord>().unwrap().unwrap();
assert_eq!(tuple_struct.array.0, 1);
assert_eq!(tuple_struct.array.1, 2);
assert_eq!(tuple_struct.array.2, 3);
}
#[test]
fn test_decoder_rejects_tuple_length_mismatch() {
#[allow(dead_code)]
#[derive(Debug, serde::Deserialize)]
struct TupleRecord {
array: (u32, u32),
}
#[allow(dead_code)]
#[derive(Debug, serde::Deserialize)]
struct TupleStructRecord {
array: TupleStruct,
}
#[allow(dead_code)]
#[derive(Debug, serde::Deserialize)]
struct TupleStruct(u32, u32);
let reader =
Reader::open_readfile("test-data/test-data/MaxMind-DB-test-decoder.mmdb").unwrap();
let lookup = reader.lookup("1.1.1.0".parse().unwrap()).unwrap();
let tuple_err = lookup.decode::<TupleRecord>().unwrap_err();
assert!(tuple_err
.to_string()
.contains("expected tuple of length 2, got array of length 3"));
let tuple_struct_err = lookup.decode::<TupleStructRecord>().unwrap_err();
assert!(tuple_struct_err
.to_string()
.contains("expected tuple of length 2, got array of length 3"));
}
#[test]
fn test_skip_value_for_verification_rejects_truncated_pointer_payload() {
let mut decoder = Decoder::new(&[0x28], 0);
let err = decoder
.skip_value_for_verification(&mut VerificationState::new(decoder.limit))
.unwrap_err();
assert!(matches!(err, MaxMindDbError::InvalidDatabase { .. }));
}
#[cfg(not(feature = "unsafe-str-decode"))]
#[test]
fn ascii_check_covers_every_byte_at_word_boundaries() {
for len in 0..=80 {
let mut storage = vec![0x7f; len + 7];
for offset in 0..8 {
let bytes = &mut storage[offset..offset + len];
assert!(super::is_ascii(bytes));
for index in 0..len {
for byte in 0x80..=0xff {
bytes[index] = byte;
assert!(
!super::is_ascii(bytes),
"accepted non-ASCII byte {byte} at index {index}, length {len}, offset {offset}"
);
}
bytes[index] = 0x7f;
}
}
}
}
#[test]
fn ignored_any_skips_pointer_targets_but_verification_follows_them() {
for (pointer_size, control) in [(1, 0x20), (2, 0x28), (3, 0x30), (4, 0x38)] {
let mut encoded = vec![control];
encoded.resize(pointer_size + 1, 0xff);
let continuation = encoded.len();
encoded.extend([0xa1, 42]); let mut decoder = Decoder::new(&encoded, 0);
serde::de::IgnoredAny::deserialize(&mut decoder).unwrap();
assert_eq!(decoder.offset(), continuation);
assert_eq!(u16::deserialize(&mut decoder).unwrap(), 42);
assert_eq!(decoder.offset(), encoded.len());
let mut decoder = Decoder::new(&encoded, 0);
let err = decoder
.skip_value_for_verification(&mut VerificationState::new(decoder.limit))
.unwrap_err();
assert!(matches!(err, MaxMindDbError::InvalidDatabase { .. }));
for limit in 1..continuation {
for buf in [&encoded[..limit], &encoded[..]] {
let mut decoder = Decoder::new_with_limit(buf, 0, limit);
let err = serde::de::IgnoredAny::deserialize(&mut decoder).unwrap_err();
assert!(matches!(
err,
MaxMindDbError::InvalidDatabase { message, offset: Some(1) }
if message == format!("pointer of size {pointer_size}")
));
assert_eq!(decoder.offset(), 1);
}
}
}
}
#[test]
fn test_decoder_caps_impossible_container_size_hint() {
let mut decoder = Decoder::new(&[0x1d, 0x04, 0xff], 0);
let err = Vec::<serde::de::IgnoredAny>::deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::InvalidDatabase { .. }));
assert!(err.to_string().contains("unexpected end of buffer"));
}
#[test]
fn ignored_any_rejects_excessive_inline_container_work() {
let mut decoder = Decoder::new(&[0x1e, 0x04, 0xfe, 0xe3], 0);
let err = serde::de::IgnoredAny::deserialize(&mut decoder).unwrap_err();
assert!(err
.to_string()
.contains("maximum number of data structure values"));
}
#[test]
fn ignored_any_rejects_excessive_inline_map_work() {
let mut decoder = Decoder::new(&[0xfe, 0x7e, 0xe3], 0);
let err = serde::de::IgnoredAny::deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert!(err
.to_string()
.contains("maximum number of data structure values"));
}
#[test]
fn navigation_rejects_excessive_declared_container_work() {
let mut decoder = Decoder::new(&[0x1e, 0x04, 0xfe, 0xe3], 0);
let err = decoder.consume_container_header().unwrap_err();
assert!(err
.to_string()
.contains("maximum number of data structure values"));
}
#[test]
fn oversized_containers_fail_before_visitor_entry() {
use std::cell::Cell;
struct EntryVisitor<'a>(&'a Cell<bool>);
impl<'de> Visitor<'de> for EntryVisitor<'_> {
type Value = ();
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("a container")
}
fn visit_seq<A: SeqAccess<'de>>(self, _: A) -> Result<(), A::Error> {
self.0.set(true);
Ok(())
}
fn visit_map<A: MapAccess<'de>>(self, _: A) -> Result<(), A::Error> {
self.0.set(true);
Ok(())
}
}
let mut array = vec![0x1e, 0x04, 0xfe, 0xe3]; array.extend([0x00, 0x07].repeat(65_536)); let mut map = vec![0xfe, 0x7e, 0xe3]; map.extend([0x40, 0x00, 0x07].repeat(32_768));
for (is_map, bytes) in [(false, array), (true, map)] {
for dynamic in [false, true] {
let entered = Cell::new(false);
let mut decoder = Decoder::new(&bytes, 0);
let visitor = EntryVisitor(&entered);
let result = if dynamic {
decoder.deserialize_any(visitor)
} else if is_map {
decoder.deserialize_map(visitor)
} else {
decoder.deserialize_seq(visitor)
};
assert!(matches!(result, Err(MaxMindDbError::ResourceLimit { .. })));
assert!(
!entered.get(),
"visitor entered: map={is_map}, dynamic={dynamic}"
);
}
}
}
#[test]
fn city_subdivisions_reject_declared_size_before_allocating() {
let mut decoder = Decoder::new(&[0x1e, 0x04, 0xfe, 0xe3], 0);
let err =
Vec::<crate::geoip2::city::Subdivision<'_>>::deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert!(err
.to_string()
.contains("maximum number of data structure values"));
}
#[test]
fn concrete_map_rejects_declared_size_before_allocating() {
let mut decoder = Decoder::new(&[0xfe, 0x7e, 0xe3], 0);
let err =
std::collections::HashMap::<String, serde::de::IgnoredAny>::deserialize(&mut decoder)
.unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert!(err
.to_string()
.contains("maximum number of data structure values"));
}
#[test]
fn concrete_sequence_charges_repeated_shared_struct_maps() {
#[derive(Debug, Deserialize)]
struct EmptyRecord {}
let mut buf = vec![0xfd, 171]; for _ in 0..200 {
buf.extend_from_slice(&[0x40, 0xe0]); }
let array_offset = buf.len();
buf.extend_from_slice(&[0x1e, 0x04, 0x00, 0x0f]); for _ in 0..300 {
append_pointer(&mut buf, 0);
}
let mut decoder = Decoder::new(&buf, array_offset);
let err = Vec::<EmptyRecord>::deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert!(err
.to_string()
.contains("maximum number of data structure values"));
}
#[test]
fn concrete_sequence_charges_repeated_string_payloads() {
let (buf, array_offset) = repeated_4k_payload_array(0x5e, 600);
let mut decoder = Decoder::new(&buf, array_offset);
let err = Vec::<&str>::deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert!(err
.to_string()
.contains("maximum size of data structure string and bytes"));
}
#[test]
fn direct_byte_buf_decoding_charges_repeated_payloads() {
let (buf, array_offset) = repeated_4k_payload_array(0x9e, 600);
let mut decoder = Decoder::new(&buf, array_offset);
let err = Vec::<OwnedBytes>::deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert!(err
.to_string()
.contains("maximum size of data structure string and bytes"));
}
#[test]
fn raw_string_decoding_charges_repeated_payloads() {
let (buf, array_offset) = repeated_4k_payload_array(0x5e, 600);
let mut decoder = Decoder::new(&buf, array_offset);
let err = RawValueSeed.deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert!(err
.to_string()
.contains("maximum size of data structure string and bytes"));
}
#[allow(dead_code)]
#[derive(Debug, Deserialize)]
enum StandaloneScalarEnum {
Known,
}
fn large_string(size: usize) -> Vec<u8> {
assert!((65_821..=65_821 + 0xff_ff_ff).contains(&size));
let encoded_size = (size - 65_821) as u32;
let mut buf = vec![0x5f]; buf.extend_from_slice(&encoded_size.to_be_bytes()[1..]);
buf.resize(buf.len() + size, b'a');
buf
}
#[test]
fn partial_struct_skips_unknown_inline_payload_over_budget() {
#[derive(Deserialize)]
struct PartialRecord {
known: bool,
}
let payload_size = super::MAXIMUM_DATA_STRUCTURE_BYTES + 1;
for control in [0x5f, 0x9f] {
let mut payload = large_string(payload_size);
payload[0] = control;
let mut buf = vec![0xe2, 0x47]; buf.extend_from_slice(b"unknown");
buf.extend_from_slice(&payload);
buf.extend_from_slice(&[0x45]); buf.extend_from_slice(b"known");
buf.extend_from_slice(&[0x01, 0x07]);
let mut decoder = Decoder::new(&buf, 0);
let decoded = PartialRecord::deserialize(&mut decoder).unwrap();
assert!(decoded.known);
}
}
#[test]
fn budgeted_standalone_scalar_entry_points_enforce_payload_limit() {
let size =
super::MAXIMUM_DATA_STRUCTURE_BYTES + super::MAXIMUM_UNCHARGED_IDENTIFIER_BYTES + 1;
let buf = large_string(size);
let mut decoder = Decoder::new(&buf, 0);
let err = serde_json::Value::deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
let mut decoder = Decoder::new(&buf, 0);
let err = RawValueSeed.deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
let mut decoder = Decoder::new(&buf, 0);
let err = StandaloneScalarEnum::deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
let mut decoder = Decoder::new(&buf, 0);
let decoded = <&str>::deserialize(&mut decoder).unwrap();
assert_eq!(decoded.len(), size);
}
#[test]
fn standalone_scalar_retains_unbudgeted_fast_path() {
let size = super::MAXIMUM_DATA_STRUCTURE_BYTES + 1;
let buf = large_string(size);
let mut decoder = Decoder::new(&buf, 0);
let decoded = <&str>::deserialize(&mut decoder).unwrap();
assert_eq!(decoded.len(), size);
}
fn repeated_4k_payload_array(control: u8, element_count: usize) -> (Vec<u8>, usize) {
assert!(matches!(control, 0x5e | 0x9e));
let mut buf = vec![control, 0x0e, 0xe3]; buf.resize(buf.len() + 4096, b'a');
let array_offset = buf.len();
assert!((285..=u16::MAX as usize + 285).contains(&element_count));
buf.extend_from_slice(&[0x1e, 0x04]); buf.extend_from_slice(&((element_count - 285) as u16).to_be_bytes());
for _ in 0..element_count {
append_pointer(&mut buf, 0);
}
(buf, array_offset)
}
fn pointer_key_map(entry_count: usize) -> (Vec<u8>, usize) {
let mut buf = vec![0x5e, 0x0e, 0xe3]; buf.resize(buf.len() + 4096, b'k');
let map_offset = buf.len();
assert!((285..=u16::MAX as usize + 285).contains(&entry_count));
buf.push(0xfe); buf.extend_from_slice(&((entry_count - 285) as u16).to_be_bytes());
for _ in 0..entry_count {
append_pointer(&mut buf, 0);
buf.extend_from_slice(&[0x00, 0x07]); }
(buf, map_offset)
}
fn repeated_inline_key_map_array(element_count: usize) -> (Vec<u8>, usize) {
let mut buf = vec![0xe1, 0x5e, 0x0e, 0xe3]; buf.resize(buf.len() + 4096, b'k');
buf.extend_from_slice(&[0x00, 0x07]); let array_offset = buf.len();
assert!((285..=u16::MAX as usize + 285).contains(&element_count));
buf.extend_from_slice(&[0x1e, 0x04]); buf.extend_from_slice(&((element_count - 285) as u16).to_be_bytes());
for _ in 0..element_count {
append_pointer(&mut buf, 0);
}
(buf, array_offset)
}
#[test]
fn dynamic_map_keys_are_charged_once() {
let (buf, map_offset) = pointer_key_map(512);
let mut decoder = Decoder::new(&buf, map_offset);
let value = RawValueSeed.deserialize(&mut decoder).unwrap();
let RawValue::Map(entries) = value else {
panic!("expected map");
};
assert_eq!(entries.len(), 512);
let (buf, map_offset) = pointer_key_map(513);
let mut decoder = Decoder::new(&buf, map_offset);
let err = RawValueSeed.deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert!(err
.to_string()
.contains("maximum size of data structure string and bytes"));
}
#[test]
fn dynamic_any_map_keys_are_charged_once() {
let (buf, map_offset) = pointer_key_map(512);
let mut decoder = Decoder::new(&buf, map_offset);
assert_eq!(
decoder.deserialize_any(AnyKeyMapVisitor).unwrap(),
AnyKeyMap(512)
);
let (buf, map_offset) = pointer_key_map(513);
let mut decoder = Decoder::new(&buf, map_offset);
let err = decoder.deserialize_any(AnyKeyMapVisitor).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert!(err
.to_string()
.contains("maximum size of data structure string and bytes"));
}
#[test]
fn navigation_charges_repeated_pointer_keys() {
let (buf, map_offset) = pointer_key_map(513);
let mut decoder = Decoder::new(&buf, map_offset);
let (size, type_num) = decoder.consume_container_header().unwrap();
assert_eq!((size, type_num), (513, super::TYPE_MAP));
for _ in 0..512 {
assert_eq!(decoder.read_str_as_bytes().unwrap().len(), 4096);
decoder.skip_value().unwrap();
}
let err = decoder.read_str_as_bytes().unwrap_err();
assert!(err
.to_string()
.contains("maximum size of data structure string and bytes"));
}
#[test]
fn flattened_pointer_keys_cannot_bypass_payload_budget() {
#[derive(Debug, Deserialize)]
struct Flattened {
#[serde(flatten)]
fields: std::collections::HashMap<String, serde::de::IgnoredAny>,
}
let (buf, map_offset) = pointer_key_map(516);
let mut decoder = Decoder::new(&buf, map_offset);
let decoded = Flattened::deserialize(&mut decoder).unwrap();
assert_eq!(decoded.fields.len(), 1);
let (buf, map_offset) = pointer_key_map(517);
let mut decoder = Decoder::new(&buf, map_offset);
let err = Flattened::deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert!(err
.to_string()
.contains("maximum size of data structure string and bytes"));
}
#[test]
fn repeated_maps_with_inline_keys_cannot_bypass_payload_budget() {
#[derive(Debug, Deserialize)]
struct Flattened {
#[serde(flatten)]
fields: std::collections::HashMap<String, serde::de::IgnoredAny>,
}
let (buf, array_offset) = repeated_inline_key_map_array(516);
let mut decoder = Decoder::new(&buf, array_offset);
let decoded = Vec::<Flattened>::deserialize(&mut decoder).unwrap();
assert_eq!(decoded.len(), 516);
assert!(decoded.iter().all(|value| value.fields.len() == 1));
let (buf, array_offset) = repeated_inline_key_map_array(517);
let mut decoder = Decoder::new(&buf, array_offset);
let err = Vec::<Flattened>::deserialize(&mut decoder).unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert!(err
.to_string()
.contains("maximum size of data structure string and bytes"));
}
#[allow(dead_code)]
#[derive(Debug, Deserialize)]
enum RecursiveEnum {
Next(Box<RecursiveEnum>),
End,
}
fn recursive_enum_chain(in_array: bool) -> (Vec<u8>, usize) {
let mut buf = vec![0x44, b'N', b'e', b'x', b't'];
let start = buf.len();
if in_array {
buf.extend_from_slice(&[0x01, 0x04]); }
for _ in 0..=super::MAXIMUM_DATA_STRUCTURE_DEPTH {
append_pointer(&mut buf, 0);
}
buf.extend_from_slice(&[0x43, b'E', b'n', b'd']);
(buf, start)
}
#[test]
fn recursive_newtype_enum_is_depth_bounded() {
let (buf, start) = recursive_enum_chain(false);
let mut decoder = Decoder::new(&buf, start);
let err = RecursiveEnum::deserialize(&mut decoder).unwrap_err();
assert!(err
.to_string()
.contains("exceeded maximum data structure depth"));
}
#[test]
fn recursive_newtype_enum_inside_array_is_depth_bounded() {
let (buf, start) = recursive_enum_chain(true);
let mut decoder = Decoder::new(&buf, start);
let err = Vec::<RecursiveEnum>::deserialize(&mut decoder).unwrap_err();
assert!(err
.to_string()
.contains("exceeded maximum data structure depth"));
}
#[test]
fn verification_bounds_overlapping_string_scans() {
let size = 65_821 + 65_536;
let mut buf = [0x5f, 0x01, 0x00, 0x00].repeat(64);
buf.resize(buf.len() + size, b'a');
let mut state = VerificationState::new(buf.len());
let allowance = state.work_remaining;
for index in 0..8 {
Decoder::new(&buf, index * 4)
.skip_value_for_verification(&mut state)
.unwrap();
}
assert_eq!(state.work_remaining, allowance - 8 * (size + 1));
let mut decoder = Decoder::new(&buf, 8 * 4);
let err = decoder.skip_value_for_verification(&mut state).unwrap_err();
assert!(matches!(
err,
MaxMindDbError::ResourceLimit {
offset: Some(36),
..
}
));
assert_eq!(state.work_remaining, allowance - 8 * (size + 1) - 1);
assert_eq!(decoder.offset(), 36);
assert_eq!(state.validated.len(), 8);
assert!(state.active.is_empty());
}
#[test]
fn verification_bounds_repeated_inline_traversal_across_roots() {
let mut buf = [0x01, 0x04].repeat(64); buf.extend_from_slice(&[0x00, 0x07]); let mut state = VerificationState::new(buf.len());
let err = (0..64)
.try_for_each(|index| {
Decoder::new(&buf, index * 2).skip_value_for_verification(&mut state)
})
.unwrap_err();
assert!(matches!(err, MaxMindDbError::ResourceLimit { .. }));
assert_eq!(state.work_remaining, 0);
assert!(state.active.is_empty());
}
#[test]
fn verification_allows_large_nonoverlapping_payloads() {
let size = super::MAXIMUM_DATA_STRUCTURE_BYTES + 1;
let mut buf = vec![0x5f];
buf.extend_from_slice(&((size - 65_821) as u32).to_be_bytes()[1..]);
buf.resize(buf.len() + size, b'a');
let mut state = VerificationState::new(buf.len());
Decoder::new(&buf, 0)
.skip_value_for_verification(&mut state)
.unwrap();
}
#[test]
fn verification_work_allowance_does_not_overflow() {
let mut state = VerificationState::new(usize::MAX);
assert_eq!(state.work_remaining, usize::MAX);
state.charge(usize::MAX, 0).unwrap();
assert!(matches!(
state.charge(1, 0),
Err(MaxMindDbError::ResourceLimit { .. })
));
assert_eq!(state.work_remaining, 0);
}
#[test]
fn test_verification_rejects_invalid_bool_size() {
let mut decoder = Decoder::new(&[0x02, 0x07], 0);
let err = decoder
.skip_value_for_verification(&mut VerificationState::new(decoder.limit))
.unwrap_err();
assert!(matches!(err, MaxMindDbError::InvalidDatabase { .. }));
}
#[test]
fn test_verification_rejects_and_does_not_cache_invalid_utf8() {
let buf = [0x41, 0xff];
let mut state = VerificationState::new(buf.len());
for _ in 0..2 {
let mut decoder = Decoder::new(&buf, 0);
let err = decoder.skip_value_for_verification(&mut state).unwrap_err();
assert!(matches!(err, MaxMindDbError::InvalidDatabase { .. }));
assert!(err.to_string().contains("invalid UTF-8"));
assert!(state.validated.is_empty());
assert!(state.active.is_empty());
}
#[cfg(not(feature = "unsafe-str-decode"))]
{
let mut decoder = Decoder::new(&buf, 0);
let err = String::deserialize(&mut decoder).unwrap_err();
assert!(err.to_string().contains("invalid UTF-8"));
}
}
fn append_pointer(buf: &mut Vec<u8>, target: usize) {
assert!(target < 2048);
buf.push(0x20 | ((target >> 8) as u8));
buf.push(target as u8);
}
#[test]
fn test_verification_caches_shared_pointer_targets() {
let mut buf = vec![0x00, 0x07];
let mut target = 0;
const LEVELS: usize = 20;
for _ in 0..LEVELS {
let array = buf.len();
buf.extend_from_slice(&[0x02, 0x04]);
append_pointer(&mut buf, target);
append_pointer(&mut buf, target);
target = array;
}
let mut decoder = Decoder::new(&buf, target);
let mut state = VerificationState::new(buf.len());
decoder.skip_value_for_verification(&mut state).unwrap();
assert_eq!(state.validated.len(), LEVELS + 1);
assert!(state.active.is_empty());
}
#[test]
fn test_verification_rejects_data_pointer_cycles() {
let mut buf = vec![0x01, 0x04];
append_pointer(&mut buf, 4);
buf.extend_from_slice(&[0x01, 0x04]);
append_pointer(&mut buf, 0);
let mut decoder = Decoder::new(&buf, 0);
let err = decoder
.skip_value_for_verification(&mut VerificationState::new(decoder.limit))
.unwrap_err();
assert!(matches!(err, MaxMindDbError::InvalidDatabase { .. }));
assert!(err.to_string().contains("cyclic data pointer"));
}
}