From a0ef988bd5606690f62ca2e09abe27c9cfdd2f40 Mon Sep 17 00:00:00 2001 From: yw1573 Date: Thu, 9 Apr 2026 16:51:22 +0800 Subject: [PATCH] fix: resolve critical issues and improve task management - Fix path traversal vulnerability in downloadVideo handler by adding download directory whitelist validation - Add graceful shutdown with signal handling for task persistence - Fix division by zero panic in progressBar.percent() when total <= 0 - Add GetDB() function that returns error instead of using log.Fatal - Change deleteTask to only remove database records, preserve downloaded files - Add paused status for interrupted tasks on shutdown Co-Authored-By: Claude --- cmd/bilidown/main.go | 62 ++++++++++++++++++++++++++++++++++------ internal/router/task.go | 28 ++---------------- internal/router/video.go | 39 +++++++++++++++++++++++-- internal/task/task.go | 16 +++++++++++ internal/util/db.go | 27 +++++++++++------ 5 files changed, 125 insertions(+), 47 deletions(-) diff --git a/cmd/bilidown/main.go b/cmd/bilidown/main.go index 0826533..66a619b 100644 --- a/cmd/bilidown/main.go +++ b/cmd/bilidown/main.go @@ -1,15 +1,20 @@ package main import ( + "context" "database/sql" "embed" "fmt" "io/fs" - "log" "net/http" + "os" + "os/signal" + "syscall" + "time" "bilidown/internal/logger" "bilidown/internal/router" + "bilidown/internal/task" "bilidown/internal/util" _ "modernc.org/sqlite" @@ -21,11 +26,12 @@ var staticFiles embed.FS const ( HTTP_PORT = 8098 // HTTP 服务器端口 HTTP_HOST = "" // HTTP 服务器主机 - VERSION = "v2.1.1" // 软件版本号 + VERSION = "v2.1.2" // 软件版本号 ) var urlLocal = fmt.Sprintf("http://127.0.0.1:%d", HTTP_PORT) var ffmpegAvailable bool +var server *http.Server func main() { ffmpegAvailable = checkFFmpeg() @@ -33,7 +39,14 @@ func main() { mustInitTables() mustRunServer() logger.ServerStarted(HTTP_PORT, VERSION) - select {} // 保持运行 + + // 等待中断信号,实现优雅退出 + quit := make(chan os.Signal, 1) + signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM) + <-quit + + logger.Info("正在关闭服务...") + gracefulShutdown() } // checkFFmpeg 检测 ffmpeg 的安装情况,返回是否可用 @@ -59,14 +72,44 @@ func mustRunServer() { }) http.Handle("/api/", http.StripPrefix("/api", apiRouter)) // 启动 HTTP 服务器 + server = &http.Server{ + Addr: fmt.Sprintf("%s:%d", HTTP_HOST, HTTP_PORT), + Handler: nil, + } go func() { - err := http.ListenAndServe(fmt.Sprintf("%s:%d", HTTP_HOST, HTTP_PORT), nil) - if err != nil { - log.Fatal("http.ListenAndServe:", err) + if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed { + logger.Fatal("http.ListenAndServe: " + err.Error()) } }() } +// gracefulShutdown 优雅退出:保存任务状态并关闭服务 +func gracefulShutdown() { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + // 等待活跃任务完成或超时 + task.GlobalTaskMux.Lock() + activeCount := len(task.GlobalTaskList) + task.GlobalTaskMux.Unlock() + + if activeCount > 0 { + logger.Infof("等待 %d 个活跃任务完成(最多10秒)...", activeCount) + // 标记所有运行中的任务为暂停状态 + task.MarkAllTasksPaused() + // 等待下载信号量释放 + task.GlobalDownloadSem.Wait() + task.GlobalMergeSem.Wait() + } + + // 关闭 HTTP 服务器 + if err := server.Shutdown(ctx); err != nil { + logger.Error("服务器关闭失败: " + err.Error()) + } + + logger.Info("服务已关闭") +} + // mustInitTables 初始化数据表 func mustInitTables() { db := util.MustGetDB() @@ -129,17 +172,18 @@ func addMissingColumns(db *sql.DB) error { return nil } -// initHistoryTask 将上一次程序运行时未完成的任务状态变为 error +// initHistoryTask 将上一次程序运行时未完成的任务状态变为 paused,支持恢复 func initHistoryTask(db *sql.DB) error { util.SqliteLock.Lock() - result, err := db.Exec(`UPDATE "task" SET "status" = 'error' WHERE "status" IN ('waiting', 'running')`) + // 将 waiting 和 running 状态改为 paused,而非 error + result, err := db.Exec(`UPDATE "task" SET "status" = 'paused' WHERE "status" IN ('waiting', 'running')`) util.SqliteLock.Unlock() if err != nil { return err } rowsAffected, _ := result.RowsAffected() if rowsAffected > 0 { - logger.Infof("重置 %d 个未完成任务状态为 error", rowsAffected) + logger.Infof("发现 %d 个未完成任务,状态已设为 paused", rowsAffected) } return nil } \ No newline at end of file diff --git a/internal/router/task.go b/internal/router/task.go index 0e1c550..b7194f1 100644 --- a/internal/router/task.go +++ b/internal/router/task.go @@ -1,11 +1,9 @@ package router import ( - "database/sql" "encoding/json" "fmt" "net/http" - "os" "os/exec" "runtime" "strconv" @@ -177,30 +175,8 @@ func deleteTask(w http.ResponseWriter, r *http.Request) { }{} for _, taskID := range taskIDs { - _task, err := task.GetTask(db, taskID) - if err == sql.ErrNoRows { - // 数据库中没有该条记录,忽略 - successCount++ - continue - } - if err != nil { - failedTasks = append(failedTasks, struct { - ID int - Error string - }{ID: taskID, Error: fmt.Sprintf("获取任务失败: %v", err)}) - continue - } - filePath := _task.FilePath() - err = os.Remove(filePath) - if err != nil && !os.IsNotExist(err) { - failedTasks = append(failedTasks, struct { - ID int - Error string - }{ID: taskID, Error: fmt.Sprintf("文件删除失败: %v", err)}) - continue - } - - err = task.DeleteTask(db, taskID) + // 只删除数据库记录,不删除已下载的视频文件 + err := task.DeleteTask(db, taskID) if err != nil { failedTasks = append(failedTasks, struct { ID int diff --git a/internal/router/video.go b/internal/router/video.go index 4d0db05..d974c47 100644 --- a/internal/router/video.go +++ b/internal/router/video.go @@ -133,9 +133,42 @@ func getPopularVideos(w http.ResponseWriter, r *http.Request) { var downloadVideo = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { path := r.URL.Query().Get("path") + if path == "" { + res_error.Send(w, res_error.ParamError) + return + } + + // 获取下载目录作为白名单 + db := util.MustGetDB() + defer db.Close() + downloadFolder, err := util.GetCurrentFolder(db) + if err != nil { + res_error.Send(w, fmt.Sprintf("获取下载目录失败: %v", err)) + return + } + + // 清理路径并转换为绝对路径 safePath := filepath.Clean(path) - safePath = strings.ReplaceAll(safePath, "\\", "/") - http.ServeFile(w, r, safePath) + absPath, err := filepath.Abs(safePath) + if err != nil { + res_error.Send(w, "无效的文件路径") + return + } + + // 获取下载目录的绝对路径 + absDownloadFolder, err := filepath.Abs(downloadFolder) + if err != nil { + res_error.Send(w, "无效的下载目录") + return + } + + // 安全检查:确保请求的文件在下载目录内 + if !strings.HasPrefix(absPath, absDownloadFolder) { + res_error.Send(w, "禁止访问该路径") + return + } + + http.ServeFile(w, r, absPath) }) var getSeasonsArchivesListFirstBvid = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -179,4 +212,4 @@ var getFavList = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { return } util.Res{Success: true, Message: "获取成功", Data: favList}.Write(w) -}) +}) \ No newline at end of file diff --git a/internal/task/task.go b/internal/task/task.go index 4c60180..3811738 100644 --- a/internal/task/task.go +++ b/internal/task/task.go @@ -108,6 +108,18 @@ func CancelTask(taskID int64) bool { 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") @@ -463,6 +475,10 @@ func (p *progressBar) add(n int) { } func (p *progressBar) percent() float64 { + // 防止除零错误:total 为 0 或负数时返回 0 + if p.total <= 0 { + return 0 + } return float64(p.current) / float64(p.total) } diff --git a/internal/util/db.go b/internal/util/db.go index b521f75..85d6462 100644 --- a/internal/util/db.go +++ b/internal/util/db.go @@ -4,7 +4,6 @@ import ( "database/sql" "errors" "fmt" - "log" "strings" "sync" ) @@ -115,18 +114,28 @@ func SaveDownloadFolder(db *sql.DB, downloadFolder string) error { var SqliteLock sync.Mutex -func MustGetDB(path ...string) *sql.DB { - pathStr := "" - if len(path) == 0 { - pathStr = "./data.db" - } else if len(path) > 1 { - log.Fatalln(errors.New("len(path) <= 1")) - } else { +// GetDB 获取数据库连接,返回错误而非直接退出 +func GetDB(path ...string) (*sql.DB, error) { + pathStr := "./data.db" + if len(path) > 0 { + if len(path) > 1 { + return nil, errors.New("len(path) must be <= 1") + } pathStr = path[0] } db, err := sql.Open("sqlite", pathStr) if err != nil { - log.Fatalln("sql.Open:", err) + return nil, fmt.Errorf("sql.Open: %w", err) + } + return db, nil +} + +// MustGetDB 获取数据库连接,失败时 panic(用于初始化阶段) +// 推荐在 main 函数中使用 GetDB 并处理错误 +func MustGetDB(path ...string) *sql.DB { + db, err := GetDB(path...) + if err != nil { + panic(fmt.Sprintf("MustGetDB: %v", err)) } return db }