use anyhow::Result;
use clap::Parser;
use futures::future::try_join_all;
use hickory_resolver::{Resolver, name_server::TokioConnectionProvider};
use indicatif::{HumanCount, ProgressBar, ProgressStyle};
use ip2asn::IpAsnMap;
use itertools::izip;
use std::{fs::File, iter::repeat_with, path::PathBuf, sync::Arc, time::SystemTime};
use tokio::{sync::mpsc, task::spawn};
use tracing::{Level, event};
use webinfo::{
IpInfo,
ipinfo::OriginRecord,
utils::{chunked, count_lines, get_resolver, open_asn_db},
};
fn get_writer(output: Option<PathBuf>) -> Box<dyn std::io::Write + Send> {
match output {
Some(path) => {
let file = File::create(path);
match file {
Err(e) => {
event!(Level::ERROR, "Failed to create output file: {}", e);
Box::new(std::io::stdout())
}
Ok(file) => Box::new(file),
}
}
None => Box::new(std::io::stdout()),
}
}
fn process_batch_of_records(
chunk: Vec<Result<OriginRecord, csv::Error>>,
resolver: &Resolver<TokioConnectionProvider>,
ip2asn_map: &Arc<IpAsnMap>,
tx: &mpsc::Sender<Result<IpInfo>>,
) -> Vec<tokio::task::JoinHandle<()>> {
let mut handles = Vec::new();
let resolver_iter = repeat_with(|| resolver.clone()).take(chunk.len());
let ip2asn_iter = repeat_with(|| ip2asn_map.clone()).take(chunk.len());
let tx_iter = repeat_with(|| tx.clone()).take(chunk.len());
for (record, r, ip2asn, sender) in izip!(chunk, resolver_iter, ip2asn_iter, tx_iter) {
let record = match record {
Ok(record) => record,
Err(e) => {
event!(Level::ERROR, "{}", e);
continue;
}
};
let handle = spawn(async move {
let ip_info = IpInfo::runner(record)
.with_resolver(r)
.with_ip2asn_map(ip2asn)
.run()
.await;
let _ = sender.send(ip_info).await;
});
handles.push(handle);
}
handles
}
#[derive(Parser)]
#[command(version, about, long_about = None, author = "Vincent Gauthier <vg@luxbulb.org>")]
struct Cli {
#[arg(short, long)]
csv: PathBuf,
#[arg(short = 's', long = "size", default_value_t = 5)]
chunk_size: usize,
#[arg(short = 'd', long = "dns")]
dns: Option<String>,
#[arg(short = 'l', long = "logfile", default_value = "./webinfo.log")]
logfile: PathBuf,
#[arg(short = 'o', long = "output")]
output: Option<PathBuf>,
}
async fn process_all_records(
mut rdr: csv::Reader<File>,
chunk_size: usize,
total_lines: usize,
custom_dns: Option<String>,
output: Option<PathBuf>,
) -> Result<()> {
let (tx, rx) = mpsc::channel::<Result<webinfo::IpInfo>>(chunk_size);
handle_result(rx, output);
let resolver = get_resolver(custom_dns)
.map_err(|_| anyhow::anyhow!("Failed to create DNS resolver with default configuration"))?;
let ip2asn_map = open_asn_db()
.await
.map_err(|e| anyhow::anyhow!("Failed to open ASN database: {}", e))?;
let ip2asn_map = Arc::new(ip2asn_map);
let bar = ProgressBar::new(total_lines as u64);
bar.set_style(ProgressStyle::with_template("[{bar:50.cyan/blue}] {msg}")?.progress_chars("= "));
let mut progress = 0;
for chunk in chunked(rdr.deserialize::<OriginRecord>(), chunk_size) {
let now = SystemTime::now();
let handles = process_batch_of_records(chunk, &resolver, &ip2asn_map, &tx);
let _ = try_join_all(handles).await?;
bar.inc(chunk_size as u64);
progress += chunk_size;
bar.set_message(format!(
"{}/{}, {} records processed in {:.2} seconds",
HumanCount(progress.try_into()?),
HumanCount(total_lines.try_into()?),
chunk_size,
now.elapsed().unwrap().as_secs_f64()
));
}
bar.finish();
Ok(())
}
fn handle_result(mut rx: mpsc::Receiver<Result<webinfo::IpInfo>>, output: Option<PathBuf>) {
let mut writer = get_writer(output);
tokio::spawn(async move {
while let Some(result) = rx.recv().await {
match result {
Ok(info) => {
writeln!(writer, "{}", serde_json::to_string_pretty(&info).unwrap())
.expect("Failed to write to output");
}
Err(e) => event!(Level::ERROR, "{}", e),
}
}
});
}
#[tokio::main]
async fn main() -> Result<()> {
let timer = tracing_subscriber::fmt::time::SystemTime;
let cli = Cli::parse();
let file_appender = tracing_appender::rolling::daily(
cli.logfile.parent().unwrap(),
cli.logfile.file_name().unwrap(),
);
let (non_blocking, _guard) = tracing_appender::non_blocking(file_appender);
let subscriber = tracing_subscriber::FmtSubscriber::builder()
.compact()
.with_timer(timer)
.with_writer(non_blocking)
.with_ansi(false)
.finish();
tracing::subscriber::set_global_default(subscriber)
.map_err(|_| anyhow::anyhow!("Failed to set global default subscriber"))?;
let csv_path = cli.csv;
let csv_path_str = csv_path
.to_str()
.ok_or_else(|| anyhow::anyhow!("Failed to convert CSV path to string"))?;
let line_count = count_lines(csv_path_str)?;
event!(
Level::INFO,
"Starting processing file: {:?} with {} lines",
csv_path,
line_count
);
let rdr = csv::Reader::from_path(&csv_path)?;
process_all_records(rdr, cli.chunk_size, line_count, cli.dns, cli.output).await?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use assert_fs::prelude::*;
#[tokio::test]
async fn test_process_batch_of_records() {
let resolver = Resolver::builder_tokio().unwrap().build();
let ip2asn_map = open_asn_db().await.unwrap();
let ip2asn_map = Arc::new(ip2asn_map);
let file = assert_fs::NamedTempFile::new("sample.txt").unwrap();
file.write_str(
"origin,popularity,date,country\nhttps://www.google.fr,1000,2025-08-28,FR\n",
)
.unwrap();
let mut rdr = csv::Reader::from_path(file.path()).unwrap();
let records = rdr.deserialize::<OriginRecord>().collect::<Vec<_>>();
let handles =
process_batch_of_records(records, &resolver, &ip2asn_map, &mpsc::channel(1).0);
assert_eq!(handles.len(), 1);
}
}