use std::collections::VecDeque;
use std::io;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll, Waker};
use bytes::Bytes;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, ReadBuf};
use crate::bus_timing::BusTiming;
use crate::wire_tap::{now_micros, WireTap};
const DEFAULT_BUF_SIZE: usize = 4096;
pub const DEFAULT_DATA_CHANNEL_CAPACITY: usize = 1024 * 1024;
pub const DEFAULT_TAP_CHANNEL_CAPACITY: usize = 256 * 1024;
struct SniffChannel {
inner: Mutex<SniffChannelInner>,
}
struct SniffChannelInner {
queue: VecDeque<Bytes>,
total_bytes: usize,
capacity: usize,
dropped_chunks: u64,
dropped_bytes: u64,
closed: bool,
error: Option<io::Error>,
waker: Option<Waker>,
}
impl SniffChannel {
fn new(capacity: usize) -> Self {
Self {
inner: Mutex::new(SniffChannelInner {
queue: VecDeque::new(),
total_bytes: 0,
capacity,
dropped_chunks: 0,
dropped_bytes: 0,
closed: false,
error: None,
waker: None,
}),
}
}
fn send(&self, chunk: Bytes) {
let mut inner = self.inner.lock().unwrap_or_else(|e| e.into_inner());
while inner.total_bytes + chunk.len() > inner.capacity {
if let Some(old) = inner.queue.pop_front() {
inner.total_bytes -= old.len();
inner.dropped_chunks += 1;
inner.dropped_bytes += old.len() as u64;
} else {
inner.dropped_chunks += 1;
inner.dropped_bytes += chunk.len() as u64;
return;
}
}
let chunk_len = chunk.len();
inner.queue.push_back(chunk);
inner.total_bytes += chunk_len;
if let Some(w) = inner.waker.take() {
w.wake();
}
}
fn send_error(&self, error: io::Error) {
let mut inner = self.inner.lock().unwrap_or_else(|e| e.into_inner());
inner.error = Some(error);
if let Some(w) = inner.waker.take() {
w.wake();
}
}
fn set_capacity(&self, new_cap: usize) {
let mut inner = self.inner.lock().unwrap_or_else(|e| e.into_inner());
inner.capacity = new_cap;
while inner.total_bytes > new_cap {
if let Some(old) = inner.queue.pop_front() {
inner.total_bytes -= old.len();
inner.dropped_chunks += 1;
inner.dropped_bytes += old.len() as u64;
} else {
break;
}
}
}
fn is_closed(&self) -> bool {
self.inner.lock().unwrap_or_else(|e| e.into_inner()).closed
}
fn close(&self) {
let mut inner = self.inner.lock().unwrap_or_else(|e| e.into_inner());
inner.closed = true;
if let Some(w) = inner.waker.take() {
w.wake();
}
}
fn poll_recv(&self, cx: &mut Context<'_>) -> io::Result<Option<Bytes>> {
let mut inner = self.inner.lock().unwrap_or_else(|e| e.into_inner());
if let Some(data) = inner.queue.pop_front() {
inner.total_bytes -= data.len();
return Ok(Some(data));
}
if let Some(e) = inner.error.take() {
return Err(e);
}
if inner.closed {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"channel closed",
));
}
inner.waker = Some(cx.waker().clone());
Ok(None)
}
async fn recv(&self) -> Option<Bytes> {
use std::future::poll_fn;
poll_fn(|cx| match self.poll_recv(cx) {
Ok(Some(data)) => Poll::Ready(Some(data)),
Ok(None) => Poll::Pending,
Err(_) => Poll::Ready(None),
})
.await
}
fn capacity(&self) -> usize {
self.inner
.lock()
.unwrap_or_else(|e| e.into_inner())
.capacity
}
fn dropped_chunks(&self) -> u64 {
self.inner
.lock()
.unwrap_or_else(|e| e.into_inner())
.dropped_chunks
}
fn dropped_bytes(&self) -> u64 {
self.inner
.lock()
.unwrap_or_else(|e| e.into_inner())
.dropped_bytes
}
}
pub struct SniffIo<T> {
tap: Option<Arc<dyn WireTap>>,
buf_size: usize,
data_capacity: usize,
tap_capacity: usize,
timing: Option<Arc<BusTiming>>,
state: SniffState<T>,
}
enum SniffState<T> {
Passthrough(T),
Sniffing {
data_channel: Arc<SniffChannel>,
tap_channel: Arc<SniffChannel>,
writer: tokio::io::WriteHalf<T>,
reader_handle: tokio::task::JoinHandle<()>,
_tap_handle: tokio::task::JoinHandle<()>,
pending: Option<Bytes>,
},
}
impl<T: AsyncRead + AsyncWrite + Unpin + Send + 'static> SniffIo<T> {
pub fn new(inner: T, tap: Option<Arc<dyn WireTap>>, timing: Option<Arc<BusTiming>>) -> Self {
match tap {
Some(t) => Self::start_sniffing(
inner,
t,
timing,
DEFAULT_BUF_SIZE,
DEFAULT_DATA_CHANNEL_CAPACITY,
DEFAULT_TAP_CHANNEL_CAPACITY,
),
None => Self {
tap: None,
buf_size: DEFAULT_BUF_SIZE,
data_capacity: DEFAULT_DATA_CHANNEL_CAPACITY,
tap_capacity: DEFAULT_TAP_CHANNEL_CAPACITY,
timing,
state: SniffState::Passthrough(inner),
},
}
}
pub fn with_buf_size(mut self, size: usize) -> Self {
self.buf_size = size;
self
}
pub fn with_channel_capacity(mut self, bytes: usize) -> Self {
self.data_capacity = bytes;
if let SniffState::Sniffing { data_channel, .. } = &self.state {
data_channel.set_capacity(bytes);
}
self
}
pub fn with_tap_channel_capacity(mut self, bytes: usize) -> Self {
self.tap_capacity = bytes;
if let SniffState::Sniffing { tap_channel, .. } = &self.state {
tap_channel.set_capacity(bytes);
}
self
}
pub fn data_dropped_chunks(&self) -> u64 {
match &self.state {
SniffState::Sniffing { data_channel, .. } => data_channel.dropped_chunks(),
SniffState::Passthrough(_) => 0,
}
}
pub fn data_dropped_bytes(&self) -> u64 {
match &self.state {
SniffState::Sniffing { data_channel, .. } => data_channel.dropped_bytes(),
SniffState::Passthrough(_) => 0,
}
}
pub fn tap_dropped_chunks(&self) -> u64 {
match &self.state {
SniffState::Sniffing { tap_channel, .. } => tap_channel.dropped_chunks(),
SniffState::Passthrough(_) => 0,
}
}
pub fn tap_dropped_bytes(&self) -> u64 {
match &self.state {
SniffState::Sniffing { tap_channel, .. } => tap_channel.dropped_bytes(),
SniffState::Passthrough(_) => 0,
}
}
pub fn total_dropped_bytes(&self) -> u64 {
self.data_dropped_bytes() + self.tap_dropped_bytes()
}
pub fn channel_capacity(&self) -> usize {
match &self.state {
SniffState::Sniffing { data_channel, .. } => data_channel.capacity(),
SniffState::Passthrough(_) => 0,
}
}
fn start_sniffing(
inner: T,
tap: Arc<dyn WireTap>,
timing: Option<Arc<BusTiming>>,
buf_size: usize,
data_capacity: usize,
tap_capacity: usize,
) -> Self {
let (read_half, write_half) = tokio::io::split(inner);
let data_channel = Arc::new(SniffChannel::new(data_capacity));
let tap_channel = Arc::new(SniffChannel::new(tap_capacity));
let reader_handle = tokio::spawn(active_reader(
read_half,
data_channel.clone(),
tap_channel.clone(),
buf_size,
));
let tap_handle = tokio::spawn(tap_forwarder(tap_channel.clone(), tap.clone()));
Self {
tap: Some(tap),
buf_size,
data_capacity,
tap_capacity,
timing,
state: SniffState::Sniffing {
data_channel,
tap_channel,
writer: write_half,
reader_handle,
_tap_handle: tap_handle,
pending: None,
},
}
}
pub fn prepare_send(&self) -> impl std::future::Future<Output = ()> + Send + '_ {
let timing = self.timing.clone();
async move {
if let Some(t) = timing {
t.wait_if_needed().await;
}
}
}
pub fn replace_inner(&mut self, new: T) {
let tap = self.tap.clone();
let timing = self.timing.clone();
let buf = self.buf_size;
let data_cap = self.data_capacity;
let tap_cap = self.tap_capacity;
*self = match tap {
Some(t) => Self::start_sniffing(new, t, timing, buf, data_cap, tap_cap),
None => Self {
tap: None,
buf_size: buf,
data_capacity: data_cap,
tap_capacity: tap_cap,
timing,
state: SniffState::Passthrough(new),
},
};
}
}
impl<T> Drop for SniffIo<T> {
fn drop(&mut self) {
if let SniffState::Sniffing {
data_channel,
tap_channel,
reader_handle,
_tap_handle,
..
} = &self.state
{
data_channel.close();
tap_channel.close();
reader_handle.abort();
_tap_handle.abort();
}
}
}
async fn active_reader<T: AsyncRead + Unpin>(
mut read_half: tokio::io::ReadHalf<T>,
data_channel: Arc<SniffChannel>,
tap_channel: Arc<SniffChannel>,
buf_size: usize,
) {
let mut buf = bytes::BytesMut::with_capacity(buf_size);
loop {
match read_half.read_buf(&mut buf).await {
Ok(0) => {
data_channel.send_error(io::Error::new(
io::ErrorKind::UnexpectedEof,
"transport closed",
));
break;
}
Ok(_n) => {
let chunk = buf.split().freeze();
data_channel.send(chunk.clone());
tap_channel.send(chunk);
if data_channel.is_closed() {
break;
}
}
Err(e) => {
data_channel.send_error(e);
break;
}
}
}
}
async fn tap_forwarder(channel: Arc<SniffChannel>, tap: Arc<dyn WireTap>) {
while let Some(chunk) = channel.recv().await {
tap.on_read(&chunk, now_micros());
}
}
impl<T: AsyncRead + AsyncWrite + Unpin> AsyncRead for SniffIo<T> {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let result = match &mut self.state {
SniffState::Passthrough(inner) => Pin::new(inner).poll_read(cx, buf),
SniffState::Sniffing {
data_channel,
reader_handle,
pending,
..
} => {
if let Some(p) = pending {
let n = p.len().min(buf.remaining());
buf.put_slice(&p[..n]);
if n < p.len() {
*pending = Some(p.slice(n..));
} else {
*pending = None;
}
return Poll::Ready(Ok(()));
}
match data_channel.poll_recv(cx) {
Ok(Some(data)) => {
let n = data.len().min(buf.remaining());
buf.put_slice(&data[..n]);
if n < data.len() {
*pending = Some(data.slice(n..));
}
Poll::Ready(Ok(()))
}
Ok(None) => {
if reader_handle.is_finished() {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"background reader terminated",
)));
}
Poll::Pending
}
Err(e) => Poll::Ready(Err(e)),
}
}
};
if let Poll::Ready(Ok(())) = &result {
if !buf.filled().is_empty() {
if let Some(t) = &self.timing {
t.touch();
}
}
}
result
}
}
impl<T: AsyncWrite + AsyncRead + Unpin> AsyncWrite for SniffIo<T> {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
let result = match &mut self.state {
SniffState::Passthrough(inner) => Pin::new(inner).poll_write(cx, buf),
SniffState::Sniffing { writer, .. } => Pin::new(writer).poll_write(cx, buf),
};
if let Poll::Ready(Ok(n)) = &result {
if *n > 0 {
if let Some(tap) = &self.tap {
tap.on_write(&buf[..*n], now_micros());
}
if let Some(t) = &self.timing {
t.touch();
}
}
}
result
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match &mut self.state {
SniffState::Passthrough(inner) => Pin::new(inner).poll_flush(cx),
SniffState::Sniffing { writer, .. } => Pin::new(writer).poll_flush(cx),
}
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match &mut self.state {
SniffState::Passthrough(inner) => Pin::new(inner).poll_shutdown(cx),
SniffState::Sniffing { writer, .. } => Pin::new(writer).poll_shutdown(cx),
}
}
}