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
This commit is contained in:
2026-04-09 16:51:22 +08:00
parent 449b5fe3a2
commit a0ef988bd5
5 changed files with 125 additions and 47 deletions
+53 -9
View File
@@ -1,15 +1,20 @@
package main package main
import ( import (
"context"
"database/sql" "database/sql"
"embed" "embed"
"fmt" "fmt"
"io/fs" "io/fs"
"log"
"net/http" "net/http"
"os"
"os/signal"
"syscall"
"time"
"bilidown/internal/logger" "bilidown/internal/logger"
"bilidown/internal/router" "bilidown/internal/router"
"bilidown/internal/task"
"bilidown/internal/util" "bilidown/internal/util"
_ "modernc.org/sqlite" _ "modernc.org/sqlite"
@@ -21,11 +26,12 @@ var staticFiles embed.FS
const ( const (
HTTP_PORT = 8098 // HTTP 服务器端口 HTTP_PORT = 8098 // HTTP 服务器端口
HTTP_HOST = "" // 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 urlLocal = fmt.Sprintf("http://127.0.0.1:%d", HTTP_PORT)
var ffmpegAvailable bool var ffmpegAvailable bool
var server *http.Server
func main() { func main() {
ffmpegAvailable = checkFFmpeg() ffmpegAvailable = checkFFmpeg()
@@ -33,7 +39,14 @@ func main() {
mustInitTables() mustInitTables()
mustRunServer() mustRunServer()
logger.ServerStarted(HTTP_PORT, VERSION) 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 的安装情况,返回是否可用 // checkFFmpeg 检测 ffmpeg 的安装情况,返回是否可用
@@ -59,14 +72,44 @@ func mustRunServer() {
}) })
http.Handle("/api/", http.StripPrefix("/api", apiRouter)) http.Handle("/api/", http.StripPrefix("/api", apiRouter))
// 启动 HTTP 服务器 // 启动 HTTP 服务器
server = &http.Server{
Addr: fmt.Sprintf("%s:%d", HTTP_HOST, HTTP_PORT),
Handler: nil,
}
go func() { go func() {
err := http.ListenAndServe(fmt.Sprintf("%s:%d", HTTP_HOST, HTTP_PORT), nil) if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed {
if err != nil { logger.Fatal("http.ListenAndServe: " + err.Error())
log.Fatal("http.ListenAndServe:", err)
} }
}() }()
} }
// 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 初始化数据表 // mustInitTables 初始化数据表
func mustInitTables() { func mustInitTables() {
db := util.MustGetDB() db := util.MustGetDB()
@@ -129,17 +172,18 @@ func addMissingColumns(db *sql.DB) error {
return nil return nil
} }
// initHistoryTask 将上一次程序运行时未完成的任务状态变为 error // initHistoryTask 将上一次程序运行时未完成的任务状态变为 paused,支持恢复
func initHistoryTask(db *sql.DB) error { func initHistoryTask(db *sql.DB) error {
util.SqliteLock.Lock() util.SqliteLock.Lock()
result, err := db.Exec(`UPDATE "task" SET "status" = 'error' WHERE "status" IN ('waiting', 'running')`) // 将 waitingrunning 状态改为 paused,而非 error
result, err := db.Exec(`UPDATE "task" SET "status" = 'paused' WHERE "status" IN ('waiting', 'running')`)
util.SqliteLock.Unlock() util.SqliteLock.Unlock()
if err != nil { if err != nil {
return err return err
} }
rowsAffected, _ := result.RowsAffected() rowsAffected, _ := result.RowsAffected()
if rowsAffected > 0 { if rowsAffected > 0 {
logger.Infof("重置 %d 个未完成任务状态为 error", rowsAffected) logger.Infof("发现 %d 个未完成任务状态已设为 paused", rowsAffected)
} }
return nil return nil
} }
+2 -26
View File
@@ -1,11 +1,9 @@
package router package router
import ( import (
"database/sql"
"encoding/json" "encoding/json"
"fmt" "fmt"
"net/http" "net/http"
"os"
"os/exec" "os/exec"
"runtime" "runtime"
"strconv" "strconv"
@@ -177,30 +175,8 @@ func deleteTask(w http.ResponseWriter, r *http.Request) {
}{} }{}
for _, taskID := range taskIDs { for _, taskID := range taskIDs {
_task, err := task.GetTask(db, taskID) // 只删除数据库记录,不删除已下载的视频文件
if err == sql.ErrNoRows { err := task.DeleteTask(db, taskID)
// 数据库中没有该条记录,忽略
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)
if err != nil { if err != nil {
failedTasks = append(failedTasks, struct { failedTasks = append(failedTasks, struct {
ID int ID int
+35 -2
View File
@@ -133,9 +133,42 @@ func getPopularVideos(w http.ResponseWriter, r *http.Request) {
var downloadVideo = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { var downloadVideo = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
path := r.URL.Query().Get("path") 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 := filepath.Clean(path)
safePath = strings.ReplaceAll(safePath, "\\", "/") absPath, err := filepath.Abs(safePath)
http.ServeFile(w, r, 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) { var getSeasonsArchivesListFirstBvid = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+16
View File
@@ -108,6 +108,18 @@ func CancelTask(taskID int64) bool {
return false 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 { func (task *Task) Create(db *sql.DB) error {
util.SqliteLock.Lock() util.SqliteLock.Lock()
result, err := db.Exec(`INSERT INTO "task" ("bvid", "cid", "format", "title", "owner", "cover", "status", "folder", "duration", "download_type") 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 { func (p *progressBar) percent() float64 {
// 防止除零错误:total 为 0 或负数时返回 0
if p.total <= 0 {
return 0
}
return float64(p.current) / float64(p.total) return float64(p.current) / float64(p.total)
} }
+18 -9
View File
@@ -4,7 +4,6 @@ import (
"database/sql" "database/sql"
"errors" "errors"
"fmt" "fmt"
"log"
"strings" "strings"
"sync" "sync"
) )
@@ -115,18 +114,28 @@ func SaveDownloadFolder(db *sql.DB, downloadFolder string) error {
var SqliteLock sync.Mutex var SqliteLock sync.Mutex
func MustGetDB(path ...string) *sql.DB { // GetDB 获取数据库连接,返回错误而非直接退出
pathStr := "" func GetDB(path ...string) (*sql.DB, error) {
if len(path) == 0 { pathStr := "./data.db"
pathStr = "./data.db" if len(path) > 0 {
} else if len(path) > 1 { if len(path) > 1 {
log.Fatalln(errors.New("len(path) <= 1")) return nil, errors.New("len(path) must be <= 1")
} else { }
pathStr = path[0] pathStr = path[0]
} }
db, err := sql.Open("sqlite", pathStr) db, err := sql.Open("sqlite", pathStr)
if err != nil { 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 return db
} }