1mod marker;
10mod tree;
11
12use memmap2::Mmap;
13use serde::de::DeserializeOwned;
14use std::net::IpAddr;
15use std::path::Path;
16use std::sync::OnceLock;
17
18use crate::decoder::{Decoder, RawDecoder};
19use crate::{Error, Metadata, MmdbDecode, Result, ValueRef};
20use marker::{METADATA_MARKER, find_metadata_marker};
21use tree::PreparedTree;
22
23#[cold]
27fn compute_ipv4_start_node(bytes: &[u8], node_count: u64, record_size: u16) -> Result<Option<u64>> {
28 let mut node = 0_u64;
29 match record_size {
30 28 => {
31 for _ in 0..96 {
32 if node >= node_count {
33 return Ok(None);
34 }
35 let offset = node as usize * 7;
36 if offset + 7 > bytes.len() {
37 return Err(Error::UnexpectedEof);
38 }
39 let p = unsafe { bytes.as_ptr().add(offset) };
41 node = unsafe {
43 (u64::from(*p.add(3) >> 4) << 24)
44 | (u64::from(*p) << 16)
45 | (u64::from(*p.add(1)) << 8)
46 | u64::from(*p.add(2))
47 };
48 if node >= node_count {
49 return Ok(None);
50 }
51 }
52 Ok(Some(node))
53 }
54 32 => {
55 for _ in 0..96 {
56 if node >= node_count {
57 return Ok(None);
58 }
59 let offset = node as usize * 8;
60 if offset + 8 > bytes.len() {
61 return Err(Error::UnexpectedEof);
62 }
63 let p = unsafe { bytes.as_ptr().add(offset) };
65 node = u64::from(u32::from_be(unsafe {
67 core::ptr::read_unaligned(p.cast::<u32>())
68 }));
69 if node >= node_count {
70 return Ok(None);
71 }
72 }
73 Ok(Some(node))
74 }
75 24 => {
76 for _ in 0..96 {
77 if node >= node_count {
78 return Ok(None);
79 }
80 let offset = node as usize * 6;
81 if offset + 6 > bytes.len() {
82 return Err(Error::UnexpectedEof);
83 }
84 let p = unsafe { bytes.as_ptr().add(offset) };
86 node = unsafe {
88 (u64::from(*p) << 16) | (u64::from(*p.add(1)) << 8) | u64::from(*p.add(2))
89 };
90 if node >= node_count {
91 return Ok(None);
92 }
93 }
94 Ok(Some(node))
95 }
96 _ => {
97 for _ in 0..96 {
99 if node >= node_count {
100 return Ok(None);
101 }
102 node = match read_record_static(bytes, node as usize, record_size)? {
103 Some(next) if next < node_count => next,
104 _ => return Ok(None),
105 };
106 }
107 Ok(Some(node))
108 }
109 }
110}
111
112#[cold]
114fn read_record_static(bytes: &[u8], node: usize, record_size: u16) -> Result<Option<u64>> {
115 let node_size = usize::from(record_size) / 4;
116 let offset = node
117 .checked_mul(node_size)
118 .ok_or(Error::InvalidNode(node as u64))?;
119 let slice = bytes
120 .get(offset..offset + node_size)
121 .ok_or(Error::UnexpectedEof)?;
122 Ok(Some(read_packed_record_static(
125 slice,
126 usize::from(record_size),
127 0,
128 )?))
129}
130
131#[cold]
133fn read_packed_record_static(bytes: &[u8], bits: usize, side: usize) -> Result<u64> {
134 let start = side * bits;
135 let mut value = 0_u64;
136 for bit_index in start..start + bits {
137 let byte = *bytes.get(bit_index / 8).ok_or(Error::UnexpectedEof)?;
138 let bit = (byte >> (7 - (bit_index % 8))) & 1;
139 value = (value << 1) | u64::from(bit);
140 }
141 Ok(value)
142}
143
144#[derive(Debug)]
149enum Source<'a> {
150 Borrowed(&'a [u8]),
151 Mmap(Mmap),
152 Owned(Vec<u8>),
153}
154
155impl Source<'_> {
156 #[inline(always)]
157 fn bytes(&self) -> &[u8] {
158 match self {
159 Self::Borrowed(v) => v,
160 Self::Mmap(v) => v,
161 Self::Owned(v) => v,
162 }
163 }
164}
165
166#[derive(Debug)]
179pub struct Reader<'a> {
180 source: Source<'a>,
181 metadata: Metadata,
182 data_pointer_bias: usize,
183 data_section_start: usize,
184 metadata_marker: usize,
185 ipv4_start_node: Option<u64>,
186 prepared_tree: PreparedTree,
188}
189
190impl Reader<'static> {
191 pub fn open(path: impl AsRef<Path>) -> Result<Self> {
208 Self::from_vec(std::fs::read(path)?)
209 }
210
211 pub unsafe fn open_mmap(path: impl AsRef<Path>) -> Result<Self> {
232 let file = std::fs::File::open(path)?;
233 let mmap = unsafe { Mmap::map(&file)? };
235 Self::from_source(Source::Mmap(mmap))
236 }
237
238 pub fn from_vec(data: Vec<u8>) -> Result<Self> {
255 Self::from_source(Source::Owned(data))
256 }
257}
258
259impl<'a> Reader<'a> {
260 #[inline]
276 pub fn from_bytes(data: &'a [u8]) -> Result<Self> {
277 Self::from_source(Source::Borrowed(data))
278 }
279
280 fn from_source(source: Source<'a>) -> Result<Self> {
281 let bytes = source.bytes();
282 let marker = find_metadata_marker(bytes)?;
283 let metadata_start = marker + METADATA_MARKER.len();
284 let decoder = Decoder::new(bytes, metadata_start, bytes.len());
285 let (value, _) = decoder.decode_at(metadata_start)?;
286 let metadata = Metadata::from_value(&value)?;
287
288 if metadata.binary_format_major_version != 2 {
289 return Err(Error::InvalidMetadata(
290 "unsupported binary format major version",
291 ));
292 }
293 if !matches!(metadata.ip_version, 4 | 6) {
294 return Err(Error::InvalidIpVersion(metadata.ip_version));
295 }
296 if metadata.record_size < 24 || metadata.record_size % 4 != 0 || metadata.record_size > 64 {
297 return Err(Error::InvalidMetadata(
298 "record_size must be a multiple of 4 between 24 and 64",
299 ));
300 }
301
302 let node_bytes = usize::from(metadata.record_size) / 4;
303 let node_count_usize = usize::try_from(metadata.node_count)
304 .map_err(|_| Error::InvalidMetadata("node count exceeds address space"))?;
305 let search_tree_size = node_count_usize
306 .checked_mul(node_bytes)
307 .ok_or(Error::InvalidMetadata("search tree size overflow"))?;
308 let data_pointer_bias = search_tree_size
311 .checked_sub(node_count_usize)
312 .ok_or(Error::InvalidMetadata("invalid search tree geometry"))?;
313 let data_section_start = search_tree_size
314 .checked_add(16)
315 .ok_or(Error::InvalidMetadata("data section offset overflow"))?;
316
317 if data_section_start > marker || data_section_start > bytes.len() {
318 return Err(Error::InvalidDatabase(
319 "search tree overlaps metadata or exceeds file",
320 ));
321 }
322 if bytes.get(search_tree_size..data_section_start) != Some(&[0_u8; 16][..]) {
323 return Err(Error::InvalidDatabase(
324 "missing 16-byte data section separator",
325 ));
326 }
327
328 let record_size_u8 = u8::try_from(metadata.record_size)
331 .map_err(|_| Error::InvalidMetadata("record_size must fit in u8"))?;
332
333 let ipv4_start_node = if metadata.ip_version == 6 {
336 compute_ipv4_start_node(bytes, metadata.node_count, metadata.record_size)?
337 } else {
338 None
339 };
340
341 let prepared_tree = PreparedTree::build(
347 &bytes[..data_section_start],
348 record_size_u8,
349 metadata.node_count,
350 ipv4_start_node,
351 )?;
352
353 Ok(Self {
354 source,
355 metadata,
356 data_pointer_bias,
357 data_section_start,
358 metadata_marker: marker,
359 ipv4_start_node,
360 prepared_tree,
361 })
362 }
363
364 #[must_use]
378 pub const fn metadata(&self) -> &Metadata {
379 &self.metadata
380 }
381
382 #[must_use]
396 pub fn as_bytes(&self) -> &[u8] {
397 self.source.bytes()
398 }
399
400 #[inline]
416 pub fn lookup_value(&self, ip: IpAddr) -> Result<ValueRef<'_>> {
417 self.lookup_value_with_prefix(ip).map(|(v, _)| v)
418 }
419
420 #[inline]
437 pub fn lookup_value_with_prefix(&self, ip: IpAddr) -> Result<(ValueRef<'_>, u8)> {
438 let (offset, prefix) = self.resolve_offset(ip)?;
439 let decoder = Decoder::new(
440 self.source.bytes(),
441 self.data_section_start,
442 self.metadata_marker,
443 );
444 let (value, _) = decoder.decode_at(offset)?;
445 Ok((value, prefix))
446 }
447
448 #[inline(always)]
483 pub fn lookup_borrowed<'s, T>(&'s self, ip: IpAddr) -> Result<T>
484 where
485 T: MmdbDecode<'s>,
486 {
487 let (offset, _) = self.resolve_offset(ip)?;
488 decode_borrowed_at::<T>(self, offset)
489 }
490
491 #[inline(always)]
529 pub fn lookup_borrowed_opt<'s, T>(&'s self, ip: IpAddr) -> Option<T>
530 where
531 T: MmdbDecode<'s>,
532 {
533 let (offset, _) = self.resolve_offset_opt(ip)?;
534 decode_borrowed_at_opt::<T>(self, offset)
535 }
536
537 #[inline(always)]
561 pub fn lookup_borrowed_map<'s, T, R>(
562 &'s self,
563 ip: IpAddr,
564 on_hit: impl FnOnce(T) -> R,
565 ) -> Result<Option<R>>
566 where
567 T: MmdbDecode<'s>,
568 {
569 let traversed = match (self.metadata.ip_version, ip) {
570 (4, IpAddr::V4(v4)) => self.traverse_ipv4(0, &v4.octets(), 0),
571 (4, IpAddr::V6(_)) => None,
572 (6, IpAddr::V6(v6)) => self.traverse_ipv6(0, &v6.octets(), 0),
573 (6, IpAddr::V4(v4)) => self
574 .ipv4_start_node
575 .and_then(|start| self.traverse_ipv4(start, &v4.octets(), 0)),
576 (v, _) => return Err(Error::InvalidIpVersion(v)),
577 };
578 let Some((record, _)) = traversed else {
579 return Ok(None);
580 };
581 let offset = self.record_to_file_offset(record)?;
582 decode_borrowed_map_at::<T, R>(self, offset, on_hit).map(Some)
583 }
584
585 pub fn lookup<T: DeserializeOwned>(&self, ip: IpAddr) -> Result<T> {
602 let value = self.lookup_value(ip)?;
603 Ok(serde_json::from_value(value.to_json())?)
604 }
605
606 pub fn lookup_many(&self, ips: &[IpAddr]) -> Vec<Result<ValueRef<'_>>> {
627 const PARALLEL_MIN: usize = 4_096;
628 static WORKERS: OnceLock<usize> = OnceLock::new();
629 let workers =
630 *WORKERS.get_or_init(|| std::thread::available_parallelism().map_or(1, |c| c.get()));
631
632 if workers <= 1 || ips.len() < PARALLEL_MIN {
633 return ips.iter().map(|&ip| self.lookup_value(ip)).collect();
634 }
635
636 let chunk = ips.len().div_ceil(workers.min(ips.len()));
637 let mut results: Vec<Result<ValueRef<'_>>> =
638 (0..ips.len()).map(|_| Err(Error::NotFound)).collect();
639
640 std::thread::scope(|scope| {
641 for (slice, slots) in ips.chunks(chunk).zip(results.chunks_mut(chunk)) {
642 scope.spawn(move || {
643 for (ip, slot) in slice.iter().zip(slots) {
644 *slot = self.lookup_value(*ip);
645 }
646 });
647 }
648 });
649
650 results
651 }
652
653 #[inline(always)]
676 #[allow(clippy::unnecessary_lazy_evaluations)]
677 fn resolve_offset(&self, ip: IpAddr) -> Result<(usize, u8)> {
678 let (record, prefix) = match (self.metadata.ip_version, ip) {
679 (4, IpAddr::V4(v4)) => self
680 .traverse_ipv4(0, &v4.octets(), 0)
681 .ok_or(Error::NotFound)?,
682 (4, IpAddr::V6(_)) => return Err(Error::NotFound),
683 (6, IpAddr::V6(v6)) => self
684 .traverse_ipv6(0, &v6.octets(), 0)
685 .ok_or(Error::NotFound)?,
686 (6, IpAddr::V4(v4)) => {
687 let start = self.ipv4_start_node.ok_or_else(|| Error::NotFound)?;
688 self.traverse_ipv4(start, &v4.octets(), 0)
689 .ok_or(Error::NotFound)?
690 }
691 (v, _) => return Err(Error::InvalidIpVersion(v)),
692 };
693 let offset = self.record_to_file_offset(record)?;
694 Ok((offset, prefix))
695 }
696
697 #[inline(always)]
704 fn resolve_offset_opt(&self, ip: IpAddr) -> Option<(usize, u8)> {
705 let (record, prefix) = match (self.metadata.ip_version, ip) {
706 (4, IpAddr::V4(v4)) => self.traverse_ipv4(0, &v4.octets(), 0)?,
707 (4, IpAddr::V6(_)) => return None,
708 (6, IpAddr::V6(v6)) => self.traverse_ipv6(0, &v6.octets(), 0)?,
709 (6, IpAddr::V4(v4)) => {
710 let start = self.ipv4_start_node?;
711 self.traverse_ipv4(start, &v4.octets(), 0)?
712 }
713 _ => return None,
714 };
715 let offset = self.record_to_file_offset_opt(record)?;
716 Some((offset, prefix))
717 }
718
719 #[inline]
736 pub fn lookup_exists(&self, ip: IpAddr) -> bool {
737 self.resolve_offset(ip).is_ok()
738 }
739
740 #[inline(always)]
741 fn traverse_ipv4(&self, node: u64, octets: &[u8; 4], prefix_base: u8) -> Option<(u64, u8)> {
742 self.prepared_tree.traverse_ipv4(node, octets, prefix_base)
743 }
744
745 #[inline(always)]
746 fn traverse_ipv6(&self, node: u64, octets: &[u8; 16], prefix_base: u8) -> Option<(u64, u8)> {
747 self.prepared_tree.traverse_ipv6(node, octets, prefix_base)
748 }
749
750 #[allow(dead_code, clippy::unnecessary_lazy_evaluations)]
755 #[cold]
756 fn read_record(&self, node: u64, side: usize) -> Result<u64> {
757 if node >= self.metadata.node_count || side > 1 {
758 return Err(Error::InvalidNode(node));
759 }
760 let node_size = usize::from(self.metadata.record_size) / 4;
761 let offset = usize::try_from(node)
762 .ok()
763 .and_then(|n| n.checked_mul(node_size))
764 .ok_or_else(|| Error::InvalidNode(node))?;
765 let bytes = self
766 .source
767 .bytes()
768 .get(offset..offset + node_size)
769 .ok_or_else(|| Error::UnexpectedEof)?;
770
771 match self.metadata.record_size {
772 24 => {
773 let base = side * 3;
774 Ok((u64::from(unsafe { *bytes.get_unchecked(base) }) << 16)
775 | (u64::from(unsafe { *bytes.get_unchecked(base + 1) }) << 8)
776 | u64::from(unsafe { *bytes.get_unchecked(base + 2) }))
777 }
778 28 => {
779 if side == 0 {
780 Ok((u64::from(unsafe { *bytes.get_unchecked(3) } >> 4) << 24)
781 | (u64::from(unsafe { *bytes.get_unchecked(0) }) << 16)
782 | (u64::from(unsafe { *bytes.get_unchecked(1) }) << 8)
783 | u64::from(unsafe { *bytes.get_unchecked(2) }))
784 } else {
785 Ok((u64::from(unsafe { *bytes.get_unchecked(3) } & 0x0f) << 24)
786 | (u64::from(unsafe { *bytes.get_unchecked(4) }) << 16)
787 | (u64::from(unsafe { *bytes.get_unchecked(5) }) << 8)
788 | u64::from(unsafe { *bytes.get_unchecked(6) }))
789 }
790 }
791 32 => {
792 let base = side * 4;
793 Ok(u64::from(u32::from_be_bytes(
794 unsafe { bytes.get_unchecked(base..base + 4) }
795 .try_into()
796 .expect("length checked"),
797 )))
798 }
799 bits => read_packed_record(bytes, usize::from(bits), side),
800 }
801 }
802
803 #[inline]
804 #[allow(clippy::unnecessary_lazy_evaluations)]
805 fn record_to_file_offset(&self, record: u64) -> Result<usize> {
806 let node_count = self.metadata.node_count;
807 if record < node_count.saturating_add(16) {
808 return Err(Error::InvalidOffset(record as usize));
809 }
810 let record = usize::try_from(record).map_err(|_| Error::InvalidOffset(usize::MAX))?;
811 let offset = record
812 .checked_add(self.data_pointer_bias)
813 .ok_or(Error::InvalidOffset(record))?;
814
815 if offset >= self.metadata_marker {
816 return Err(Error::InvalidOffset(offset));
817 }
818 Ok(offset)
819 }
820
821 #[inline(always)]
825 fn record_to_file_offset_opt(&self, record: u64) -> Option<usize> {
826 let node_count = self.metadata.node_count;
827 if record < node_count.saturating_add(16) {
828 return None;
829 }
830 let record = usize::try_from(record).ok()?;
831 let offset = record.checked_add(self.data_pointer_bias)?;
832 if offset >= self.metadata_marker {
833 return None;
834 }
835 Some(offset)
836 }
837}
838
839#[cold]
850#[inline(never)]
851fn decode_borrowed_at<'s, T>(reader: &'s Reader<'_>, offset: usize) -> Result<T>
852where
853 T: MmdbDecode<'s>,
854{
855 let mut decoder = RawDecoder::new(
856 reader.source.bytes(),
857 reader.data_section_start,
858 reader.metadata_marker,
859 offset,
860 );
861 T::decode_raw(&mut decoder)
862}
863
864#[cold]
867#[inline(never)]
868fn decode_borrowed_map_at<'s, T, R>(
869 reader: &'s Reader<'_>,
870 offset: usize,
871 on_hit: impl FnOnce(T) -> R,
872) -> Result<R>
873where
874 T: MmdbDecode<'s>,
875{
876 let mut decoder = RawDecoder::new(
877 reader.source.bytes(),
878 reader.data_section_start,
879 reader.metadata_marker,
880 offset,
881 );
882 let value = T::decode_raw(&mut decoder)?;
883 Ok(on_hit(value))
884}
885
886#[cold]
893#[inline(never)]
894fn decode_borrowed_at_opt<'s, T>(reader: &'s Reader<'_>, offset: usize) -> Option<T>
895where
896 T: MmdbDecode<'s>,
897{
898 let mut decoder = RawDecoder::new(
899 reader.source.bytes(),
900 reader.data_section_start,
901 reader.metadata_marker,
902 offset,
903 );
904 T::decode_raw(&mut decoder).ok()
905}
906
907#[allow(clippy::unnecessary_lazy_evaluations)]
912fn read_packed_record(bytes: &[u8], bits: usize, side: usize) -> Result<u64> {
913 let start = side * bits;
914 let mut value = 0_u64;
915 for bit_index in start..start + bits {
916 let byte = *bytes.get(bit_index / 8).ok_or(Error::UnexpectedEof)?;
917 let bit = (byte >> (7 - (bit_index % 8))) & 1;
918 value = (value << 1) | u64::from(bit);
919 }
920 Ok(value)
921}
922
923#[cfg(all(test, feature = "writer"))]
928mod reader_tests {
929 use super::*;
930 use crate::writer::Writer;
931 use crate::{MetadataBuilder, Value};
932
933 fn raw_node(record_size: u8, left: u64, right: u64) -> Vec<u8> {
936 match record_size {
937 24 => {
938 let mut n = Vec::with_capacity(6);
939 n.extend_from_slice(&(left as u32).to_be_bytes()[1..=3]);
940 n.extend_from_slice(&(right as u32).to_be_bytes()[1..=3]);
941 n
942 }
943 28 => vec![
944 (left >> 16) as u8,
945 (left >> 8) as u8,
946 left as u8,
947 ((((left >> 24) & 0x0f) << 4) | ((right >> 24) & 0x0f)) as u8,
948 (right >> 16) as u8,
949 (right >> 8) as u8,
950 right as u8,
951 ],
952 32 => {
953 let mut n = Vec::with_capacity(8);
954 n.extend_from_slice(&(left as u32).to_be_bytes());
955 n.extend_from_slice(&(right as u32).to_be_bytes());
956 n
957 }
958 36..=64 if record_size.is_multiple_of(4) => {
961 let bits = usize::from(record_size);
962 let mut n = vec![0_u8; bits / 4];
963 for (side, value) in [left, right].into_iter().enumerate() {
964 for bit in 0..bits {
965 let stream_bit = side * bits + bit;
966 n[stream_bit / 8] |=
967 (((value >> (bits - bit - 1)) & 1) as u8) << (7 - stream_bit % 8);
968 }
969 }
970 n
971 }
972 _ => panic!("unexpected record size {record_size}"),
973 }
974 }
975
976 fn node_stream(record_size: u8, nodes: &[(u64, u64)]) -> Vec<u8> {
977 let mut out = Vec::new();
978 for &(l, r) in nodes {
979 out.extend_from_slice(&raw_node(record_size, l, r));
980 }
981 out
982 }
983
984 fn crafted_reader(
987 record_size: u16,
988 node_count: u64,
989 tree: Vec<u8>,
990 data: Vec<u8>,
991 ip_version: u16,
992 ipv4_start: Option<u64>,
993 ) -> Reader<'static> {
994 let mut file = tree;
995 file.extend_from_slice(&[0_u8; 16]);
996 file.extend_from_slice(&data);
997 let search_tree_size = (node_count as usize) * (usize::from(record_size) / 4);
998 let data_section_start = search_tree_size + 16;
999 let prepared_tree = PreparedTree::build(&file, record_size as u8, node_count, ipv4_start)
1000 .expect("crafted tree must be preparable");
1001 Reader {
1002 source: Source::Owned(file),
1003 metadata: Metadata {
1004 node_count,
1005 record_size,
1006 ip_version,
1007 database_type: "test".into(),
1008 languages: vec!["en".into()],
1009 binary_format_major_version: 2,
1010 binary_format_minor_version: 0,
1011 build_epoch: 0,
1012 description: Default::default(),
1013 },
1014 data_pointer_bias: search_tree_size - (node_count as usize),
1015 data_section_start,
1016 metadata_marker: data_section_start + data.len(),
1017 ipv4_start_node: ipv4_start,
1018 prepared_tree,
1019 }
1020 }
1021
1022 fn scalar_reader(record_size: u16) -> Reader<'static> {
1025 let node_count = 1;
1026 let data_pointer = node_count + 16;
1027 crafted_reader(
1028 record_size,
1029 node_count,
1030 node_stream(record_size as u8, &[(data_pointer, data_pointer)]),
1031 vec![0x42, b'a', b'b'],
1032 6,
1033 Some(0),
1034 )
1035 }
1036
1037 fn ip(s: &str) -> IpAddr {
1038 s.parse().unwrap()
1039 }
1040
1041 #[test]
1042 fn ipv4_subtree_start_handles_early_leaves_truncation_and_packed_fallback() {
1043 for size in [24_u8, 28, 32] {
1044 assert_eq!(
1045 compute_ipv4_start_node(&raw_node(size, 1, 1), 1, u16::from(size)).unwrap(),
1046 None
1047 );
1048 assert!(matches!(
1049 compute_ipv4_start_node(&[], 1, u16::from(size)),
1050 Err(Error::UnexpectedEof)
1051 ));
1052 assert_eq!(
1053 compute_ipv4_start_node(&[], 0, u16::from(size)).unwrap(),
1054 None
1055 );
1056 }
1057 assert_eq!(
1058 compute_ipv4_start_node(&raw_node(40, 1, 1), 1, 40).unwrap(),
1059 None
1060 );
1061 assert!(matches!(
1062 compute_ipv4_start_node(&[], 1, 40),
1063 Err(Error::UnexpectedEof)
1064 ));
1065 assert!(matches!(
1066 read_record_static(&[], 0, 40),
1067 Err(Error::UnexpectedEof)
1068 ));
1069 assert!(matches!(
1070 read_packed_record_static(&[], 40, 0),
1071 Err(Error::UnexpectedEof)
1072 ));
1073
1074 let nodes: Vec<_> = (0..97_u64).map(|i| ((i + 1).min(96), 97)).collect();
1075 assert_eq!(
1076 compute_ipv4_start_node(&node_stream(40, &nodes), 97, 40).unwrap(),
1077 Some(96)
1078 );
1079 }
1080
1081 #[test]
1082 fn optional_and_mapped_lookups_preserve_hit_miss_and_decode_error_semantics() {
1083 struct Text<'a>(&'a str);
1084 impl<'a> MmdbDecode<'a> for Text<'a> {
1085 fn decode(value: &ValueRef<'a>) -> Result<Self> {
1086 match value {
1087 ValueRef::Utf8(v) => Ok(Self(v)),
1088 _ => Err(Error::DecodingError("expected text".into())),
1089 }
1090 }
1091 }
1092 struct Number;
1093 impl<'a> MmdbDecode<'a> for Number {
1094 fn decode(_value: &ValueRef<'a>) -> Result<Self> {
1095 Err(Error::DecodingError("expected number".into()))
1096 }
1097 }
1098
1099 let reader = scalar_reader(24);
1100 let address = ip("2001:db8::1");
1101 assert_eq!(
1102 reader.lookup_borrowed_opt::<Text<'_>>(address).unwrap().0,
1103 "ab"
1104 );
1105 assert_eq!(
1106 reader
1107 .lookup_borrowed_map(address, |record: Text<'_>| record.0)
1108 .unwrap(),
1109 Some("ab")
1110 );
1111 assert!(reader.lookup_borrowed_opt::<Number>(address).is_none());
1112 assert!(
1113 reader
1114 .lookup_borrowed_map(address, |_record: Number| ())
1115 .is_err()
1116 );
1117 assert!(reader.lookup_exists(address));
1118 assert!(reader.record_to_file_offset_opt(17).is_some());
1119 assert!(reader.record_to_file_offset_opt(1).is_none());
1120 assert!(
1121 reader
1122 .record_to_file_offset_opt(usize::MAX as u64)
1123 .is_none()
1124 );
1125
1126 let miss = crafted_reader(24, 1, raw_node(24, 1, 1), vec![0x42, b'a', b'b'], 4, None);
1127 let ipv4 = ip("203.0.113.1");
1128 assert!(miss.lookup_borrowed_opt::<Text<'_>>(ipv4).is_none());
1129 assert_eq!(
1130 miss.lookup_borrowed_map(ipv4, |record: Text<'_>| record.0)
1131 .unwrap(),
1132 None
1133 );
1134 assert!(!miss.lookup_exists(ipv4));
1135 assert!(miss.lookup_borrowed_opt::<Text<'_>>(address).is_none());
1136
1137 let mut ipv6 = scalar_reader(24);
1138 assert_eq!(ipv6.lookup_borrowed_opt::<Text<'_>>(ipv4).unwrap().0, "ab");
1139 assert_eq!(
1140 ipv6.lookup_borrowed_map(ipv4, |record: Text<'_>| record.0)
1141 .unwrap(),
1142 Some("ab")
1143 );
1144 ipv6.metadata.ip_version = 9;
1145 assert!(ipv6.lookup_borrowed_opt::<Text<'_>>(ipv4).is_none());
1146 assert!(matches!(
1147 ipv6.lookup_borrowed_map(ipv4, |_record: Text<'_>| ()),
1148 Err(Error::InvalidIpVersion(9))
1149 ));
1150
1151 let mut offset_reader = scalar_reader(24);
1152 offset_reader.data_pointer_bias = usize::MAX;
1153 assert!(offset_reader.record_to_file_offset_opt(17).is_none());
1154 offset_reader.data_pointer_bias = 5;
1155 offset_reader.metadata_marker = 20;
1156 assert!(offset_reader.record_to_file_offset_opt(17).is_none());
1157 }
1158
1159 #[test]
1160 fn scalar_traversal_all_record_sizes() {
1161 for record_size in [24, 28, 32] {
1162 let reader = scalar_reader(record_size);
1163 let (value, prefix) = reader.lookup_value_with_prefix(ip("2001:db8::1")).unwrap();
1164 assert_eq!(value, ValueRef::Utf8("ab"));
1165 assert_eq!(prefix, 1);
1166 }
1167 }
1168
1169 #[test]
1170 fn scalar_traversal_packed_fallback() {
1171 let reader = scalar_reader(40);
1174 let value = reader.lookup_value(ip("127.0.0.1")).unwrap();
1175 assert_eq!(value, ValueRef::Utf8("ab"));
1176 }
1177
1178 #[test]
1179 fn scalar_traversal_reports_not_found() {
1180 let node_count = 1;
1181 let reader = crafted_reader(
1183 24,
1184 node_count,
1185 node_stream(24, &[(node_count, 17)]),
1186 vec![0x42, b'a', b'b'],
1187 6,
1188 Some(0),
1189 );
1190 let err = reader.lookup_value(ip("2001:db8::1")).unwrap_err();
1191 assert!(matches!(err, Error::NotFound));
1192 }
1193
1194 #[test]
1195 fn prepared_traversal_preserves_bits_nodes_and_prefixes() {
1196 let octets = [
1199 0xa5, 0x5a, 0x93, 0x6c, 0x81, 0x7e, 0xc3, 0x3c, 0xf0, 0x0f, 0x96, 0x69, 0x87, 0x78,
1200 0xaa, 0x55,
1201 ];
1202 for record_size in [24, 28, 32, 36, 40, 44, 48, 52, 56, 60, 64] {
1203 for ip_version in [4, 6] {
1204 let query = if ip_version == 4 {
1205 IpAddr::from([octets[0], octets[1], octets[2], octets[3]])
1206 } else {
1207 IpAddr::from(octets)
1208 };
1209 let query_bytes = &octets[..if ip_version == 4 { 4 } else { 16 }];
1210 for prefix in [1, 2, 7, 8, 9, 17, 31, 32, 63, 64, 65, 127, 128] {
1211 if prefix > query_bytes.len() * 8 {
1212 continue;
1213 }
1214 let count = prefix as u64;
1215 let data_pointer = count + 16;
1216 let mut nodes = Vec::new();
1217 for bit in 0..prefix {
1218 let side = (query_bytes[bit / 8] >> (7 - bit % 8)) & 1;
1219 let next = if bit + 1 == prefix {
1220 data_pointer
1221 } else {
1222 (bit + 1) as u64
1223 };
1224 nodes.push(if side == 0 {
1225 (next, count)
1226 } else {
1227 (count, next)
1228 });
1229 }
1230 let reader = crafted_reader(
1231 record_size,
1232 count,
1233 node_stream(record_size as u8, &nodes),
1234 vec![0x42, b'a', b'b'],
1235 ip_version,
1236 None,
1237 );
1238 let traverse = |bytes: &[u8]| {
1239 if ip_version == 4 {
1240 reader.traverse_ipv4(0, bytes.try_into().unwrap(), 0)
1241 } else {
1242 reader.traverse_ipv6(0, bytes.try_into().unwrap(), 0)
1243 }
1244 };
1245 assert_eq!(traverse(query_bytes), Some((data_pointer, prefix as u8)));
1246 assert_eq!(
1247 reader.lookup_value_with_prefix(query).unwrap(),
1248 (ValueRef::Utf8("ab"), prefix as u8)
1249 );
1250 for bit in [0, prefix / 2, prefix - 1] {
1253 let mut miss = query_bytes.to_vec();
1254 miss[bit / 8] ^= 0x80 >> (bit % 8);
1255 assert_eq!(traverse(&miss), None);
1256 }
1257 }
1258 }
1259 }
1260 }
1261
1262 #[test]
1263 fn prepared_traversal_terminal_records_and_cycles_do_not_load_children() {
1264 for record_size in [24, 28, 32, 36, 40, 44, 48, 52, 56, 60, 64] {
1265 let max_record = if record_size == 64 {
1266 u64::MAX
1267 } else {
1268 (1_u64 << record_size) - 1
1269 };
1270 let reader = crafted_reader(
1271 record_size,
1272 3,
1273 node_stream(record_size as u8, &[(1, 3), (2, 3), (19, max_record)]),
1274 vec![0x42, b'a', b'b'],
1275 6,
1276 Some(0),
1277 );
1278 assert_eq!(reader.traverse_ipv6(1, &[0; 16], 10), Some((19, 12)));
1279 assert_eq!(
1280 reader.traverse_ipv6(2, &[0x80; 16], 0),
1281 Some((max_record, 1))
1282 );
1283 assert_eq!(
1284 reader.traverse_ipv6(max_record, &[0; 16], 7),
1285 Some((max_record, 7))
1286 );
1287 assert_eq!(reader.traverse_ipv6(3, &[0; 16], 0), None);
1288 assert!(matches!(
1289 reader.resolve_offset(ip("2000::")),
1290 Err(Error::InvalidOffset(_))
1291 ));
1292
1293 let reserved = crafted_reader(
1295 record_size,
1296 1,
1297 node_stream(record_size as u8, &[(2, 2)]),
1298 vec![0x42, b'a', b'b'],
1299 6,
1300 Some(0),
1301 );
1302 assert!(matches!(
1303 reserved.resolve_offset(ip("::")),
1304 Err(Error::InvalidOffset(_))
1305 ));
1306
1307 let cyclic = crafted_reader(
1308 record_size,
1309 1,
1310 node_stream(record_size as u8, &[(0, 0)]),
1311 Vec::new(),
1312 6,
1313 Some(0),
1314 );
1315 assert_eq!(cyclic.traverse_ipv6(0, &[0xa5; 16], 0), None);
1316 }
1317 }
1318
1319 #[test]
1320 fn read_record_all_sizes_and_errors() {
1321 for record_size in [24, 28, 32] {
1322 let nodes = [(1, 2), (19, 19), (3, 19)];
1325 let reader = crafted_reader(
1326 record_size,
1327 3,
1328 node_stream(record_size as u8, &nodes),
1329 vec![0x42, b'a', b'b'],
1330 6,
1331 Some(0),
1332 );
1333 assert_eq!(reader.read_record(0, 0).unwrap(), 1);
1334 assert_eq!(reader.read_record(0, 1).unwrap(), 2);
1335 assert_eq!(reader.read_record(1, 0).unwrap(), 19);
1336 assert_eq!(reader.read_record(1, 1).unwrap(), 19);
1337 assert_eq!(reader.read_record(2, 0).unwrap(), 3);
1338 assert_eq!(reader.read_record(2, 1).unwrap(), 19);
1339 assert!(matches!(
1340 reader.read_record(3, 0).unwrap_err(),
1341 Error::InvalidNode(3)
1342 ));
1343 assert!(matches!(
1344 reader.read_record(0, 2).unwrap_err(),
1345 Error::InvalidNode(_)
1346 ));
1347 }
1348 let reader = crafted_reader(
1350 40,
1351 1,
1352 node_stream(40, &[(17, 17)]),
1353 vec![0x42, b'a', b'b'],
1354 6,
1355 Some(0),
1356 );
1357 assert_eq!(reader.read_record(0, 0).unwrap(), 17);
1358 assert_eq!(reader.read_record(0, 1).unwrap(), 17);
1359 }
1360
1361 #[test]
1362 fn read_record_rejects_out_of_bounds_node() {
1363 let reader = crafted_reader(
1365 24,
1366 2,
1367 node_stream(24, &[(0, 0), (0, 0)]),
1368 vec![0x42, b'a', b'b'],
1369 6,
1370 Some(0),
1371 );
1372 assert!(matches!(
1373 reader.read_record(2, 0).unwrap_err(),
1374 Error::InvalidNode(2)
1375 ));
1376 }
1377
1378 #[test]
1379 fn record_to_file_offset_error_branches() {
1380 let reader = crafted_reader(
1381 24,
1382 3,
1383 node_stream(24, &[(0, 0), (0, 0), (0, 0)]),
1384 vec![0x42, b'a', b'b'],
1385 6,
1386 Some(0),
1387 );
1388 assert!(matches!(
1390 reader.record_to_file_offset(3).unwrap_err(),
1391 Error::InvalidOffset(3)
1392 ));
1393 assert!(matches!(
1395 reader.record_to_file_offset(u64::MAX).unwrap_err(),
1396 Error::InvalidOffset(usize::MAX)
1397 ));
1398 assert!(matches!(
1400 reader
1401 .record_to_file_offset(0xffff_ffff_ffff_ff00)
1402 .unwrap_err(),
1403 Error::InvalidOffset(_)
1404 ));
1405 let reader = crafted_reader(
1407 24,
1408 3,
1409 node_stream(24, &[(0, 0), (0, 0), (0, 0)]),
1410 vec![],
1411 6,
1412 Some(0),
1413 );
1414 assert!(matches!(
1416 reader.record_to_file_offset(24).unwrap_err(),
1417 Error::InvalidOffset(39)
1418 ));
1419 }
1420
1421 #[test]
1422 fn lookup_rejects_v6_in_v4_database() {
1423 let reader = crafted_reader(24, 1, Vec::new(), Vec::new(), 4, None);
1424 assert!(matches!(
1425 reader.lookup_value(ip("::1")).unwrap_err(),
1426 Error::NotFound
1427 ));
1428 }
1429
1430 #[test]
1431 fn lookup_rejects_invalid_ip_version() {
1432 let reader = crafted_reader(24, 1, Vec::new(), Vec::new(), 5, None);
1433 assert!(matches!(
1434 reader.lookup_value(ip("1.2.3.4")).unwrap_err(),
1435 Error::InvalidIpVersion(5)
1436 ));
1437 }
1438
1439 #[test]
1440 fn lookup_v4_in_v6_missing_start_node() {
1441 let reader = crafted_reader(24, 1, Vec::new(), Vec::new(), 6, None);
1443 assert!(matches!(
1444 reader.lookup_value(ip("1.2.3.4")).unwrap_err(),
1445 Error::NotFound
1446 ));
1447 }
1448
1449 fn metadata_db(ip_version: u16) -> Vec<u8> {
1450 let mut writer = Writer::with_metadata(
1451 MetadataBuilder::new()
1452 .database_type("reader-tests")
1453 .ip_version(ip_version)
1454 .build()
1455 .unwrap(),
1456 );
1457 writer
1458 .insert_value(
1459 if ip_version == 6 {
1460 "2001:db8::/32".parse().unwrap()
1461 } else {
1462 "10.0.0.0/8".parse().unwrap()
1463 },
1464 Value::Map(Into::into({
1465 let mut m = std::collections::BTreeMap::new();
1466 m.insert("name".into(), Value::Utf8("x".into()));
1467 m
1468 })),
1469 )
1470 .unwrap();
1471 writer.finish().unwrap()
1472 }
1473
1474 fn patch_uint16(db: &mut [u8], key: &str, value: u8) {
1475 let index = db
1476 .windows(key.len())
1477 .position(|w| w == key.as_bytes())
1478 .expect("key present in metadata");
1479 debug_assert_eq!(db[index + key.len()], 0xA1, "u16 encoded as control+1");
1480 db[index + key.len() + 1] = value;
1481 }
1482
1483 #[test]
1484 fn from_source_rejects_bad_metadata_values() {
1485 let mut bad_major = metadata_db(6);
1486 patch_uint16(&mut bad_major, "binary_format_major_version", 3);
1487 assert!(matches!(
1488 Reader::from_bytes(&bad_major).unwrap_err(),
1489 Error::InvalidMetadata("unsupported binary format major version")
1490 ));
1491
1492 let mut bad_ip = metadata_db(6);
1493 patch_uint16(&mut bad_ip, "ip_version", 5);
1494 assert!(matches!(
1495 Reader::from_bytes(&bad_ip).unwrap_err(),
1496 Error::InvalidIpVersion(5)
1497 ));
1498
1499 let mut bad_record_size = metadata_db(6);
1500 patch_uint16(&mut bad_record_size, "record_size", 16);
1501 assert!(matches!(
1502 Reader::from_bytes(&bad_record_size).unwrap_err(),
1503 Error::InvalidMetadata(msg) if msg.contains("record_size")
1504 ));
1505 }
1506
1507 #[test]
1508 fn from_source_rejects_broken_separator() {
1509 let db = metadata_db(6);
1510 let reader = Reader::from_bytes(&db).unwrap();
1511 let mut broken = db.clone();
1512 broken[reader.data_section_start - 1] = 0x01;
1513 assert!(matches!(
1514 Reader::from_bytes(&broken).unwrap_err(),
1515 Error::InvalidDatabase(msg) if msg.contains("separator")
1516 ));
1517 }
1518
1519 #[test]
1520 fn from_source_and_as_bytes_and_serde_lookup() {
1521 let db = metadata_db(6);
1522 assert_eq!(Reader::from_bytes(&db).unwrap().metadata().ip_version, 6);
1523 let reader = Reader::from_bytes(&db).unwrap();
1524 let prepared = &reader.prepared_tree as *const PreparedTree;
1525 assert_eq!(reader.as_bytes(), &db);
1526 let map: std::collections::BTreeMap<String, String> =
1527 reader.lookup(ip("2001:db8::1")).unwrap();
1528 assert_eq!(map["name"], "x");
1529 assert_eq!(&reader.prepared_tree as *const PreparedTree, prepared);
1530 }
1531
1532 #[test]
1533 fn open_owned_and_mmap_sources() {
1534 let dir = tempfile::tempdir().unwrap();
1535 let path = dir.path().join("db.mmdb");
1536 std::fs::write(&path, metadata_db(6)).unwrap();
1537 let owned = Reader::open(&path).unwrap();
1538 assert!(matches!(
1539 owned.lookup_value(ip("2001:db8::1")).unwrap(),
1540 ValueRef::Map(_)
1541 ));
1542 let mapped = unsafe { Reader::open_mmap(&path) }.unwrap();
1544 let (value, _) = mapped.lookup_value_with_prefix(ip("2001:db8::1")).unwrap();
1545 assert!(matches!(value, ValueRef::Map(_)));
1546 }
1547}