Skip to main content

gosh_dl/http/
segment.rs

1//! Segmented Download Support
2//!
3//! This module provides multi-connection segmented downloads for faster
4//! HTTP/HTTPS transfers. It splits files into segments and downloads
5//! them in parallel using multiple connections.
6
7use 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
29/// Minimum segment size (1 MiB)
30pub const MIN_SEGMENT_SIZE: u64 = 1024 * 1024;
31
32/// Default number of connections per download
33pub const DEFAULT_CONNECTIONS: usize = 16;
34
35/// Progress update interval
36const PROGRESS_INTERVAL: Duration = Duration::from_millis(250);
37
38/// Persistence interval for segment state
39const 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
61/// Shared state for a segmented download
62struct SharedState {
63    /// Total bytes downloaded across all segments
64    downloaded: AtomicU64,
65    /// Current download speed (bytes/sec)
66    speed: AtomicU64,
67    /// Number of active connections
68    active_connections: AtomicU64,
69    /// Whether download is paused
70    paused: AtomicBool,
71    /// Per-segment downloaded bytes (for tracking progress)
72    segment_progress: RwLock<Vec<u64>>,
73    /// Last persistence time
74    last_persistence: RwLock<Instant>,
75}
76
77/// Segmented download manager
78pub struct SegmentedDownload {
79    /// URL to download from
80    url: String,
81    /// Total file size
82    total_size: u64,
83    /// Path to save the file
84    save_path: PathBuf,
85    /// Segments
86    segments: Vec<Segment>,
87    /// Whether server supports range requests (stored for resume validation)
88    #[allow(dead_code)]
89    supports_range: bool,
90    /// ETag for validation
91    etag: Option<String>,
92    /// Last-Modified for validation
93    last_modified: Option<String>,
94    /// Shared state (wrapped in Arc for task sharing)
95    state: Arc<SharedState>,
96}
97
98/// Server capabilities determined from HEAD request
99#[derive(Debug, Clone)]
100pub struct ServerCapabilities {
101    /// Content-Length header value
102    pub content_length: Option<u64>,
103    /// Whether server supports Range requests
104    pub supports_range: bool,
105    /// ETag header for validation
106    pub etag: Option<String>,
107    /// Last-Modified header for validation
108    pub last_modified: Option<String>,
109    /// Suggested filename from Content-Disposition
110    pub suggested_filename: Option<String>,
111}
112
113impl SegmentedDownload {
114    /// Create a new segmented download
115    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    /// Initialize segments for a new download
143    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        // Initialize segment progress tracking
160        *self.state.segment_progress.write() = vec![0u64; num_segments];
161
162        self.segments = segments;
163    }
164
165    /// Restore segments from saved state
166    pub fn restore_segments(&mut self, saved_segments: Vec<Segment>) {
167        // Calculate total already downloaded
168        let downloaded: u64 = saved_segments.iter().map(|s| s.downloaded).sum();
169        self.state.downloaded.store(downloaded, Ordering::Relaxed);
170
171        // Initialize segment progress tracking with saved values
172        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    /// Get current segments
179    pub fn segments(&self) -> &[Segment] {
180        &self.segments
181    }
182
183    /// Get segments with current progress updated
184    ///
185    /// This creates a snapshot of the current segment state for persistence.
186    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    /// Start the segmented download
207    #[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        // Create/open the file and pre-allocate space
251        let file = self.prepare_file().await?;
252        let file = Arc::new(tokio::sync::Mutex::new(file));
253
254        // Create semaphore for connection limiting
255        let semaphore = Arc::new(Semaphore::new(max_connections));
256
257        // Child cancel token: cancelled on fatal (non-retryable) segment errors
258        // so sibling segments stop promptly instead of wasting bandwidth
259        let fatal_cancel = cancel_token.child_token();
260
261        // Shared state for progress tracking
262        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        // Clone segments data for tasks
267        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        // Spawn tasks for each pending segment
276        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                // Acquire permit
299                let _permit = semaphore
300                    .acquire()
301                    .await
302                    .map_err(|_| EngineError::Shutdown)?;
303
304                // Check cancellation
305                if cancel_token.is_cancelled() {
306                    return Ok(());
307                }
308
309                // Check if paused
310                if state.paused.load(Ordering::Relaxed) {
311                    return Ok(());
312                }
313
314                state.active_connections.fetch_add(1, Ordering::Relaxed);
315
316                // Persistent state across retries
317                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                // Check if already complete before entering retry loop
324                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                    // Check cancellation between retries
331                    if cancel_token.is_cancelled() {
332                        break 'retry Ok(());
333                    }
334
335                    // Calculate resume position from current progress
336                    let resume_start = start + segment_bytes;
337                    if resume_start > end {
338                        break 'retry Ok(());
339                    }
340
341                    // Build request with Range header
342                    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                    // Add ETag for validation if available
347                    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                    // Add custom headers
352                    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                    // Send request
358                    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                    // Handle 416 Range Not Satisfiable — not retryable
389                    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                    // Handle server errors (5xx) with retry
401                    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                    // Stream data to file
456                    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                        // Check pause
464                        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                        // Write to file at correct offset
491                        {
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                        // Update segment progress for persistence
523                        {
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                        // Update global counters
531                        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                        // Update speed calculation
546                        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                        // Emit progress at intervals
557                        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                    // Stream completed successfully (or was paused/cancelled)
597                    break 'retry Ok(());
598                };
599
600                state.active_connections.fetch_sub(1, Ordering::Relaxed);
601
602                // Return the result from the retry loop
603                result
604            });
605
606            handles.push(handle);
607        }
608
609        // Wait for all segment tasks to complete and collect errors
610        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                    // Task panicked
617                    tracing::error!("Segment {} task panicked: {:?}", idx, e);
618                    segment_errors.push(format!("Segment {} panicked: {:?}", idx, e));
619                }
620                Ok(Err(e)) => {
621                    // Task returned an error
622                    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                    // Task completed successfully
634                }
635            }
636        }
637
638        // If any segments failed, return error
639        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            // Preserve retryability: if any segment had a retryable error,
650            // the aggregate should also be retryable so the engine can retry
651            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        // Sync file to disk
667        {
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        // Final progress update
686        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        // Check if complete
710        if total_downloaded >= self.total_size {
711            // Rename from .part to final name
712            self.finalize().await?;
713        }
714
715        Ok(())
716    }
717
718    /// Check if persistence is due based on the time interval.
719    ///
720    /// Returns true if enough time has passed since the last persistence,
721    /// and resets the timer if so.
722    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    /// Force mark persistence as done (call after successful save).
734    pub fn mark_persisted(&self) {
735        *self.state.last_persistence.write() = Instant::now();
736    }
737
738    /// Prepare the output file
739    async fn prepare_file(&self) -> Result<File> {
740        // Use .part extension during download
741        let part_path = self.part_path();
742
743        // Ensure parent directory exists
744        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        // Check if file exists (for resume)
755        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            // Create new file and pre-allocate
770            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            // Pre-allocate space
779            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    /// Get the .part file path
794    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    /// Rename .part file to final name
804    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    /// Pause the download
821    pub fn pause(&self) {
822        self.state.paused.store(true, Ordering::Relaxed);
823    }
824
825    /// Check if download is complete
826    pub fn is_complete(&self) -> bool {
827        self.state.downloaded.load(Ordering::Relaxed) >= self.total_size
828    }
829
830    /// Get current progress
831    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
853/// Calculate optimal number of segments based on file size and constraints
854pub 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    // Calculate maximum segments based on min_segment_size
864    let max_segments_by_size = (total_size / min_segment_size) as usize;
865
866    // Use the smaller of max_connections and max_segments_by_size
867    let num_segments = max_connections.min(max_segments_by_size.max(1));
868
869    // Ensure at least 1 segment
870    num_segments.max(1)
871}
872
873/// Probe server capabilities with a HEAD request
874pub 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
931/// Parse filename from Content-Disposition header
932fn parse_content_disposition(header: &str) -> Option<String> {
933    // Look for filename="..." or filename*=UTF-8''...
934    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        // 100MB file, 16 connections, 1MB min
966        assert_eq!(
967            calculate_segment_count(100 * 1024 * 1024, 16, 1024 * 1024),
968            16
969        );
970
971        // 10MB file, 16 connections, 1MB min -> only 10 segments
972        assert_eq!(
973            calculate_segment_count(10 * 1024 * 1024, 16, 1024 * 1024),
974            10
975        );
976
977        // 500KB file, 16 connections, 1MB min -> 1 segment
978        assert_eq!(calculate_segment_count(512 * 1024, 16, 1024 * 1024), 1);
979
980        // Empty file
981        assert_eq!(calculate_segment_count(0, 16, 1024 * 1024), 1);
982
983        // Very large file
984        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, // 100MB
995            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        // Check segment boundaries
1007        assert_eq!(segments[0].start, 0);
1008        assert_eq!(segments[15].end, 100 * 1024 * 1024 - 1);
1009
1010        // Check segments are contiguous
1011        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}