Files
BiliDown/internal/task/task.go
T
yw1573 60b8a98a74 fix: cleanup temp files when task is cancelled
任务取消时清理临时文件(.audio/.video),避免残留

Co-Authored-By: AI
2026-04-09 11:20:14 +08:00

590 lines
14 KiB
Go
Raw 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"`
}
// 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(3)
var GlobalMergeSem = util.NewSemaphore(3)
// 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
}
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
}
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 {
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
}