use crate::ip::IpPacket;
use crate::sock::{Raw, RSockErr};
use std::sync::atomic::AtomicBool;
use std::sync::atomic::Ordering;
use std::sync::mpsc::{channel, Receiver, RecvTimeoutError, Sender};
use std::sync::Arc;
use std::thread::{spawn, JoinHandle};
use std::time::Duration;
use thiserror::Error;
#[derive(Debug, Error)]
pub enum CaptureError {
#[error("Failed to join one of the threads")]
JoinError,
#[error("Failed to do a socket operation")]
SockError(#[from] RSockErr),
}
pub trait Filter: Send {
fn filter(&mut self, packet: &IpPacket) -> bool;
}
pub struct NoFilter;
impl Filter for NoFilter {
fn filter(&mut self, _packet: &IpPacket) -> bool {
true
}
}
pub struct Ready {
filter: Box<dyn Filter>,
interface: Option<String>,
}
pub struct Capturing {
snooper: JoinHandle<()>,
store: JoinHandle<Vec<IpPacket>>,
cond: Arc<AtomicBool>,
}
pub struct Capture<T> {
state: T,
}
impl Capture<Ready> {
#[must_use]
pub fn new(filter: Box<dyn Filter>, interface: Option<String>) -> Self {
Self {
state: Ready { filter, interface },
}
}
pub fn start(self) -> Result<Capture<Capturing>, CaptureError> {
let (send, recv) = channel::<IpPacket>();
let signal = Arc::new(AtomicBool::new(false));
let cap_signal = Arc::clone(&signal);
let capture = spawn(move || {
let capture = Store {
chan: recv,
cond: cap_signal,
store: Vec::with_capacity(1000),
filter: self.state.filter,
};
capture.start()
});
let sock = Raw::new();
let mut sock = match sock {
Ok(s) => s,
Err(e) => {
signal.store(true, Ordering::Relaxed);
capture.join().map_err(|_| CaptureError::JoinError)?;
return Err(CaptureError::SockError(e));
}
};
if let Some(iface) = self.state.interface {
sock.bind_interface(&iface)
.map_err(CaptureError::SockError)?;
};
let snoop_signal = Arc::clone(&signal);
let snoop = spawn(move || {
let snooper = Snooper {
chan: send,
cond: snoop_signal,
sock,
};
snooper.start();
});
let cap_state = Capturing {
snooper: snoop,
store: capture,
cond: signal,
};
Ok(Capture { state: cap_state })
}
}
impl Capture<Capturing> {
pub fn end(self) -> Result<Vec<IpPacket>, CaptureError> {
self.state.cond.store(true, Ordering::Relaxed);
let _cap_res = self.state.snooper.join();
let store_res = self.state.store.join();
store_res.map_err(|_| CaptureError::JoinError)
}
}
struct Snooper {
chan: Sender<IpPacket>,
cond: Arc<AtomicBool>,
sock: Raw,
}
impl Snooper {
fn start(mut self) {
const RECV_WAIT: Duration = Duration::new(0, 50_000_000);
let mut buf = Box::new([0u8; 65536]);
'recv_loop: loop {
if self.cond.load(Ordering::Relaxed) {
break 'recv_loop;
}
let try_recv = self.sock.read_timeout(&mut buf[..], &RECV_WAIT);
if let Ok(n) = try_recv {
let packet = IpPacket::parse_from_bytes(&&buf[0..n], None);
if let Ok(p) = packet {
self.chan.send(p).unwrap();
}
}
}
}
}
struct Store {
chan: Receiver<IpPacket>,
cond: Arc<AtomicBool>,
store: Vec<IpPacket>,
filter: Box<dyn Filter>,
}
impl Store {
fn start(mut self) -> Vec<IpPacket> {
const RECV_WAIT: Duration = Duration::new(0, 50_000_000);
'recv_loop: loop {
if self.cond.load(Ordering::Relaxed) {
break 'recv_loop;
}
let packet = self.chan.recv_timeout(RECV_WAIT);
match packet {
Err(RecvTimeoutError::Timeout) => {
},
Err(RecvTimeoutError::Disconnected) => break 'recv_loop,
Ok(ip_packet) => {
if self.filter.filter(&ip_packet) {
self.store.push(ip_packet);
}
}
}
}
self.store
}
}