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:
+2
-26
@@ -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
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
+18
-9
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user