1use super::connection::RetryPolicy;
8use super::resume::{
9 should_restart_without_ranges, validate_ranged_response, RangedResponseContext,
10};
11use super::ACCEPT_ENCODING_IDENTITY;
12use crate::error::{EngineError, NetworkErrorKind, ProtocolErrorKind, Result, StorageErrorKind};
13use crate::storage::Segment;
14use crate::types::DownloadProgress;
15
16use bytes::Bytes;
17use futures::stream::StreamExt;
18use parking_lot::RwLock;
19use reqwest::Client;
20use std::path::PathBuf;
21use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
22use std::sync::Arc;
23use std::time::{Duration, Instant};
24use tokio::fs::{File, OpenOptions};
25use tokio::io::{AsyncSeekExt, AsyncWriteExt, SeekFrom};
26use tokio::sync::Semaphore;
27use tokio_util::sync::CancellationToken;
28
29pub const MIN_SEGMENT_SIZE: u64 = 1024 * 1024;
31
32pub const DEFAULT_CONNECTIONS: usize = 16;
34
35const PROGRESS_INTERVAL: Duration = Duration::from_millis(250);
37
38const PERSISTENCE_INTERVAL: Duration = Duration::from_secs(5);
40
41fn log_progress_invariant(context: &str, progress: &DownloadProgress) {
42 if let Some(total_size) = progress.total_size {
43 if progress.completed_size > total_size {
44 debug_assert!(
45 progress.completed_size <= total_size,
46 "{} progress exceeded total size: {} > {}",
47 context,
48 progress.completed_size,
49 total_size
50 );
51 tracing::warn!(
52 "{} progress exceeded total size: {} > {}",
53 context,
54 progress.completed_size,
55 total_size
56 );
57 }
58 }
59}
60
61struct SharedState {
63 downloaded: AtomicU64,
65 speed: AtomicU64,
67 active_connections: AtomicU64,
69 paused: AtomicBool,
71 segment_progress: RwLock<Vec<u64>>,
73 last_persistence: RwLock<Instant>,
75}
76
77pub struct SegmentedDownload {
79 url: String,
81 total_size: u64,
83 save_path: PathBuf,
85 segments: Vec<Segment>,
87 #[allow(dead_code)]
89 supports_range: bool,
90 etag: Option<String>,
92 last_modified: Option<String>,
94 state: Arc<SharedState>,
96}
97
98#[derive(Debug, Clone)]
100pub struct ServerCapabilities {
101 pub content_length: Option<u64>,
103 pub supports_range: bool,
105 pub etag: Option<String>,
107 pub last_modified: Option<String>,
109 pub suggested_filename: Option<String>,
111}
112
113impl SegmentedDownload {
114 pub fn new(
116 url: String,
117 total_size: u64,
118 save_path: PathBuf,
119 supports_range: bool,
120 etag: Option<String>,
121 last_modified: Option<String>,
122 ) -> Self {
123 Self {
124 url,
125 total_size,
126 save_path,
127 segments: Vec::new(),
128 supports_range,
129 etag,
130 last_modified,
131 state: Arc::new(SharedState {
132 downloaded: AtomicU64::new(0),
133 speed: AtomicU64::new(0),
134 active_connections: AtomicU64::new(0),
135 paused: AtomicBool::new(false),
136 segment_progress: RwLock::new(Vec::new()),
137 last_persistence: RwLock::new(Instant::now()),
138 }),
139 }
140 }
141
142 pub fn init_segments(&mut self, max_connections: usize, min_segment_size: u64) {
144 let num_segments =
145 calculate_segment_count(self.total_size, max_connections, min_segment_size);
146 let segment_size = self.total_size / num_segments as u64;
147
148 let mut segments = Vec::with_capacity(num_segments);
149 for i in 0..num_segments {
150 let start = i as u64 * segment_size;
151 let end = if i == num_segments - 1 {
152 self.total_size - 1
153 } else {
154 (i as u64 + 1) * segment_size - 1
155 };
156 segments.push(Segment::new(i, start, end));
157 }
158
159 *self.state.segment_progress.write() = vec![0u64; num_segments];
161
162 self.segments = segments;
163 }
164
165 pub fn restore_segments(&mut self, saved_segments: Vec<Segment>) {
167 let downloaded: u64 = saved_segments.iter().map(|s| s.downloaded).sum();
169 self.state.downloaded.store(downloaded, Ordering::Relaxed);
170
171 let progress: Vec<u64> = saved_segments.iter().map(|s| s.downloaded).collect();
173 *self.state.segment_progress.write() = progress;
174
175 self.segments = saved_segments;
176 }
177
178 pub fn segments(&self) -> &[Segment] {
180 &self.segments
181 }
182
183 pub fn segments_with_progress(&self) -> Vec<Segment> {
187 let progress = self.state.segment_progress.read();
188 self.segments
189 .iter()
190 .enumerate()
191 .map(|(idx, s)| {
192 let mut segment = s.clone();
193 if let Some(&downloaded) = progress.get(idx) {
194 segment.downloaded = downloaded;
195 if segment.downloaded >= segment.size() {
196 segment.state = crate::storage::SegmentState::Completed;
197 } else if segment.downloaded > 0 {
198 segment.state = crate::storage::SegmentState::Downloading;
199 }
200 }
201 segment
202 })
203 .collect()
204 }
205
206 #[allow(clippy::too_many_arguments)]
208 pub async fn start<F>(
209 &self,
210 client: &Client,
211 user_agent: &str,
212 headers: &[(String, String)],
213 max_connections: usize,
214 retry_policy: &RetryPolicy,
215 cancel_token: CancellationToken,
216 progress_callback: F,
217 ) -> Result<()>
218 where
219 F: Fn(DownloadProgress) + Send + Sync + 'static,
220 {
221 self.start_with_scope(
222 client,
223 user_agent,
224 headers,
225 max_connections,
226 retry_policy,
227 #[cfg(feature = "recursive-http")]
228 None,
229 cancel_token,
230 progress_callback,
231 )
232 .await
233 }
234
235 #[allow(clippy::too_many_arguments)]
236 pub(crate) async fn start_with_scope<F>(
237 &self,
238 client: &Client,
239 user_agent: &str,
240 headers: &[(String, String)],
241 max_connections: usize,
242 retry_policy: &RetryPolicy,
243 #[cfg(feature = "recursive-http")] redirect_scope: Option<super::crawl::RedirectScope>,
244 cancel_token: CancellationToken,
245 progress_callback: F,
246 ) -> Result<()>
247 where
248 F: Fn(DownloadProgress) + Send + Sync + 'static,
249 {
250 let file = self.prepare_file().await?;
252 let file = Arc::new(tokio::sync::Mutex::new(file));
253
254 let semaphore = Arc::new(Semaphore::new(max_connections));
256
257 let fatal_cancel = cancel_token.child_token();
260
261 let progress_callback = Arc::new(progress_callback);
263 let last_progress = Arc::new(RwLock::new(Instant::now()));
264 let bytes_since_progress = Arc::new(AtomicU64::new(0));
265
266 let segments_data: Vec<_> = self
268 .segments
269 .iter()
270 .enumerate()
271 .filter(|(_, s)| !s.is_complete())
272 .map(|(idx, s)| (idx, s.start, s.end, s.downloaded))
273 .collect();
274
275 let mut handles = Vec::new();
277
278 for (segment_idx, start, end, already_downloaded) in segments_data {
279 let client = client.clone();
280 let url = self.url.clone();
281 let user_agent = user_agent.to_string();
282 let headers = headers.to_vec();
283 let file = Arc::clone(&file);
284 let semaphore = Arc::clone(&semaphore);
285 let cancel_token = fatal_cancel.clone();
286 let etag = self.etag.clone();
287 let last_modified = self.last_modified.clone();
288 let state = Arc::clone(&self.state);
289 let progress_callback = Arc::clone(&progress_callback);
290 let last_progress = Arc::clone(&last_progress);
291 let bytes_since_progress = Arc::clone(&bytes_since_progress);
292 let total_size = self.total_size;
293 let retry_policy = retry_policy.clone();
294 #[cfg(feature = "recursive-http")]
295 let redirect_scope = redirect_scope.clone();
296
297 let handle = tokio::spawn(async move {
298 let _permit = semaphore
300 .acquire()
301 .await
302 .map_err(|_| EngineError::Shutdown)?;
303
304 if cancel_token.is_cancelled() {
306 return Ok(());
307 }
308
309 if state.paused.load(Ordering::Relaxed) {
311 return Ok(());
312 }
313
314 state.active_connections.fetch_add(1, Ordering::Relaxed);
315
316 let mut segment_bytes: u64 = already_downloaded;
318 let expected_segment_size = end - start + 1;
319 let mut last_speed_update = Instant::now();
320 let mut bytes_for_speed: u64 = 0;
321 let mut attempt = 0u32;
322
323 if start + segment_bytes > end {
325 state.active_connections.fetch_sub(1, Ordering::Relaxed);
326 return Ok(());
327 }
328
329 let result: Result<()> = 'retry: loop {
330 if cancel_token.is_cancelled() {
332 break 'retry Ok(());
333 }
334
335 let resume_start = start + segment_bytes;
337 if resume_start > end {
338 break 'retry Ok(());
339 }
340
341 let mut request = client.get(&url);
343 request = request.header("User-Agent", &user_agent);
344 request = request.header("Range", format!("bytes={}-{}", resume_start, end));
345
346 if let Some(if_range_val) = etag.as_deref().or(last_modified.as_deref()) {
348 request = request.header("If-Range", if_range_val);
349 }
350
351 for (name, value) in &headers {
353 request = request.header(name.as_str(), value.as_str());
354 }
355 request = request.header("Accept-Encoding", ACCEPT_ENCODING_IDENTITY);
356
357 let response = match request.send().await {
359 Ok(r) => r,
360 Err(e) => {
361 let err: EngineError = e.into();
362 attempt += 1;
363 if retry_policy.should_retry(attempt - 1, &err) {
364 tracing::warn!(
365 "Segment {} request failed (attempt {}/{}), retrying: {}",
366 segment_idx,
367 attempt,
368 retry_policy.max_attempts,
369 err
370 );
371 let delay = retry_policy.delay_for_attempt(attempt - 1);
372 tokio::time::sleep(delay).await;
373 continue 'retry;
374 }
375 if !err.is_retryable() {
376 cancel_token.cancel();
377 }
378 break 'retry Err(err);
379 }
380 };
381 #[cfg(feature = "recursive-http")]
382 if let Some(scope) = redirect_scope.as_ref() {
383 super::crawl::validate_redirect_scope(response.url(), scope)?;
384 }
385
386 let status = response.status();
387
388 if status == reqwest::StatusCode::RANGE_NOT_SATISFIABLE {
390 cancel_token.cancel();
391 break 'retry Err(EngineError::network(
392 NetworkErrorKind::HttpStatus(416),
393 format!(
394 "Segment {} range not satisfiable (file may have changed on server)",
395 segment_idx
396 ),
397 ));
398 }
399
400 if status.is_server_error() {
402 let err = EngineError::network(
403 NetworkErrorKind::HttpStatus(status.as_u16()),
404 format!("Segment {} server error: {}", segment_idx, status),
405 );
406 attempt += 1;
407 if retry_policy.should_retry(attempt - 1, &err) {
408 tracing::warn!(
409 "Segment {} server error (attempt {}/{}), retrying: {}",
410 segment_idx,
411 attempt,
412 retry_policy.max_attempts,
413 status
414 );
415 let delay = retry_policy.delay_for_attempt(attempt - 1);
416 tokio::time::sleep(delay).await;
417 continue 'retry;
418 }
419 break 'retry Err(err);
420 }
421
422 if !status.is_success() && status != reqwest::StatusCode::PARTIAL_CONTENT {
423 cancel_token.cancel();
424 break 'retry Err(EngineError::network(
425 NetworkErrorKind::HttpStatus(status.as_u16()),
426 format!("Segment {} HTTP error: {}", segment_idx, status),
427 ));
428 }
429
430 if let Err(e) = validate_ranged_response(
431 resume_start,
432 Some(end),
433 status,
434 response
435 .headers()
436 .get("content-range")
437 .and_then(|v| v.to_str().ok()),
438 RangedResponseContext {
439 sent_if_range: etag.is_some() || last_modified.is_some(),
440 expected_etag: etag.as_deref(),
441 expected_last_modified: last_modified.as_deref(),
442 response_etag: response
443 .headers()
444 .get("etag")
445 .and_then(|v| v.to_str().ok()),
446 response_last_modified: response
447 .headers()
448 .get("last-modified")
449 .and_then(|v| v.to_str().ok()),
450 },
451 ) {
452 break 'retry Err(e);
453 }
454
455 let mut stream = response.bytes_stream();
457 let mut stream_failed = false;
458
459 while let Some(chunk_result) = tokio::select! {
460 chunk = stream.next() => chunk,
461 _ = cancel_token.cancelled() => None,
462 } {
463 if state.paused.load(Ordering::Relaxed) {
465 break;
466 }
467
468 let chunk: Bytes = match chunk_result {
469 Ok(c) => c,
470 Err(e) => {
471 let err: EngineError = e.into();
472 attempt += 1;
473 if retry_policy.should_retry(attempt - 1, &err) {
474 tracing::warn!(
475 "Segment {} stream error (attempt {}/{}), retrying from byte {}: {}",
476 segment_idx, attempt, retry_policy.max_attempts, segment_bytes, err
477 );
478 stream_failed = true;
479 break;
480 }
481 if !err.is_retryable() {
482 cancel_token.cancel();
483 }
484 break 'retry Err(err);
485 }
486 };
487
488 let chunk_len = chunk.len() as u64;
489
490 {
492 let mut file = file.lock().await;
493 file.seek(SeekFrom::Start(start + segment_bytes))
494 .await
495 .map_err(|e| {
496 EngineError::storage(
497 StorageErrorKind::Io,
498 PathBuf::new(),
499 format!("Seek failed: {}", e),
500 )
501 })?;
502 file.write_all(&chunk).await.map_err(|e| {
503 EngineError::storage(
504 StorageErrorKind::Io,
505 PathBuf::new(),
506 format!("Write failed: {}", e),
507 )
508 })?;
509 }
510
511 segment_bytes += chunk_len;
512 if segment_bytes > expected_segment_size {
513 break 'retry Err(EngineError::protocol(
514 ProtocolErrorKind::InvalidResponse,
515 format!(
516 "Segment {} exceeded expected size: received {} bytes, expected {} bytes",
517 segment_idx, segment_bytes, expected_segment_size
518 ),
519 ));
520 }
521
522 {
524 let mut progress = state.segment_progress.write();
525 if let Some(p) = progress.get_mut(segment_idx) {
526 *p = segment_bytes;
527 }
528 }
529
530 let total_downloaded =
532 state.downloaded.fetch_add(chunk_len, Ordering::Relaxed) + chunk_len;
533 if total_downloaded > total_size {
534 break 'retry Err(EngineError::protocol(
535 ProtocolErrorKind::InvalidResponse,
536 format!(
537 "Download exceeded expected size: received {} bytes, expected {} bytes",
538 total_downloaded, total_size
539 ),
540 ));
541 }
542 bytes_since_progress.fetch_add(chunk_len, Ordering::Relaxed);
543 bytes_for_speed += chunk_len;
544
545 let now = Instant::now();
547 let speed_elapsed = now.duration_since(last_speed_update);
548 if speed_elapsed >= Duration::from_millis(500) {
549 let current_speed =
550 (bytes_for_speed as f64 / speed_elapsed.as_secs_f64()) as u64;
551 state.speed.store(current_speed, Ordering::Relaxed);
552 bytes_for_speed = 0;
553 last_speed_update = now;
554 }
555
556 let should_emit = {
558 let mut last = last_progress.write();
559 if now.duration_since(*last) >= PROGRESS_INTERVAL {
560 *last = now;
561 bytes_since_progress.store(0, Ordering::Relaxed);
562 true
563 } else {
564 false
565 }
566 };
567
568 if should_emit {
569 let current_speed = state.speed.load(Ordering::Relaxed);
570 let connections =
571 state.active_connections.load(Ordering::Relaxed) as u32;
572
573 let progress = DownloadProgress {
574 total_size: Some(total_size),
575 completed_size: total_downloaded,
576 download_speed: current_speed,
577 upload_speed: 0,
578 connections,
579 seeders: 0,
580 peers: 0,
581 eta_seconds: total_size
582 .saturating_sub(total_downloaded)
583 .checked_div(current_speed),
584 };
585 log_progress_invariant("segmented http download", &progress);
586 progress_callback(progress);
587 }
588 }
589
590 if stream_failed {
591 let delay = retry_policy.delay_for_attempt(attempt - 1);
592 tokio::time::sleep(delay).await;
593 continue 'retry;
594 }
595
596 break 'retry Ok(());
598 };
599
600 state.active_connections.fetch_sub(1, Ordering::Relaxed);
601
602 result
604 });
605
606 handles.push(handle);
607 }
608
609 let mut segment_errors: Vec<String> = Vec::new();
611 let mut any_retryable = false;
612 let mut restart_without_ranges_reason: Option<String> = None;
613 for (idx, handle) in handles.into_iter().enumerate() {
614 match handle.await {
615 Err(e) => {
616 tracing::error!("Segment {} task panicked: {:?}", idx, e);
618 segment_errors.push(format!("Segment {} panicked: {:?}", idx, e));
619 }
620 Ok(Err(e)) => {
621 tracing::error!("Segment {} failed: {:?}", idx, e);
623 if e.is_retryable() {
624 any_retryable = true;
625 }
626 if restart_without_ranges_reason.is_none() && should_restart_without_ranges(&e)
627 {
628 restart_without_ranges_reason = Some(e.to_string());
629 }
630 segment_errors.push(format!("Segment {} failed: {}", idx, e));
631 }
632 Ok(Ok(())) => {
633 }
635 }
636 }
637
638 if !segment_errors.is_empty() {
640 if let Some(reason) = restart_without_ranges_reason {
641 return Err(EngineError::protocol(
642 ProtocolErrorKind::RangeNotSupported,
643 format!(
644 "Segmented download requires restart without ranges: {}",
645 reason
646 ),
647 ));
648 }
649 let kind = if any_retryable {
652 NetworkErrorKind::ConnectionReset
653 } else {
654 NetworkErrorKind::Other
655 };
656 return Err(EngineError::network(
657 kind,
658 format!(
659 "Download failed: {} segment(s) failed: {}",
660 segment_errors.len(),
661 segment_errors.join("; ")
662 ),
663 ));
664 }
665
666 {
668 let mut file = file.lock().await;
669 file.flush().await.map_err(|e| {
670 EngineError::storage(
671 StorageErrorKind::Io,
672 &self.save_path,
673 format!("Flush failed: {}", e),
674 )
675 })?;
676 file.sync_all().await.map_err(|e| {
677 EngineError::storage(
678 StorageErrorKind::Io,
679 &self.save_path,
680 format!("Sync failed: {}", e),
681 )
682 })?;
683 }
684
685 let total_downloaded = self.state.downloaded.load(Ordering::Relaxed);
687 if total_downloaded != self.total_size {
688 return Err(EngineError::protocol(
689 ProtocolErrorKind::InvalidResponse,
690 format!(
691 "Segmented download size mismatch: received {} bytes, expected {} bytes",
692 total_downloaded, self.total_size
693 ),
694 ));
695 }
696 let progress = DownloadProgress {
697 total_size: Some(self.total_size),
698 completed_size: total_downloaded,
699 download_speed: 0,
700 upload_speed: 0,
701 connections: 0,
702 seeders: 0,
703 peers: 0,
704 eta_seconds: None,
705 };
706 log_progress_invariant("segmented http download", &progress);
707 progress_callback(progress);
708
709 if total_downloaded >= self.total_size {
711 self.finalize().await?;
713 }
714
715 Ok(())
716 }
717
718 pub fn should_persist(&self) -> bool {
723 let mut last = self.state.last_persistence.write();
724 let now = Instant::now();
725 if now.duration_since(*last) >= PERSISTENCE_INTERVAL {
726 *last = now;
727 true
728 } else {
729 false
730 }
731 }
732
733 pub fn mark_persisted(&self) {
735 *self.state.last_persistence.write() = Instant::now();
736 }
737
738 async fn prepare_file(&self) -> Result<File> {
740 let part_path = self.part_path();
742
743 if let Some(parent) = part_path.parent() {
745 tokio::fs::create_dir_all(parent).await.map_err(|e| {
746 EngineError::storage(
747 StorageErrorKind::Io,
748 parent,
749 format!("Create dir failed: {}", e),
750 )
751 })?;
752 }
753
754 let file = if part_path.exists() {
756 OpenOptions::new()
757 .write(true)
758 .read(true)
759 .open(&part_path)
760 .await
761 .map_err(|e| {
762 EngineError::storage(
763 StorageErrorKind::Io,
764 &part_path,
765 format!("Open failed: {}", e),
766 )
767 })?
768 } else {
769 let file = File::create(&part_path).await.map_err(|e| {
771 EngineError::storage(
772 StorageErrorKind::Io,
773 &part_path,
774 format!("Create failed: {}", e),
775 )
776 })?;
777
778 file.set_len(self.total_size).await.map_err(|e| {
780 EngineError::storage(
781 StorageErrorKind::Io,
782 &part_path,
783 format!("Pre-allocate failed: {}", e),
784 )
785 })?;
786
787 file
788 };
789
790 Ok(file)
791 }
792
793 fn part_path(&self) -> PathBuf {
795 let ext = self
796 .save_path
797 .extension()
798 .map(|e| format!("{}.part", e.to_string_lossy()))
799 .unwrap_or_else(|| "part".to_string());
800 self.save_path.with_extension(ext)
801 }
802
803 async fn finalize(&self) -> Result<()> {
805 let part_path = self.part_path();
806 if part_path.exists() {
807 tokio::fs::rename(&part_path, &self.save_path)
808 .await
809 .map_err(|e| {
810 EngineError::storage(
811 StorageErrorKind::Io,
812 &self.save_path,
813 format!("Rename failed: {}", e),
814 )
815 })?;
816 }
817 Ok(())
818 }
819
820 pub fn pause(&self) {
822 self.state.paused.store(true, Ordering::Relaxed);
823 }
824
825 pub fn is_complete(&self) -> bool {
827 self.state.downloaded.load(Ordering::Relaxed) >= self.total_size
828 }
829
830 pub fn progress(&self) -> DownloadProgress {
832 let progress = DownloadProgress {
833 total_size: Some(self.total_size),
834 completed_size: self.state.downloaded.load(Ordering::Relaxed),
835 download_speed: self.state.speed.load(Ordering::Relaxed),
836 upload_speed: 0,
837 connections: self.state.active_connections.load(Ordering::Relaxed) as u32,
838 seeders: 0,
839 peers: 0,
840 eta_seconds: {
841 let speed = self.state.speed.load(Ordering::Relaxed);
842 let remaining = self
843 .total_size
844 .saturating_sub(self.state.downloaded.load(Ordering::Relaxed));
845 remaining.checked_div(speed)
846 },
847 };
848 log_progress_invariant("segmented http download", &progress);
849 progress
850 }
851}
852
853pub fn calculate_segment_count(
855 total_size: u64,
856 max_connections: usize,
857 min_segment_size: u64,
858) -> usize {
859 if total_size == 0 {
860 return 1;
861 }
862
863 let max_segments_by_size = (total_size / min_segment_size) as usize;
865
866 let num_segments = max_connections.min(max_segments_by_size.max(1));
868
869 num_segments.max(1)
871}
872
873pub async fn probe_server(
875 client: &Client,
876 url: &str,
877 user_agent: &str,
878) -> Result<ServerCapabilities> {
879 let response = client
880 .head(url)
881 .header("User-Agent", user_agent)
882 .header("Accept-Encoding", ACCEPT_ENCODING_IDENTITY)
883 .send()
884 .await
885 .map_err(EngineError::from)?;
886
887 if !response.status().is_success() {
888 return Err(EngineError::network(
889 NetworkErrorKind::HttpStatus(response.status().as_u16()),
890 format!("HEAD request returned: {}", response.status()),
891 ));
892 }
893
894 let headers = response.headers();
895
896 let content_length = headers
897 .get("content-length")
898 .and_then(|v| v.to_str().ok())
899 .and_then(|s| s.parse::<u64>().ok());
900
901 let supports_range = headers
902 .get("accept-ranges")
903 .and_then(|v| v.to_str().ok())
904 .map(|v| v.contains("bytes"))
905 .unwrap_or(false);
906
907 let etag = headers
908 .get("etag")
909 .and_then(|v| v.to_str().ok())
910 .map(|s| s.to_string());
911
912 let last_modified = headers
913 .get("last-modified")
914 .and_then(|v| v.to_str().ok())
915 .map(|s| s.to_string());
916
917 let suggested_filename = headers
918 .get("content-disposition")
919 .and_then(|v| v.to_str().ok())
920 .and_then(parse_content_disposition);
921
922 Ok(ServerCapabilities {
923 content_length,
924 supports_range,
925 etag,
926 last_modified,
927 suggested_filename,
928 })
929}
930
931fn parse_content_disposition(header: &str) -> Option<String> {
933 if let Some(start) = header.find("filename=") {
935 let rest = &header[start + 9..];
936 if let Some(stripped) = rest.strip_prefix('"') {
937 let end = stripped.find('"')?;
938 return Some(stripped[..end].to_string());
939 } else {
940 let end = rest.find(';').unwrap_or(rest.len());
941 return Some(rest[..end].trim().to_string());
942 }
943 }
944
945 if let Some(start) = header.find("filename*=") {
946 let rest = &header[start + 10..];
947 if let Some(quote_start) = rest.find("''") {
948 let encoded = &rest[quote_start + 2..];
949 let end = encoded.find(';').unwrap_or(encoded.len());
950 if let Ok(decoded) = urlencoding::decode(&encoded[..end]) {
951 return Some(decoded.to_string());
952 }
953 }
954 }
955
956 None
957}
958
959#[cfg(test)]
960mod tests {
961 use super::*;
962
963 #[test]
964 fn test_calculate_segment_count() {
965 assert_eq!(
967 calculate_segment_count(100 * 1024 * 1024, 16, 1024 * 1024),
968 16
969 );
970
971 assert_eq!(
973 calculate_segment_count(10 * 1024 * 1024, 16, 1024 * 1024),
974 10
975 );
976
977 assert_eq!(calculate_segment_count(512 * 1024, 16, 1024 * 1024), 1);
979
980 assert_eq!(calculate_segment_count(0, 16, 1024 * 1024), 1);
982
983 assert_eq!(
985 calculate_segment_count(10 * 1024 * 1024 * 1024, 16, 1024 * 1024),
986 16
987 );
988 }
989
990 #[test]
991 fn test_segment_init() {
992 let mut download = SegmentedDownload::new(
993 "https://example.com/file.zip".to_string(),
994 100 * 1024 * 1024, PathBuf::from("/tmp/file.zip"),
996 true,
997 None,
998 None,
999 );
1000
1001 download.init_segments(16, 1024 * 1024);
1002
1003 let segments = download.segments();
1004 assert_eq!(segments.len(), 16);
1005
1006 assert_eq!(segments[0].start, 0);
1008 assert_eq!(segments[15].end, 100 * 1024 * 1024 - 1);
1009
1010 for i in 0..15 {
1012 assert_eq!(segments[i].end + 1, segments[i + 1].start);
1013 }
1014 }
1015
1016 #[test]
1017 fn test_parse_content_disposition() {
1018 assert_eq!(
1019 parse_content_disposition("attachment; filename=\"test.zip\""),
1020 Some("test.zip".to_string())
1021 );
1022
1023 assert_eq!(
1024 parse_content_disposition("attachment; filename=test.zip"),
1025 Some("test.zip".to_string())
1026 );
1027
1028 assert_eq!(
1029 parse_content_disposition("attachment; filename*=UTF-8''test%20file.zip"),
1030 Some("test file.zip".to_string())
1031 );
1032 }
1033}