Files
yw1573 a0ef988bd5 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
2026-04-09 16:51:22 +08:00

142 lines
3.1 KiB
Go

package util
import (
"database/sql"
"errors"
"fmt"
"strings"
"sync"
)
func CreateLog(db *sql.DB, content string) error {
SqliteLock.Lock()
_, err := db.Exec(`INSERT INTO "log" ("content") VALUES (?)`, content)
SqliteLock.Unlock()
return err
}
func GetFields(db *sql.DB, names ...string) (map[string]string, error) {
if len(names) == 0 {
return nil, nil
}
placeholders := make([]string, len(names))
for i := 0; i < len(names); i++ {
placeholders[i] = "?"
}
query := fmt.Sprintf(`SELECT "name", "value" FROM "field" WHERE "name" IN (%s)`, strings.Join(placeholders, ","))
values := make([]interface{}, len(names))
for i := 0; i < len(names); i++ {
values[i] = names[i]
}
SqliteLock.Lock()
row, err := db.Query(query, values...)
SqliteLock.Unlock()
if err != nil {
return nil, err
}
defer row.Close()
var name, value string
fields := make(map[string]string)
for row.Next() {
if err := row.Scan(&name, &value); err != nil {
return nil, err
}
fields[name] = value
}
return fields, nil
}
func SaveFields(db *sql.DB, data [][2]string) error {
if len(data) == 0 {
return nil
}
tx, err := db.Begin()
if err != nil {
return err
}
defer func() {
if err != nil {
tx.Rollback()
} else {
tx.Commit()
}
}()
stmt, err := tx.Prepare(`INSERT OR REPLACE INTO "field" ("name", "value") VALUES (?, ?)`)
if err != nil {
return err
}
defer stmt.Close()
for _, d := range data {
SqliteLock.Lock()
_, err = stmt.Exec(d[0], d[1])
SqliteLock.Unlock()
if err != nil {
return err
}
}
return nil
}
// GetCurrentFolder 获取数据库中的下载保存路径,如果不存在则将默认路径保存到数据库
func GetCurrentFolder(db *sql.DB) (string, error) {
var folder string
SqliteLock.Lock()
err := db.QueryRow(`SELECT "value" FROM "field" WHERE "name" = 'download_folder'`).Scan(&folder)
SqliteLock.Unlock()
if err != nil && err == sql.ErrNoRows {
folder, err = GetDefaultDownloadFolder()
if err != nil {
return "", err
}
err = SaveDownloadFolder(db, folder)
if err != nil {
return "", err
}
return folder, nil
}
return folder, err
}
// SaveDownloadFolder 保存下载路径(不自动创建目录)
func SaveDownloadFolder(db *sql.DB, downloadFolder string) error {
SqliteLock.Lock()
_, err := db.Exec(`INSERT OR REPLACE INTO "field" ("name", "value") VALUES ('download_folder', ?)`, downloadFolder)
SqliteLock.Unlock()
return err
}
var SqliteLock sync.Mutex
// 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 {
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
}