use std::collections::VecDeque;
use std::pin::Pin;
use std::task::{Context, Poll};
use futures_core::Stream;
use crate::async_adapters::tokio_adapter::AsyncCapture;
use crate::dedup::Dedup;
use crate::error::Error;
use crate::packet::OwnedPacket;
use crate::traits::PacketSource;
pub struct DedupStream<S>
where
S: PacketSource + std::os::unix::io::AsRawFd,
{
cap: AsyncCapture<S>,
dedup: Dedup,
pending: VecDeque<OwnedPacket>,
#[cfg(feature = "pcap")]
tap: Option<crate::pcap_tap::PcapTap>,
}
impl<S> DedupStream<S>
where
S: PacketSource + std::os::unix::io::AsRawFd,
{
pub(crate) fn new(cap: AsyncCapture<S>, dedup: Dedup) -> Self {
Self {
cap,
dedup,
pending: VecDeque::new(),
#[cfg(feature = "pcap")]
tap: None,
}
}
pub fn dedup(&self) -> &Dedup {
&self.dedup
}
pub fn dedup_mut(&mut self) -> &mut Dedup {
&mut self.dedup
}
#[cfg(feature = "pcap")]
pub fn with_pcap_tap<W>(self, writer: crate::pcap::CaptureWriter<W>) -> Self
where
W: std::io::Write + Send + 'static,
{
self.with_pcap_tap_policy(writer, crate::pcap_tap::TapErrorPolicy::default())
}
#[cfg(feature = "pcap")]
pub fn with_pcap_tap_policy<W>(
mut self,
writer: crate::pcap::CaptureWriter<W>,
policy: crate::pcap_tap::TapErrorPolicy,
) -> Self
where
W: std::io::Write + Send + 'static,
{
self.tap = Some(crate::pcap_tap::PcapTap::new(writer, policy));
self
}
}
impl<S> Stream for DedupStream<S>
where
S: PacketSource + std::os::unix::io::AsRawFd + Unpin,
{
type Item = Result<OwnedPacket, Error>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
loop {
if let Some(pkt) = this.pending.pop_front() {
return Poll::Ready(Some(Ok(pkt)));
}
let mut guard = match this.cap.poll_read_ready_mut(cx) {
Poll::Ready(Ok(g)) => g,
Poll::Ready(Err(e)) => return Poll::Ready(Some(Err(Error::Io(e)))),
Poll::Pending => return Poll::Pending,
};
let got_batch = {
let inner = guard.get_inner_mut();
if let Some(batch) = inner.next_batch() {
#[cfg(feature = "pcap")]
let mut tap_error: Option<Error> = None;
for pkt in &batch {
if this.dedup.keep(&pkt) {
#[cfg(feature = "pcap")]
if let Some(tap) = this.tap.as_mut()
&& let Some(err) = tap.write_or_handle(&pkt)
{
tap_error = Some(err);
break;
}
this.pending.push_back(pkt.to_owned());
}
}
drop(batch);
#[cfg(feature = "pcap")]
if let Some(err) = tap_error {
return Poll::Ready(Some(Err(err)));
}
true
} else {
false
}
};
if !got_batch {
guard.clear_ready();
}
}
}
}
impl<S> AsyncCapture<S>
where
S: PacketSource + std::os::unix::io::AsRawFd,
{
pub fn dedup_stream(self, dedup: Dedup) -> DedupStream<S> {
DedupStream::new(self, dedup)
}
}
use crate::async_adapters::stream_capture::{Sealed, StreamCapture};
impl<S> Sealed for DedupStream<S> where S: PacketSource + std::os::unix::io::AsRawFd {}
impl<S> StreamCapture for DedupStream<S>
where
S: PacketSource + std::os::unix::io::AsRawFd,
{
type Source = S;
fn capture(&self) -> &AsyncCapture<S> {
&self.cap
}
}