use std::{collections::HashMap, sync::Mutex, task::Poll};
#[cfg(feature = "pcap")]
use std::path::PathBuf;
use futures::{StreamExt, stream::FuturesUnordered};
use tokio::{
io::{AsyncRead, AsyncWrite},
sync::{mpsc, oneshot},
};
use tracing::trace;
use crate::adapter::ConnectionStatus;
use crate::time::timeout;
enum HandleMessage {
ConnectToPort {
target: u16,
res: oneshot::Sender<
Result<(u16, mpsc::UnboundedReceiver<Result<Vec<u8>, std::io::Error>>), std::io::Error>,
>,
},
Close {
host_port: u16,
},
Send {
host_port: u16,
data: Vec<u8>,
res: oneshot::Sender<Result<(), std::io::Error>>,
},
#[cfg(feature = "pcap")]
Pcap {
path: PathBuf,
res: oneshot::Sender<Result<(), std::io::Error>>,
},
Die,
}
#[derive(Debug, Clone)]
pub struct AdapterHandle {
sender: mpsc::UnboundedSender<HandleMessage>,
}
impl AdapterHandle {
pub fn new(mut adapter: crate::adapter::Adapter) -> Self {
let (tx, mut rx) = mpsc::unbounded_channel::<HandleMessage>();
crate::spawn(async move {
let mut handles: HashMap<u16, mpsc::UnboundedSender<Result<Vec<u8>, std::io::Error>>> =
HashMap::new();
let mut tick = crate::time::interval(std::time::Duration::from_millis(1));
loop {
tokio::select! {
msg = rx.recv() => {
match msg {
Some(m) => match m {
HandleMessage::ConnectToPort { target, res } => {
let connect_response = match adapter.connect(target).await {
Ok(c) => {
let (ptx, prx) = mpsc::unbounded_channel();
handles.insert(c, ptx);
Ok((c, prx))
}
Err(e) => Err(e),
};
res.send(connect_response).ok();
}
HandleMessage::Close { host_port } => {
handles.remove(&host_port);
adapter.close(host_port).await.ok();
}
HandleMessage::Send {
host_port,
data,
res,
} => {
if let Err(e) = adapter.queue_send(&data, host_port) {
res.send(Err(e)).ok();
} else {
let response = adapter.write_buffer_flush().await;
res.send(response).ok();
}
}
#[cfg(feature = "pcap")]
HandleMessage::Pcap {
path,
res
} => {
res.send(adapter.pcap(path).await).ok();
},
HandleMessage::Die => {
break;
}
},
None => {
break;
},
}
}
r = adapter.process_tcp_packet() => {
if let Err(e) = r {
for (hp, tx) in handles.drain() {
let _ = tx.send(Err(e.kind().into())); let _ = adapter.close(hp).await;
}
break;
}
let mut dead = Vec::new();
for (&hp, tx) in &handles {
match adapter.uncache_all(hp) {
Ok(buf) if !buf.is_empty() => {
if tx.send(Ok(buf)).is_err() {
dead.push(hp);
}
}
Err(e) => {
let _ = tx.send(Err(e));
dead.push(hp);
}
_ => {}
}
}
for hp in dead {
handles.remove(&hp);
let _ = adapter.close(hp).await;
}
let mut to_close = Vec::new();
for (&hp, tx) in &handles {
if let Ok(ConnectionStatus::Error(kind)) = adapter.get_status(hp) {
if kind == std::io::ErrorKind::UnexpectedEof {
to_close.push(hp);
} else {
let _ = tx.send(Err(std::io::Error::from(kind)));
to_close.push(hp);
}
}
}
for hp in to_close {
handles.remove(&hp);
let _ = adapter.close(hp).await;
}
}
_ = tick.tick() => {
let _ = adapter.write_buffer_flush().await;
}
}
}
});
Self { sender: tx }
}
pub async fn connect(&mut self, port: u16) -> Result<StreamHandle, std::io::Error> {
let (res_tx, res_rx) = oneshot::channel();
if self
.sender
.send(HandleMessage::ConnectToPort {
target: port,
res: res_tx,
})
.is_err()
{
return Err(std::io::Error::new(
std::io::ErrorKind::NetworkUnreachable,
"adapter closed",
));
}
match timeout(std::time::Duration::from_secs(8), res_rx).await {
Ok(Ok(r)) => {
let (host_port, recv_channel) = r?;
Ok(StreamHandle {
host_port,
recv_channel: Mutex::new(recv_channel),
send_channel: self.sender.clone(),
read_buffer: Vec::new(),
pending_writes: FuturesUnordered::new(),
})
}
Ok(Err(_)) => Err(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"adapter closed",
)),
Err(_) => Err(std::io::Error::new(
std::io::ErrorKind::TimedOut,
"channel recv timeout",
)),
}
}
#[cfg(feature = "pcap")]
pub async fn pcap(&mut self, path: impl Into<PathBuf>) -> Result<(), std::io::Error> {
let (res_tx, res_rx) = oneshot::channel();
let path: PathBuf = path.into();
if self
.sender
.send(HandleMessage::Pcap { path, res: res_tx })
.is_err()
{
return Err(std::io::Error::new(
std::io::ErrorKind::NetworkUnreachable,
"adapter closed",
));
}
match res_rx.await {
Ok(r) => r,
Err(_) => Err(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"adapter closed",
)),
}
}
pub async fn close(&mut self) -> Result<(), std::io::Error> {
if self.sender.send(HandleMessage::Die).is_err() {
return Err(std::io::Error::new(
std::io::ErrorKind::NetworkUnreachable,
"adapter closed",
));
}
Ok(())
}
}
#[derive(Debug)]
pub struct StreamHandle {
host_port: u16,
recv_channel: Mutex<mpsc::UnboundedReceiver<Result<Vec<u8>, std::io::Error>>>,
send_channel: mpsc::UnboundedSender<HandleMessage>,
read_buffer: Vec<u8>,
pending_writes: FuturesUnordered<oneshot::Receiver<Result<(), std::io::Error>>>,
}
impl StreamHandle {
pub fn close(&mut self) {
let _ = self.send_channel.send(HandleMessage::Close {
host_port: self.host_port,
});
}
}
impl AsyncRead for StreamHandle {
fn poll_read(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
if !self.read_buffer.is_empty() {
let n = buf.remaining().min(self.read_buffer.len());
buf.put_slice(&self.read_buffer[..n]);
self.read_buffer.drain(..n); return Poll::Ready(Ok(()));
}
let mut lock = self
.recv_channel
.lock()
.expect("somehow the mutex was poisoned");
let mut extend_slice = Vec::new();
let res = match lock.poll_recv(cx) {
Poll::Pending => Poll::Pending,
Poll::Ready(None) => Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"channel closed",
))),
Poll::Ready(Some(res)) => match res {
Ok(data) => {
let n = buf.remaining().min(data.len());
buf.put_slice(&data[..n]);
if n < data.len() {
extend_slice = data[n..].to_vec();
}
Poll::Ready(Ok(()))
}
Err(e) => Poll::Ready(Err(e)),
},
};
std::mem::drop(lock);
self.read_buffer.extend(extend_slice);
res
}
}
impl AsyncWrite for StreamHandle {
fn poll_write(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
buf: &[u8],
) -> std::task::Poll<Result<usize, std::io::Error>> {
trace!("poll psh {}", buf.len());
let (tx, rx) = oneshot::channel();
self.send_channel
.send(HandleMessage::Send {
host_port: self.host_port,
data: buf.to_vec(),
res: tx,
})
.map_err(|_| std::io::Error::new(std::io::ErrorKind::BrokenPipe, "channel closed"))?;
self.pending_writes.push(rx);
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
while let Poll::Ready(maybe) = self.pending_writes.poll_next_unpin(cx) {
match maybe {
Some(Ok(Ok(()))) => {}
Some(Ok(Err(e))) => return Poll::Ready(Err(e)),
Some(Err(_canceled)) => {
return Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"channel closed",
)));
}
None => break, }
}
if self.pending_writes.is_empty() {
Poll::Ready(Ok(()))
} else {
Poll::Pending
}
}
fn poll_shutdown(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), std::io::Error>> {
std::task::Poll::Ready(Ok(()))
}
}
impl Drop for StreamHandle {
fn drop(&mut self) {
let _ = self.send_channel.send(HandleMessage::Close {
host_port: self.host_port,
});
}
}