use core::fmt;
use std::io::{self, BufWriter, Write};
use std::path::Path;
use super::{IpVersion, METADATA_MARKER, RecordSize};
use ahash::HashMap;
use ipnet::IpNet;
#[derive(Debug)]
#[non_exhaustive]
pub enum MmdbWriteError {
ZeroPrefix,
FamilyMismatch,
OverlappingNetwork,
TooLarge,
Io(io::Error),
}
impl fmt::Display for MmdbWriteError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::ZeroPrefix => f.write_str("mmdb writer: zero-length prefix is not supported"),
Self::FamilyMismatch => {
f.write_str("mmdb writer: ip family does not match database ip_version")
}
Self::OverlappingNetwork => f.write_str("mmdb writer: overlapping networks"),
Self::TooLarge => f.write_str("mmdb writer: database exceeds 4 GiB addressing limit"),
Self::Io(err) => write!(f, "mmdb writer: i/o error: {err}"),
}
}
}
impl core::error::Error for MmdbWriteError {
fn source(&self) -> Option<&(dyn core::error::Error + 'static)> {
match self {
Self::Io(err) => Some(err),
_ => None,
}
}
}
impl From<io::Error> for MmdbWriteError {
fn from(err: io::Error) -> Self {
Self::Io(err)
}
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub(crate) enum MmdbValue {
Map(Vec<(String, Self)>),
Array(Vec<Self>),
String(String),
Double(f64),
U16(u16),
U32(u32),
U64(u64),
}
impl MmdbValue {
pub(crate) fn map<I, K>(pairs: I) -> Self
where
I: IntoIterator<Item = (K, Self)>,
K: Into<String>,
{
Self::Map(pairs.into_iter().map(|(k, v)| (k.into(), v)).collect())
}
pub(crate) fn string(s: impl Into<String>) -> Self {
Self::String(s.into())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
enum Record {
#[default]
Empty,
Node(u32),
Data(u32),
}
#[derive(Debug, Clone, Copy, Default)]
struct Node {
left: Record,
right: Record,
}
impl Node {
fn get(&self, bit: u8) -> Record {
if bit == 0 { self.left } else { self.right }
}
fn set(&mut self, bit: u8, rec: Record) {
if bit == 0 {
self.left = rec;
} else {
self.right = rec;
}
}
}
#[derive(Debug, Clone)]
pub struct MmdbBuilder {
ip_version: IpVersion,
record_size: RecordSize,
database_type: String,
languages: Vec<String>,
build_epoch: u64,
nodes: Vec<Node>,
data: Vec<u8>,
dedup: HashMap<Box<[u8]>, usize>,
}
impl MmdbBuilder {
#[must_use]
pub fn new(ip_version: IpVersion, database_type: impl Into<String>) -> Self {
Self {
ip_version,
record_size: RecordSize::Bits32,
database_type: database_type.into(),
languages: Vec::new(),
build_epoch: 0,
nodes: vec![Node::default()],
data: Vec::new(),
dedup: HashMap::default(),
}
}
#[must_use]
pub fn with_languages<I, S>(mut self, langs: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
self.languages = langs.into_iter().map(Into::into).collect();
self
}
#[must_use]
pub fn with_build_epoch(mut self, epoch: u64) -> Self {
self.build_epoch = epoch;
self
}
pub(crate) fn insert_value(
&mut self,
net: IpNet,
value: impl Into<MmdbValue>,
) -> Result<(), MmdbWriteError> {
let mut octets = [0u8; 16];
let nbits = match (self.ip_version, net) {
(IpVersion::V4, IpNet::V4(n)) => {
octets[..4].copy_from_slice(&n.addr().octets());
n.prefix_len() as usize
}
(IpVersion::V6, IpNet::V6(n)) => {
octets.copy_from_slice(&n.addr().octets());
n.prefix_len() as usize
}
(IpVersion::V6, IpNet::V4(n)) => {
octets[12..16].copy_from_slice(&n.addr().octets());
96 + n.prefix_len() as usize
}
(IpVersion::V4, IpNet::V6(_)) => return Err(MmdbWriteError::FamilyMismatch),
};
if nbits == 0 {
return Err(MmdbWriteError::ZeroPrefix);
}
let mut node = 0usize;
for i in 0..nbits - 1 {
let bit = (octets[i / 8] >> (7 - (i % 8))) & 1;
node = self.follow_or_create(node, bit)?;
}
let last = nbits - 1;
let final_bit = (octets[last / 8] >> (7 - (last % 8))) & 1;
if !matches!(self.nodes[node].get(final_bit), Record::Empty) {
return Err(MmdbWriteError::OverlappingNetwork);
}
let data_offset = self.append_data(&value.into())?;
let data_offset = u32::try_from(data_offset)
.ok()
.ok_or(MmdbWriteError::TooLarge)?;
self.nodes[node].set(final_bit, Record::Data(data_offset));
Ok(())
}
fn follow_or_create(&mut self, node: usize, bit: u8) -> Result<usize, MmdbWriteError> {
match self.nodes[node].get(bit) {
Record::Node(idx) => Ok(idx as usize),
Record::Empty => {
if self.nodes.len() >= u32::MAX as usize {
return Err(MmdbWriteError::TooLarge);
}
let new = self.nodes.len() as u32;
self.nodes.push(Node::default());
self.nodes[node].set(bit, Record::Node(new));
Ok(new as usize)
}
Record::Data(_) => Err(MmdbWriteError::OverlappingNetwork),
}
}
fn append_data(&mut self, value: &MmdbValue) -> Result<usize, MmdbWriteError> {
let mut encoded = Vec::new();
encode_inline(value, &mut encoded)?;
if let Some(&offset) = self.dedup.get(encoded.as_slice()) {
return Ok(offset);
}
let offset = self.data.len();
self.data.extend_from_slice(&encoded);
self.dedup.insert(encoded.into_boxed_slice(), offset);
Ok(offset)
}
pub fn build(&self) -> Result<Vec<u8>, MmdbWriteError> {
let mut out = Vec::new();
self.serialize_to(&mut out)?;
Ok(out)
}
fn serialize_to<W: Write>(&self, w: &mut W) -> Result<(), MmdbWriteError> {
let node_count = self.nodes.len() as u32;
for node in &self.nodes {
write_record(w, node.left, node_count)?;
write_record(w, node.right, node_count)?;
}
w.write_all(&[0u8; 16])?; w.write_all(&self.data)?;
w.write_all(METADATA_MARKER)?;
w.write_all(&self.encode_metadata(node_count)?)?;
Ok(())
}
pub fn write_to<W: Write>(&self, w: W) -> Result<(), MmdbWriteError> {
let mut w = BufWriter::new(w);
self.serialize_to(&mut w)?;
w.flush()?;
Ok(())
}
pub fn write_to_file(&self, path: impl AsRef<Path>) -> Result<(), MmdbWriteError> {
let file = std::fs::File::create(path)?;
let mut writer = BufWriter::new(file);
self.serialize_to(&mut writer)?;
writer.flush()?;
Ok(())
}
fn encode_metadata(&self, node_count: u32) -> Result<Vec<u8>, MmdbWriteError> {
let mut pairs = vec![
("node_count".to_owned(), MmdbValue::U32(node_count)),
(
"record_size".to_owned(),
MmdbValue::U16(self.record_size.bits()),
),
(
"ip_version".to_owned(),
MmdbValue::U16(self.ip_version.number()),
),
(
"database_type".to_owned(),
MmdbValue::String(self.database_type.clone()),
),
("binary_format_major_version".to_owned(), MmdbValue::U16(2)),
("binary_format_minor_version".to_owned(), MmdbValue::U16(0)),
("build_epoch".to_owned(), MmdbValue::U64(self.build_epoch)),
];
if !self.languages.is_empty() {
pairs.push((
"languages".to_owned(),
MmdbValue::Array(
self.languages
.iter()
.map(|l| MmdbValue::String(l.clone()))
.collect(),
),
));
}
let mut out = Vec::new();
encode_inline(&MmdbValue::Map(pairs), &mut out)?;
Ok(out)
}
}
fn write_record<W: Write>(w: &mut W, rec: Record, node_count: u32) -> Result<(), MmdbWriteError> {
let value: u32 = match rec {
Record::Node(idx) => idx,
Record::Empty => node_count,
Record::Data(off) => node_count
.checked_add(16)
.and_then(|v| v.checked_add(off))
.ok_or(MmdbWriteError::TooLarge)?,
};
w.write_all(&value.to_be_bytes())?;
Ok(())
}
fn encode_inline(value: &MmdbValue, out: &mut Vec<u8>) -> Result<(), MmdbWriteError> {
match value {
MmdbValue::Map(pairs) => {
encode_header(7, pairs.len(), out)?;
for (k, v) in pairs {
encode_string(k, out)?;
encode_inline(v, out)?;
}
}
MmdbValue::Array(items) => {
encode_header(11, items.len(), out)?;
for v in items {
encode_inline(v, out)?;
}
}
MmdbValue::String(s) => encode_string(s, out)?,
MmdbValue::Double(f) => {
encode_header(3, 8, out)?;
out.extend_from_slice(&f.to_be_bytes());
}
MmdbValue::U16(n) => encode_uint(5, u128::from(*n), out)?,
MmdbValue::U32(n) => encode_uint(6, u128::from(*n), out)?,
MmdbValue::U64(n) => encode_uint(9, u128::from(*n), out)?,
}
Ok(())
}
fn encode_string(s: &str, out: &mut Vec<u8>) -> Result<(), MmdbWriteError> {
encode_header(2, s.len(), out)?;
out.extend_from_slice(s.as_bytes());
Ok(())
}
fn encode_uint(type_num: u8, value: u128, out: &mut Vec<u8>) -> Result<(), MmdbWriteError> {
let bytes = min_be_bytes(value);
encode_header(type_num, bytes.len(), out)?;
out.extend_from_slice(&bytes);
Ok(())
}
fn min_be_bytes(value: u128) -> Vec<u8> {
if value == 0 {
return Vec::new();
}
let full = value.to_be_bytes();
let first = full.iter().position(|&b| b != 0).unwrap_or(full.len());
full[first..].to_vec()
}
fn encode_header(type_num: u8, size: usize, out: &mut Vec<u8>) -> Result<(), MmdbWriteError> {
let type_bits = if type_num <= 7 { type_num } else { 0 };
let (low5, ext): (u8, Vec<u8>) = if size <= 28 {
(size as u8, Vec::new())
} else if size <= 284 {
(29, vec![(size - 29) as u8])
} else if size <= 65820 {
(30, ((size - 285) as u16).to_be_bytes().to_vec())
} else {
let s = size - 65821;
if s > 0xFF_FFFF {
return Err(MmdbWriteError::TooLarge);
}
let s = s as u32;
(31, vec![(s >> 16) as u8, (s >> 8) as u8, s as u8])
};
out.push((type_bits << 5) | low5);
if type_num > 7 {
out.push(type_num - 7);
}
out.extend_from_slice(&ext);
Ok(())
}