use std::{
fmt::Write as _,
io::{IsTerminal, Write as _},
pin::Pin,
sync::{
Arc,
atomic::{AtomicU64, AtomicUsize, Ordering},
},
task::{Context, Poll},
time::{Duration, Instant},
};
use serde::{Deserialize, Serialize};
use tocat_api::Direction;
use tokio::{
io::{AsyncRead, ReadBuf},
sync::oneshot,
task::JoinHandle,
time::MissedTickBehavior,
};
use tracing_subscriber::fmt::MakeWriter;
use crate::endpoint::{EndpointSpec, ReadHalf};
const REDRAW: Duration = Duration::from_millis(100);
const SMOOTHING: f64 = 0.25;
const FALLBACK_WIDTH: usize = 80;
const MIN_BAR: usize = 10;
const MAX_ETA: f64 = 359_999.0;
const BYTE_UNITS: &[&str] = &["B", "KiB", "MiB", "GiB", "TiB", "PiB"];
static ON_SCREEN: AtomicUsize = AtomicUsize::new(0);
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize, Serialize, clap::ValueEnum)]
#[serde(rename_all = "lowercase")]
pub enum ProgressMode {
#[default]
Never,
Auto,
Always,
}
pub struct Meter {
counts: [AtomicU64; 2],
connections: AtomicUsize,
started: Instant,
expected: Option<u64>,
}
fn slot(direction: Direction) -> usize {
match direction {
Direction::SourceToSink => 0,
Direction::SinkToSource => 1,
}
}
impl Meter {
#[must_use]
pub fn counter(self: &Arc<Self>, direction: Direction) -> Counter {
Counter {
meter: Arc::clone(self),
slot: slot(direction),
}
}
#[must_use]
pub fn connected(self: &Arc<Self>) -> ConnectionGuard {
self.connections.fetch_add(1, Ordering::Relaxed);
ConnectionGuard(Arc::clone(self))
}
fn read(&self) -> (u64, u64) {
(
self.counts[0].load(Ordering::Relaxed),
self.counts[1].load(Ordering::Relaxed),
)
}
}
#[derive(Clone)]
pub struct Counter {
meter: Arc<Meter>,
slot: usize,
}
impl Counter {
pub fn add(&self, bytes: u64) {
self.meter.counts[self.slot].fetch_add(bytes, Ordering::Relaxed);
}
}
pub struct ConnectionGuard(Arc<Meter>);
impl Drop for ConnectionGuard {
fn drop(&mut self) {
self.0.connections.fetch_sub(1, Ordering::Relaxed);
}
}
pub struct Counted<R> {
inner: R,
counter: Counter,
}
impl<R: AsyncRead + Unpin> AsyncRead for Counted<R> {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
let before = buf.filled().len();
let this = self.get_mut();
let poll = Pin::new(&mut this.inner).poll_read(cx, buf);
if poll.is_ready() {
let read = buf.filled().len().saturating_sub(before);
if read > 0 {
this.counter.add(read as u64);
}
}
poll
}
}
pub fn count(
meter: Option<&Arc<Meter>>,
half: ReadHalf,
direction: Direction,
) -> (ReadHalf, Option<Counter>) {
let Some(meter) = meter else {
return (half, None);
};
let counter = meter.counter(direction);
match half {
ReadHalf::Stream(reader) => (
ReadHalf::Stream(Box::new(Counted {
inner: reader,
counter,
})),
None,
),
ReadHalf::Datagram(socket) => (ReadHalf::Datagram(socket), Some(counter)),
}
}
pub struct Progress {
meter: Arc<Meter>,
stop: oneshot::Sender<()>,
painter: JoinHandle<()>,
}
impl Progress {
#[must_use]
pub fn meter(&self) -> Arc<Meter> {
Arc::clone(&self.meter)
}
pub async fn finish(self) {
let _ = self.stop.send(());
let _ = self.painter.await;
let (forward, reverse) = self.meter.read();
let moved = forward + reverse;
let mut out = std::io::stderr().lock();
let _ = erase(&mut out);
if moved > 0 {
let elapsed = self.meter.started.elapsed();
let seconds = elapsed.as_secs_f64();
let average = if seconds > 0.0 {
moved as f64 / seconds
} else {
0.0
};
let _ = writeln!(
out,
"{} in {} ({}/s)",
transferred(forward, reverse).trim_start(),
hms(elapsed),
bytes(average),
);
}
let _ = out.flush();
}
}
#[must_use]
pub fn start(mode: ProgressMode, source: &EndpointSpec, sink: &EndpointSpec) -> Option<Progress> {
let enabled = match mode {
ProgressMode::Never => false,
ProgressMode::Auto => std::io::stderr().is_terminal(),
ProgressMode::Always => true,
};
if !enabled {
return None;
}
let meter = Arc::new(Meter {
counts: [AtomicU64::new(0), AtomicU64::new(0)],
connections: AtomicUsize::new(0),
started: Instant::now(),
expected: expected_size(source, sink),
});
let (stop, halt) = oneshot::channel();
let painter = tokio::spawn(paint(Arc::clone(&meter), halt));
Some(Progress {
meter,
stop,
painter,
})
}
fn expected_size(source: &EndpointSpec, sink: &EndpointSpec) -> Option<u64> {
if source.is_fork() || sink.is_fork() {
return None;
}
match source {
EndpointSpec::File(e) => {
let metadata = std::fs::metadata(&e.path).ok()?;
metadata.is_file().then_some(metadata.len())
}
_ => None,
}
}
async fn paint(meter: Arc<Meter>, mut halt: oneshot::Receiver<()>) {
let mut painter = Painter::new(meter);
let mut ticker = tokio::time::interval(REDRAW);
ticker.set_missed_tick_behavior(MissedTickBehavior::Delay);
loop {
tokio::select! {
biased;
_ = &mut halt => break,
_ = ticker.tick() => painter.draw(),
}
}
}
struct Painter {
meter: Arc<Meter>,
rate: f64,
last: (Instant, u64),
line: String,
seen: bool,
}
impl Painter {
fn new(meter: Arc<Meter>) -> Self {
let started = meter.started;
Self {
meter,
rate: 0.0,
last: (started, 0),
line: String::new(),
seen: false,
}
}
fn draw(&mut self) {
let now = Instant::now();
let (forward, reverse) = self.meter.read();
let moved = forward + reverse;
let window = now.duration_since(self.last.0).as_secs_f64();
if window > 0.0 {
let instant = moved.saturating_sub(self.last.1) as f64 / window;
self.rate = if self.seen {
self.rate * (1.0 - SMOOTHING) + instant * SMOOTHING
} else {
instant
};
self.seen = moved > 0;
self.last = (now, moved);
}
let width = terminal_width();
self.compose(forward, reverse, now, width);
self.line.truncate(width);
let mut out = std::io::stderr().lock();
let previous = ON_SCREEN.swap(self.line.len(), Ordering::AcqRel);
let padding = previous.saturating_sub(self.line.len());
let _ = write!(out, "\r{}{:padding$}", self.line, "");
let _ = out.flush();
}
fn compose(&mut self, forward: u64, reverse: u64, now: Instant, width: usize) {
let elapsed = now.duration_since(self.meter.started);
self.line.clear();
let _ = write!(
self.line,
"{} {} [{:>9}/s]",
transferred(forward, reverse),
hms(elapsed),
bytes(self.rate),
);
let connections = self.meter.connections.load(Ordering::Relaxed);
if connections > 1 {
let _ = write!(self.line, " {connections} conns");
}
let Some(expected) = self.meter.expected.filter(|expected| *expected > 0) else {
return;
};
let done = forward.min(expected);
let fraction = done as f64 / expected as f64;
let tail = format!(
" {:>3.0}% ETA {}",
fraction * 100.0,
eta(expected - done, self.rate),
);
let room = width.saturating_sub(self.line.len() + tail.len() + 3);
if room >= MIN_BAR {
self.line.push(' ');
bar(&mut self.line, fraction, room);
}
self.line.push_str(&tail);
}
}
fn erase(out: &mut impl std::io::Write) -> std::io::Result<()> {
let width = ON_SCREEN.swap(0, Ordering::AcqRel);
if width > 0 {
write!(out, "\r{:width$}\r", "")?;
}
Ok(())
}
#[derive(Debug, Clone, Copy, Default)]
pub struct LogWriter;
impl std::io::Write for LogWriter {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
let mut out = std::io::stderr().lock();
erase(&mut out)?;
out.write(buf)
}
fn flush(&mut self) -> std::io::Result<()> {
std::io::stderr().lock().flush()
}
}
impl<'a> MakeWriter<'a> for LogWriter {
type Writer = LogWriter;
fn make_writer(&'a self) -> Self::Writer {
*self
}
}
fn transferred(forward: u64, reverse: u64) -> String {
if reverse == 0 {
format!("{:>9}", bytes(forward as f64))
} else {
format!(
"{:>9} out {:>9} in",
bytes(forward as f64),
bytes(reverse as f64)
)
}
}
fn bar(out: &mut String, fraction: f64, width: usize) {
let filled = ((fraction.clamp(0.0, 1.0) * width as f64).round() as usize).min(width);
out.push('[');
for cell in 0..width {
out.push(if cell + 1 < filled || filled == width {
'='
} else if cell < filled {
'>'
} else {
' '
});
}
out.push(']');
}
fn eta(remaining: u64, rate: f64) -> String {
if rate <= 1.0 {
return "--:--:--".to_string();
}
let seconds = remaining as f64 / rate;
if !seconds.is_finite() || seconds > MAX_ETA {
return "--:--:--".to_string();
}
hms(Duration::from_secs_f64(seconds))
}
fn bytes(value: f64) -> String {
let mut value = if value.is_finite() {
value.max(0.0)
} else {
0.0
};
let mut unit = 0;
while value >= 1024.0 && unit + 1 < BYTE_UNITS.len() {
value /= 1024.0;
unit += 1;
}
let digits = if unit == 0 || value >= 100.0 {
0
} else if value >= 10.0 {
1
} else {
2
};
format!("{value:.digits$}{}", BYTE_UNITS[unit])
}
fn hms(duration: Duration) -> String {
let seconds = duration.as_secs();
format!(
"{}:{:02}:{:02}",
seconds / 3600,
(seconds % 3600) / 60,
seconds % 60
)
}
#[cfg(unix)]
fn terminal_width() -> usize {
match rustix::termios::tcgetwinsize(std::io::stderr()) {
Ok(size) if size.ws_col > 0 => usize::from(size.ws_col),
_ => FALLBACK_WIDTH,
}
}
#[cfg(not(unix))]
fn terminal_width() -> usize {
FALLBACK_WIDTH
}
#[cfg(test)]
mod tests {
use super::*;
fn meter(expected: Option<u64>) -> Arc<Meter> {
Arc::new(Meter {
counts: [AtomicU64::new(0), AtomicU64::new(0)],
connections: AtomicUsize::new(0),
started: Instant::now(),
expected,
})
}
fn line(meter: &Arc<Meter>, rate: f64, width: usize) -> String {
let mut painter = Painter::new(Arc::clone(meter));
painter.rate = rate;
painter.compose(
meter.counts[0].load(Ordering::Relaxed),
meter.counts[1].load(Ordering::Relaxed),
Instant::now(),
width,
);
painter.line.clone()
}
#[test]
fn scales_byte_counts() {
assert_eq!(bytes(0.0), "0B");
assert_eq!(bytes(512.0), "512B");
assert_eq!(bytes(2048.0), "2.00KiB");
assert_eq!(bytes(1_073_741_824.0), "1.00GiB");
assert_eq!(bytes(f64::INFINITY), "0B");
}
#[test]
fn one_direction_shows_one_figure() {
let meter = meter(None);
meter.counts[0].store(2048, Ordering::Relaxed);
let line = line(&meter, 1024.0, 80);
assert!(line.contains("2.00KiB"), "{line}");
assert!(!line.contains(" in"), "nothing came back: {line}");
}
#[test]
fn both_directions_are_labelled() {
let meter = meter(None);
meter.counts[0].store(2048, Ordering::Relaxed);
meter.counts[1].store(1024, Ordering::Relaxed);
let line = line(&meter, 1024.0, 80);
assert!(line.contains("2.00KiB out"), "{line}");
assert!(line.contains("1.00KiB in"), "{line}");
}
#[test]
fn a_known_total_adds_a_bar_and_an_eta() {
let meter = meter(Some(1000));
meter.counts[0].store(500, Ordering::Relaxed);
let line = line(&meter, 100.0, 100);
assert!(line.contains(" 50%"), "{line}");
assert!(line.contains("ETA 0:00:05"), "{line}");
assert!(line.contains('='), "expected a bar: {line}");
}
#[test]
fn a_narrow_terminal_drops_the_bar_not_the_numbers() {
let meter = meter(Some(1000));
meter.counts[0].store(500, Ordering::Relaxed);
let line = line(&meter, 100.0, 44);
assert!(!line.contains('='), "no room for a bar: {line}");
assert!(line.contains(" 50%"), "{line}");
assert!(line.contains("100B/s"), "the rate survives: {line}");
}
#[test]
fn an_unknown_total_has_no_bar() {
let meter = meter(None);
meter.counts[0].store(500, Ordering::Relaxed);
let line = line(&meter, 100.0, 100);
assert!(!line.contains('%'), "{line}");
assert!(!line.contains("ETA"), "{line}");
}
#[test]
fn a_stalled_transfer_has_no_estimate() {
assert_eq!(eta(1000, 0.0), "--:--:--");
assert_eq!(eta(u64::MAX, 2.0), "--:--:--");
assert_eq!(eta(1000, 100.0), "0:00:10");
}
#[test]
fn the_bar_fills() {
let mut out = String::new();
bar(&mut out, 0.0, 4);
assert_eq!(out, "[ ]");
out.clear();
bar(&mut out, 0.5, 4);
assert_eq!(out, "[=> ]");
out.clear();
bar(&mut out, 1.0, 4);
assert_eq!(out, "[====]");
}
#[test]
fn frames_are_ascii() {
let meter = meter(Some(4096));
meter.counts[0].store(2048, Ordering::Relaxed);
meter.counts[1].store(64, Ordering::Relaxed);
let line = line(&meter, 1024.0, 100);
assert!(line.is_ascii(), "{line}");
}
}