use std::{
io::{Read as _, Write as _},
sync::Arc,
};
use anyhow::Context;
use tocat_api::{Chain, ChannelTarget, Direction as Flow, PluginSpec, Registry};
use tokio::{
net::{TcpListener, UnixListener},
sync::Semaphore,
};
use tracing::{Instrument, debug, error, info, warn};
use crate::{
buffer::Buffer,
endpoint::{Direction, EndpointSpec, EndpointStream, PathGuard, SyncRead, SyncWrite},
host::{ChannelPlan, Channels},
progress::{self, Counter, Meter},
pump::pump,
shutdown::Shutdown,
};
enum Listener {
Tcp(TcpListener),
Unix(UnixListener),
}
impl Listener {
async fn bind(spec: &EndpointSpec) -> anyhow::Result<(Self, Option<PathGuard>)> {
match spec {
EndpointSpec::TcpListen(e) => {
let l = e.bind().await?;
info!(local = %l.local_addr()?, "listening");
Ok((Listener::Tcp(l), None))
}
EndpointSpec::UnixListen(e) => {
let l = e.bind().await?;
info!(path = %e.path.display(), "listening");
Ok((Listener::Unix(l), Some(PathGuard(e.path.clone()))))
}
_ => anyhow::bail!("fork is only supported on listening endpoints"),
}
}
async fn accept(&self) -> std::io::Result<(EndpointStream, String)> {
match self {
Listener::Tcp(l) => {
let (s, peer) = l.accept().await?;
Ok((EndpointStream::tcp(s), peer.to_string()))
}
Listener::Unix(l) => {
let (s, peer) = l.accept().await?;
let label = peer
.as_pathname()
.map(|p| p.display().to_string())
.unwrap_or_else(|| "unnamed".to_string());
Ok((EndpointStream::unix(s), label))
}
}
}
}
fn is_fatal_accept(e: &std::io::Error) -> bool {
!matches!(
e.kind(),
std::io::ErrorKind::ConnectionAborted
| std::io::ErrorKind::Interrupted
| std::io::ErrorKind::WouldBlock
)
}
fn copy_sync(
mut reader: SyncRead,
mut writer: SyncWrite,
shutdown: &Shutdown,
buffer: usize,
counter: Option<Counter>,
) -> anyhow::Result<u64> {
let mut buf = Buffer::new(buffer);
let mut total = 0u64;
loop {
if shutdown.is_triggered() {
info!(bytes = total, "interrupted");
break;
}
let n = reader.read(&mut buf)?;
if n == 0 {
break;
}
writer.write_all(&buf[..n])?;
total += n as u64;
if let Some(counter) = &counter {
counter.add(n as u64);
}
}
writer.flush()?;
Ok(total)
}
async fn relay_streams(
src_stream: EndpointStream,
sink_stream: EndpointStream,
forward: Chain,
reverse: Chain,
channels: Arc<Channels>,
buffer: usize,
meter: Option<Arc<Meter>>,
) -> anyhow::Result<()> {
let (src_stream, sink_stream) = if forward.is_empty() && reverse.is_empty() && meter.is_none() {
match (src_stream, sink_stream) {
(EndpointStream::Duplex(mut a), EndpointStream::Duplex(mut b)) => {
let (to_sink, to_source) =
tokio::io::copy_bidirectional_with_sizes(&mut a, &mut b, buffer, buffer)
.await?;
info!(bytes = to_sink + to_source, "relay finished");
return Ok(());
}
pair => pair,
}
} else {
(src_stream, sink_stream)
};
let (src_read, src_write) = src_stream.into_halves();
let (sink_read, sink_write) = sink_stream.into_halves();
let (src_read, forward_count) = progress::count(meter.as_ref(), src_read, Flow::SourceToSink);
let (sink_read, reverse_count) = progress::count(meter.as_ref(), sink_read, Flow::SinkToSource);
let (a, b) = tokio::try_join!(
pump(
src_read,
sink_write,
forward,
channels.clone(),
buffer,
forward_count
),
pump(
sink_read,
src_write,
reverse,
channels.clone(),
buffer,
reverse_count
),
)?;
info!(bytes = a + b, "relay finished");
channels.flush().await?;
Ok(())
}
pub struct Relay {
source: EndpointSpec,
sink: EndpointSpec,
plugins: Vec<PluginSpec>,
registry: Registry,
plan: ChannelPlan,
channels: Arc<Channels>,
buffer: usize,
progress: Option<Arc<Meter>>,
}
impl Relay {
pub async fn new(
source: EndpointSpec,
sink: EndpointSpec,
plugins: Vec<PluginSpec>,
registry: Registry,
buffer: usize,
progress: Option<Arc<Meter>>,
) -> anyhow::Result<Self> {
let mut plan = ChannelPlan::new();
let (forward, reverse) =
registry.build_pair(&plugins, &source.name(), &sink.name(), None, &mut plan)?;
debug!(
forward = ?forward.stage_names(),
reverse = ?reverse.stage_names(),
forward_segments = forward.segments().len(),
reverse_segments = reverse.segments().len(),
channels = plan.targets().len(),
"plugin chains resolved",
);
for (chain, downstream, direction) in [
(&forward, &sink, "source-to-sink"),
(&reverse, &source, "sink-to-source"),
] {
if downstream.is_datagram()
&& let Some(stage) = chain.datagram_hazard()
{
warn!(
stage, direction, endpoint = %downstream.name(),
"stage may not preserve message boundaries; datagrams send to this endpoint \
may be split, merged, or malformed",
);
}
}
plan.freeze();
if progress.is_some()
&& plan
.targets()
.iter()
.any(|target| matches!(target, ChannelTarget::Stderr))
{
warn!(
"a plugin is dumping to stderr while the progress line is drawn there; send the \
dump to a file, or drop --progress",
);
}
let channels = Channels::open(plan.targets()).await?;
let peak = buffer.saturating_mul(source.max_connections().get().max(1)) * 2;
if peak > 1024 * 1024 * 1024 {
warn!(
buffer,
"buffer size and connection ceiling allow over 1 GiB of copy buffers"
);
}
Ok(Self {
source,
sink,
plugins,
registry,
plan,
channels,
buffer,
progress,
})
}
pub async fn run(self, shutdown: Shutdown) -> anyhow::Result<()> {
let this = Arc::new(self);
let channels = this.channels.clone();
let result = this.dispatch(shutdown).await;
if let Err(e) = channels.flush().await {
warn!(error = %e, "flushing plugin channels failed");
}
result
}
async fn dispatch(self: Arc<Self>, mut shutdown: Shutdown) -> anyhow::Result<()> {
if self.source.is_fork() {
self.serve(Direction::Sink, shutdown).await
} else if self.sink.is_fork() {
self.serve(Direction::Source, shutdown).await
} else {
let watcher = shutdown.clone();
tokio::select! {
res = self.run_once(watcher) => res,
_ = shutdown.recv() => {
info!("interrupted");
Ok(())
}
}
}
}
fn chains(
&self,
src_name: &str,
sink_name: &str,
peer: Option<&str>,
) -> anyhow::Result<(Chain, Chain)> {
let mut plan = self.plan.clone();
Ok(self
.registry
.build_pair(&self.plugins, src_name, sink_name, peer, &mut plan)?)
}
fn prefers_sync(&self) -> bool {
self.plugins.is_empty()
&& self.source.is_blocking_backed()
&& self.sink.is_blocking_backed()
}
async fn run_sync(&self, shutdown: Shutdown) -> anyhow::Result<()> {
let mut source = self.source.connect_sync(Direction::Source, self.buffer)?;
let mut sink = self.sink.connect_sync(Direction::Sink, self.buffer)?;
let _guards = (source.guard.take(), sink.guard.take());
let directions = [
(source.reader, sink.writer, Flow::SourceToSink),
(sink.reader, source.writer, Flow::SinkToSource),
];
let mut running = Vec::new();
for (reader, writer, path) in directions {
if let (Some(reader), Some(writer)) = (reader, writer) {
let shutdown = shutdown.clone();
let buffer = self.buffer;
let counter = self.progress.as_ref().map(|meter| meter.counter(path));
running.push(tokio::task::spawn_blocking(move || {
copy_sync(reader, writer, &shutdown, buffer, counter)
}));
}
}
let mut total = 0u64;
for task in running {
total += task.await.context("blocking copy task panicked")??;
}
info!(bytes = total, "relay finished");
Ok(())
}
async fn run_once(&self, shutdown: Shutdown) -> anyhow::Result<()> {
if self.prefers_sync() {
return self.run_sync(shutdown).await;
}
let (src_conn, sink_conn) = if self.source.is_listen() && self.sink.is_listen() {
tokio::try_join!(
self.source.connect(Direction::Source, self.buffer),
self.sink.connect(Direction::Sink, self.buffer)
)?
} else {
(
self.source.connect(Direction::Source, self.buffer).await?,
self.sink.connect(Direction::Sink, self.buffer).await?,
)
};
let _guards = (src_conn.guard, sink_conn.guard);
let (forward, reverse) = self.chains(&self.source.name(), &self.sink.name(), None)?;
relay_streams(
src_conn.stream,
sink_conn.stream,
forward,
reverse,
self.channels.clone(),
self.buffer,
self.progress.clone(),
)
.await
}
fn listening(&self, peer_dir: Direction) -> &EndpointSpec {
match peer_dir {
Direction::Sink => &self.source,
Direction::Source => &self.sink,
}
}
async fn serve(
self: Arc<Self>,
peer_dir: Direction,
mut shutdown: Shutdown,
) -> anyhow::Result<()> {
let listen = self.listening(peer_dir);
let max = listen.max_connections();
let (listener, _socket_guard) = Listener::bind(listen).await?;
info!(max = max.get(), "accepting connections");
let permits = Arc::new(Semaphore::new(max.get()));
let tracker = tokio_util::task::TaskTracker::new();
loop {
let permit = tokio::select! {
biased;
_ = shutdown.recv() => break,
p = permits.clone().acquire_owned() => p?,
};
let (stream, peer) = tokio::select! {
biased;
_ = shutdown.recv() => break,
conn = listener.accept() => match conn {
Ok(conn) => conn,
Err(e) if is_fatal_accept(&e) => return Err(e).context("accept"),
Err(e) => {
warn!("Accept error: {e}");
continue;
}
},
};
let this = Arc::clone(&self);
let span = tracing::info_span!("conn", %peer);
tracker.spawn(
async move {
let _permit = permit;
match this.handle_client(stream, &peer, peer_dir).await {
Ok(()) => info!("closed cleanly"),
Err(err) => error!(error = ?err, "terminated with error"),
}
}
.instrument(span),
);
}
tracker.close();
info!(active = tracker.len(), "waiting for connections to drain");
tracker.wait().await;
info!("drained");
Ok(())
}
async fn handle_client(
&self,
accepted: EndpointStream,
peer: &str,
peer_dir: Direction,
) -> anyhow::Result<()> {
let _connection = self.progress.as_ref().map(|meter| meter.connected());
let listen = self.listening(peer_dir);
let peer_spec = match peer_dir {
Direction::Sink => &self.sink,
Direction::Source => &self.source,
};
let dialled = peer_spec.connect(peer_dir, self.buffer).await?;
let _guard = dialled.guard;
let (src_stream, sink_stream, src_spec, sink_spec) = match peer_dir {
Direction::Source => (dialled.stream, accepted, peer_spec, listen),
Direction::Sink => (accepted, dialled.stream, listen, peer_spec),
};
let src_name = if peer_dir == Direction::Sink {
format!("{}_{}", src_spec.name(), peer)
} else {
src_spec.name()
};
let sink_name = if peer_dir == Direction::Source {
format!("{}_{}", sink_spec.name(), peer)
} else {
sink_spec.name()
};
let (forward, reverse) = self.chains(&src_name, &sink_name, Some(peer))?;
relay_streams(
src_stream,
sink_stream,
forward,
reverse,
self.channels.clone(),
self.buffer,
self.progress.clone(),
)
.await
}
}