use std::{
io::{Read as _, Write as _},
sync::Arc,
};
use anyhow::Context;
use tocat_api::{Chain, ChannelTarget, Direction as Flow, PluginSpec, Registry, Side};
use tokio::{
net::{TcpListener, UnixListener},
sync::Semaphore,
};
use tokio_seqpacket::UnixSeqpacketListener;
use tracing::{Instrument, debug, error, info, warn};
use crate::{
buffer::Buffer,
endpoint::{
Demux, Direction, EndpointSpec, EndpointStream, PathGuard, ReadHalf, SyncRead, SyncWrite,
WriteHalf,
},
host::{ChannelPlan, Channels},
progress::{self, Counter, Meter},
pump::pump,
shutdown::Shutdown,
};
enum Listener {
Tcp(TcpListener),
Unix(UnixListener),
Seqpacket(UnixSeqpacketListener),
Datagram(Demux),
}
impl Listener {
async fn bind(
spec: &EndpointSpec,
buffer: usize,
shutdown: Shutdown,
) -> 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, "listening");
Ok((Listener::Unix(l), e.path.guard()))
}
EndpointSpec::UnixSeqpacketListen(e) => {
let l = e.bind().await?;
info!(path = %e.path, "listening");
Ok((Listener::Seqpacket(l), e.path.guard()))
}
EndpointSpec::UnixDgramListen(e) => Ok((
Listener::Datagram(e.demux(buffer, shutdown).await?),
e.path.guard(),
)),
EndpointSpec::UdpListen(e) => {
Ok((Listener::Datagram(e.demux(buffer, shutdown).await?), None))
}
_ => anyhow::bail!("fork is only supported on listening endpoints"),
}
}
async fn accept(&mut 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))
}
Listener::Seqpacket(l) => {
let s = l.accept().await?;
Ok((EndpointStream::seqpacket(s), "unnamed".to_string()))
}
Listener::Datagram(d) => d.accept().await,
}
}
}
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)
}
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",
);
let mut faults = Vec::new();
for (chain, upstream, downstream, direction) in [
(&forward, &source, &sink, "source-to-sink"),
(&reverse, &sink, &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 sent to this endpoint \
may be split, merged, or malformed",
);
}
for fault in chain.boundary_faults(upstream.is_datagram(), downstream.is_datagram()) {
let (want, remedy, end) = match fault.side {
Side::Upstream => ("arriving", "an unframe above it", upstream),
Side::Downstream => ("to survive", "a frame below it", downstream),
};
let cause = match fault.cause {
Some(stage) => format!("{stage} does not carry them"),
None => format!(
"the {} {} is a byte stream",
fault.side.endpoint_role(),
end.name(),
),
};
faults.push(format!(
"{} on {direction} needs message boundaries {want}, and {cause}: put {remedy}, \
or use a stage that does not need them",
fault.stage,
));
}
}
if !faults.is_empty() {
anyhow::bail!("{}", faults.join("\n"));
}
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 if self.prefers_sync() {
let watcher = shutdown.clone();
tokio::select! {
res = self.run_sync(watcher) => res,
() = shutdown.recv() => {
info!("interrupted");
Ok(())
}
}
} else {
self.run_once(shutdown).await
}
}
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 pump_direction(
&self,
reader: Option<ReadHalf>,
writer: Option<WriteHalf>,
chain: Chain,
flow: Flow,
shutdown: Shutdown,
) -> anyhow::Result<u64> {
let (Some(reader), Some(writer)) = (reader, writer) else {
if chain.is_empty() {
debug!(?flow, "one way, skipping the direction that does not exist");
} else {
warn!(?flow, "no such direction, so its stages will not run");
}
return Ok(0);
};
let (reader, counter) = progress::count(self.progress.as_ref(), reader, flow);
pump(
reader,
writer,
chain,
self.channels.clone(),
self.buffer,
counter,
shutdown,
)
.await
}
async fn relay_streams(
&self,
src_stream: EndpointStream,
sink_stream: EndpointStream,
forward: Chain,
reverse: Chain,
mut shutdown: Shutdown,
) -> anyhow::Result<()> {
let buffer = self.buffer;
let (src_stream, sink_stream) =
if forward.is_empty() && reverse.is_empty() && self.progress.is_none() {
match (src_stream, sink_stream) {
(EndpointStream::Duplex(mut a), EndpointStream::Duplex(mut b)) => {
tokio::select! {
copied = tokio::io::copy_bidirectional_with_sizes(
&mut a, &mut b, buffer, buffer,
) => {
let (to_sink, to_source) = copied?;
info!(bytes = to_sink + to_source, "relay finished");
}
() = shutdown.recv() => info!("interrupted"),
}
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 (a, b) = tokio::try_join!(
self.pump_direction(
src_read,
sink_write,
forward,
Flow::SourceToSink,
shutdown.clone(),
),
self.pump_direction(
sink_read,
src_write,
reverse,
Flow::SinkToSource,
shutdown.clone(),
),
)?;
info!(bytes = a + b, "relay finished");
self.channels.flush().await?;
Ok(())
}
async fn run_once(&self, shutdown: Shutdown) -> anyhow::Result<()> {
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)?;
self.relay_streams(
src_conn.stream,
sink_conn.stream,
forward,
reverse,
shutdown,
)
.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 (mut listener, _socket_guard) =
Listener::bind(listen, self.buffer, shutdown.clone()).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);
let shutdown = shutdown.clone();
tracker.spawn(
async move {
let _permit = permit;
match this.handle_client(stream, &peer, peer_dir, shutdown).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,
shutdown: Shutdown,
) -> 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))?;
self.relay_streams(src_stream, sink_stream, forward, reverse, shutdown)
.await
}
}