use anyhow::Result;
use chrono::{DateTime, Utc};
use polars::prelude::*;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::path::Path;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PacketFeatures {
pub timestamp: DateTime<Utc>,
pub src_ip: String,
pub dst_ip: String,
pub src_port: Option<u16>,
pub dst_port: Option<u16>,
pub protocol: String,
pub packet_size: u32,
pub flags: HashMap<String, bool>,
pub payload_size: u32,
pub ttl: Option<u8>,
pub window_size: Option<u16>,
pub custom_features: HashMap<String, serde_json::Value>,
}
pub struct FeatureStore {
features: Vec<PacketFeatures>,
dataframe: Option<DataFrame>,
}
impl FeatureStore {
pub fn new() -> Self {
Self {
features: Vec::new(),
dataframe: None,
}
}
pub fn add_features(&mut self, features: PacketFeatures) {
self.features.push(features);
}
pub fn build_dataframe(&mut self) -> Result<()> {
if self.features.is_empty() {
return Ok(());
}
let mut timestamps = Vec::new();
let mut src_ips = Vec::new();
let mut dst_ips = Vec::new();
let mut src_ports = Vec::new();
let mut dst_ports = Vec::new();
let mut protocols = Vec::new();
let mut packet_sizes = Vec::new();
let mut payload_sizes = Vec::new();
for feature in &self.features {
timestamps.push(feature.timestamp.timestamp_millis());
src_ips.push(feature.src_ip.clone());
dst_ips.push(feature.dst_ip.clone());
src_ports.push(feature.src_port);
dst_ports.push(feature.dst_port);
protocols.push(feature.protocol.clone());
packet_sizes.push(feature.packet_size);
payload_sizes.push(feature.payload_size);
}
let df = DataFrame::new(vec![
Column::new("timestamp".into(), timestamps),
Column::new("src_ip".into(), src_ips),
Column::new("dst_ip".into(), dst_ips),
Column::new("src_port".into(), src_ports.into_iter().map(|p| p.unwrap_or(0) as u32).collect::<Vec<_>>()),
Column::new("dst_port".into(), dst_ports.into_iter().map(|p| p.unwrap_or(0) as u32).collect::<Vec<_>>()),
Column::new("protocol".into(), protocols),
Column::new("packet_size".into(), packet_sizes),
Column::new("payload_size".into(), payload_sizes),
])?;
self.dataframe = Some(df);
Ok(())
}
pub fn save_parquet<P: AsRef<Path>>(&mut self, path: P) -> Result<()> {
if self.dataframe.is_none() {
self.build_dataframe()?;
}
if let Some(df) = &mut self.dataframe {
let file = std::fs::File::create(path)?;
ParquetWriter::new(file).finish(df)?;
}
Ok(())
}
pub fn save_json<P: AsRef<Path>>(&self, path: P) -> Result<()> {
let file = std::fs::File::create(path)?;
serde_json::to_writer_pretty(file, &self.features)?;
Ok(())
}
pub fn save_csv<P: AsRef<Path>>(&mut self, path: P) -> Result<()> {
if self.dataframe.is_none() {
self.build_dataframe()?;
}
if let Some(df) = &mut self.dataframe {
let mut file = std::fs::File::create(path)?;
CsvWriter::new(&mut file).finish(df)?;
}
Ok(())
}
pub fn get_statistics(&self) -> FeatureStatistics {
let total_packets = self.features.len();
let mut protocol_counts: HashMap<String, usize> = HashMap::new();
let mut total_bytes = 0u64;
for feature in &self.features {
*protocol_counts.entry(feature.protocol.clone()).or_insert(0) += 1;
total_bytes += feature.packet_size as u64;
}
FeatureStatistics {
total_packets,
total_bytes,
protocol_distribution: protocol_counts,
unique_src_ips: self.features.iter()
.map(|f| f.src_ip.clone())
.collect::<std::collections::HashSet<_>>()
.len(),
unique_dst_ips: self.features.iter()
.map(|f| f.dst_ip.clone())
.collect::<std::collections::HashSet<_>>()
.len(),
}
}
}
#[derive(Debug, Serialize)]
pub struct FeatureStatistics {
pub total_packets: usize,
pub total_bytes: u64,
pub protocol_distribution: HashMap<String, usize>,
pub unique_src_ips: usize,
pub unique_dst_ips: usize,
}