use super::{DecodeBudget, Decoder, MAX_CONTAINER_ITEMS, MAX_DEPTH, utf8};
use crate::{Error, Result, ValueRef};
const INLINE: usize = usize::MAX;
const T_UTF8: u8 = 2;
const T_DOUBLE: u8 = 3;
const T_BYTES: u8 = 4;
const T_U16: u8 = 5;
const T_U32: u8 = 6;
const T_MAP: u8 = 7;
const T_I32: u8 = 8;
const T_U64: u8 = 9;
const T_U128: u8 = 10;
const T_ARRAY: u8 = 11;
const T_BOOL: u8 = 14;
const T_FLOAT: u8 = 15;
#[doc(hidden)]
pub struct RawDecoder<'a> {
dec: Decoder<'a>,
pos: usize,
depth: usize,
budget: DecodeBudget,
}
#[doc(hidden)]
#[derive(Clone, Copy, Debug)]
#[must_use]
pub struct Container {
len: usize,
resume: usize,
}
impl Container {
#[inline(always)]
pub fn len(self) -> usize {
self.len
}
#[inline(always)]
pub fn is_empty(self) -> bool {
self.len == 0
}
}
#[cold]
#[inline(never)]
fn mismatch(message: &'static str) -> Error {
Error::DecodingError(message.into())
}
#[cold]
#[inline(never)]
fn eof() -> Error {
Error::UnexpectedEof
}
#[cold]
#[inline(never)]
fn too_deep() -> Error {
Error::ResourceLimit("maximum MMDB nesting depth exceeded")
}
impl<'a> RawDecoder<'a> {
#[inline]
pub(crate) fn new(data: &'a [u8], pointer_base: usize, limit: usize, offset: usize) -> Self {
Self {
dec: Decoder::new(data, pointer_base, limit),
pos: offset,
depth: 0,
budget: DecodeBudget::default(),
}
}
#[inline(always)]
fn byte(&self, pos: usize) -> Result<u8> {
match self.dec.data.get(pos) {
Some(&byte) => Ok(byte),
None => Err(eof()),
}
}
#[inline(always)]
fn header(&mut self) -> Result<(u8, usize, usize)> {
self.budget.charge_value()?;
let mut pos = self.pos;
let mut control = self.byte(pos)?;
let mut resume = INLINE;
if control >> 5 == 1 {
(pos, control, resume) = self.follow(pos, control)?;
}
pos += 1;
let mut kind = control >> 5;
if kind == 0 {
let ext = self.byte(pos)?;
pos += 1;
if ext > 248 {
return Err(Error::InvalidDataType(ext));
}
kind = ext + 7;
}
let (size, extra) = self.dec.decode_size(control & 0x1f, pos)?;
self.pos = pos + extra;
Ok((kind, size, resume))
}
#[inline(never)]
fn follow(&mut self, pos: usize, control: u8) -> Result<(usize, u8, usize)> {
let (mut pointer, consumed) = self.dec.decode_pointer(control, pos + 1)?;
let resume = pos + 1 + consumed;
let mut hops = 0_usize;
loop {
hops += 1;
if self.depth + hops > MAX_DEPTH {
return Err(too_deep());
}
self.budget.charge_value()?;
let target = self
.dec
.pointer_base
.checked_add(pointer)
.ok_or(Error::InvalidOffset(pointer))?;
let control = self.byte(target)?;
if control >> 5 != 1 {
return Ok((target, control, resume));
}
pointer = self.dec.decode_pointer(control, target + 1)?.0;
}
}
#[inline(always)]
fn payload(&mut self, size: usize, resume: usize) -> Result<&'a [u8]> {
let start = self.pos;
if !self.dec.in_range(start, size) {
return Err(eof());
}
self.pos = if resume == INLINE {
start + size
} else {
resume
};
Ok(self.dec.slice(start, size))
}
#[inline(always)]
fn done(&mut self, resume: usize) {
if resume != INLINE {
self.pos = resume;
}
}
#[inline(always)]
fn uint_payload(&mut self, size: usize, resume: usize) -> Result<u64> {
debug_assert!(size <= 8, "callers validate integer widths");
let start = self.pos;
let data = self.dec.data;
let value = if start <= data.len() && data.len() - start >= 8 {
let word = unsafe { core::ptr::read_unaligned(data.as_ptr().add(start).cast::<u64>()) };
u64::from_be(word)
.checked_shr(64 - 8 * size as u32)
.unwrap_or(0)
} else {
if !self.dec.in_range(start, size) {
return Err(eof());
}
self.dec
.slice(start, size)
.iter()
.fold(0_u64, |acc, &b| (acc << 8) | u64::from(b))
};
self.pos = if resume == INLINE {
start + size
} else {
resume
};
Ok(value)
}
#[inline]
pub fn read_str(&mut self) -> Result<&'a str> {
let (kind, size, resume) = self.header()?;
if kind != T_UTF8 {
return Err(mismatch("expected UTF-8 string"));
}
self.budget.charge_bytes(size)?;
utf8(self.payload(size, resume)?)
}
#[inline]
pub fn read_bytes(&mut self) -> Result<&'a [u8]> {
let (kind, size, resume) = self.header()?;
if kind != T_BYTES {
return Err(mismatch("expected byte array"));
}
self.budget.charge_bytes(size)?;
self.payload(size, resume)
}
#[inline]
pub fn read_u64(&mut self) -> Result<u64> {
let (kind, size, resume) = self.header()?;
let max = match kind {
T_U16 => 2,
T_U32 => 4,
T_U64 => 8,
_ => return Err(mismatch("expected numeric value")),
};
if size > max {
return Err(mismatch("integer payload is too large"));
}
self.uint_payload(size, resume)
}
#[inline]
pub fn read_u128(&mut self) -> Result<u128> {
let (kind, size, resume) = self.header()?;
let max = match kind {
T_U16 => 2,
T_U32 => 4,
T_U64 => 8,
T_U128 => 16,
_ => return Err(mismatch("expected unsigned integer")),
};
if size > max {
return Err(mismatch("integer payload is too large"));
}
if size <= 8 {
return self.uint_payload(size, resume).map(u128::from);
}
let bytes = self.payload(size, resume)?;
Ok(bytes
.iter()
.fold(0_u128, |acc, &b| (acc << 8) | u128::from(b)))
}
#[inline]
pub fn read_i32(&mut self) -> Result<i32> {
let (kind, size, resume) = self.header()?;
if kind != T_I32 {
return Err(mismatch("expected int32"));
}
if size > 4 {
return Err(mismatch("int32 is longer than 4 bytes"));
}
let raw = self.uint_payload(size, resume)? as u32;
Ok(raw as i32)
}
#[inline]
pub fn read_f64(&mut self) -> Result<f64> {
let (kind, size, resume) = self.header()?;
match kind {
T_DOUBLE => {
if size != 8 {
return Err(mismatch("double must contain 8 bytes"));
}
let bytes = self.payload(8, resume)?;
Ok(f64::from_be_bytes(bytes.try_into().map_err(|_| eof())?))
}
T_FLOAT => {
if size != 4 {
return Err(mismatch("float must contain 4 bytes"));
}
let bytes = self.payload(4, resume)?;
Ok(f64::from(f32::from_be_bytes(
bytes.try_into().map_err(|_| eof())?,
)))
}
_ => Err(mismatch("expected float/double")),
}
}
#[inline]
pub fn read_f32(&mut self) -> Result<f32> {
let (kind, size, resume) = self.header()?;
if kind != T_FLOAT {
return Err(mismatch("expected float"));
}
if size != 4 {
return Err(mismatch("float must contain 4 bytes"));
}
let bytes = self.payload(4, resume)?;
Ok(f32::from_be_bytes(bytes.try_into().map_err(|_| eof())?))
}
#[inline]
pub fn read_bool(&mut self) -> Result<bool> {
let (kind, size, resume) = self.header()?;
if kind != T_BOOL {
return Err(mismatch("expected boolean"));
}
if size > 1 {
return Err(mismatch("boolean size must be zero or one"));
}
self.done(resume);
Ok(size == 1)
}
#[inline(always)]
fn enter(&mut self, expected: u8, message: &'static str) -> Result<Container> {
let (kind, size, resume) = self.header()?;
if kind != expected {
return Err(mismatch(message));
}
if size > MAX_CONTAINER_ITEMS {
return Err(Error::ResourceLimit("MMDB container item limit exceeded"));
}
if self.depth >= MAX_DEPTH {
return Err(too_deep());
}
self.depth += 1;
Ok(Container { len: size, resume })
}
#[inline]
pub fn enter_map(&mut self, message: &'static str) -> Result<Container> {
self.enter(T_MAP, message)
}
#[inline]
pub fn enter_array(&mut self, message: &'static str) -> Result<Container> {
self.enter(T_ARRAY, message)
}
#[inline]
pub fn capacity_hint(&self, container: Container) -> usize {
let remaining = self.dec.data.len().saturating_sub(self.pos);
container.len.min(remaining).min(1024)
}
#[inline(always)]
pub fn leave(&mut self, container: Container) {
self.depth -= 1;
if container.resume != INLINE {
self.pos = container.resume;
}
}
#[inline]
pub fn finish_map(&mut self, container: Container, remaining: usize) -> Result<()> {
self.depth -= 1;
if container.resume != INLINE {
self.pos = container.resume;
Ok(())
} else if remaining == 0 || self.depth == 0 {
Ok(())
} else {
self.skip_values(remaining.saturating_mul(2))
}
}
#[inline]
pub fn read_key(&mut self) -> Result<&'a [u8]> {
let (kind, size, resume) = self.header()?;
if kind != T_UTF8 {
return Err(mismatch("MMDB map key is not UTF-8"));
}
self.budget.charge_bytes(size)?;
self.payload(size, resume)
}
#[inline]
pub fn read_key_str(&mut self) -> Result<&'a str> {
utf8(self.read_key()?)
}
#[inline]
pub fn skip_value(&mut self) -> Result<()> {
self.skip_values(1)
}
#[inline(never)]
fn skip_values(&mut self, mut pending: usize) -> Result<()> {
let mut pos = self.pos;
while pending != 0 {
pending -= 1;
self.budget.charge_value()?;
let control = self.byte(pos)?;
pos += 1;
let mut kind = control >> 5;
if kind == 1 {
pos += 1 + usize::from((control >> 3) & 0x03);
continue;
}
if kind == 0 {
let ext = self.byte(pos)?;
pos += 1;
if ext > 248 {
return Err(Error::InvalidDataType(ext));
}
kind = ext + 7;
}
let (size, extra) = self.dec.decode_size(control & 0x1f, pos)?;
pos += extra;
let max_payload = match kind {
T_MAP | T_ARRAY => {
if size > MAX_CONTAINER_ITEMS {
return Err(Error::ResourceLimit("MMDB container item limit exceeded"));
}
let children = if kind == T_MAP { size * 2 } else { size };
pending = pending.checked_add(children).ok_or_else(eof)?;
continue;
}
T_UTF8 | T_BYTES => usize::MAX,
T_DOUBLE if size == 8 => 8,
T_FLOAT if size == 4 => 4,
T_U16 => 2,
T_U32 | T_I32 => 4,
T_U64 => 8,
T_U128 => 16,
T_BOOL if size <= 1 => {
continue;
}
T_DOUBLE | T_FLOAT | T_BOOL => {
return Err(mismatch("invalid scalar size"));
}
other => return Err(Error::InvalidDataType(other)),
};
if size > max_payload {
return Err(mismatch("integer payload is too large"));
}
pos += size;
if pos > self.dec.data.len() {
return Err(eof());
}
}
if pos > self.dec.data.len() {
return Err(eof());
}
self.pos = pos;
Ok(())
}
pub fn read_value(&mut self) -> Result<ValueRef<'a>> {
let (value, next) = self
.dec
.decode_inner(self.pos, self.depth, &mut self.budget)?;
self.pos = next;
Ok(value)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(feature = "writer")]
use crate::{Value, encoder::encode_value};
#[cfg(feature = "writer")]
use std::collections::BTreeMap;
fn decoder(data: &[u8]) -> RawDecoder<'_> {
RawDecoder::new(data, 0, data.len(), 0)
}
#[cfg(feature = "writer")]
fn encoded(value: Value) -> Vec<u8> {
let mut data = Vec::new();
encode_value(&value, &mut data).unwrap();
data
}
#[test]
fn scalars_inline_and_through_pointers() {
let data = [0x42, b'a', b'b', 0xc2, 0x01, 0x02, 0x20, 0x00, 0x20, 0x03];
let mut d = decoder(&data);
assert_eq!(d.read_str().unwrap(), "ab");
assert_eq!(d.read_u64().unwrap(), 0x0102);
assert_eq!(d.read_str().unwrap(), "ab");
assert_eq!(d.pos, 8);
assert_eq!(d.read_u64().unwrap(), 0x0102);
assert_eq!(d.pos, 10);
}
#[test]
fn integer_payload_widths_match_generic_decoder() {
for size in 0..=8_usize {
let mut data = vec![size as u8, 2];
data.extend((0..size).map(|i| 0x10 + i as u8));
let expected = match Decoder::new(&data, 0, data.len()).decode_at(0).unwrap().0 {
ValueRef::Uint64(v) => v,
other => panic!("unexpected {other:?}"),
};
assert_eq!(decoder(&data).read_u64().unwrap(), expected, "size {size}");
let mut padded = data.clone();
padded.extend_from_slice(&[0xff; 8]);
assert_eq!(
decoder(&padded).read_u64().unwrap(),
expected,
"size {size}"
);
}
}
#[test]
fn skip_crosses_nested_containers_without_following_pointers() {
let data = [
0xe1, 0x41, b'k', 0x03, 0x04, 0xa1, 0x01, 0x20, 0x63, 0xe1, 0x41, b'x', 0x01, 0x07,
0x41, b'z',
];
let mut d = decoder(&data);
d.skip_value().unwrap();
assert_eq!(d.read_str().unwrap(), "z");
}
#[test]
fn truncation_and_bad_types_are_errors() {
assert!(decoder(&[0x42, b'a']).read_str().is_err());
assert!(decoder(&[0x42, b'a']).skip_value().is_err());
assert!(decoder(&[0xe2, 0x41]).skip_value().is_err());
assert!(matches!(
decoder(&[0x00, 0x06]).skip_value(),
Err(Error::InvalidDataType(13))
));
assert!(decoder(&[0x41, b'a']).read_u64().is_err());
}
#[test]
fn pointer_cycles_hit_the_depth_limit() {
let data = [0x20, 0x00];
assert!(matches!(
decoder(&data).read_str(),
Err(Error::ResourceLimit(_))
));
}
#[test]
fn early_finished_inline_map_skips_remaining_entries_when_nested() {
let data = [
0x02, 0x04, 0xe2, 0x41, b'a', 0xa1, 0x01, 0x41, b'b', 0xa1, 0x02, 0x44, b't', b'a',
b'i', b'l',
];
let mut d = decoder(&data);
let array = d.enter_array("expected array").unwrap();
assert_eq!(array.len(), 2);
let map = d.enter_map("expected map").unwrap();
assert_eq!(d.read_key().unwrap(), b"a");
assert_eq!(d.read_u64().unwrap(), 1);
d.finish_map(map, 1).unwrap();
assert_eq!(d.read_str().unwrap(), "tail");
d.leave(array);
}
#[cfg(feature = "writer")]
#[test]
fn typed_scalars_match_the_generic_decoder() {
let bytes = encoded(Value::Bytes(vec![0, 127, 255]));
assert_eq!(decoder(&bytes).read_bytes().unwrap(), &[0, 127, 255]);
for (value, expected) in [
(Value::Uint16(50), 50_u64),
(Value::Uint32(65_000), 65_000),
(Value::Uint64(1 << 40), 1 << 40),
] {
assert_eq!(decoder(&encoded(value)).read_u64().unwrap(), expected);
}
for (value, expected) in [
(Value::Uint16(50), 50_u128),
(Value::Uint64(1 << 40), 1 << 40),
(Value::Uint128(1 << 100), 1 << 100),
] {
assert_eq!(decoder(&encoded(value)).read_u128().unwrap(), expected);
}
for n in [-123_i32, 0, 123] {
assert_eq!(decoder(&encoded(Value::Int32(n))).read_i32().unwrap(), n);
}
assert_eq!(
decoder(&encoded(Value::Double(1.25))).read_f64().unwrap(),
1.25
);
assert_eq!(
decoder(&encoded(Value::Float(1.25))).read_f64().unwrap(),
1.25
);
assert_eq!(
decoder(&encoded(Value::Float(1.25))).read_f32().unwrap(),
1.25
);
for value in [false, true] {
assert_eq!(
decoder(&encoded(Value::Bool(value))).read_bool().unwrap(),
value
);
}
}
#[cfg(feature = "writer")]
#[test]
fn containers_keys_and_materialization_follow_cursor_rules() {
let map = Value::Map(BTreeMap::from([("name".into(), Value::Utf8("ok".into()))]));
let bytes = encoded(map.clone());
let mut d = decoder(&bytes);
let container = d.enter_map("expected map").unwrap();
assert_eq!(container.len(), 1);
assert!(!container.is_empty());
assert_eq!(d.capacity_hint(container), 1);
assert_eq!(d.read_key_str().unwrap(), "name");
assert_eq!(d.read_str().unwrap(), "ok");
d.leave(container);
assert_eq!(decoder(&bytes).read_value().unwrap().to_owned_value(), map);
let bytes = encoded(Value::Array(vec![]));
let mut d = decoder(&bytes);
let container = d.enter_array("expected array").unwrap();
assert!(container.is_empty());
assert_eq!(d.capacity_hint(container), 0);
d.finish_map(container, 0).unwrap();
let data = [0x20, 0x04, 0x41, b'z', 0x01, 0x04, 0xa1, 1];
let mut d = decoder(&data);
let container = d.enter_array("expected array").unwrap();
assert_eq!(d.read_u64().unwrap(), 1);
d.leave(container);
assert_eq!(d.read_str().unwrap(), "z");
let mut d = decoder(&data);
let container = d.enter_array("expected array").unwrap();
d.finish_map(container, 1).unwrap();
assert_eq!(d.read_str().unwrap(), "z");
}
#[cfg(feature = "writer")]
#[test]
fn typed_reads_reject_mismatches_and_invalid_widths() {
let text = encoded(Value::Utf8("x".into()));
assert!(decoder(&text).read_bytes().is_err());
assert!(decoder(&text).read_u128().is_err());
assert!(decoder(&text).read_i32().is_err());
assert!(decoder(&text).read_f64().is_err());
assert!(decoder(&text).read_f32().is_err());
assert!(decoder(&text).read_bool().is_err());
assert!(decoder(&text).enter_map("expected map").is_err());
assert!(decoder(&text).enter_array("expected array").is_err());
assert!(decoder(&text).read_key().is_ok());
assert!(decoder(&[0x41, 0xff]).read_key_str().is_err());
assert!(decoder(&[]).read_str().is_err());
assert!(decoder(&[0x00, 249]).read_str().is_err());
for bytes in [
vec![0xa3, 0, 0, 1], vec![0xc5, 0, 0, 0, 0, 1], vec![0x09, 2], ] {
assert!(decoder(&bytes).read_u64().is_err());
assert!(decoder(&bytes).read_u128().is_err());
}
assert!(decoder(&[0x11, 3]).read_u128().is_err()); assert!(decoder(&[0x05, 1]).read_i32().is_err()); assert!(decoder(&[0x07, 8]).read_f64().is_err()); assert!(decoder(&[0x05, 8]).read_f64().is_err()); assert!(decoder(&[0x03, 8]).read_f32().is_err()); assert!(decoder(&[0x02, 7]).read_bool().is_err()); assert!(decoder(&[0x43, b'a']).read_str().is_err()); }
#[cfg(feature = "writer")]
#[test]
fn skip_validates_types_widths_and_boundaries() {
for value in [
Value::Utf8("ok".into()),
Value::Bytes(vec![1, 2]),
Value::Double(1.0),
Value::Float(1.0),
Value::Uint16(1),
Value::Uint32(1),
Value::Int32(-1),
Value::Uint64(1),
Value::Uint128(1),
Value::Bool(false),
Value::Bool(true),
Value::Array(vec![Value::Uint16(1)]),
Value::Map(BTreeMap::from([("k".into(), Value::Bool(true))])),
] {
let bytes = encoded(value);
let mut d = decoder(&bytes);
d.skip_value().unwrap();
assert_eq!(d.pos, bytes.len());
}
for bytes in [
vec![0x68], vec![0x03, 8], vec![0x02, 7], vec![0xa3], vec![0x00, 249], vec![0x20], vec![0x42, b'a'], ] {
assert!(decoder(&bytes).skip_value().is_err(), "{bytes:?}");
}
}
#[cfg(feature = "writer")]
#[test]
fn pointer_booleans_restore_cursor_and_malformed_keys_fail() {
let data = [0x20, 0x04, 0x41, b'z', 0x01, 0x07];
let mut d = decoder(&data);
assert!(d.read_bool().unwrap());
assert_eq!(d.read_str().unwrap(), "z");
assert!(decoder(&[0xa1, 1]).read_key().is_err());
assert!(decoder(&[0x42, b'a']).read_value().is_err());
assert!(decoder(&[0x20]).read_value().is_err());
assert!(decoder(&[0x20, 0x00]).skip_value().is_ok());
assert!(decoder(&[0x2f]).skip_value().is_err());
let bytes = encoded(Value::Array(vec![]));
let mut d = decoder(&bytes);
d.depth = MAX_DEPTH;
assert!(matches!(
d.enter_array("expected array"),
Err(Error::ResourceLimit(_))
));
let chain = [0x20, 0x04, 0x41, b'z', 0x20, 0x06, 0x41, b'a'];
let mut d = decoder(&chain);
assert_eq!(d.read_str().unwrap(), "a");
assert_eq!(d.read_str().unwrap(), "z");
assert!(matches!(
decoder(&[0x20, 0x63]).read_str(),
Err(Error::UnexpectedEof)
));
}
}