1use std::collections::HashMap;
4use std::hash::{BuildHasher, BuildHasherDefault, Hasher};
5use std::path::Path;
6
7use crate::encoder::encode_value;
8use crate::{Error, IpNetwork, Metadata, MmdbEncode, MmdbRecord, Result, Value};
9
10const METADATA_MARKER: &[u8] = b"\xab\xcd\xefMaxMind.com";
11
12#[derive(Default)]
18struct FxHasher {
19 hash: u64,
20}
21
22type FxHashMap<K, V> = HashMap<K, V, BuildHasherDefault<FxHasher>>;
23
24const FX_SEED: u64 = 0x51_7c_c1_b7_27_22_0a_95;
25
26impl FxHasher {
27 #[inline]
28 fn add(&mut self, value: u64) {
29 self.hash = (self.hash.rotate_left(5) ^ value).wrapping_mul(FX_SEED);
30 }
31}
32
33impl Hasher for FxHasher {
34 #[inline]
35 fn finish(&self) -> u64 {
36 self.hash
37 }
38
39 #[inline]
40 fn write(&mut self, bytes: &[u8]) {
41 let (chunks, remainder) = bytes.as_chunks::<8>();
42 for chunk in chunks {
43 self.add(u64::from_ne_bytes(*chunk));
44 }
45 if !remainder.is_empty() {
46 let mut tail = [0u8; 8];
47 tail[..remainder.len()].copy_from_slice(remainder);
48 self.add(u64::from_ne_bytes(tail));
49 }
50 }
51
52 #[inline]
53 fn write_u8(&mut self, value: u8) {
54 self.add(u64::from(value));
55 }
56
57 #[inline]
58 fn write_u16(&mut self, value: u16) {
59 self.add(u64::from(value));
60 }
61
62 #[inline]
63 fn write_u32(&mut self, value: u32) {
64 self.add(u64::from(value));
65 }
66
67 #[inline]
68 fn write_u64(&mut self, value: u64) {
69 self.add(value);
70 }
71
72 #[inline]
73 fn write_usize(&mut self, value: usize) {
74 self.add(value as u64);
75 }
76}
77
78#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
80pub enum MergeStrategy {
81 #[default]
83 Replace,
84 Append,
86 AppendUnique,
88 DeepMerge,
90}
91
92const NO_NODE: u32 = u32::MAX;
93
94#[derive(Debug, Clone, PartialEq)]
95struct TrieNode {
96 children: [u32; 2],
97 value: Option<std::sync::Arc<Value>>,
98 index: Option<u32>,
99}
100
101impl Default for TrieNode {
102 #[inline]
103 fn default() -> Self {
104 Self {
105 children: [NO_NODE, NO_NODE],
106 value: None,
107 index: None,
108 }
109 }
110}
111
112impl TrieNode {
113 #[inline]
114 fn has_children(&self) -> bool {
115 self.children[0] != NO_NODE || self.children[1] != NO_NODE
116 }
117}
118
119#[derive(Debug)]
129pub struct Writer {
130 nodes: Vec<TrieNode>,
133 root: u32,
135 metadata: Metadata,
136 merge_strategy: MergeStrategy,
137}
138
139impl Writer {
140 #[must_use]
154 pub fn with_metadata(metadata: Metadata) -> Self {
155 Self::with_metadata_and_capacity(metadata, 1024)
156 }
157
158 #[must_use]
173 pub fn with_metadata_and_capacity(metadata: Metadata, capacity: usize) -> Self {
174 let mut nodes = Vec::with_capacity(capacity.max(1));
175 nodes.push(TrieNode::default());
176 Self {
177 nodes,
178 root: 0,
179 metadata,
180 merge_strategy: MergeStrategy::Replace,
181 }
182 }
183
184 #[must_use]
200 pub fn merge_strategy(mut self, strategy: MergeStrategy) -> Self {
201 self.merge_strategy = strategy;
202 self
203 }
204
205 pub fn insert<T: serde::Serialize + ?Sized>(
220 &mut self,
221 network: IpNetwork,
222 value: &T,
223 ) -> Result<()> {
224 self.insert_value(network, Value::from_serialize(value)?)
225 }
226
227 pub fn insert_encoded<T: MmdbEncode + ?Sized>(
246 &mut self,
247 network: IpNetwork,
248 value: &T,
249 ) -> Result<()> {
250 self.insert_value(network, value.encode()?)
251 }
252
253 pub fn insert_entry<T: MmdbRecord + ?Sized>(&mut self, entry: &T) -> Result<()> {
280 self.insert_value(entry.network(), entry.encode()?)
281 }
282
283 pub fn insert_value(&mut self, network: IpNetwork, value: Value) -> Result<()> {
298 let node_index = descend_network(
299 &mut self.nodes,
300 self.root,
301 network,
302 self.metadata.ip_version,
303 )?;
304 let strategy = self.merge_strategy;
305 let node = &mut self.nodes[node_index as usize];
306 node.value = Some(if strategy == MergeStrategy::Replace {
307 std::sync::Arc::new(value)
308 } else {
309 match node.value.take() {
310 Some(old) => std::sync::Arc::new(merge_values((*old).clone(), value, strategy)),
311 None => std::sync::Arc::new(value),
312 }
313 });
314 Ok(())
315 }
316
317 #[inline(always)]
335 pub fn insert_value_arc(
336 &mut self,
337 network: IpNetwork,
338 value: std::sync::Arc<Value>,
339 ) -> Result<()> {
340 self.insert_value_shared(network, value)
341 }
342
343 pub fn insert_value_shared(
361 &mut self,
362 network: IpNetwork,
363 value: std::sync::Arc<Value>,
364 ) -> Result<()> {
365 let node_index = descend_network(
366 &mut self.nodes,
367 self.root,
368 network,
369 self.metadata.ip_version,
370 )?;
371 let strategy = self.merge_strategy;
372 let node = &mut self.nodes[node_index as usize];
373 node.value = Some(if strategy == MergeStrategy::Replace {
374 value
375 } else {
376 match node.value.take() {
377 Some(old) => {
378 std::sync::Arc::new(merge_values((*old).clone(), (*value).clone(), strategy))
379 }
380 None => value,
381 }
382 });
383 Ok(())
384 }
385
386 pub fn insert_batch_shared<I>(
403 &mut self,
404 networks: I,
405 value: &std::sync::Arc<Value>,
406 ) -> Result<()>
407 where
408 I: IntoIterator<Item = IpNetwork>,
409 {
410 let iter = networks.into_iter();
411 let (lower, _) = iter.size_hint();
412 if lower > 0 {
413 self.nodes.reserve(lower.saturating_mul(4).min(2_000_000));
414 }
415 for net in iter {
416 self.insert_value_shared(net, std::sync::Arc::clone(value))?;
417 }
418 Ok(())
419 }
420
421 pub fn insert_batch<I>(&mut self, entries: I) -> Result<()>
436 where
437 I: IntoIterator<Item = (IpNetwork, Value)>,
438 {
439 for (net, val) in entries {
440 self.insert_value(net, val)?;
441 }
442 Ok(())
443 }
444
445 pub fn finish(mut self) -> Result<Vec<u8>> {
462 let mut next = 0_u64;
464 assign_indices(&mut self.nodes, self.root, &mut next);
465 let node_count = next;
466
467 if node_count > u64::from(u32::MAX) {
468 return Err(Error::EncodingError(
469 "MMDB node_count exceeds uint32".into(),
470 ));
471 }
472
473 let mut pool = DataPool::default();
474 let mut records: Vec<[u32; 2]> = vec![[0, 0]; node_count as usize];
475 fill_records(
476 &self.nodes,
477 self.root,
478 None,
479 &mut None,
480 &mut records,
481 node_count,
482 &mut pool,
483 )?;
484 drop(std::mem::take(&mut self.nodes));
487
488 let max_pointer = node_count
489 .checked_add(16)
490 .and_then(|v| v.checked_add(pool.data.len() as u64))
491 .ok_or_else(|| Error::EncodingError("MMDB pointer overflow".into()))?;
492 let record_size = if max_pointer < (1_u64 << 24) {
493 24
494 } else if max_pointer < (1_u64 << 28) {
495 28
496 } else if max_pointer <= u64::from(u32::MAX) {
497 32
498 } else {
499 return Err(Error::EncodingError(
500 "MMDB data section exceeds 32-bit pointer space".into(),
501 ));
502 };
503
504 self.metadata.node_count = node_count;
505 self.metadata.record_size = record_size;
506
507 let mut out = Vec::with_capacity(
508 records.len() * (record_size as usize / 4) + 16 + pool.data.len() + 256,
509 );
510 for [left, right] in records {
511 encode_node(left as u64, right as u64, record_size, &mut out)?;
512 }
513 out.extend_from_slice(&[0_u8; 16]);
514 out.extend_from_slice(&pool.data);
515 out.extend_from_slice(METADATA_MARKER);
516 encode_value(&self.metadata.to_value(), &mut out)?;
517 Ok(out)
518 }
519
520 pub fn write_to_file(self, path: impl AsRef<Path>) -> Result<()> {
537 std::fs::write(path, self.finish()?)?;
538 Ok(())
539 }
540}
541
542#[derive(Default)]
543struct DataPool {
544 data: Vec<u8>,
545 offsets: FxHashMap<u64, usize>,
547 entries: Vec<PoolEntry>,
548 arc_offsets: FxHashMap<usize, usize>,
549}
550
551struct PoolEntry {
552 offset: usize,
553 len: usize,
554 next: Option<usize>,
555}
556
557impl DataPool {
558 fn intern(&mut self, value: &Value) -> Result<usize> {
559 let start = self.data.len();
560 if let Err(error) = encode_value(value, &mut self.data) {
561 self.data.truncate(start);
562 return Err(error);
563 }
564 let hash = self.offsets.hasher().hash_one(&self.data[start..]);
565 Ok(self.intern_encoded(start, hash))
566 }
567
568 fn intern_arc(&mut self, arc: &std::sync::Arc<Value>) -> Result<usize> {
569 let ptr = std::sync::Arc::as_ptr(arc) as usize;
570 if let Some(&offset) = self.arc_offsets.get(&ptr) {
571 return Ok(offset);
572 }
573 let offset = self.intern(arc.as_ref())?;
574 self.arc_offsets.insert(ptr, offset);
575 Ok(offset)
576 }
577
578 fn intern_encoded(&mut self, start: usize, hash: u64) -> usize {
579 let head = self.offsets.entry(hash).or_insert(self.entries.len());
580 let mut candidate = Some(*head);
581 while let Some(index) = candidate {
582 let Some(entry) = self.entries.get(index) else {
583 break;
584 };
585 if self.data[entry.offset..entry.offset + entry.len] == self.data[start..] {
586 let offset = entry.offset;
587 self.data.truncate(start);
589 return offset;
590 }
591 candidate = entry.next;
592 }
593 let next = (*head < self.entries.len()).then_some(*head);
594 *head = self.entries.len();
595 self.entries.push(PoolEntry {
596 offset: start,
597 len: self.data.len() - start,
598 next,
599 });
600 start
601 }
602}
603
604#[cold]
609fn assign_indices(nodes: &mut [TrieNode], node_index: u32, next: &mut u64) {
610 if node_index == 0 || nodes[node_index as usize].has_children() {
611 nodes[node_index as usize].index = Some(*next as u32);
612 *next += 1;
613 let c0 = nodes[node_index as usize].children[0];
614 let c1 = nodes[node_index as usize].children[1];
615 if c0 != NO_NODE {
616 assign_indices(nodes, c0, next);
617 }
618 if c1 != NO_NODE {
619 assign_indices(nodes, c1, next);
620 }
621 }
622}
623
624#[cold]
628fn fill_records(
629 nodes: &[TrieNode],
630 node_index: u32,
631 inherited: Option<&std::sync::Arc<Value>>,
632 inherited_offset: &mut Option<usize>,
633 records: &mut [[u32; 2]],
634 node_count: u64,
635 pool: &mut DataPool,
636) -> Result<()> {
637 let node = &nodes[node_index as usize];
638 let index =
639 node.index
640 .ok_or_else(|| Error::EncodingError("indexed node expected".into()))? as usize;
641 let mut own_offset = None;
642 let (effective, cached_offset) = match node.value.as_ref() {
643 Some(value) => (Some(value), &mut own_offset),
644 None => (inherited, inherited_offset),
645 };
646 let mut targets = [node_count as u32, node_count as u32];
647 let c0 = node.children[0];
648 let c1 = node.children[1];
649 for (side, child_idx) in [(0, c0), (1, c1)] {
650 targets[side] = if child_idx != NO_NODE {
651 let child = &nodes[child_idx as usize];
652 if child.has_children() {
653 fill_records(
654 nodes,
655 child_idx,
656 effective,
657 cached_offset,
658 records,
659 node_count,
660 pool,
661 )?;
662 child
663 .index
664 .ok_or_else(|| Error::EncodingError("child index missing".into()))?
665 } else {
666 match child.value.as_ref() {
668 Some(value) => {
669 let offset = pool.intern_arc(value)?;
670 resolve_data_pointer(offset, node_count)?
671 }
672 None => match effective {
673 Some(value) => {
674 let offset = match *cached_offset {
675 Some(offset) => offset,
676 None => {
677 let offset = pool.intern_arc(value)?;
678 *cached_offset = Some(offset);
679 offset
680 }
681 };
682 resolve_data_pointer(offset, node_count)?
683 }
684 None => node_count as u32,
685 },
686 }
687 }
688 } else {
689 match effective {
690 Some(value) => {
691 let offset = match *cached_offset {
692 Some(offset) => offset,
693 None => {
694 let offset = pool.intern_arc(value)?;
695 *cached_offset = Some(offset);
696 offset
697 }
698 };
699 resolve_data_pointer(offset, node_count)?
700 }
701 None => node_count as u32,
702 }
703 };
704 }
705 records[index] = targets;
706 Ok(())
707}
708
709#[inline(always)]
710fn resolve_data_pointer(offset: usize, node_count: u64) -> Result<u32> {
711 let ptr = node_count
712 .checked_add(16)
713 .and_then(|v| v.checked_add(offset as u64))
714 .ok_or_else(|| Error::EncodingError("MMDB data pointer overflow".into()))?;
715 if ptr > u64::from(u32::MAX) {
716 return Err(Error::EncodingError(
717 "MMDB data section exceeds 32-bit pointer space".into(),
718 ));
719 }
720 Ok(ptr as u32)
721}
722
723fn encode_node(left: u64, right: u64, record_size: u16, out: &mut Vec<u8>) -> Result<()> {
724 match record_size {
725 24 => {
726 if left >= 1 << 24 || right >= 1 << 24 {
727 return Err(Error::EncodingError("24-bit tree pointer overflow".into()));
728 }
729 let bytes = [
730 (left >> 16) as u8,
731 (left >> 8) as u8,
732 left as u8,
733 (right >> 16) as u8,
734 (right >> 8) as u8,
735 right as u8,
736 ];
737 out.extend_from_slice(&bytes);
738 }
739 28 => {
740 if left >= 1 << 28 || right >= 1 << 28 {
741 return Err(Error::EncodingError("28-bit tree pointer overflow".into()));
742 }
743 let bytes = [
744 (left >> 16) as u8,
745 (left >> 8) as u8,
746 left as u8,
747 (((left >> 24) & 0x0f) << 4 | ((right >> 24) & 0x0f)) as u8,
748 (right >> 16) as u8,
749 (right >> 8) as u8,
750 right as u8,
751 ];
752 out.extend_from_slice(&bytes);
753 }
754 32 => {
755 out.extend_from_slice(&(left as u32).to_be_bytes());
756 out.extend_from_slice(&(right as u32).to_be_bytes());
757 }
758 _ => {
759 return Err(Error::EncodingError(
760 "writer supports 24/28/32-bit records".into(),
761 ));
762 }
763 }
764 Ok(())
765}
766
767#[inline(never)]
770fn descend_network(
771 nodes: &mut Vec<TrieNode>,
772 root: u32,
773 network: IpNetwork,
774 db_ip_version: u16,
775) -> Result<u32> {
776 match (db_ip_version, network) {
777 (4, IpNetwork::V4(net)) => {
778 let ip = u32::from(net.network());
779 Ok(descend_prefix_u32(
780 nodes,
781 root,
782 ip,
783 usize::from(net.prefix_len()),
784 ))
785 }
786 (4, IpNetwork::V6(_)) => Err(Error::InvalidIpVersion(6)),
787 (6, IpNetwork::V6(net)) => {
788 let ip = u128::from(net.network());
789 Ok(descend_prefix_u128(
790 nodes,
791 root,
792 ip,
793 usize::from(net.prefix_len()),
794 ))
795 }
796 (6, IpNetwork::V4(net)) => {
797 let mut node_index = root;
799 for _ in 0..96 {
800 let bit = 0; let child = nodes[node_index as usize].children[bit];
802 node_index = if child != NO_NODE {
803 child
804 } else {
805 let new_idx = nodes.len() as u32;
806 nodes.push(TrieNode::default());
807 nodes[node_index as usize].children[bit] = new_idx;
808 new_idx
809 };
810 }
811 let ip = u32::from(net.network());
812 Ok(descend_prefix_u32(
813 nodes,
814 node_index,
815 ip,
816 usize::from(net.prefix_len()),
817 ))
818 }
819 (version, _) => Err(Error::InvalidIpVersion(version)),
820 }
821}
822
823#[inline(always)]
824fn descend_prefix_u32(
825 nodes: &mut Vec<TrieNode>,
826 mut node_index: u32,
827 ip: u32,
828 prefix: usize,
829) -> u32 {
830 for shift in (32 - prefix..32).rev() {
831 let bit = ((ip >> shift) & 1) as usize;
832 let child = nodes[node_index as usize].children[bit];
833 node_index = if child != NO_NODE {
834 child
835 } else {
836 let new_idx = nodes.len() as u32;
837 nodes.push(TrieNode::default());
838 nodes[node_index as usize].children[bit] = new_idx;
839 new_idx
840 };
841 }
842 node_index
843}
844
845#[inline(always)]
846fn descend_prefix_u128(
847 nodes: &mut Vec<TrieNode>,
848 mut node_index: u32,
849 ip: u128,
850 prefix: usize,
851) -> u32 {
852 for shift in (128 - prefix..128).rev() {
853 let bit = ((ip >> shift) & 1) as usize;
854 let child = nodes[node_index as usize].children[bit];
855 node_index = if child != NO_NODE {
856 child
857 } else {
858 let new_idx = nodes.len() as u32;
859 nodes.push(TrieNode::default());
860 nodes[node_index as usize].children[bit] = new_idx;
861 new_idx
862 };
863 }
864 node_index
865}
866
867fn append_unique(left: &mut Vec<Value>, right: Vec<Value>) {
873 const SCAN_LIMIT: usize = 64;
874 if left.len().saturating_mul(right.len()) <= SCAN_LIMIT {
876 for value in right {
877 if !left.contains(&value) {
878 left.push(value);
879 }
880 }
881 return;
882 }
883
884 let mut seen: FxHashMap<u64, Vec<usize>> = FxHashMap::default();
887 seen.reserve(left.len() + right.len());
888
889 for (index, value) in left.iter().enumerate() {
891 seen.entry(hash_value(value)).or_default().push(index);
892 }
893
894 for value in right {
896 let hash = hash_value(&value);
897 let duplicate = seen
899 .get(&hash)
900 .is_some_and(|indices| indices.iter().any(|&index| left[index] == value));
901 if duplicate {
902 continue;
903 }
904 seen.entry(hash).or_default().push(left.len());
906 left.push(value);
907 }
908}
909
910fn hash_value(value: &Value) -> u64 {
911 let mut hasher = FxHasher::default();
912 hash_value_into(value, &mut hasher);
913 hasher.finish()
914}
915
916fn hash_value_into<H: Hasher>(value: &Value, state: &mut H) {
917 match value {
918 Value::Utf8(v) => {
919 state.write_u8(0);
920 state.write(v.as_bytes());
921 }
922 Value::Bytes(v) => {
923 state.write_u8(1);
924 state.write(v);
925 }
926 Value::Double(v) => {
927 state.write_u8(2);
928 state.write_u64(v.to_bits());
929 }
930 Value::Float(v) => {
931 state.write_u8(3);
932 state.write_u32(v.to_bits());
933 }
934 Value::Uint16(v) => {
935 state.write_u8(4);
936 state.write_u16(*v);
937 }
938 Value::Uint32(v) => {
939 state.write_u8(5);
940 state.write_u32(*v);
941 }
942 Value::Int32(v) => {
943 state.write_u8(6);
944 state.write_i32(*v);
945 }
946 Value::Uint64(v) => {
947 state.write_u8(7);
948 state.write_u64(*v);
949 }
950 Value::Uint128(v) => {
951 state.write_u8(8);
952 state.write_u128(*v);
953 }
954 Value::Bool(v) => {
955 state.write_u8(9);
956 state.write_u8(u8::from(*v));
957 }
958 Value::Array(values) => {
959 state.write_u8(10);
960 state.write_usize(values.len());
961 for value in values {
962 hash_value_into(value, state);
963 }
964 }
965 Value::Map(entries) => {
966 state.write_u8(11);
967 state.write_usize(entries.len());
968 for (key, value) in entries {
969 state.write(key.as_bytes());
970 hash_value_into(value, state);
971 }
972 }
973 }
974}
975
976fn merge_values(old: Value, new: Value, strategy: MergeStrategy) -> Value {
977 match strategy {
978 MergeStrategy::Replace => new,
979 MergeStrategy::Append => match (old, new) {
980 (Value::Array(mut left), Value::Array(right)) => {
981 left.extend(right);
982 Value::Array(left)
983 }
984 (_, new) => new,
985 },
986 MergeStrategy::AppendUnique => match (old, new) {
987 (Value::Array(mut left), Value::Array(right)) => {
988 append_unique(&mut left, right);
989 Value::Array(left)
990 }
991 (_, new) => new,
992 },
993 MergeStrategy::DeepMerge => match (old, new) {
994 (Value::Map(mut left), Value::Map(right)) => {
995 for (key, value) in right {
996 match left.remove(&key) {
997 Some(previous) => {
998 left.insert(
999 key,
1000 merge_values(previous, value, MergeStrategy::DeepMerge),
1001 );
1002 }
1003 None => {
1004 left.insert(key, value);
1005 }
1006 }
1007 }
1008 Value::Map(left)
1009 }
1010 (Value::Array(mut left), Value::Array(right)) => {
1011 left.extend(right);
1012 Value::Array(left)
1013 }
1014 (_, new) => new,
1015 },
1016 }
1017}
1018
1019#[cfg(test)]
1020mod tests {
1021 use std::collections::BTreeMap;
1022
1023 use super::*;
1024
1025 #[test]
1026 fn replace_duplicate_reuses_the_new_shared_value() {
1027 let metadata = crate::MetadataBuilder::new().ip_version(4).build().unwrap();
1028 let mut writer = Writer::with_metadata(metadata);
1029 let network = "198.51.100.0/24".parse().unwrap();
1030 let old = std::sync::Arc::new(Value::Utf8("old".into()));
1031 let new = std::sync::Arc::new(Value::Utf8("new".into()));
1032
1033 writer.insert_value_shared(network, old.clone()).unwrap();
1034 writer.insert_value_shared(network, new.clone()).unwrap();
1035
1036 assert_eq!(std::sync::Arc::strong_count(&old), 1);
1037 assert!(
1038 writer
1039 .nodes
1040 .iter()
1041 .filter_map(|node| node.value.as_ref())
1042 .any(|stored| std::sync::Arc::ptr_eq(stored, &new))
1043 );
1044 assert!(!writer.finish().unwrap().is_empty());
1045 }
1046
1047 #[test]
1048 fn inherited_cache_preserves_first_use_order_and_tree_records() {
1049 fn reference(
1050 nodes: &[TrieNode],
1051 node_index: u32,
1052 inherited: Option<&std::sync::Arc<Value>>,
1053 records: &mut [[u32; 2]],
1054 node_count: u64,
1055 pool: &mut DataPool,
1056 ) {
1057 let node = &nodes[node_index as usize];
1058 let effective = node.value.as_ref().or(inherited);
1059 let mut targets = [node_count as u32, node_count as u32];
1060 let child_indices = [(0, node.children[0]), (1, node.children[1])];
1061 for (side, child_idx) in child_indices {
1062 targets[side] = if child_idx != NO_NODE {
1063 let child = &nodes[child_idx as usize];
1064 if child.has_children() {
1065 reference(nodes, child_idx, effective, records, node_count, pool);
1066 child.index.unwrap()
1067 } else {
1068 match child.value.as_ref().or(effective) {
1069 Some(value) => {
1070 let offset = pool.intern_arc(value).unwrap();
1071 (node_count + 16 + offset as u64) as u32
1072 }
1073 None => node_count as u32,
1074 }
1075 }
1076 } else {
1077 match effective {
1078 Some(value) => {
1079 let offset = pool.intern_arc(value).unwrap();
1080 (node_count + 16 + offset as u64) as u32
1081 }
1082 None => node_count as u32,
1083 }
1084 };
1085 }
1086 records[node.index.unwrap() as usize] = targets;
1087 }
1088 let mut writer = Writer::with_metadata(crate::MetadataBuilder::new().build().unwrap());
1089 for (network, text) in [
1090 ("::/0", "root"),
1091 ("2001:db8::/32", "v6"),
1092 ("2001:db8:1::/48", "leaf"),
1093 ("10.0.0.0/8", "parent"),
1094 ("10.1.0.0/16", "child"),
1095 ("10.1.2.0/24", "leaf"),
1096 ("10.2.3.4/32", "leaf"),
1097 ("10.2.3.5/32", "parent"),
1098 ] {
1099 writer
1100 .insert_value(network.parse().unwrap(), Value::Utf8(text.into()))
1101 .unwrap();
1102 }
1103 let mut count = 0;
1104 assign_indices(&mut writer.nodes, writer.root, &mut count);
1105 let mut expected = vec![[0, 0]; count as usize];
1106 let mut actual = expected.clone();
1107 let mut old_pool = DataPool::default();
1108 let mut new_pool = DataPool::default();
1109 reference(
1110 &writer.nodes,
1111 writer.root,
1112 None,
1113 &mut expected,
1114 count,
1115 &mut old_pool,
1116 );
1117 fill_records(
1118 &writer.nodes,
1119 writer.root,
1120 None,
1121 &mut None,
1122 &mut actual,
1123 count,
1124 &mut new_pool,
1125 )
1126 .unwrap();
1127 assert_eq!(actual, expected);
1128 assert_eq!(new_pool.data, old_pool.data);
1129 }
1130
1131 #[test]
1132 fn pool_compares_bytes_even_when_hashes_collide() {
1133 let mut pool = DataPool::default();
1134 for (text, expected) in [("first", 0), ("second", 6), ("first", 0), ("second", 6)] {
1135 let start = pool.data.len();
1136 encode_value(&Value::Utf8(text.into()), &mut pool.data).unwrap();
1137 assert_eq!(pool.intern_encoded(start, 42), expected);
1138 }
1139 assert_eq!(pool.entries.len(), 2);
1140 assert_eq!(pool.data.len(), 13);
1141 }
1142
1143 #[test]
1144 fn pool_reuses_duplicates_and_preserves_distinct_integer_types() {
1145 let mut pool = DataPool::default();
1146 let a = pool.intern(&Value::Uint16(42)).unwrap();
1147 let b = pool.intern(&Value::Uint32(42)).unwrap();
1148 assert_ne!(a, b);
1149 assert_eq!(pool.intern(&Value::Uint16(42)).unwrap(), a);
1150 assert_eq!(pool.entries.len(), 2);
1151 }
1152
1153 #[test]
1154 fn packs_28_bit_nodes_per_mmdb_layout() {
1155 let mut out = Vec::new();
1156 encode_node(0x0123_4567, 0x0abc_def0, 28, &mut out).unwrap();
1157 assert_eq!(out, vec![0x23, 0x45, 0x67, 0x1a, 0xbc, 0xde, 0xf0]);
1158 }
1159
1160 #[test]
1161 fn append_unique_matches_scan_for_small_and_large_inputs() {
1162 fn merge(left: Vec<Value>, right: Vec<Value>) -> Vec<Value> {
1163 let Value::Array(merged) = merge_values(
1164 Value::Array(left),
1165 Value::Array(right),
1166 MergeStrategy::AppendUnique,
1167 ) else {
1168 panic!()
1169 };
1170 merged
1171 }
1172
1173 let small_left = vec![Value::Uint64(1), Value::Utf8("a".into())];
1174 let small_right = vec![Value::Uint64(1), Value::Uint64(2), Value::Utf8("a".into())];
1175 assert_eq!(
1176 merge(small_left.clone(), small_right.clone()),
1177 vec![Value::Uint64(1), Value::Utf8("a".into()), Value::Uint64(2)]
1178 );
1179
1180 let large_left: Vec<Value> = (0..50).map(Value::Uint64).collect();
1181 let mut large_right: Vec<Value> = (25..75).map(Value::Uint64).collect();
1182 large_right.push(Value::Utf8("new".into()));
1183 let merged = merge(large_left, large_right);
1184 assert_eq!(merged.len(), 76);
1185 assert_eq!(merged[49], Value::Uint64(49));
1186 assert_eq!(merged[50], Value::Uint64(50));
1187 assert_eq!(merged[75], Value::Utf8("new".into()));
1188 }
1189
1190 #[test]
1191 fn value_hash_distinguishes_variants_and_float_bits() {
1192 assert_ne!(hash_value(&Value::Uint16(1)), hash_value(&Value::Uint32(1)));
1193 assert_ne!(hash_value(&Value::Int32(-0)), hash_value(&Value::Uint64(0)));
1194 assert_ne!(
1195 hash_value(&Value::Double(0.0)),
1196 hash_value(&Value::Double(-0.0))
1197 );
1198 assert_eq!(
1199 hash_value(&Value::Double(0.0)),
1200 hash_value(&Value::Double(0.0))
1201 );
1202 for value in [
1204 Value::Bytes(vec![0, 255]),
1205 Value::Float(1.25),
1206 Value::Uint128(1 << 100),
1207 Value::Bool(true),
1208 Value::Array(vec![Value::Uint32(1)]),
1209 Value::Map(BTreeMap::from([("k".into(), Value::Uint16(1))])),
1210 ] {
1211 assert_eq!(hash_value(&value), hash_value(&value));
1212 }
1213 }
1214
1215 #[test]
1216 fn inherited_value_fills_empty_leaf_and_missing_sibling_once() {
1217 use std::sync::Arc;
1218
1219 let mut nodes = vec![TrieNode::default(), TrieNode::default()];
1220 nodes[0].index = Some(0);
1221 nodes[0].children[0] = 1;
1222 nodes[0].value = Some(Arc::new(Value::Utf8("parent".into())));
1223 let mut records = [[0, 0]];
1224 let mut pool = DataPool::default();
1225 fill_records(&nodes, 0, None, &mut None, &mut records, 1, &mut pool).unwrap();
1226 assert_eq!(records[0][0], records[0][1]);
1227 assert_eq!(pool.entries.len(), 1);
1228
1229 nodes[0].value = None;
1230 fill_records(
1231 &nodes,
1232 0,
1233 None,
1234 &mut None,
1235 &mut records,
1236 1,
1237 &mut DataPool::default(),
1238 )
1239 .unwrap();
1240 assert_eq!(records[0], [1, 1]);
1241 }
1242
1243 #[test]
1244 fn deep_merge_maps_and_arrays() {
1245 let old = Value::Map(BTreeMap::from([
1246 ("country".into(), Value::Utf8("FR".into())),
1247 (
1248 "categories".into(),
1249 Value::Array(vec![Value::Utf8("abuse".into())]),
1250 ),
1251 ]));
1252 let new = Value::Map(BTreeMap::from([
1253 ("city".into(), Value::Utf8("Paris".into())),
1254 (
1255 "categories".into(),
1256 Value::Array(vec![Value::Utf8("proxy".into())]),
1257 ),
1258 ]));
1259 let Value::Map(merged) = merge_values(old, new, MergeStrategy::DeepMerge) else {
1260 panic!()
1261 };
1262 assert_eq!(merged.get("country"), Some(&Value::Utf8("FR".into())));
1263 assert_eq!(merged.get("city"), Some(&Value::Utf8("Paris".into())));
1264 assert_eq!(
1265 merged.get("categories"),
1266 Some(&Value::Array(vec![
1267 Value::Utf8("abuse".into()),
1268 Value::Utf8("proxy".into())
1269 ]))
1270 );
1271 }
1272
1273 #[cfg(feature = "reader")]
1274 #[test]
1275 fn insertion_apis_preserve_distinct_networks_and_shared_values() {
1276 use crate::{MmdbEncode, MmdbRecord, Reader};
1277 use std::{net::IpAddr, sync::Arc};
1278
1279 struct Entry {
1280 network: IpNetwork,
1281 value: u32,
1282 }
1283 impl MmdbEncode for Entry {
1284 fn encode(&self) -> Result<Value> {
1285 Ok(Value::Uint32(self.value))
1286 }
1287 }
1288 impl MmdbRecord for Entry {
1289 fn network(&self) -> IpNetwork {
1290 self.network
1291 }
1292 }
1293
1294 let metadata = crate::MetadataBuilder::new().ip_version(4).build().unwrap();
1295 let mut writer = Writer::with_metadata(metadata);
1296 writer
1297 .insert("198.51.100.0/24".parse().unwrap(), &"serde")
1298 .unwrap();
1299 writer
1300 .insert_encoded(
1301 "198.51.101.0/24".parse().unwrap(),
1302 &Entry {
1303 network: "198.51.101.0/24".parse().unwrap(),
1304 value: 1,
1305 },
1306 )
1307 .unwrap();
1308 writer
1309 .insert_entry(&Entry {
1310 network: "198.51.102.0/24".parse().unwrap(),
1311 value: 2,
1312 })
1313 .unwrap();
1314 writer
1315 .insert_batch([("198.51.103.0/24".parse().unwrap(), Value::Uint32(3))])
1316 .unwrap();
1317 writer
1318 .insert_value("198.51.103.0/24".parse().unwrap(), Value::Uint32(3))
1319 .unwrap();
1320 let shared = Arc::new(Value::Uint32(4));
1321 writer
1322 .insert_batch_shared(
1323 [
1324 "198.51.104.0/24".parse().unwrap(),
1325 "198.51.105.0/24".parse().unwrap(),
1326 ],
1327 &shared,
1328 )
1329 .unwrap();
1330 writer
1331 .insert_value_arc("198.51.106.0/24".parse().unwrap(), Arc::clone(&shared))
1332 .unwrap();
1333 writer
1334 .insert_value_shared("198.51.107.0/24".parse().unwrap(), Arc::clone(&shared))
1335 .unwrap();
1336 writer
1337 .insert_value_shared(
1338 "198.51.107.0/24".parse().unwrap(),
1339 Arc::new(Value::Uint32(5)),
1340 )
1341 .unwrap();
1342
1343 let bytes = writer.finish().unwrap();
1344 let reader = Reader::from_bytes(&bytes).unwrap();
1345 for (last_octet, expected) in [
1346 (100, "serde"),
1347 (101, "1"),
1348 (102, "2"),
1349 (103, "3"),
1350 (104, "4"),
1351 (105, "4"),
1352 (106, "4"),
1353 (107, "5"),
1354 ] {
1355 let ip: IpAddr = format!("198.51.{last_octet}.1").parse().unwrap();
1356 let actual = reader.lookup_value(ip).unwrap().to_json();
1357 assert_eq!(actual.to_string().trim_matches('"'), expected);
1358 }
1359 }
1360
1361 #[test]
1362 fn batches_stop_on_family_mismatch_and_file_write_propagates_io_errors() {
1363 use std::sync::Arc;
1364
1365 let metadata = crate::MetadataBuilder::new().ip_version(4).build().unwrap();
1366 let mut writer = Writer::with_metadata(metadata.clone());
1367 let ipv4 = "198.51.100.0/24".parse().unwrap();
1368 let ipv6 = "2001:db8::/32".parse().unwrap();
1369 assert!(matches!(
1370 writer.insert_batch([(ipv4, Value::Bool(true)), (ipv6, Value::Bool(false)),]),
1371 Err(Error::InvalidIpVersion(6))
1372 ));
1373 assert!(matches!(
1374 writer.insert_batch_shared([ipv6], &Arc::new(Value::Bool(true))),
1375 Err(Error::InvalidIpVersion(6))
1376 ));
1377 writer
1378 .insert_batch_shared(std::iter::empty(), &Arc::new(Value::Bool(true)))
1379 .unwrap();
1380 let file = tempfile::tempdir().unwrap();
1381 let path = file.path().join("generated.mmdb");
1382 Writer::with_metadata(metadata.clone())
1383 .write_to_file(&path)
1384 .unwrap();
1385 assert!(!std::fs::read(&path).unwrap().is_empty());
1386 assert!(
1387 Writer::with_metadata(metadata)
1388 .write_to_file(file.path())
1389 .is_err()
1390 );
1391 }
1392
1393 #[test]
1394 fn merge_strategies_and_node_encodings_reject_invalid_inputs() {
1395 let writer = Writer::with_metadata(crate::MetadataBuilder::new().build().unwrap())
1396 .merge_strategy(MergeStrategy::DeepMerge);
1397 assert_eq!(writer.merge_strategy, MergeStrategy::DeepMerge);
1398 let left = Value::Array(vec![Value::Uint16(1), Value::Uint16(2)]);
1399 let right = Value::Array(vec![Value::Uint16(2), Value::Uint16(3)]);
1400 assert_eq!(
1401 merge_values(left.clone(), right.clone(), MergeStrategy::Replace),
1402 right
1403 );
1404 assert_eq!(
1405 merge_values(left.clone(), right.clone(), MergeStrategy::Append),
1406 Value::Array(vec![
1407 Value::Uint16(1),
1408 Value::Uint16(2),
1409 Value::Uint16(2),
1410 Value::Uint16(3)
1411 ])
1412 );
1413 assert_eq!(
1414 merge_values(left, right, MergeStrategy::AppendUnique),
1415 Value::Array(vec![Value::Uint16(1), Value::Uint16(2), Value::Uint16(3)])
1416 );
1417 assert_eq!(
1418 merge_values(
1419 Value::Bool(false),
1420 Value::Bool(true),
1421 MergeStrategy::DeepMerge
1422 ),
1423 Value::Bool(true)
1424 );
1425 assert_eq!(
1426 merge_values(
1427 Value::Bool(false),
1428 Value::Utf8("x".into()),
1429 MergeStrategy::Append
1430 ),
1431 Value::Utf8("x".into())
1432 );
1433 assert_eq!(
1434 merge_values(
1435 Value::Bool(false),
1436 Value::Utf8("x".into()),
1437 MergeStrategy::AppendUnique
1438 ),
1439 Value::Utf8("x".into())
1440 );
1441
1442 assert!(encode_node(1 << 24, 0, 24, &mut Vec::new()).is_err());
1443 assert!(encode_node(0, 1 << 28, 28, &mut Vec::new()).is_err());
1444 assert!(encode_node(0, 0, 20, &mut Vec::new()).is_err());
1445 let mut bytes = Vec::new();
1446 encode_node(0x1234_5678, 0x9abc_def0, 32, &mut bytes).unwrap();
1447 assert_eq!(bytes, [0x12, 0x34, 0x56, 0x78, 0x9a, 0xbc, 0xde, 0xf0]);
1448 assert!(resolve_data_pointer(usize::MAX, 1).is_err());
1449 assert!(resolve_data_pointer(u32::MAX as usize, 1).is_err());
1450 }
1451}