#![cfg_attr(not(feature = "reader"), allow(dead_code))]
mod ascii;
mod raw;
pub use raw::{Container, RawDecoder};
use crate::{Error, Result, ValueRef};
const MAX_DEPTH: usize = 512;
const MAX_CONTAINER_ITEMS: usize = 16_843_036;
const MAX_DECODED_VALUES: usize = 1_000_000;
const MAX_EXPANDED_BYTES: usize = 64 * 1024 * 1024;
#[derive(Default)]
struct DecodeBudget {
values: usize,
expanded_bytes: usize,
}
impl DecodeBudget {
#[inline(always)]
fn charge_value(&mut self) -> Result<()> {
self.values += 1;
if self.values > MAX_DECODED_VALUES {
return Err(Error::ResourceLimit("MMDB decoded-value budget exceeded"));
}
Ok(())
}
#[inline(always)]
fn charge_bytes(&mut self, bytes: usize) -> Result<()> {
self.expanded_bytes += bytes;
if self.expanded_bytes > MAX_EXPANDED_BYTES {
return Err(Error::ResourceLimit("MMDB expanded-byte budget exceeded"));
}
Ok(())
}
}
pub(crate) struct Decoder<'a> {
data: &'a [u8],
pointer_base: usize,
}
macro_rules! leaf_decode {
($self:expr, $control:expr, $pos:expr, $budget:expr) => {{
let control: u8 = $control;
let data_type = control >> 5;
match data_type {
2 | 3 | 4 | 5 | 6 | 8 | 9 | 10 | 14 | 15 => {
$budget.charge_value()?;
let (size, size_bytes) = $self.decode_size(control & 0x1f, $pos)?;
let (value, next) =
$self.decode_scalar(data_type, size, $pos + size_bytes, $budget)?;
Ok(Some((value, next)))
}
1 => {
if let Ok((pointer, consumed)) = $self.decode_pointer(control, $pos) {
if let Some(target) = $self.pointer_base.checked_add(pointer) {
if let Some(&target_control) = $self.data.get(target) {
let target_type = target_control >> 5;
match target_type {
2 | 3 | 4 | 5 | 6 | 8 | 9 | 10 | 14 | 15 => {
$budget.charge_value()?;
$budget.charge_value()?;
let (size, size_bytes) =
$self.decode_size(target_control & 0x1f, target + 1)?;
let (value, _) = $self.decode_scalar(
target_type,
size,
target + 1 + size_bytes,
$budget,
)?;
let next = $pos + consumed;
Ok(Some((value, next)))
}
_ => Ok(None),
}
} else {
Ok(None)
}
} else {
Ok(None)
}
} else {
Ok(None)
}
}
_ => Ok::<Option<(ValueRef<'a>, usize)>, Error>(None),
}
}};
}
impl<'a> Decoder<'a> {
pub(crate) fn new(data: &'a [u8], pointer_base: usize, limit: usize) -> Self {
Self {
data: &data[..limit.min(data.len())],
pointer_base,
}
}
pub(crate) fn decode_at(&self, offset: usize) -> Result<(ValueRef<'a>, usize)> {
let mut budget = DecodeBudget::default();
self.decode_inner(offset, 0, &mut budget)
}
#[allow(clippy::unnecessary_lazy_evaluations)]
fn decode_inner(
&self,
mut offset: usize,
depth: usize,
budget: &mut DecodeBudget,
) -> Result<(ValueRef<'a>, usize)> {
if depth > MAX_DEPTH {
return Err(Error::ResourceLimit("maximum MMDB nesting depth exceeded"));
}
budget.charge_value()?;
if !self.in_range(offset, 1) {
return Err(Error::UnexpectedEof);
}
let control = unsafe { *self.data.get_unchecked(offset) };
offset += 1;
let mut data_type = control >> 5;
if data_type == 1 {
let (pointer, consumed) = self.decode_pointer(control, offset)?;
let target = self
.pointer_base
.checked_add(pointer)
.ok_or_else(|| Error::InvalidOffset(pointer))?;
let (value, _) = self.decode_inner(target, depth + 1, budget)?;
return Ok((value, offset + consumed));
}
if data_type == 0 {
if !self.in_range(offset, 1) {
return Err(Error::UnexpectedEof);
}
let ext = unsafe { *self.data.get_unchecked(offset) };
offset += 1;
if ext > 248 {
return Err(Error::InvalidDataType(ext));
}
data_type = ext + 7;
}
let (size, size_bytes) = self.decode_size(control & 0x1f, offset)?;
offset += size_bytes;
if size > MAX_CONTAINER_ITEMS && matches!(data_type, 7 | 11) {
return Err(Error::ResourceLimit("MMDB container item limit exceeded"));
}
let (value, next) = match data_type {
7 => self.decode_container(true, size, offset, depth, budget),
11 => self.decode_container(false, size, offset, depth, budget),
_ => self.decode_scalar(data_type, size, offset, budget),
}?;
Ok((value, next))
}
#[inline(never)]
#[allow(clippy::unnecessary_lazy_evaluations)]
fn decode_container(
&self,
is_map: bool,
size: usize,
offset: usize,
depth: usize,
budget: &mut DecodeBudget,
) -> Result<(ValueRef<'a>, usize)> {
if is_map {
let mut entries = Vec::with_capacity(size.min(256));
let mut cursor = offset;
let leaf_depth_ok = depth < MAX_DEPTH;
#[cfg(all(target_arch = "x86_64", feature = "simd"))]
{
if size > 2 && cursor < self.data.len() {
unsafe {
core::arch::x86_64::_mm_prefetch(
self.data
.as_ptr()
.wrapping_add(cursor.saturating_add(64))
.cast(),
core::arch::x86_64::_MM_HINT_T0,
);
}
}
}
for _ in 0..size {
let (key, next) = self.decode_key(cursor, depth + 1, budget)?;
cursor = next;
if leaf_depth_ok
&& self.in_range(cursor, 1)
&& let &control = unsafe {self.data.get_unchecked(cursor)}
&& let Some((value, next)) =
leaf_decode!(self, control, cursor + 1, budget)?
{
cursor = next;
entries.push((key, value));
continue;
}
let (value, next) = self.decode_inner(cursor, depth + 1, budget)?;
cursor = next;
entries.push((key, value));
}
Ok((ValueRef::Map(entries), cursor))
} else {
let mut values = Vec::with_capacity(size.min(256));
let mut cursor = offset;
let leaf_depth_ok = depth < MAX_DEPTH;
#[cfg(all(target_arch = "x86_64", feature = "simd"))]
{
if size > 2 && cursor < self.data.len() {
unsafe {
core::arch::x86_64::_mm_prefetch(
self.data
.as_ptr()
.wrapping_add(cursor.saturating_add(64))
.cast(),
core::arch::x86_64::_MM_HINT_T0,
);
}
}
}
for _ in 0..size {
if leaf_depth_ok
&& self.in_range(cursor, 1)
&& let &control = unsafe {self.data.get_unchecked(cursor)}
&& let Some((value, next)) =
leaf_decode!(self, control, cursor + 1, budget)?
{
cursor = next;
values.push(value);
continue;
}
let (value, next) = self.decode_inner(cursor, depth + 1, budget)?;
values.push(value);
cursor = next;
}
Ok((ValueRef::Array(values), cursor))
}
}
#[allow(clippy::collapsible_if)]
#[inline(always)]
fn decode_key(
&self,
offset: usize,
depth: usize,
budget: &mut DecodeBudget,
) -> Result<(&'a str, usize)> {
if depth > MAX_DEPTH {
return Err(Error::ResourceLimit("maximum MMDB nesting depth exceeded"));
}
if let Some(&control) = self.data.get(offset) {
let data_type = control >> 5;
if data_type == 2 {
budget.charge_value()?;
let offset = offset + 1;
let (size, extra) = self.decode_size(control & 0x1f, offset)?;
let offset = offset + extra;
budget.charge_bytes(size)?;
if !self.in_range(offset, size) {
return Err(Error::UnexpectedEof);
}
return Ok((utf8(self.slice(offset, size))?, offset + size));
} else if data_type == 1
&& let Ok((pointer, consumed)) = self.decode_pointer(control, offset + 1)
&& let Some(target) = self.pointer_base.checked_add(pointer)
&& let Some(&target_control) = self.data.get(target)
&& target_control >> 5 == 2
{
budget.charge_value()?;
if depth + 1 > MAX_DEPTH {
return Err(Error::ResourceLimit("maximum MMDB nesting depth exceeded"));
}
budget.charge_value()?;
let t_offset = target + 1;
let (size, extra) = self.decode_size(target_control & 0x1f, t_offset)?;
let t_offset = t_offset + extra;
budget.charge_bytes(size)?;
if !self.in_range(t_offset, size) {
return Err(Error::UnexpectedEof);
}
return Ok((utf8(self.slice(t_offset, size))?, offset + 1 + consumed));
}
}
let (key, next) = self.decode_inner(offset, depth, budget)?;
match key {
ValueRef::Utf8(key) => Ok((key, next)),
_ => Err(Error::DecodingError("MMDB map key is not UTF-8".into())),
}
}
#[cfg_attr(not(debug_assertions), inline(always))]
fn decode_scalar(
&self,
data_type: u8,
size: usize,
offset: usize,
budget: &mut DecodeBudget,
) -> Result<(ValueRef<'a>, usize)> {
match data_type {
2 => {
budget.charge_bytes(size)?;
if !self.in_range(offset, size) {
return Err(Error::UnexpectedEof);
}
let bytes = self.slice(offset, size);
let text = utf8(bytes)?;
Ok((ValueRef::Utf8(text), offset + size))
}
3 => {
if size != 8 {
return Err(Error::DecodingError("double must contain 8 bytes".into()));
}
if !self.in_range(offset, 8) {
return Err(Error::UnexpectedEof);
}
let bytes: [u8; 8] = self.slice(offset, 8).try_into().expect("length checked");
Ok((ValueRef::Double(f64::from_be_bytes(bytes)), offset + 8))
}
4 => {
budget.charge_bytes(size)?;
if !self.in_range(offset, size) {
return Err(Error::UnexpectedEof);
}
Ok((ValueRef::Bytes(self.slice(offset, size)), offset + size))
}
5 => {
let v = self.read_uint(offset, size, 2)? as u16;
Ok((ValueRef::Uint16(v), offset + size))
}
6 => {
let v = self.read_uint(offset, size, 4)? as u32;
Ok((ValueRef::Uint32(v), offset + size))
}
8 => {
if size > 4 {
return Err(Error::DecodingError("int32 is longer than 4 bytes".into()));
}
let raw = self.read_uint(offset, size, 4)? as u32;
let value = if size == 4 {
i32::from_be_bytes(raw.to_be_bytes())
} else {
raw as i32
};
Ok((ValueRef::Int32(value), offset + size))
}
9 => {
let v = self.read_uint(offset, size, 8)? as u64;
Ok((ValueRef::Uint64(v), offset + size))
}
10 => {
let v = self.read_uint(offset, size, 16)?;
Ok((ValueRef::Uint128(v), offset + size))
}
13 => Err(Error::InvalidDataType(13)),
14 => {
if size > 1 {
return Err(Error::DecodingError(
"boolean size must be zero or one".into(),
));
}
Ok((ValueRef::Bool(size == 1), offset))
}
15 => {
if size != 4 {
return Err(Error::DecodingError("float must contain 4 bytes".into()));
}
if !self.in_range(offset, 4) {
return Err(Error::UnexpectedEof);
}
let bytes: [u8; 4] = self.slice(offset, 4).try_into().expect("length checked");
Ok((ValueRef::Float(f32::from_be_bytes(bytes)), offset + 4))
}
other => Err(Error::InvalidDataType(other)),
}
}
#[inline(always)]
fn decode_size(&self, size: u8, offset: usize) -> Result<(usize, usize)> {
if size <= 28 {
Ok((usize::from(size), 0))
} else {
match size {
29 => {
if !self.in_range(offset, 1) {
return Err(Error::UnexpectedEof);
}
Ok((
29 + usize::from(unsafe { *self.data.get_unchecked(offset) }),
1,
))
}
30 => {
if !self.in_range(offset, 2) {
return Err(Error::UnexpectedEof);
}
let bytes: [u8; 2] = self.slice(offset, 2).try_into().expect("length checked");
Ok((285 + usize::from(u16::from_be_bytes(bytes)), 2))
}
31 => {
if !self.in_range(offset, 3) {
return Err(Error::UnexpectedEof);
}
let b = self.slice(offset, 3);
let n = (usize::from(unsafe { *b.get_unchecked(0) }) << 16)
| (usize::from(b[1]) << 8)
| usize::from(b[2]);
Ok((65_821 + n, 3))
}
_ => unreachable!(),
}
}
}
#[allow(clippy::unnecessary_lazy_evaluations)]
#[inline(always)]
fn decode_pointer(&self, control: u8, offset: usize) -> Result<(usize, usize)> {
let selector = (control >> 3) & 0x03;
let high = usize::from(control & 0x07);
match selector {
0 => {
if !self.in_range(offset, 1) {
return Err(Error::UnexpectedEof);
}
let low = unsafe { *self.data.get_unchecked(offset) };
Ok(((high << 8) | usize::from(low), 1))
}
1 => {
if !self.in_range(offset, 2) {
return Err(Error::UnexpectedEof);
}
let b = self.slice(offset, 2);
let byte0 = unsafe { *b.get_unchecked(0) };
let byte1 = unsafe { *b.get_unchecked(1) };
Ok((
((high << 16) | (usize::from(byte0) << 8) | usize::from(byte1)) + 2_048,
2,
))
}
2 => {
if !self.in_range(offset, 3) {
return Err(Error::UnexpectedEof);
}
let b = self.slice(offset, 3);
let byte0 = unsafe { *b.get_unchecked(0) };
let byte1 = unsafe { *b.get_unchecked(1) };
let byte2 = unsafe { *b.get_unchecked(2) };
let raw = (high << 24)
| (usize::from(byte0) << 16)
| (usize::from(byte1) << 8)
| usize::from(byte2);
Ok((raw + 526_336, 3))
}
3 => {
if !self.in_range(offset, 4) {
return Err(Error::UnexpectedEof);
}
let bytes = self.slice(offset, 4);
Ok((u32::from_be_bytes(bytes.try_into().unwrap()) as usize, 4))
}
_ => unreachable!(),
}
}
#[allow(clippy::unnecessary_lazy_evaluations)]
#[inline]
fn slice(&self, offset: usize, len: usize) -> &'a [u8] {
let end = offset + len;
unsafe { self.data.get_unchecked(offset..end) }
}
#[inline(always)]
fn in_range(&self, offset: usize, len: usize) -> bool {
offset <= self.data.len() && len <= self.data.len() - offset
}
#[cfg_attr(not(debug_assertions), inline(always))]
fn read_uint(&self, offset: usize, len: usize, max: usize) -> Result<u128> {
if len > max {
return Err(Error::DecodingError("integer payload is too large".into()));
}
if !self.in_range(offset, len) {
return Err(Error::UnexpectedEof);
}
unsafe {
let ptr = self.data.as_ptr().add(offset);
match len {
0 => Ok(0),
1 => Ok(u128::from(*ptr)),
2 => Ok(u128::from(u16::from_be(core::ptr::read_unaligned(
ptr.cast::<u16>(),
)))),
3 => {
let a =
u128::from(u16::from_be(core::ptr::read_unaligned(ptr.cast::<u16>()))) << 8;
let b = u128::from(*ptr.add(2));
Ok(a | b)
}
4 => Ok(u128::from(u32::from_be(core::ptr::read_unaligned(
ptr.cast::<u32>(),
)))),
5 => {
let a =
u128::from(u32::from_be(core::ptr::read_unaligned(ptr.cast::<u32>()))) << 8;
let b = u128::from(*ptr.add(4));
Ok(a | b)
}
6 => {
let a = u128::from(u32::from_be(core::ptr::read_unaligned(ptr.cast::<u32>())))
<< 16;
let b = u128::from(u16::from_be(core::ptr::read_unaligned(
ptr.add(4).cast::<u16>(),
)));
Ok(a | b)
}
7 => {
let a = u128::from(u32::from_be(core::ptr::read_unaligned(ptr.cast::<u32>())))
<< 24;
let b = u128::from(u16::from_be(core::ptr::read_unaligned(
ptr.add(4).cast::<u16>(),
))) << 8;
let c = u128::from(*ptr.add(6));
Ok(a | b | c)
}
8 => Ok(u128::from(u64::from_be(core::ptr::read_unaligned(
ptr.cast::<u64>(),
)))),
9..=16 => {
let mut out = 0_u128;
for i in 0..len {
out = (out << 8) | u128::from(*ptr.add(i));
}
Ok(out)
}
_ => {
let mut out = 0_u128;
for i in 0..len {
out = (out << 8) | u128::from(*ptr.add(i));
}
Ok(out)
}
}
}
}
}
#[inline]
fn utf8(bytes: &[u8]) -> Result<&str> {
if bytes.is_empty() {
return Ok("");
}
if ascii::is_ascii(bytes) {
Ok(unsafe { std::str::from_utf8_unchecked(bytes) })
} else {
utf8_validated(bytes)
}
}
#[cold]
#[inline(never)]
fn utf8_validated(bytes: &[u8]) -> Result<&str> {
std::str::from_utf8(bytes).map_err(|_| Error::DecodingError("invalid UTF-8 string".into()))
}
#[cfg(test)]
mod tests {
use super::*;
fn probe(data: &[u8], start: usize) {
let decoder = Decoder::new(data, 0, data.len());
let _ = decoder.decode_at(start);
}
fn probe_deep(data: &[u8], start: usize) {
let owned = data.to_vec();
std::thread::Builder::new()
.name("decoder-depth".into())
.stack_size(64 * 1024 * 1024)
.spawn(move || probe(&owned, start))
.unwrap()
.join()
.unwrap();
}
#[test]
fn pointers_containers_cycles_and_truncations_are_bounded() {
let data = [
0x41, b'k', 0x43, b'v', b'a', b'l', 0xe1, 0x20, 0, 3, 4, 0x20, 2, 0xe0, 0, 4,
];
probe(&data, 6);
for end in 0..data.len() {
probe(&data[..end], 6);
}
let mut nested = [1, 4].repeat(MAX_DEPTH);
nested.push(0x40);
probe_deep(&nested, 0);
nested.splice(0..0, [1, 4]);
probe_deep(&nested, 0);
for data in [
&[0x20, 0][..],
&[0xe1, 0x20, 0][..],
&[0xe1, 0xa0, 0x40][..],
&[0x5f, 255, 255, 255][..],
&[0, 255][..],
] {
probe_deep(data, 0);
}
}
#[test]
fn deterministic_corruption_replay_does_not_panic() {
let mut state = 0xcafe_f00d_u64;
for len in 0..128 {
for _ in 0..64 {
let mut data = vec![0; len];
for byte in &mut data {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
*byte = state as u8;
}
probe(&data, 0);
}
}
}
#[test]
fn decoding_keeps_depth_value_and_byte_budgets() {
let mut dag = vec![0x40];
let mut previous = 0;
for _ in 0..20 {
let offset = dag.len();
dag.extend_from_slice(&[2, 4, 0x20, previous as u8, 0x20, previous as u8]);
previous = offset;
}
let decoder = Decoder::new(&dag, 0, dag.len());
assert!(matches!(
decoder.decode_at(previous),
Err(Error::ResourceLimit(_))
));
let mut budget = DecodeBudget {
values: MAX_DECODED_VALUES,
expanded_bytes: 0,
};
assert!(decoder.decode_inner(0, 0, &mut budget).is_err());
let mut budget = DecodeBudget {
values: 0,
expanded_bytes: MAX_EXPANDED_BYTES,
};
let decoder = Decoder::new(&[0x41, b'x'], 0, 2);
assert!(matches!(
decoder.decode_inner(0, 0, &mut budget),
Err(Error::ResourceLimit(_))
));
assert!(matches!(
decoder.decode_inner(0, MAX_DEPTH + 1, &mut DecodeBudget::default()),
Err(Error::ResourceLimit(_))
));
}
#[test]
fn fast_keys_preserve_general_decoder_budgets_and_errors() {
for data in [
&[0x41, b'x'][..],
&[0x20, 2, 0x41, b'x'][..],
&[0x41, 255][..],
&[0x5d, 0][..],
&[0x20, 0][..],
&[0x20, 255][..],
&[][..],
] {
for depth in [0, MAX_DEPTH, MAX_DEPTH + 1] {
for values in [0, MAX_DECODED_VALUES] {
for expanded_bytes in [0, MAX_EXPANDED_BYTES] {
let decoder = Decoder::new(data, 0, data.len());
let mut a = DecodeBudget {
values,
expanded_bytes,
};
let mut b = DecodeBudget {
values,
expanded_bytes,
};
let expected = decoder
.decode_inner(0, depth, &mut a)
.map(|(value, next)| {
let ValueRef::Utf8(text) = value else {
panic!("string fixture")
};
(text, next)
})
.map_err(|e| e.to_string());
let actual = decoder
.decode_key(0, depth, &mut b)
.map_err(|e| e.to_string());
assert_eq!(actual, expected);
assert_eq!((a.values, a.expanded_bytes), (b.values, b.expanded_bytes));
}
}
}
}
}
#[test]
fn truncated_scalar_payloads_fail_before_any_slice_is_borrowed() {
for bytes in [
&[0x68, 0x01][..], &[0x83, 0x01, 0x02][..], &[0x04, 0x08, 0x01][..], &[0x42, b'a'][..], ] {
assert!(matches!(
Decoder::new(bytes, 0, bytes.len()).decode_at(0),
Err(Error::UnexpectedEof)
));
}
let invalid_key = [0xe1, 0xa1, 1, 0x41, b'x'];
assert!(matches!(
Decoder::new(&invalid_key, 0, invalid_key.len()).decode_at(0),
Err(Error::DecodingError(_))
));
let pointer_key = [0xe1, 0x20, 0x05, 0x41, b'v', 0x41, b'k'];
let (value, _) = Decoder::new(&pointer_key, 0, pointer_key.len())
.decode_at(0)
.unwrap();
assert_eq!(value, ValueRef::Map(vec![("k", ValueRef::Utf8("v"))]));
for broken in [
[0xe1, 0x20, 0x05, 0x41, b'v', 0x42, b'k'],
[0xe1, 0x20, 0x05, 0x41, b'v', 0x5d, 0],
] {
assert!(matches!(
Decoder::new(&broken, 0, broken.len()).decode_at(0),
Err(Error::UnexpectedEof)
));
}
}
#[test]
fn utf8_validation_matches_standard_library() {
for prefix in [0, 7, 15, 31, 127, 255] {
for byte in 0..=255 {
let mut bytes = vec![b'x'; prefix];
bytes.push(byte);
assert_eq!(utf8(&bytes).ok(), std::str::from_utf8(&bytes).ok());
}
for text in [
"Paris 東京 café 🦀",
"\u{7ff}\u{800}\u{ffff}\u{10000}\u{10ffff}",
] {
let mut bytes = vec![b'x'; prefix];
bytes.extend_from_slice(text.as_bytes());
for end in prefix..=bytes.len() {
assert_eq!(
utf8(&bytes[..end]).ok(),
std::str::from_utf8(&bytes[..end]).ok()
);
}
}
}
}
}