如何编写多线程加速文件下载脚本

wen 实用脚本 31

本文目录导读:

如何编写多线程加速文件下载脚本

  1. Python实现(使用requests + threading)
  2. 使用aria2c命令行工具(推荐)
  3. Python实现(使用urllib + asyncio)
  4. Shell脚本(使用wget和curl)
  5. Go语言实现
  6. 使用建议

Python实现(使用requests + threading)

import requests
import threading
import os
import time
from concurrent.futures import ThreadPoolExecutor
class MultiThreadDownloader:
    def __init__(self, url, filename=None, num_threads=4):
        self.url = url
        self.filename = filename or url.split('/')[-1]
        self.num_threads = num_threads
        self.file_size = 0
    def get_file_size(self):
        """获取文件大小"""
        response = requests.head(self.url)
        self.file_size = int(response.headers.get('content-length', 0))
        return self.file_size
    def download_range(self, start, end, part_num):
        """下载文件指定范围"""
        headers = {'Range': f'bytes={start}-{end}'}
        response = requests.get(self.url, headers=headers, stream=True)
        # 写入临时文件
        temp_filename = f"{self.filename}.part{part_num}"
        with open(temp_filename, 'wb') as f:
            for chunk in response.iter_content(chunk_size=8192):
                if chunk:
                    f.write(chunk)
        print(f"Part {part_num} downloaded: {start}-{end}")
        return temp_filename
    def merge_files(self, part_files):
        """合并下载的文件"""
        with open(self.filename, 'wb') as output:
            for part_file in part_files:
                with open(part_file, 'rb') as f:
                    output.write(f.read())
                os.remove(part_file)
        print(f"File merged successfully: {self.filename}")
    def download(self):
        """多线程下载主函数"""
        # 获取文件大小
        file_size = self.get_file_size()
        if file_size == 0:
            print("无法获取文件大小")
            return
        # 计算每个线程的下载范围
        part_size = file_size // self.num_threads
        ranges = []
        for i in range(self.num_threads):
            start = i * part_size
            end = start + part_size - 1 if i < self.num_threads - 1 else file_size - 1
            ranges.append((start, end))
        # 多线程下载
        print(f"开始下载,文件大小: {file_size} bytes,线程数: {self.num_threads}")
        start_time = time.time()
        with ThreadPoolExecutor(max_workers=self.num_threads) as executor:
            futures = []
            for i, (start, end) in enumerate(ranges):
                future = executor.submit(self.download_range, start, end, i)
                futures.append(future)
            # 等待所有线程完成
            part_files = [future.result() for future in futures]
        # 合并文件
        self.merge_files(part_files)
        elapsed_time = time.time() - start_time
        print(f"下载完成,耗时: {elapsed_time:.2f}秒")
# 使用示例
if __name__ == "__main__":
    url = "https://example.com/large-file.zip"
    downloader = MultiThreadDownloader(url, num_threads=8)
    downloader.download()

使用aria2c命令行工具(推荐)

#!/bin/bash
# 多线程下载脚本
download_url="https://example.com/large-file.zip"
output_file="large-file.zip"
# 使用aria2c下载,16个连接
aria2c \
    --max-connection-per-server=16 \
    --split=16 \
    --min-split-size=1M \
    --continue=true \
    --max-concurrent-downloads=5 \
    --dir="./downloads" \
    --out="$output_file" \
    "$download_url"
echo "下载完成"
# 高级用法:批量下载
aria2c \
    --max-connection-per-server=16 \
    --split=16 \
    --input-file=urls.txt \
    --max-concurrent-downloads=5

Python实现(使用urllib + asyncio)

import asyncio
import aiohttp
import aiofiles
import os
from typing import List
class AsyncMultiThreadDownloader:
    def __init__(self, url: str, filename: str = None, num_connections: int = 4):
        self.url = url
        self.filename = filename or url.split('/')[-1]
        self.num_connections = num_connections
        self.file_size = 0
    async def get_file_size(self, session: aiohttp.ClientSession) -> int:
        """异步获取文件大小"""
        async with session.head(self.url) as response:
            return int(response.headers.get('content-length', 0))
    async def download_part(self, session: aiohttp.ClientSession, 
                           start: int, end: int, part_num: int):
        """异步下载文件部分"""
        headers = {'Range': f'bytes={start}-{end}'}
        temp_filename = f"{self.filename}.part{part_num}"
        async with session.get(self.url, headers=headers) as response:
            async with aiofiles.open(temp_filename, 'wb') as f:
                async for chunk in response.content.iter_chunked(8192):
                    await f.write(chunk)
        print(f"Part {part_num} downloaded: {start}-{end}")
        return temp_filename
    async def merge_parts(self, part_files: List[str]):
        """异步合并文件"""
        async with aiofiles.open(self.filename, 'wb') as output:
            for part_file in part_files:
                async with aiofiles.open(part_file, 'rb') as f:
                    content = await f.read()
                    await output.write(content)
                os.remove(part_file)
        print(f"File merged: {self.filename}")
    async def download(self):
        """异步下载主函数"""
        async with aiohttp.ClientSession() as session:
            # 获取文件大小
            self.file_size = await self.get_file_size(session)
            if self.file_size == 0:
                print("无法获取文件大小")
                return
            # 计算范围
            part_size = self.file_size // self.num_connections
            tasks = []
            print(f"开始下载: {self.file_size} bytes,{self.num_connections}个连接")
            for i in range(self.num_connections):
                start = i * part_size
                end = start + part_size - 1 if i < self.num_connections - 1 else self.file_size - 1
                task = self.download_part(session, start, end, i)
                tasks.append(task)
            # 并发执行所有下载任务
            part_files = await asyncio.gather(*tasks)
            # 合并文件
            await self.merge_parts(part_files)
            print("下载完成!")
# 使用示例
async def main():
    url = "https://example.com/large-file.zip"
    downloader = AsyncMultiThreadDownloader(url, num_connections=8)
    await downloader.download()
if __name__ == "__main__":
    asyncio.run(main())

Shell脚本(使用wget和curl)

#!/bin/bash
# 多线程下载函数
multi_thread_download() {
    local url=$1
    local filename=$2
    local threads=${3:-4}
    # 获取文件大小
    file_size=$(curl -sI "$url" | grep -i content-length | awk '{print $2}' | tr -d '\r')
    if [ -z "$file_size" ]; then
        echo "无法获取文件大小"
        return 1
    fi
    echo "文件大小: $file_size bytes"
    # 计算每个线程的块大小
    block_size=$((file_size / threads))
    # 启动多个后台进程
    for ((i=0; i<threads; i++)); do
        start=$((i * block_size))
        if [ $i -eq $((threads - 1)) ]; then
            end=$((file_size - 1))
        else
            end=$(((i + 1) * block_size - 1))
        fi
        # 使用curl下载指定范围
        curl -r "$start-$end" -o "${filename}.part$i" "$url" &
    done
    # 等待所有下载完成
    wait
    # 合并文件
    for ((i=0; i<threads; i++)); do
        cat "${filename}.part$i" >> "$filename"
        rm "${filename}.part$i"
    done
    echo "下载完成: $filename"
}
# 使用示例
multi_thread_download "https://example.com/large-file.zip" "output.zip" 8

Go语言实现

package main
import (
    "fmt"
    "io"
    "net/http"
    "os"
    "strconv"
    "sync"
    "time"
)
type Downloader struct {
    url        string
    filename   string
    threads    int
    fileSize   int64
}
func NewDownloader(url, filename string, threads int) *Downloader {
    return &Downloader{
        url:      url,
        filename: filename,
        threads:  threads,
    }
}
func (d *Downloader) getFileSize() (int64, error) {
    resp, err := http.Head(d.url)
    if err != nil {
        return 0, err
    }
    defer resp.Body.Close()
    size, err := strconv.ParseInt(resp.Header.Get("Content-Length"), 10, 64)
    if err != nil {
        return 0, err
    }
    return size, nil
}
func (d *Downloader) downloadPart(start, end int64, partNum int, wg *sync.WaitGroup) {
    defer wg.Done()
    client := &http.Client{}
    req, err := http.NewRequest("GET", d.url, nil)
    if err != nil {
        fmt.Printf("Part %d: Error creating request: %v\n", partNum, err)
        return
    }
    req.Header.Set("Range", fmt.Sprintf("bytes=%d-%d", start, end))
    resp, err := client.Do(req)
    if err != nil {
        fmt.Printf("Part %d: Error downloading: %v\n", partNum, err)
        return
    }
    defer resp.Body.Close()
    partFile := fmt.Sprintf("%s.part%d", d.filename, partNum)
    file, err := os.Create(partFile)
    if err != nil {
        fmt.Printf("Part %d: Error creating file: %v\n", partNum, err)
        return
    }
    defer file.Close()
    written, err := io.Copy(file, resp.Body)
    if err != nil {
        fmt.Printf("Part %d: Error writing file: %v\n", partNum, err)
        return
    }
    fmt.Printf("Part %d downloaded: %d bytes\n", partNum, written)
}
func (d *Downloader) Download() error {
    // 获取文件大小
    fileSize, err := d.getFileSize()
    if err != nil {
        return fmt.Errorf("Error getting file size: %v", err)
    }
    d.fileSize = fileSize
    fmt.Printf("文件大小: %d bytes, 线程数: %d\n", fileSize, d.threads)
    // 计算每块大小
    partSize := fileSize / int64(d.threads)
    var wg sync.WaitGroup
    // 启动下载线程
    for i := 0; i < d.threads; i++ {
        start := int64(i) * partSize
        end := start + partSize - 1
        if i == d.threads-1 {
            end = fileSize - 1
        }
        wg.Add(1)
        go d.downloadPart(start, end, i, &wg)
    }
    wg.Wait()
    // 合并文件
    outputFile, err := os.Create(d.filename)
    if err != nil {
        return fmt.Errorf("Error creating output file: %v", err)
    }
    defer outputFile.Close()
    for i := 0; i < d.threads; i++ {
        partFile := fmt.Sprintf("%s.part%d", d.filename, i)
        f, err := os.Open(partFile)
        if err != nil {
            return fmt.Errorf("Error opening part file: %v", err)
        }
        io.Copy(outputFile, f)
        f.Close()
        os.Remove(partFile)
    }
    return nil
}
func main() {
    startTime := time.Now()
    downloader := NewDownloader(
        "https://example.com/large-file.zip",
        "output.zip",
        8, // 8个并发
    )
    if err := downloader.Download(); err != nil {
        fmt.Printf("Error: %v\n", err)
        return
    }
    elapsedTime := time.Since(startTime)
    fmt.Printf("下载完成,耗时: %v\n", elapsedTime)
}

使用建议

  1. 选择合适的工具

    • 简单使用:aria2c(命令行最快)
    • 需要集成到Python项目:使用requests + threading
    • 高性能需求:Go或Rust实现
  2. 注意事项

    • 服务器是否支持范围请求(Range)
    • 适当的线程数(通常4-16个)
    • 错误处理和断点续传
    • 网络带宽和服务器限制
  3. 优化建议

    • 动态调整线程数
    • 实现断点续传
    • 添加进度显示
    • 处理网络波动

选择哪种实现取决于你的具体需求和环境,aria2c是最简单高效的方案,Python方案适合集成到其他应用中。

抱歉,评论功能暂时关闭!