Skip to main content

aria2_core/engine/
download_command.rs

1use async_trait::async_trait;
2use std::sync::Arc;
3use std::time::Duration;
4use tokio::io::AsyncReadExt;
5use tokio::sync::mpsc;
6use tracing::{debug, info, warn};
7
8use crate::checksum::checksum::Checksum;
9use crate::checksum::message_digest::HashType;
10use crate::constants;
11use crate::engine::active_output_registry::global_registry;
12use crate::engine::command::{Command, CommandStatus, ProgressUpdate};
13use crate::engine::concurrent_download::{ConcurrentDownloadResult, ConcurrentDownloader};
14use crate::engine::download_cookie::CookieHelper;
15use crate::engine::download_progress::ProgressUpdater;
16use crate::engine::range_prober::RangeProber;
17use crate::engine::retry_policy::RetryPolicy;
18use crate::engine::sequential_download::SequentialDownloader;
19use crate::error::{Aria2Error, Result};
20use crate::filesystem::file_allocation;
21use crate::filesystem::resume_helper::ResumeHelper;
22use crate::http::client_pool;
23use crate::http::cookie::Cookie;
24use crate::http::cookie_storage::CookieStorage;
25use crate::http::socks_connector::{NoProxyMatcher, ProxyUrl};
26use crate::request::request_group::{DownloadOptions, GroupId, RequestGroup};
27use crate::selector::server_stat_man::ServerStatMan;
28use crate::util::perf_monitor::{AtomicMetrics, Metrics, PerformanceMonitor};
29
30pub struct DownloadCommand {
31    group: Arc<tokio::sync::RwLock<RequestGroup>>,
32    client: Arc<reqwest::Client>,
33    output_path: std::path::PathBuf,
34    started: bool,
35    completed: bool,
36    completed_bytes: u64,
37    file_allocation: String,
38    mmap_threshold: u64,
39    secure_falloc: bool,
40    cookie_storage: Arc<CookieStorage>,
41    cookie_file: Option<String>,
42    no_proxy_matcher: Option<NoProxyMatcher>,
43    stat_man: Arc<ServerStatMan>,
44    perf_monitor: Option<Arc<PerformanceMonitor>>,
45    atomic_metrics: Arc<AtomicMetrics>,
46    headers: Vec<(String, String)>,
47    progress_sender: Option<mpsc::UnboundedSender<ProgressUpdate>>,
48    progress_receiver: Option<mpsc::UnboundedReceiver<ProgressUpdate>>,
49    progress_aggregator_handle: Option<tokio::task::JoinHandle<()>>,
50}
51
52impl DownloadCommand {
53    pub fn new(
54        gid: GroupId,
55        uri: &str,
56        options: &DownloadOptions,
57        output_dir: Option<&str>,
58        output_name: Option<&str>,
59    ) -> Result<Self> {
60        let group = Arc::new(tokio::sync::RwLock::new(RequestGroup::new(
61            gid,
62            vec![uri.to_string()],
63            options.clone(),
64        )));
65        Self::new_with_group(group, uri, options, output_dir, output_name)
66    }
67
68    pub fn new_with_group(
69        group: Arc<tokio::sync::RwLock<RequestGroup>>,
70        uri: &str,
71        options: &DownloadOptions,
72        output_dir: Option<&str>,
73        output_name: Option<&str>,
74    ) -> Result<Self> {
75        let dir = output_dir
76            .map(|d| d.to_string())
77            .or_else(|| options.dir.clone())
78            .unwrap_or_else(|| constants::DEFAULT_OUTPUT_DIR.to_string());
79
80        let filename = output_name
81            .map(|n| n.to_string())
82            .or_else(|| Self::extract_filename(uri))
83            .unwrap_or_else(|| constants::DEFAULT_FILENAME.to_string());
84
85        let path = std::path::PathBuf::from(&dir).join(&filename);
86        let headers = options.parsed_headers();
87
88        let no_proxy = options.http_proxy.is_none() && options.all_proxy.is_none();
89        let client = if no_proxy {
90            client_pool::get_global_client()
91        } else {
92            let mut builder = reqwest::Client::builder()
93                .connect_timeout(Duration::from_secs(
94                    constants::HTTP_DEFAULT_CONNECT_TIMEOUT_SECS,
95                ))
96                .timeout(Duration::from_secs(
97                    constants::HTTP_DEFAULT_OVERALL_TIMEOUT_SECS,
98                ))
99                .user_agent(constants::USER_AGENT)
100                .redirect(reqwest::redirect::Policy::limited(
101                    constants::HTTP_DEFAULT_MAX_REDIRECTS,
102                ))
103                .pool_max_idle_per_host(constants::HTTP_DEFAULT_POOL_MAX_IDLE_PER_HOST)
104                .pool_idle_timeout(Some(std::time::Duration::from_secs(
105                    constants::HTTP_DEFAULT_POOL_IDLE_TIMEOUT_SECS,
106                )))
107                .tcp_keepalive(Some(std::time::Duration::from_secs(
108                    constants::HTTP_DEFAULT_TCP_KEEPALIVE_SECS,
109                )));
110
111            if let Some(ref proxy) = options.http_proxy
112                && let Ok(proxy_url) = proxy.parse::<reqwest::Url>()
113                && let Ok(p) = reqwest::Proxy::all(proxy_url.to_string())
114            {
115                builder = builder.proxy(p);
116            }
117
118            if options.http_proxy.is_none()
119                && let Some(ref all_proxy) = options.all_proxy
120            {
121                match ProxyUrl::parse(all_proxy) {
122                    Ok(parsed) => match parsed.protocol {
123                        crate::http::socks_connector::ProxyProtocol::Http
124                        | crate::http::socks_connector::ProxyProtocol::Https => {
125                            if let Ok(p) = reqwest::Proxy::all(all_proxy.to_string()) {
126                                builder = builder.proxy(p);
127                            }
128                        }
129                        _ => {
130                            tracing::info!(
131                                "SOCKS proxy configured ({}) - use SocksConnector for direct TCP connections",
132                                all_proxy
133                            );
134                        }
135                    },
136                    Err(e) => {
137                        warn!("Failed to parse all-proxy URL '{}': {}", all_proxy, e);
138                    }
139                }
140            }
141
142            let client = builder.build().map_err(|e| {
143                Aria2Error::Fatal(crate::error::FatalError::Config(format!(
144                    "Failed to build HTTP client: {}",
145                    e
146                )))
147            })?;
148
149            Arc::new(client)
150        };
151
152        info!("DownloadCommand created: {} -> {}", uri, path.display());
153
154        let cookie_file = options.cookie_file.clone();
155        let cookie_storage = Arc::new(CookieStorage::new());
156
157        Self::load_cookies(&cookie_storage, &cookie_file, uri, options);
158
159        let (progress_tx, progress_rx) = mpsc::unbounded_channel::<ProgressUpdate>();
160
161        Ok(Self {
162            group,
163            client,
164            output_path: path,
165            started: false,
166            completed: false,
167            completed_bytes: 0,
168            file_allocation: options
169                .file_allocation
170                .clone()
171                .unwrap_or_else(|| constants::DEFAULT_FILE_ALLOCATION.to_string()),
172            mmap_threshold: options.mmap_threshold.unwrap_or(256 * 1024 * 1024),
173            secure_falloc: options.secure_falloc,
174            cookie_storage,
175            cookie_file,
176            no_proxy_matcher: options
177                .no_proxy
178                .as_ref()
179                .map(|np| NoProxyMatcher::from_env_value(np)),
180            stat_man: Arc::new(ServerStatMan::new()),
181            perf_monitor: None,
182            atomic_metrics: Arc::new(AtomicMetrics::new()),
183            headers,
184            progress_sender: Some(progress_tx),
185            progress_receiver: Some(progress_rx),
186            progress_aggregator_handle: None,
187        })
188    }
189
190    fn load_cookies(
191        cookie_storage: &Arc<CookieStorage>,
192        cookie_file: &Option<String>,
193        uri: &str,
194        options: &DownloadOptions,
195    ) {
196        if let Some(cf) = cookie_file {
197            let p = std::path::Path::new(cf);
198            if p.exists() {
199                match cookie_storage.load_file(p) {
200                    Ok(n) => info!("Loaded {} cookies from file: {}", n, cf),
201                    Err(e) => warn!("Failed to load cookie file {}: {}", cf, e),
202                }
203            }
204        }
205
206        if let Some(ref cookies_str) = options.cookies {
207            let domain = Self::extract_host(uri);
208            for pair in cookies_str.split(';') {
209                let pair = pair.trim();
210                if pair.is_empty() {
211                    continue;
212                }
213                if let Some((name, value)) = pair.split_once('=') {
214                    let name = name.trim();
215                    let value = value.trim();
216                    if !name.is_empty() {
217                        cookie_storage.add(Cookie::new(name, value, &domain));
218                    }
219                }
220            }
221            if !cookie_storage.is_empty() {
222                info!("Manually set {} cookies", cookie_storage.count());
223            }
224        }
225    }
226
227    pub fn new_with_client(
228        gid: GroupId,
229        uri: &str,
230        options: &DownloadOptions,
231        output_dir: Option<&str>,
232        output_name: Option<&str>,
233        client: Arc<reqwest::Client>,
234    ) -> Result<Self> {
235        let group = Arc::new(tokio::sync::RwLock::new(RequestGroup::new(
236            gid,
237            vec![uri.to_string()],
238            options.clone(),
239        )));
240        Self::new_with_group_and_client(group, uri, options, output_dir, output_name, client)
241    }
242
243    pub fn new_with_group_and_client(
244        group: Arc<tokio::sync::RwLock<RequestGroup>>,
245        uri: &str,
246        options: &DownloadOptions,
247        output_dir: Option<&str>,
248        output_name: Option<&str>,
249        client: Arc<reqwest::Client>,
250    ) -> Result<Self> {
251        let dir = output_dir
252            .map(|d| d.to_string())
253            .or_else(|| options.dir.clone())
254            .unwrap_or_else(|| constants::DEFAULT_OUTPUT_DIR.to_string());
255
256        let filename = output_name
257            .map(|n| n.to_string())
258            .or_else(|| Self::extract_filename(uri))
259            .unwrap_or_else(|| constants::DEFAULT_FILENAME.to_string());
260
261        let path = std::path::PathBuf::from(&dir).join(&filename);
262
263        let headers = options.parsed_headers();
264        info!(
265            "DownloadCommand created (shared client): {} -> {}",
266            uri,
267            path.display()
268        );
269
270        let cookie_file = options.cookie_file.clone();
271        let cookie_storage = Arc::new(CookieStorage::new());
272
273        Self::load_cookies(&cookie_storage, &cookie_file, uri, options);
274
275        let (progress_tx, progress_rx) = mpsc::unbounded_channel::<ProgressUpdate>();
276
277        Ok(Self {
278            group,
279            client,
280            output_path: path,
281            started: false,
282            completed: false,
283            completed_bytes: 0,
284            file_allocation: options
285                .file_allocation
286                .clone()
287                .unwrap_or_else(|| constants::DEFAULT_FILE_ALLOCATION.to_string()),
288            mmap_threshold: options.mmap_threshold.unwrap_or(256 * 1024 * 1024),
289            secure_falloc: options.secure_falloc,
290            cookie_storage,
291            cookie_file,
292            no_proxy_matcher: options
293                .no_proxy
294                .as_ref()
295                .map(|np| NoProxyMatcher::from_env_value(np)),
296            stat_man: Arc::new(ServerStatMan::new()),
297            perf_monitor: None,
298            atomic_metrics: Arc::new(AtomicMetrics::new()),
299            headers,
300            progress_sender: Some(progress_tx),
301            progress_receiver: Some(progress_rx),
302            progress_aggregator_handle: None,
303        })
304    }
305
306    pub fn new_with_stat_man(
307        gid: GroupId,
308        uri: &str,
309        options: &DownloadOptions,
310        output_dir: Option<&str>,
311        output_name: Option<&str>,
312        stat_man: Arc<ServerStatMan>,
313    ) -> Result<Self> {
314        let mut cmd = Self::new(gid, uri, options, output_dir, output_name)?;
315        cmd.stat_man = stat_man;
316        Ok(cmd)
317    }
318
319    pub fn enable_perf_monitor(&mut self) {
320        self.perf_monitor = Some(Arc::new(PerformanceMonitor::new()));
321    }
322
323    #[allow(dead_code)]
324    pub(crate) fn with_progress_sender(
325        mut self,
326        sender: mpsc::UnboundedSender<ProgressUpdate>,
327    ) -> Self {
328        self.progress_sender = Some(sender);
329        self.progress_receiver = None;
330        self
331    }
332
333    pub(crate) fn spawn_progress_aggregator(&mut self) {
334        if self.progress_aggregator_handle.is_some() {
335            return;
336        }
337        if let Some(rx) = self.progress_receiver.take() {
338            let handle = crate::engine::download_engine::DownloadEngine::spawn_progress_aggregator(
339                Arc::clone(&self.group),
340                rx,
341            );
342            self.progress_aggregator_handle = Some(handle);
343        }
344    }
345
346    pub(crate) async fn drain_progress_aggregator(&mut self) {
347        self.progress_sender = None;
348        if let Some(handle) = self.progress_aggregator_handle.take()
349            && let Err(e) = handle.await
350        {
351            warn!("Progress aggregator task ended unexpectedly: {}", e);
352        }
353    }
354
355    pub fn get_perf_metrics(&self) -> Metrics {
356        self.atomic_metrics.snapshot()
357    }
358
359    pub fn get_perf_report(&self) -> Option<String> {
360        self.perf_monitor.as_ref().map(|m| m.export_text())
361    }
362
363    pub fn get_perf_report_json(&self) -> Option<String> {
364        self.perf_monitor.as_ref().map(|m| m.export_json())
365    }
366
367    fn extract_filename(uri: &str) -> Option<String> {
368        uri.rsplit('/')
369            .next()
370            .filter(|s| !s.is_empty() && *s != "/")
371            .map(|s| s.split('?').next().unwrap_or(s).to_string())
372    }
373
374    fn extract_host(uri: &str) -> String {
375        reqwest::Url::parse(uri)
376            .ok()
377            .and_then(|u| u.host_str().map(|h| h.to_string()))
378            .unwrap_or_else(|| constants::DEFAULT_HOST.to_string())
379    }
380
381    pub async fn group(&self) -> tokio::sync::RwLockReadGuard<'_, RequestGroup> {
382        self.group.read().await
383    }
384
385    pub async fn group_mut(&self) -> tokio::sync::RwLockWriteGuard<'_, RequestGroup> {
386        self.group.write().await
387    }
388
389    pub fn no_proxy_matcher(&self) -> Option<&NoProxyMatcher> {
390        self.no_proxy_matcher.as_ref()
391    }
392
393    fn should_use_concurrent(&self, total_length: u64, supports_range: bool, split: u16) -> bool {
394        if !supports_range {
395            return false;
396        }
397        if total_length < constants::CONCURRENT_MIN_FILE_SIZE as u64 {
398            return false;
399        }
400        split > 1
401    }
402
403    fn create_cookie_helper(&self) -> CookieHelper {
404        CookieHelper::new(Arc::clone(&self.cookie_storage), self.cookie_file.clone())
405    }
406
407    fn create_progress_updater(&self) -> ProgressUpdater {
408        ProgressUpdater::new(
409            self.progress_sender.clone(),
410            Arc::clone(&self.group),
411            Arc::clone(&self.atomic_metrics),
412            self.perf_monitor.clone(),
413        )
414    }
415}
416
417#[async_trait]
418impl Command for DownloadCommand {
419    async fn execute(&mut self) -> Result<()> {
420        if !self.started {
421            self.group.write().await.start().await?;
422            self.started = true;
423        }
424
425        let uri = {
426            let g = self.group.read().await;
427            g.uris().first().cloned().unwrap_or_default()
428        };
429
430        if uri.is_empty() {
431            return Err(Aria2Error::Fatal(crate::error::FatalError::Config(
432                "Download URI is empty".into(),
433            )));
434        }
435
436        debug!(
437            "Starting download: {} -> {}",
438            uri,
439            self.output_path.display()
440        );
441
442        if let Some(parent) = self.output_path.parent()
443            && !parent.exists()
444        {
445            std::fs::create_dir_all(parent).map_err(|e| {
446                Aria2Error::Fatal(crate::error::FatalError::Config(format!(
447                    "Failed to create directory: {}",
448                    e
449                )))
450            })?;
451        }
452
453        let original_path = self.output_path.clone();
454        self.output_path = global_registry().resolve(&original_path).await;
455        if self.output_path != original_path {
456            info!(
457                "Filename collision resolved: '{}' -> '{}'",
458                original_path.display(),
459                self.output_path.display()
460            );
461        }
462
463        let release_path = |path: &std::path::Path| {
464            let p = path.to_path_buf();
465            #[allow(clippy::let_underscore_future)]
466            let _ = tokio::spawn(async move {
467                global_registry().release(&p).await;
468            });
469        };
470
471        let url_for_head = reqwest::Url::parse(&uri).ok();
472        let cookie_hdr_head = if let Some(ref url) = url_for_head {
473            CookieHelper::new(Arc::clone(&self.cookie_storage), self.cookie_file.clone())
474                .build_cookie_header_from_url(url)
475        } else {
476            String::new()
477        };
478        let mut head_req = self.client.head(&uri);
479        if !cookie_hdr_head.is_empty() {
480            head_req = head_req.header("Cookie", &cookie_hdr_head);
481        }
482        for (name, value) in &self.headers {
483            head_req = head_req.header(name, value);
484        }
485        let head_resp = head_req.send().await.ok();
486        let (total_length, head_supports_range) = if let Some(ref resp) = head_resp {
487            let tl = resp
488                .headers()
489                .get(reqwest::header::CONTENT_LENGTH)
490                .and_then(|v| v.to_str().ok())
491                .and_then(|v| v.parse::<u64>().ok())
492                .unwrap_or(0);
493            let sr = resp
494                .headers()
495                .get("Accept-Ranges")
496                .and_then(|v| v.to_str().ok())
497                .is_some_and(|v| v.to_lowercase().contains("bytes"));
498            (tl, sr)
499        } else {
500            (0, false)
501        };
502
503        let supports_range = if head_supports_range {
504            true
505        } else if total_length > constants::CONCURRENT_MIN_FILE_SIZE as u64 {
506            let prober = RangeProber::new(Arc::clone(&self.client), self.headers.clone());
507            prober.probe_range_support(&uri, total_length).await
508        } else {
509            false
510        };
511
512        let resume_helper = ResumeHelper::new(&self.output_path, true);
513        let resume_state = resume_helper.detect(total_length).await?;
514
515        if resume_state.is_complete {
516            info!(
517                "File already exists completely, skipping download: {} ({} bytes)",
518                self.output_path.display(),
519                resume_state.existing_length
520            );
521            self.completed_bytes = resume_state.existing_length;
522            let mut g = self.group.write().await;
523            g.set_total_length(self.completed_bytes).await;
524            g.update_progress(self.completed_bytes).await;
525            g.set_total_length_atomic(self.completed_bytes);
526            g.set_completed_length(self.completed_bytes);
527            g.complete().await?;
528            self.completed = true;
529            release_path(&self.output_path);
530            return Ok(());
531        }
532
533        self.spawn_progress_aggregator();
534
535        let download_result: Result<()> = async {
536            if total_length > 0 {
537                file_allocation::preallocate_file(
538                    &self.output_path,
539                    total_length,
540                    &self.file_allocation,
541                    self.secure_falloc,
542                )
543                .await?;
544            }
545
546            let options = self.group.read().await.options().clone();
547            let split = options.split.unwrap_or(constants::DEFAULT_SPLIT);
548
549            let cookie_helper = self.create_cookie_helper();
550            let progress_updater = self.create_progress_updater();
551
552            if self.should_use_concurrent(total_length, supports_range, split) {
553                if resume_state.should_resume {
554                    info!(
555                        "Concurrent mode + resume: existing {} bytes, continuing from offset {}",
556                        resume_state.existing_length, resume_state.start_offset
557                    );
558                }
559                let max_retries = options.max_retries;
560                let mut concurrent_downloader = ConcurrentDownloader::new(
561                    Arc::clone(&self.client),
562                    self.output_path.clone(),
563                    self.headers.clone(),
564                    cookie_helper.clone(),
565                    progress_updater.clone(),
566                    Arc::clone(&self.group),
567                    self.mmap_threshold,
568                    self.file_allocation.clone(),
569                );
570                match concurrent_downloader.execute_with_retry(
571                    &uri,
572                    total_length,
573                    &resume_state,
574                    max_retries,
575                ).await {
576                    Ok(ConcurrentDownloadResult::Complete) => return Ok(()),
577                    Ok(ConcurrentDownloadResult::Fallback { completed_ranges }) => {
578                        warn!(
579                            "Concurrent download falling back to sequential mode, preserving {} completed ranges",
580                            completed_ranges.len()
581                        );
582                        let retry_policy = RetryPolicy::new(options.max_retries, options.retry_wait * 1000);
583                        let mut sequential_downloader = SequentialDownloader::new(
584                            Arc::clone(&self.client),
585                            self.output_path.clone(),
586                            self.headers.clone(),
587                            cookie_helper,
588                            progress_updater,
589                            Arc::clone(&self.group),
590                        );
591                        return sequential_downloader.execute_with_gaps_with_retry(
592                            &uri,
593                            total_length,
594                            &completed_ranges,
595                            &retry_policy,
596                        ).await;
597                    }
598                    Err(e) => return Err(e),
599                }
600            }
601
602            let retry_policy = RetryPolicy::new(options.max_retries, options.retry_wait * 1000);
603            let mut sequential_downloader = SequentialDownloader::new(
604                Arc::clone(&self.client),
605                self.output_path.clone(),
606                self.headers.clone(),
607                cookie_helper,
608                progress_updater,
609                Arc::clone(&self.group),
610            );
611            sequential_downloader.execute_with_retry(
612                &uri,
613                &resume_state,
614                total_length,
615                &retry_policy,
616            ).await
617        }
618        .await;
619
620        self.drain_progress_aggregator().await;
621
622        if download_result.is_ok() {
623            // Verify checksum if configured
624            {
625                let g = self.group.read().await;
626                if let Some((ref algo, ref expected)) = g.options().checksum
627                    && let Some(ht) = HashType::from_str(algo)
628                {
629                    let cs = Checksum::new(ht, expected)?;
630                    let file = tokio::fs::File::open(&self.output_path)
631                        .await
632                        .map_err(|e| {
633                            Aria2Error::Io(format!(
634                                "Failed to open file for checksum verification: {}",
635                                e
636                            ))
637                        })?;
638                    let mut reader = tokio::io::BufReader::with_capacity(65536, file);
639                    let mut validator = cs.create_validator();
640                    let mut buf = vec![0u8; 65536];
641                    loop {
642                        let n = reader.read(&mut buf).await.map_err(|e| {
643                            Aria2Error::Io(format!(
644                                "Read error during checksum verification: {}",
645                                e
646                            ))
647                        })?;
648                        if n == 0 {
649                            break;
650                        }
651                        validator.update(&buf[..n]);
652                    }
653                    if !validator.finalize()? {
654                        tracing::error!(
655                            algo = %algo,
656                            path = %self.output_path.display(),
657                            "Checksum mismatch"
658                        );
659                        return Err(Aria2Error::Checksum(format!(
660                            "{} checksum mismatch for {}",
661                            algo,
662                            self.output_path.display()
663                        )));
664                    }
665                    tracing::info!(
666                        algo = %algo,
667                        path = %self.output_path.display(),
668                        "Checksum verified successfully"
669                    );
670                }
671            }
672            self.completed = true;
673            let g = self.group.write().await;
674            let total = g.total_length();
675            g.update_progress(total).await;
676            g.set_completed_length(total);
677        }
678        release_path(&self.output_path);
679        download_result
680    }
681
682    fn status(&self) -> CommandStatus {
683        if self.completed {
684            CommandStatus::Completed
685        } else if self.completed_bytes > 0 {
686            CommandStatus::Running
687        } else {
688            CommandStatus::Pending
689        }
690    }
691
692    fn timeout(&self) -> Option<Duration> {
693        Some(Duration::from_secs(
694            constants::HTTP_DEFAULT_COMMAND_TIMEOUT_SECS,
695        ))
696    }
697}
698
699#[cfg(test)]
700mod tests {
701    use super::*;
702    use crate::request::request_group::{DownloadOptions, GroupId, RequestGroup};
703
704    impl DownloadCommand {
705        fn has_progress_sender(&self) -> bool {
706            self.progress_sender.is_some()
707        }
708
709        fn has_progress_receiver(&self) -> bool {
710            self.progress_receiver.is_some()
711        }
712
713        fn has_progress_aggregator_handle(&self) -> bool {
714            self.progress_aggregator_handle.is_some()
715        }
716
717        fn send_progress_update(&self, update: ProgressUpdate) {
718            if let Some(ref sender) = self.progress_sender {
719                let _ = sender.send(update);
720            } else {
721                panic!("test called send_progress_update but no sender is set");
722            }
723        }
724    }
725
726    #[test]
727    fn test_progress_channel_auto_created() {
728        let cmd = DownloadCommand::new(
729            GroupId::new(1),
730            "http://example.com/file.bin",
731            &DownloadOptions::default(),
732            None,
733            None,
734        )
735        .expect("DownloadCommand::new should succeed with a valid HTTP URI");
736
737        assert!(
738            cmd.has_progress_sender(),
739            "progress_sender should be Some after construction (auto-created)"
740        );
741        assert!(
742            cmd.has_progress_receiver(),
743            "progress_receiver should be Some after construction (held for lazy spawn)"
744        );
745        assert!(
746            !cmd.has_progress_aggregator_handle(),
747            "progress_aggregator_handle should be None until execute() spawns it"
748        );
749    }
750
751    #[tokio::test]
752    async fn test_progress_updates_flow_through_channel() {
753        let group = Arc::new(tokio::sync::RwLock::new(RequestGroup::new(
754            GroupId::new(2),
755            vec!["http://example.com/file.bin".to_string()],
756            DownloadOptions::default(),
757        )));
758        let group_clone = Arc::clone(&group);
759
760        let mut cmd = DownloadCommand::new_with_group(
761            group,
762            "http://example.com/file.bin",
763            &DownloadOptions::default(),
764            None,
765            None,
766        )
767        .expect("DownloadCommand::new_with_group should succeed");
768
769        assert!(cmd.has_progress_sender());
770        assert!(cmd.has_progress_receiver());
771
772        cmd.spawn_progress_aggregator();
773        assert!(cmd.has_progress_aggregator_handle());
774        assert!(!cmd.has_progress_receiver());
775
776        cmd.send_progress_update(ProgressUpdate {
777            completed_bytes: 4096,
778            download_speed: 0,
779            upload_speed: 0,
780        });
781
782        cmd.drain_progress_aggregator().await;
783        assert!(!cmd.has_progress_sender());
784        assert!(!cmd.has_progress_aggregator_handle());
785
786        let completed = { group_clone.read().await.get_completed_length() };
787        assert_eq!(
788            completed, 4096,
789            "aggregator should have applied the progress update to RequestGroup"
790        );
791    }
792}