Files
goget/internal/metalink/downloader.go
T

704 lines
17 KiB
Go

//go:build linux || freebsd
// +build linux freebsd
package metalink
import (
"context"
"fmt"
"io"
"net/http"
"net/url"
"os"
"path/filepath"
"sync"
"sync/atomic"
"time"
"codeberg.org/petrbalvin/goget/internal/core"
"codeberg.org/petrbalvin/goget/internal/crypto"
format "codeberg.org/petrbalvin/goget/internal/format"
)
// MultiSourceDownloader downloads from multiple sources simultaneously
type MultiSourceDownloader struct {
config *DownloaderConfig
}
// DownloaderConfig configuration for multi-source download
type DownloaderConfig struct {
// Maximum number of parallel sources
MaxSources int
// Timeout for each source
Timeout time.Duration
// Minimum speed (bytes/s) before switching to another source
MinSpeed int64
// Buffer size for reading
BufferSize int
// Maximum bytes to accept from a single source (0 = no limit)
MaxDownloadSize int64
// Verbose mode
Verbose bool
// Callback for progress
ProgressCallback func(current, total int64, speed float64)
}
// DefaultDownloaderConfig returns default configuration
func DefaultDownloaderConfig() *DownloaderConfig {
return &DownloaderConfig{
MaxSources: 4,
Timeout: 30 * time.Second,
MinSpeed: 1024, // 1 KB/s
BufferSize: 32 * 1024,
MaxDownloadSize: 1 << 30, // 1 GiB safety cap per source
Verbose: false,
}
}
// NewMultiSourceDownloader creates new downloader
func NewMultiSourceDownloader(cfg *DownloaderConfig) *MultiSourceDownloader {
if cfg == nil {
cfg = DefaultDownloaderConfig()
}
return &MultiSourceDownloader{
config: cfg,
}
}
// DownloadResult represents the result of a multi-source download
type DownloadResult struct {
BytesDownloaded int64
TotalSize int64
SourcesUsed int
Duration time.Duration
Speed float64
OutputPath string
Verified bool
Hash string
}
// Download downloads a file from multiple sources simultaneously
func (d *MultiSourceDownloader) Download(ctx context.Context, file *File, outputPath string) (*DownloadResult, error) {
startTime := time.Now()
if len(file.URLs) == 0 {
return nil, fmt.Errorf("no URLs available for download")
}
// Prepare output file
if err := os.MkdirAll(filepath.Dir(outputPath), 0755); err != nil {
return nil, fmt.Errorf("failed to create output directory: %w", err)
}
outFile, err := os.Create(outputPath)
if err != nil {
return nil, fmt.Errorf("failed to create output file: %w", err)
}
defer outFile.Close()
// Get sorted URLs by priority
urls := file.GetURLs()
if len(urls) > d.config.MaxSources {
urls = urls[:d.config.MaxSources]
}
// Create channels for coordination
type chunkResult struct {
written int64
source string
}
chunkChan := make(chan chunkResult, len(urls))
errChan := make(chan error, len(urls))
doneChan := make(chan struct{})
var downloadedBytes int64
var sourcesUsed int32
var mu sync.Mutex
// Start download goroutines for each URL
var wg sync.WaitGroup
for i, u := range urls {
wg.Add(1)
go func(idx int, u URL) {
defer wg.Done()
select {
case <-ctx.Done():
return
default:
}
// Stream this source directly to the output file. All sources
// currently write from offset 0 (simple failover: last writer
// wins), serialised by &mu. The chunkResult carries only the
// byte count, not the payload, so no per-source body lingers
// in RAM while we wait for other sources to finish.
written, err := d.downloadChunk(ctx, u.URL, outFile, 0, &mu)
if err != nil {
if d.config.Verbose {
fmt.Fprintf(os.Stderr, "[metalink] Source %s failed: %v\n", u.URL, err)
}
errChan <- err
return
}
atomic.AddInt32(&sourcesUsed, 1)
chunkChan <- chunkResult{
written: written,
source: u.URL,
}
}(i, u)
}
// Collect results
go func() {
wg.Wait()
close(doneChan)
}()
// Track downloaded bytes and progress. The output file has already
// been written by each downloadChunk under &mu; we only need to
// aggregate counts here.
chunks:
for {
select {
case <-ctx.Done():
return nil, ctx.Err()
case result := <-chunkChan:
mu.Lock()
downloadedBytes += result.written
mu.Unlock()
if d.config.ProgressCallback != nil {
speed := float64(downloadedBytes) / time.Since(startTime).Seconds()
d.config.ProgressCallback(downloadedBytes, file.Size, speed)
}
case <-doneChan:
break chunks
case err := <-errChan:
if d.config.Verbose {
fmt.Fprintf(os.Stderr, "[metalink] Source error: %v\n", err)
}
// Continue with other sources
}
}
// Check if we got any data
if downloadedBytes == 0 {
return nil, fmt.Errorf("no data downloaded from any source")
}
duration := time.Since(startTime)
// Verify hash if available
verified := false
hash := ""
if file.HasHash("sha-256") {
expectedHash := file.GetSHA256()
actualHash, err := crypto.ComputeFileChecksum(outputPath, crypto.SHA256)
if err == nil {
verified = actualHash == expectedHash
hash = actualHash
if d.config.Verbose {
fmt.Fprintf(os.Stderr, "[metalink] Hash verification: %v (expected: %s, got: %s)\n",
verified, expectedHash, hash)
}
}
}
return &DownloadResult{
BytesDownloaded: downloadedBytes,
TotalSize: file.Size,
SourcesUsed: int(sourcesUsed),
Duration: duration,
Speed: float64(downloadedBytes) / duration.Seconds(),
OutputPath: outputPath,
Verified: verified,
Hash: hash,
}, nil
}
// downloadChunk downloads a chunk from a single source and streams it
// directly to outFile at the given offset, instead of buffering the
// entire body in memory. The fileMu mutex serialises the Seek + io.Copy
// sequence because the file position pointer is shared state and
// concurrent writes would race even though *os.File methods are
// individually safe.
func (d *MultiSourceDownloader) downloadChunk(ctx context.Context, rawURL string, outFile *os.File, offset int64, fileMu *sync.Mutex) (int64, error) {
parsedURL, err := url.Parse(rawURL)
if err != nil {
return 0, fmt.Errorf("invalid url: %w", err)
}
req, err := http.NewRequestWithContext(ctx, "GET", parsedURL.String(), nil)
if err != nil {
return 0, fmt.Errorf("failed to create request: %w", err)
}
req.Header.Set("User-Agent", "Goget/"+core.Version+" (Metalink)")
req.Header.Set("Accept-Encoding", "identity")
client := &http.Client{
Timeout: d.config.Timeout,
}
resp, err := client.Do(req)
if err != nil {
return 0, fmt.Errorf("request failed: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return 0, fmt.Errorf("http error: %s", resp.Status)
}
// Cap the response body to prevent disk exhaustion from a malicious
// server. The cap is applied before any data is buffered, so the
// multi-source download path no longer loads up to MaxDownloadSize
// bytes into RAM per source.
body := resp.Body
if d.config.MaxDownloadSize > 0 {
if resp.ContentLength > d.config.MaxDownloadSize {
return 0, fmt.Errorf("response size %d exceeds max download size %d", resp.ContentLength, d.config.MaxDownloadSize)
}
body = io.NopCloser(io.LimitReader(resp.Body, d.config.MaxDownloadSize))
}
// Serialise Seek + io.Copy: the file position is shared state, so
// concurrent writes from multiple sources would otherwise interleave
// or corrupt the output.
fileMu.Lock()
defer fileMu.Unlock()
if _, err := outFile.Seek(offset, io.SeekStart); err != nil {
return 0, fmt.Errorf("seek failed: %w", err)
}
written, err := io.Copy(outFile, body)
if err != nil {
return written, fmt.Errorf("failed to stream response: %w", err)
}
return written, nil
}
// DownloadWithFailover downloads a file with automatic failover to another source on error
func (d *MultiSourceDownloader) DownloadWithFailover(ctx context.Context, file *File, outputPath string) (*DownloadResult, error) {
startTime := time.Now()
urls := file.GetURLs()
if len(urls) == 0 {
return nil, fmt.Errorf("no URLs available")
}
var lastErr error
var downloaded int64
for i, u := range urls {
if d.config.Verbose {
fmt.Fprintf(os.Stderr, "[metalink] Trying source %d/%d: %s\n", i+1, len(urls), u.URL)
}
result, err := d.downloadFromSource(ctx, u.URL, outputPath, downloaded)
if err == nil {
// Success!
duration := time.Since(startTime)
return &DownloadResult{
BytesDownloaded: result.bytes,
TotalSize: file.Size,
SourcesUsed: 1,
Duration: duration,
Speed: float64(result.bytes) / duration.Seconds(),
OutputPath: outputPath,
Verified: result.verified,
Hash: result.hash,
}, nil
}
lastErr = err
downloaded = result.bytes
if d.config.Verbose {
fmt.Fprintf(os.Stderr, "[metalink] Source failed, trying next: %v\n", err)
}
}
return nil, fmt.Errorf("all sources failed: %w", lastErr)
}
type downloadResult struct {
bytes int64
verified bool
hash string
}
func (d *MultiSourceDownloader) downloadFromSource(ctx context.Context, rawURL, outputPath string, resumeOffset int64) (*downloadResult, error) {
parsedURL, err := url.Parse(rawURL)
if err != nil {
return nil, fmt.Errorf("invalid url: %w", err)
}
req, err := http.NewRequestWithContext(ctx, "GET", parsedURL.String(), nil)
if err != nil {
return nil, fmt.Errorf("failed to create request: %w", err)
}
req.Header.Set("User-Agent", "Goget/"+core.Version+" (Metalink)")
req.Header.Set("Accept-Encoding", "identity")
if resumeOffset > 0 {
req.Header.Set("Range", fmt.Sprintf("bytes=%d-", resumeOffset))
}
client := &http.Client{
Timeout: d.config.Timeout,
}
resp, err := client.Do(req)
if err != nil {
return &downloadResult{bytes: resumeOffset}, fmt.Errorf("request failed: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusPartialContent {
return &downloadResult{bytes: resumeOffset}, fmt.Errorf("http error: %s", resp.Status)
}
// Open file for writing
mode := os.O_CREATE | os.O_WRONLY
if resumeOffset > 0 {
mode |= os.O_APPEND
}
outFile, err := os.OpenFile(outputPath, mode, 0644)
if err != nil {
return &downloadResult{bytes: resumeOffset}, fmt.Errorf("failed to open file: %w", err)
}
defer outFile.Close()
buf := make([]byte, d.config.BufferSize)
var written int64
for {
if ctx.Err() != nil {
return &downloadResult{bytes: resumeOffset + written}, ctx.Err()
}
n, err := resp.Body.Read(buf)
if n > 0 {
_, wErr := outFile.Write(buf[:n])
if wErr != nil {
return &downloadResult{bytes: resumeOffset + written}, fmt.Errorf("write failed: %w", wErr)
}
written += int64(n)
}
if err == io.EOF {
break
}
if err != nil {
return &downloadResult{bytes: resumeOffset + written}, fmt.Errorf("read failed: %w", err)
}
}
return &downloadResult{
bytes: resumeOffset + written,
verified: false,
hash: "",
}, nil
}
// DownloadQueueItem represents an item for batch download
type DownloadQueueItem struct {
File *File
OutputPath string
Priority int
}
// DownloadBatch downloads multiple files from a metalink
func (d *MultiSourceDownloader) DownloadBatch(ctx context.Context, metalink *Metalink, outputDir string, maxParallel int) (*BatchResult, error) {
if maxParallel <= 0 {
maxParallel = 3
}
startTime := time.Now()
var totalBytes int64
var totalFiles int
var failedFiles int
sem := make(chan struct{}, maxParallel)
var wg sync.WaitGroup
var mu sync.Mutex
var firstError error
results := make(map[string]*DownloadResult)
for _, file := range metalink.Files {
wg.Add(1)
go func(f File) {
defer wg.Done()
sem <- struct{}{}
defer func() { <-sem }()
outputPath := filepath.Join(outputDir, f.Name)
result, err := d.DownloadWithFailover(ctx, &f, outputPath)
mu.Lock()
defer mu.Unlock()
if err != nil {
failedFiles++
if firstError == nil {
firstError = err
}
if d.config.Verbose {
fmt.Fprintf(os.Stderr, "[metalink] Failed to download %s: %v\n", f.Name, err)
}
} else {
totalBytes += result.BytesDownloaded
totalFiles++
results[f.Name] = result
if d.config.Verbose {
fmt.Fprintf(os.Stderr, "[metalink] Downloaded %s (%s)\n", f.Name, format.Bytes(result.BytesDownloaded))
}
}
}(file)
}
wg.Wait()
duration := time.Since(startTime)
return &BatchResult{
TotalFiles: len(metalink.Files),
Downloaded: totalFiles,
Failed: failedFiles,
TotalBytes: totalBytes,
Duration: duration,
Speed: float64(totalBytes) / duration.Seconds(),
Results: results,
FirstError: firstError,
}, nil
}
// BatchResult represents the result of a batch download
type BatchResult struct {
TotalFiles int
Downloaded int
Failed int
TotalBytes int64
Duration time.Duration
Speed float64
Results map[string]*DownloadResult
FirstError error
}
// SegmentDownloader downloads different segments of a file from different sources
type SegmentDownloader struct {
config *DownloaderConfig
}
// NewSegmentDownloader creates segment downloader
func NewSegmentDownloader(cfg *DownloaderConfig) *SegmentDownloader {
if cfg == nil {
cfg = DefaultDownloaderConfig()
}
return &SegmentDownloader{config: cfg}
}
// DownloadSegmented downloads a file split into segments from different sources
func (d *SegmentDownloader) DownloadSegmented(ctx context.Context, file *File, outputPath string) (*DownloadResult, error) {
startTime := time.Now()
if len(file.URLs) < 2 {
// Not enough sources for segmented download
msd := NewMultiSourceDownloader(d.config)
return msd.DownloadWithFailover(ctx, file, outputPath)
}
urls := file.GetURLs()
segmentSize := file.Size / int64(len(urls))
// Create temp files for segments
tempDir, err := os.MkdirTemp("", "goget-metalink-*")
if err != nil {
return nil, fmt.Errorf("failed to create temp dir: %w", err)
}
defer os.RemoveAll(tempDir)
var wg sync.WaitGroup
var mu sync.Mutex
var firstErr error
var downloadedBytes int64
results := make([]segmentResult, len(urls))
for i, u := range urls {
wg.Add(1)
go func(idx int, u URL) {
defer wg.Done()
start := int64(idx) * segmentSize
end := start + segmentSize
if idx == len(urls)-1 {
end = file.Size // Last segment gets remainder
}
tempPath := filepath.Join(tempDir, fmt.Sprintf("segment_%03d", idx))
result, err := d.downloadSegment(ctx, u.URL, start, end, tempPath)
mu.Lock()
defer mu.Unlock()
if err != nil {
if firstErr == nil {
firstErr = err
}
} else {
results[idx] = segmentResult{
path: tempPath,
offset: start,
size: result,
}
downloadedBytes += result
}
}(i, u)
}
wg.Wait()
if firstErr != nil && downloadedBytes == 0 {
return nil, firstErr
}
// Merge segments
if err := mergeSegments(outputPath, results, file.Size); err != nil {
return nil, fmt.Errorf("failed to merge segments: %w", err)
}
duration := time.Since(startTime)
// Verify hash
verified := false
hash := ""
if file.HasHash("sha-256") {
expectedHash := file.GetSHA256()
actualHash, err := crypto.ComputeFileChecksum(outputPath, crypto.SHA256)
if err == nil {
verified = actualHash == expectedHash
hash = actualHash
}
}
return &DownloadResult{
BytesDownloaded: downloadedBytes,
TotalSize: file.Size,
SourcesUsed: len(urls),
Duration: duration,
Speed: float64(downloadedBytes) / duration.Seconds(),
OutputPath: outputPath,
Verified: verified,
Hash: hash,
}, nil
}
func (d *SegmentDownloader) downloadSegment(ctx context.Context, rawURL string, start, end int64, outputPath string) (int64, error) {
parsedURL, err := url.Parse(rawURL)
if err != nil {
return 0, err
}
req, err := http.NewRequestWithContext(ctx, "GET", parsedURL.String(), nil)
if err != nil {
return 0, err
}
req.Header.Set("User-Agent", "Goget/"+core.Version+" (Metalink/Segmented)")
req.Header.Set("Range", fmt.Sprintf("bytes=%d-%d", start, end-1))
client := &http.Client{Timeout: d.config.Timeout}
resp, err := client.Do(req)
if err != nil {
return 0, err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusPartialContent {
return 0, fmt.Errorf("server doesn't support range requests")
}
outFile, err := os.Create(outputPath)
if err != nil {
return 0, err
}
defer outFile.Close()
buf := make([]byte, d.config.BufferSize)
var written int64
for {
n, err := resp.Body.Read(buf)
if n > 0 {
_, wErr := outFile.Write(buf[:n])
if wErr != nil {
return written, wErr
}
written += int64(n)
}
if err == io.EOF {
break
}
if err != nil {
return written, err
}
}
return written, nil
}
// segmentResult represents a downloaded segment
type segmentResult struct {
path string
offset int64
size int64
}
func mergeSegments(outputPath string, segments []segmentResult, totalSize int64) error {
outFile, err := os.Create(outputPath)
if err != nil {
return err
}
defer outFile.Close()
// Pre-allocate file
if err := outFile.Truncate(totalSize); err != nil {
return err
}
for _, seg := range segments {
if seg.path == "" {
continue
}
data, err := os.ReadFile(seg.path)
if err != nil {
return err
}
_, err = outFile.WriteAt(data, seg.offset)
if err != nil {
return err
}
}
return nil
}