mod marker;
mod tree;
use memmap2::Mmap;
use serde::de::DeserializeOwned;
use std::net::IpAddr;
use std::path::Path;
use std::sync::OnceLock;
use crate::decoder::{Decoder, RawDecoder};
use crate::{Error, Metadata, MmdbDecode, Result, ValueRef};
use marker::{METADATA_MARKER, find_metadata_marker};
use tree::PreparedTree;
#[cold]
fn compute_ipv4_start_node(bytes: &[u8], node_count: u64, record_size: u16) -> Result<Option<u64>> {
let mut node = 0_u64;
match record_size {
28 => {
for _ in 0..96 {
if node >= node_count {
return Ok(None);
}
let offset = node as usize * 7;
if offset + 7 > bytes.len() {
return Err(Error::UnexpectedEof);
}
let p = unsafe { bytes.as_ptr().add(offset) };
node = unsafe {
(u64::from(*p.add(3) >> 4) << 24)
| (u64::from(*p) << 16)
| (u64::from(*p.add(1)) << 8)
| u64::from(*p.add(2))
};
if node >= node_count {
return Ok(None);
}
}
Ok(Some(node))
}
32 => {
for _ in 0..96 {
if node >= node_count {
return Ok(None);
}
let offset = node as usize * 8;
if offset + 8 > bytes.len() {
return Err(Error::UnexpectedEof);
}
let p = unsafe { bytes.as_ptr().add(offset) };
node = u64::from(u32::from_be(unsafe {
core::ptr::read_unaligned(p.cast::<u32>())
}));
if node >= node_count {
return Ok(None);
}
}
Ok(Some(node))
}
24 => {
for _ in 0..96 {
if node >= node_count {
return Ok(None);
}
let offset = node as usize * 6;
if offset + 6 > bytes.len() {
return Err(Error::UnexpectedEof);
}
let p = unsafe { bytes.as_ptr().add(offset) };
node = unsafe {
(u64::from(*p) << 16) | (u64::from(*p.add(1)) << 8) | u64::from(*p.add(2))
};
if node >= node_count {
return Ok(None);
}
}
Ok(Some(node))
}
_ => {
for _ in 0..96 {
if node >= node_count {
return Ok(None);
}
node = match read_record_static(bytes, node as usize, record_size)? {
Some(next) if next < node_count => next,
_ => return Ok(None),
};
}
Ok(Some(node))
}
}
}
#[cold]
fn read_record_static(bytes: &[u8], node: usize, record_size: u16) -> Result<Option<u64>> {
let node_size = usize::from(record_size) / 4;
let offset = node
.checked_mul(node_size)
.ok_or(Error::InvalidNode(node as u64))?;
let slice = bytes
.get(offset..offset + node_size)
.ok_or(Error::UnexpectedEof)?;
Ok(Some(read_packed_record_static(
slice,
usize::from(record_size),
0,
)?))
}
#[cold]
fn read_packed_record_static(bytes: &[u8], bits: usize, side: usize) -> Result<u64> {
let start = side * bits;
let mut value = 0_u64;
for bit_index in start..start + bits {
let byte = *bytes.get(bit_index / 8).ok_or(Error::UnexpectedEof)?;
let bit = (byte >> (7 - (bit_index % 8))) & 1;
value = (value << 1) | u64::from(bit);
}
Ok(value)
}
#[derive(Debug)]
enum Source<'a> {
Borrowed(&'a [u8]),
Mmap(Mmap),
Owned(Vec<u8>),
}
impl Source<'_> {
#[inline(always)]
fn bytes(&self) -> &[u8] {
match self {
Self::Borrowed(v) => v,
Self::Mmap(v) => v,
Self::Owned(v) => v,
}
}
}
#[derive(Debug)]
pub struct Reader<'a> {
source: Source<'a>,
metadata: Metadata,
data_pointer_bias: usize,
data_section_start: usize,
metadata_marker: usize,
ipv4_start_node: Option<u64>,
prepared_tree: PreparedTree,
}
impl Reader<'static> {
pub fn open(path: impl AsRef<Path>) -> Result<Self> {
Self::from_vec(std::fs::read(path)?)
}
pub unsafe fn open_mmap(path: impl AsRef<Path>) -> Result<Self> {
let file = std::fs::File::open(path)?;
let mmap = unsafe { Mmap::map(&file)? };
Self::from_source(Source::Mmap(mmap))
}
pub fn from_vec(data: Vec<u8>) -> Result<Self> {
Self::from_source(Source::Owned(data))
}
}
impl<'a> Reader<'a> {
#[inline]
pub fn from_bytes(data: &'a [u8]) -> Result<Self> {
Self::from_source(Source::Borrowed(data))
}
fn from_source(source: Source<'a>) -> Result<Self> {
let bytes = source.bytes();
let marker = find_metadata_marker(bytes)?;
let metadata_start = marker + METADATA_MARKER.len();
let decoder = Decoder::new(bytes, metadata_start, bytes.len());
let (value, _) = decoder.decode_at(metadata_start)?;
let metadata = Metadata::from_value(&value)?;
if metadata.binary_format_major_version != 2 {
return Err(Error::InvalidMetadata(
"unsupported binary format major version",
));
}
if !matches!(metadata.ip_version, 4 | 6) {
return Err(Error::InvalidIpVersion(metadata.ip_version));
}
if metadata.record_size < 24 || metadata.record_size % 4 != 0 || metadata.record_size > 64 {
return Err(Error::InvalidMetadata(
"record_size must be a multiple of 4 between 24 and 64",
));
}
let node_bytes = usize::from(metadata.record_size) / 4;
let node_count_usize = usize::try_from(metadata.node_count)
.map_err(|_| Error::InvalidMetadata("node count exceeds address space"))?;
let search_tree_size = node_count_usize
.checked_mul(node_bytes)
.ok_or(Error::InvalidMetadata("search tree size overflow"))?;
let data_pointer_bias = search_tree_size
.checked_sub(node_count_usize)
.ok_or(Error::InvalidMetadata("invalid search tree geometry"))?;
let data_section_start = search_tree_size
.checked_add(16)
.ok_or(Error::InvalidMetadata("data section offset overflow"))?;
if data_section_start > marker || data_section_start > bytes.len() {
return Err(Error::InvalidDatabase(
"search tree overlaps metadata or exceeds file",
));
}
if bytes.get(search_tree_size..data_section_start) != Some(&[0_u8; 16][..]) {
return Err(Error::InvalidDatabase(
"missing 16-byte data section separator",
));
}
let record_size_u8 = u8::try_from(metadata.record_size)
.map_err(|_| Error::InvalidMetadata("record_size must fit in u8"))?;
let ipv4_start_node = if metadata.ip_version == 6 {
compute_ipv4_start_node(bytes, metadata.node_count, metadata.record_size)?
} else {
None
};
let prepared_tree = PreparedTree::build(
&bytes[..data_section_start],
record_size_u8,
metadata.node_count,
ipv4_start_node,
)?;
Ok(Self {
source,
metadata,
data_pointer_bias,
data_section_start,
metadata_marker: marker,
ipv4_start_node,
prepared_tree,
})
}
#[must_use]
pub const fn metadata(&self) -> &Metadata {
&self.metadata
}
#[must_use]
pub fn as_bytes(&self) -> &[u8] {
self.source.bytes()
}
#[inline]
pub fn lookup_value(&self, ip: IpAddr) -> Result<ValueRef<'_>> {
self.lookup_value_with_prefix(ip).map(|(v, _)| v)
}
#[inline]
pub fn lookup_value_with_prefix(&self, ip: IpAddr) -> Result<(ValueRef<'_>, u8)> {
let (offset, prefix) = self.resolve_offset(ip)?;
let decoder = Decoder::new(
self.source.bytes(),
self.data_section_start,
self.metadata_marker,
);
let (value, _) = decoder.decode_at(offset)?;
Ok((value, prefix))
}
#[inline(always)]
pub fn lookup_borrowed<'s, T>(&'s self, ip: IpAddr) -> Result<T>
where
T: MmdbDecode<'s>,
{
let (offset, _) = self.resolve_offset(ip)?;
decode_borrowed_at::<T>(self, offset)
}
#[inline(always)]
pub fn lookup_borrowed_opt<'s, T>(&'s self, ip: IpAddr) -> Option<T>
where
T: MmdbDecode<'s>,
{
let (offset, _) = self.resolve_offset_opt(ip)?;
decode_borrowed_at_opt::<T>(self, offset)
}
#[inline(always)]
pub fn lookup_borrowed_map<'s, T, R>(
&'s self,
ip: IpAddr,
on_hit: impl FnOnce(T) -> R,
) -> Result<Option<R>>
where
T: MmdbDecode<'s>,
{
let traversed = match (self.metadata.ip_version, ip) {
(4, IpAddr::V4(v4)) => self.traverse_ipv4(0, &v4.octets(), 0),
(4, IpAddr::V6(_)) => None,
(6, IpAddr::V6(v6)) => self.traverse_ipv6(0, &v6.octets(), 0),
(6, IpAddr::V4(v4)) => self
.ipv4_start_node
.and_then(|start| self.traverse_ipv4(start, &v4.octets(), 0)),
(v, _) => return Err(Error::InvalidIpVersion(v)),
};
let Some((record, _)) = traversed else {
return Ok(None);
};
let offset = self.record_to_file_offset(record)?;
decode_borrowed_map_at::<T, R>(self, offset, on_hit).map(Some)
}
pub fn lookup<T: DeserializeOwned>(&self, ip: IpAddr) -> Result<T> {
let value = self.lookup_value(ip)?;
Ok(serde_json::from_value(value.to_json())?)
}
pub fn lookup_many(&self, ips: &[IpAddr]) -> Vec<Result<ValueRef<'_>>> {
const PARALLEL_MIN: usize = 4_096;
static WORKERS: OnceLock<usize> = OnceLock::new();
let workers =
*WORKERS.get_or_init(|| std::thread::available_parallelism().map_or(1, |c| c.get()));
if workers <= 1 || ips.len() < PARALLEL_MIN {
return ips.iter().map(|&ip| self.lookup_value(ip)).collect();
}
let chunk = ips.len().div_ceil(workers.min(ips.len()));
let mut results: Vec<Result<ValueRef<'_>>> =
(0..ips.len()).map(|_| Err(Error::NotFound)).collect();
std::thread::scope(|scope| {
for (slice, slots) in ips.chunks(chunk).zip(results.chunks_mut(chunk)) {
scope.spawn(move || {
for (ip, slot) in slice.iter().zip(slots) {
*slot = self.lookup_value(*ip);
}
});
}
});
results
}
#[inline(always)]
#[allow(clippy::unnecessary_lazy_evaluations)]
fn resolve_offset(&self, ip: IpAddr) -> Result<(usize, u8)> {
let (record, prefix) = match (self.metadata.ip_version, ip) {
(4, IpAddr::V4(v4)) => self
.traverse_ipv4(0, &v4.octets(), 0)
.ok_or(Error::NotFound)?,
(4, IpAddr::V6(_)) => return Err(Error::NotFound),
(6, IpAddr::V6(v6)) => self
.traverse_ipv6(0, &v6.octets(), 0)
.ok_or(Error::NotFound)?,
(6, IpAddr::V4(v4)) => {
let start = self.ipv4_start_node.ok_or_else(|| Error::NotFound)?;
self.traverse_ipv4(start, &v4.octets(), 0)
.ok_or(Error::NotFound)?
}
(v, _) => return Err(Error::InvalidIpVersion(v)),
};
let offset = self.record_to_file_offset(record)?;
Ok((offset, prefix))
}
#[inline(always)]
fn resolve_offset_opt(&self, ip: IpAddr) -> Option<(usize, u8)> {
let (record, prefix) = match (self.metadata.ip_version, ip) {
(4, IpAddr::V4(v4)) => self.traverse_ipv4(0, &v4.octets(), 0)?,
(4, IpAddr::V6(_)) => return None,
(6, IpAddr::V6(v6)) => self.traverse_ipv6(0, &v6.octets(), 0)?,
(6, IpAddr::V4(v4)) => {
let start = self.ipv4_start_node?;
self.traverse_ipv4(start, &v4.octets(), 0)?
}
_ => return None,
};
let offset = self.record_to_file_offset_opt(record)?;
Some((offset, prefix))
}
#[inline]
pub fn lookup_exists(&self, ip: IpAddr) -> bool {
self.resolve_offset(ip).is_ok()
}
#[inline(always)]
fn traverse_ipv4(&self, node: u64, octets: &[u8; 4], prefix_base: u8) -> Option<(u64, u8)> {
self.prepared_tree.traverse_ipv4(node, octets, prefix_base)
}
#[inline(always)]
fn traverse_ipv6(&self, node: u64, octets: &[u8; 16], prefix_base: u8) -> Option<(u64, u8)> {
self.prepared_tree.traverse_ipv6(node, octets, prefix_base)
}
#[allow(dead_code, clippy::unnecessary_lazy_evaluations)]
#[cold]
fn read_record(&self, node: u64, side: usize) -> Result<u64> {
if node >= self.metadata.node_count || side > 1 {
return Err(Error::InvalidNode(node));
}
let node_size = usize::from(self.metadata.record_size) / 4;
let offset = usize::try_from(node)
.ok()
.and_then(|n| n.checked_mul(node_size))
.ok_or_else(|| Error::InvalidNode(node))?;
let bytes = self
.source
.bytes()
.get(offset..offset + node_size)
.ok_or_else(|| Error::UnexpectedEof)?;
match self.metadata.record_size {
24 => {
let base = side * 3;
Ok((u64::from(unsafe { *bytes.get_unchecked(base) }) << 16)
| (u64::from(unsafe { *bytes.get_unchecked(base + 1) }) << 8)
| u64::from(unsafe { *bytes.get_unchecked(base + 2) }))
}
28 => {
if side == 0 {
Ok((u64::from(unsafe { *bytes.get_unchecked(3) } >> 4) << 24)
| (u64::from(unsafe { *bytes.get_unchecked(0) }) << 16)
| (u64::from(unsafe { *bytes.get_unchecked(1) }) << 8)
| u64::from(unsafe { *bytes.get_unchecked(2) }))
} else {
Ok((u64::from(unsafe { *bytes.get_unchecked(3) } & 0x0f) << 24)
| (u64::from(unsafe { *bytes.get_unchecked(4) }) << 16)
| (u64::from(unsafe { *bytes.get_unchecked(5) }) << 8)
| u64::from(unsafe { *bytes.get_unchecked(6) }))
}
}
32 => {
let base = side * 4;
Ok(u64::from(u32::from_be_bytes(
unsafe { bytes.get_unchecked(base..base + 4) }
.try_into()
.expect("length checked"),
)))
}
bits => read_packed_record(bytes, usize::from(bits), side),
}
}
#[inline]
#[allow(clippy::unnecessary_lazy_evaluations)]
fn record_to_file_offset(&self, record: u64) -> Result<usize> {
let node_count = self.metadata.node_count;
if record < node_count.saturating_add(16) {
return Err(Error::InvalidOffset(record as usize));
}
let record = usize::try_from(record).map_err(|_| Error::InvalidOffset(usize::MAX))?;
let offset = record
.checked_add(self.data_pointer_bias)
.ok_or(Error::InvalidOffset(record))?;
if offset >= self.metadata_marker {
return Err(Error::InvalidOffset(offset));
}
Ok(offset)
}
#[inline(always)]
fn record_to_file_offset_opt(&self, record: u64) -> Option<usize> {
let node_count = self.metadata.node_count;
if record < node_count.saturating_add(16) {
return None;
}
let record = usize::try_from(record).ok()?;
let offset = record.checked_add(self.data_pointer_bias)?;
if offset >= self.metadata_marker {
return None;
}
Some(offset)
}
}
#[cold]
#[inline(never)]
fn decode_borrowed_at<'s, T>(reader: &'s Reader<'_>, offset: usize) -> Result<T>
where
T: MmdbDecode<'s>,
{
let mut decoder = RawDecoder::new(
reader.source.bytes(),
reader.data_section_start,
reader.metadata_marker,
offset,
);
T::decode_raw(&mut decoder)
}
#[cold]
#[inline(never)]
fn decode_borrowed_map_at<'s, T, R>(
reader: &'s Reader<'_>,
offset: usize,
on_hit: impl FnOnce(T) -> R,
) -> Result<R>
where
T: MmdbDecode<'s>,
{
let mut decoder = RawDecoder::new(
reader.source.bytes(),
reader.data_section_start,
reader.metadata_marker,
offset,
);
let value = T::decode_raw(&mut decoder)?;
Ok(on_hit(value))
}
#[cold]
#[inline(never)]
fn decode_borrowed_at_opt<'s, T>(reader: &'s Reader<'_>, offset: usize) -> Option<T>
where
T: MmdbDecode<'s>,
{
let mut decoder = RawDecoder::new(
reader.source.bytes(),
reader.data_section_start,
reader.metadata_marker,
offset,
);
T::decode_raw(&mut decoder).ok()
}
#[allow(clippy::unnecessary_lazy_evaluations)]
fn read_packed_record(bytes: &[u8], bits: usize, side: usize) -> Result<u64> {
let start = side * bits;
let mut value = 0_u64;
for bit_index in start..start + bits {
let byte = *bytes.get(bit_index / 8).ok_or(Error::UnexpectedEof)?;
let bit = (byte >> (7 - (bit_index % 8))) & 1;
value = (value << 1) | u64::from(bit);
}
Ok(value)
}
#[cfg(all(test, feature = "writer"))]
mod reader_tests {
use super::*;
use crate::writer::Writer;
use crate::{MetadataBuilder, Value};
fn raw_node(record_size: u8, left: u64, right: u64) -> Vec<u8> {
match record_size {
24 => {
let mut n = Vec::with_capacity(6);
n.extend_from_slice(&(left as u32).to_be_bytes()[1..=3]);
n.extend_from_slice(&(right as u32).to_be_bytes()[1..=3]);
n
}
28 => vec![
(left >> 16) as u8,
(left >> 8) as u8,
left as u8,
((((left >> 24) & 0x0f) << 4) | ((right >> 24) & 0x0f)) as u8,
(right >> 16) as u8,
(right >> 8) as u8,
right as u8,
],
32 => {
let mut n = Vec::with_capacity(8);
n.extend_from_slice(&(left as u32).to_be_bytes());
n.extend_from_slice(&(right as u32).to_be_bytes());
n
}
36..=64 if record_size.is_multiple_of(4) => {
let bits = usize::from(record_size);
let mut n = vec![0_u8; bits / 4];
for (side, value) in [left, right].into_iter().enumerate() {
for bit in 0..bits {
let stream_bit = side * bits + bit;
n[stream_bit / 8] |=
(((value >> (bits - bit - 1)) & 1) as u8) << (7 - stream_bit % 8);
}
}
n
}
_ => panic!("unexpected record size {record_size}"),
}
}
fn node_stream(record_size: u8, nodes: &[(u64, u64)]) -> Vec<u8> {
let mut out = Vec::new();
for &(l, r) in nodes {
out.extend_from_slice(&raw_node(record_size, l, r));
}
out
}
fn crafted_reader(
record_size: u16,
node_count: u64,
tree: Vec<u8>,
data: Vec<u8>,
ip_version: u16,
ipv4_start: Option<u64>,
) -> Reader<'static> {
let mut file = tree;
file.extend_from_slice(&[0_u8; 16]);
file.extend_from_slice(&data);
let search_tree_size = (node_count as usize) * (usize::from(record_size) / 4);
let data_section_start = search_tree_size + 16;
let prepared_tree = PreparedTree::build(&file, record_size as u8, node_count, ipv4_start)
.expect("crafted tree must be preparable");
Reader {
source: Source::Owned(file),
metadata: Metadata {
node_count,
record_size,
ip_version,
database_type: "test".into(),
languages: vec!["en".into()],
binary_format_major_version: 2,
binary_format_minor_version: 0,
build_epoch: 0,
description: Default::default(),
},
data_pointer_bias: search_tree_size - (node_count as usize),
data_section_start,
metadata_marker: data_section_start + data.len(),
ipv4_start_node: ipv4_start,
prepared_tree,
}
}
fn scalar_reader(record_size: u16) -> Reader<'static> {
let node_count = 1;
let data_pointer = node_count + 16;
crafted_reader(
record_size,
node_count,
node_stream(record_size as u8, &[(data_pointer, data_pointer)]),
vec![0x42, b'a', b'b'],
6,
Some(0),
)
}
fn ip(s: &str) -> IpAddr {
s.parse().unwrap()
}
#[test]
fn ipv4_subtree_start_handles_early_leaves_truncation_and_packed_fallback() {
for size in [24_u8, 28, 32] {
assert_eq!(
compute_ipv4_start_node(&raw_node(size, 1, 1), 1, u16::from(size)).unwrap(),
None
);
assert!(matches!(
compute_ipv4_start_node(&[], 1, u16::from(size)),
Err(Error::UnexpectedEof)
));
assert_eq!(
compute_ipv4_start_node(&[], 0, u16::from(size)).unwrap(),
None
);
}
assert_eq!(
compute_ipv4_start_node(&raw_node(40, 1, 1), 1, 40).unwrap(),
None
);
assert!(matches!(
compute_ipv4_start_node(&[], 1, 40),
Err(Error::UnexpectedEof)
));
assert!(matches!(
read_record_static(&[], 0, 40),
Err(Error::UnexpectedEof)
));
assert!(matches!(
read_packed_record_static(&[], 40, 0),
Err(Error::UnexpectedEof)
));
let nodes: Vec<_> = (0..97_u64).map(|i| ((i + 1).min(96), 97)).collect();
assert_eq!(
compute_ipv4_start_node(&node_stream(40, &nodes), 97, 40).unwrap(),
Some(96)
);
}
#[test]
fn optional_and_mapped_lookups_preserve_hit_miss_and_decode_error_semantics() {
struct Text<'a>(&'a str);
impl<'a> MmdbDecode<'a> for Text<'a> {
fn decode(value: &ValueRef<'a>) -> Result<Self> {
match value {
ValueRef::Utf8(v) => Ok(Self(v)),
_ => Err(Error::DecodingError("expected text".into())),
}
}
}
struct Number;
impl<'a> MmdbDecode<'a> for Number {
fn decode(_value: &ValueRef<'a>) -> Result<Self> {
Err(Error::DecodingError("expected number".into()))
}
}
let reader = scalar_reader(24);
let address = ip("2001:db8::1");
assert_eq!(
reader.lookup_borrowed_opt::<Text<'_>>(address).unwrap().0,
"ab"
);
assert_eq!(
reader
.lookup_borrowed_map(address, |record: Text<'_>| record.0)
.unwrap(),
Some("ab")
);
assert!(reader.lookup_borrowed_opt::<Number>(address).is_none());
assert!(
reader
.lookup_borrowed_map(address, |_record: Number| ())
.is_err()
);
assert!(reader.lookup_exists(address));
assert!(reader.record_to_file_offset_opt(17).is_some());
assert!(reader.record_to_file_offset_opt(1).is_none());
assert!(
reader
.record_to_file_offset_opt(usize::MAX as u64)
.is_none()
);
let miss = crafted_reader(24, 1, raw_node(24, 1, 1), vec![0x42, b'a', b'b'], 4, None);
let ipv4 = ip("203.0.113.1");
assert!(miss.lookup_borrowed_opt::<Text<'_>>(ipv4).is_none());
assert_eq!(
miss.lookup_borrowed_map(ipv4, |record: Text<'_>| record.0)
.unwrap(),
None
);
assert!(!miss.lookup_exists(ipv4));
assert!(miss.lookup_borrowed_opt::<Text<'_>>(address).is_none());
let mut ipv6 = scalar_reader(24);
assert_eq!(ipv6.lookup_borrowed_opt::<Text<'_>>(ipv4).unwrap().0, "ab");
assert_eq!(
ipv6.lookup_borrowed_map(ipv4, |record: Text<'_>| record.0)
.unwrap(),
Some("ab")
);
ipv6.metadata.ip_version = 9;
assert!(ipv6.lookup_borrowed_opt::<Text<'_>>(ipv4).is_none());
assert!(matches!(
ipv6.lookup_borrowed_map(ipv4, |_record: Text<'_>| ()),
Err(Error::InvalidIpVersion(9))
));
let mut offset_reader = scalar_reader(24);
offset_reader.data_pointer_bias = usize::MAX;
assert!(offset_reader.record_to_file_offset_opt(17).is_none());
offset_reader.data_pointer_bias = 5;
offset_reader.metadata_marker = 20;
assert!(offset_reader.record_to_file_offset_opt(17).is_none());
}
#[test]
fn scalar_traversal_all_record_sizes() {
for record_size in [24, 28, 32] {
let reader = scalar_reader(record_size);
let (value, prefix) = reader.lookup_value_with_prefix(ip("2001:db8::1")).unwrap();
assert_eq!(value, ValueRef::Utf8("ab"));
assert_eq!(prefix, 1);
}
}
#[test]
fn scalar_traversal_packed_fallback() {
let reader = scalar_reader(40);
let value = reader.lookup_value(ip("127.0.0.1")).unwrap();
assert_eq!(value, ValueRef::Utf8("ab"));
}
#[test]
fn scalar_traversal_reports_not_found() {
let node_count = 1;
let reader = crafted_reader(
24,
node_count,
node_stream(24, &[(node_count, 17)]),
vec![0x42, b'a', b'b'],
6,
Some(0),
);
let err = reader.lookup_value(ip("2001:db8::1")).unwrap_err();
assert!(matches!(err, Error::NotFound));
}
#[test]
fn prepared_traversal_preserves_bits_nodes_and_prefixes() {
let octets = [
0xa5, 0x5a, 0x93, 0x6c, 0x81, 0x7e, 0xc3, 0x3c, 0xf0, 0x0f, 0x96, 0x69, 0x87, 0x78,
0xaa, 0x55,
];
for record_size in [24, 28, 32, 36, 40, 44, 48, 52, 56, 60, 64] {
for ip_version in [4, 6] {
let query = if ip_version == 4 {
IpAddr::from([octets[0], octets[1], octets[2], octets[3]])
} else {
IpAddr::from(octets)
};
let query_bytes = &octets[..if ip_version == 4 { 4 } else { 16 }];
for prefix in [1, 2, 7, 8, 9, 17, 31, 32, 63, 64, 65, 127, 128] {
if prefix > query_bytes.len() * 8 {
continue;
}
let count = prefix as u64;
let data_pointer = count + 16;
let mut nodes = Vec::new();
for bit in 0..prefix {
let side = (query_bytes[bit / 8] >> (7 - bit % 8)) & 1;
let next = if bit + 1 == prefix {
data_pointer
} else {
(bit + 1) as u64
};
nodes.push(if side == 0 {
(next, count)
} else {
(count, next)
});
}
let reader = crafted_reader(
record_size,
count,
node_stream(record_size as u8, &nodes),
vec![0x42, b'a', b'b'],
ip_version,
None,
);
let traverse = |bytes: &[u8]| {
if ip_version == 4 {
reader.traverse_ipv4(0, bytes.try_into().unwrap(), 0)
} else {
reader.traverse_ipv6(0, bytes.try_into().unwrap(), 0)
}
};
assert_eq!(traverse(query_bytes), Some((data_pointer, prefix as u8)));
assert_eq!(
reader.lookup_value_with_prefix(query).unwrap(),
(ValueRef::Utf8("ab"), prefix as u8)
);
for bit in [0, prefix / 2, prefix - 1] {
let mut miss = query_bytes.to_vec();
miss[bit / 8] ^= 0x80 >> (bit % 8);
assert_eq!(traverse(&miss), None);
}
}
}
}
}
#[test]
fn prepared_traversal_terminal_records_and_cycles_do_not_load_children() {
for record_size in [24, 28, 32, 36, 40, 44, 48, 52, 56, 60, 64] {
let max_record = if record_size == 64 {
u64::MAX
} else {
(1_u64 << record_size) - 1
};
let reader = crafted_reader(
record_size,
3,
node_stream(record_size as u8, &[(1, 3), (2, 3), (19, max_record)]),
vec![0x42, b'a', b'b'],
6,
Some(0),
);
assert_eq!(reader.traverse_ipv6(1, &[0; 16], 10), Some((19, 12)));
assert_eq!(
reader.traverse_ipv6(2, &[0x80; 16], 0),
Some((max_record, 1))
);
assert_eq!(
reader.traverse_ipv6(max_record, &[0; 16], 7),
Some((max_record, 7))
);
assert_eq!(reader.traverse_ipv6(3, &[0; 16], 0), None);
assert!(matches!(
reader.resolve_offset(ip("2000::")),
Err(Error::InvalidOffset(_))
));
let reserved = crafted_reader(
record_size,
1,
node_stream(record_size as u8, &[(2, 2)]),
vec![0x42, b'a', b'b'],
6,
Some(0),
);
assert!(matches!(
reserved.resolve_offset(ip("::")),
Err(Error::InvalidOffset(_))
));
let cyclic = crafted_reader(
record_size,
1,
node_stream(record_size as u8, &[(0, 0)]),
Vec::new(),
6,
Some(0),
);
assert_eq!(cyclic.traverse_ipv6(0, &[0xa5; 16], 0), None);
}
}
#[test]
fn read_record_all_sizes_and_errors() {
for record_size in [24, 28, 32] {
let nodes = [(1, 2), (19, 19), (3, 19)];
let reader = crafted_reader(
record_size,
3,
node_stream(record_size as u8, &nodes),
vec![0x42, b'a', b'b'],
6,
Some(0),
);
assert_eq!(reader.read_record(0, 0).unwrap(), 1);
assert_eq!(reader.read_record(0, 1).unwrap(), 2);
assert_eq!(reader.read_record(1, 0).unwrap(), 19);
assert_eq!(reader.read_record(1, 1).unwrap(), 19);
assert_eq!(reader.read_record(2, 0).unwrap(), 3);
assert_eq!(reader.read_record(2, 1).unwrap(), 19);
assert!(matches!(
reader.read_record(3, 0).unwrap_err(),
Error::InvalidNode(3)
));
assert!(matches!(
reader.read_record(0, 2).unwrap_err(),
Error::InvalidNode(_)
));
}
let reader = crafted_reader(
40,
1,
node_stream(40, &[(17, 17)]),
vec![0x42, b'a', b'b'],
6,
Some(0),
);
assert_eq!(reader.read_record(0, 0).unwrap(), 17);
assert_eq!(reader.read_record(0, 1).unwrap(), 17);
}
#[test]
fn read_record_rejects_out_of_bounds_node() {
let reader = crafted_reader(
24,
2,
node_stream(24, &[(0, 0), (0, 0)]),
vec![0x42, b'a', b'b'],
6,
Some(0),
);
assert!(matches!(
reader.read_record(2, 0).unwrap_err(),
Error::InvalidNode(2)
));
}
#[test]
fn record_to_file_offset_error_branches() {
let reader = crafted_reader(
24,
3,
node_stream(24, &[(0, 0), (0, 0), (0, 0)]),
vec![0x42, b'a', b'b'],
6,
Some(0),
);
assert!(matches!(
reader.record_to_file_offset(3).unwrap_err(),
Error::InvalidOffset(3)
));
assert!(matches!(
reader.record_to_file_offset(u64::MAX).unwrap_err(),
Error::InvalidOffset(usize::MAX)
));
assert!(matches!(
reader
.record_to_file_offset(0xffff_ffff_ffff_ff00)
.unwrap_err(),
Error::InvalidOffset(_)
));
let reader = crafted_reader(
24,
3,
node_stream(24, &[(0, 0), (0, 0), (0, 0)]),
vec![],
6,
Some(0),
);
assert!(matches!(
reader.record_to_file_offset(24).unwrap_err(),
Error::InvalidOffset(39)
));
}
#[test]
fn lookup_rejects_v6_in_v4_database() {
let reader = crafted_reader(24, 1, Vec::new(), Vec::new(), 4, None);
assert!(matches!(
reader.lookup_value(ip("::1")).unwrap_err(),
Error::NotFound
));
}
#[test]
fn lookup_rejects_invalid_ip_version() {
let reader = crafted_reader(24, 1, Vec::new(), Vec::new(), 5, None);
assert!(matches!(
reader.lookup_value(ip("1.2.3.4")).unwrap_err(),
Error::InvalidIpVersion(5)
));
}
#[test]
fn lookup_v4_in_v6_missing_start_node() {
let reader = crafted_reader(24, 1, Vec::new(), Vec::new(), 6, None);
assert!(matches!(
reader.lookup_value(ip("1.2.3.4")).unwrap_err(),
Error::NotFound
));
}
fn metadata_db(ip_version: u16) -> Vec<u8> {
let mut writer = Writer::with_metadata(
MetadataBuilder::new()
.database_type("reader-tests")
.ip_version(ip_version)
.build()
.unwrap(),
);
writer
.insert_value(
if ip_version == 6 {
"2001:db8::/32".parse().unwrap()
} else {
"10.0.0.0/8".parse().unwrap()
},
Value::Map(Into::into({
let mut m = std::collections::BTreeMap::new();
m.insert("name".into(), Value::Utf8("x".into()));
m
})),
)
.unwrap();
writer.finish().unwrap()
}
fn patch_uint16(db: &mut [u8], key: &str, value: u8) {
let index = db
.windows(key.len())
.position(|w| w == key.as_bytes())
.expect("key present in metadata");
debug_assert_eq!(db[index + key.len()], 0xA1, "u16 encoded as control+1");
db[index + key.len() + 1] = value;
}
#[test]
fn from_source_rejects_bad_metadata_values() {
let mut bad_major = metadata_db(6);
patch_uint16(&mut bad_major, "binary_format_major_version", 3);
assert!(matches!(
Reader::from_bytes(&bad_major).unwrap_err(),
Error::InvalidMetadata("unsupported binary format major version")
));
let mut bad_ip = metadata_db(6);
patch_uint16(&mut bad_ip, "ip_version", 5);
assert!(matches!(
Reader::from_bytes(&bad_ip).unwrap_err(),
Error::InvalidIpVersion(5)
));
let mut bad_record_size = metadata_db(6);
patch_uint16(&mut bad_record_size, "record_size", 16);
assert!(matches!(
Reader::from_bytes(&bad_record_size).unwrap_err(),
Error::InvalidMetadata(msg) if msg.contains("record_size")
));
}
#[test]
fn from_source_rejects_broken_separator() {
let db = metadata_db(6);
let reader = Reader::from_bytes(&db).unwrap();
let mut broken = db.clone();
broken[reader.data_section_start - 1] = 0x01;
assert!(matches!(
Reader::from_bytes(&broken).unwrap_err(),
Error::InvalidDatabase(msg) if msg.contains("separator")
));
}
#[test]
fn from_source_and_as_bytes_and_serde_lookup() {
let db = metadata_db(6);
assert_eq!(Reader::from_bytes(&db).unwrap().metadata().ip_version, 6);
let reader = Reader::from_bytes(&db).unwrap();
let prepared = &reader.prepared_tree as *const PreparedTree;
assert_eq!(reader.as_bytes(), &db);
let map: std::collections::BTreeMap<String, String> =
reader.lookup(ip("2001:db8::1")).unwrap();
assert_eq!(map["name"], "x");
assert_eq!(&reader.prepared_tree as *const PreparedTree, prepared);
}
#[test]
fn open_owned_and_mmap_sources() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("db.mmdb");
std::fs::write(&path, metadata_db(6)).unwrap();
let owned = Reader::open(&path).unwrap();
assert!(matches!(
owned.lookup_value(ip("2001:db8::1")).unwrap(),
ValueRef::Map(_)
));
let mapped = unsafe { Reader::open_mmap(&path) }.unwrap();
let (value, _) = mapped.lookup_value_with_prefix(ip("2001:db8::1")).unwrap();
assert!(matches!(value, ValueRef::Map(_)));
}
}