use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use indicatif::{MultiProgress, ProgressBar, ProgressStyle};
use tokio::sync::{broadcast, mpsc, RwLock};
use tokio::task::JoinHandle;
use tracing::debug;
use crate::app::coordinator::DownloadStats;
use crate::errors::{DownloadError, DownloadResult};
#[derive(Debug, Clone)]
pub struct ProgressConfig {
pub enable_progress_bars: bool,
pub update_interval: Duration,
pub show_worker_details: bool,
pub enable_colors: bool,
pub show_download_rate: bool,
pub show_eta: bool,
pub max_filename_width: usize,
pub enable_spinner: bool,
pub compact_mode: bool,
}
impl Default for ProgressConfig {
fn default() -> Self {
Self {
enable_progress_bars: true,
update_interval: Duration::from_millis(100),
show_worker_details: true,
enable_colors: true,
show_download_rate: true,
show_eta: true,
max_filename_width: 40,
enable_spinner: true,
compact_mode: false,
}
}
}
#[derive(Debug, Clone)]
pub enum ProgressEvent {
FileCompleted {
worker_id: u32,
bytes_downloaded: u64,
file_name: String,
},
FileFailed {
worker_id: u32,
file_name: String,
error: String,
},
WorkerStatusChanged {
worker_id: u32,
status: String,
current_file: Option<String>,
},
StatsUpdate { stats: DownloadStats },
SessionCompleted {
total_files: usize,
successful: usize,
failed: usize,
duration: Duration,
},
}
#[derive(Debug, Clone)]
struct WorkerProgress {
#[allow(dead_code)]
worker_id: u32,
status: String,
current_file: Option<String>,
files_completed: usize,
bytes_downloaded: u64,
last_update: Instant,
}
pub struct ProgressDisplay {
config: ProgressConfig,
multi_progress: Option<MultiProgress>,
main_progress: Option<ProgressBar>,
worker_progress_bars: HashMap<u32, ProgressBar>,
worker_progress: Arc<RwLock<HashMap<u32, WorkerProgress>>>,
stats: Arc<RwLock<DownloadStats>>,
update_task: Option<JoinHandle<()>>,
event_tx: Option<mpsc::UnboundedSender<ProgressEvent>>,
shutdown_tx: Option<broadcast::Sender<()>>,
is_terminal: bool,
rate_tracking: Arc<RwLock<RateTracker>>,
}
#[derive(Debug)]
struct RateTracker {
rate_samples: Vec<(Instant, u64)>,
}
impl RateTracker {
fn new() -> Self {
Self {
rate_samples: Vec::new(),
}
}
fn update(&mut self, current_position: u64) -> u32 {
let now = Instant::now();
self.rate_samples.push((now, current_position));
let cutoff = now - Duration::from_secs(5);
self.rate_samples.retain(|(time, _)| *time > cutoff);
if self.rate_samples.len() >= 2 {
let oldest = &self.rate_samples[0];
let newest = &self.rate_samples[self.rate_samples.len() - 1];
let time_diff = newest.0.duration_since(oldest.0).as_secs_f64();
let position_diff = newest.1.saturating_sub(oldest.1);
if time_diff > 0.0 {
(position_diff as f64 / time_diff).round() as u32
} else {
0
}
} else {
0
}
}
}
impl ProgressDisplay {
pub fn new(config: ProgressConfig) -> Self {
let is_terminal = atty::is(atty::Stream::Stderr);
Self {
config,
multi_progress: None,
main_progress: None,
worker_progress_bars: HashMap::new(),
worker_progress: Arc::new(RwLock::new(HashMap::new())),
stats: Arc::new(RwLock::new(DownloadStats::default())),
update_task: None,
event_tx: None,
shutdown_tx: None,
is_terminal,
rate_tracking: Arc::new(RwLock::new(RateTracker::new())),
}
}
pub async fn start(&mut self, total_files: usize, worker_count: usize) -> DownloadResult<()> {
if !self.config.enable_progress_bars || !self.is_terminal {
return self.start_text_mode(total_files, worker_count).await;
}
let multi = MultiProgress::new();
let main_pb = multi.add(ProgressBar::new(total_files as u64));
main_pb.set_style(
ProgressStyle::default_bar()
.template(if self.config.show_eta {
"{spinner:.green} [{elapsed_precise}] [{bar:40.cyan/blue}] {pos}/{len} (ETA: {eta}) {msg}"
} else {
"{spinner:.green} [{elapsed_precise}] [{bar:40.cyan/blue}] {pos}/{len} {msg}"
})
.map_err(|e| DownloadError::Other(format!("Progress bar template error: {}", e)))?
.progress_chars("##-")
);
main_pb.enable_steady_tick(std::time::Duration::from_millis(100));
let mut worker_bars = HashMap::new();
if self.config.show_worker_details && !self.config.compact_mode {
for i in 0..worker_count {
let worker_pb = multi.add(ProgressBar::new_spinner());
worker_pb.set_style(
ProgressStyle::default_spinner()
.template(" Worker {prefix}: {spinner:.blue} {msg}")
.map_err(|e| {
DownloadError::Other(format!("Worker progress template error: {}", e))
})?,
);
worker_pb.set_prefix(match i {
0 => "1",
1 => "2",
2 => "3",
3 => "4",
4 => "5",
5 => "6",
6 => "7",
7 => "8",
_ => "N",
});
worker_pb.set_message("Initializing...");
worker_bars.insert(i as u32 + 1, worker_pb);
}
}
let (event_tx, event_rx) = mpsc::unbounded_channel();
let (shutdown_tx, shutdown_rx) = broadcast::channel(1);
{
let mut stats = self.stats.write().await;
stats.total_files = total_files;
stats.active_workers = worker_count;
}
let update_task = self.start_update_task(event_rx, shutdown_rx).await;
self.multi_progress = Some(multi);
self.main_progress = Some(main_pb);
self.worker_progress_bars = worker_bars;
self.event_tx = Some(event_tx);
self.shutdown_tx = Some(shutdown_tx);
self.update_task = Some(update_task);
debug!(
"Progress display started for {} files with {} workers",
total_files, worker_count
);
Ok(())
}
async fn start_text_mode(
&mut self,
total_files: usize,
worker_count: usize,
) -> DownloadResult<()> {
let (event_tx, mut event_rx) = mpsc::unbounded_channel();
let (shutdown_tx, mut shutdown_rx) = broadcast::channel(1);
{
let mut stats = self.stats.write().await;
stats.total_files = total_files;
stats.active_workers = worker_count;
}
let stats = self.stats.clone();
let update_interval = self.config.update_interval;
let update_task = tokio::spawn(async move {
let mut last_report = Instant::now();
let report_interval = Duration::from_secs(10);
loop {
tokio::select! {
event = event_rx.recv() => {
match event {
Some(ProgressEvent::FileCompleted { file_name: _, .. }) => {
if last_report.elapsed() >= report_interval {
let stats_guard = stats.read().await;
eprintln!("Progress: {}/{} files completed ({:.1}%)",
stats_guard.files_completed,
stats_guard.total_files,
(stats_guard.files_completed as f64 / stats_guard.total_files as f64) * 100.0);
last_report = Instant::now();
}
}
Some(ProgressEvent::SessionCompleted { successful, failed, duration, .. }) => {
eprintln!("Download completed: {} successful, {} failed in {:?}", successful, failed, duration);
break;
}
Some(_) => {} None => break,
}
}
_ = shutdown_rx.recv() => {
break;
}
_ = tokio::time::sleep(update_interval) => {
}
}
}
});
self.event_tx = Some(event_tx);
self.shutdown_tx = Some(shutdown_tx);
self.update_task = Some(update_task);
eprintln!(
"Starting download of {} files with {} workers...",
total_files, worker_count
);
Ok(())
}
pub async fn update(&self, event: ProgressEvent) -> DownloadResult<()> {
if let Some(tx) = &self.event_tx {
tx.send(event).map_err(|e| {
DownloadError::Other(format!("Failed to send progress event: {}", e))
})?;
}
Ok(())
}
pub async fn update_with_stats(&self, completed: usize, failed: usize) -> DownloadResult<()> {
if let Some(main_pb) = &self.main_progress {
main_pb.set_position(completed as u64);
let rate = {
let mut tracker = self.rate_tracking.write().await;
tracker.update(completed as u64)
};
if self.config.show_download_rate {
let message = if failed > 0 {
format!("{}/s {} failed", rate, failed)
} else {
format!("{}/s", rate)
};
main_pb.set_message(message);
} else if failed > 0 {
main_pb.set_message(format!("{} failed", failed));
}
}
Ok(())
}
pub async fn finish(&mut self) -> DownloadResult<()> {
debug!("Finishing progress display");
if let Some(tx) = &self.shutdown_tx {
let _ = tx.send(());
}
if let Some(task) = self.update_task.take() {
let _ = task.await;
}
if self.config.enable_progress_bars && self.is_terminal {
if let Some(main_pb) = &self.main_progress {
main_pb.finish_with_message("Download completed");
}
for worker_pb in self.worker_progress_bars.values() {
worker_pb.finish_and_clear();
}
}
Ok(())
}
async fn start_update_task(
&self,
mut event_rx: mpsc::UnboundedReceiver<ProgressEvent>,
mut shutdown_rx: broadcast::Receiver<()>,
) -> JoinHandle<()> {
let main_pb = self.main_progress.clone();
let worker_bars = self.worker_progress_bars.clone();
let worker_progress = self.worker_progress.clone();
let stats = self.stats.clone();
let update_interval = self.config.update_interval;
let config = self.config.clone();
tokio::spawn(async move {
let mut last_update = Instant::now();
loop {
tokio::select! {
event = event_rx.recv() => {
match event {
Some(event) => {
Self::handle_progress_event(
&event,
&main_pb,
&worker_bars,
&worker_progress,
&stats,
&config
).await;
}
None => {
debug!("Progress event channel closed");
break;
}
}
}
_ = shutdown_rx.recv() => {
debug!("Progress display received shutdown signal");
break;
}
_ = tokio::time::sleep(update_interval) => {
if last_update.elapsed() >= update_interval {
Self::periodic_update(&main_pb, &stats).await;
last_update = Instant::now();
}
}
}
}
})
}
async fn handle_progress_event(
event: &ProgressEvent,
main_pb: &Option<ProgressBar>,
worker_bars: &HashMap<u32, ProgressBar>,
worker_progress: &Arc<RwLock<HashMap<u32, WorkerProgress>>>,
stats: &Arc<RwLock<DownloadStats>>,
config: &ProgressConfig,
) {
match event {
ProgressEvent::FileCompleted {
worker_id,
bytes_downloaded,
file_name,
} => {
if let Some(pb) = main_pb {
pb.inc(1);
}
{
let mut worker_map = worker_progress.write().await;
let worker = worker_map
.entry(*worker_id)
.or_insert_with(|| WorkerProgress {
worker_id: *worker_id,
status: "Working".to_string(),
current_file: None,
files_completed: 0,
bytes_downloaded: 0,
last_update: Instant::now(),
});
worker.files_completed += 1;
worker.bytes_downloaded += bytes_downloaded;
worker.current_file = None;
worker.last_update = Instant::now();
}
if let Some(worker_pb) = worker_bars.get(worker_id) {
let truncated_name = if file_name.len() > config.max_filename_width {
format!(
"...{}",
&file_name[file_name.len() - config.max_filename_width + 3..]
)
} else {
file_name.clone()
};
worker_pb.set_message(format!("✅ Completed: {}", truncated_name));
}
{
let mut stats_guard = stats.write().await;
stats_guard.files_completed += 1;
stats_guard.total_bytes_downloaded += bytes_downloaded;
}
}
ProgressEvent::FileFailed {
worker_id,
file_name,
error,
} => {
if let Some(worker_pb) = worker_bars.get(worker_id) {
let truncated_name = if file_name.len() > config.max_filename_width {
format!(
"...{}",
&file_name[file_name.len() - config.max_filename_width + 3..]
)
} else {
file_name.clone()
};
worker_pb.set_message(format!("❌ Failed: {} ({})", truncated_name, error));
}
{
let mut stats_guard = stats.write().await;
stats_guard.files_failed += 1;
}
}
ProgressEvent::WorkerStatusChanged {
worker_id,
status,
current_file,
} => {
{
let mut worker_map = worker_progress.write().await;
let worker = worker_map
.entry(*worker_id)
.or_insert_with(|| WorkerProgress {
worker_id: *worker_id,
status: status.clone(),
current_file: current_file.clone(),
files_completed: 0,
bytes_downloaded: 0,
last_update: Instant::now(),
});
worker.status = status.clone();
worker.current_file = current_file.clone();
worker.last_update = Instant::now();
}
if let Some(worker_pb) = worker_bars.get(worker_id) {
let message = if let Some(file) = current_file {
let truncated_name = if file.len() > config.max_filename_width {
format!("...{}", &file[file.len() - config.max_filename_width + 3..])
} else {
file.clone()
};
format!("{}: {}", status, truncated_name)
} else {
status.clone()
};
worker_pb.set_message(message);
}
}
ProgressEvent::StatsUpdate { stats: new_stats } => {
{
let mut stats_guard = stats.write().await;
*stats_guard = new_stats.clone();
}
if let Some(pb) = main_pb {
pb.set_position(new_stats.files_completed as u64);
}
}
ProgressEvent::SessionCompleted {
total_files: _,
successful,
failed,
duration,
} => {
if let Some(pb) = main_pb {
pb.finish_with_message(format!(
"✅ Completed: {} successful, {} failed in {:?}",
successful, failed, duration
));
}
for worker_pb in worker_bars.values() {
worker_pb.finish_and_clear();
}
}
}
}
async fn periodic_update(main_pb: &Option<ProgressBar>, stats: &Arc<RwLock<DownloadStats>>) {
if let Some(pb) = main_pb {
let stats_guard = stats.read().await;
pb.set_position(stats_guard.files_completed as u64);
}
}
}
impl Drop for ProgressDisplay {
fn drop(&mut self) {
}
}
#[cfg(test)]
mod tests {
use super::*;
fn create_test_config() -> ProgressConfig {
ProgressConfig {
enable_progress_bars: false, update_interval: Duration::from_millis(1),
show_worker_details: true,
enable_colors: false,
show_download_rate: true,
show_eta: true,
max_filename_width: 20,
enable_spinner: false,
compact_mode: true,
}
}
#[tokio::test]
async fn test_progress_display_creation() {
let config = create_test_config();
let display = ProgressDisplay::new(config.clone());
assert_eq!(display.config.update_interval, config.update_interval);
assert_eq!(
display.config.show_worker_details,
config.show_worker_details
);
assert!(display.multi_progress.is_none());
assert!(display.event_tx.is_none());
}
#[tokio::test]
async fn test_progress_events() {
let config = create_test_config();
let mut display = ProgressDisplay::new(config);
display.start(10, 2).await.unwrap();
let event = ProgressEvent::FileCompleted {
worker_id: 1,
bytes_downloaded: 1024,
file_name: "test.csv".to_string(),
};
display.update(event).await.unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
{
let stats = display.stats.read().await;
assert_eq!(stats.total_files, 10);
}
display.finish().await.unwrap();
}
#[tokio::test]
async fn test_text_mode_fallback() {
let mut config = create_test_config();
config.enable_progress_bars = false;
let mut display = ProgressDisplay::new(config);
let result = display.start(5, 1).await;
assert!(result.is_ok());
let event = ProgressEvent::FileCompleted {
worker_id: 1,
bytes_downloaded: 512,
file_name: "test.txt".to_string(),
};
let result = display.update(event).await;
assert!(result.is_ok());
let result = display.finish().await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_progress_config_defaults() {
let config = ProgressConfig::default();
assert!(config.enable_progress_bars);
assert!(config.show_worker_details);
assert!(config.enable_colors);
assert!(config.show_download_rate);
assert!(config.show_eta);
assert!(config.update_interval > Duration::ZERO);
assert!(config.max_filename_width > 0);
}
#[tokio::test]
async fn test_filename_truncation() {
let config = ProgressConfig {
max_filename_width: 10,
..create_test_config()
};
let mut display = ProgressDisplay::new(config);
display.start(1, 1).await.unwrap();
let event = ProgressEvent::FileCompleted {
worker_id: 1,
bytes_downloaded: 1024,
file_name: "very_long_filename_that_should_be_truncated.csv".to_string(),
};
let result = display.update(event).await;
assert!(result.is_ok());
display.finish().await.unwrap();
}
#[tokio::test]
async fn test_session_completion() {
let config = create_test_config();
let mut display = ProgressDisplay::new(config);
display.start(3, 1).await.unwrap();
let event = ProgressEvent::SessionCompleted {
total_files: 3,
successful: 2,
failed: 1,
duration: Duration::from_secs(30),
};
let result = display.update(event).await;
assert!(result.is_ok());
display.finish().await.unwrap();
}
}