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:
+53
-9
@@ -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')`)
|
// 将 waiting 和 running 状态改为 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
@@ -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
|
||||||
|
|||||||
@@ -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) {
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user