use std::collections::{BTreeMap, HashMap};
use std::io;
use std::net::IpAddr;
use std::time::{SystemTime, UNIX_EPOCH};
use ipnet::IpNet;
use crate::data_section::{DataOffset, DataSection};
use crate::error::Error;
use crate::metadata::Metadata;
use crate::net::{IpVersion, alias_networks, range_to_networks, to_tree_prefix};
use crate::options::{Ipv4Aliasing, MergeStrategy, MetadataPointers, ReservedNetworks};
use crate::pool::{ValueId, ValuePool};
use crate::record_size::RecordSize;
use crate::reserved;
use crate::tree::Tree;
use crate::value::Value;
const DATA_SECTION_SEPARATOR: [u8; 16] = [0; 16];
const METADATA_MARKER: &[u8; 14] = b"\xab\xcd\xefMaxMind.com";
fn default_languages() -> Vec<String> {
vec!["en".to_owned()]
}
#[derive(Debug)]
pub struct Writer {
database_type: String,
description: BTreeMap<String, String>,
languages: Vec<String>,
ip_version: IpVersion,
record_size: Option<RecordSize>,
ipv4_aliasing: Ipv4Aliasing,
reserved_networks: ReservedNetworks,
metadata_pointers: MetadataPointers,
build_epoch: Option<SystemTime>,
tree: Tree,
pool: ValuePool,
}
#[bon::bon]
impl Writer {
#[builder(builder_type = WriterBuilder, finish_fn = build)]
pub fn builder(
#[builder(start_fn, into)] database_type: String,
#[builder(default = default_languages(), with = |langs: impl IntoIterator<Item: Into<String>>| langs.into_iter().map(Into::into).collect())]
languages: Vec<String>,
#[builder(default, with = |entries: &[(&str, &str)]| entries
.iter()
.map(|(lang, text)| ((*lang).to_owned(), (*text).to_owned()))
.collect())]
description: BTreeMap<String, String>,
#[builder(default)]
ip_version: IpVersion,
record_size: Option<RecordSize>,
#[builder(default)]
ipv4_aliasing: Ipv4Aliasing,
#[builder(default)]
reserved_networks: ReservedNetworks,
#[builder(default)]
metadata_pointers: MetadataPointers,
build_epoch: Option<SystemTime>,
) -> Self {
let mut tree = Tree::new();
if reserved_networks.is_excluded() {
for network in reserved::networks(ip_version) {
let (bits, prefix_len) =
to_tree_prefix(network, ip_version).expect("reserved network fits the tree");
tree.paint_reserved(bits, prefix_len)
.expect("reserved network fits the tree");
}
}
Self {
database_type,
description,
languages,
ip_version,
record_size,
ipv4_aliasing,
reserved_networks,
metadata_pointers,
build_epoch,
tree,
pool: ValuePool::new(),
}
}
}
impl Writer {
#[must_use]
pub fn new(database_type: impl Into<String>) -> Self {
Self::builder(database_type).build()
}
#[cfg(feature = "serde")]
pub fn insert<N: Into<IpNet>, T: serde::Serialize + ?Sized>(
&mut self,
network: N,
value: &T,
) -> Result<(), Error> {
let value = crate::ser::to_value(value)?;
self.insert_value(network, value)
}
pub fn insert_value<N: Into<IpNet>>(&mut self, network: N, value: Value) -> Result<(), Error> {
let net = network.into();
self.ensure_insertable(net)?;
let (bits, prefix_len) = to_tree_prefix(net, self.ip_version)?;
let id = self.pool.intern(value);
self.tree.insert(bits, prefix_len, &mut |_| Some(id))?;
Ok(())
}
fn ensure_insertable(&self, net: IpNet) -> Result<(), Error> {
use ipnet::IpNet as N;
let contains = |outer: &N, inner: &N| outer.contains(inner);
if self.ip_version == IpVersion::V6
&& self.ipv4_aliasing.is_enabled()
&& alias_networks().iter().any(|a| contains(a, &net))
{
return Err(Error::AliasedNetwork(net));
}
if self.reserved_networks.is_excluded()
&& reserved::networks(self.ip_version)
.iter()
.any(|r| contains(r, &net))
{
return Err(Error::ReservedNetwork(net));
}
Ok(())
}
pub fn insert_with<N, F>(&mut self, network: N, mut op: F) -> Result<(), Error>
where
N: Into<IpNet>,
F: FnMut(Option<&Value>) -> Option<Value>,
{
let net = network.into();
self.ensure_insertable(net)?;
let (bits, prefix_len) = to_tree_prefix(net, self.ip_version)?;
let tree = &mut self.tree;
let pool = &mut self.pool;
tree.insert(bits, prefix_len, &mut |old_id| {
let old_value = old_id.map(|id| pool.get(id).clone());
op(old_value.as_ref()).map(|new_value| pool.intern(new_value))
})
}
pub fn insert_value_merged<N: Into<IpNet>>(
&mut self,
network: N,
value: Value,
strategy: MergeStrategy,
) -> Result<(), Error> {
match strategy {
MergeStrategy::Replace => self.insert_value(network, value),
MergeStrategy::TopLevelMerge => self.insert_with(network, |existing| {
Some(match existing {
Some(old) => Value::merge_top_level(old, &value),
None => value.clone(),
})
}),
MergeStrategy::DeepMerge => self.insert_with(network, |existing| {
Some(match existing {
Some(old) => Value::merge_deep(old, &value),
None => value.clone(),
})
}),
}
}
#[cfg(feature = "serde")]
pub fn insert_merged<N: Into<IpNet>, T: serde::Serialize + ?Sized>(
&mut self,
network: N,
value: &T,
strategy: MergeStrategy,
) -> Result<(), Error> {
let value = crate::ser::to_value(value)?;
self.insert_value_merged(network, value, strategy)
}
pub fn insert_range(&mut self, start: IpAddr, end: IpAddr, value: &Value) -> Result<(), Error> {
for network in range_to_networks(start, end)? {
self.insert_value(network, value.clone())?;
}
Ok(())
}
#[must_use]
pub fn get(&self, ip: IpAddr) -> Option<&Value> {
let (bits, _) = to_tree_prefix(IpNet::from(ip), self.ip_version).ok()?;
let id = self.tree.get(bits, self.ip_version.tree_depth())?;
Some(self.pool.get(id))
}
pub fn to_bytes(&self) -> Result<Vec<u8>, Error> {
self.build()
}
pub fn write_to<W: io::Write>(&self, mut writer: W) -> Result<(), Error> {
let bytes = self.build()?;
writer.write_all(&bytes)?;
Ok(())
}
fn build(&self) -> Result<Vec<u8>, Error> {
let mut tree = self.tree.clone();
if self.ip_version == IpVersion::V6 && self.ipv4_aliasing.is_enabled() {
tree.install_ipv4_aliases()?;
}
tree.compact();
let ids = tree.reachable_data_ids();
let mut data = DataSection::new();
let mut id_to_offset: HashMap<ValueId, DataOffset> = HashMap::with_capacity(ids.len());
for id in ids {
let offset = data.push(self.pool.get(id));
id_to_offset.insert(id, offset);
}
let data_len = data.len();
let data_bytes = data.into_bytes();
let record_size = pick_record_size(&tree, data_len, self.record_size)?;
let tree_bytes = tree.serialize(record_size, &id_to_offset)?;
let metadata = Metadata {
database_type: &self.database_type,
description: &self.description,
languages: &self.languages,
ip_version: self.ip_version,
record_size,
node_count: tree.node_count(),
build_epoch: self.build_epoch_secs(),
disable_pointers: self.metadata_pointers.is_disabled(),
};
let metadata_bytes = metadata.to_bytes();
let mut out = Vec::with_capacity(
tree_bytes.len()
+ DATA_SECTION_SEPARATOR.len()
+ data_bytes.len()
+ METADATA_MARKER.len()
+ metadata_bytes.len(),
);
out.extend_from_slice(&tree_bytes);
out.extend_from_slice(&DATA_SECTION_SEPARATOR);
out.extend_from_slice(&data_bytes);
out.extend_from_slice(METADATA_MARKER);
out.extend_from_slice(&metadata_bytes);
Ok(out)
}
fn build_epoch_secs(&self) -> u64 {
self.build_epoch
.unwrap_or_else(SystemTime::now)
.duration_since(UNIX_EPOCH)
.map_or(0, |d| d.as_secs())
}
}
fn pick_record_size(
tree: &Tree,
data_len: usize,
requested: Option<RecordSize>,
) -> Result<RecordSize, Error> {
if let Some(size) = requested {
return Ok(size);
}
for candidate in RecordSize::ASCENDING {
if tree.fits_record_size(candidate, data_len) {
return Ok(candidate);
}
}
Err(Error::TreeTooLarge {
node_count: tree.node_count(),
max: RecordSize::Bits32.max_value(),
record_size: RecordSize::Bits32,
})
}