use std::{sync::Arc, time::Instant};
use anyhow::{Context, bail};
use tocat_api::{Chain, Emitted, ExternalStage, Pipeline, Segment};
use tokio::{
io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader},
sync::mpsc,
time::{MissedTickBehavior, sleep},
};
use tracing::{info, warn};
use crate::{
buffer::Buffer,
child,
endpoint::{BoxRead, BoxWrite, DatagramSocket, ReadHalf, WriteHalf},
host::{Channels, Effects},
progress::Counter,
};
const LINK_DEPTH: usize = 2;
pub async fn pump(
reader: ReadHalf,
writer: WriteHalf,
chain: Chain,
channels: Arc<Channels>,
buffer: usize,
counter: Option<Counter>,
) -> anyhow::Result<u64> {
let meta = chain.meta().clone();
let mut segments = chain.into_segments();
match segments.len() {
0 if matches!(reader, ReadHalf::Stream(_)) && matches!(writer, WriteHalf::Stream(_)) => {
copy_direct(reader, writer, buffer).await
}
0 => {
run_pipeline(
Source::new(Upstream::Stream(reader), buffer, counter),
Downstream::Stream(writer).into(),
Pipeline::new(meta, Vec::new()),
channels,
)
.await
}
1 => {
run_segment(
segments.pop().expect("one segment"),
Upstream::Stream(reader),
Downstream::Stream(writer),
channels,
buffer,
counter,
)
.await
}
_ => run_segmented(reader, writer, segments, channels, buffer, counter).await,
}
}
enum Upstream {
Stream(ReadHalf),
Link(Inlet),
}
enum Downstream {
Stream(WriteHalf),
Link(Outlet),
}
async fn run_segment(
segment: Segment,
input: Upstream,
output: Downstream,
channels: Arc<Channels>,
buffer: usize,
counter: Option<Counter>,
) -> anyhow::Result<u64> {
match segment {
Segment::Inline(pipeline) => {
run_pipeline(
Source::new(input, buffer, counter),
output.into(),
pipeline,
channels,
)
.await
}
Segment::Process(external) => run_process(external, input, output, buffer, counter).await,
}
}
async fn copy_direct(reader: ReadHalf, writer: WriteHalf, buffer: usize) -> anyhow::Result<u64> {
let (ReadHalf::Stream(reader), WriteHalf::Stream(mut writer)) = (reader, writer) else {
unreachable!("checked by the caller");
};
let mut reader = BufReader::with_capacity(buffer, reader);
let total = tokio::io::copy_buf(&mut reader, &mut writer).await?;
writer.flush().await?;
let _ = writer.shutdown().await;
Ok(total)
}
async fn run_segmented(
reader: ReadHalf,
writer: WriteHalf,
mut segments: Vec<Segment>,
channels: Arc<Channels>,
buffer: usize,
counter: Option<Counter>,
) -> anyhow::Result<u64> {
let n = segments.len();
let mut inlets: Vec<Option<Inlet>> = (0..n).map(|_| None).collect();
let mut outlets: Vec<Option<Outlet>> = (0..n).map(|_| None).collect();
for i in 0..n - 1 {
let (outlet, inlet) = link();
outlets[i] = Some(outlet);
inlets[i + 1] = Some(inlet);
}
let mut writer = Some(writer);
let mut spawned = Vec::with_capacity(n - 1);
for i in (1..n).rev() {
let segment = segments.pop().expect("index is in range");
let inlet = inlets[i]
.take()
.expect("every non-head segment has an inlet");
let output = if i == n - 1 {
Downstream::Stream(writer.take().expect("the tail owns the writer"))
} else {
Downstream::Link(outlets[i].take().expect("a middle segment has an outlet"))
};
let channels = channels.clone();
spawned.push(tokio::spawn(async move {
run_segment(
segment,
Upstream::Link(inlet),
output,
channels,
buffer,
None,
)
.await
.map(|_| ())
}));
}
let head = segments.pop().expect("at least one segment");
let output = match outlets[0].take() {
Some(outlet) => Downstream::Link(outlet),
None => Downstream::Stream(writer.take().expect("a lone segment owns the writer")),
};
let total = run_segment(
head,
Upstream::Stream(reader),
output,
channels,
buffer,
counter,
)
.await?;
for handle in spawned {
handle.await.context("segment task panicked")??;
}
Ok(total)
}
async fn run_process(
external: ExternalStage,
input: Upstream,
output: Downstream,
buffer: usize,
counter: Option<Counter>,
) -> anyhow::Result<u64> {
let (program, args) = external
.argv
.split_first()
.expect("the factory rejects an empty argv");
let mut parts = child::spawn(program, args, external.shell, external.stderr, buffer)?;
let name = external.name.clone();
let feed = async {
let mut stdin = parts.stdin;
let total = match input {
Upstream::Stream(ReadHalf::Stream(reader)) => {
let mut reader = BufReader::with_capacity(buffer, reader);
tokio::io::copy_buf(&mut reader, &mut stdin).await?
}
Upstream::Stream(ReadHalf::Datagram(socket)) => {
let mut buf = Buffer::new(buffer);
loop {
let n = socket.recv(&mut buf).await?;
if let Some(counter) = &counter {
counter.add(n as u64);
}
stdin.write_all(&buf[..n]).await?;
stdin.flush().await?;
}
}
Upstream::Link(mut inlet) => {
let mut total = 0u64;
while let Some(parcel) = inlet.recv().await {
total += parcel.len() as u64;
stdin.write_all(&parcel).await?;
inlet.release(parcel);
}
total
}
};
stdin.flush().await?;
drop(stdin);
Ok::<u64, anyhow::Error>(total)
};
let drain = async {
let mut stdout = parts.stdout;
match output {
Downstream::Stream(WriteHalf::Datagram(socket)) => {
let mut buf = Buffer::new(buffer);
loop {
let n = stdout.read(&mut buf).await?;
if n == 0 {
break;
}
socket.send(&buf[..n]).await?;
}
}
Downstream::Stream(WriteHalf::Stream(mut writer)) => {
let mut stdout = BufReader::with_capacity(buffer, stdout);
tokio::io::copy_buf(&mut stdout, &mut writer).await?;
writer.flush().await?;
let _ = writer.shutdown().await;
}
Downstream::Link(mut outlet) => {
let mut buf = Buffer::new(buffer);
loop {
let n = stdout.read(&mut buf).await?;
if n == 0 {
break;
}
outlet.send(&buf[..n]).await?;
}
drop(outlet);
}
}
Ok::<(), anyhow::Error>(())
};
let diagnostics = async {
if let Some(stderr) = parts.stderr.take() {
let mut lines = BufReader::new(stderr).lines();
while let Some(line) = lines.next_line().await? {
warn!(stage = %name, "{line}");
}
}
Ok::<(), anyhow::Error>(())
};
let (total, (), ()) = tokio::try_join!(feed, drain, diagnostics)?;
let status = parts.child.wait().await.context("waiting on child")?;
if !status.success() {
bail!(
"stage `{}` exited {status}; its output is incomplete",
external.name
);
}
Ok(total)
}
enum Source {
Stream {
reader: BoxRead,
buf: Buffer,
},
Datagram {
socket: DatagramSocket,
buf: Buffer,
counter: Option<Counter>,
},
Link {
inlet: Inlet,
spent: Option<Vec<u8>>,
},
}
impl Source {
fn new(upstream: Upstream, buffer: usize, counter: Option<Counter>) -> Self {
match upstream {
Upstream::Stream(ReadHalf::Stream(reader)) => Source::Stream {
reader,
buf: Buffer::new(buffer),
},
Upstream::Stream(ReadHalf::Datagram(socket)) => Source::Datagram {
socket,
buf: Buffer::new(buffer),
counter,
},
Upstream::Link(inlet) => Source::Link { inlet, spent: None },
}
}
}
impl Source {
async fn next(&mut self) -> anyhow::Result<Option<&[u8]>> {
match self {
Source::Stream { reader, buf } => {
let n = reader.read(&mut buf[..]).await?;
Ok((n > 0).then(|| &buf[..n]))
}
Source::Datagram {
socket,
buf,
counter,
} => {
let n = socket.recv(&mut buf[..]).await?;
if let Some(counter) = counter {
counter.add(n as u64);
}
Ok(Some(&buf[..n]))
}
Source::Link { inlet, spent } => {
if let Some(parcel) = spent.take() {
inlet.release(parcel);
}
match inlet.recv().await {
Some(parcel) => Ok(Some(&spent.insert(parcel)[..])),
None => Ok(None),
}
}
}
}
}
enum Dest {
Stream(BoxWrite),
Datagram(DatagramSocket),
Link(Outlet),
}
impl From<Downstream> for Dest {
fn from(downstream: Downstream) -> Self {
match downstream {
Downstream::Stream(WriteHalf::Stream(writer)) => Dest::Stream(writer),
Downstream::Stream(WriteHalf::Datagram(socket)) => Dest::Datagram(socket),
Downstream::Link(outlet) => Dest::Link(outlet),
}
}
}
impl Dest {
async fn send(&mut self, emitted: Emitted<'_>) -> anyhow::Result<bool> {
match self {
Dest::Stream(writer) => {
if !emitted.is_empty() {
writer.write_all(emitted.bytes()).await?;
}
}
Dest::Datagram(socket) => {
for unit in emitted.units() {
if !unit.is_empty() {
socket.send(unit).await?;
}
}
}
Dest::Link(outlet) => {
for unit in emitted.units() {
if !outlet.send(unit).await? {
return Ok(false);
}
}
}
}
Ok(true)
}
async fn finish(self) -> anyhow::Result<()> {
match self {
Dest::Stream(mut writer) => {
writer.flush().await?;
let _ = writer.shutdown().await;
}
Dest::Datagram(_) => {}
Dest::Link(outlet) => drop(outlet),
}
Ok(())
}
}
async fn run_pipeline(
mut input: Source,
mut output: Dest,
mut pipeline: Pipeline,
channels: Arc<Channels>,
) -> anyhow::Result<u64> {
let mut effects = Effects::new(&channels);
let mut total = 0u64;
let mut ticker = pipeline.tick_interval().map(|period| {
let mut ticker = tokio::time::interval(period);
ticker.set_missed_tick_behavior(MissedTickBehavior::Delay);
ticker
});
loop {
let arrived = match ticker.as_mut() {
Some(ticker) => {
tokio::select! {
biased;
chunk = input.next() => chunk?,
_ = ticker.tick() => {
let alive =
drive_ticks(&mut pipeline, &mut output, &channels, &mut effects)
.await?;
if !alive || !flow(&mut effects).await {
break;
}
continue;
}
}
}
None => input.next().await?,
};
let Some(chunk) = arrived else {
break;
};
total += chunk.len() as u64;
let emitted = pipeline.process(chunk, &mut effects)?;
let alive = if effects.is_empty() {
output.send(emitted).await?
} else {
tokio::try_join!(output.send(emitted), channels.apply(&mut effects))?.0
};
if !alive || !flow(&mut effects).await {
break;
}
}
let emitted = pipeline.finish(&mut effects)?;
output.send(emitted).await?;
channels.apply(&mut effects).await?;
output.finish().await?;
Ok(total)
}
async fn flow(effects: &mut Effects) -> bool {
if let Some(reason) = effects.take_halt() {
info!("{reason}");
return false;
}
let pace = effects.take_pace();
if !pace.is_zero() {
sleep(pace).await;
}
true
}
async fn drive_ticks(
pipeline: &mut Pipeline,
output: &mut Dest,
channels: &Channels,
effects: &mut Effects,
) -> anyhow::Result<bool> {
let now = Instant::now();
let mut alive = true;
while let Some(emitted) = pipeline.tick(now, &mut *effects)? {
alive &= output.send(emitted).await?;
}
if !effects.is_empty() {
channels.apply(effects).await?;
}
Ok(alive)
}
struct Outlet {
data: mpsc::Sender<Vec<u8>>,
back: mpsc::Receiver<Vec<u8>>,
}
struct Inlet {
data: mpsc::Receiver<Vec<u8>>,
back: mpsc::Sender<Vec<u8>>,
}
fn link() -> (Outlet, Inlet) {
let (data_tx, data_rx) = mpsc::channel(LINK_DEPTH);
let (back_tx, back_rx) = mpsc::channel(LINK_DEPTH + 1);
(
Outlet {
data: data_tx,
back: back_rx,
},
Inlet {
data: data_rx,
back: back_tx,
},
)
}
impl Outlet {
async fn send(&mut self, bytes: &[u8]) -> anyhow::Result<bool> {
if bytes.is_empty() {
return Ok(true);
}
let mut buf = self
.back
.try_recv()
.ok()
.unwrap_or_else(|| Vec::with_capacity(bytes.len()));
buf.clear();
buf.extend_from_slice(bytes);
Ok(self.data.send(buf).await.is_ok())
}
}
impl Inlet {
async fn recv(&mut self) -> Option<Vec<u8>> {
self.data.recv().await
}
fn release(&self, buf: Vec<u8>) {
let _ = self.back.try_send(buf);
}
}