Files
yw1573 be925abfdb perf: increase download and merge concurrency to 5
将并发下载数和合并数从 3 提升到 5,加快批量任务处理速度。

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-04-09 18:05:17 +08:00

617 lines
15 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package task
import (
"bufio"
"database/sql"
"encoding/base64"
"errors"
"fmt"
"io"
"net/http"
"os"
"os/exec"
"path/filepath"
"regexp"
"strconv"
"strings"
"sync"
"time"
"bilidown/internal/bilibili"
"bilidown/internal/common"
"bilidown/internal/logger"
"bilidown/internal/util"
)
// TaskInitOption 创建任务时需要从 POST 请求获取的参数
type TaskInitOption struct {
Bvid string `json:"bvid"`
Cid int `json:"cid"`
Format common.MediaFormat `json:"format"`
Title string `json:"title"`
Owner string `json:"owner"`
Cover string `json:"cover"`
Status TaskStatus `json:"status"`
Folder string `json:"folder"`
Audio string `json:"audio"`
Video string `json:"video"`
Duration int `json:"duration"`
DownloadType string `json:"downloadType"`
SeasonTitle string `json:"seasonTitle"` // 合集标题,用于创建子目录
}
// TaskInDB 任务数据库中的数据
type TaskInDB struct {
TaskInitOption
ID int64 `json:"id"`
CreateAt time.Time `json:"createAt"`
}
func (task *TaskInDB) FilePath() string {
ext := ".mp4"
if task.DownloadType == "audio" {
ext = ".m4a"
}
return filepath.Join(task.Folder,
fmt.Sprintf("%s %s%s", task.Title,
strings.Replace(base64.StdEncoding.EncodeToString([]byte(strconv.FormatInt(task.ID, 10))), "=", "", -1),
ext,
),
)
}
// done | waiting | running | error | cancelled
type TaskStatus string
type Task struct {
TaskInDB
AudioProgress float64 `json:"audioProgress"`
VideoProgress float64 `json:"videoProgress"`
MergeProgress float64 `json:"mergeProgress"`
Cancelled bool `json:"cancelled"`
}
var GlobalTaskList = []*Task{}
var GlobalTaskMux = &sync.Mutex{}
var GlobalDownloadSem = util.NewSemaphore(5)
var GlobalMergeSem = util.NewSemaphore(5)
// cleanupTempFiles 清理任务的临时文件
func (task *Task) cleanupTempFiles() {
if task.Folder == "" || task.ID == 0 {
return
}
tempFiles := []string{
filepath.Join(task.Folder, strconv.FormatInt(task.ID, 10)+".audio"),
filepath.Join(task.Folder, strconv.FormatInt(task.ID, 10)+".video"),
}
for _, f := range tempFiles {
if _, err := os.Stat(f); err == nil {
os.Remove(f)
logger.Debugf("清理临时文件: %s", f)
}
}
}
// CancelTask 取消指定任务
func CancelTask(taskID int64) bool {
GlobalTaskMux.Lock()
defer GlobalTaskMux.Unlock()
for _, task := range GlobalTaskList {
if task.ID == taskID && (task.Status == "waiting" || task.Status == "running") {
task.Cancelled = true
task.Status = "error"
logger.TaskCancelled(task.ID, task.Title)
return true
}
}
return false
}
// MarkAllTasksPaused 标记所有活跃任务为暂停状态,用于程序退出时保存状态
func MarkAllTasksPaused() {
GlobalTaskMux.Lock()
defer GlobalTaskMux.Unlock()
for _, task := range GlobalTaskList {
if task.Status == "waiting" || task.Status == "running" {
task.Cancelled = true
task.Status = "paused"
}
}
}
func (task *Task) Create(db *sql.DB) error {
util.SqliteLock.Lock()
result, err := db.Exec(`INSERT INTO "task" ("bvid", "cid", "format", "title", "owner", "cover", "status", "folder", "duration", "download_type")
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
task.Bvid,
task.Cid,
task.Format,
task.Title,
task.Owner,
task.Cover,
task.Status,
task.Folder,
task.Duration,
task.DownloadType,
)
util.SqliteLock.Unlock()
if err != nil {
return err
}
task.ID, err = result.LastInsertId()
task.CreateAt = time.Now()
logger.TaskCreated(task.ID, task.Title, task.Bvid)
return err
}
// Create 创建任务,并将任务加入全局任务列表
func (task *Task) Start() {
if task.DownloadType == "" {
task.DownloadType = "merge"
}
GlobalTaskMux.Lock()
GlobalTaskList = append(GlobalTaskList, task)
GlobalTaskMux.Unlock()
db := util.MustGetDB()
defer db.Close()
// 检查是否已取消
if task.Cancelled {
task.cleanupTempFiles()
task.cleanupTempFiles()
task.UpdateStatus(db, "error", fmt.Errorf("任务已取消"))
return
}
sessdata, err := bilibili.GetSessdata(db)
if err != nil {
task.UpdateStatus(db, "error", fmt.Errorf("bilibili.GetSessdata: %v", err))
return
}
// 确保下载目录存在
if err := os.MkdirAll(task.Folder, os.ModePerm); err != nil {
task.UpdateStatus(db, "error", fmt.Errorf("创建下载目录失败: %v", err))
return
}
// 如果有合集标题,创建子目录
if task.SeasonTitle != "" {
seasonFolder := filepath.Join(task.Folder, task.SeasonTitle)
if err := os.MkdirAll(seasonFolder, os.ModePerm); err != nil {
task.UpdateStatus(db, "error", fmt.Errorf("创建合集目录失败: %v", err))
return
}
task.Folder = seasonFolder
}
client := &bilibili.BiliClient{SESSDATA: sessdata}
GlobalDownloadSem.Acquire()
// 再次检查是否已取消
if task.Cancelled {
GlobalDownloadSem.Release()
task.cleanupTempFiles()
task.UpdateStatus(db, "error", fmt.Errorf("任务已取消"))
return
}
task.UpdateStatus(db, "running")
logger.TaskStarted(task.ID, task.Title)
if task.DownloadType == "audio" {
// 仅音频模式:只下载音频,重命名音频文件为输出文件
err = DownloadMedia(client, task.Audio, task, "audio")
if err != nil || task.Cancelled {
GlobalDownloadSem.Release()
if task.Cancelled {
task.cleanupTempFiles()
task.UpdateStatus(db, "error", fmt.Errorf("任务已取消"))
} else {
task.UpdateStatus(db, "error", fmt.Errorf("DownloadMedia: %v", err))
}
return
}
GlobalDownloadSem.Release()
outputPath := task.TaskInDB.FilePath()
audioPath := filepath.Join(task.Folder, strconv.FormatInt(task.ID, 10)+".audio")
err = os.Rename(audioPath, outputPath)
if err != nil {
task.UpdateStatus(db, "error", fmt.Errorf("os.Rename: %v", err))
return
}
// 添加元数据
if err := task.addMetadata(outputPath); err != nil {
logger.Warnf("添加元数据失败 #%d: %v", task.ID, err)
}
task.UpdateStatus(db, "done")
logger.TaskCompleted(task.ID, task.Title, time.Since(task.CreateAt))
return
} else if task.DownloadType == "video" {
// 仅视频模式:只下载视频,重命名视频文件为输出文件
err = DownloadMedia(client, task.Video, task, "video")
if err != nil || task.Cancelled {
GlobalDownloadSem.Release()
if task.Cancelled {
task.cleanupTempFiles()
task.UpdateStatus(db, "error", fmt.Errorf("任务已取消"))
} else {
task.UpdateStatus(db, "error", fmt.Errorf("DownloadMedia: %v", err))
}
return
}
GlobalDownloadSem.Release()
outputPath := task.TaskInDB.FilePath()
videoPath := filepath.Join(task.Folder, strconv.FormatInt(task.ID, 10)+".video")
err = os.Rename(videoPath, outputPath)
if err != nil {
task.UpdateStatus(db, "error", fmt.Errorf("os.Rename: %v", err))
return
}
// 添加元数据
if err := task.addMetadata(outputPath); err != nil {
logger.Warnf("添加元数据失败 #%d: %v", task.ID, err)
}
task.UpdateStatus(db, "done")
logger.TaskCompleted(task.ID, task.Title, time.Since(task.CreateAt))
return
} else {
// 合并模式:下载音频和视频,然后合并
err = DownloadMedia(client, task.Audio, task, "audio")
if err != nil || task.Cancelled {
GlobalDownloadSem.Release()
if task.Cancelled {
task.cleanupTempFiles()
task.UpdateStatus(db, "error", fmt.Errorf("任务已取消"))
} else {
task.UpdateStatus(db, "error", fmt.Errorf("DownloadMedia: %v", err))
}
return
}
err = DownloadMedia(client, task.Video, task, "video")
if err != nil || task.Cancelled {
GlobalDownloadSem.Release()
if task.Cancelled {
task.cleanupTempFiles()
task.UpdateStatus(db, "error", fmt.Errorf("任务已取消"))
} else {
task.UpdateStatus(db, "error", fmt.Errorf("DownloadMedia: %v", err))
}
return
}
GlobalDownloadSem.Release()
if task.Cancelled {
task.cleanupTempFiles()
task.UpdateStatus(db, "error", fmt.Errorf("任务已取消"))
return
}
outputPath := task.TaskInDB.FilePath()
videoPath := filepath.Join(task.Folder, strconv.FormatInt(task.ID, 10)+".video")
audioPath := filepath.Join(task.Folder, strconv.FormatInt(task.ID, 10)+".audio")
GlobalMergeSem.Acquire()
err = task.MergeMedia(outputPath, videoPath, audioPath)
if err != nil {
GlobalMergeSem.Release()
task.UpdateStatus(db, "error", fmt.Errorf("task.MergeMedia: %v", err))
return
}
err = os.Remove(videoPath)
if err != nil {
GlobalMergeSem.Release()
task.UpdateStatus(db, "error", fmt.Errorf("os.Remove: %v", err))
return
}
err = os.Remove(audioPath)
if err != nil {
GlobalMergeSem.Release()
task.UpdateStatus(db, "error", fmt.Errorf("os.Remove: %v", err))
return
}
GlobalMergeSem.Release()
// 添加元数据
if err := task.addMetadata(outputPath); err != nil {
logger.Warnf("添加元数据失败 #%d: %v", task.ID, err)
}
task.UpdateStatus(db, "done")
}
}
// 合并音视频
func (task *Task) MergeMedia(outputPath string, inputPaths ...string) error {
inputs := []string{}
for _, path := range inputPaths {
inputs = append(inputs, "-i", path)
}
ffmpegPath, err := util.GetFFmpegPath()
if err != nil {
return err
}
cmd := exec.Command(ffmpegPath, append(inputs, "-c:v", "copy", "-c:a", "copy", "-progress", "pipe:1", "-strict", "-2", outputPath)...)
stdout, err := cmd.StdoutPipe()
if err != nil {
return err
}
if err := cmd.Start(); err != nil {
return err
}
scanner := bufio.NewScanner(stdout)
progress := newProgressBar(int64(task.Duration))
outTimeRegex := regexp.MustCompile(`out_time_ms=(\d+)`) // 毫秒
for scanner.Scan() {
line := scanner.Text()
match := outTimeRegex.FindStringSubmatch(line)
if len(match) == 2 {
outTime, err := strconv.ParseInt(match[1], 10, 64)
if err != nil {
return err
}
progress.current = outTime / 1000000
task.MergeProgress = progress.percent()
}
}
if err := scanner.Err(); err != nil {
return err
}
if err := cmd.Wait(); err != nil {
return err
}
task.MergeProgress = 1
return nil
}
func GetVideoURL(medias []bilibili.Media, format common.MediaFormat) (string, error) {
for _, code := range []int{12, 7, 13} {
for _, item := range medias {
if item.ID == format && item.Codecid == code {
return item.BaseURL, nil
}
}
}
return "", errors.New("未找到对应视频分辨率格式")
}
func GetAudioURL(dash *bilibili.Dash) string {
if dash.Flac != nil {
return dash.Flac.Audio.BaseURL
}
var maxAudioID common.MediaFormat
var audioURL string
for _, item := range dash.Audio {
if item.ID > maxAudioID {
maxAudioID = item.ID
audioURL = item.BaseURL
}
}
return audioURL
}
func (task *Task) UpdateStatus(db *sql.DB, status TaskStatus, errs ...error) error {
util.SqliteLock.Lock()
_, err := db.Exec(`UPDATE "task" SET "status" = ? WHERE "id" = ?`, status, task.ID)
util.SqliteLock.Unlock()
if err != nil {
return nil
}
for _, err := range errs {
if err != nil {
err = util.CreateLog(db, fmt.Sprintf("Task-%d-Error: %v", task.ID, err))
if err != nil {
logger.Fatal("CreateLog: " + err.Error())
}
logger.TaskFailed(task.ID, task.Title, err)
}
}
task.Status = status
return err
}
func DownloadMedia(client *bilibili.BiliClient, _url string, task *Task, mediaType string) error {
var resp *http.Response
var err error
for i := 0; i < 5; i++ {
resp, err = client.SimpleGET(_url, nil)
if err == nil {
break
}
}
if err != nil {
return err
}
filename := strconv.FormatInt(task.ID, 10) + "." + mediaType
filepath := filepath.Join(task.Folder, filename)
progress := newProgressBar(resp.ContentLength)
file, err := os.Create(filepath)
if err != nil {
return err
}
defer file.Close()
reader := io.TeeReader(resp.Body, file)
buf := make([]byte, 1024)
for {
n, err := reader.Read(buf)
if err != nil && err != io.EOF {
return err
}
if n == 0 {
break
}
progress.add(n)
GlobalTaskMux.Lock()
if mediaType == "video" {
task.VideoProgress = progress.percent()
} else {
task.AudioProgress = progress.percent()
}
GlobalTaskMux.Unlock()
}
return nil
}
type progressBar struct {
total int64
current int64
}
func (p *progressBar) add(n int) {
p.current += int64(n)
}
func (p *progressBar) percent() float64 {
// 防止除零错误:total 为 0 或负数时返回 0
if p.total <= 0 {
return 0
}
return float64(p.current) / float64(p.total)
}
func newProgressBar(total int64) *progressBar {
return &progressBar{
total: total,
}
}
func GetTaskList(db *sql.DB, page int, pageSize int) ([]TaskInDB, error) {
tasks := []TaskInDB{}
util.SqliteLock.Lock()
rows, err := db.Query(`SELECT
"id", "bvid", "cid", "format", "title",
"owner", "cover", "status", "folder", "duration", "download_type", "create_at"
FROM "task" ORDER BY "id" DESC LIMIT ?, ?`,
page*pageSize, pageSize,
)
util.SqliteLock.Unlock()
if err != nil {
return nil, err
}
createAt := ""
for rows.Next() {
task := TaskInDB{}
err = rows.Scan(
&task.ID,
&task.Bvid,
&task.Cid,
&task.Format,
&task.Title,
&task.Owner,
&task.Cover,
&task.Status,
&task.Folder,
&task.Duration,
&task.DownloadType,
&createAt,
)
if err != nil {
return nil, err
}
task.CreateAt, err = time.Parse("2006-01-02 15:04:05", createAt)
if err != nil {
return nil, err
}
tasks = append(tasks, task)
}
return tasks, nil
}
func DeleteTask(db *sql.DB, taskID int) error {
util.SqliteLock.Lock()
_, err := db.Exec(`DELETE FROM "task" WHERE "id" = ?`, taskID)
util.SqliteLock.Unlock()
return err
}
func GetTask(db *sql.DB, taskID int) (*TaskInDB, error) {
task := TaskInDB{}
createAt := ""
util.SqliteLock.Lock()
err := db.QueryRow(`SELECT
"id", "bvid", "cid", "format", "title",
"owner", "cover", "status", "folder", "duration", "download_type", "create_at"
FROM "task" WHERE "id" = ?`,
taskID,
).Scan(
&task.ID,
&task.Bvid,
&task.Cid,
&task.Format,
&task.Title,
&task.Owner,
&task.Cover,
&task.Status,
&task.Folder,
&task.Duration,
&task.DownloadType,
&createAt,
)
util.SqliteLock.Unlock()
if err != nil {
return nil, err
}
task.CreateAt, err = time.Parse("2006-01-02 15:04:05", createAt)
if err != nil {
return nil, err
}
return &task, nil
}
// addMetadata 使用 ffmpeg 给输出文件添加元数据(description 和 artist
func (task *Task) addMetadata(filePath string) error {
ffmpegPath, err := util.GetFFmpegPath()
if err != nil {
return err
}
desc := task.Bvid
if desc == "" {
desc = ""
}
author := task.Owner
// 临时文件加上 .mp4 扩展名
tempPath := filePath + ".tmp.mp4"
// 使用双引号包裹文件路径,避免特殊字符
cmd := exec.Command(ffmpegPath,
"-i", filePath,
"-metadata", "description="+desc,
"-metadata", "artist="+author,
"-codec", "copy",
"-y",
tempPath,
)
output, err := cmd.CombinedOutput()
if err != nil {
return fmt.Errorf("ffmpeg添加元数据失败: %v, 输出: %s", err, string(output))
}
if err := os.Remove(filePath); err != nil {
return fmt.Errorf("删除原文件失败: %v", err)
}
if err := os.Rename(tempPath, filePath); err != nil {
return fmt.Errorf("重命名临时文件失败: %v", err)
}
return nil
}