Compare commits

...

56 Commits

Author SHA1 Message Date
github-actions[bot] e365250440 chore: bump version to 0.0.50 [skip ci] 2026-08-28 10:28:29 +00:00
truewhile 47d10e1f58 优化 2026-08-28 18:28:09 +08:00
github-actions[bot] e6473300a7 chore: bump version to 0.0.49 [skip ci] 2026-08-28 08:27:32 +00:00
truewhile 994f64f753 优化 2026-08-28 16:27:11 +08:00
github-actions[bot] 90064a5480 chore: bump version to 0.0.48 [skip ci] 2026-08-27 15:40:31 +00:00
truewhile 6e8eac9887 优化 2026-08-27 23:40:12 +08:00
github-actions[bot] d3051eaffe chore: bump version to 0.0.47 [skip ci] 2026-08-27 13:24:18 +00:00
truewhile 496a897782 优化 2026-08-27 21:24:01 +08:00
github-actions[bot] 22b7290ee1 chore: bump version to 0.0.46 [skip ci] 2026-08-27 07:42:16 +00:00
truewhile 14037d5dea 5 2026-08-27 15:42:00 +08:00
github-actions[bot] 152db3fb9f chore: bump version to 0.0.45 [skip ci] 2026-08-27 06:20:25 +00:00
truewhile 0413d123da 4 2026-08-27 14:20:05 +08:00
github-actions[bot] 5f6bd7b5cd chore: bump version to 0.0.44 [skip ci] 2026-08-27 05:41:04 +00:00
truewhile 3c25bb5d61 3 2026-08-27 13:40:46 +08:00
github-actions[bot] 659b91b000 chore: bump version to 0.0.43 [skip ci] 2026-08-27 03:34:05 +00:00
truewhile fe5b3bd56a 2 2026-08-27 11:33:50 +08:00
github-actions[bot] 82bbb116ae chore: bump version to 0.0.42 [skip ci] 2026-08-27 03:07:35 +00:00
truewhile 9f5ff7e6f0 1 2026-08-27 11:07:18 +08:00
github-actions[bot] a00504080a chore: bump version to 0.0.41 [skip ci] 2026-08-27 02:09:16 +00:00
truewhile fc6e2e6f10 优化 strm 同步记录:支持删除记录并展示上传数量统计 2026-08-27 10:08:54 +08:00
github-actions[bot] 65c5f3e4bf chore: bump version to 0.0.40 [skip ci] 2026-08-26 15:35:37 +00:00
truewhile 07e340251b 优化 2026-08-26 23:32:02 +08:00
github-actions[bot] 9b956b928b chore: bump version to 0.0.39 [skip ci] 2026-08-26 13:41:25 +00:00
truewhile c3187f6e3f 优化上传逻辑
优化上传逻辑
2026-08-26 21:41:06 +08:00
github-actions[bot] 60c815a8b3 chore: bump version to 0.0.38 [skip ci] 2026-08-26 09:06:04 +00:00
truewhile 41b155ea31 Merge branch 'main' of https://github.com/truewhile/MMTL 2026-08-26 17:05:45 +08:00
truewhile 2888ae8bf7 8 2026-08-26 17:05:41 +08:00
github-actions[bot] 7363064d89 chore: bump version to 0.0.37 [skip ci] 2026-08-26 08:16:56 +00:00
truewhile 1d53bf2ae1 7 2026-08-26 16:16:39 +08:00
github-actions[bot] 618165ec31 chore: bump version to 0.0.36 [skip ci] 2026-08-26 07:51:14 +00:00
truewhile 87c66a9b8c 6 2026-08-26 15:50:59 +08:00
github-actions[bot] 1ea4724261 chore: bump version to 0.0.35 [skip ci] 2026-08-26 06:50:29 +00:00
truewhile 3d372f039e Merge branch 'main' of https://github.com/truewhile/MMTL 2026-08-26 14:50:10 +08:00
truewhile 0384017e98 6 2026-08-26 14:50:06 +08:00
github-actions[bot] 0332579d5f chore: bump version to 0.0.34 [skip ci] 2026-08-26 04:52:28 +00:00
truewhile 6aefe18caa 5 2026-08-26 12:52:10 +08:00
github-actions[bot] ef72fc8d83 chore: bump version to 0.0.33 [skip ci] 2026-08-26 04:12:13 +00:00
truewhile 4764c09572 4 2026-08-26 12:11:56 +08:00
github-actions[bot] 98ca766a37 chore: bump version to 0.0.32 [skip ci] 2026-08-26 03:38:14 +00:00
truewhile 13c9035b76 3 2026-08-26 11:37:58 +08:00
github-actions[bot] ad6d0ba21d chore: bump version to 0.0.31 [skip ci] 2026-08-26 03:19:09 +00:00
truewhile 431f7f088b 2 2026-08-26 11:18:53 +08:00
github-actions[bot] 3f13ed1113 chore: bump version to 0.0.30 [skip ci] 2026-08-26 01:51:54 +00:00
truewhile 9d359c40dd 1 2026-08-26 09:51:27 +08:00
github-actions[bot] 0e7dbd6215 chore: bump version to 0.0.29 [skip ci] 2026-08-26 00:43:54 +00:00
truewhile 7fd8de91cb Merge branch 'main' of https://github.com/truewhile/MMTL 2026-08-26 08:43:40 +08:00
truewhile 5f323eb2ce 优化
yo 优化
2026-08-26 08:43:36 +08:00
github-actions[bot] c0ac8bf11a chore: bump version to 0.0.28 [skip ci] 2026-08-25 16:26:45 +00:00
truewhile b676733af7 优化strm同步
优化strm同步
2026-08-26 00:26:28 +08:00
github-actions[bot] 7a2027a3a7 chore: bump version to 0.0.27 [skip ci] 2026-08-25 15:19:28 +00:00
truewhile 9ffb74adce 优化
优化
2026-08-25 23:19:07 +08:00
github-actions[bot] 13faff7078 chore: bump version to 0.0.26 [skip ci] 2026-08-25 14:55:08 +00:00
truewhile a8a4e88d86 Merge branch 'main' of https://github.com/truewhile/MMTL 2026-08-25 22:54:50 +08:00
truewhile 8fa5db88ff 优化
优化
2026-08-25 22:54:46 +08:00
github-actions[bot] 3c325f81c8 chore: bump version to 0.0.25 [skip ci] 2026-08-25 14:12:30 +00:00
truewhile 585434010c 优化
优化
2026-08-25 22:12:05 +08:00
111 changed files with 8363 additions and 1151 deletions
+117
View File
@@ -27,6 +27,9 @@ permissions:
jobs:
version-and-publish:
runs-on: ubuntu-latest
outputs:
new_version: ${{ steps.bump_version.outputs.new_version }}
tag: ${{ steps.bump_version.outputs.tag }}
steps:
- uses: actions/checkout@v4
with:
@@ -149,3 +152,117 @@ jobs:
VERSION=${{ steps.bump_version.outputs.new_version }}
cache-from: type=gha
cache-to: type=gha,mode=max
# 单文件可执行构建:把前端打包进二进制(go:embed),交叉编译 Windows /
# Linux / macOS 的 amd64 / arm64 产物,作为 GitHub Release 附件发布。
build-frontend:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- uses: actions/setup-node@v4
with:
node-version: '20'
cache: 'npm'
cache-dependency-path: web/package-lock.json
- name: Install
working-directory: web
run: npm ci
- name: Build SPA
working-directory: web
run: npm run build
- name: Upload web/dist
uses: actions/upload-artifact@v4
with:
name: web-dist
path: web/dist
retention-days: 1
# 先创建(幂等)空的 GitHub Release,供后续 build-binaries 并行上传附件,
# 也避免矩阵各 job 并发 upload 时 release 尚不存在而互相竞争。
publish-create-release:
needs: [version-and-publish]
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- uses: actions/checkout@v4
- name: Create release
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
RELEASE_TAG: ${{ needs.version-and-publish.outputs.tag }}
run: |
set -eux
# tag 已由 version-and-publish 推送;若 release 已存在则忽略(--verify-tag 幂等)
gh release create "$RELEASE_TAG" \
--title "MMTL ${{ needs.version-and-publish.outputs.new_version }}" \
--notes "自动化发布 ${{ needs.version-and-publish.outputs.new_version }}" \
--verify-tag --latest || true
build-binaries:
needs: [version-and-publish, build-frontend, publish-create-release]
runs-on: ubuntu-latest
permissions:
contents: write
strategy:
fail-fast: false
matrix:
include:
- goos: linux
goarch: amd64
ext: ""
- goos: linux
goarch: arm64
ext: ""
- goos: windows
goarch: amd64
ext: .exe
- goos: windows
goarch: arm64
ext: .exe
- goos: darwin
goarch: amd64
ext: ""
- goos: darwin
goarch: arm64
ext: ""
steps:
- uses: actions/checkout@v4
- uses: actions/setup-go@v5
with:
go-version: '1.25'
cache: true
- name: Download web/dist
uses: actions/download-artifact@v4
with:
name: web-dist
path: web/dist
- name: Build binary
run: |
CGO_ENABLED=0 GOOS=${{ matrix.goos }} GOARCH=${{ matrix.goarch }} \
go build -trimpath -ldflags="-s -w -X main.version=${{ needs.version-and-publish.outputs.tag }}" \
-o "dist/mmtl-${{ matrix.goos }}-${{ matrix.goarch }}${{ matrix.ext }}" ./cmd/server
- name: Package
run: |
mkdir -p package/mmtl
cp "dist/mmtl-${{ matrix.goos }}-${{ matrix.goarch }}${{ matrix.ext }}" package/mmtl/mmtl${{ matrix.ext }}
cp README.md package/mmtl/ 2>/dev/null || true
if [ "${{ matrix.goos }}" = "windows" ]; then
(cd package && zip -r "../mmtl_${{ matrix.goos }}_${{ matrix.goarch }}.zip" mmtl)
else
tar -czf "mmtl_${{ matrix.goos }}_${{ matrix.goarch }}.tar.gz" -C package mmtl
fi
- name: Upload to GitHub Release
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
RELEASE_TAG: ${{ needs.version-and-publish.outputs.tag }}
run: |
set -eux
PKG="mmtl_${{ matrix.goos }}_${{ matrix.goarch }}.zip"
TAR="mmtl_${{ matrix.goos }}_${{ matrix.goarch }}.tar.gz"
# 并发上传到同一 release 各自文件,--clobber 幂等覆盖
if [ -f "$PKG" ]; then
for i in 1 2 3; do gh release upload "$RELEASE_TAG" "$PKG" --clobber && break || sleep 5; done
fi
if [ -f "$TAR" ]; then
for i in 1 2 3; do gh release upload "$RELEASE_TAG" "$TAR" --clobber && break || sleep 5; done
fi
+13
View File
@@ -18,6 +18,19 @@ jobs:
go-version: '1.25'
cache: true
# The binary embeds the SPA (web/dist) via go:embed, so the dist must exist
# before the Go toolchain touches the `web` package.
- uses: actions/setup-node@v4
with:
node-version: '20'
cache: 'npm'
cache-dependency-path: web/package-lock.json
- name: Build SPA
working-directory: web
run: |
npm ci
npm run build
- name: go vet
run: go vet ./...
+22 -5
View File
@@ -360,15 +360,17 @@ docker compose -f docker-compose.search.yml up -d --no-deps mmtl
本地开发需要 Go、Node.js 和 npm。
后端会将 `web/dist` 通过 `go:embed` 编进二进制,因此**在编译 / 运行后端之前要先构建前端**,否则 `web` 包会因为缺少嵌入资源而编译失败。
```bash
# 前端依赖与构建(必须先做,产物被 go:embed 打进二进制)
npm --prefix web ci
npm --prefix web run build
# 后端测试
go test ./...
# 前端依赖与构建
npm --prefix web install
npm --prefix web run build
# 本地运行后端
# 本地运行后端(二进制自带前端界面,无需额外 web 目录)
go run ./cmd/server
# 本地运行前端开发服务器
@@ -387,6 +389,21 @@ http://127.0.0.1:3000
http://127.0.0.1:8080/api/health
```
### 交叉编译单文件发布物
CI(`.github/workflows/Auto-docker-publish.yml`)每次发布会自动为 Windows / Linux(含 Debian) / macOS 交叉编译 amd64 + arm64 的单文件可执行程序,并上传到对应的 GitHub Release。你可以在 Releases 页面下载 `.zip`(Windows)或 `.tar.gz`(Linux / macOS)附件,解压后直接运行其中的 `mmtl`(Windows 为 `mmtl.exe`),无需额外携带前端目录。
本地手动交叉编译某个平台:
```bash
# 先构建前端
npm --prefix web ci && npm --prefix web run build
# 例如:构建 Linux amd64 单文件
CGO_ENABLED=0 GOOS=linux GOARCH=amd64 \
go build -trimpath -ldflags="-s -w" -o mmtl-linux-amd64 ./cmd/server
```
## 贡献与反馈
提交 Bug、功能建议或 Pull Request 前,请先阅读 [贡献规范](CONTRIBUTING.md)。
+1 -1
View File
@@ -1 +1 @@
0.0.24
0.0.50
+3 -3
View File
@@ -52,7 +52,7 @@ func TestServeSPANoCachesIndexAndServesRoutes(t *testing.T) {
}
router := gin.New()
serveSPA(router, webDir)
serveSPA(router, os.DirFS(webDir))
for _, path := range []string{"/", "/login", "/library/e1c3507e-2878-40ae-a0e1-6b6e44b7fa7a", "/media/abc"} {
req := httptest.NewRequest(http.MethodGet, path, nil)
@@ -93,7 +93,7 @@ func TestServeSPAServesAssetsImmutableAndBypassesAPIRoutes(t *testing.T) {
}
router := gin.New()
serveSPA(router, webDir)
serveSPA(router, os.DirFS(webDir))
assetReq := httptest.NewRequest(http.MethodGet, "/assets/app.js", nil)
assetResp := httptest.NewRecorder()
@@ -155,7 +155,7 @@ func TestServeSPAServesAssetsImmutableAndBypassesAPIRoutes(t *testing.T) {
func TestServeSPAMissingIndexReportsExplicit404(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
serveSPA(router, t.TempDir())
serveSPA(router, os.DirFS(t.TempDir()))
req := httptest.NewRequest(http.MethodGet, "/", nil)
w := httptest.NewRecorder()
+59 -23
View File
@@ -1,6 +1,8 @@
package main
import (
"io/fs"
"mime"
"net/http"
"os"
"path/filepath"
@@ -13,6 +15,8 @@ import (
"github.com/ShukeBta/MMTL/internal/handler"
"github.com/ShukeBta/MMTL/internal/middleware"
"github.com/ShukeBta/MMTL/internal/service"
"github.com/ShukeBta/MMTL/web"
)
func buildRouter(cfg *config.Config, logger *zap.Logger, svc *service.Container) *gin.Engine {
@@ -29,31 +33,41 @@ func buildRouter(cfg *config.Config, logger *zap.Logger, svc *service.Container)
handler.Register(r, cfg, logger, svc)
if cfg.App.WebDir != "" {
serveSPA(r, cfg.App.WebDir)
// Prefer a directory on disk when configured explicitly (e.g. the Docker image
// mounts web/dist from the build stage, or an operator overrides app.web_dir
// with a custom skin). Otherwise fall back to the SPA embedded into the binary,
// which is what makes the cross-platform single-file artifacts work.
uiFS := webui.DistFS()
if dir := cfg.App.WebDir; dir != "" {
disk := os.DirFS(dir)
if index, err := fs.Stat(disk, "index.html"); err == nil && !index.IsDir() {
uiFS = disk
}
}
serveSPA(r, uiFS)
return r
}
// serveSPA serves the React build artifacts and falls back to index.html for
// non-API, non-asset paths so client-side routing keeps working.
func serveSPA(r *gin.Engine, webDir string) {
// non-API, non-asset paths so client-side routing keeps working. The UI tree
// comes from root, which is either the compiled-in SPA or an on-disk web dir.
func serveSPA(r *gin.Engine, root fs.FS) {
assets := r.Group("/assets")
assets.Use(func(c *gin.Context) {
c.Header("Cache-Control", "public, max-age=31536000, immutable")
c.Next()
})
assets.Static("/", filepath.Join(webDir, "assets"))
assets.GET("/*filepath", serveFSDir(root, "assets"))
brand := r.Group("/brand")
brand.Use(func(c *gin.Context) {
setNoCacheHeaders(c)
c.Next()
})
brand.Static("/", filepath.Join(webDir, "brand"))
brand.GET("/*filepath", serveFSDir(root, "brand"))
for _, rootFile := range []string{"/favicon.ico", "/favicon.svg", "/artwork-cache-sw.js"} {
filePath := filepath.Join(webDir, strings.TrimPrefix(rootFile, "/"))
r.GET(rootFile, serveNoCacheFile(filePath))
r.HEAD(rootFile, serveNoCacheFile(filePath))
name := strings.TrimPrefix(rootFile, "/")
r.GET(rootFile, serveFSFile(root, name))
r.HEAD(rootFile, serveFSFile(root, name))
}
r.NoRoute(func(c *gin.Context) {
path := c.Request.URL.Path
@@ -61,28 +75,50 @@ func serveSPA(r *gin.Engine, webDir string) {
c.Status(http.StatusNotFound)
return
}
serveSPAIndex(c, filepath.Join(webDir, "index.html"))
setNoCacheHeaders(c)
data, err := fs.ReadFile(root, "index.html")
if err != nil {
c.String(http.StatusNotFound, "MMTL web UI not found")
return
}
c.Data(http.StatusOK, "text/html; charset=utf-8", data)
})
}
func serveNoCacheFile(filePath string) gin.HandlerFunc {
// serveFSDir serves a static subdirectory of root. A missing asset returns 404.
func serveFSDir(root fs.FS, dir string) gin.HandlerFunc {
sub, err := fs.Sub(root, dir)
if err != nil {
return func(c *gin.Context) { c.Status(http.StatusNotFound) }
}
handler := http.StripPrefix("/"+dir, http.FileServerFS(sub))
return func(c *gin.Context) {
setNoCacheHeaders(c)
if _, err := os.Stat(filePath); err != nil {
c.Status(http.StatusNotFound)
return
}
c.File(filePath)
handler.ServeHTTP(c.Writer, c.Request)
}
}
func serveSPAIndex(c *gin.Context, indexPath string) {
setNoCacheHeaders(c)
if _, err := os.Stat(indexPath); err != nil {
c.String(http.StatusNotFound, "MMTL web UI not found: %s", indexPath)
return
// serveFSFile serves a single root-level file (favicon / service worker) with
// no-cache headers. It reads from root, which may be the embedded SPA or disk.
func serveFSFile(root fs.FS, name string) gin.HandlerFunc {
return func(c *gin.Context) {
setNoCacheHeaders(c)
data, err := fs.ReadFile(root, name)
if err != nil {
c.Status(http.StatusNotFound)
return
}
c.Data(http.StatusOK, mimeTypeByName(name), data)
}
}
// mimeTypeByName returns an HTTP content type guessed from a file extension.
func mimeTypeByName(name string) string {
switch mime.TypeByExtension(filepath.Ext(name)) {
case "":
return "application/octet-stream"
default:
return mime.TypeByExtension(filepath.Ext(name))
}
c.File(indexPath)
}
func setNoCacheHeaders(c *gin.Context) {
+2 -1
View File
@@ -3,6 +3,7 @@ module github.com/ShukeBta/MMTL
go 1.25.0
require (
github.com/aliyun/alibabacloud-oss-go-sdk-v2 v1.5.2
github.com/fsnotify/fsnotify v1.7.0
github.com/gin-gonic/gin v1.9.1
github.com/glebarez/sqlite v1.11.0
@@ -16,6 +17,7 @@ require (
go.uber.org/zap v1.27.0
golang.org/x/crypto v0.21.0
golang.org/x/sys v0.20.0
golang.org/x/time v0.15.0
gorm.io/driver/postgres v1.5.7
gorm.io/gorm v1.30.0
)
@@ -72,7 +74,6 @@ require (
golang.org/x/exp v0.0.0-20230905200255-921286631fa9 // indirect
golang.org/x/net v0.21.0 // indirect
golang.org/x/text v0.20.0 // indirect
golang.org/x/time v0.15.0 // indirect
google.golang.org/protobuf v1.31.0 // indirect
gopkg.in/ini.v1 v1.67.0 // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect
+2
View File
@@ -1,3 +1,5 @@
github.com/aliyun/alibabacloud-oss-go-sdk-v2 v1.5.2 h1:40yUSXwdkWN851BHCq6uiDhleh7A4+0yIBS+IUAqZVY=
github.com/aliyun/alibabacloud-oss-go-sdk-v2 v1.5.2/go.mod h1:FTzydeQVmR24FI0D6XWUOMKckjXehM/jgMn1xC+DA9M=
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
+2 -2
View File
@@ -3,8 +3,8 @@ package config
import "github.com/spf13/viper"
const (
defaultDatabaseMaxOpenConns = 4
defaultDatabaseMaxIdleConns = 2
defaultDatabaseMaxOpenConns = 16
defaultDatabaseMaxIdleConns = 4
defaultLicenseServerURL = "https://mgosever.3jzs.com"
defaultLicensePublicKey = "MCowBQYDK2VwAyEABRXnXy+urjrbKit6Yu/HiezWgP0NdsZW3tsegJWRrtI="
)
+41
View File
@@ -0,0 +1,41 @@
package config
import (
"fmt"
"os"
"gopkg.in/yaml.v3"
)
// SaveDatabaseConfig updates or creates config.yaml with the specified database configuration.
func SaveDatabaseConfig(dbType, dsn string) error {
configPath := "config.yaml"
data := make(map[string]any)
content, err := os.ReadFile(configPath)
if err == nil {
if err := yaml.Unmarshal(content, &data); err != nil {
data = make(map[string]any)
}
} else if !os.IsNotExist(err) {
return fmt.Errorf("read config.yaml: %w", err)
}
dbSection, ok := data["database"].(map[string]any)
if !ok {
dbSection = make(map[string]any)
}
dbSection["type"] = dbType
dbSection["dsn"] = dsn
data["database"] = dbSection
out, err := yaml.Marshal(data)
if err != nil {
return fmt.Errorf("marshal config.yaml: %w", err)
}
if err := os.WriteFile(configPath, out, 0644); err != nil {
return fmt.Errorf("write config.yaml: %w", err)
}
return nil
}
+36
View File
@@ -0,0 +1,36 @@
package config
import (
"os"
"path/filepath"
"testing"
)
func TestSaveDatabaseConfig(t *testing.T) {
dir := t.TempDir()
wd, _ := os.Getwd()
defer func() { _ = os.Chdir(wd) }()
if err := os.Chdir(dir); err != nil {
t.Fatalf("chdir: %v", err)
}
dsn := "postgres://admin:pass@127.0.0.1:5432/mmtl?sslmode=disable"
if err := SaveDatabaseConfig("postgres", dsn); err != nil {
t.Fatalf("SaveDatabaseConfig error: %v", err)
}
if _, err := os.Stat(filepath.Join(dir, "config.yaml")); err != nil {
t.Fatalf("expected config.yaml to exist: %v", err)
}
loaded, err := Load()
if err != nil {
t.Fatalf("Load error: %v", err)
}
if loaded.Database.Type != "postgres" {
t.Fatalf("expected database.type=postgres, got %s", loaded.Database.Type)
}
if loaded.Database.DSN != dsn {
t.Fatalf("expected dsn=%s, got %s", dsn, loaded.Database.DSN)
}
}
+223
View File
@@ -0,0 +1,223 @@
package database
import (
"context"
"fmt"
"net/url"
"strings"
"time"
"go.uber.org/zap"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"github.com/ShukeBta/MMTL/internal/config"
"github.com/ShukeBta/MMTL/internal/model"
)
// DatabaseStatus describes the currently active database engine and runtime metrics.
type DatabaseStatus struct {
Type string `json:"type"`
DSN string `json:"dsn,omitempty"`
DBPath string `json:"db_path,omitempty"`
OpenConns int `json:"open_conns"`
InUse int `json:"in_use"`
Idle int `json:"idle"`
MaxOpenConns int `json:"max_open_conns"`
TableCounts map[string]int64 `json:"table_counts"`
}
// PostgresTestResult returns latency and version info after testing connection.
type PostgresTestResult struct {
Success bool `json:"success"`
LatencyMS int64 `json:"latency_ms"`
Version string `json:"version,omitempty"`
Message string `json:"message,omitempty"`
Error string `json:"error,omitempty"`
}
// DatabaseMigrationResult returns row counts and execution duration of migration.
type DatabaseMigrationResult struct {
Success bool `json:"success"`
TotalRows int64 `json:"total_rows"`
TableRows map[string]int64 `json:"table_rows"`
DurationMS int64 `json:"duration_ms"`
Message string `json:"message,omitempty"`
Error string `json:"error,omitempty"`
}
// InspectDatabaseStatus queries the currently active database for metrics and table rows.
func InspectDatabaseStatus(db *gorm.DB, cfg *config.Config) *DatabaseStatus {
st := &DatabaseStatus{
Type: "sqlite",
TableCounts: make(map[string]int64),
}
if cfg != nil {
st.DBPath = cfg.Database.DBPath
if cfg.Database.Type == "postgres" || (cfg.Database.Type == "auto" && strings.TrimSpace(cfg.Database.DSN) != "") {
st.Type = "postgres"
st.DSN = MaskDSN(cfg.Database.DSN)
}
}
if isPostgres(db) {
st.Type = "postgres"
}
if db != nil {
if sqlDB, err := db.DB(); err == nil {
stats := sqlDB.Stats()
st.OpenConns = stats.OpenConnections
st.InUse = stats.InUse
st.Idle = stats.Idle
st.MaxOpenConns = stats.MaxOpenConnections
}
// Count rows for major model tables
for _, m := range model.AllModels() {
if tbl, err := modelTableName(db, m); err == nil {
if db.Migrator().HasTable(tbl) {
var count int64
if err := db.Raw("SELECT COUNT(1) FROM " + quoteIdent(tbl)).Scan(&count).Error; err == nil {
st.TableCounts[tbl] = count
}
}
}
}
}
return st
}
// TestPostgres establishes a temporary connection to verify reachability and permissions.
func TestPostgres(dsn string) (*PostgresTestResult, error) {
dsn = strings.TrimSpace(dsn)
if dsn == "" {
return &PostgresTestResult{
Success: false,
Error: "PostgreSQL DSN 不能为空",
}, nil
}
start := time.Now()
testDB, err := gorm.Open(postgres.Open(dsn), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
return &PostgresTestResult{
Success: false,
Error: fmt.Sprintf("连接失败: %v", err),
}, nil
}
sqlDB, err := testDB.DB()
if err != nil {
return &PostgresTestResult{
Success: false,
Error: fmt.Sprintf("获取底层连接失败: %v", err),
}, nil
}
defer sqlDB.Close()
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := sqlDB.PingContext(ctx); err != nil {
return &PostgresTestResult{
Success: false,
Error: fmt.Sprintf("Ping 超时或失败: %v", err),
}, nil
}
var version string
if err := testDB.WithContext(ctx).Raw("SELECT version()").Scan(&version).Error; err != nil {
version = "PostgreSQL (unknown version)"
}
latency := time.Since(start).Milliseconds()
return &PostgresTestResult{
Success: true,
LatencyMS: latency,
Version: version,
Message: "连接成功",
}, nil
}
// MigrateCurrentToPostgres performs schema initialization and full table data copy into target PostgreSQL.
func MigrateCurrentToPostgres(src *gorm.DB, targetDSN string, batchSize int, log *zap.Logger) (*DatabaseMigrationResult, error) {
targetDSN = strings.TrimSpace(targetDSN)
if targetDSN == "" {
return nil, fmt.Errorf("target PostgreSQL DSN cannot be empty")
}
if src == nil {
return nil, fmt.Errorf("current database is not available")
}
started := time.Now()
targetDB, err := gorm.Open(postgres.Open(targetDSN), &gorm.Config{
Logger: newGormLogger(log),
})
if err != nil {
return nil, fmt.Errorf("open target PostgreSQL: %w", err)
}
targetSQLDB, err := targetDB.DB()
if err == nil {
defer targetSQLDB.Close()
}
// 1. 初始化目标库 Schema、类型与索引
if err := AutoMigrate(targetDB); err != nil {
return nil, fmt.Errorf("auto migrate target PostgreSQL: %w", err)
}
// 2. 安全重置目标数据库的初始默认数据
if err := resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, targetDB, log); err != nil {
return nil, fmt.Errorf("reset target bootstrap data: %w", err)
}
// 3. 执行数据批量复制
tableRows, totalRows, err := copyModelTables(src, targetDB, batchSize)
if err != nil {
return nil, fmt.Errorf("copy tables: %w", err)
}
// 4. 标记迁移完成
if err := markSQLiteMigrationComplete(targetDB); err != nil {
return nil, fmt.Errorf("mark migration complete: %w", err)
}
duration := time.Since(started).Milliseconds()
return &DatabaseMigrationResult{
Success: true,
TotalRows: totalRows,
TableRows: tableRows,
DurationMS: duration,
Message: fmt.Sprintf("成功迁移 %d 条记录至 PostgreSQL", totalRows),
}, nil
}
// MaskDSN masks the password in a connection string for safe API responses.
func MaskDSN(rawDSN string) string {
rawDSN = strings.TrimSpace(rawDSN)
if rawDSN == "" {
return ""
}
if u, err := url.Parse(rawDSN); err == nil && u.User != nil {
if pass, hasPassword := u.User.Password(); hasPassword && pass != "" {
rawUserPass := u.User.String()
user := u.User.Username()
maskedUserPass := user + ":******"
return strings.Replace(rawDSN, rawUserPass+"@", maskedUserPass+"@", 1)
}
}
// Fallback for keyword-style DSN (e.g. host=... password=...)
if strings.Contains(rawDSN, "password=") {
parts := strings.Fields(rawDSN)
for i, p := range parts {
if strings.HasPrefix(p, "password=") {
parts[i] = "password=******"
}
}
return strings.Join(parts, " ")
}
return rawDSN
}
+71
View File
@@ -0,0 +1,71 @@
package database
import (
"testing"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
"github.com/ShukeBta/MMTL/internal/config"
"github.com/ShukeBta/MMTL/internal/model"
)
func TestMaskDSN(t *testing.T) {
cases := []struct {
in string
want string
}{
{
in: "postgres://admin:secret123@localhost:5432/mmtl?sslmode=disable",
want: "postgres://admin:******@localhost:5432/mmtl?sslmode=disable",
},
{
in: "host=localhost port=5432 user=admin password=secret dbname=mmtl sslmode=disable",
want: "host=localhost port=5432 user=admin password=****** dbname=mmtl sslmode=disable",
},
{
in: "sqlite://data/mmtl.db",
want: "sqlite://data/mmtl.db",
},
{
in: "",
want: "",
},
}
for _, c := range cases {
got := MaskDSN(c.in)
if got != c.want {
t.Errorf("MaskDSN(%q) = %q, want %q", c.in, got, c.want)
}
}
}
func TestInspectDatabaseStatus(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.User{}, &model.Media{}); err != nil {
t.Fatal(err)
}
_ = db.Create(&model.User{Username: "testuser", PasswordHash: "h", Role: "user"}).Error
cfg := &config.Config{}
cfg.Database.Type = "sqlite"
cfg.Database.DBPath = "./data/mmtl.db"
st := InspectDatabaseStatus(db, cfg)
if st == nil {
t.Fatal("expected non-nil DatabaseStatus")
}
if st.Type != "sqlite" {
t.Fatalf("expected sqlite, got %s", st.Type)
}
if st.DBPath != "./data/mmtl.db" {
t.Fatalf("expected db_path, got %s", st.DBPath)
}
if st.TableCounts["users"] != 1 {
t.Fatalf("expected 1 user, got %d", st.TableCounts["users"])
}
}
+11 -11
View File
@@ -159,7 +159,7 @@ func TestCopyModelTablesMigratesExistingSQLiteRows(t *testing.T) {
t.Fatal(err)
}
copied, err := copyModelTables(src, dst, 2)
_, copied, err := copyModelTables(src, dst, 2)
if err != nil {
t.Fatal(err)
}
@@ -222,7 +222,7 @@ func TestCopyModelTablesResumesPartialSQLiteMigration(t *testing.T) {
t.Fatal(err)
}
copied, err := copyModelTables(src, dst, 2)
_, copied, err := copyModelTables(src, dst, 2)
if err != nil {
t.Fatal(err)
}
@@ -240,7 +240,7 @@ func TestCopyModelTablesResumesPartialSQLiteMigration(t *testing.T) {
t.Fatalf("genres = %q, want %q", got.Genres, media.Genres)
}
copied, err = copyModelTables(src, dst, 2)
_, copied, err = copyModelTables(src, dst, 2)
if err != nil {
t.Fatal(err)
}
@@ -332,7 +332,7 @@ func TestSQLiteMigrationFallsBackToDataDirDefaultPath(t *testing.T) {
if err := resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src2, dst, nil); err != nil {
t.Fatal(err)
}
copied, err := copyModelTables(src2, dst, 2)
_, copied, err := copyModelTables(src2, dst, 2)
if err != nil {
t.Fatal(err)
}
@@ -408,13 +408,13 @@ func TestOpenSQLiteMigrationSourceUsesFallbackSourcePath(t *testing.T) {
_ = sqlDB2.Close()
}
}()
copied, err := copyModelTables(src2, dst, 2)
if err != nil {
t.Fatal(err)
}
if copied != 2 {
t.Fatalf("copied rows = %d, want 2", copied)
}
_, copied, err := copyModelTables(src2, dst, 2)
if err != nil {
t.Fatal(err)
}
if copied != 2 {
t.Fatalf("copied rows = %d, want 2", copied)
}
var userCount int64
if err := dst.Model(&model.User{}).Where("username = ?", "real-admin").Count(&userCount).Error; err != nil {
t.Fatal(err)
+1 -1
View File
@@ -48,7 +48,7 @@ func MigrateSQLiteToCurrentIfNeeded(cfg *config.Config, target *gorm.DB, log *za
if err := resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, target, log); err != nil {
return err
}
copied, err := copyModelTables(src, target, 500)
_, copied, err := copyModelTables(src, target, 500)
if err != nil {
return err
}
+16 -13
View File
@@ -13,52 +13,53 @@ import (
"github.com/ShukeBta/MMTL/internal/model"
)
func copyModelTables(src, target *gorm.DB, batchSize int) (int64, error) {
func copyModelTables(src, target *gorm.DB, batchSize int) (map[string]int64, int64, error) {
if batchSize <= 0 {
batchSize = 500
}
var copied int64
tableCounts := make(map[string]int64)
var totalCopied int64
for _, m := range model.AllModels() {
table, err := modelTableName(src, m)
if err != nil {
return copied, err
return tableCounts, totalCopied, err
}
primaryColumns, err := modelPrimaryColumns(src, m)
if err != nil {
return copied, fmt.Errorf("inspect model %T primary keys: %w", m, err)
return tableCounts, totalCopied, fmt.Errorf("inspect model %T primary keys: %w", m, err)
}
exists, err := sqliteTableExists(src, table)
if err != nil {
return copied, err
return tableCounts, totalCopied, err
}
if !exists {
continue
}
var sourceCount int64
if err := src.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&sourceCount).Error; err != nil {
return copied, fmt.Errorf("count sqlite table %s: %w", table, err)
return tableCounts, totalCopied, fmt.Errorf("count sqlite table %s: %w", table, err)
}
if sourceCount == 0 {
continue
}
var targetCount int64
if err := target.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&targetCount).Error; err != nil {
return copied, fmt.Errorf("count target table %s: %w", table, err)
return tableCounts, totalCopied, fmt.Errorf("count target table %s: %w", table, err)
}
modelType := reflect.TypeOf(m)
if modelType.Kind() != reflect.Ptr {
return copied, fmt.Errorf("model %T is not a pointer", m)
return tableCounts, totalCopied, fmt.Errorf("model %T is not a pointer", m)
}
sliceType := reflect.SliceOf(modelType.Elem())
slicePtr := reflect.New(sliceType)
if err := src.Unscoped().Find(slicePtr.Interface()).Error; err != nil {
return copied, fmt.Errorf("read sqlite table %s: %w", table, err)
return tableCounts, totalCopied, fmt.Errorf("read sqlite table %s: %w", table, err)
}
filtered := slicePtr.Elem()
if targetCount > 0 {
primaryKeySet, err := targetPrimaryKeySet(target, table, primaryColumns)
if err != nil {
return copied, err
return tableCounts, totalCopied, err
}
filtered = filterRowsMissingInTarget(target, table, primaryColumns, filtered, primaryKeySet)
}
@@ -68,11 +69,13 @@ func copyModelTables(src, target *gorm.DB, batchSize int) (int64, error) {
filteredPtr := reflect.New(filtered.Type())
filteredPtr.Elem().Set(filtered)
if err := target.Clauses(clause.OnConflict{DoNothing: true}).CreateInBatches(filteredPtr.Interface(), batchSize).Error; err != nil {
return copied, fmt.Errorf("copy sqlite table %s: %w", table, err)
return tableCounts, totalCopied, fmt.Errorf("copy sqlite table %s: %w", table, err)
}
copied += int64(filtered.Len())
copiedForTable := int64(filtered.Len())
tableCounts[table] = copiedForTable
totalCopied += copiedForTable
}
return copied, nil
return tableCounts, totalCopied, nil
}
func modelPrimaryColumns(db *gorm.DB, m any) ([]string, error) {
+25 -2
View File
@@ -4,6 +4,7 @@ import (
"context"
"fmt"
"path/filepath"
"strings"
"gorm.io/gorm"
@@ -32,16 +33,37 @@ func installSQLiteWriteGate(db *gorm.DB) {
gate.Unlock()
}
}
rawLock := func(tx *gorm.DB) {
if tx.Statement != nil && isReadOnlySQL(tx.Statement.SQL.String()) {
return
}
lock(tx)
}
_ = db.Callback().Create().Before("gorm:create").Register("mmtl:sqlite_write_lock", lock)
_ = db.Callback().Create().After("gorm:create").Register("mmtl:sqlite_write_unlock", unlock)
_ = db.Callback().Update().Before("gorm:update").Register("mmtl:sqlite_write_lock", lock)
_ = db.Callback().Update().After("gorm:update").Register("mmtl:sqlite_write_unlock", unlock)
_ = db.Callback().Delete().Before("gorm:delete").Register("mmtl:sqlite_write_lock", lock)
_ = db.Callback().Delete().After("gorm:delete").Register("mmtl:sqlite_write_unlock", unlock)
_ = db.Callback().Raw().Before("gorm:raw").Register("mmtl:sqlite_write_lock", lock)
_ = db.Callback().Raw().Before("gorm:raw").Register("mmtl:sqlite_write_lock", rawLock)
_ = db.Callback().Raw().After("gorm:raw").Register("mmtl:sqlite_write_unlock", unlock)
}
func isReadOnlySQL(sql string) bool {
trimmed := strings.TrimSpace(sql)
if len(trimmed) == 0 {
return false
}
upper := strings.ToUpper(trimmed)
if strings.HasPrefix(upper, "SELECT") || strings.HasPrefix(upper, "EXPLAIN") {
return true
}
if strings.HasPrefix(upper, "WITH") && !strings.Contains(upper, "INSERT") && !strings.Contains(upper, "UPDATE") && !strings.Contains(upper, "DELETE") {
return true
}
return false
}
// sqliteWriteGate serializes in-process SQLite writes while respecting the
// statement context, so request cancellation can break out of a queued write.
type sqliteWriteGate struct {
@@ -84,7 +106,7 @@ func buildSQLiteDSN(cfg *config.Config) string {
}
dsn := dbPath + "?_pragma=foreign_keys(1)"
if cfg.Database.WALMode {
dsn += "&_pragma=journal_mode(WAL)"
dsn += "&_pragma=journal_mode(WAL)&_pragma=synchronous(NORMAL)"
}
if cfg.Database.BusyTimeout > 0 {
dsn += fmt.Sprintf("&_pragma=busy_timeout(%d)", cfg.Database.BusyTimeout)
@@ -92,6 +114,7 @@ func buildSQLiteDSN(cfg *config.Config) string {
if cfg.Database.CacheSize != 0 {
dsn += fmt.Sprintf("&_pragma=cache_size(%d)", cfg.Database.CacheSize)
}
dsn += "&_pragma=temp_store(MEMORY)&_pragma=mmap_size(268435456)"
return dsn
}
+145
View File
@@ -0,0 +1,145 @@
package handler
import (
"fmt"
"net/http"
"net/url"
"strings"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MMTL/internal/service"
)
type DatabaseConnectionPayload struct {
Type string `json:"type"`
DSN string `json:"dsn"`
Host string `json:"host"`
Port int `json:"port"`
User string `json:"user"`
Password string `json:"password"`
DBName string `json:"dbname"`
SSLMode string `json:"sslmode"`
}
func (p *DatabaseConnectionPayload) BuildDSN() string {
raw := strings.TrimSpace(p.DSN)
if raw != "" {
return raw
}
host := strings.TrimSpace(p.Host)
if host == "" {
return ""
}
port := p.Port
if port <= 0 {
port = 5432
}
user := strings.TrimSpace(p.User)
dbname := strings.TrimSpace(p.DBName)
if dbname == "" {
dbname = "mmtl"
}
sslmode := strings.TrimSpace(p.SSLMode)
if sslmode == "" {
sslmode = "disable"
}
userInfo := url.User(user)
if p.Password != "" {
userInfo = url.UserPassword(user, p.Password)
}
u := url.URL{
Scheme: "postgres",
User: userInfo,
Host: fmt.Sprintf("%s:%d", host, port),
Path: "/" + dbname,
RawQuery: "sslmode=" + url.QueryEscape(sslmode),
}
return u.String()
}
func getDatabaseStatusHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if svc.Database == nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "database service unavailable"})
return
}
status := svc.Database.GetStatus(c.Request.Context())
c.JSON(http.StatusOK, status)
}
}
func testDatabaseHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req DatabaseConnectionPayload
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid payload: " + err.Error()})
return
}
dsn := req.BuildDSN()
if dsn == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "请提供有效的 PostgreSQL 连接信息或 DSN"})
return
}
res, err := svc.Database.TestPostgres(c.Request.Context(), dsn)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, res)
}
}
func migrateDatabaseHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req DatabaseConnectionPayload
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid payload: " + err.Error()})
return
}
dsn := req.BuildDSN()
if dsn == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "请提供目标 PostgreSQL 连接信息或 DSN"})
return
}
res, err := svc.Database.MigrateToPostgres(c.Request.Context(), dsn)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "迁移失败: " + err.Error()})
return
}
c.JSON(http.StatusOK, res)
}
}
func saveDatabaseConfigHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req DatabaseConnectionPayload
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid payload: " + err.Error()})
return
}
dbType := strings.ToLower(strings.TrimSpace(req.Type))
if dbType == "" {
dbType = "postgres"
}
var dsn string
if dbType == "postgres" {
dsn = req.BuildDSN()
if dsn == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "请提供有效的 PostgreSQL 连接信息或 DSN"})
return
}
}
if err := svc.Database.SaveConfig(c.Request.Context(), dbType, dsn); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{
"message": "数据库配置已成功保存至配置文件,重启服务后将以新数据库运行",
"type": dbType,
})
}
}
+101
View File
@@ -0,0 +1,101 @@
package handler
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MMTL/internal/config"
"github.com/ShukeBta/MMTL/internal/service"
)
func TestBuildDSN(t *testing.T) {
cases := []struct {
payload DatabaseConnectionPayload
want string
}{
{
payload: DatabaseConnectionPayload{
DSN: "postgres://myuser:mypass@10.0.0.1:5432/mydb?sslmode=require",
},
want: "postgres://myuser:mypass@10.0.0.1:5432/mydb?sslmode=require",
},
{
payload: DatabaseConnectionPayload{
Host: "127.0.0.1",
Port: 5432,
User: "postgres",
Password: "secretpassword",
DBName: "mmtl_prod",
SSLMode: "disable",
},
want: "postgres://postgres:secretpassword@127.0.0.1:5432/mmtl_prod?sslmode=disable",
},
}
for _, c := range cases {
got := c.payload.BuildDSN()
if got != c.want {
t.Errorf("BuildDSN() = %q, want %q", got, c.want)
}
}
}
func TestGetDatabaseStatusHandler(t *testing.T) {
gin.SetMode(gin.TestMode)
cfg := &config.Config{}
cfg.Database.Type = "sqlite"
cfg.Database.DBPath = "./data/mmtl.db"
svc := &service.Container{
Database: service.NewDatabaseAdminService(cfg, nil, nil, nil),
}
r := gin.New()
r.GET("/api/admin/database/status", getDatabaseStatusHandler(svc))
req := httptest.NewRequest(http.MethodGet, "/api/admin/database/status", nil)
rec := httptest.NewRecorder()
r.ServeHTTP(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d: %s", rec.Code, rec.Body.String())
}
var resp map[string]any
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
t.Fatalf("unmarshal response: %v", err)
}
if resp["type"] != "sqlite" {
t.Fatalf("expected type=sqlite, got %v", resp["type"])
}
}
func TestSaveDatabaseConfigHandler(t *testing.T) {
gin.SetMode(gin.TestMode)
dir := t.TempDir()
cfg := &config.Config{}
cfg.App.DataDir = dir
cfg.Database.Type = "sqlite"
svc := &service.Container{
Database: service.NewDatabaseAdminService(cfg, nil, nil, nil),
}
r := gin.New()
r.POST("/api/admin/database/save-config", saveDatabaseConfigHandler(svc))
body := bytes.NewBufferString(`{"type":"postgres","host":"localhost","port":5432,"user":"admin","password":"pwd","dbname":"mmtl"}`)
req := httptest.NewRequest(http.MethodPost, "/api/admin/database/save-config", body)
req.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
r.ServeHTTP(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d: %s", rec.Code, rec.Body.String())
}
}
+73 -14
View File
@@ -16,12 +16,13 @@ import (
)
type createLibraryReq struct {
Name string `json:"name" binding:"required"`
Path string `json:"path"`
Paths []string `json:"paths"`
Roots []service.LibraryRootInput `json:"roots"`
Type string `json:"type"`
CoverURL string `json:"cover_url"`
Name string `json:"name"`
Path string `json:"path"`
Paths []string `json:"paths"`
Roots []service.LibraryRootInput `json:"roots"`
Type string `json:"type"`
CoverURL string `json:"cover_url"`
CreatePerSubfolder bool `json:"create_per_subfolder"`
}
func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
@@ -32,7 +33,7 @@ func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
return
}
role, _ := c.Get(middleware.CtxUserRole)
includeHidden := role == "admin" && (c.Query("include_hidden") == "1" || c.Query("all") == "1")
includeHidden := role == "admin" && (c.Query("include_hidden") == "1" || c.Query("include_hidden") == "true" || c.Query("all") == "1")
if !includeHidden {
libs = service.FilterDisplayCloudLibraries(c.Request.Context(), svc.Repo, libs)
visibility := mediaVisibilityForRequest(c, svc)
@@ -60,7 +61,7 @@ func getLibraryHandler(svc *service.Container) gin.HandlerFunc {
return
}
role, _ := c.Get(middleware.CtxUserRole)
includeHidden := role == "admin" && (c.Query("include_hidden") == "1" || c.Query("all") == "1")
includeHidden := role == "admin" && (c.Query("include_hidden") == "1" || c.Query("include_hidden") == "true" || c.Query("all") == "1")
if !includeHidden {
libs := service.FilterDisplayCloudLibraries(c.Request.Context(), svc.Repo, []model.Library{*lib})
if len(libs) == 0 || !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, libs[0], mediaVisibilityForRequest(c, svc)) {
@@ -88,9 +89,38 @@ func createLibraryHandler(svc *service.Container) gin.HandlerFunc {
}
}
if len(roots) == 0 && strings.TrimSpace(req.Path) != "" {
roots = append(roots, service.LibraryRootInput{Path: req.Path})
roots = append(roots, service.LibraryRootInput{Path: req.Path})
}
var l *model.Library
if req.CreatePerSubfolder {
parent := ""
if len(roots) > 0 {
parent = roots[0].Path
} else if strings.TrimSpace(req.Path) != "" {
parent = req.Path
}
l, err := svc.Media.CreateLibraryWithRootsAndCover(c.Request.Context(), req.Name, req.Type, req.CoverURL, roots)
created, err := svc.Media.CreateLibrariesPerSubfolder(c.Request.Context(), parent, req.Type, req.CoverURL)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
uid, _ := c.Get("ctx_user_id")
for i := range created {
lib := &created[i]
svc.Audit.Record(c.Request.Context(), toString(uid), "library.create", lib.ID, c.ClientIP(), lib.Path)
if svc.Watcher != nil {
go func() { _ = svc.Watcher.Refresh(context.Background()) }()
}
for _, root := range lib.Roots {
if root.Enabled {
queueLibraryRootScan(svc, lib.ID, root.ID)
}
}
}
c.JSON(http.StatusCreated, gin.H{"libraries": created})
return
}
l, err := svc.Media.CreateLibraryWithRootsAndCover(c.Request.Context(), req.Name, req.Type, req.CoverURL, roots)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -117,7 +147,9 @@ func createLibraryHandler(svc *service.Container) gin.HandlerFunc {
}
type updateLibraryReq struct {
CoverURL string `json:"cover_url"`
CoverURL *string `json:"cover_url"`
SortOrder *int `json:"sort_order"`
CarouselEnabled *bool `json:"carousel_enabled"`
}
func updateLibraryHandler(svc *service.Container) gin.HandlerFunc {
@@ -127,9 +159,17 @@ func updateLibraryHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if err := svc.Media.UpdateLibraryCover(c.Request.Context(), c.Param("id"), req.CoverURL); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
if req.CoverURL != nil {
if err := svc.Media.UpdateLibraryCover(c.Request.Context(), c.Param("id"), *req.CoverURL); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
}
if req.SortOrder != nil || req.CarouselEnabled != nil {
if err := svc.Media.UpdateLibraryFields(c.Request.Context(), c.Param("id"), req.SortOrder, req.CarouselEnabled); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
}
lib, err := svc.Repo.Library.FindByID(c.Request.Context(), c.Param("id"))
if err != nil || lib == nil {
@@ -140,6 +180,25 @@ func updateLibraryHandler(svc *service.Container) gin.HandlerFunc {
}
}
type reorderLibrariesReq struct {
IDs []string `json:"ids" binding:"required"`
}
func reorderLibrariesHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req reorderLibrariesReq
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if err := svc.Media.ReorderLibraries(c.Request.Context(), req.IDs); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"updated": len(req.IDs)})
}
}
func deleteLibraryHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
id := c.Param("id")
+35
View File
@@ -20,6 +20,25 @@ func registerAdminRoutes(api *gin.RouterGroup, cfg *config.Config, svc *service.
registerAdminAPIConfigRoutes(admin, svc)
registerAdminRecognitionWordRoutes(admin, svc)
registerAdminStrmRoutes(admin, svc)
registerAdminScraperRoutes(admin, svc)
registerAdminDatabaseRoutes(admin, svc)
}
func registerAdminScraperRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.GET("/scraper/queue", listScrapeQueueHandler(svc))
admin.POST("/scraper/queue/:id/cancel", cancelScrapeTaskHandler(svc))
admin.POST("/scraper/queue/:id/retry", retryScrapeTaskHandler(svc))
admin.DELETE("/scraper/queue/:id", deleteScrapeTaskHandler(svc))
admin.POST("/scraper/queue/batch", batchActionScrapeTasksHandler(svc))
admin.POST("/scraper/queue/clear-done", clearDoneScrapeTasksHandler(svc))
admin.POST("/scraper/queue/clear-finished", clearFinishedScrapeTasksHandler(svc))
admin.POST("/scraper/queue/clear-canceled", clearCanceledScrapeTasksHandler(svc))
admin.POST("/scraper/queue/retry-failed", retryAllFailedScrapeTasksHandler(svc))
admin.POST("/scraper/queue/cancel-pending", cancelPendingScrapeTasksHandler(svc))
admin.POST("/scraper/queue/enqueue-library/:id", enqueueLibraryScrapeHandler(svc))
admin.POST("/scraper/queue/enqueue-all", enqueueAllScrapeHandler(svc))
admin.POST("/media/repair-rescrape", enqueueAllScrapeHandler(svc))
admin.POST("/libraries/:id/repair-rescrape", enqueueLibraryScrapeHandler(svc))
}
func registerAdminStrmRoutes(admin *gin.RouterGroup, svc *service.Container) {
@@ -43,18 +62,27 @@ func registerAdminStrmRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.POST("/strm/paths/:id/sync", startStrmSyncHandler(svc))
admin.POST("/strm/paths/:id/cancel", cancelStrmSyncHandler(svc))
admin.GET("/strm/records", listStrmSyncRecordsHandler(svc))
admin.DELETE("/strm/records/:id", deleteStrmSyncRecordHandler(svc))
admin.DELETE("/strm/records", clearStrmSyncRecordsHandler(svc))
admin.GET("/strm/local-dirs", listStrmLocalDirsHandler(svc))
admin.GET("/strm/downloads", downloadQueueHandler(svc))
admin.POST("/strm/downloads/:id/cancel", cancelStrmDownloadHandler(svc))
admin.POST("/strm/downloads/:id/retry", retryStrmDownloadHandler(svc))
admin.DELETE("/strm/downloads/:id", deleteStrmDownloadHandler(svc))
admin.POST("/strm/downloads/batch", batchActionDownloadsHandler(svc))
admin.POST("/strm/downloads/clear-done", clearDoneDownloadsHandler(svc))
admin.POST("/strm/downloads/clear-finished", clearFinishedDownloadsHandler(svc))
admin.POST("/strm/downloads/clear-canceled", clearCanceledDownloadsHandler(svc))
admin.POST("/strm/downloads/retry-failed", retryAllFailedDownloadsHandler(svc))
admin.POST("/strm/downloads/cancel-pending", cancelPendingDownloadsHandler(svc))
admin.GET("/strm/uploads", uploadQueueHandler(svc))
admin.POST("/strm/uploads/:id/cancel", cancelStrmUploadHandler(svc))
admin.POST("/strm/uploads/:id/retry", retryStrmUploadHandler(svc))
admin.DELETE("/strm/uploads/:id", deleteStrmUploadHandler(svc))
admin.POST("/strm/uploads/batch", batchActionUploadsHandler(svc))
admin.POST("/strm/uploads/cancel-pending", cancelPendingUploadsHandler(svc))
admin.POST("/strm/uploads/clear-canceled", clearCanceledUploadsHandler(svc))
}
func registerAdminUserRoutes(admin *gin.RouterGroup, svc *service.Container) {
@@ -111,3 +139,10 @@ func registerAdminRecognitionWordRoutes(admin *gin.RouterGroup, svc *service.Con
admin.POST("/recognition-words/sync", syncRecognitionWordsHandler(svc))
admin.POST("/recognition-words/test", testRecognitionWordsHandler(svc))
}
func registerAdminDatabaseRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.GET("/database/status", getDatabaseStatusHandler(svc))
admin.POST("/database/test", testDatabaseHandler(svc))
admin.POST("/database/migrate", migrateDatabaseHandler(svc))
admin.POST("/database/save-config", saveDatabaseConfigHandler(svc))
}
@@ -21,6 +21,7 @@ func registerAuthedLibraryRoutes(authed *gin.RouterGroup, svc *service.Container
authed.POST("/libraries", middleware.AdminRequired(), createLibraryHandler(svc))
authed.GET("/libraries/:id", getLibraryHandler(svc))
authed.PATCH("/libraries/:id", middleware.AdminRequired(), updateLibraryHandler(svc))
authed.PUT("/libraries/reorder", middleware.AdminRequired(), reorderLibrariesHandler(svc))
authed.DELETE("/libraries/:id", middleware.AdminRequired(), deleteLibraryHandler(svc))
authed.GET("/libraries/:id/roots", middleware.AdminRequired(), listLibraryRootsHandler(svc))
authed.POST("/libraries/:id/roots", middleware.AdminRequired(), createLibraryRootHandler(svc))
+167
View File
@@ -0,0 +1,167 @@
package handler
import (
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MMTL/internal/service"
)
func listScrapeQueueHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "50"))
snap, err := svc.Scraper.ScrapeQueueSnapshot(c.Request.Context(), c.Query("status"), page, pageSize)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, snap)
}
}
func cancelScrapeTaskHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if err := svc.Scraper.CancelScrapeTask(c.Request.Context(), c.Param("id")); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
}
func retryScrapeTaskHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if err := svc.Scraper.RetryScrapeTask(c.Request.Context(), c.Param("id")); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
}
func deleteScrapeTaskHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if err := svc.Scraper.DeleteScrapeTask(c.Request.Context(), c.Param("id")); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
}
func batchActionScrapeTasksHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req queueBatchReq
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
n, err := svc.Scraper.BatchActionScrapeTasks(c.Request.Context(), req.Action, req.IDs)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"affected": n, "action": req.Action})
}
}
func clearDoneScrapeTasksHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
n, err := svc.Scraper.ClearDoneScrapeTasks(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"deleted": n})
}
}
func clearFinishedScrapeTasksHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
n, err := svc.Scraper.ClearFinishedScrapeTasks(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"deleted": n})
}
}
func clearCanceledScrapeTasksHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
n, err := svc.Scraper.ClearCanceledScrapeTasks(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"deleted": n})
}
}
func retryAllFailedScrapeTasksHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
n, err := svc.Scraper.RetryAllFailedScrapeTasks(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"retried": n})
}
}
func cancelPendingScrapeTasksHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
n, err := svc.Scraper.CancelPendingScrapeTasks(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"canceled": n})
}
}
type enqueueScrapeReq struct {
EpisodeImages bool `json:"episode_images"`
EpisodeArtwork bool `json:"episode_artwork"`
RefreshMatched bool `json:"refresh_matched"`
IncludeMatched bool `json:"include_matched"`
}
func (r enqueueScrapeReq) toOptions() service.ScrapeOptions {
epArtwork := r.EpisodeImages || r.EpisodeArtwork
return service.ScrapeOptions{
EpisodeArtwork: &epArtwork,
IncludeMatched: r.IncludeMatched || r.RefreshMatched,
RetryNoMatch: true,
}
}
func enqueueLibraryScrapeHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req enqueueScrapeReq
_ = c.ShouldBindJSON(&req)
libID := c.Param("id")
n, err := svc.Scraper.EnqueueLibrary(c.Request.Context(), libID, req.toOptions())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"enqueued": n})
}
}
func enqueueAllScrapeHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req enqueueScrapeReq
_ = c.ShouldBindJSON(&req)
n, err := svc.Scraper.EnqueueAll(c.Request.Context(), req.toOptions())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"enqueued": n})
}
}
+8 -53
View File
@@ -2,7 +2,6 @@
package handler
import (
"context"
"errors"
"io"
"net/http"
@@ -135,28 +134,12 @@ func scrapeOneHandler(svc *service.Container) gin.HandlerFunc {
return
}
options.IncludeMatched = true
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
if err != nil || m == nil {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
}
task := startScrapeHTTPTask(svc, "手动刮削媒体", m.Title, m.Path)
if err := svc.Scraper.EnrichOneWithOptions(c.Request.Context(), m, options); err != nil {
finishHTTPTask(task, err, "scrape", "手动刮削媒体失败", nil, nil)
task, err := svc.Scraper.EnqueueMedia(c.Request.Context(), c.Param("id"), options)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
reclassified := reclassifyMediaAfterScrape(c.Request.Context(), svc, m.ID)
refreshed, _ := svc.Repo.Media.FindByID(c.Request.Context(), m.ID)
metrics := map[string]int64{"processed": 1}
if refreshed != nil && refreshed.ScrapeStatus == "matched" {
metrics["matched"] = 1
}
if reclassified > 0 {
metrics["reclassified"] = int64(reclassified)
}
finishHTTPTask(task, nil, "completed", "手动刮削媒体结束", metrics, nil)
c.JSON(http.StatusOK, refreshed)
c.JSON(http.StatusOK, task)
}
}
@@ -170,40 +153,12 @@ func scrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
return
}
options.IncludeMatched = true
var task *service.TaskHandle
if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), libID); err == nil && lib != nil {
task = startScrapeHTTPTask(svc, "手动刮削媒体库", lib.Name, lib.Path)
} else {
task = startScrapeHTTPTask(svc, "手动刮削媒体库", libID, "")
n, err := svc.Scraper.EnqueueLibrary(c.Request.Context(), libID, options)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
// Run in the background so HTTP returns instantly; the WS hub
// pushes per-item progress on the "scrape" topic.
go func(libID string, task *service.TaskHandle, options service.ScrapeOptions) {
result, err := svc.Scraper.EnrichLibraryDetailedWithOptions(context.Background(), libID, options)
reclassified := 0
if result.Processed > 0 {
reclassified = reclassifyLibraryAfterScrape(context.Background(), svc, libID)
}
metrics := map[string]int64{
"matched": int64(result.Matched),
"processed": int64(result.Processed),
"candidates": int64(result.Candidates),
}
if reclassified > 0 {
metrics["reclassified"] = int64(reclassified)
}
if result.Failed > 0 {
metrics["errors"] = int64(result.Failed)
}
stage := "completed"
message := "手动刮削媒体库结束"
if err != nil {
stage = "scrape"
message = "手动刮削媒体库失败"
}
finishHTTPTask(task, err, stage, message, metrics, nil)
}(libID, task, options)
c.JSON(http.StatusAccepted, gin.H{"status": "scraping"})
c.JSON(http.StatusOK, gin.H{"status": "queued", "enqueued": n})
}
}
+127 -1
View File
@@ -178,6 +178,7 @@ type strmSyncPathReq struct {
DeleteDir *bool `json:"delete_dir"`
Cron string `json:"cron"`
EnableCron *bool `json:"enable_cron"`
SyncMode string `json:"sync_mode"`
Enabled *bool `json:"enabled"`
}
@@ -261,7 +262,16 @@ func deleteStrmSyncPathHandler(svc *service.Container) gin.HandlerFunc {
func startStrmSyncHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if err := svc.Strm.StartSync(c.Request.Context(), c.Param("id")); err != nil {
mode := c.Query("mode")
if mode == "" {
var body struct {
Mode string `json:"mode"`
}
if err := c.ShouldBindJSON(&body); err == nil && body.Mode != "" {
mode = body.Mode
}
}
if err := svc.Strm.StartSync(c.Request.Context(), c.Param("id"), mode); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
@@ -290,6 +300,31 @@ func listStrmSyncRecordsHandler(svc *service.Container) gin.HandlerFunc {
}
}
func deleteStrmSyncRecordHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if c.Param("id") == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "缺少记录 ID"})
return
}
if err := svc.Strm.DeleteSyncRecord(c.Request.Context(), c.Param("id")); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
}
func clearStrmSyncRecordsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
deleted, err := svc.Strm.ClearSyncRecords(c.Request.Context(), c.Query("path_id"))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true, "deleted": deleted})
}
}
// ─── 下载/上传队列 ─────────────────────────────────────────────────────────────
func downloadQueueHandler(svc *service.Container) gin.HandlerFunc {
@@ -358,6 +393,63 @@ func retryStrmUploadHandler(svc *service.Container) gin.HandlerFunc {
}
}
func deleteStrmDownloadHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if err := svc.Strm.DeleteDownloadTask(c.Request.Context(), c.Param("id")); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
}
func deleteStrmUploadHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if err := svc.Strm.DeleteUploadTask(c.Request.Context(), c.Param("id")); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
}
type queueBatchReq struct {
Action string `json:"action" binding:"required"`
IDs []string `json:"ids" binding:"required"`
}
func batchActionDownloadsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req queueBatchReq
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
n, err := svc.Strm.BatchActionDownloadTasks(c.Request.Context(), req.Action, req.IDs)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"affected": n, "action": req.Action})
}
}
func batchActionUploadsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req queueBatchReq
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
n, err := svc.Strm.BatchActionUploadTasks(c.Request.Context(), req.Action, req.IDs)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"affected": n, "action": req.Action})
}
}
// ─── 下载队列批量操作 ─────────────────────────────────────────────────────────
func clearDoneDownloadsHandler(svc *service.Container) gin.HandlerFunc {
@@ -382,6 +474,28 @@ func clearFinishedDownloadsHandler(svc *service.Container) gin.HandlerFunc {
}
}
func clearCanceledDownloadsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
n, err := svc.Strm.ClearCanceledDownloadTasks(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"deleted": n})
}
}
func clearCanceledUploadsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
n, err := svc.Strm.ClearCanceledUploadTasks(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"deleted": n})
}
}
func retryAllFailedDownloadsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
n, err := svc.Strm.RetryAllFailedDownloadTasks(c.Request.Context())
@@ -404,6 +518,17 @@ func cancelPendingDownloadsHandler(svc *service.Container) gin.HandlerFunc {
}
}
func cancelPendingUploadsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
n, err := svc.Strm.CancelPendingUploadTasks(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"canceled": n})
}
}
// ─── 公开播放端点 ──────────────────────────────────────────────────────────────
// strmPlayHandler 处理 strm 文件指向的播放请求(Emby/Infuse 直接请求,无 JWT)。
@@ -458,6 +583,7 @@ func strmSyncPathFromReq(req strmSyncPathReq) *model.StrmSyncPath {
DeleteDir: boolValue(req.DeleteDir, false),
Cron: strings.TrimSpace(req.Cron),
EnableCron: boolValue(req.EnableCron, false),
SyncMode: strings.TrimSpace(req.SyncMode),
Enabled: boolValue(req.Enabled, true),
}
}
+10 -4
View File
@@ -50,10 +50,16 @@ func TestStrmAdminRoutesAreRegistered(t *testing.T) {
"GET /api/admin/strm/downloads",
"POST /api/admin/strm/downloads/:id/cancel",
"POST /api/admin/strm/downloads/:id/retry",
"GET /api/admin/strm/uploads",
"POST /api/admin/strm/uploads/:id/cancel",
"POST /api/admin/strm/uploads/:id/retry",
"GET /api/strm/play/:provider/:file",
"POST /api/admin/strm/downloads/clear-finished",
"POST /api/admin/strm/downloads/clear-canceled",
"POST /api/admin/strm/downloads/retry-failed",
"POST /api/admin/strm/downloads/cancel-pending",
"GET /api/admin/strm/uploads",
"POST /api/admin/strm/uploads/:id/cancel",
"POST /api/admin/strm/uploads/:id/retry",
"POST /api/admin/strm/uploads/cancel-pending",
"POST /api/admin/strm/uploads/clear-canceled",
"GET /api/strm/play/:provider/:file",
} {
if !routes[want] {
t.Fatalf("%s route is not registered", want)
+1 -1
View File
@@ -163,7 +163,7 @@ func historyDeleteHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusBadRequest, gin.H{"error": "status must be completed or incomplete"})
return
}
res := q.Delete(&model.PlaybackHistory{})
res := q.Unscoped().Delete(&model.PlaybackHistory{})
if err := res.Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
+8 -6
View File
@@ -3,12 +3,14 @@ package model
// Library 表示一个逻辑媒体库。Path 保留为兼容字段,指向第一条 LibraryRoot。
type Library struct {
Base
Name string `gorm:"size:128;not null" json:"name"`
Path string `gorm:"size:1024;not null" json:"path"`
Type string `gorm:"size:16;not null;default:movie" json:"type"` // movie / tv / anime / music
CoverURL string `gorm:"size:1024" json:"cover_url,omitempty"`
Enabled bool `gorm:"default:true" json:"enabled"`
Roots []LibraryRoot `gorm:"foreignKey:LibraryID" json:"roots,omitempty"`
Name string `gorm:"size:128;not null" json:"name"`
Path string `gorm:"size:1024;not null" json:"path"`
Type string `gorm:"size:16;not null;default:movie" json:"type"` // movie / tv / anime / music
CoverURL string `gorm:"size:1024" json:"cover_url,omitempty"`
Enabled bool `gorm:"default:true" json:"enabled"`
SortOrder int `gorm:"index;default:0" json:"sort_order"` // 手动拖拽排序用,越小越靠前
CarouselEnabled bool `gorm:"default:false" json:"carousel_enabled"` // 是否参与首页海报轮播(默认不参与)
Roots []LibraryRoot `gorm:"foreignKey:LibraryID" json:"roots,omitempty"`
}
// LibraryRoot 是逻辑媒体库下的一条真实物理/挂载路径。
+5 -3
View File
@@ -54,7 +54,9 @@ func AllModels() []interface{} {
&StrmAccount{},
&StrmSyncPath{},
&StrmSyncRecord{},
&StrmDownloadTask{},
&StrmUploadTask{},
}
&StrmDownloadTask{},
&StrmUploadTask{},
&StrmDirCache{},
&ScrapeTask{},
}
}
+34
View File
@@ -0,0 +1,34 @@
package model
import "time"
const (
ScrapeTaskPending = "pending"
ScrapeTaskRunning = "running"
ScrapeTaskDone = "done"
ScrapeTaskFailed = "failed"
ScrapeTaskCanceled = "canceled"
)
// ScrapeTask 表示一条持久化的媒体刮削任务。
type ScrapeTask struct {
Base
MediaID string `gorm:"index;size:36" json:"media_id"`
LibraryID string `gorm:"index;size:36" json:"library_id"`
LibraryName string `gorm:"size:128" json:"library_name"`
MediaTitle string `gorm:"size:255;not null" json:"media_title"`
MediaPath string `gorm:"size:1024;not null" json:"media_path"`
MediaType string `gorm:"size:16" json:"media_type"` // movie / tv / anime / adult
Provider string `gorm:"size:32" json:"provider"` // tmdb / douban / bangumi / thetvdb / metatube
MatchedTitle string `gorm:"size:255" json:"matched_title"`
MatchedYear int `json:"matched_year"`
PosterURL string `gorm:"size:1024" json:"poster_url"`
BackdropURL string `gorm:"size:1024" json:"backdrop_url"`
Status string `gorm:"index;size:16;default:pending" json:"status"` // pending / running / done / failed / canceled
Error string `gorm:"type:text" json:"error"`
RetryCount int `gorm:"default:0" json:"retry_count"`
EpisodeImages bool `gorm:"default:true" json:"episode_images"`
RefreshMatched bool `gorm:"default:false" json:"refresh_matched"`
StartedAt *time.Time `json:"started_at,omitempty"`
FinishedAt *time.Time `json:"finished_at,omitempty"`
}
+17
View File
@@ -48,12 +48,19 @@ type StrmSyncPath struct {
DeleteDir bool `json:"delete_dir"` // 清理多余文件时删除空目录
Cron string `gorm:"size:128" json:"cron"` // 5 段 cron 表达式(可选)
EnableCron bool `json:"enable_cron"` // 是否按 Cron 定时同步
SyncMode string `gorm:"size:32;default:'incremental'" json:"sync_mode"` // 默认同步模式:incremental / full
Enabled bool `gorm:"default:true" json:"enabled"`
LastSyncAt *time.Time `json:"last_sync_at"`
LastSyncStatus string `gorm:"size:16" json:"last_sync_status"` // idle/running/ok/error/canceled
LastSyncMessage string `gorm:"size:1024" json:"last_sync_message"`
}
// STRM 同步类型。
const (
StrmSyncTypeIncremental = "incremental"
StrmSyncTypeFull = "full"
)
// StrmSyncRecord 是一次同步执行的记录。
const (
StrmSyncRecordPending = "pending"
@@ -66,6 +73,7 @@ const (
type StrmSyncRecord struct {
Base
SyncPathID string `gorm:"size:36;index" json:"sync_path_id"`
SyncType string `gorm:"size:32;default:'incremental'" json:"sync_type"` // incremental / full
Status string `gorm:"size:16;index" json:"status"`
Total int64 `json:"total"` // 远端发现的文件总数
NewStrm int64 `json:"new_strm"` // 本次新建/更新的 strm 数
@@ -123,3 +131,12 @@ type StrmUploadTask struct {
StartedAt *time.Time `json:"started_at"`
FinishedAt *time.Time `json:"finished_at"`
}
// StrmDirCache 缓存远端网盘目录 ID 与相对路径映射(支持 115 增量同步秒级寻址)。
type StrmDirCache struct {
Base
SyncPathID string `gorm:"size:36;index:idx_strm_dir_cache,priority:1" json:"sync_path_id"`
DirID string `gorm:"size:128;index:idx_strm_dir_cache,priority:2" json:"dir_id"`
Path string `gorm:"size:1024" json:"path"` // 相对根目录的路径
}
+2 -2
View File
@@ -62,9 +62,9 @@ func (r *ApiConfigRepository) Update(ctx context.Context, c *model.ApiConfig) er
}).Error
}
// Delete removes an API config.
// Delete 物理删除 API 配置。
func (r *ApiConfigRepository) Delete(ctx context.Context, provider string) error {
return r.db.WithContext(ctx).Where("provider = ?", provider).Delete(&model.ApiConfig{}).Error
return r.db.WithContext(ctx).Unscoped().Where("provider = ?", provider).Delete(&model.ApiConfig{}).Error
}
// UpdateTestResult 更新测试结果。
+1 -1
View File
@@ -23,7 +23,7 @@ func (r *FavoriteRepository) Toggle(ctx context.Context, userID, mediaID string)
if err != nil {
return false, err
}
return false, r.db.WithContext(ctx).Delete(&f).Error
return false, r.db.WithContext(ctx).Unscoped().Delete(&f).Error
}
// ListByUser returns all favourite media IDs for a user.
+31 -5
View File
@@ -15,6 +15,11 @@ type LibraryRepository struct{ db *gorm.DB }
// Create persists a new library row.
func (r *LibraryRepository) Create(ctx context.Context, l *model.Library) error {
if l != nil && l.SortOrder == 0 {
var maxSort int
_ = r.db.WithContext(ctx).Model(&model.Library{}).Select("COALESCE(MAX(sort_order), -1)").Scan(&maxSort)
l.SortOrder = maxSort + 1
}
return r.db.WithContext(ctx).Create(l).Error
}
@@ -23,6 +28,11 @@ func (r *LibraryRepository) CreateWithRoots(ctx context.Context, l *model.Librar
return r.Create(ctx, l)
}
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if l != nil && l.SortOrder == 0 {
var maxSort int
_ = tx.Model(&model.Library{}).Select("COALESCE(MAX(sort_order), -1)").Scan(&maxSort)
l.SortOrder = maxSort + 1
}
if err := tx.Create(l).Error; err != nil {
return err
}
@@ -50,7 +60,7 @@ func (r *LibraryRepository) CreateWithRoots(ctx context.Context, l *model.Librar
// List returns all enabled+disabled libraries.
func (r *LibraryRepository) List(ctx context.Context) ([]model.Library, error) {
var ls []model.Library
q := r.db.WithContext(ctx).Order("created_at asc")
q := r.db.WithContext(ctx).Order("sort_order asc, created_at asc")
if r.hasLibraryRootsTable() {
q = q.Preload("Roots", func(db *gorm.DB) *gorm.DB {
return db.Order("sort_order asc, created_at asc")
@@ -60,6 +70,23 @@ func (r *LibraryRepository) List(ctx context.Context) ([]model.Library, error) {
return ls, err
}
// SetSortOrder assigns sort_order to libraries, preserving position order for
// any library not present in the provided map.
func (r *LibraryRepository) SetSortOrder(ctx context.Context, ids []string) error {
if len(ids) == 0 {
return nil
}
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
for i, id := range ids {
if err := tx.Model(&model.Library{}).Where("id = ?", id).
Update("sort_order", i).Error; err != nil {
return err
}
}
return nil
})
}
// FindByID returns the library, or (nil, nil) when missing.
func (r *LibraryRepository) FindByID(ctx context.Context, id string) (*model.Library, error) {
var l model.Library
@@ -79,10 +106,9 @@ func (r *LibraryRepository) FindByID(ctx context.Context, id string) (*model.Lib
return &l, nil
}
// Delete removes a library and (soft) cascades to its media via repository
// callers; we do not run CASCADE here to keep this method narrow.
// Delete 物理删除媒体库。
func (r *LibraryRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Delete(&model.Library{}, "id = ?", id).Error
return r.db.WithContext(ctx).Unscoped().Delete(&model.Library{}, "id = ?", id).Error
}
func (r *LibraryRepository) ListRoots(ctx context.Context, libraryID string) ([]model.LibraryRoot, error) {
@@ -149,7 +175,7 @@ func (r *LibraryRepository) DeleteRoot(ctx context.Context, libraryID, rootID st
if !r.hasLibraryRootsTable() {
return nil
}
return r.db.WithContext(ctx).Where("library_id = ?", libraryID).Delete(&model.LibraryRoot{}, "id = ?", rootID).Error
return r.db.WithContext(ctx).Unscoped().Where("library_id = ?", libraryID).Delete(&model.LibraryRoot{}, "id = ?", rootID).Error
}
func (r *LibraryRepository) hasLibraryRootsTable() bool {
+3 -3
View File
@@ -116,12 +116,12 @@ func (r *MediaRepository) ListByLibrariesFiltered(ctx context.Context, libraryID
// DeleteByLibrary purges all media tied to a library.
func (r *MediaRepository) DeleteByLibrary(ctx context.Context, libraryID string) error {
// FTS 行由 media 表上的触发器同步清理(软删/硬删都覆盖)。
return r.db.WithContext(ctx).Where("library_id = ?", libraryID).Delete(&model.Media{}).Error
// FTS 行由 media 表上的触发器同步清理(物理删除触发 FTS 清理)。
return r.db.WithContext(ctx).Unscoped().Where("library_id = ?", libraryID).Delete(&model.Media{}).Error
}
func (r *MediaRepository) DeleteByLibraryRoot(ctx context.Context, libraryID, rootID string) error {
return r.db.WithContext(ctx).
return r.db.WithContext(ctx).Unscoped().
Where("library_id = ? AND library_root_id = ?", libraryID, rootID).
Delete(&model.Media{}).Error
}
+2 -2
View File
@@ -51,9 +51,9 @@ func (r *PermissionRepository) Upsert(ctx context.Context, p *model.UserPermissi
})
}
// Delete removes a permission record.
// Delete 物理删除权限记录。
func (r *PermissionRepository) Delete(ctx context.Context, userID string) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Where("user_id = ?", userID).Delete(&model.UserPermission{}).Error
return r.db.WithContext(ctx).Unscoped().Where("user_id = ?", userID).Delete(&model.UserPermission{}).Error
})
}
+2 -2
View File
@@ -59,9 +59,9 @@ func (r *PlayProfileRepository) Update(ctx context.Context, id string, patch map
Where("id = ?", id).Updates(patch).Error
}
// Delete soft-deletes a profile.
// Delete 物理删除播放档案。
func (r *PlayProfileRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Delete(&model.PlayProfile{}, "id = ?", id).Error
return r.db.WithContext(ctx).Unscoped().Delete(&model.PlayProfile{}, "id = ?", id).Error
}
// ClearDefaultsFor resets is_default for all of a user's profiles.
@@ -72,10 +72,10 @@ func (r *RefreshTokenRepository) RevokeOldestActiveByUserID(ctx context.Context,
})
}
// DeleteExpired removes all expired refresh tokens.
// DeleteExpired 物理清理所有过期的 refresh tokens。
func (r *RefreshTokenRepository) DeleteExpired(ctx context.Context) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Where("expires_at < ?", time.Now()).Delete(&model.RefreshToken{}).Error
return r.db.WithContext(ctx).Unscoped().Where("expires_at < ?", time.Now()).Delete(&model.RefreshToken{}).Error
})
}
+6 -2
View File
@@ -30,8 +30,10 @@ type Container struct {
StrmSyncPath *StrmSyncPathRepository
StrmSyncRecord *StrmSyncRecordRepository
StrmDownload *StrmDownloadTaskRepository
StrmUpload *StrmUploadTaskRepository
}
StrmUpload *StrmUploadTaskRepository
StrmDirCache *StrmDirCacheRepository
ScrapeTask *ScrapeTaskRepository
}
// New 将每个 repository 连接到单个 *gorm.DB。
func New(db *gorm.DB) *Container {
@@ -58,5 +60,7 @@ func New(db *gorm.DB) *Container {
StrmSyncRecord: &StrmSyncRecordRepository{db: db},
StrmDownload: &StrmDownloadTaskRepository{db: db},
StrmUpload: &StrmUploadTaskRepository{db: db},
StrmDirCache: &StrmDirCacheRepository{db: db},
ScrapeTask: &ScrapeTaskRepository{db: db},
}
}
@@ -0,0 +1,274 @@
package repository
import (
"context"
"errors"
"strings"
"sync"
"time"
"gorm.io/gorm"
"github.com/ShukeBta/MMTL/internal/model"
)
var scrapeClaimMu sync.Mutex
// ScrapeTaskRepository persists model.ScrapeTask.
type ScrapeTaskRepository struct{ db *gorm.DB }
func (r *ScrapeTaskRepository) Create(ctx context.Context, t *model.ScrapeTask) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Create(t).Error
})
}
func (r *ScrapeTaskRepository) CreateBatch(ctx context.Context, tasks []model.ScrapeTask) error {
if len(tasks) == 0 {
return nil
}
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).CreateInBatches(tasks, 100).Error
})
}
func (r *ScrapeTaskRepository) FindByID(ctx context.Context, id string) (*model.ScrapeTask, error) {
var t model.ScrapeTask
err := r.db.WithContext(ctx).Where("id = ?", id).First(&t).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
return &t, err
}
func (r *ScrapeTaskRepository) FindActiveByMediaID(ctx context.Context, mediaID string) (*model.ScrapeTask, error) {
var t model.ScrapeTask
err := r.db.WithContext(ctx).
Where("media_id = ? AND status IN ?", mediaID, []string{model.ScrapeTaskPending, model.ScrapeTaskRunning}).
First(&t).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
return &t, err
}
func (r *ScrapeTaskRepository) List(ctx context.Context, status string, page, pageSize int) ([]model.ScrapeTask, int64, error) {
if page < 1 {
page = 1
}
if pageSize < 1 || pageSize > 200 {
pageSize = 50
}
q := r.db.WithContext(ctx).Model(&model.ScrapeTask{})
if strings.TrimSpace(status) != "" && status != "all" {
q = q.Where("status = ?", strings.TrimSpace(status))
}
var total int64
if err := q.Count(&total).Error; err != nil {
return nil, 0, err
}
var rows []model.ScrapeTask
err := q.Order("created_at desc").
Offset((page - 1) * pageSize).
Limit(pageSize).
Find(&rows).Error
return rows, total, err
}
func (r *ScrapeTaskRepository) CountByStatus(ctx context.Context) (map[string]int64, error) {
var rows []struct {
Status string
Count int64
}
err := r.db.WithContext(ctx).Model(&model.ScrapeTask{}).
Select("status, count(*) as count").
Group("status").Scan(&rows).Error
if err != nil {
return nil, err
}
out := map[string]int64{}
for _, row := range rows {
out[row.Status] = row.Count
}
return out, nil
}
// ClaimPending picks pending scrape tasks and marks them running.
func (r *ScrapeTaskRepository) ClaimPending(ctx context.Context, limit int) ([]model.ScrapeTask, error) {
scrapeClaimMu.Lock()
defer scrapeClaimMu.Unlock()
var rows []model.ScrapeTask
err := withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("status = ?", model.ScrapeTaskPending).
Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil {
return err
}
if len(rows) == 0 {
return nil
}
ids := make([]string, 0, len(rows))
now := time.Now()
for i := range rows {
ids = append(ids, rows[i].ID)
rows[i].Status = model.ScrapeTaskRunning
rows[i].StartedAt = &now
}
return tx.Model(&model.ScrapeTask{}).Where("id IN ?", ids).
Updates(map[string]any{"status": model.ScrapeTaskRunning, "started_at": now}).Error
})
})
if err != nil {
return nil, err
}
return rows, nil
}
func (r *ScrapeTaskRepository) Update(ctx context.Context, t *model.ScrapeTask) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Model(&model.ScrapeTask{}).Where("id = ?", t.ID).Updates(map[string]any{
"status": t.Status,
"error": t.Error,
"provider": t.Provider,
"matched_title": t.MatchedTitle,
"matched_year": t.MatchedYear,
"poster_url": t.PosterURL,
"backdrop_url": t.BackdropURL,
"retry_count": t.RetryCount,
"started_at": t.StartedAt,
"finished_at": t.FinishedAt,
"updated_at": time.Now(),
}).Error
})
}
func (r *ScrapeTaskRepository) Delete(ctx context.Context, id string) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.ScrapeTask{}).Error
})
}
func (r *ScrapeTaskRepository) DeleteBatch(ctx context.Context, ids []string) (int64, error) {
if len(ids) == 0 {
return 0, nil
}
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Unscoped().Where("id IN ?", ids).Delete(&model.ScrapeTask{})
count = res.RowsAffected
return res.Error
})
return count, err
}
func (r *ScrapeTaskRepository) RetryBatch(ctx context.Context, ids []string) (int64, error) {
if len(ids) == 0 {
return 0, nil
}
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Model(&model.ScrapeTask{}).
Where("id IN ? AND status IN ?", ids, []string{model.ScrapeTaskFailed, model.ScrapeTaskCanceled}).
Updates(map[string]any{
"status": model.ScrapeTaskPending,
"error": "",
"retry_count": 0,
"started_at": nil,
"finished_at": nil,
"updated_at": time.Now(),
})
count = res.RowsAffected
return res.Error
})
return count, err
}
func (r *ScrapeTaskRepository) CancelBatch(ctx context.Context, ids []string) (int64, error) {
if len(ids) == 0 {
return 0, nil
}
now := time.Now()
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Model(&model.ScrapeTask{}).
Where("id IN ? AND status IN ?", ids, []string{model.ScrapeTaskPending, model.ScrapeTaskRunning}).
Updates(map[string]any{
"status": model.ScrapeTaskCanceled,
"error": "已批量取消",
"finished_at": now,
"updated_at": now,
})
count = res.RowsAffected
return res.Error
})
return count, err
}
func (r *ScrapeTaskRepository) ClearDone(ctx context.Context) (int64, error) {
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Unscoped().Where("status = ?", model.ScrapeTaskDone).Delete(&model.ScrapeTask{})
count = res.RowsAffected
return res.Error
})
return count, err
}
func (r *ScrapeTaskRepository) ClearFinished(ctx context.Context) (int64, error) {
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Unscoped().Where("status IN ?", []string{model.ScrapeTaskDone, model.ScrapeTaskFailed, model.ScrapeTaskCanceled}).
Delete(&model.ScrapeTask{})
count = res.RowsAffected
return res.Error
})
return count, err
}
func (r *ScrapeTaskRepository) ClearCanceled(ctx context.Context) (int64, error) {
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Unscoped().Where("status = ?", model.ScrapeTaskCanceled).Delete(&model.ScrapeTask{})
count = res.RowsAffected
return res.Error
})
return count, err
}
func (r *ScrapeTaskRepository) RetryAllFailed(ctx context.Context) (int64, error) {
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Model(&model.ScrapeTask{}).
Where("status = ?", model.ScrapeTaskFailed).
Updates(map[string]any{
"status": model.ScrapeTaskPending,
"error": "",
"retry_count": 0,
"started_at": nil,
"finished_at": nil,
"updated_at": time.Now(),
})
count = res.RowsAffected
return res.Error
})
return count, err
}
func (r *ScrapeTaskRepository) CancelPending(ctx context.Context) (int64, error) {
now := time.Now()
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Model(&model.ScrapeTask{}).
Where("status IN ?", []string{model.ScrapeTaskPending, model.ScrapeTaskRunning}).
Updates(map[string]any{
"status": model.ScrapeTaskCanceled,
"error": "已批量取消",
"finished_at": now,
"updated_at": now,
})
count = res.RowsAffected
return res.Error
})
return count, err
}
+2 -2
View File
@@ -29,9 +29,9 @@ func (r *SettingRepository) Set(ctx context.Context, key, value string) error {
return r.db.WithContext(ctx).Save(&s).Error
}
// Delete removes a setting key.
// Delete 物理删除设置键。
func (r *SettingRepository) Delete(ctx context.Context, key string) error {
return r.db.WithContext(ctx).Where("key = ?", key).Delete(&model.Setting{}).Error
return r.db.WithContext(ctx).Unscoped().Where("key = ?", key).Delete(&model.Setting{}).Error
}
// All returns every key/value pair (used by the admin UI).
+2 -2
View File
@@ -66,9 +66,9 @@ func (r *StorageConfigRepository) Upsert(ctx context.Context, c *model.StorageCo
}).Error
}
// Delete removes a storage config by ID.
// Delete 物理删除存储配置。
func (r *StorageConfigRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StorageConfig{}).Error
return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StorageConfig{}).Error
}
// FindByID returns a storage config by ID.
+484 -141
View File
@@ -3,6 +3,7 @@ package repository
import (
"context"
"errors"
"sync"
"time"
"gorm.io/gorm"
@@ -10,13 +11,17 @@ import (
"github.com/ShukeBta/MMTL/internal/model"
)
var strmClaimMu sync.Mutex
// ─── StrmAccount ───────────────────────────────────────────────────────────────
// StrmAccountRepository persists model.StrmAccount.
type StrmAccountRepository struct{ db *gorm.DB }
func (r *StrmAccountRepository) Create(ctx context.Context, a *model.StrmAccount) error {
return r.db.WithContext(ctx).Create(a).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Create(a).Error
})
}
func (r *StrmAccountRepository) FindByID(ctx context.Context, id string) (*model.StrmAccount, error) {
@@ -38,20 +43,24 @@ func (r *StrmAccountRepository) List(ctx context.Context) ([]model.StrmAccount,
}
func (r *StrmAccountRepository) Update(ctx context.Context, a *model.StrmAccount) error {
return r.db.WithContext(ctx).Model(&model.StrmAccount{}).Where("id = ?", a.ID).Updates(map[string]any{
"name": a.Name,
"provider": a.Provider,
"config": a.Config,
"enabled": a.Enabled,
"last_test_at": a.LastTestAt,
"last_test_result": a.LastTestResult,
"last_test_ok": a.LastTestOK,
"updated_at": time.Now(),
}).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Model(&model.StrmAccount{}).Where("id = ?", a.ID).Updates(map[string]any{
"name": a.Name,
"provider": a.Provider,
"config": a.Config,
"enabled": a.Enabled,
"last_test_at": a.LastTestAt,
"last_test_result": a.LastTestResult,
"last_test_ok": a.LastTestOK,
"updated_at": time.Now(),
}).Error
})
}
func (r *StrmAccountRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmAccount{}).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmAccount{}).Error
})
}
// ─── StrmSyncPath ──────────────────────────────────────────────────────────────
@@ -60,7 +69,9 @@ func (r *StrmAccountRepository) Delete(ctx context.Context, id string) error {
type StrmSyncPathRepository struct{ db *gorm.DB }
func (r *StrmSyncPathRepository) Create(ctx context.Context, p *model.StrmSyncPath) error {
return r.db.WithContext(ctx).Create(p).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Create(p).Error
})
}
func (r *StrmSyncPathRepository) FindByID(ctx context.Context, id string) (*model.StrmSyncPath, error) {
@@ -82,33 +93,38 @@ func (r *StrmSyncPathRepository) List(ctx context.Context) ([]model.StrmSyncPath
}
func (r *StrmSyncPathRepository) Update(ctx context.Context, p *model.StrmSyncPath) error {
return r.db.WithContext(ctx).Model(&model.StrmSyncPath{}).Where("id = ?", p.ID).Updates(map[string]any{
"name": p.Name,
"account_id": p.AccountID,
"provider": p.Provider,
"remote_path": p.RemotePath,
"local_path": p.LocalPath,
"strm_base_url": p.StrmBaseURL,
"video_ext": p.VideoExt,
"meta_ext": p.MetaExt,
"exclude_name": p.ExcludeName,
"min_video_size_mb": p.MinVideoSizeMB,
"add_path": p.AddPath,
"download_meta": p.DownloadMeta,
"upload_meta": p.UploadMeta,
"delete_dir": p.DeleteDir,
"cron": p.Cron,
"enable_cron": p.EnableCron,
"enabled": p.Enabled,
"last_sync_at": p.LastSyncAt,
"last_sync_status": p.LastSyncStatus,
"last_sync_message": p.LastSyncMessage,
"updated_at": time.Now(),
}).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Model(&model.StrmSyncPath{}).Where("id = ?", p.ID).Updates(map[string]any{
"name": p.Name,
"account_id": p.AccountID,
"provider": p.Provider,
"remote_path": p.RemotePath,
"local_path": p.LocalPath,
"strm_base_url": p.StrmBaseURL,
"video_ext": p.VideoExt,
"meta_ext": p.MetaExt,
"exclude_name": p.ExcludeName,
"min_video_size_mb": p.MinVideoSizeMB,
"add_path": p.AddPath,
"download_meta": p.DownloadMeta,
"upload_meta": p.UploadMeta,
"delete_dir": p.DeleteDir,
"cron": p.Cron,
"enable_cron": p.EnableCron,
"sync_mode": p.SyncMode,
"enabled": p.Enabled,
"last_sync_at": p.LastSyncAt,
"last_sync_status": p.LastSyncStatus,
"last_sync_message": p.LastSyncMessage,
"updated_at": time.Now(),
}).Error
})
}
func (r *StrmSyncPathRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmSyncPath{}).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmSyncPath{}).Error
})
}
// ─── StrmSyncRecord ────────────────────────────────────────────────────────────
@@ -117,23 +133,28 @@ func (r *StrmSyncPathRepository) Delete(ctx context.Context, id string) error {
type StrmSyncRecordRepository struct{ db *gorm.DB }
func (r *StrmSyncRecordRepository) Create(ctx context.Context, rec *model.StrmSyncRecord) error {
return r.db.WithContext(ctx).Create(rec).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Create(rec).Error
})
}
func (r *StrmSyncRecordRepository) Update(ctx context.Context, rec *model.StrmSyncRecord) error {
return r.db.WithContext(ctx).Model(&model.StrmSyncRecord{}).Where("id = ?", rec.ID).Updates(map[string]any{
"status": rec.Status,
"total": rec.Total,
"new_strm": rec.NewStrm,
"new_meta": rec.NewMeta,
"uploaded": rec.Uploaded,
"pruned": rec.Pruned,
"skipped": rec.Skipped,
"message": rec.Message,
"started_at": rec.StartedAt,
"finished_at": rec.FinishedAt,
"updated_at": time.Now(),
}).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Model(&model.StrmSyncRecord{}).Where("id = ?", rec.ID).Updates(map[string]any{
"sync_type": rec.SyncType,
"status": rec.Status,
"total": rec.Total,
"new_strm": rec.NewStrm,
"new_meta": rec.NewMeta,
"uploaded": rec.Uploaded,
"pruned": rec.Pruned,
"skipped": rec.Skipped,
"message": rec.Message,
"started_at": rec.StartedAt,
"finished_at": rec.FinishedAt,
"updated_at": time.Now(),
}).Error
})
}
func (r *StrmSyncRecordRepository) List(ctx context.Context, syncPathID string, limit int) ([]model.StrmSyncRecord, error) {
@@ -149,13 +170,45 @@ func (r *StrmSyncRecordRepository) List(ctx context.Context, syncPathID string,
return rows, err
}
// Delete 删除单条同步记录(物理删除)。
func (r *StrmSyncRecordRepository) Delete(ctx context.Context, id string) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmSyncRecord{}).Error
})
}
// DeleteBySyncPathID 删除某同步目录下的全部同步记录(删除同步目录时级联清理)。
func (r *StrmSyncRecordRepository) DeleteBySyncPathID(ctx context.Context, syncPathID string) (int64, error) {
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Unscoped().Where("sync_path_id = ?", syncPathID).Delete(&model.StrmSyncRecord{})
count = res.RowsAffected
return res.Error
})
return count, err
}
// ─── StrmDownloadTask ──────────────────────────────────────────────────────────
// StrmDownloadTaskRepository persists model.StrmDownloadTask.
type StrmDownloadTaskRepository struct{ db *gorm.DB }
func (r *StrmDownloadTaskRepository) Create(ctx context.Context, t *model.StrmDownloadTask) error {
return r.db.WithContext(ctx).Create(t).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Create(t).Error
})
}
func (r *StrmDownloadTaskRepository) CreateInBatches(ctx context.Context, tasks []*model.StrmDownloadTask, batchSize int) error {
if len(tasks) == 0 {
return nil
}
if batchSize <= 0 {
batchSize = 100
}
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).CreateInBatches(tasks, batchSize).Error
})
}
func (r *StrmDownloadTaskRepository) FindByID(ctx context.Context, id string) (*model.StrmDownloadTask, error) {
@@ -205,24 +258,29 @@ func (r *StrmDownloadTaskRepository) CountByStatus(ctx context.Context) (map[str
// ClaimPendingDownload picks the oldest pending task and marks it running.
// Returns (nil, nil) when the queue is empty.
func (r *StrmDownloadTaskRepository) ClaimPendingDownload(ctx context.Context, limit int) ([]model.StrmDownloadTask, error) {
strmClaimMu.Lock()
defer strmClaimMu.Unlock()
var rows []model.StrmDownloadTask
err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()).
Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil {
return err
}
if len(rows) == 0 {
return nil
}
ids := make([]string, 0, len(rows))
now := time.Now()
for i := range rows {
ids = append(ids, rows[i].ID)
rows[i].Status = model.StrmTaskRunning
rows[i].StartedAt = &now
}
return tx.Model(&model.StrmDownloadTask{}).Where("id IN ?", ids).
Updates(map[string]any{"status": model.StrmTaskRunning, "started_at": now}).Error
err := withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()).
Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil {
return err
}
if len(rows) == 0 {
return nil
}
ids := make([]string, 0, len(rows))
now := time.Now()
for i := range rows {
ids = append(ids, rows[i].ID)
rows[i].Status = model.StrmTaskRunning
rows[i].StartedAt = &now
}
return tx.Model(&model.StrmDownloadTask{}).Where("id IN ?", ids).
Updates(map[string]any{"status": model.StrmTaskRunning, "started_at": now}).Error
})
})
if err != nil {
return nil, err
@@ -231,62 +289,157 @@ func (r *StrmDownloadTaskRepository) ClaimPendingDownload(ctx context.Context, l
}
func (r *StrmDownloadTaskRepository) Update(ctx context.Context, t *model.StrmDownloadTask) error {
return r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).Where("id = ?", t.ID).Updates(map[string]any{
"status": t.Status,
"error": t.Error,
"retry_count": t.RetryCount,
"next_try_at": t.NextTryAt,
"started_at": t.StartedAt,
"finished_at": t.FinishedAt,
"updated_at": time.Now(),
}).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).Where("id = ?", t.ID).Updates(map[string]any{
"status": t.Status,
"error": t.Error,
"retry_count": t.RetryCount,
"next_try_at": t.NextTryAt,
"started_at": t.StartedAt,
"finished_at": t.FinishedAt,
"updated_at": time.Now(),
}).Error
})
}
func (r *StrmDownloadTaskRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmDownloadTask{}).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmDownloadTask{}).Error
})
}
// DeleteBatch 批量删除指定 ID 的下载任务。
func (r *StrmDownloadTaskRepository) DeleteBatch(ctx context.Context, ids []string) (int64, error) {
if len(ids) == 0 {
return 0, nil
}
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Unscoped().Where("id IN ?", ids).Delete(&model.StrmDownloadTask{})
count = res.RowsAffected
return res.Error
})
return count, err
}
// RetryBatch 批量重试指定 ID 的失败/已取消下载任务。
func (r *StrmDownloadTaskRepository) RetryBatch(ctx context.Context, ids []string) (int64, error) {
if len(ids) == 0 {
return 0, nil
}
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).
Where("id IN ? AND status IN ?", ids, []string{model.StrmTaskFailed, model.StrmTaskCanceled}).
Updates(map[string]any{
"status": model.StrmTaskPending,
"error": "",
"retry_count": 0,
"next_try_at": nil,
"started_at": nil,
"finished_at": nil,
"updated_at": time.Now(),
})
count = res.RowsAffected
return res.Error
})
return count, err
}
// CancelBatch 批量取消指定 ID 的排队/进行中下载任务。
func (r *StrmDownloadTaskRepository) CancelBatch(ctx context.Context, ids []string) (int64, error) {
if len(ids) == 0 {
return 0, nil
}
now := time.Now()
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).
Where("id IN ? AND status IN ?", ids, []string{model.StrmTaskPending, model.StrmTaskRunning}).
Updates(map[string]any{
"status": model.StrmTaskCanceled,
"error": "已批量取消",
"finished_at": now,
"updated_at": now,
})
count = res.RowsAffected
return res.Error
})
return count, err
}
// ClearDone 清空全部已完成下载任务。
func (r *StrmDownloadTaskRepository) ClearDone(ctx context.Context) (int64, error) {
res := r.db.WithContext(ctx).Where("status = ?", model.StrmTaskDone).Delete(&model.StrmDownloadTask{})
return res.RowsAffected, res.Error
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Unscoped().Where("status = ?", model.StrmTaskDone).Delete(&model.StrmDownloadTask{})
count = res.RowsAffected
return res.Error
})
return count, err
}
// ClearFinished 清空全部已完成与失败下载任务。
func (r *StrmDownloadTaskRepository) ClearFinished(ctx context.Context) (int64, error) {
res := r.db.WithContext(ctx).Where("status IN ?", []string{model.StrmTaskDone, model.StrmTaskFailed}).
Delete(&model.StrmDownloadTask{})
return res.RowsAffected, res.Error
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Unscoped().Where("status IN ?", []string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}).
Delete(&model.StrmDownloadTask{})
count = res.RowsAffected
return res.Error
})
return count, err
}
// ClearCanceled 清空全部已取消下载任务。
func (r *StrmDownloadTaskRepository) ClearCanceled(ctx context.Context) (int64, error) {
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Unscoped().Where("status = ?", model.StrmTaskCanceled).Delete(&model.StrmDownloadTask{})
count = res.RowsAffected
return res.Error
})
return count, err
}
// RetryAllFailed 把所有失败任务重置回待处理,清空错误与重试计数。
func (r *StrmDownloadTaskRepository) RetryAllFailed(ctx context.Context) (int64, error) {
res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).
Where("status = ?", model.StrmTaskFailed).
Updates(map[string]any{
"status": model.StrmTaskPending,
"error": "",
"retry_count": 0,
"next_try_at": nil,
"started_at": nil,
"finished_at": nil,
"updated_at": time.Now(),
})
return res.RowsAffected, res.Error
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).
Where("status = ?", model.StrmTaskFailed).
Updates(map[string]any{
"status": model.StrmTaskPending,
"error": "",
"retry_count": 0,
"next_try_at": nil,
"started_at": nil,
"finished_at": nil,
"updated_at": time.Now(),
})
count = res.RowsAffected
return res.Error
})
return count, err
}
// CancelPending 批量取消所有排队中的任务。
// CancelPending 批量取消所有排队中和进行中的任务。
func (r *StrmDownloadTaskRepository) CancelPending(ctx context.Context) (int64, error) {
now := time.Now()
res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).
Where("status = ?", model.StrmTaskPending).
Updates(map[string]any{
"status": model.StrmTaskCanceled,
"error": "已批量取消",
"finished_at": now,
"updated_at": now,
})
return res.RowsAffected, res.Error
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).
Where("status IN ?", []string{model.StrmTaskPending, model.StrmTaskRunning}).
Updates(map[string]any{
"status": model.StrmTaskCanceled,
"error": "已批量取消",
"finished_at": now,
"updated_at": now,
})
count = res.RowsAffected
return res.Error
})
return count, err
}
// CountActive 统计某同步目录下目标仍在排队/进行的任务数(用于去重)。
@@ -299,10 +452,28 @@ func (r *StrmDownloadTaskRepository) CountActive(ctx context.Context, syncPathID
return count
}
// GetActiveLocalPathMap 一次性获取某同步目录下正在排队或执行中的 local_path 集合,供同步时 O(1) 内存去重。
func (r *StrmDownloadTaskRepository) GetActiveLocalPathMap(ctx context.Context, syncPathID string) (map[string]bool, error) {
var paths []string
err := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).
Where("sync_path_id = ? AND status IN ?", syncPathID, []string{model.StrmTaskPending, model.StrmTaskRunning}).
Pluck("local_path", &paths).Error
if err != nil {
return nil, err
}
out := make(map[string]bool, len(paths))
for _, p := range paths {
out[p] = true
}
return out, nil
}
func (r *StrmDownloadTaskRepository) DeleteFinishedOlderThan(ctx context.Context, before time.Time) error {
return r.db.WithContext(ctx).Where("status IN ? AND finished_at < ?",
[]string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before).
Delete(&model.StrmDownloadTask{}).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Unscoped().Where("status IN ? AND finished_at < ?",
[]string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before).
Delete(&model.StrmDownloadTask{}).Error
})
}
// ─── StrmUploadTask ────────────────────────────────────────────────────────────
@@ -311,7 +482,21 @@ func (r *StrmDownloadTaskRepository) DeleteFinishedOlderThan(ctx context.Context
type StrmUploadTaskRepository struct{ db *gorm.DB }
func (r *StrmUploadTaskRepository) Create(ctx context.Context, t *model.StrmUploadTask) error {
return r.db.WithContext(ctx).Create(t).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Create(t).Error
})
}
func (r *StrmUploadTaskRepository) CreateInBatches(ctx context.Context, tasks []*model.StrmUploadTask, batchSize int) error {
if len(tasks) == 0 {
return nil
}
if batchSize <= 0 {
batchSize = 100
}
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).CreateInBatches(tasks, batchSize).Error
})
}
func (r *StrmUploadTaskRepository) FindByID(ctx context.Context, id string) (*model.StrmUploadTask, error) {
@@ -380,24 +565,29 @@ func (r *StrmUploadTaskRepository) CountByStatus(ctx context.Context) (map[strin
// ClaimPendingUpload picks the oldest pending task and marks it running.
// Returns (nil, nil) when the queue is empty.
func (r *StrmUploadTaskRepository) ClaimPendingUpload(ctx context.Context, limit int) ([]model.StrmUploadTask, error) {
strmClaimMu.Lock()
defer strmClaimMu.Unlock()
var rows []model.StrmUploadTask
err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()).
Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil {
return err
}
if len(rows) == 0 {
return nil
}
ids := make([]string, 0, len(rows))
now := time.Now()
for i := range rows {
ids = append(ids, rows[i].ID)
rows[i].Status = model.StrmTaskRunning
rows[i].StartedAt = &now
}
return tx.Model(&model.StrmUploadTask{}).Where("id IN ?", ids).
Updates(map[string]any{"status": model.StrmTaskRunning, "started_at": now}).Error
err := withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()).
Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil {
return err
}
if len(rows) == 0 {
return nil
}
ids := make([]string, 0, len(rows))
now := time.Now()
for i := range rows {
ids = append(ids, rows[i].ID)
rows[i].Status = model.StrmTaskRunning
rows[i].StartedAt = &now
}
return tx.Model(&model.StrmUploadTask{}).Where("id IN ?", ids).
Updates(map[string]any{"status": model.StrmTaskRunning, "started_at": now}).Error
})
})
if err != nil {
return nil, err
@@ -406,19 +596,113 @@ func (r *StrmUploadTaskRepository) ClaimPendingUpload(ctx context.Context, limit
}
func (r *StrmUploadTaskRepository) Update(ctx context.Context, t *model.StrmUploadTask) error {
return r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).Where("id = ?", t.ID).Updates(map[string]any{
"status": t.Status,
"error": t.Error,
"retry_count": t.RetryCount,
"next_try_at": t.NextTryAt,
"started_at": t.StartedAt,
"finished_at": t.FinishedAt,
"updated_at": time.Now(),
}).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).Where("id = ?", t.ID).Updates(map[string]any{
"status": t.Status,
"error": t.Error,
"retry_count": t.RetryCount,
"next_try_at": t.NextTryAt,
"started_at": t.StartedAt,
"finished_at": t.FinishedAt,
"updated_at": time.Now(),
}).Error
})
}
func (r *StrmUploadTaskRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmUploadTask{}).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmUploadTask{}).Error
})
}
// DeleteBatch 批量删除指定 ID 的上传任务。
func (r *StrmUploadTaskRepository) DeleteBatch(ctx context.Context, ids []string) (int64, error) {
if len(ids) == 0 {
return 0, nil
}
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Unscoped().Where("id IN ?", ids).Delete(&model.StrmUploadTask{})
count = res.RowsAffected
return res.Error
})
return count, err
}
// RetryBatch 批量重试指定 ID 的失败/已取消上传任务。
func (r *StrmUploadTaskRepository) RetryBatch(ctx context.Context, ids []string) (int64, error) {
if len(ids) == 0 {
return 0, nil
}
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).
Where("id IN ? AND status IN ?", ids, []string{model.StrmTaskFailed, model.StrmTaskCanceled}).
Updates(map[string]any{
"status": model.StrmTaskPending,
"error": "",
"retry_count": 0,
"next_try_at": nil,
"started_at": nil,
"finished_at": nil,
"updated_at": time.Now(),
})
count = res.RowsAffected
return res.Error
})
return count, err
}
// CancelBatch 批量取消指定 ID 的排队/进行中上传任务。
func (r *StrmUploadTaskRepository) CancelBatch(ctx context.Context, ids []string) (int64, error) {
if len(ids) == 0 {
return 0, nil
}
now := time.Now()
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).
Where("id IN ? AND status IN ?", ids, []string{model.StrmTaskPending, model.StrmTaskRunning}).
Updates(map[string]any{
"status": model.StrmTaskCanceled,
"error": "已批量取消",
"finished_at": now,
"updated_at": now,
})
count = res.RowsAffected
return res.Error
})
return count, err
}
// ClearCanceled 清空全部已取消上传任务。
func (r *StrmUploadTaskRepository) ClearCanceled(ctx context.Context) (int64, error) {
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Unscoped().Where("status = ?", model.StrmTaskCanceled).Delete(&model.StrmUploadTask{})
count = res.RowsAffected
return res.Error
})
return count, err
}
// CancelPending 批量取消所有排队中和进行中的任务。
func (r *StrmUploadTaskRepository) CancelPending(ctx context.Context) (int64, error) {
now := time.Now()
var count int64
err := withSQLiteBusyRetry(ctx, func() error {
res := r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).
Where("status IN ?", []string{model.StrmTaskPending, model.StrmTaskRunning}).
Updates(map[string]any{
"status": model.StrmTaskCanceled,
"error": "已批量取消",
"finished_at": now,
"updated_at": now,
})
count = res.RowsAffected
return res.Error
})
return count, err
}
// CountActive 统计某同步目录下目标仍在排队/进行的任务数(用于去重)。
@@ -431,8 +715,67 @@ func (r *StrmUploadTaskRepository) CountActive(ctx context.Context, syncPathID,
return count
}
func (r *StrmUploadTaskRepository) DeleteFinishedOlderThan(ctx context.Context, before time.Time) error {
return r.db.WithContext(ctx).Where("status IN ? AND finished_at < ?",
[]string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before).
Delete(&model.StrmUploadTask{}).Error
// GetActiveLocalPathMap 一次性获取某同步目录下正在排队或执行中的 local_path 集合,供同步时 O(1) 内存去重。
func (r *StrmUploadTaskRepository) GetActiveLocalPathMap(ctx context.Context, syncPathID string) (map[string]bool, error) {
var paths []string
err := r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).
Where("sync_path_id = ? AND status IN ?", syncPathID, []string{model.StrmTaskPending, model.StrmTaskRunning}).
Pluck("local_path", &paths).Error
if err != nil {
return nil, err
}
out := make(map[string]bool, len(paths))
for _, p := range paths {
out[p] = true
}
return out, nil
}
func (r *StrmUploadTaskRepository) DeleteFinishedOlderThan(ctx context.Context, before time.Time) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Unscoped().Where("status IN ? AND finished_at < ?",
[]string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before).
Delete(&model.StrmUploadTask{}).Error
})
}
// ─── StrmDirCache ─────────────────────────────────────────────────────────────
// StrmDirCacheRepository persists model.StrmDirCache.
type StrmDirCacheRepository struct{ db *gorm.DB }
func (r *StrmDirCacheRepository) ListBySyncPathID(ctx context.Context, syncPathID string) ([]model.StrmDirCache, error) {
var rows []model.StrmDirCache
err := r.db.WithContext(ctx).Where("sync_path_id = ?", syncPathID).Find(&rows).Error
return rows, err
}
func (r *StrmDirCacheRepository) Set(ctx context.Context, syncPathID, dirID, path string) error {
return withSQLiteBusyRetry(ctx, func() error {
var row model.StrmDirCache
err := r.db.WithContext(ctx).Where("sync_path_id = ? AND dir_id = ?", syncPathID, dirID).First(&row).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
row = model.StrmDirCache{
SyncPathID: syncPathID,
DirID: dirID,
Path: path,
}
return r.db.WithContext(ctx).Create(&row).Error
}
if err != nil {
return err
}
return r.db.WithContext(ctx).Model(&model.StrmDirCache{}).Where("id = ?", row.ID).Updates(map[string]any{
"path": path,
"updated_at": time.Now(),
}).Error
})
}
func (r *StrmDirCacheRepository) DeleteBySyncPathID(ctx context.Context, syncPathID string) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Unscoped().Where("sync_path_id = ?", syncPathID).Delete(&model.StrmDirCache{}).Error
})
}
+11 -20
View File
@@ -133,26 +133,17 @@ func (r *UserRepository) TouchLogin(ctx context.Context, id string) error {
})
}
// Delete removes a user (soft-delete via gorm.DeletedAt), releases the unique
// username, and drops Telegram bindings so future re-created users bind cleanly.
// Delete 物理删除用户并级联清理其关联记录。
func (r *UserRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var user model.User
if err := tx.Where("id = ?", id).First(&user).Error; err != nil {
return err
}
released := user.Username + "__deleted__" + time.Now().Format("20060102150405.000000000")
if len(released) > 64 {
sum := sha256.Sum256([]byte(user.ID + user.Username))
base := user.Username
if len(base) > 43 {
base = base[:43]
}
released = base + "__deleted__" + hex.EncodeToString(sum[:])[:10]
}
if err := tx.Model(&model.User{}).Where("id = ?", id).Update("username", released).Error; err != nil {
return err
}
return tx.Delete(&model.User{}, "id = ?", id).Error
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
_ = tx.Unscoped().Where("user_id = ?", id).Delete(&model.RefreshToken{})
_ = tx.Unscoped().Where("user_id = ?", id).Delete(&model.UserPermission{})
_ = tx.Unscoped().Where("user_id = ?", id).Delete(&model.PlayProfile{})
_ = tx.Unscoped().Where("user_id = ?", id).Delete(&model.PlaybackHistory{})
_ = tx.Unscoped().Where("user_id = ?", id).Delete(&model.Favorite{})
_ = tx.Unscoped().Where("user_id = ?", id).Delete(&model.UserDevice{})
return tx.Unscoped().Delete(&model.User{}, "id = ?", id).Error
})
})
}
+5 -4
View File
@@ -37,10 +37,11 @@ var ErrUnsupported = errors.New("unsupported cloud provider")
// FileEntry is one item in a cloud directory listing.
type FileEntry struct {
ID string `json:"id"` // provider-native file id
Name string `json:"name"`
IsDir bool `json:"is_dir"`
Size int64 `json:"size"`
ID string `json:"id"` // provider-native file id
Name string `json:"name"`
IsDir bool `json:"is_dir"`
Size int64 `json:"size"`
MTime int64 `json:"mtime,omitempty"`
// PickCode is 115-specific; other providers use ID directly.
PickCode string `json:"pick_code,omitempty"`
}
+43
View File
@@ -17,11 +17,20 @@ package cloud
import (
"context"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"github.com/ShukeBta/MMTL/internal/service/cloud115"
)
// OpenAPI115Provider 暴露 115 开放平台驱动接口。
type OpenAPI115Provider interface {
Provider
OpenClient() *cloud115.OpenClient
}
// openAPI115Provider 实现 Provider 接口:List 列目录、Resolve 用 pickcode
// 换下载直链(302 offload,无需代理)、Ping 探测根目录。
type openAPI115Provider struct {
@@ -61,6 +70,7 @@ func (p *openAPI115Provider) List(ctx context.Context, dirID string) ([]FileEntr
Name: f.FileName,
IsDir: f.Category == cloud115.TypeDir,
Size: f.FileSize,
MTime: f.Utime,
PickCode: f.PickCode,
})
}
@@ -98,6 +108,39 @@ func (p *openAPI115Provider) ResolveWithUA(ctx context.Context, fileRef, ua stri
// OpenClient 暴露底层客户端(token 刷新用)。
func (p *openAPI115Provider) OpenClient() *cloud115.OpenClient { return p.c }
// PutFileNamed 把本地元数据上传到 115 指定父目录(parentCID 为父目录 cid)。
// io.Reader 无法携带文件名,因此走独立的 named 上传接口。将内容落为临时文件后
// 重命名为目标文件名,再交给 115 上传(/open/upload/init 的 file_name 取真实文件名)。
func (p *openAPI115Provider) PutFileNamed(ctx context.Context, parentCID, fileName string, r io.Reader) error {
tmp, err := os.CreateTemp("", "mmtl-upload-*")
if err != nil {
return fmt.Errorf("115: 创建临时文件失败:%w", err)
}
tmpPath := tmp.Name()
defer func() {
_ = tmp.Close()
_ = os.Remove(tmpPath)
}()
if _, err := io.Copy(tmp, r); err != nil {
return fmt.Errorf("115: 写入临时文件失败:%w", err)
}
if err := tmp.Close(); err != nil {
return fmt.Errorf("115: 关闭临时文件失败:%w", err)
}
// 重命名为目标文件名,保证上传到 115 后保留原始文件名
if fileName != "" && fileName != filepath.Base(tmpPath) {
namedPath := filepath.Join(filepath.Dir(tmpPath), fileName)
if err := os.Rename(tmpPath, namedPath); err == nil {
tmpPath = namedPath
}
}
_, err = p.c.Upload(ctx, tmpPath, parentCID, "", "")
if err != nil {
return err
}
return nil
}
// RefreshToken 刷新访问令牌并返回新令牌;refresh_token 失效时返回
// cloud115.IsRefreshTokenDead(err) 为 true 的错误。
func (p *openAPI115Provider) RefreshToken(refreshToken string) (*cloud115.TokenData, error) {
+34 -9
View File
@@ -90,6 +90,7 @@ type RespBase struct {
Errno int `json:"errno"`
Message string `json:"message"`
Error string `json:"error"`
Count int64 `json:"count"`
Data json.RawMessage `json:"data"`
Raw json.RawMessage `json:"-"` // 原始响应体(外层附加字段用)
}
@@ -184,7 +185,7 @@ func (c *OpenClient) doJSON(ctx context.Context, method, rawURL string, form map
// refresh_token 刷新后重试一次。刷新失败或重试后仍失败才返回,
// 避免长时间同步因 token 过期而整体失败。
if isTokenCode(base.Code) {
if access && c.tryRefreshTokenLocked() {
if access && c.tryRefreshTokenLocked(ctx) {
continue
}
if access {
@@ -258,19 +259,43 @@ func (c *OpenClient) doAuthJSONWithUA(ctx context.Context, method, rawURL string
}
// tryRefreshTokenLocked 并发安全地刷新 access_token;成功返回 true(调用方
// 应使用内存中的新 token 重试原请求)。refresh_token 已失效时也会清空内存 token。
func (c *OpenClient) tryRefreshTokenLocked() bool {
// 应使用内存中的新 token 重试原请求)。
//
// 对"refresh_token 本身已失效/被吊销"(IsRefreshTokenDead,如 40140114/116/119/120)
// 这类不可恢复的错误直接放弃并清空内存 token(提示需重新授权)。
// 对其它失败(网络瞬时抖动、刷新接口可重试错误码等)做指数退避重试几次再放弃,
// 避免同步长任务中途 token 到期时恰好撞上一个短暂的刷新失败就整体失败。
func (c *OpenClient) tryRefreshTokenLocked(ctx context.Context) bool {
c.tokenMu.Lock()
defer c.tokenMu.Unlock()
token, err := c.RefreshToken(c.RefreshTokenStr)
if err != nil {
for attempt := 0; attempt < refreshAttempts; attempt++ {
token, err := c.RefreshToken(c.RefreshTokenStr)
if err == nil {
c.SetAuthToken(token.AccessToken, token.RefreshToken)
return true
}
if IsRefreshTokenDead(err) {
c.SetAuthToken("", "")
return false
}
// 可恢复失败:退避后重试。ctx 取消时立即放弃。
if attempt < refreshAttempts-1 {
select {
case <-ctx.Done():
return false
case <-time.After(refreshBackoff(attempt)):
}
}
return false
}
c.SetAuthToken(token.AccessToken, token.RefreshToken)
return true
return false
}
// refreshAttempts 是刷新 access_token 失败时的最大尝试次数(含首次)。
const refreshAttempts = 3
// refreshBackoff 返回第 attempt 次(从 0 计)刷新失败后的退避时长(指数退避)。
func refreshBackoff(attempt int) time.Duration {
return time.Duration(200*(1<<attempt)) * time.Millisecond // 200ms, 400ms
}
// IsThrottleCode 判断是否为限流错误码。
@@ -280,7 +305,7 @@ func IsThrottleCode(code int) bool {
func isTokenCode(code int) bool {
switch code {
case AccessTokenAuthFail, AccessAuthInvalid, AccessTokenExpiryCode, RefreshTokenInvalid:
case AccessTokenAuthFail, AccessAuthInvalid, AccessTokenExpiryCode, AccessTokenFormatInvalid, RefreshTokenInvalid:
return true
}
return false
@@ -375,3 +375,100 @@ func TestThrottleCodeHandling(t *testing.T) {
t.Fatal("code 770004 should trigger throttle status")
}
}
func TestRemoteFileDetailRelativePath(t *testing.T) {
rootCID := "3238787832374488117" // 影视库
// 场景 1:115 目录 paths 中只有祖先目录链,不包含自身
d1 := &RemoteFileDetail{
FileId: "3251154147730910635",
FileName: "出包王女",
Paths: []struct {
FileId string `json:"file_id"`
Name string `json:"file_name"`
}{
{FileId: "0", Name: "根目录"},
{FileId: "3238787832374488117", Name: "影视库"},
{FileId: "3238787913223892116", Name: "动漫"},
},
}
if got := d1.RelativePath(rootCID); got != "动漫/出包王女" {
t.Errorf("d1.RelativePath = %q, want %q", got, "动漫/出包王女")
}
// 场景 2:祖先中间目录,自身在 paths 末尾
d2 := &RemoteFileDetail{
FileId: "3238787913223892116",
FileName: "动漫",
Paths: []struct {
FileId string `json:"file_id"`
Name string `json:"file_name"`
}{
{FileId: "0", Name: "根目录"},
{FileId: "3238787832374488117", Name: "影视库"},
{FileId: "3238787913223892116", Name: "动漫"},
},
}
if got := d2.RelativePath(rootCID); got != "动漫" {
t.Errorf("d2.RelativePath = %q, want %q", got, "动漫")
}
// 场景 3:根同步目录自身
d3 := &RemoteFileDetail{
FileId: rootCID,
FileName: "影视库",
Paths: []struct {
FileId string `json:"file_id"`
Name string `json:"file_name"`
}{
{FileId: "0", Name: "根目录"},
{FileId: rootCID, Name: "影视库"},
},
}
if got := d3.RelativePath(rootCID); got != "" {
t.Errorf("d3.RelativePath = %q, want %q", got, "")
}
}
// TestFsListRefreshContinue 验证 access_token 在请求中途过期(40140126)时:
// 自动用 refresh_token 刷新得到新 token,然后对原请求重试成功(同步得以继续)。
func TestFsListRefreshContinue(t *testing.T) {
var filesCalls int
var refreshCalls int
mockAPI(t, func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/open/refreshToken":
refreshCalls++
w.Write([]byte(`{"state":true,"data":{"access_token":"at2","refresh_token":"rt2","expires_in":7200}}`))
case "/open/ufile/files":
filesCalls++
switch filesCalls {
case 1:
// 第一次用旧 access_token,返回过期错误,应触发刷新
w.Write([]byte(`{"state":false,"code":40140126,"message":"access_token 校验失败"}`))
default:
// 刷新后续请求应使用新 access_token
if got := r.Header.Get("Authorization"); got != "Bearer at2" {
t.Errorf("retried request auth = %q, want Bearer at2", got)
}
w.Write([]byte(`{"state":true,"path":[],"data":[{"fid":"200","fc":"1","fn":"a.mkv","fs":123,"pc":"pickA"}]}`))
}
default:
t.Errorf("unexpected path %s", r.URL.Path)
}
})
c := NewOpenClient("100195125", "at1", "rt1")
files, _, err := c.GetFsList(context.Background(), "0", 0, 100)
if err != nil {
t.Fatalf("expected sync to continue after refresh, got error: %v", err)
}
if filesCalls != 2 {
t.Fatalf("want 2 files calls (original + retried), got %d", filesCalls)
}
if refreshCalls == 0 {
t.Fatal("expected refresh_token to be used once")
}
if len(files) != 1 {
t.Fatalf("want 1 file, got %d", len(files))
}
}
+8 -7
View File
@@ -21,13 +21,14 @@ var (
const (
// 业务错误码
AccessTokenAuthFail = 40140126 // 访问过期,需刷新
AccessTokenExpiryCode = 40140125 // 访问过期,需刷新
AccessAuthInvalid = 40140124 // 访问无效,需刷新
RefreshTokenInvalid = 40140116 // 需重新授权
TokenRefreshFail = 40140121 // 刷新失败,可重试
RequestMaxLimitCode = 770004 // 访问频率过高
RequestRateLimitCode = 406 // 达到访问上限
AccessTokenAuthFail = 40140126 // 访问过期,需刷新
AccessTokenExpiryCode = 40140125 // 访问过期,需刷新
AccessAuthInvalid = 40140124 // 访问无效,需刷新
AccessTokenFormatInvalid = 40140123 // access_token 格式错误,需刷新
RefreshTokenInvalid = 40140116 // 需重新授权
TokenRefreshFail = 40140121 // 刷新失败,可重试
RequestMaxLimitCode = 770004 // 访问频率过高
RequestRateLimitCode = 406 // 达到访问上限
// 刷新 token 的错误码
RefreshTokenFormatInvalid = 40140114
+72
View File
@@ -89,6 +89,36 @@ func (c *OpenClient) GetFsList(ctx context.Context, cid string, offset, limit in
return files, strings.Join(pathStr, "/"), nil
}
// GetFsListFlat 递归扁平化列出 cid 下的所有文件(跨越所有子目录,不包含文件夹节点),并返回文件列表与该树下的总文件数。
// 类似于 QMediaSync 的 115 扁平化批量拉取机制,极大地降低多层级子目录下的 API 请求次数。
func (c *OpenClient) GetFsListFlat(ctx context.Context, cid string, offset, limit int) ([]RemoteFile, int64, error) {
if cid == "" {
cid = "0"
}
if limit <= 0 {
limit = 1150
}
params := map[string]string{
"cid": cid,
"limit": fmt.Sprint(limit),
"offset": fmt.Sprint(offset),
"cur": "0",
"show_dir": "0",
}
resp, err := c.doAuthJSON(ctx, "GET", ProAPIBase+"/open/ufile/files", params, 2)
if err != nil {
return nil, 0, err
}
if !resp.State {
return nil, 0, NewOpenAPIResponseError(resp.Code, resp.Errno, resp.Message, resp.Error, "115 接口调用失败")
}
files, err := openList[RemoteFile](resp.Data)
if err != nil {
return nil, 0, fmt.Errorf("115: 解析文件列表失败:%w", err)
}
return files, resp.Count, nil
}
// GetFsDetailByCid 查询文件(夹)详情。
func (c *OpenClient) GetFsDetailByCid(ctx context.Context, fileId string) (*RemoteFileDetail, error) {
params := map[string]string{"file_id": fileId}
@@ -113,6 +143,48 @@ type RemoteFileDetail struct {
} `json:"paths"`
}
// RelativePath 计算该目录相对于根同步目录(rootCID)的相对路径。
func (d *RemoteFileDetail) RelativePath(rootCID string) string {
if d == nil {
return ""
}
if rootCID == "" {
rootCID = "0"
}
if d.FileId == rootCID {
return ""
}
rootIdx := -1
for i, p := range d.Paths {
if p.FileId == rootCID {
rootIdx = i
break
}
}
var segments []string
start := 0
if rootIdx >= 0 {
start = rootIdx + 1
} else if len(d.Paths) > 0 && (d.Paths[0].FileId == "0" || d.Paths[0].FileId == "") {
start = 1
}
hasSelf := false
for i := start; i < len(d.Paths); i++ {
if d.Paths[i].FileId == d.FileId {
hasSelf = true
}
name := strings.TrimSpace(d.Paths[i].Name)
if name != "" {
segments = append(segments, name)
}
}
// 若 115 返回的 paths 祖先链未包含当前目录自身,则将其自身目录名 FileName 补在末尾
if !hasSelf && strings.TrimSpace(d.FileName) != "" && d.FileId != rootCID {
segments = append(segments, strings.TrimSpace(d.FileName))
}
return strings.Join(segments, "/")
}
// ─── 下载直链 ──────────────────────────────────────────────────────────────────
type downloadURLData struct {
+367
View File
@@ -0,0 +1,367 @@
// 阿里云 OSS multipart 分片上传(用于 115 元数据上传直传)。
// 使用 115 下发的临时 STS 凭证,将本地文件分片上传到 OSS,并经 complete 回调
// 通知 115 完成落盘。参考 QMediaSync 的 OSSMultipartUploader 实现。
package cloud115
import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"sort"
"github.com/aliyun/alibabacloud-oss-go-sdk-v2/oss"
"github.com/aliyun/alibabacloud-oss-go-sdk-v2/oss/credentials"
)
const (
defaultMultipartPartSize int64 = 32 * 1024 * 1024
multipartPartAlign int64 = 1024 * 1024
maxMultipartParts int64 = 9999
maxMultipartPartSize int64 = 5 * 1024 * 1024 * 1024
)
type ossMultipartClient interface {
InitiateMultipartUpload(context.Context, *oss.InitiateMultipartUploadRequest, ...func(*oss.Options)) (*oss.InitiateMultipartUploadResult, error)
UploadPart(context.Context, *oss.UploadPartRequest, ...func(*oss.Options)) (*oss.UploadPartResult, error)
ListParts(context.Context, *oss.ListPartsRequest, ...func(*oss.Options)) (*oss.ListPartsResult, error)
CompleteMultipartUpload(context.Context, *oss.CompleteMultipartUploadRequest, ...func(*oss.Options)) (*oss.CompleteMultipartUploadResult, error)
AbortMultipartUpload(context.Context, *oss.AbortMultipartUploadRequest, ...func(*oss.Options)) (*oss.AbortMultipartUploadResult, error)
}
// OSSMultipartUploader 封装 OSS multipart 上传。
type OSSMultipartUploader struct {
client ossMultipartClient
}
// OSSMultipartUploadInput 是 multipart 上传输入。
type OSSMultipartUploadInput struct {
Bucket string
Object string
Callback string
CallbackVar string
FilePath string
FileSize int64
UploadId string
PartSize int64
PartRetryMax int
refreshClient func(context.Context) (ossMultipartClient, error)
}
// OSSMultipartUploadResult 是 multipart 上传后的结果。
type OSSMultipartUploadResult struct {
CallbackResult map[string]any
UploadId string
PartSize int64
TotalParts int
UploadedBytes int64
UploadedParts int
}
// CalculateMultipartPartSize 计算 OSS multipart 分片大小与分片数量。
func CalculateMultipartPartSize(fileSize int64) (int64, int, error) {
if fileSize < 0 {
return 0, 0, fmt.Errorf("文件大小不能为负数:%d", fileSize)
}
partSize := defaultMultipartPartSize
minPartSize := ceilDiv(fileSize, maxMultipartParts)
if minPartSize > partSize {
partSize = roundUp(minPartSize, multipartPartAlign)
}
if partSize > maxMultipartPartSize {
return 0, 0, fmt.Errorf("文件过大,所需分片大小 %d 超过 OSS 上限 %d", partSize, maxMultipartPartSize)
}
totalParts := int(ceilDiv(fileSize, partSize))
if totalParts == 0 {
totalParts = 1
}
if int64(totalParts) > maxMultipartParts {
return 0, 0, fmt.Errorf("分片数量 %d 超过上限 %d", totalParts, maxMultipartParts)
}
return partSize, totalParts, nil
}
// NewOSSMultipartUploader 创建 OSS multipart 上传器。
func NewOSSMultipartUploader(endpoint, accessKeyId, accessKeySecret, securityToken string) *OSSMultipartUploader {
return &OSSMultipartUploader{client: newOSSMultipartClient(endpoint, accessKeyId, accessKeySecret, securityToken)}
}
func newOSSMultipartClient(endpoint, accessKeyId, accessKeySecret, securityToken string) ossMultipartClient {
cfg := oss.LoadDefaultConfig().
WithCredentialsProvider(credentials.NewStaticCredentialsProvider(accessKeyId, accessKeySecret, securityToken)).
WithRegion("cn-shenzhen").
WithEndpoint(endpoint)
return oss.NewClient(cfg)
}
// UploadFile 上传文件并完成 OSS multipart,返回 complete callback 结果。
func (u *OSSMultipartUploader) UploadFile(ctx context.Context, input OSSMultipartUploadInput) (map[string]any, error) {
result, err := u.UploadFileWithResult(ctx, input)
if err != nil {
return nil, err
}
return result.CallbackResult, nil
}
// UploadFileWithResult 上传文件并返回 multipart 结果。
func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input OSSMultipartUploadInput) (OSSMultipartUploadResult, error) {
if input.PartRetryMax <= 0 {
input.PartRetryMax = 3
}
partSize := input.PartSize
totalParts := 0
var err error
if partSize <= 0 {
partSize, totalParts, err = CalculateMultipartPartSize(input.FileSize)
if err != nil {
return OSSMultipartUploadResult{}, err
}
} else {
totalParts = int(ceilDiv(input.FileSize, partSize))
}
uploadId := input.UploadId
if uploadId == "" {
initResult, err := u.client.InitiateMultipartUpload(ctx, &oss.InitiateMultipartUploadRequest{
Bucket: oss.Ptr(input.Bucket),
Key: oss.Ptr(input.Object),
RequestCommon: oss.RequestCommon{
Parameters: map[string]string{"sequential": "1"},
},
})
if err != nil {
return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 失败:%w", err)
}
if initResult.UploadId == nil || *initResult.UploadId == "" {
return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 返回空 upload_id")
}
uploadId = *initResult.UploadId
}
existingPartMap := make(map[int32]int64)
existingParts, err := u.ListUploadedParts(ctx, input.Bucket, input.Object, uploadId)
if err == nil {
for _, part := range existingParts {
existingPartMap[part.PartNumber] = part.Size
}
}
file, err := os.Open(input.FilePath)
if err != nil {
return OSSMultipartUploadResult{}, fmt.Errorf("打开待上传文件失败:%w", err)
}
defer file.Close()
var uploadedBytes int64
uploadedParts := 0
completeParts := make([]oss.UploadPart, 0, totalParts)
for partNumber := 1; partNumber <= totalParts; partNumber++ {
offset := int64(partNumber-1) * partSize
length := minInt64(partSize, input.FileSize-offset)
if length < 0 {
length = 0
}
if existingSize, ok := existingPartMap[int32(partNumber)]; ok && existingSize == length {
uploadedBytes += length
uploadedParts++
}
etag, err := u.uploadPartWithRetry(ctx, input, uploadId, int32(partNumber), file, offset, length)
if err != nil {
return OSSMultipartUploadResult{}, err
}
uploadedBytes += length
uploadedParts++
completeParts = append(completeParts, oss.UploadPart{
PartNumber: int32(partNumber),
ETag: oss.Ptr(etag),
})
}
sort.Slice(completeParts, func(i, j int) bool {
return completeParts[i].PartNumber < completeParts[j].PartNumber
})
// 115 下发的 callback / callback_var 是 JSON 字符串,而 OSS CompleteMultipartUpload
// 要求 callback 参数为 Base64 编码后的 JSON,否则报 "The callback configuration is
// not base64 encoded"。这里把两者转为 Base64 后再提交(参考 QMediaSync 的
// BuildOSSCallbackHeaders)。
cb := input.Callback
cbVar := input.CallbackVar
if cb == "" {
return OSSMultipartUploadResult{}, errors.New("OSS callback 为空")
}
if !json.Valid([]byte(cb)) {
return OSSMultipartUploadResult{}, errors.New("解析 callback 失败:不是合法 JSON")
}
if cbVar == "" {
cbVar = "{}"
}
if !json.Valid([]byte(cbVar)) {
return OSSMultipartUploadResult{}, errors.New("解析 callback_var 失败:不是合法 JSON")
}
completeResult, err := u.client.CompleteMultipartUpload(ctx, &oss.CompleteMultipartUploadRequest{
Bucket: oss.Ptr(input.Bucket),
Key: oss.Ptr(input.Object),
UploadId: oss.Ptr(uploadId),
CompleteMultipartUpload: &oss.CompleteMultipartUpload{
Parts: completeParts,
},
Callback: oss.Ptr(base64.StdEncoding.EncodeToString([]byte(cb))),
CallbackVar: oss.Ptr(base64.StdEncoding.EncodeToString([]byte(cbVar))),
})
if err != nil {
return OSSMultipartUploadResult{}, fmt.Errorf("完成 OSS multipart 失败:%w", err)
}
return OSSMultipartUploadResult{
CallbackResult: completeResult.CallbackResult,
UploadId: uploadId,
PartSize: partSize,
TotalParts: totalParts,
UploadedBytes: uploadedBytes,
UploadedParts: uploadedParts,
}, nil
}
// ListUploadedParts 查询 OSS 已上传分片。
func (u *OSSMultipartUploader) ListUploadedParts(ctx context.Context, bucket, object, uploadId string) ([]struct {
PartNumber int32
Size int64
}, error) {
parts := []struct {
PartNumber int32
Size int64
}{}
result, err := u.client.ListParts(ctx, &oss.ListPartsRequest{
Bucket: oss.Ptr(bucket),
Key: oss.Ptr(object),
UploadId: oss.Ptr(uploadId),
MaxParts: 1000,
})
if err != nil {
return nil, fmt.Errorf("查询 OSS 已上传分片失败:%w", err)
}
for _, part := range result.Parts {
parts = append(parts, struct {
PartNumber int32
Size int64
}{PartNumber: part.PartNumber, Size: part.Size})
}
return parts, nil
}
func (u *OSSMultipartUploader) uploadPartWithRetry(
ctx context.Context,
input OSSMultipartUploadInput,
uploadId string,
partNumber int32,
file *os.File,
offset, length int64,
) (string, error) {
var lastErr error
for attempt := 0; attempt < input.PartRetryMax; attempt++ {
reader := io.NewSectionReader(file, offset, length)
result, err := u.client.UploadPart(ctx, &oss.UploadPartRequest{
Bucket: oss.Ptr(input.Bucket),
Key: oss.Ptr(input.Object),
PartNumber: partNumber,
UploadId: oss.Ptr(uploadId),
Body: reader,
ContentLength: oss.Ptr(length),
})
if err == nil {
if result.ETag == nil || *result.ETag == "" {
return "", fmt.Errorf("OSS part %d 返回空 ETag", partNumber)
}
return *result.ETag, nil
}
lastErr = err
if attempt < input.PartRetryMax-1 && input.refreshClient != nil {
refreshed, refreshErr := input.refreshClient(ctx)
if refreshErr != nil {
lastErr = refreshErr
continue
}
u.client = refreshed
}
}
return "", fmt.Errorf("上传 OSS part %d 失败:%w", partNumber, lastErr)
}
// ParseCompleteCallbackResult 校验并解析 OSS complete 后的 115 callback 结果。
func ParseCompleteCallbackResult(result map[string]any) (UploadCompleteResult, error) {
if result == nil {
return UploadCompleteResult{}, errors.New("OSS complete callback 结果为空")
}
if state, ok := result["state"].(bool); ok && !state {
return UploadCompleteResult{}, fmt.Errorf("115 callback 返回失败:%s", anyToString(result["message"]))
}
if message := anyToString(result["message"]); message != "" {
return UploadCompleteResult{}, fmt.Errorf("115 callback 返回错误:%s", message)
}
data, ok := result["data"].(map[string]any)
if !ok {
return UploadCompleteResult{}, errors.New("115 callback 缺少 data")
}
complete := UploadCompleteResult{
FileId: anyToString(data["file_id"]),
PickCode: anyToString(data["pick_code"]),
ParentId: anyToString(data["parent_id"]),
Sha1: anyToString(data["sha1"]),
Size: anyToInt64(data["size"]),
Mtime: anyToInt64(data["mtime"]),
}
if complete.FileId == "" || complete.PickCode == "" {
return UploadCompleteResult{}, errors.New("115 callback 缺少 file_id/pick_code")
}
return complete, nil
}
func ceilDiv(n, d int64) int64 {
if d <= 0 {
return 0
}
if n <= 0 {
return 0
}
return (n + d - 1) / d
}
func roundUp(n, align int64) int64 {
if align <= 0 {
return n
}
return ceilDiv(n, align) * align
}
func minInt64(a, b int64) int64 {
if a < b {
return a
}
return b
}
func anyToInt64(v any) int64 {
switch t := v.(type) {
case string:
var n int64
fmt.Sscanf(t, "%d", &n)
return n
case float64:
return int64(t)
case int64:
return t
case int:
return int64(t)
default:
return 0
}
}
func anyToString(v any) string {
if s, ok := v.(string); ok {
return s
}
return fmt.Sprintf("%v", v)
}
+7 -2
View File
@@ -26,10 +26,15 @@ var (
executorOnce sync.Once
)
// GetGlobalExecutor 获取全局队列执行器单例(默认 QPS=2, QPM=120, QPH=6000,保障 115 API 调用安全不超频)。
// GetGlobalExecutor 获取全局队列执行器单例(默认 QPS=3, QPM=200, QPH=12000,保障 115 API 调用安全不超频)。
//
// 历史教训:QPS 提到 8 后,下载换直链接口(/open/ufile/downurl,WAF 重点盯防对象)
// 瞬时突发撞上 115 风控,返回阿里云 405 阻断页(HTTP 405),导致全量同步失败。
// 因此回调到 3——这是经过实测的安全上限:宁慢勿触发风控,一旦 405 冷却 180 秒,
// 整体吞吐反而更低。下载实际走 CDN 不受此限速影响,瓶颈仅在换链环节。
func GetGlobalExecutor() *QueueExecutor {
executorOnce.Do(func() {
globalExecutor = NewQueueExecutor(2, 120, 6000)
globalExecutor = NewQueueExecutor(3, 200, 12000)
})
return globalExecutor
}
+49
View File
@@ -0,0 +1,49 @@
package cloud115
import (
"crypto/sha1"
"encoding/hex"
"io"
"os"
)
// FileSHA1 计算文件完整 SHA1(小写 hex)。
func FileSHA1(path string) (string, error) {
f, err := os.Open(path)
if err != nil {
return "", err
}
defer f.Close()
h := sha1.New()
if _, err := io.Copy(h, f); err != nil {
return "", err
}
return hex.EncodeToString(h.Sum(nil)), nil
}
// FileSHA1Partial 计算文件 [start,end](含)字节区间的 SHA1(小写 hex)。
// 用于 115 上传二次签名按 sign_check 指定的区间重算哈希。
func FileSHA1Partial(path string, start, end int64) (string, error) {
if start < 0 {
start = 0
}
if end < start {
end = start
}
f, err := os.Open(path)
if err != nil {
return "", err
}
defer f.Close()
if _, err := f.Seek(start, io.SeekStart); err != nil {
return "", err
}
length := end - start + 1
h := sha1.New()
// io.CopyN 在文件不足 length 字节时会返回 io.EOF,导致小文件(如小于 128 KiB 的
// 元数据图片)无法上传。这里只拷贝实际读到的字节,文件尾对齐到区间终点即可。
if _, err := io.CopyN(h, f, length); err != nil && err != io.EOF {
return "", err
}
return hex.EncodeToString(h.Sum(nil)), nil
}
+355
View File
@@ -0,0 +1,355 @@
// 115 网盘元数据上传能力:115 开放平台调度 + 阿里云 OSS 直传。
// 参考 QMediaSync 的上传流程实现:
//
// POST /open/upload/init 上传初始化/秒传调度(含二次签名)
// GET /open/upload/get_token 获取 OSS 临时上传凭证(STS)
// OSS multipart 分片直传 + callback 完成
//
// 上传目标父目录为 115 目录 ID(cid),而非路径字符串。
package cloud115
import (
"context"
"encoding/json"
"errors"
"fmt"
"os"
"strconv"
"strings"
)
// 115 上传状态码。
const (
UploadInitStatusNeedUpload = 1 // 需要真实上传
UploadInitStatusRapidUploaded = 2 // 秒传成功
UploadInitStatusSignFailed = 6 // 签名验证失败
UploadInitStatusNeedSign = 7 // 需要二次签名
UploadInitStatusSignRejected = 8 // 签名认证失败
)
// UploadInitRequest 是 /open/upload/init 的结构化请求。
type UploadInitRequest struct {
FileName string
FileSize int64
ParentFileId string
FileSha1 string
Preid string
PickCode string
TopUpload string
SignKey string
SignVal string
}
// UploadInitResult 是 /open/upload/init 的调度结果。
type UploadInitResult struct {
PickCode string
Status int
FileId string
Target string
Bucket string
Object string
SignKey string
SignCheck string
Callback UploadResultCallBack
}
type uploadScheduleAPIResult struct {
PickCode string `json:"pick_code"`
Status int `json:"status"`
FileId string `json:"file_id"`
Target string `json:"target"`
Version string `json:"version"`
Bucket string `json:"bucket"`
Object string `json:"object"`
SignKey string `json:"sign_key"`
SignCheck string `json:"sign_check"`
Callback json.RawMessage `json:"callback"`
}
// UploadResultCallBack 是 init 返回给 OSS complete 使用的 callback 内容。
type UploadResultCallBack struct {
Callback string `json:"callback"`
CallbackVar string `json:"callback_var"`
}
// UploadToken 是 /open/upload/get_token 返回的 OSS STS 临时凭证。
type UploadToken struct {
Endpoint string `json:"endpoint"`
AccessKeySecret string `json:"AccessKeySecret"`
AccessKeySecrett string `json:"AccessKeySecrett"`
SecurityToken string `json:"SecurityToken"`
Expiration string `json:"Expiration"`
AccessKeyId string `json:"AccessKeyId"`
}
func (token *UploadToken) normalize() {
if token == nil {
return
}
if token.AccessKeySecret == "" {
token.AccessKeySecret = token.AccessKeySecrett
}
}
// UploadCompleteResult 是 OSS complete callback 成功后的远端文件定位结果。
type UploadCompleteResult struct {
FileId string
PickCode string
ParentId string
Sha1 string
Size int64
Mtime int64
}
// SignCheckRange 是 115 二次认证要求的闭区间 [start,end]。
type SignCheckRange struct {
Start int64
End int64
}
// UploadInit 调用 115 上传初始化/秒传调度接口。
func (c *OpenClient) UploadInit(ctx context.Context, input UploadInitRequest) (*UploadInitResult, error) {
params := buildUploadInitForm(input)
resp, err := c.doAuthJSON(ctx, "POST", ProAPIBase+"/open/upload/init", params, 2)
if err != nil {
return nil, err
}
var raw uploadScheduleAPIResult
if err := json.Unmarshal(resp.Data, &raw); err != nil {
return nil, fmt.Errorf("115: 解析 upload/init 结果失败:%w", err)
}
callback, err := decodeUploadCallback(raw.Callback)
if err != nil {
return nil, err
}
return &UploadInitResult{
PickCode: raw.PickCode,
Status: raw.Status,
FileId: raw.FileId,
Target: raw.Target,
Bucket: raw.Bucket,
Object: raw.Object,
SignKey: raw.SignKey,
SignCheck: raw.SignCheck,
Callback: callback,
}, nil
}
func buildUploadInitForm(input UploadInitRequest) map[string]string {
topUpload := input.TopUpload
if topUpload == "" {
topUpload = "0"
}
params := map[string]string{
"file_name": input.FileName,
"file_size": strconv.FormatInt(input.FileSize, 10),
"target": fmt.Sprintf("U_1_%s", input.ParentFileId),
"fileid": input.FileSha1,
"preid": input.Preid,
"topupload": topUpload,
}
if input.PickCode != "" {
params["pick_code"] = input.PickCode
}
if input.SignKey != "" && input.SignVal != "" {
params["sign_key"] = input.SignKey
params["sign_val"] = input.SignVal
}
return params
}
func decodeUploadCallback(raw json.RawMessage) (UploadResultCallBack, error) {
if len(raw) == 0 || string(raw) == "null" {
return UploadResultCallBack{}, nil
}
if raw[0] == '[' {
var callbacks []UploadResultCallBack
if err := json.Unmarshal(raw, &callbacks); err != nil {
return UploadResultCallBack{}, err
}
if len(callbacks) == 0 {
return UploadResultCallBack{}, nil
}
return callbacks[0], nil
}
var callback UploadResultCallBack
if err := json.Unmarshal(raw, &callback); err != nil {
return UploadResultCallBack{}, err
}
return callback, nil
}
func parseSignCheckRange(value string) (SignCheckRange, error) {
parts := strings.Split(value, "-")
if len(parts) != 2 {
return SignCheckRange{}, fmt.Errorf("sign_check 格式错误:%s", value)
}
start, err := strconv.ParseInt(strings.TrimSpace(parts[0]), 10, 64)
if err != nil {
return SignCheckRange{}, err
}
end, err := strconv.ParseInt(strings.TrimSpace(parts[1]), 10, 64)
if err != nil {
return SignCheckRange{}, err
}
if start < 0 || end < start {
return SignCheckRange{}, fmt.Errorf("sign_check 范围非法:%s", value)
}
return SignCheckRange{Start: start, End: end}, nil
}
// GetUploadToken 获取 115 下发的 OSS 临时上传凭证。
func (c *OpenClient) GetUploadToken(ctx context.Context) (*UploadToken, error) {
resp, err := c.doAuthJSON(ctx, "GET", ProAPIBase+"/open/upload/get_token", nil, 2)
if err != nil {
return nil, err
}
var token UploadToken
if err := json.Unmarshal(resp.Data, &token); err != nil {
return nil, fmt.Errorf("115: 解析 get_token 结果失败:%w", err)
}
token.normalize()
return &token, nil
}
// Upload 上传单个本地文件到 115 指定父目录(cid),返回成功后的远端文件信息。
// filePath 必须是落到磁盘的真实文件路径(调用方负责把 io.Reader 落盘为临时文件)。
func (c *OpenClient) Upload(ctx context.Context, filePath, parentCID, signKey, signVal string) (*UploadCompleteResult, error) {
fileSize := fileSizeOf(filePath)
if fileSize < 0 {
return nil, fmt.Errorf("115: 无法获取文件大小:%s", filePath)
}
fileSha1, err := FileSHA1(filePath)
if err != nil {
return nil, fmt.Errorf("115: 计算文件 SHA1 失败:%w", err)
}
preSha1, err := FileSHA1Partial(filePath, 0, 128*1024-1)
if err != nil {
return nil, fmt.Errorf("115: 计算文件前 128 KiB SHA1 失败:%w", err)
}
request := UploadInitRequest{
FileName: baseNameOf(filePath),
FileSize: fileSize,
ParentFileId: parentCID,
FileSha1: fileSha1,
Preid: preSha1,
TopUpload: "0",
SignKey: signKey,
SignVal: signVal,
}
initResult, err := c.UploadInit(ctx, request)
if err != nil {
return nil, fmt.Errorf("115: 上传初始化失败:%w", err)
}
status := initResult.Status
if status == UploadInitStatusNeedSign {
// 二次签名:按 sign_check 指定区间重算 sha1
rng, err := parseSignCheckRange(initResult.SignCheck)
if err != nil {
return nil, err
}
signValue, err := FileSHA1Partial(filePath, rng.Start, rng.End)
if err != nil {
return nil, err
}
request.SignKey = initResult.SignKey
request.SignVal = signValue
initResult, err = c.UploadInit(ctx, request)
if err != nil {
return nil, fmt.Errorf("115: 上传二次签名失败:%w", err)
}
status = initResult.Status
}
switch status {
case UploadInitStatusRapidUploaded:
// 秒传成功
return &UploadCompleteResult{FileId: initResult.FileId, PickCode: initResult.PickCode}, nil
case UploadInitStatusSignFailed:
return nil, fmt.Errorf("115: 签名验证后失败")
case UploadInitStatusSignRejected:
return nil, fmt.Errorf("115: 签名认证失败")
case UploadInitStatusNeedUpload:
// 真实上传:OSS multipart
default:
return &UploadCompleteResult{FileId: initResult.FileId, PickCode: initResult.PickCode}, nil
}
if initResult.Bucket == "" || initResult.Object == "" {
return nil, fmt.Errorf("115: upload/init 缺少 bucket/object 信息")
}
token, err := c.GetUploadToken(ctx)
if err != nil {
return nil, fmt.Errorf("115: 获取上传凭证失败:%w", err)
}
if token == nil || token.Endpoint == "" || token.AccessKeyId == "" || token.AccessKeySecret == "" {
return nil, fmt.Errorf("115: 上传凭证不完整")
}
uploader := NewOSSMultipartUploader(token.Endpoint, token.AccessKeyId, token.AccessKeySecret, token.SecurityToken)
result, err := uploader.UploadFile(ctx, OSSMultipartUploadInput{
Bucket: initResult.Bucket,
Object: initResult.Object,
Callback: initResult.Callback.Callback,
CallbackVar: initResult.Callback.CallbackVar,
FilePath: filePath,
FileSize: fileSize,
refreshClient: func(ctx context.Context) (ossMultipartClient, error) {
refreshed, rerr := c.GetUploadToken(ctx)
if rerr != nil || refreshed == nil {
return nil, rerr
}
return newOSSMultipartClient(refreshed.Endpoint, refreshed.AccessKeyId, refreshed.AccessKeySecret, refreshed.SecurityToken), nil
},
})
if err != nil {
return nil, fmt.Errorf("115: OSS 上传失败:%w", err)
}
complete, err := ParseCompleteCallbackResult(result)
if err != nil {
return nil, err
}
return &complete, nil
}
// MkDir 在 115 的 parentCid 下创建目录,返回新目录 cid。
func (c *OpenClient) MkDir(ctx context.Context, parentCID, name string) (string, error) {
params := map[string]string{
"cname": name,
"pid": parentCID,
}
resp, err := c.doAuthJSON(ctx, "POST", ProAPIBase+"/open/folder/add", params, 2)
if err != nil {
return "", err
}
// /open/folder/add 结构:{ aid, cid, fid, name, pid, ... },单一对象
var r struct {
Cid string `json:"cid"`
}
if err := json.Unmarshal(resp.Data, &r); err != nil {
return "", fmt.Errorf("115: 解析 folder/add 结果失败:%w", err)
}
if r.Cid == "" {
return "", errors.New("115: folder/add 未返回 cid")
}
return r.Cid, nil
}
func fileSizeOf(path string) int64 {
info, err := os.Stat(path)
if err != nil {
return -1
}
if info.IsDir() {
return -1
}
return info.Size()
}
func baseNameOf(path string) string {
s := path
for i := len(s) - 1; i >= 0; i-- {
if s[i] == '/' || s[i] == '\\' {
return s[i+1:]
}
}
return s
}
+197
View File
@@ -0,0 +1,197 @@
package cloud115
import (
"context"
"encoding/base64"
"io"
"os"
"path/filepath"
"testing"
"github.com/aliyun/alibabacloud-oss-go-sdk-v2/oss"
)
func TestFileSHA1(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "a.txt")
if err := os.WriteFile(path, []byte("hello"), 0o644); err != nil {
t.Fatal(err)
}
sum, err := FileSHA1(path)
if err != nil {
t.Fatal(err)
}
// sha1("hello") = aaf4c61ddcc5e8a2dabede0f3b482cd9aea9434d
if sum != "aaf4c61ddcc5e8a2dabede0f3b482cd9aea9434d" {
t.Errorf("unexpected sha1: %s", sum)
}
}
func TestFileSHA1Partial(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "b.txt")
// 10 bytes: "0123456789"
if err := os.WriteFile(path, []byte("0123456789"), 0o644); err != nil {
t.Fatal(err)
}
// bytes [2,4] = "234"
sum, err := FileSHA1Partial(path, 2, 4)
if err != nil {
t.Fatal(err)
}
if sum != "0ec09ef9836da03f1add21e3ef607627e687e790" {
t.Errorf("unexpected partial sha1: %s", sum)
}
}
// TestFileSHA1PartialSmallerThanWindow 回归测试:经典 bug 是 io.CopyN 在文件不足
// length 字节时返回 io.EOF。115 上传固定用 [0,128*1024-1] 窗口计算 preid,导致所有
// 小于 128 KiB 的元数据文件(如海报/缩略图)上传必然失败。
func TestFileSHA1PartialSmallerThanWindow(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "small.bin")
// 6 字节小文件,不足 128 KiB 窗口
if err := os.WriteFile(path, []byte("abcdef"), 0o644); err != nil {
t.Fatal(err)
}
sum, err := FileSHA1Partial(path, 0, 128*1024-1)
if err != nil {
t.Fatalf("compute partial sha1 for small file should not fail: %v", err)
}
// 应等于整个文件(6 字节)的 sha1
if sum != "1f8ac10f23c5b5bc1167bda84b833e5c057a77d2" {
t.Errorf("unexpected partial sha1: %s", sum)
}
}
func TestParseSignCheckRange(t *testing.T) {
rng, err := parseSignCheckRange("0-131071")
if err != nil {
t.Fatal(err)
}
if rng.Start != 0 || rng.End != 131071 {
t.Errorf("unexpected range: %+v", rng)
}
if _, err := parseSignCheckRange("bad"); err == nil {
t.Error("expected error for bad range")
}
if _, err := parseSignCheckRange("100-50"); err == nil {
t.Error("expected error for end<start")
}
}
func TestCalculateMultipartPartSize(t *testing.T) {
// small file: 1 MiB -> partSize 32MiB, 1 part
ps, parts, err := CalculateMultipartPartSize(1 << 20)
if err != nil {
t.Fatal(err)
}
if ps != defaultMultipartPartSize {
t.Errorf("partSize=%d, want %d", ps, defaultMultipartPartSize)
}
if parts != 1 {
t.Errorf("parts=%d, want 1", parts)
}
// zero-size -> 1 part
_, parts, err = CalculateMultipartPartSize(0)
if err != nil {
t.Fatal(err)
}
if parts != 1 {
t.Errorf("zero-size parts=%d, want 1", parts)
}
// negative -> error
if _, _, err := CalculateMultipartPartSize(-1); err == nil {
t.Error("expected error for negative size")
}
}
func TestBaseNameOf(t *testing.T) {
if got := baseNameOf("/a/b/file.nfo"); got != "file.nfo" {
t.Errorf("got %s", got)
}
if got := baseNameOf("a\\b\\c.jpg"); got != "c.jpg" {
t.Errorf("got %s", got)
}
if got := baseNameOf("top.txt"); got != "top.txt" {
t.Errorf("got %s", got)
}
}
// fakeCallbackOSSClient 捕获 CompleteMultipartUpload 收到的 callback / callback_var,
// 用于断言已经 Base64 编码(116 要求 callback 必须是 Base64 后的 JSON,否则报
// "The callback configuration is not base64 encoded")。
type fakeCallbackOSSClient struct {
capturedCallback string
capturedCallbackVar string
}
func (c *fakeCallbackOSSClient) InitiateMultipartUpload(_ context.Context, _ *oss.InitiateMultipartUploadRequest, _ ...func(*oss.Options)) (*oss.InitiateMultipartUploadResult, error) {
return &oss.InitiateMultipartUploadResult{UploadId: oss.Ptr("upload-new")}, nil
}
func (c *fakeCallbackOSSClient) UploadPart(_ context.Context, r *oss.UploadPartRequest, _ ...func(*oss.Options)) (*oss.UploadPartResult, error) {
if r.Body != nil {
_, _ = io.Copy(io.Discard, r.Body)
}
return &oss.UploadPartResult{ETag: oss.Ptr("etag-1")}, nil
}
func (c *fakeCallbackOSSClient) ListParts(context.Context, *oss.ListPartsRequest, ...func(*oss.Options)) (*oss.ListPartsResult, error) {
return &oss.ListPartsResult{}, nil
}
func (c *fakeCallbackOSSClient) CompleteMultipartUpload(_ context.Context, r *oss.CompleteMultipartUploadRequest, _ ...func(*oss.Options)) (*oss.CompleteMultipartUploadResult, error) {
c.capturedCallback = *r.Callback
c.capturedCallbackVar = *r.CallbackVar
return &oss.CompleteMultipartUploadResult{
CallbackResult: map[string]any{
"state": true,
"data": map[string]any{"file_id": "file-1", "pick_code": "pick-1"},
},
}, nil
}
func (c *fakeCallbackOSSClient) AbortMultipartUpload(context.Context, *oss.AbortMultipartUploadRequest, ...func(*oss.Options)) (*oss.AbortMultipartUploadResult, error) {
return &oss.AbortMultipartUploadResult{}, nil
}
// TestCompleteMultipartUploadCallbackBase64 回归测试:OSS CompleteMultipartUpload 的
// callback 必须 Base64 编码,否则报 "The callback configuration is not base64 encoded",
// 导致大于 128 KiB 的元数据文件上传失败。
func TestCompleteMultipartUploadCallbackBase64(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "big.bin")
data := make([]byte, 8) // 8 字节,PartSize=8 → 1 part
if err := os.WriteFile(path, data, 0o644); err != nil {
t.Fatal(err)
}
fake := &fakeCallbackOSSClient{}
uploader := &OSSMultipartUploader{client: fake}
callback := `{"callbackUrl":"http://uplb.115.com/3.0/completeupload.php"}`
callbackVar := `{"x:pick_code":"abc"}`
_, err := uploader.UploadFileWithResult(context.Background(), OSSMultipartUploadInput{
Bucket: "bucket-1",
Object: "object-1",
Callback: callback,
CallbackVar: callbackVar,
FilePath: path,
FileSize: int64(len(data)),
PartSize: 8,
})
if err != nil {
t.Fatalf("multipart 上传失败:%v", err)
}
// 捕获的 callback 必须是合法 Base64,且解码后与原 JSON 一致
cbBytes, err := base64.StdEncoding.DecodeString(fake.capturedCallback)
if err != nil {
t.Fatalf("callback 未 Base64 编码:%v (raw=%q)", err, fake.capturedCallback)
}
if string(cbBytes) != callback {
t.Errorf("callback 解码后 = %s,期望 %s", cbBytes, callback)
}
cbvBytes, err := base64.StdEncoding.DecodeString(fake.capturedCallbackVar)
if err != nil {
t.Fatalf("callback_var 未 Base64 编码:%v (raw=%q)", err, fake.capturedCallbackVar)
}
if string(cbvBytes) != callbackVar {
t.Errorf("callback_var 解码后 = %s,期望 %s", cbvBytes, callbackVar)
}
}
+65 -14
View File
@@ -278,18 +278,69 @@ func TestDanmakuSameBase(t *testing.T) {
require.False(t, sameDanmakuBase("", "https://api.dandanplay.net"))
}
// fetchCommentWithFallback:配置源与官方同源时不重复请求;
// 全失败时带出最后一跳错误。
func TestDanmakuFetchCommentWithFallback(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
}))
t.Cleanup(srv.Close)
// fetchCommentWithFallback:配置源与官方同源时不重复请求;
// 全失败时带出最后一跳错误。
func TestDanmakuFetchCommentWithFallback(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
}))
t.Cleanup(srv.Close)
svc := newDanmakuTestService(t)
ctx := context.Background()
raw, st, err := svc.fetchCommentWithFallback(ctx, srv.URL, srv.URL, "25484")
require.Error(t, err)
require.Empty(t, raw)
require.Equal(t, "auto", st)
}
svc := newDanmakuTestService(t)
ctx := context.Background()
raw, st, err := svc.fetchCommentWithFallback(ctx, srv.URL, srv.URL, "25484")
require.Error(t, err)
require.Empty(t, raw)
require.Equal(t, "auto", st)
}
// 视频即便能命中 Hash 自动识别,当用户传入手动搜索关键词时应跳过 Hash 匹配,走关键词搜索。
func TestDanmakuFetchHashMatchSkippedOnManualKeyword(t *testing.T) {
videoPath, _ := writeDanmakuTestVideo(t, "测试动画.第01话.mkv")
// 官方服务同时提供 match 和 search:
// match 会返回 episodeId=25484(动画A)
// search 会根据关键词返回 episodeId=99999(动画B)
mux := http.NewServeMux()
var matchCalled bool
mux.HandleFunc("/api/v2/match", func(w http.ResponseWriter, r *http.Request) {
matchCalled = true
w.Header().Set("Content-Type", "application/json")
fmt.Fprint(w, `{"success":true,"isMatched":true,"matches":[{"episodeId":25484,"animeId":1001,"animeTitle":"自动识别动画A","episodeTitle":"第1话"}]}`)
})
mux.HandleFunc("/api/v2/search/episodes", func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
fmt.Fprint(w, `{"hasMore":false,"animes":[{"animeId":2002,"animeTitle":"手动搜索动画B","episodes":[{"episodeId":99999,"episodeTitle":"第1话"}]}]}`)
})
mux.HandleFunc("/api/v2/comment/25484", func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/xml")
fmt.Fprint(w, `<?xml version="1.0"?><i><d p="0.5,1,16777215,user1">自动识别弹幕</d></i>`)
})
mux.HandleFunc("/api/v2/comment/99999", func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/xml")
fmt.Fprint(w, `<?xml version="1.0"?><i><d p="0.5,1,16777215,user2">手动搜索弹幕</d></i>`)
})
official := httptest.NewServer(mux)
t.Cleanup(official.Close)
overrideDanmakuOfficialBase(t, official.URL)
svc := newDanmakuTestService(t)
ctx := context.Background()
seedDanmakuVideoMedia(t, svc, "mManual", "自动识别动画A", videoPath, 32000, 1)
// 1) 默认自动识别:命中 Hash 识别
resAuto, err := svc.Fetch(ctx, "mManual", "", "")
require.NoError(t, err)
require.True(t, matchCalled)
require.Equal(t, "hash", resAuto.MatchMode)
require.Equal(t, int64(25484), resAuto.EpisodeID)
require.Contains(t, resAuto.Raw, "自动识别弹幕")
// 2) 用户传入手动搜索关键词:跳过 Hash 识别,命中搜索结果动画B
resManual, err := svc.Fetch(ctx, "mManual", "手动搜索动画B", "")
require.NoError(t, err)
require.Equal(t, "search", resManual.MatchMode)
require.Equal(t, int64(99999), resManual.EpisodeID)
require.Equal(t, "手动搜索动画B", resManual.AnimeTitle)
require.Contains(t, resManual.Raw, "手动搜索弹幕")
}
+2 -2
View File
@@ -219,8 +219,8 @@ func (s *DanmakuService) Fetch(ctx context.Context, mediaID, keyword, episodeID
target := ""
// 1) hash 识别:始终走官方 /api/v2/match。
if media != nil && media.Path != "" {
// 1) hash 识别:始终走官方 /api/v2/match(keyword 手动覆盖时跳过,直接走第 3 层)。
if target == "" && !manualKeyword && media != nil && media.Path != "" {
if hash, ok := s.mediaHash(ctx, media); ok {
fileSize := media.SizeBytes
if strings.EqualFold(filepath.Ext(media.Path), ".strm") {
+90
View File
@@ -0,0 +1,90 @@
package service
import (
"context"
"fmt"
"strings"
"go.uber.org/zap"
"gorm.io/gorm"
"github.com/ShukeBta/MMTL/internal/config"
"github.com/ShukeBta/MMTL/internal/database"
"github.com/ShukeBta/MMTL/internal/repository"
)
// DatabaseAdminService manages database configuration, connectivity testing, and migration.
type DatabaseAdminService struct {
cfg *config.Config
log *zap.Logger
repos *repository.Container
db *gorm.DB
}
// NewDatabaseAdminService creates a new DatabaseAdminService.
func NewDatabaseAdminService(cfg *config.Config, log *zap.Logger, repos *repository.Container, db *gorm.DB) *DatabaseAdminService {
if log == nil {
log = zap.NewNop()
}
return &DatabaseAdminService{
cfg: cfg,
log: log,
repos: repos,
db: db,
}
}
// GetStatus returns the status of the currently active database.
func (s *DatabaseAdminService) GetStatus(ctx context.Context) *database.DatabaseStatus {
return database.InspectDatabaseStatus(s.db, s.cfg)
}
// TestPostgres verifies connectivity and permissions to the specified PostgreSQL DSN.
func (s *DatabaseAdminService) TestPostgres(ctx context.Context, dsn string) (*database.PostgresTestResult, error) {
return database.TestPostgres(dsn)
}
// MigrateToPostgres copies all records from the current active database to the target PostgreSQL database.
func (s *DatabaseAdminService) MigrateToPostgres(ctx context.Context, targetDSN string) (*database.DatabaseMigrationResult, error) {
s.log.Info("starting user-initiated database migration to PostgreSQL", zap.String("target", database.MaskDSN(targetDSN)))
res, err := database.MigrateCurrentToPostgres(s.db, targetDSN, 500, s.log)
if err != nil {
s.log.Error("database migration to PostgreSQL failed", zap.Error(err))
return nil, err
}
s.log.Info("database migration to PostgreSQL completed successfully",
zap.Int64("total_rows", res.TotalRows),
zap.Int64("duration_ms", res.DurationMS),
)
return res, nil
}
// SaveConfig persists the database configuration to config.yaml and the database settings table.
func (s *DatabaseAdminService) SaveConfig(ctx context.Context, dbType, dsn string) error {
dbType = strings.TrimSpace(dbType)
dsn = strings.TrimSpace(dsn)
if dbType == "" {
dbType = "postgres"
}
if dbType == "postgres" && dsn == "" {
return fmt.Errorf("PostgreSQL DSN 不能为空")
}
// 1. 保存到本地 config.yaml
if err := config.SaveDatabaseConfig(dbType, dsn); err != nil {
return fmt.Errorf("保存配置文件失败: %w", err)
}
// 2. 更新内存配置
s.cfg.Database.Type = dbType
s.cfg.Database.DSN = dsn
// 3. 同时更新 settings 存储库作为副本
if s.repos != nil && s.repos.Setting != nil {
_ = s.repos.Setting.Set(ctx, "database.type", dbType)
_ = s.repos.Setting.Set(ctx, "database.dsn", dsn)
}
s.log.Info("database configuration saved", zap.String("type", dbType), zap.String("dsn", database.MaskDSN(dsn)))
return nil
}
+1
View File
@@ -111,6 +111,7 @@ const (
var (
embySeasonDirRE = regexp.MustCompile(`(?i)^(season[\s._-]*\d+|s\d+|specials?|sp|ova|oad|extra|extras|第\s*[0-9一二三四五六七八九十百零两]+\s*季|特别篇|特別篇|番外|特典)$`)
embySeasonSuffixRE = regexp.MustCompile(`(?i)(?:[\s._-]+(?:season[\s._-]*\d+|s\d+|第\s*[0-9一二三四五六七八九十百零两]+\s*季|specials?|sp|ova|oad|extra|extras|特别篇|特別篇|番外|特典)|\s*第\s*[0-9一二三四五六七八九十百零两]+\s*季)\s*$`)
embyYearSuffixRE = regexp.MustCompile(`\s*[\((\[]\d{4}[\))\]]\s*$`)
embyEpisodeTitleRE = regexp.MustCompile(`(?i)\s*[-_ ]*s\d{1,2}e\d{1,3}.*$`)
)
@@ -334,3 +334,105 @@ func TestEmbyCloudAnimeUsesSeriesNameFromChineseSeasonFolder(t *testing.T) {
t.Fatalf("cloud anime should be grouped as one series named 剑来, got %#v", items)
}
}
func TestEmbySeriesGroupingWithPrefixedSeasonFolders(t *testing.T) {
svc := newTestEmbyService(t)
lib := model.Library{Name: "动漫", Path: `/media/动漫`, Type: "anime", Enabled: true}
if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
for season := 1; season <= 5; season++ {
for ep := 1; ep <= 3; ep++ {
media := model.Media{
Base: model.Base{ID: fmt.Sprintf("shokugeki-s%02de%02d", season, ep)},
LibraryID: lib.ID,
Title: "食戟之灵",
OriginalName: "食戟のソーマ",
ScrapeStatus: "matched",
TMDbID: 62273,
BangumiID: 116461,
Path: fmt.Sprintf(`/media/动漫/食戟之灵/食戟之灵 S%02d/食戟之灵 S%02dE%02d.strm`, season, season, ep),
SeasonNum: season,
EpisodeNum: ep,
}
if err := svc.repo.DB.Create(&media).Error; err != nil {
t.Fatalf("create media: %v", err)
}
}
}
root, err := svc.Items(t.Context(), ItemsParams{ParentID: lib.ID, Limit: 50})
if err != nil {
t.Fatalf("library items: %v", err)
}
rootItems := root["Items"].([]map[string]any)
if len(rootItems) != 1 {
t.Fatalf("expected 1 series card for 食戟之灵 across 5 seasons, got %d cards: %#v", len(rootItems), rootItems)
}
if rootItems[0]["Name"] != "食戟之灵" || rootItems[0]["Type"] != "Series" {
t.Fatalf("unexpected series item: %#v", rootItems[0])
}
seriesID := rootItems[0]["Id"].(string)
seasons, err := svc.Items(t.Context(), ItemsParams{ParentID: seriesID, Limit: 50})
if err != nil {
t.Fatalf("series seasons: %v", err)
}
seasonItems := seasons["Items"].([]map[string]any)
if len(seasonItems) != 5 {
t.Fatalf("expected 5 seasons, got %d: %#v", len(seasonItems), seasonItems)
}
for i, s := range seasonItems {
wantSeasonNum := i + 1
if s["Type"] != "Season" || s["IndexNumber"] != wantSeasonNum {
t.Errorf("season [%d] = %#v, want IndexNumber=%d", i, s, wantSeasonNum)
}
}
counts, err := svc.ItemCounts(t.Context(), "user-1")
if err != nil {
t.Fatalf("item counts: %v", err)
}
if counts["SeriesCount"] != 1 || counts["EpisodeCount"] != int64(15) {
t.Fatalf("counts = %#v, want 1 series and 15 episodes", counts)
}
}
func TestInferSeriesNameFromPath(t *testing.T) {
tests := []struct {
path string
want string
}{
{
path: `/media/动漫/食戟之灵/食戟之灵 S01/食戟之灵 S01E01.strm`,
want: "食戟之灵",
},
{
path: `/media/动漫/食戟之灵/食戟之灵 S05/食戟之灵 S05E12.strm`,
want: "食戟之灵",
},
{
path: `/media/动漫/食戟之灵/Season 02/01.mkv`,
want: "食戟之灵",
},
{
path: `/media/动漫/进击的巨人 第2季/01.mkv`,
want: "进击的巨人",
},
{
path: `cloud://openlist/国漫/剑来/第二季/04.mkv`,
want: "剑来",
},
{
path: `/media/tv/间谍过家家 (2022)/Specials/S00E01.mkv`,
want: "间谍过家家",
},
}
for _, tc := range tests {
got := inferSeriesNameFromPath(tc.path)
if got != tc.want {
t.Errorf("inferSeriesNameFromPath(%q) = %q, want %q", tc.path, got, tc.want)
}
}
}
+30 -1
View File
@@ -31,6 +31,14 @@ func (e *EmbyService) seriesNameForMedia(m *model.Media) string {
return series.Title
}
}
if strings.EqualFold(strings.TrimSpace(m.ScrapeStatus), "matched") && strings.TrimSpace(m.Title) != "" {
name := strings.TrimSpace(m.Title)
name = embyEpisodeTitleRE.ReplaceAllString(name, "")
name = embyYearSuffixRE.ReplaceAllString(name, "")
if name != "" {
return name
}
}
if name := inferSeriesNameFromPath(m.Path); name != "" {
return name
}
@@ -53,14 +61,35 @@ func inferSeriesNameFromPath(path string) string {
if embySeasonDirRE.MatchString(base) {
dir = filepath.Dir(dir)
base = filepath.Base(dir)
} else if stripped := strings.TrimSpace(embySeasonSuffixRE.ReplaceAllString(base, "")); stripped != "" && stripped != base {
parentDir := filepath.Dir(dir)
parentBase := filepath.Base(parentDir)
if parentBase != "." && parentBase != string(filepath.Separator) && !isEmbyGenericContainer(parentBase) {
dir = parentDir
base = parentBase
} else {
base = stripped
}
}
base = strings.TrimSpace(embyYearSuffixRE.ReplaceAllString(base, ""))
if base == "." || base == string(filepath.Separator) {
if base == "." || base == string(filepath.Separator) || isEmbyGenericContainer(base) {
return ""
}
return base
}
func isEmbyGenericContainer(name string) bool {
switch strings.ToLower(strings.TrimSpace(name)) {
case "movie", "movies", "film", "films", "tv", "series", "show", "shows", "anime", "animation", "variety",
"电视剧", "剧集", "连续剧", "短剧", "国产剧", "国剧", "欧美剧", "美剧", "英剧", "日韩剧", "日剧", "韩剧", "港剧", "台剧", "港台剧",
"综艺", "纪录片", "儿童", "动漫", "番剧", "国漫", "日番", "韩漫", "美漫", "欧美动漫", "欧美动画", "其他动漫", "电影", "成人", "未分类",
"media", "downloads", "download", "videos", "video", "share", "shares":
return true
default:
return false
}
}
func stableEmbyID(prefix string, parts ...string) string {
h := sha256.New()
for _, part := range parts {
+77 -21
View File
@@ -49,7 +49,7 @@ func findMovieNFO(mediaPath, libraryRoot string) (*nfoDocument, string, error) {
continue
}
seen[key] = struct{}{}
if doc, _, err := readNFO(path); err == nil {
if doc, _, err := decodeNFOFile(path); err == nil && doc != nil {
return doc, path, nil
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, "", err
@@ -62,7 +62,7 @@ func findMovieNFO(mediaPath, libraryRoot string) (*nfoDocument, string, error) {
for _, match := range matches {
baseKey := strings.ToLower(strings.ReplaceAll(strings.TrimSuffix(filepath.Base(match), filepath.Ext(match)), "-", ""))
if strings.Contains(baseKey, codeKey) || strings.Contains(codeKey, baseKey) {
if doc, _, err := readNFO(match); err == nil {
if doc, _, err := decodeNFOFile(match); err == nil && doc != nil {
return doc, match, nil
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, "", err
@@ -71,7 +71,7 @@ func findMovieNFO(mediaPath, libraryRoot string) (*nfoDocument, string, error) {
}
}
if len(matches) == 1 {
if doc, _, err := readNFO(matches[0]); err == nil {
if doc, _, err := decodeNFOFile(matches[0]); err == nil && doc != nil {
return doc, matches[0], nil
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, "", err
@@ -84,21 +84,33 @@ func findMovieNFO(mediaPath, libraryRoot string) (*nfoDocument, string, error) {
func readSeriesMetadata(mediaPath, libraryRoot string) (*LocalMetadata, error) {
var meta *LocalMetadata
showBaseDir := ""
if showDoc, showPath, err := findShowNFO(mediaPath, libraryRoot); err == nil && showDoc != nil {
if showDoc, showPath, partial, err := findShowNFO(mediaPath, libraryRoot); err == nil && showDoc != nil {
showBaseDir = filepath.Dir(showPath)
meta = metadataFromDoc(showDoc, showBaseDir, true)
// A truncated show NFO still yields a usable title/poster (decodePartialNFO
// only surfaces docs with recoverable fields). Keep HasNFO so the recovered
// title participates in series grouping; a fully partial episode match
// below can still demote it.
meta.HasNFO = !partial || meta.HasNFO
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, err
// A damaged show NFO must not discard the whole series: readLocalScanMetadata
// treats this as a hard failure and leaves every episode without local
// metadata. When the show level fails we still merge the episode NFO and
// local artwork below instead of returning early.
meta = nil
}
if episodeDoc, episodePath, err := readNFO(nfoPath(mediaPath)); err == nil {
if episodeDoc, episodePath, episodePartial, err := readEpisodeNFO(nfoPath(mediaPath)); err == nil && episodeDoc != nil {
episodeMeta := metadataFromDoc(episodeDoc, filepath.Dir(episodePath), true)
if meta == nil {
meta = &LocalMetadata{}
}
mergeEpisodeMetadata(meta, episodeMeta, episodeDoc)
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, err
// Only a fully valid episode NFO can confirm the match; a truncated one
// keeps whatever the show NFO supplied but must not force scrape_status.
if episodePartial {
meta.HasNFO = false
}
}
if meta == nil {
meta = metadataFromArtwork(mediaPath, showBaseDir)
@@ -108,19 +120,49 @@ func readSeriesMetadata(mediaPath, libraryRoot string) (*LocalMetadata, error) {
return meta, nil
}
func readNFO(path string) (*nfoDocument, string, error) {
// readEpisodeNFO reads the sidecar NFO next to an episode file and reports both
// the recovered document and whether it came from a truncated/malformed file.
func readEpisodeNFO(path string) (*nfoDocument, string, bool, error) {
body, err := os.ReadFile(path) // #nosec G304 -- path is a discovered NFO sidecar under the configured library root.
if err != nil {
return nil, "", err
return nil, "", false, err
}
var doc nfoDocument
if err := xml.Unmarshal(body, &doc); err != nil {
return nil, "", err
doc, partial, err := decodePartialNFO(body)
if err != nil {
return nil, "", false, err
}
return &doc, path, nil
return doc, path, partial, nil
}
func findShowNFO(mediaPath, libraryRoot string) (*nfoDocument, string, error) {
func decodePartialNFO(body []byte) (*nfoDocument, bool, error) {
var doc nfoDocument
err := xml.Unmarshal(body, &doc)
if err == nil {
return &doc, false, nil
}
if !isLikelyTruncatedXMLError(err) {
// A real parse error unrelated to truncation: don't trust partial fields.
return nil, false, err
}
// Unmarshal still fills the elements it closed before hitting the cut; treat
// those as partial show metadata instead of dropping everything.
if doc.Title == "" && doc.OriginalTitle == "" && len(doc.Thumbs) == 0 &&
doc.Premiered == "" && doc.Plot == "" {
return nil, false, err
}
return &doc, true, nil
}
func isLikelyTruncatedXMLError(err error) bool {
for _, msg := range []string{"unexpected EOF", "EOF"} {
if strings.Contains(strings.ToLower(err.Error()), strings.ToLower(msg)) {
return true
}
}
return false
}
func findShowNFO(mediaPath, libraryRoot string) (*nfoDocument, string, bool, error) {
dir := filepath.Dir(mediaPath)
root := filepath.Clean(libraryRoot)
for {
@@ -133,19 +175,33 @@ func findShowNFO(mediaPath, libraryRoot string) (*nfoDocument, string, error) {
names = append(names, base+".nfo")
for _, name := range names {
path := filepath.Join(dir, name)
if doc, _, err := readNFO(path); err == nil {
return doc, path, nil
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, "", err
doc, partial, err := decodeNFOFile(path)
if doc != nil {
return doc, path, partial, nil
}
if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, "", false, err
}
}
if samePath(dir, root) {
return nil, "", os.ErrNotExist
return nil, "", false, os.ErrNotExist
}
parent := filepath.Dir(dir)
if parent == dir {
return nil, "", os.ErrNotExist
return nil, "", false, os.ErrNotExist
}
dir = parent
}
}
func decodeNFOFile(path string) (*nfoDocument, bool, error) {
body, err := os.ReadFile(path) // #nosec G304 -- path is a discovered NFO sidecar under the configured library root.
if err != nil {
return nil, false, err
}
doc, partial, err := decodePartialNFO(body)
if err != nil {
return nil, false, err
}
return doc, partial, nil
}
+108
View File
@@ -267,3 +267,111 @@ func TestReadLocalMetadataWithoutNFOStillFindsArtwork(t *testing.T) {
t.Fatalf("unexpected artwork metadata: %+v", got)
}
}
// TestReadLocalMetadataRecoversFromTruncatedShowNFO mirrors a real-world issue:
// some anime tvshow.nfo files are truncated mid-URL (unexpected EOF). Before the
// fix this discarded the whole series, leaving episodes pending with per-episode
// titles and no poster. The recoverable fields (title/year) and the matching
// episode NFO + local artwork must still be applied.
func TestReadLocalMetadataRecoversFromTruncatedShowNFO(t *testing.T) {
root := t.TempDir()
showDir := filepath.Join(root, "夏日重现 (2022)")
seasonDir := filepath.Join(showDir, "Season 1")
if err := os.MkdirAll(seasonDir, 0o755); err != nil {
t.Fatal(err)
}
mediaPath := filepath.Join(seasonDir, "S01E01.mkv")
if err := os.WriteFile(mediaPath, []byte("x"), 0o644); err != nil {
t.Fatal(err)
}
// tvshow.nfo cut off inside a <thumb> URL, like the broken real files.
truncated := `<?xml version="1.0" encoding="UTF-8" standalone="yes"?>
<tvshow>
<title>夏日重现</title>
<originaltitle>サマータイムレンダ</originaltitle>
<year>2022</year>
<plot>听闻自己青梅竹马死讯。</plot>
<thumb aspect="poster">https://image.tmdb.org/t/p/original/2koyWLm6iVn5OTEExTjKzVms5Iz.jpg</thumb>
<fanart>
<thumb>https://image.tmdb.org/t/p/original/p2eZlGwd8OjkWpwD2hSoBiIlHBZ.jpg</thu`
if err := os.WriteFile(filepath.Join(showDir, "tvshow.nfo"), []byte(truncated), 0o644); err != nil {
t.Fatal(err)
}
// The episode sidecar NFO is complete and should still be merged.
if err := os.WriteFile(nfoPath(mediaPath), []byte(`<episodedetails><title>再见了夏日</title><season>1</season><episode>1</episode></episodedetails>`), 0o644); err != nil {
t.Fatal(err)
}
// Local artwork next to the show folder.
poster := filepath.Join(showDir, "poster.jpg")
if err := os.WriteFile(poster, []byte("jpg"), 0o644); err != nil {
t.Fatal(err)
}
got, err := ReadLocalMetadata(mediaPath, root, true)
if err != nil {
t.Fatalf("ReadLocalMetadata returned error on truncated show NFO: %v", err)
}
if got == nil {
t.Fatal("metadata is nil; truncated show NFO discarded the series")
}
// The recovered show title takes precedence over the episode title.
if got.Title != "夏日重现" {
t.Fatalf("Title = %q, want recovered show title 夏日重现", got.Title)
}
if got.Year != 2022 {
t.Fatalf("Year = %d, want 2022", got.Year)
}
if got.EpisodeTitle != "再见了夏日" || got.SeasonNum != 1 || got.EpisodeNum != 1 {
t.Fatalf("episode metadata not preserved: %+v", got)
}
// Prior to the fix the episode metadata was dropped entirely; the poster comes
// from the local poster.jpg next to the show.
if got.PosterURL != poster {
t.Fatalf("PosterURL = %q, want local poster %q", got.PosterURL, poster)
}
// A truncated show NFO must not be treated as an authoritative match by
// itself; since the episode NFO is valid we still mark it matched so the
// recovered series title participates in grouping.
if !got.HasNFO {
t.Fatalf("HasNFO = false, want true (episode NFO is valid)")
}
}
// TestReadLocalMetadataKeepsArtworkOnlyWhenNoUsableNFO verifies that when a show
// NFO is truncated AND yields no recoverable fields, we still fall back to local
// artwork instead of returning an error.
func TestReadLocalMetadataArtworkFallbackOnGarbageShowNFO(t *testing.T) {
root := t.TempDir()
showDir := filepath.Join(root, "Some Show")
seasonDir := filepath.Join(showDir, "Season 1")
if err := os.MkdirAll(seasonDir, 0o755); err != nil {
t.Fatal(err)
}
mediaPath := filepath.Join(seasonDir, "S01E01.mkv")
if err := os.WriteFile(mediaPath, []byte("x"), 0o644); err != nil {
t.Fatal(err)
}
// Severely truncated: no recoverable title/fields at all.
if err := os.WriteFile(filepath.Join(showDir, "tvshow.nfo"), []byte(`<tvshow><title>半截`), 0o644); err != nil {
t.Fatal(err)
}
// Episode-level backdrop sits next to the episode file.
backdrop := filepath.Join(seasonDir, "S01E01-backdrop.jpg")
if err := os.WriteFile(backdrop, []byte("jpg"), 0o644); err != nil {
t.Fatal(err)
}
got, err := ReadLocalMetadata(mediaPath, root, true)
if err != nil {
t.Fatalf("ReadLocalMetadata should not error when NFO is unusable: %v", err)
}
if got == nil {
t.Fatal("metadata is nil; artwork fallback missing")
}
if got.HasNFO {
t.Fatalf("HasNFO = true, want false for unusable NFO")
}
if got.BackdropURL != backdrop {
t.Fatalf("expected episode artwork fallback, got %+v", got)
}
}
+8 -5
View File
@@ -32,18 +32,21 @@ func (s *ScraperService) manualTMDbCandidates(ctx context.Context, query string,
for _, typ := range manualTMDbSearchTypes(mediaType) {
switch typ {
case "movie":
if matches, err := s.tmdb.SearchMovieCandidates(ctx, query, year); err == nil {
if matches, err := s.tmdb.SearchMovieCandidates(ctx, query, year); err == nil && len(matches) > 0 {
for _, match := range matches {
out = append(out, manualTMDbCandidate{MediaType: "movie", Match: match})
}
}
case "tv":
if matches, err := s.tmdb.SearchTVCandidates(ctx, query, year); err == nil {
if matches, err := s.tmdb.SearchTVCandidates(ctx, query, year); err == nil && len(matches) > 0 {
for _, match := range matches {
out = append(out, manualTMDbCandidate{MediaType: "tv", Match: match})
}
}
}
if len(out) > 0 {
break
}
}
return out
}
@@ -70,15 +73,15 @@ func manualTMDbIDSearchTypes(mediaType string) []string {
func manualTMDbSearchTypes(mediaType string) []string {
if strings.TrimSpace(mediaType) == "" {
return []string{"movie", "tv"}
return []string{"tv", "movie"}
}
switch normalizeMediaType(mediaType, "", "") {
case "tv", "anime", "variety":
return []string{"tv", "movie"}
case "movie", "adult":
return []string{"movie"}
default:
return []string{"movie", "tv"}
default:
return []string{"tv", "movie"}
}
}
+3 -3
View File
@@ -179,9 +179,9 @@ func TestManualSearchFallsBackToMovieFolderForGenericQuery(t *testing.T) {
if len(results) != 1 || results[0].TMDbID != 27205 {
t.Fatalf("manual search results=%#v, want folder fallback candidate; queries=%v", results, queries)
}
if len(queries) < 2 || queries[0] != "00000" || queries[1] != "inception" {
t.Fatalf("manual search queries=%v, want explicit query then folder fallback", queries)
}
if len(queries) < 2 || queries[0] != "00000" || queries[len(queries)-1] != "inception" {
t.Fatalf("manual search queries=%v, want explicit query then folder fallback", queries)
}
}
func TestManualSearchReturnsMovieFallbackForTVTypedTMDbSearch(t *testing.T) {
+4 -4
View File
@@ -43,10 +43,10 @@ func (s *MediaService) DeleteLibrary(ctx context.Context, id string) error {
if err := tx.Unscoped().Where("library_id = ?", id).Delete(&model.Media{}).Error; err != nil {
return err
}
if err := hardDeleteLibraryRoots(ctx, tx, id); err != nil {
return err
}
return tx.Delete(&model.Library{}, "id = ?", id).Error
if err := hardDeleteLibraryRoots(ctx, tx, id); err != nil {
return err
}
return tx.Unscoped().Delete(&model.Library{}, "id = ?", id).Error
})
if err == nil {
s.invalidateMediaCache(ctx)
+70 -1
View File
@@ -4,6 +4,7 @@ import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
"strings"
@@ -51,7 +52,14 @@ func (s *MediaService) CreateLibraryWithRootsAndCover(ctx context.Context, name,
s.invalidateMediaCache(ctx)
return lib, nil
}
lib := &model.Library{Name: strings.TrimSpace(name), Path: roots[0].Path, Type: kind, CoverURL: strings.TrimSpace(coverURL), Enabled: true}
lib := &model.Library{
Name: strings.TrimSpace(name),
Path: roots[0].Path,
Type: kind,
CoverURL: strings.TrimSpace(coverURL),
Enabled: true,
CarouselEnabled: false,
}
if err := s.repo.Library.CreateWithRoots(ctx, lib, roots); err != nil {
return nil, err
}
@@ -59,11 +67,72 @@ func (s *MediaService) CreateLibraryWithRootsAndCover(ctx context.Context, name,
return lib, nil
}
// CreateLibrariesPerSubfolder 为 parent 目录下的每个直接子目录各建一个媒体库,
// 媒体库名取子目录名,路径指向该子目录。kind 为空时按子目录名推断类型。
func (s *MediaService) CreateLibrariesPerSubfolder(ctx context.Context, parent, kind, coverURL string) ([]model.Library, error) {
parent = strings.TrimSpace(parent)
if parent == "" {
return nil, errors.New("parent path required")
}
dir, err := resolveAccessibleLibraryPath(parent)
if err != nil {
return nil, err
}
entries, err := os.ReadDir(dir)
if err != nil {
return nil, fmt.Errorf("read directory failed: %w", err)
}
subdirs := make([]string, 0, len(entries))
for _, entry := range entries {
if !entry.IsDir() {
continue
}
if strings.HasPrefix(entry.Name(), ".") {
continue
}
subdirs = append(subdirs, filepath.Join(dir, entry.Name()))
}
if len(subdirs) == 0 {
return nil, errors.New("no subfolders found")
}
created := make([]model.Library, 0, len(subdirs))
for _, subdir := range subdirs {
name := filepath.Base(subdir)
lib, err := s.CreateLibraryWithRootsAndCover(ctx, name, kind, coverURL, []LibraryRootInput{{Path: subdir}})
if err != nil {
return nil, fmt.Errorf("create library for %s: %w", subdir, err)
}
created = append(created, *lib)
}
return created, nil
}
func (s *MediaService) UpdateLibraryCover(ctx context.Context, libraryID, coverURL string) error {
return s.repo.DB.WithContext(ctx).Model(&model.Library{}).Where("id = ?", libraryID).
Update("cover_url", strings.TrimSpace(coverURL)).Error
}
// UpdateLibraryFields updates sort_order / carousel_enabled on a library.
func (s *MediaService) UpdateLibraryFields(ctx context.Context, libraryID string, sortOrder *int, carouselEnabled *bool) error {
updates := map[string]any{}
if sortOrder != nil {
updates["sort_order"] = *sortOrder
}
if carouselEnabled != nil {
updates["carousel_enabled"] = *carouselEnabled
}
if len(updates) == 0 || strings.TrimSpace(libraryID) == "" {
return nil
}
return s.repo.DB.WithContext(ctx).Model(&model.Library{}).
Where("id = ?", libraryID).Updates(updates).Error
}
// ReorderLibraries persists a full media-library ordering.
func (s *MediaService) ReorderLibraries(ctx context.Context, ids []string) error {
return s.repo.Library.SetSortOrder(ctx, ids)
}
func (s *MediaService) findLogicalLibrary(ctx context.Context, name, kind string) (*model.Library, error) {
if s == nil || s.repo == nil || s.repo.Library == nil {
return nil, nil
+2 -17
View File
@@ -10,25 +10,10 @@ import (
const maxRecycleBinRecords = 200
// SoftDelete moves a media row to the recycle bin (gorm soft delete).
// The on-disk file is kept; admins can purge it later.
// SoftDelete 物理删除媒体记录(统一硬删除以降低 SQLite 存储与索引压力)。
func (s *MediaService) SoftDelete(ctx context.Context, id string) error {
media, err := s.repo.Media.FindByID(ctx, id)
if err != nil {
return err
}
if media != nil && isCloudMediaPath(media.Path) {
err := s.repo.DB.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.Media{}).Error
if err == nil {
s.invalidateMediaCache(ctx)
}
return err
}
err = s.repo.DB.WithContext(ctx).Where("id = ?", id).Delete(&model.Media{}).Error
err := s.repo.DB.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.Media{}).Error
if err == nil {
if pruneErr := pruneRecycleBinRows(ctx, s.repo.DB, maxRecycleBinRecords); pruneErr != nil {
return pruneErr
}
s.invalidateMediaCache(ctx)
}
return err
@@ -229,9 +229,9 @@ func (o *OrganizerService) replaceVersions(ctx context.Context, src string, exis
o.log.Warn("organize replace remove existing failed",
zap.String("path", e), zap.Error(err))
}
if o.repo != nil && o.repo.DB != nil {
_ = o.repo.DB.WithContext(ctx).Where("path = ?", e).Delete(&model.Media{}).Error
}
if o.repo != nil && o.repo.DB != nil {
_ = o.repo.DB.WithContext(ctx).Unscoped().Where("path = ?", e).Delete(&model.Media{}).Error
}
}
// Move staged file + sidecars into the final path.
if err := os.Rename(stage, dst); err != nil {
@@ -97,7 +97,7 @@ func (o *OrganizerService) deleteMediaRowForPath(ctx context.Context, path strin
if o == nil || o.repo == nil || o.repo.DB == nil {
return
}
_ = o.repo.DB.WithContext(ctx).Where("path = ?", path).Delete(&model.Media{}).Error
_ = o.repo.DB.WithContext(ctx).Unscoped().Where("path = ?", path).Delete(&model.Media{}).Error
}
func (o *OrganizerService) mediaPathExists(ctx context.Context, path string) bool {
+5 -5
View File
@@ -196,18 +196,18 @@ func (p *PlaybackService) AddToPlaylist(ctx context.Context, playlistID, mediaID
return p.repo.DB.Create(item).Error
}
// RemoveFromPlaylist removes a media item from a playlist (idempotent).
// RemoveFromPlaylist 物理删除播放列表项(幂等)。
func (p *PlaybackService) RemoveFromPlaylist(ctx context.Context, playlistID, mediaID string) error {
return p.repo.DB.
return p.repo.DB.WithContext(ctx).Unscoped().
Where("playlist_id = ? AND media_id = ?", playlistID, mediaID).
Delete(&model.PlaylistItem{}).Error
}
// DeletePlaylist removes a playlist and all of its items.
// DeletePlaylist 物理删除播放列表及其全部条目。
func (p *PlaybackService) DeletePlaylist(ctx context.Context, playlistID string) error {
if err := p.repo.DB.Where("playlist_id = ?", playlistID).
if err := p.repo.DB.WithContext(ctx).Unscoped().Where("playlist_id = ?", playlistID).
Delete(&model.PlaylistItem{}).Error; err != nil {
return err
}
return p.repo.DB.Where("id = ?", playlistID).Delete(&model.Playlist{}).Error
return p.repo.DB.WithContext(ctx).Unscoped().Where("id = ?", playlistID).Delete(&model.Playlist{}).Error
}
+4 -5
View File
@@ -11,13 +11,12 @@ import (
"github.com/ShukeBta/MMTL/internal/model"
)
// RemovePath deletes the media row for a path that has disappeared from disk
// (incremental delete used by the watcher on Remove/Rename events).
// RemovePath 物理删除磁盘上已不存在的媒体记录。
func (s *ScannerService) RemovePath(ctx context.Context, path string) (int64, error) {
if _, err := os.Stat(path); err == nil {
return 0, nil // still exists; nothing to remove
}
res := s.repo.DB.WithContext(ctx).
res := s.repo.DB.WithContext(ctx).Unscoped().
Where("path = ?", path).
Delete(&model.Media{})
if res.Error == nil && res.RowsAffected > 0 {
@@ -55,7 +54,7 @@ func (s *ScannerService) pruneMissingMedia(ctx context.Context, libraryID string
}
stale = append(stale, row.ID)
}
return s.deleteMediaByIDs(ctx, stale, false)
return s.deleteMediaByIDs(ctx, stale, true)
}
func (s *ScannerService) pruneMissingMediaForRoot(ctx context.Context, libraryID, rootID, rootPath string, seen map[string]struct{}) (int64, error) {
@@ -92,7 +91,7 @@ func (s *ScannerService) pruneMissingMediaForRoot(ctx context.Context, libraryID
}
stale = append(stale, row.ID)
}
return s.deleteMediaByIDs(ctx, stale, false)
return s.deleteMediaByIDs(ctx, stale, true)
}
func pathBelongsToRoot(pathValue, rootPath string) bool {
+354
View File
@@ -0,0 +1,354 @@
package service
import (
"context"
"errors"
"fmt"
"strings"
"sync"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MMTL/internal/model"
)
type ScrapeQueueCounts struct {
Pending int64 `json:"pending"`
Running int64 `json:"running"`
Done int64 `json:"done"`
Failed int64 `json:"failed"`
Canceled int64 `json:"canceled"`
}
type ScrapeQueueSnapshot struct {
Counts ScrapeQueueCounts `json:"counts"`
Tasks []model.ScrapeTask `json:"tasks"`
Total int64 `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
}
// Start 启动刮削任务队列的后台消费者。
func (s *ScraperService) Start(ctx context.Context) {
if s == nil {
return
}
go s.queueWorker(ctx)
}
func (s *ScraperService) queueWorker(ctx context.Context) {
const claimBatch = 4
sem := make(chan struct{}, 2) // 最大并发刮削数:2
for {
select {
case <-ctx.Done():
return
default:
}
tasks, err := s.repo.ScrapeTask.ClaimPending(ctx, claimBatch)
if err != nil {
if s.log != nil {
s.log.Warn("claim pending scrape task failed", zap.Error(err))
}
sleepContext(ctx, 3*time.Second)
continue
}
if len(tasks) == 0 {
sleepContext(ctx, 2*time.Second)
continue
}
var wg sync.WaitGroup
for i := range tasks {
wg.Add(1)
go func(t *model.ScrapeTask) {
defer wg.Done()
select {
case <-ctx.Done():
return
case sem <- struct{}{}:
}
defer func() { <-sem }()
s.processScrapeTask(ctx, t)
}(&tasks[i])
}
wg.Wait()
}
}
func (s *ScraperService) processScrapeTask(ctx context.Context, task *model.ScrapeTask) {
media, err := s.repo.Media.FindByID(ctx, task.MediaID)
if err != nil || media == nil {
now := time.Now()
task.Status = model.ScrapeTaskFailed
task.Error = "媒体项已不存在或被删除"
task.FinishedAt = &now
_ = s.repo.ScrapeTask.Update(ctx, task)
return
}
epArtwork := task.EpisodeImages
options := ScrapeOptions{
EpisodeArtwork: &epArtwork,
IncludeMatched: task.RefreshMatched,
RetryNoMatch: true,
}
enrichErr := s.EnrichOneWithOptions(ctx, media, options)
now := time.Now()
task.FinishedAt = &now
refreshed, _ := s.repo.Media.FindByID(ctx, media.ID)
if refreshed != nil && refreshed.ScrapeStatus == "matched" {
task.Status = model.ScrapeTaskDone
task.Error = ""
task.MatchedTitle = refreshed.Title
task.MatchedYear = refreshed.Year
task.PosterURL = refreshed.PosterURL
task.BackdropURL = refreshed.BackdropURL
if refreshed.TMDbID > 0 {
task.Provider = "tmdb"
} else if strings.TrimSpace(refreshed.DoubanID) != "" {
task.Provider = "douban"
} else if refreshed.BangumiID > 0 {
task.Provider = "bangumi"
} else if strings.TrimSpace(refreshed.TheTVDBID) != "" {
task.Provider = "thetvdb"
} else {
task.Provider = "metatube"
}
} else {
task.Status = model.ScrapeTaskFailed
if enrichErr != nil {
task.Error = enrichErr.Error()
} else if refreshed != nil && refreshed.ScrapeStatus == "no_match" {
task.Error = "未搜索到匹配的元数据"
} else {
task.Error = "刮削未完成匹配"
}
}
_ = s.repo.ScrapeTask.Update(ctx, task)
if s.hub != nil {
s.hub.Publish("scraper_queue", map[string]any{
"task_id": task.ID,
"status": task.Status,
"title": task.MediaTitle,
})
}
}
func (s *ScraperService) mediaKind(m *model.Media, lib *model.Library) string {
if m == nil {
return ""
}
if lib != nil && lib.Type != "" {
return lib.Type
}
if mediaIsEpisodic(m, lib) {
return "tv"
}
return "movie"
}
// EnqueueMedia 把单个媒体项放入刮削队列。
func (s *ScraperService) EnqueueMedia(ctx context.Context, mediaID string, options ScrapeOptions) (*model.ScrapeTask, error) {
if s == nil || s.repo == nil {
return nil, errors.New("scraper service not initialized")
}
media, err := s.repo.Media.FindByID(ctx, mediaID)
if err != nil || media == nil {
return nil, errors.New("media not found")
}
if active, _ := s.repo.ScrapeTask.FindActiveByMediaID(ctx, mediaID); active != nil {
return active, nil
}
libName := ""
var lib *model.Library
if strings.TrimSpace(media.LibraryID) != "" {
lib, _ = s.repo.Library.FindByID(ctx, media.LibraryID)
if lib != nil {
libName = lib.Name
}
}
task := &model.ScrapeTask{
MediaID: media.ID,
LibraryID: media.LibraryID,
LibraryName: libName,
MediaTitle: media.Title,
MediaPath: media.Path,
MediaType: s.mediaKind(media, lib),
Status: model.ScrapeTaskPending,
EpisodeImages: options.episodeArtworkEnabled(),
RefreshMatched: options.IncludeMatched || options.RefreshWeakMatched,
}
if err := s.repo.ScrapeTask.Create(ctx, task); err != nil {
return nil, err
}
return task, nil
}
// EnqueueLibrary 把指定媒体库内的所有候选媒体批量推入刮削队列。
func (s *ScraperService) EnqueueLibrary(ctx context.Context, libraryID string, options ScrapeOptions) (int, error) {
if s == nil || s.repo == nil {
return 0, errors.New("scraper service not initialized")
}
lib, err := s.repo.Library.FindByID(ctx, libraryID)
if err != nil || lib == nil {
return 0, errors.New("library not found")
}
rows, err := s.scrapeCandidateRows(ctx, libraryID, options)
if err != nil {
return 0, err
}
if len(rows) == 0 {
return 0, nil
}
tasks := make([]model.ScrapeTask, 0, len(rows))
for _, m := range rows {
tasks = append(tasks, model.ScrapeTask{
MediaID: m.ID,
LibraryID: lib.ID,
LibraryName: lib.Name,
MediaTitle: m.Title,
MediaPath: m.Path,
MediaType: s.mediaKind(&m, lib),
Status: model.ScrapeTaskPending,
EpisodeImages: options.episodeArtworkEnabled(),
RefreshMatched: options.IncludeMatched || options.RefreshWeakMatched,
})
}
if err := s.repo.ScrapeTask.CreateBatch(ctx, tasks); err != nil {
return 0, err
}
return len(tasks), nil
}
// EnqueueAll 把所有已启用媒体库的媒体推入刮削队列。
func (s *ScraperService) EnqueueAll(ctx context.Context, options ScrapeOptions) (int, error) {
libs, err := s.repo.Library.List(ctx)
if err != nil {
return 0, err
}
total := 0
for _, lib := range libs {
if !lib.Enabled {
continue
}
n, err := s.EnqueueLibrary(ctx, lib.ID, options)
if err != nil {
if s.log != nil {
s.log.Warn("enqueue library for scrape failed", zap.String("library", lib.ID), zap.Error(err))
}
continue
}
total += n
}
return total, nil
}
func (s *ScraperService) ScrapeQueueSnapshot(ctx context.Context, status string, page, pageSize int) (*ScrapeQueueSnapshot, error) {
tasks, total, err := s.repo.ScrapeTask.List(ctx, status, page, pageSize)
if err != nil {
return nil, err
}
countsMap, err := s.repo.ScrapeTask.CountByStatus(ctx)
if err != nil {
return nil, err
}
snap := &ScrapeQueueSnapshot{
Counts: ScrapeQueueCounts{
Pending: countsMap[model.ScrapeTaskPending],
Running: countsMap[model.ScrapeTaskRunning],
Done: countsMap[model.ScrapeTaskDone],
Failed: countsMap[model.ScrapeTaskFailed],
Canceled: countsMap[model.ScrapeTaskCanceled],
},
Tasks: tasks,
Total: total,
Page: page,
PageSize: pageSize,
}
return snap, nil
}
func (s *ScraperService) CancelScrapeTask(ctx context.Context, id string) error {
task, err := s.repo.ScrapeTask.FindByID(ctx, id)
if err != nil || task == nil {
return errors.New("刮削任务不存在")
}
if task.Status != model.ScrapeTaskPending && task.Status != model.ScrapeTaskRunning {
return errors.New("任务已完成或已终止,无法取消")
}
now := time.Now()
task.Status = model.ScrapeTaskCanceled
task.Error = "已取消"
task.FinishedAt = &now
return s.repo.ScrapeTask.Update(ctx, task)
}
func (s *ScraperService) RetryScrapeTask(ctx context.Context, id string) error {
task, err := s.repo.ScrapeTask.FindByID(ctx, id)
if err != nil || task == nil {
return errors.New("刮削任务不存在")
}
if task.Status != model.ScrapeTaskFailed && task.Status != model.ScrapeTaskCanceled {
return errors.New("只有失败或已取消的任务可以重试")
}
task.Status = model.ScrapeTaskPending
task.Error = ""
task.RetryCount = 0
task.StartedAt = nil
task.FinishedAt = nil
return s.repo.ScrapeTask.Update(ctx, task)
}
func (s *ScraperService) DeleteScrapeTask(ctx context.Context, id string) error {
return s.repo.ScrapeTask.Delete(ctx, id)
}
func (s *ScraperService) BatchActionScrapeTasks(ctx context.Context, action string, ids []string) (int64, error) {
switch action {
case "delete":
return s.repo.ScrapeTask.DeleteBatch(ctx, ids)
case "retry":
return s.repo.ScrapeTask.RetryBatch(ctx, ids)
case "cancel":
return s.repo.ScrapeTask.CancelBatch(ctx, ids)
default:
return 0, fmt.Errorf("不支持的操作: %s", action)
}
}
func (s *ScraperService) ClearDoneScrapeTasks(ctx context.Context) (int64, error) {
return s.repo.ScrapeTask.ClearDone(ctx)
}
func (s *ScraperService) ClearFinishedScrapeTasks(ctx context.Context) (int64, error) {
return s.repo.ScrapeTask.ClearFinished(ctx)
}
func (s *ScraperService) ClearCanceledScrapeTasks(ctx context.Context) (int64, error) {
return s.repo.ScrapeTask.ClearCanceled(ctx)
}
func (s *ScraperService) RetryAllFailedScrapeTasks(ctx context.Context) (int64, error) {
return s.repo.ScrapeTask.RetryAllFailed(ctx)
}
func (s *ScraperService) CancelPendingScrapeTasks(ctx context.Context) (int64, error) {
return s.repo.ScrapeTask.CancelPending(ctx)
}
+11 -5
View File
@@ -58,11 +58,12 @@ type Container struct {
Device *DeviceService
Cache *RuntimeCacheService
Sessions *SessionTrackerService
RecognitionWords *RecognitionWordsService
Danmaku *DanmakuService
Strm *StrmService
RecognitionWords *RecognitionWordsService
Danmaku *DanmakuService
Strm *StrmService
Database *DatabaseAdminService
stopCtx context.Context
stopCtx context.Context
stopCancel context.CancelFunc
// ReloadHTTPServer 由 cmd/server 注入。HTTPS 相关设置保存后,handler
@@ -102,7 +103,12 @@ func (c *Container) Boot() {
c.Strm.Start(c.stopCtx)
}
// Mgo 保号规则巡检:默认关闭,由管理员通过 Telegram Bot 命令开启。
// 启动刮削队列后台消费者
if c.Scraper != nil {
c.Scraper.Start(c.stopCtx)
}
// Mgo 保号规则巡检:默认关闭,由管理员通过 Telegram Bot 命令开启。
// 每天触发一次评估;规则里的窗口可随机,不固定。
if c.Device != nil {
go c.runInactivitySweeper(c.stopCtx)
+1
View File
@@ -118,6 +118,7 @@ func (b *serviceContainerBuilder) initContentServices() {
func (b *serviceContainerBuilder) initAccessAndStorageServices() {
b.c.PlayProfiles = NewPlayProfileService(b.log, b.repos)
b.c.Permissions = NewPermissionService(b.log, b.repos)
b.c.Database = NewDatabaseAdminService(b.cfg, b.log, b.repos, b.repos.DB)
b.c.Emby.SetRuntimeCache(b.c.Cache)
b.c.Emby.SetSubtitleService(b.c.Subtitle)
b.c.Scheduler = NewSchedulerService(
+11 -4
View File
@@ -22,11 +22,18 @@ func TestNormalizeCloudPlayTarget(t *testing.T) {
if parsed.IsAbs() || parsed.Host != "" {
t.Fatalf("normalized target should be relative, got %q", got)
}
if parsed.Query().Get("ref") != ref {
t.Fatalf("ref round-trip failed: %q", parsed.Query().Get("ref"))
}
if parsed.Query().Get("ref") != ref {
t.Fatalf("ref round-trip failed: %q", parsed.Query().Get("ref"))
}
// 非云盘播放 URL 保持原样(WebDAV/直链等)。
strmStale := "http://bwg.linkmy.fun:1314/api/strm/play/cloud115/video.mkv?acct=abc&pickcode=123"
gotStrm := normalizeCloudPlayTarget(strmStale)
wantStrm := "/api/strm/play/cloud115/video.mkv?acct=abc&pickcode=123"
if gotStrm != wantStrm {
t.Fatalf("normalizeCloudPlayTarget(strm) = %q, want %q", gotStrm, wantStrm)
}
// 非云盘播放 URL 保持原样(WebDAV/直链等)。
passthrough := "https://dav.example.com/media/file.mkv"
if got := normalizeCloudPlayTarget(passthrough); got != passthrough {
t.Fatalf("non-cloud target should pass through, got %q", got)
+14 -3
View File
@@ -15,11 +15,22 @@ import (
// /api/cloud/play 路径,由 absoluteInternalRedirect 基于「当前请求」补全
// host,从而对历史脏数据免疫。
func normalizeCloudPlayTarget(raw string) string {
typ, ref, ok := parseCloudMediaPlaybackURL(raw)
if !ok {
raw = strings.TrimSpace(raw)
if raw == "" {
return raw
}
return BuildRelativeCloudPlayURL(typ, ref)
if typ, ref, ok := parseCloudMediaPlaybackURL(raw); ok {
return BuildRelativeCloudPlayURL(typ, ref)
}
if u, err := url.Parse(raw); err == nil {
path := strings.ToLower(u.Path)
if strings.HasPrefix(path, "/api/strm/play/") || strings.HasPrefix(path, "/api/cloud/play/") || strings.HasPrefix(path, "/api/stream/") {
u.Scheme = ""
u.Host = ""
return u.String()
}
}
return raw
}
// BuildRelativeCloudPlayURL 构造相对的云盘播放 API 路径。
+116 -3
View File
@@ -13,6 +13,7 @@ import (
"os"
"path/filepath"
"strings"
"sync"
"time"
"go.uber.org/zap"
@@ -27,7 +28,12 @@ const (
)
// downloadWorker 下载队列 worker:认领 → 解析直链 → 下载 → 落盘。
//
// 采用「批量认领 + 全局并发限流」:一次认领数个任务,用 StrmService 上的全局信号量
// 限制整个进程「同时换直链+下载」的并发数(与 115 换链风控匹配,见 strmDownloadSemCap),
// 同时让下载充分并行。换链走全局令牌桶(QPS=3)兜底,下载走 CDN 不限速。
func (s *StrmService) downloadWorker(ctx context.Context) {
const claimBatch = 12 // 每次批量认领的任务数
for {
select {
case <-ctx.Done():
@@ -42,7 +48,7 @@ func (s *StrmService) downloadWorker(ctx context.Context) {
sleepContext(ctx, left)
continue
}
tasks, err := s.repo.StrmDownload.ClaimPendingDownload(ctx, 1)
tasks, err := s.repo.StrmDownload.ClaimPendingDownload(ctx, claimBatch)
if err != nil {
s.log.Warn("claim strm download task failed", zap.Error(err))
sleepContext(ctx, 3*time.Second)
@@ -52,9 +58,21 @@ func (s *StrmService) downloadWorker(ctx context.Context) {
sleepContext(ctx, 2*time.Second)
continue
}
// 并发处理本批任务:每个任务先获取全局下载槽位,槽位内部执行换链+下载。
// 信号量与令牌桶双重限速,确保任意时刻并发换链请求不超过安全阈值。
var wg sync.WaitGroup
for i := range tasks {
s.processDownloadTask(ctx, &tasks[i])
wg.Add(1)
go func(i int) {
defer wg.Done()
if !s.acquireDownloadSlot(ctx) {
return
}
defer s.releaseDownloadSlot()
s.processDownloadTask(ctx, &tasks[i])
}(i)
}
wg.Wait()
}
}
@@ -139,7 +157,7 @@ func (s *StrmService) processUploadTask(ctx context.Context, task *model.StrmUpl
}
}
if task.Provider == model.StrmProvider115 {
finish(model.StrmTaskFailed, "115 网盘暂不支持元数据上传")
s.processUpload115(ctx, task)
return
}
acct, err := s.repo.StrmAccount.FindByID(ctx, task.AccountID)
@@ -178,6 +196,48 @@ func (s *StrmService) processUploadTask(ctx context.Context, task *model.StrmUpl
finish(model.StrmTaskDone, "")
}
// processUpload115 115 元数据上传:task.RemotePath 存的是父目录 cid,FileName 为远端文件名。
func (s *StrmService) processUpload115(ctx context.Context, task *model.StrmUploadTask) {
finish := func(status, message string) {
now := time.Now()
task.Status = status
task.Error = message
task.FinishedAt = &now
if err := s.repo.StrmUpload.Update(context.Background(), task); err != nil {
s.log.Warn("update strm upload task failed", zap.Error(err))
}
}
acct, err := s.repo.StrmAccount.FindByID(ctx, task.AccountID)
if err != nil || acct == nil {
finish(model.StrmTaskFailed, "网盘账号不存在")
return
}
provider, err := s.providerFor(ctx, acct)
if err != nil {
s.uploadTaskFailWithRetry(task, err.Error())
return
}
named, ok := provider.(interface {
PutFileNamed(ctx context.Context, parentCID, fileName string, r io.Reader) error
})
if !ok {
finish(model.StrmTaskFailed, "该网盘不支持元数据上传")
return
}
f, err := os.Open(task.LocalPath)
if err != nil {
s.uploadTaskFailWithRetry(task, "打开本地文件失败:"+err.Error())
return
}
if err := named.PutFileNamed(ctx, task.RemotePath, task.FileName, f); err != nil {
_ = f.Close()
s.uploadTaskFailWithRetry(task, "上传失败:"+err.Error())
return
}
_ = f.Close()
finish(model.StrmTaskDone, "")
}
// downloadTaskFailWithRetry 下载失败任务按退避重试,超过上限标记 failed。
func (s *StrmService) downloadTaskFailWithRetry(task *model.StrmDownloadTask, message string) {
if !retryTask(&task.RetryCount, &task.Status, &task.Error, &task.NextTryAt, &task.FinishedAt, message) {
@@ -491,6 +551,44 @@ func (s *StrmService) RetryUploadTask(ctx context.Context, id string) error {
// ─── 下载队列批量操作(handler 使用) ─────────────────────────────────────────
// DeleteDownloadTask 删除一个下载任务记录。
func (s *StrmService) DeleteDownloadTask(ctx context.Context, id string) error {
return s.repo.StrmDownload.Delete(ctx, id)
}
// DeleteUploadTask 删除一个上传任务记录。
func (s *StrmService) DeleteUploadTask(ctx context.Context, id string) error {
return s.repo.StrmUpload.Delete(ctx, id)
}
// BatchActionDownloadTasks 对选中的下载任务执行批量操作(delete / retry / cancel)。
func (s *StrmService) BatchActionDownloadTasks(ctx context.Context, action string, ids []string) (int64, error) {
switch action {
case "delete":
return s.repo.StrmDownload.DeleteBatch(ctx, ids)
case "retry":
return s.repo.StrmDownload.RetryBatch(ctx, ids)
case "cancel":
return s.repo.StrmDownload.CancelBatch(ctx, ids)
default:
return 0, fmt.Errorf("不支持的批量操作: %s", action)
}
}
// BatchActionUploadTasks 对选中的上传任务执行批量操作(delete / retry / cancel)。
func (s *StrmService) BatchActionUploadTasks(ctx context.Context, action string, ids []string) (int64, error) {
switch action {
case "delete":
return s.repo.StrmUpload.DeleteBatch(ctx, ids)
case "retry":
return s.repo.StrmUpload.RetryBatch(ctx, ids)
case "cancel":
return s.repo.StrmUpload.CancelBatch(ctx, ids)
default:
return 0, fmt.Errorf("不支持的批量操作: %s", action)
}
}
// ClearDoneDownloadTasks 清空全部已完成下载记录,返回删除数量。
func (s *StrmService) ClearDoneDownloadTasks(ctx context.Context) (int64, error) {
return s.repo.StrmDownload.ClearDone(ctx)
@@ -501,6 +599,16 @@ func (s *StrmService) ClearFinishedDownloadTasks(ctx context.Context) (int64, er
return s.repo.StrmDownload.ClearFinished(ctx)
}
// ClearCanceledDownloadTasks 清空全部已取消的下载记录,返回删除数量。
func (s *StrmService) ClearCanceledDownloadTasks(ctx context.Context) (int64, error) {
return s.repo.StrmDownload.ClearCanceled(ctx)
}
// ClearCanceledUploadTasks 清空全部已取消的上传记录,返回删除数量。
func (s *StrmService) ClearCanceledUploadTasks(ctx context.Context) (int64, error) {
return s.repo.StrmUpload.ClearCanceled(ctx)
}
// RetryAllFailedDownloadTasks 批量重试所有失败下载任务,返回重新入队数量。
func (s *StrmService) RetryAllFailedDownloadTasks(ctx context.Context) (int64, error) {
return s.repo.StrmDownload.RetryAllFailed(ctx)
@@ -511,6 +619,11 @@ func (s *StrmService) CancelPendingDownloadTasks(ctx context.Context) (int64, er
return s.repo.StrmDownload.CancelPending(ctx)
}
// CancelPendingUploadTasks 批量取消所有排队上传任务,返回取消数量。
func (s *StrmService) CancelPendingUploadTasks(ctx context.Context) (int64, error) {
return s.repo.StrmUpload.CancelPending(ctx)
}
func sleepContext(ctx context.Context, d time.Duration) {
select {
case <-ctx.Done():
+94
View File
@@ -90,11 +90,48 @@ type StrmService struct {
running map[string]context.CancelFunc // sync path id -> cancel
oauthSessions map[string]*strm115AuthSession
wafUntil time.Time // 115 风控/限流熔断截止时间(由 mu 保护)
downloadSem chan struct{} // 全局下载并发信号量:限制整个进程同时进行「换直链+下载」的并发数
downloadSemOnce sync.Once
}
// strmWAFCooldown 检测到 115 风控/限流后下载队列的全局冷却时长。
const strmWAFCooldown = 3 * time.Minute
// strmDownloadSemCap 全局同时进行「换直链+下载」的并发上限。
//
// 115 对换直链接口(/open/ufile/downurl)风控极严:过去把全局 QPS 提到 8 或让多
// worker 高并发换链,会瞬时撞上 WAF 返回 405 阻断页并触发 180 秒冷却,反而更慢。
// 因此用信号量把整个进程同时换直链的并发数压到 3,与令牌桶限速共同兜底:
// 宁可下载稍慢,也绝不触发风控。下载本身走 CDN 不限速。
const strmDownloadSemCap = 3
// ensureDownloadSem 惰性初始化全局共享的下载并发信号量。
func (s *StrmService) ensureDownloadSem() {
s.downloadSemOnce.Do(func() {
s.downloadSem = make(chan struct{}, strmDownloadSemCap)
})
}
// acquireDownloadSlot 获取一个下载并发槽位(等待/取消安全)。
func (s *StrmService) acquireDownloadSlot(ctx context.Context) bool {
s.ensureDownloadSem()
select {
case s.downloadSem <- struct{}{}:
return true
case <-ctx.Done():
return false
}
}
// releaseDownloadSlot 释放一个下载并发槽位。
func (s *StrmService) releaseDownloadSlot() {
if s.downloadSem == nil {
return
}
<-s.downloadSem
}
// NewStrmService constructs the STRM service.
func NewStrmService(cfg *config.Config, log *zap.Logger, repos *repository.Container, crypto *CryptoService) *StrmService {
return &StrmService{
@@ -113,6 +150,7 @@ func NewStrmService(cfg *config.Config, log *zap.Logger, repos *repository.Conta
// Start 启动下载/上传队列 worker、定时同步巡检、115 token 刷新与队列清理。
func (s *StrmService) Start(ctx context.Context) {
s.sync115RelayKey(ctx)
s.recoverInterruptedSyncs(ctx)
downloadThreads := s.strmIntSetting(ctx, StrmSettingDownloadThreads, 3)
if downloadThreads < 1 {
downloadThreads = 1
@@ -141,6 +179,21 @@ func (s *StrmService) Start(ctx context.Context) {
zap.Int("upload_threads", uploadThreads))
}
// recoverInterruptedSyncs 在服务启动时自愈重置因服务重启遗留的 running 状态。
func (s *StrmService) recoverInterruptedSyncs(ctx context.Context) {
paths, err := s.repo.StrmSyncPath.List(ctx)
if err == nil {
for i := range paths {
p := &paths[i]
if p.LastSyncStatus == model.StrmSyncRecordRunning {
p.LastSyncStatus = model.StrmSyncRecordCanceled
p.LastSyncMessage = "服务重启,已重置同步状态"
_ = s.repo.StrmSyncPath.Update(ctx, p)
}
}
}
}
func (s *StrmService) Stop() {
s.stopOnce.Do(func() { close(s.stopCh) })
}
@@ -396,6 +449,38 @@ func (s *StrmService) ListSyncRecords(ctx context.Context, pathID string, limit
return s.repo.StrmSyncRecord.List(ctx, pathID, limit)
}
// DeleteSyncRecord 删除单条同步记录。
func (s *StrmService) DeleteSyncRecord(ctx context.Context, id string) error {
if err := s.repo.StrmSyncRecord.Delete(ctx, id); err != nil {
return err
}
return nil
}
// ClearSyncRecords 清空某同步目录(pathID 为空则全部)的同步记录,返回删除条数。
func (s *StrmService) ClearSyncRecords(ctx context.Context, pathID string) (int64, error) {
if pathID != "" {
return s.repo.StrmSyncRecord.DeleteBySyncPathID(ctx, pathID)
}
var total int64
// 全量清空:分页拉取物理删除所有记录
for {
rows, err := s.repo.StrmSyncRecord.List(ctx, "", 200)
if err != nil {
return total, err
}
if len(rows) == 0 {
return total, nil
}
for _, rec := range rows {
if err := s.repo.StrmSyncRecord.Delete(ctx, rec.ID); err != nil {
return total, err
}
}
total += int64(len(rows))
}
}
// CreateSyncPath 校验并创建同步目录。
func (s *StrmService) CreateSyncPath(ctx context.Context, p *model.StrmSyncPath) (*model.StrmSyncPath, error) {
if err := s.validateSyncPath(ctx, p); err != nil {
@@ -404,6 +489,9 @@ func (s *StrmService) CreateSyncPath(ctx context.Context, p *model.StrmSyncPath)
if strings.TrimSpace(p.Name) == "" {
p.Name = "同步目录 " + time.Now().Format("01-02 15:04")
}
if p.SyncMode == "" {
p.SyncMode = model.StrmSyncTypeIncremental
}
if p.EnableCron && strings.TrimSpace(p.Cron) == "" {
return nil, errors.New("启用定时同步需要填写 cron 表达式")
}
@@ -430,6 +518,12 @@ func (s *StrmService) UpdateSyncPath(ctx context.Context, id string, p *model.St
p.LastSyncAt = existing.LastSyncAt
p.LastSyncStatus = existing.LastSyncStatus
p.LastSyncMessage = existing.LastSyncMessage
if p.SyncMode == "" {
p.SyncMode = existing.SyncMode
if p.SyncMode == "" {
p.SyncMode = model.StrmSyncTypeIncremental
}
}
if p.EnableCron && strings.TrimSpace(p.Cron) == "" {
return nil, errors.New("启用定时同步需要填写 cron 表达式")
}
+574 -40
View File
@@ -20,6 +20,7 @@ import (
"github.com/ShukeBta/MMTL/internal/model"
"github.com/ShukeBta/MMTL/internal/service/cloud"
"github.com/ShukeBta/MMTL/internal/service/cloud115"
)
// strmSyncState 是一次同步执行的上下文。
@@ -31,16 +32,29 @@ type strmSyncState struct {
provider cloud.Provider // local 提供方为 nil
cfg *strmPathConfig
rec *model.StrmSyncRecord
syncType string
mu sync.Mutex
processed int // 已处理文件计数(用于定期落库进度)
seenVideo map[string]bool // "v:"+去掉扩展名的相对路径 → 远端存在该视频
seenMeta map[string]bool // "m:"+相对路径 → 远端存在该元数据
remoteMeta map[string]int64 // 远端元数据大小(上传比对用)
mu sync.Mutex
processed int // 已处理文件计数(用于定期落库进度)
lastProgressFlush time.Time // 上次进度落库时间
seenVideo map[string]bool // "v:"+去掉扩展名的相对路径 → 远端存在该视频
seenMeta map[string]bool // "m:"+相对路径 → 远端存在该元数据
remoteMeta map[string]int64 // 远端元数据大小(上传比对用)
seenMetaTarget map[string]cloud.FileEntry
seenVideoTarget map[string]cloud.FileEntry
activeDownloadPaths map[string]bool // 本地已在排队/进行的下载任务路径(内存去重)
activeUploadPaths map[string]bool // 本地已在排队/进行的上传任务路径(内存去重)
pendingDownloads []*model.StrmDownloadTask
pendingUploads []*model.StrmUploadTask
dirCache sync.Map // dirID (string) -> relativePath (string)
dirPathToID map[string]string // relativePath (string) -> dirID(115 上传父目录寻址用,walk 后构建)
scanIncomplete atomic.Bool // 远端目录树/文件列表本次扫描不完整 → 禁止增量 prune 误删本地文件
}
// StartSync 启动一次同步(异步执行,同一目录同时只允许一个任务)。
func (s *StrmService) StartSync(ctx context.Context, pathID string) error {
// syncType 支持 "incremental"(默认增量)和 "full"(全量同步)。
func (s *StrmService) StartSync(ctx context.Context, pathID string, syncType ...string) error {
p, err := s.repo.StrmSyncPath.FindByID(ctx, pathID)
if err != nil || p == nil {
return errNotFoundOr(err, "同步目录不存在")
@@ -64,9 +78,20 @@ func (s *StrmService) StartSync(ctx context.Context, pathID string) error {
s.running[pathID] = cancel
s.mu.Unlock()
mode := model.StrmSyncTypeIncremental
if len(syncType) > 0 && syncType[0] != "" {
mode = syncType[0]
} else if p.SyncMode != "" {
mode = p.SyncMode
}
if mode != model.StrmSyncTypeFull {
mode = model.StrmSyncTypeIncremental
}
now := time.Now()
rec := &model.StrmSyncRecord{
SyncPathID: pathID,
SyncType: mode,
Status: model.StrmSyncRecordRunning,
StartedAt: &now,
}
@@ -84,15 +109,27 @@ func (s *StrmService) StartSync(ctx context.Context, pathID string) error {
return nil
}
// CancelSync 取消正在进行的同步。
// CancelSync 取消正在进行的同步(若为僵尸运行状态则直接自愈重置)。
func (s *StrmService) CancelSync(ctx context.Context, pathID string) error {
s.mu.Lock()
defer s.mu.Unlock()
cancel, exists := s.running[pathID]
if !exists {
return errors.New("该目录当前没有进行中的同步")
if exists {
delete(s.running, pathID)
}
s.mu.Unlock()
if exists && cancel != nil {
cancel()
}
// 无论内存中是否活跃,确保同步目录状态正确重置为已取消
if p, err := s.repo.StrmSyncPath.FindByID(ctx, pathID); err == nil && p != nil {
if p.LastSyncStatus == model.StrmSyncRecordRunning {
p.LastSyncStatus = model.StrmSyncRecordCanceled
p.LastSyncMessage = "已取消"
_ = s.repo.StrmSyncPath.Update(ctx, p)
}
}
cancel()
return nil
}
@@ -140,14 +177,17 @@ func (s *StrmService) runSync(ctx context.Context, p *model.StrmSyncPath, rec *m
return
}
st := &strmSyncState{
s: s,
ctx: ctx,
p: p,
cfg: cfg,
rec: rec,
seenVideo: map[string]bool{},
seenMeta: map[string]bool{},
remoteMeta: map[string]int64{},
s: s,
ctx: ctx,
p: p,
cfg: cfg,
rec: rec,
syncType: rec.SyncType,
seenVideo: map[string]bool{},
seenMeta: map[string]bool{},
remoteMeta: map[string]int64{},
seenMetaTarget: map[string]cloud.FileEntry{},
seenVideoTarget: map[string]cloud.FileEntry{},
}
if p.Provider != model.StrmProviderLocal {
acct, err := s.repo.StrmAccount.FindByID(ctx, p.AccountID)
@@ -191,36 +231,77 @@ func (s *StrmService) finishSync(p *model.StrmSyncPath, rec *model.StrmSyncRecor
p.LastSyncStatus = status
p.LastSyncMessage = message
if status != model.StrmSyncRecordFailed && message == "" {
p.LastSyncMessage = fmt.Sprintf("完成:新增/更新 %d 个 strm,下载 %d 个元数据,清理 %d 个文件",
rec.NewStrm, rec.NewMeta, rec.Pruned)
syncTypeLabel := "增量"
if rec.SyncType == model.StrmSyncTypeFull {
syncTypeLabel = "全量"
}
p.LastSyncMessage = fmt.Sprintf("[%s] 完成:新增/更新 %d 个 strm,跳过 %d 个,下载 %d 个元数据,上传 %d 个元数据,清理 %d 个文件",
syncTypeLabel, rec.NewStrm, rec.Skipped, rec.NewMeta, rec.Uploaded, rec.Pruned)
}
if err := s.repo.StrmSyncPath.Update(context.Background(), p); err != nil {
s.log.Warn("update strm sync path failed", zap.Error(err))
}
s.log.Info("strm sync finished",
zap.String("path_id", p.ID), zap.String("status", status),
zap.Int64("new_strm", rec.NewStrm), zap.Int64("new_meta", rec.NewMeta),
zap.Int64("pruned", rec.Pruned), zap.String("message", message))
zap.String("path_id", p.ID), zap.String("sync_type", rec.SyncType), zap.String("status", status),
zap.Int64("new_strm", rec.NewStrm), zap.Int64("skipped", rec.Skipped), zap.Int64("new_meta", rec.NewMeta),
zap.Int64("uploaded", rec.Uploaded), zap.Int64("pruned", rec.Pruned), zap.String("message", message))
}
func (st *strmSyncState) run() error {
if err := ensureLocalDir(st.p.LocalPath); err != nil {
return fmt.Errorf("创建输出目录失败:%w", err)
}
if st.cfg.DownloadMeta {
if active, err := st.s.repo.StrmDownload.GetActiveLocalPathMap(st.ctx, st.p.ID); err == nil {
st.activeDownloadPaths = active
} else {
st.activeDownloadPaths = map[string]bool{}
}
}
if st.cfg.UploadMeta {
if active, err := st.s.repo.StrmUpload.GetActiveLocalPathMap(st.ctx, st.p.ID); err == nil {
st.activeUploadPaths = active
} else {
st.activeUploadPaths = map[string]bool{}
}
}
if st.provider != nil {
if err := st.walkRemote(); err != nil {
return err
if open115, ok := st.provider.(cloud.OpenAPI115Provider); ok && st.p.Provider == model.StrmProvider115 {
if err := st.walk115Flat(open115.OpenClient()); err != nil {
return err
}
} else {
if err := st.walkRemote(); err != nil {
return err
}
}
} else {
if err := st.walkLocalSource(); err != nil {
return err
}
}
st.flushPendingDownloads()
st.flushProgress()
if st.cfg.UploadMeta && st.provider != nil && st.p.Provider != model.StrmProvider115 {
if st.cfg.UploadMeta && st.provider != nil {
// 115 上传需要父目录 cid,先用 dirCache 构建「路径 → cid」反向索引
if st.p.Provider == model.StrmProvider115 {
reversed := map[string]string{}
st.dirCache.Range(func(key, value any) bool {
path, ok := value.(string)
if ok && path != "" {
if id, ok2 := key.(string); ok2 {
reversed[path] = id
}
}
return true
})
st.dirPathToID = reversed
}
if err := st.scanLocalMetaForUpload(); err != nil {
return err
}
st.flushPendingUploads()
}
if err := st.pruneLocal(); err != nil {
return err
@@ -239,6 +320,7 @@ const strmScanWorkers = 8
// 多个 worker 并行执行 List(受全局 115 令牌桶限流约束),子目录动态
// 入队;任一目录失败则取消其余 worker 并返回错误(与旧串行版语义一致)。
func (st *strmSyncState) walkRemote() error {
defer st.flushPendingDownloads()
root := strings.TrimSpace(st.p.RemotePath)
if root == "" {
root = "/"
@@ -374,6 +456,286 @@ func (st *strmSyncState) isMetaExt(ext string) bool {
return false
}
// cleanDirRel 对 115 扁平化拉取的目录相对路径逐段套用目录级文件名清洗,
// 确保与 walkRemote / joinLocalRel(sanitizeRelativePath)使用同一套清洗规则。
// 若不清洗,目录名中的冒号等非法字符会直达 rel,而 seenVideo/seenMeta 的 key
// 与磁盘实际路径不一致,导致 pruneLocal 误删已下载的 strm / 元数据。
// 空 rel(根目录)原样返回。
func cleanDirRel(rel string) string {
if rel == "" {
return ""
}
parts := strings.Split(rel, "/")
out := make([]string, 0, len(parts))
for _, part := range parts {
if part == "" {
continue
}
clean := cleanEntryName(part, true)
if clean != "" && clean != "." && clean != ".." {
out = append(out, clean)
}
}
return strings.Join(out, "/")
}
// walk115Flat 使用 115 开放平台扁平化分页批量拉取机制与目录拓扑缓存(参考 QMediaSync)。
// 极大地降低 API 请求次数并支持毫秒级/秒级增量同步。
func (st *strmSyncState) walk115Flat(open115 *cloud115.OpenClient) error {
defer st.flushPendingDownloads()
ctx, cancel := context.WithCancel(st.ctx)
defer cancel()
rootCID := strings.TrimSpace(st.p.RemotePath)
if rootCID == "" {
rootCID = "0"
}
// 1. 目录拓扑缓存处理
st.dirCache.Store(rootCID, "")
if st.syncType == model.StrmSyncTypeFull {
// 全量同步:清空本路径的历史目录缓存
if err := st.s.repo.StrmDirCache.DeleteBySyncPathID(ctx, st.p.ID); err != nil {
st.s.log.Warn("delete strm dir cache failed", zap.Error(err))
}
} else {
// 增量同步:预加载历史目录缓存(过滤历史一对多塌陷冲突的脏数据以自愈刷新)
cached, err := st.s.repo.StrmDirCache.ListBySyncPathID(ctx, st.p.ID)
if err == nil {
pathCounts := make(map[string]int, len(cached))
for _, item := range cached {
pathCounts[item.Path]++
}
for _, item := range cached {
// 若同一个 path 对应了多个不同 dir_id,说明包含历史层级塌陷的脏数据,不预加载,让后续步骤重新向 115 获取精确路径
if pathCounts[item.Path] > 1 {
continue
}
st.dirCache.Store(item.DirID, cleanDirRel(item.Path))
}
}
}
// 2. 探测文件总数
const pageSize = 1150
firstBatch, totalCount, err := open115.GetFsListFlat(ctx, rootCID, 0, pageSize)
if err != nil {
return fmt.Errorf("115: 获取文件列表失败:%w", err)
}
st.updateSyncMessage(fmt.Sprintf("正在拉取远端文件列表 (共 %d 个文件)...", totalCount))
allFiles := make([]cloud115.RemoteFile, 0, totalCount)
allFiles = append(allFiles, firstBatch...)
// 3. 并发分页拉取剩余文件
if totalCount > int64(len(firstBatch)) {
totalPages := int((totalCount + pageSize - 1) / pageSize)
type pageTask struct {
offset int
}
pageTasks := make([]pageTask, 0, totalPages-1)
for page := 1; page < totalPages; page++ {
pageTasks = append(pageTasks, pageTask{offset: page * pageSize})
}
var (
filesMu sync.Mutex
wg sync.WaitGroup
taskCh = make(chan pageTask, len(pageTasks))
errMu sync.Mutex
fetchErr error
)
for _, t := range pageTasks {
taskCh <- t
}
close(taskCh)
workers := 8
if len(pageTasks) < workers {
workers = len(pageTasks)
}
for i := 0; i < workers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for t := range taskCh {
if ctx.Err() != nil {
return
}
files, _, err := open115.GetFsListFlat(ctx, rootCID, t.offset, pageSize)
if err != nil {
errMu.Lock()
if fetchErr == nil {
fetchErr = err
}
errMu.Unlock()
return
}
filesMu.Lock()
allFiles = append(allFiles, files...)
filesMu.Unlock()
}
}()
}
wg.Wait()
if fetchErr != nil {
return fmt.Errorf("115: 分页拉取失败:%w", fetchErr)
}
}
if ctx.Err() != nil {
return ctx.Err()
}
// 4. 收集所有未在缓存中的父目录 ID (file.Pid)
missingPids := make(map[string]struct{})
for _, f := range allFiles {
pid := f.Pid
if pid == "" || pid == rootCID {
continue
}
if _, ok := st.dirCache.Load(pid); !ok {
missingPids[pid] = struct{}{}
}
}
// 并发补全未知目录详情与祖先链
if len(missingPids) > 0 {
pidList := make([]string, 0, len(missingPids))
for pid := range missingPids {
pidList = append(pidList, pid)
}
pidCh := make(chan string, len(pidList))
for _, pid := range pidList {
pidCh <- pid
}
close(pidCh)
var (
pwg sync.WaitGroup
dirWorkers = 8
doneDirs atomic.Int64
totalDirs = len(pidList)
errMu sync.Mutex
firstErr error
)
if len(pidList) < dirWorkers {
dirWorkers = len(pidList)
}
st.updateSyncMessage(fmt.Sprintf("正在解析目录树 (0/%d)...", totalDirs))
for i := 0; i < dirWorkers; i++ {
pwg.Add(1)
go func() {
defer pwg.Done()
for pid := range pidCh {
if ctx.Err() != nil {
return
}
if _, loaded := st.dirCache.Load(pid); loaded {
if n := doneDirs.Add(1); n%20 == 0 || n == int64(totalDirs) {
st.updateSyncMessage(fmt.Sprintf("正在解析目录树 (%d/%d)...", n, totalDirs))
}
continue
}
detail, err := open115.GetFsDetailByCid(ctx, pid)
if err != nil {
// 目录详情解析失败会导致下游文件 rel 无法还原真实父路径,
// seen key 与磁盘路径对不上:增量 prune 会误删本地文件、上传会
// 误传本地未变文件、下载会重复下载。这里不是降级容错,而是
// 直接中止整个同步——宁可本次同步失败,也不带着损坏的相对路径
// 继续执行造成大规模误删/误传/重下(参考用户反馈"云盘没动却重下重传")。
errMu.Lock()
if firstErr == nil {
firstErr = fmt.Errorf("115: 解析目录树失败(file_id=%s):%w", pid, err)
}
errMu.Unlock()
st.scanIncomplete.Store(true)
cancel()
return
} else if detail != nil {
// 解析相对路径
relPath := cleanDirRel(detail.RelativePath(rootCID))
st.dirCache.Store(pid, relPath)
_ = st.s.repo.StrmDirCache.Set(ctx, st.p.ID, pid, relPath)
// 顺便解析并缓存 detail.Paths 中包含的中间各层级目录
for _, ancestor := range detail.Paths {
if ancestor.FileId == "0" || ancestor.FileId == rootCID {
continue
}
if _, loaded := st.dirCache.Load(ancestor.FileId); !loaded {
subDetail := &cloud115.RemoteFileDetail{
FileId: ancestor.FileId,
FileName: ancestor.Name,
Paths: nil,
}
for _, p := range detail.Paths {
subDetail.Paths = append(subDetail.Paths, p)
if p.FileId == ancestor.FileId {
break
}
}
ancestorRel := cleanDirRel(subDetail.RelativePath(rootCID))
st.dirCache.Store(ancestor.FileId, ancestorRel)
_ = st.s.repo.StrmDirCache.Set(ctx, st.p.ID, ancestor.FileId, ancestorRel)
}
}
}
if n := doneDirs.Add(1); n%10 == 0 || n == int64(totalDirs) {
st.updateSyncMessage(fmt.Sprintf("正在解析目录树 (%d/%d)...", n, totalDirs))
}
}
}()
}
pwg.Wait()
if firstErr != nil {
// 目录树解析失败会导致 rel 塌缩,若继续处理会让大量本地文件
// 被错误判定为"云端不存在"而重复下载/上传,并可能误删本地文件。
// 中止本次同步,避免在损坏的相对路径上执行任何写操作。
return firstErr
}
}
st.updateSyncMessage(fmt.Sprintf("正在生成 STRM 与同步文件 (共 %d 个)...", len(allFiles)))
// 5. 分类处理所有文件
for _, f := range allFiles {
if ctx.Err() != nil {
return ctx.Err()
}
cleanName := cleanEntryName(f.FileName, false)
var rel string
if f.Pid == "" || f.Pid == rootCID {
rel = cleanName
} else {
if parentVal, ok := st.dirCache.Load(f.Pid); ok && parentVal.(string) != "" {
rel = cleanDirRel(parentVal.(string)) + "/" + cleanName
} else {
// 父目录不在目录缓存,无法还原真实相对路径。若继续用塌缩后的
// 根路径处理,该文件会被错误判定,导致重复下载/上传或误删本地文件。
// 目录树不完整时宁可中止本次同步,也不带着损坏的 rel 继续执行。
return fmt.Errorf("115: 文件 %s 的父目录未解析成功,目录树不完整,中止同步以防误删/误传", cleanName)
}
}
entry := cloud.FileEntry{
ID: f.FileId,
Name: f.FileName,
IsDir: false,
Size: f.FileSize,
MTime: f.Utime,
PickCode: f.PickCode,
}
st.processRemoteFile(entry, rel)
}
return nil
}
// handleVideo 生成/更新 .strm 文件。
func (st *strmSyncState) handleVideo(entry cloud.FileEntry, rel, ext string) {
relSansExt := rel[:len(rel)-len(ext)]
@@ -391,6 +753,30 @@ func (st *strmSyncState) handleVideo(entry cloud.FileEntry, rel, ext string) {
st.s.log.Warn("strm target path out of root", zap.String("rel", targetRel), zap.Error(err))
return
}
st.mu.Lock()
if st.seenVideoTarget == nil {
st.seenVideoTarget = map[string]cloud.FileEntry{}
}
if _, exists := st.seenVideoTarget[target]; exists {
st.mu.Unlock()
st.touchProgress()
return
}
st.seenVideoTarget[target] = entry
st.mu.Unlock()
// 增量同步模式快速检查:本地 strm 文件存在、非空且修改时间与远端 mtime 一致,直接跳过无需读磁盘
if st.syncType == model.StrmSyncTypeIncremental && entry.MTime > 0 {
if info, err := os.Stat(target); err == nil && info.Size() > 0 && info.ModTime().Unix() == entry.MTime {
st.mu.Lock()
st.rec.Skipped++
st.mu.Unlock()
st.touchProgress()
return
}
}
content, err := st.strmContent(entry, rel, ext)
if err != nil {
// 并发 worker 下 rec.Message 无锁写会有数据竞争,这里仅记录日志;
@@ -403,6 +789,11 @@ func (st *strmSyncState) handleVideo(entry cloud.FileEntry, rel, ext string) {
existing = string(data)
}
if existing == content {
// 对齐本地 strm 修改时间为远端 mtime,便于后续秒级比对
if entry.MTime > 0 {
mTime := time.Unix(entry.MTime, 0)
_ = os.Chtimes(target, mTime, mTime)
}
st.mu.Lock()
st.rec.Skipped++
st.mu.Unlock()
@@ -423,6 +814,10 @@ func (st *strmSyncState) handleVideo(entry cloud.FileEntry, rel, ext string) {
st.s.log.Warn("rename strm failed", zap.String("file", target), zap.Error(err))
return
}
if entry.MTime > 0 {
mTime := time.Unix(entry.MTime, 0)
_ = os.Chtimes(target, mTime, mTime)
}
st.mu.Lock()
st.rec.NewStrm++
st.mu.Unlock()
@@ -482,6 +877,36 @@ func (st *strmSyncState) recordRemoteMeta(entry cloud.FileEntry, rel string) {
st.mu.Unlock()
}
func (st *strmSyncState) flushPendingDownloads() {
st.mu.Lock()
if len(st.pendingDownloads) == 0 {
st.mu.Unlock()
return
}
batch := st.pendingDownloads
st.pendingDownloads = nil
st.mu.Unlock()
if err := st.s.repo.StrmDownload.CreateInBatches(st.ctx, batch, 100); err != nil {
st.s.log.Warn("batch enqueue strm download tasks failed", zap.Error(err))
}
}
func (st *strmSyncState) flushPendingUploads() {
st.mu.Lock()
if len(st.pendingUploads) == 0 {
st.mu.Unlock()
return
}
batch := st.pendingUploads
st.pendingUploads = nil
st.mu.Unlock()
if err := st.s.repo.StrmUpload.CreateInBatches(st.ctx, batch, 100); err != nil {
st.s.log.Warn("batch enqueue strm upload tasks failed", zap.Error(err))
}
}
// handleMeta 元数据入下载队列(本地已存在且大小一致则跳过)。
func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) {
st.recordRemoteMeta(entry, rel)
@@ -490,14 +915,40 @@ func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) {
if err != nil {
return
}
st.mu.Lock()
if st.seenMetaTarget == nil {
st.seenMetaTarget = map[string]cloud.FileEntry{}
}
if _, exists := st.seenMetaTarget[target]; exists {
// 该本地目标路径在当前批次中已被处理(存在同名/重名冲突),直接忽略重复项,避免多份不同大小的文件在本地交替覆盖导致增量死循环
st.mu.Unlock()
st.touchProgress()
return
}
st.seenMetaTarget[target] = entry
st.mu.Unlock()
if info, err := os.Stat(target); err == nil && info.Size() == entry.Size {
st.touchProgress()
return
}
if st.taskExists("download", st.p.ID, target) {
st.mu.Lock()
if st.activeDownloadPaths == nil {
if active, err := st.s.repo.StrmDownload.GetActiveLocalPathMap(st.ctx, st.p.ID); err == nil {
st.activeDownloadPaths = active
} else {
st.activeDownloadPaths = map[string]bool{}
}
}
if st.activeDownloadPaths[target] {
st.mu.Unlock()
st.touchProgress()
return
}
st.activeDownloadPaths[target] = true
st.mu.Unlock()
task := &model.StrmDownloadTask{
SyncPathID: st.p.ID,
AccountID: st.p.AccountID,
@@ -516,13 +967,16 @@ func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) {
if st.p.Provider != model.StrmProvider115 {
task.RemoteRef = entry.ID
}
if err := st.s.repo.StrmDownload.Create(st.ctx, task); err != nil {
st.s.log.Warn("enqueue strm download task failed", zap.Error(err))
return
}
st.mu.Lock()
st.pendingDownloads = append(st.pendingDownloads, task)
shouldFlush := len(st.pendingDownloads) >= 100
st.rec.NewMeta++
st.mu.Unlock()
if shouldFlush {
st.flushPendingDownloads()
}
st.touchProgress()
}
@@ -580,10 +1034,22 @@ func (st *strmSyncState) walkLocalSource() error {
if err != nil {
return nil
}
mTime := info.ModTime()
if st.syncType == model.StrmSyncTypeIncremental {
if tInfo, err := os.Stat(target); err == nil && tInfo.Size() > 0 && tInfo.ModTime().Unix() == mTime.Unix() {
st.mu.Lock()
st.rec.Skipped++
st.mu.Unlock()
st.touchProgress()
return nil
}
}
if data, err := os.ReadFile(target); err == nil && string(data) == content {
_ = os.Chtimes(target, mTime, mTime)
st.mu.Lock()
st.rec.Skipped++
st.mu.Unlock()
st.touchProgress()
return nil
}
if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil {
@@ -592,18 +1058,28 @@ func (st *strmSyncState) walkLocalSource() error {
tmp := target + ".tmp"
if err := os.WriteFile(tmp, []byte(content), 0o644); err == nil {
_ = os.Rename(tmp, target)
_ = os.Chtimes(target, mTime, mTime)
} else {
_ = os.Remove(tmp)
}
st.mu.Lock()
st.rec.NewStrm++
st.mu.Unlock()
st.touchProgress()
return nil
})
}
// scanLocalMetaForUpload 扫描本地元数据,与远端比对后入上传队列。
func (st *strmSyncState) scanLocalMetaForUpload() error {
defer st.flushPendingUploads()
if st.activeUploadPaths == nil {
if active, err := st.s.repo.StrmUpload.GetActiveLocalPathMap(st.ctx, st.p.ID); err == nil {
st.activeUploadPaths = active
} else {
st.activeUploadPaths = map[string]bool{}
}
}
localRoot := filepath.Clean(st.p.LocalPath)
return filepath.WalkDir(localRoot, func(path string, d os.DirEntry, err error) error {
if err != nil {
@@ -640,26 +1116,35 @@ func (st *strmSyncState) scanLocalMetaForUpload() error {
// 网盘端已存在该元数据文件,跳过上传
return nil
}
if st.taskExists("upload", st.p.ID, path) {
st.mu.Lock()
if st.activeUploadPaths != nil && st.activeUploadPaths[path] {
st.mu.Unlock()
return nil
}
if st.activeUploadPaths != nil {
st.activeUploadPaths[path] = true
}
st.mu.Unlock()
task := &model.StrmUploadTask{
SyncPathID: st.p.ID,
AccountID: st.p.AccountID,
Provider: st.p.Provider,
FileName: filepath.Base(rel),
LocalPath: path,
RemotePath: st.remoteUploadPath(rel),
RemotePath: st.uploadRemoteTarget(rel),
Size: info.Size(),
Status: model.StrmTaskPending,
}
if err := st.s.repo.StrmUpload.Create(st.ctx, task); err != nil {
st.s.log.Warn("enqueue strm upload task failed", zap.Error(err))
return nil
}
st.mu.Lock()
st.pendingUploads = append(st.pendingUploads, task)
shouldFlush := len(st.pendingUploads) >= 100
st.rec.Uploaded++
st.mu.Unlock()
if shouldFlush {
st.flushPendingUploads()
}
return nil
})
}
@@ -673,6 +1158,31 @@ func (st *strmSyncState) remoteUploadPath(rel string) string {
return root + "/" + rel
}
// uploadRemoteTarget 返回上传任务的目标远端描述。
// - 115:返回父目录 cid(供 PutFileNamed 定位),基于 dirPathToID 把父目录相对路径映射到 cid。
// - 网盘桥接(clouddrive2/openlist):返回完整远端路径。
func (st *strmSyncState) uploadRemoteTarget(rel string) string {
if st.p.Provider == model.StrmProvider115 {
dir := rel
if idx := strings.LastIndexByte(dir, '/'); idx >= 0 {
dir = dir[:idx]
} else {
dir = ""
}
if dir == "" {
// 文件在同步根目录下,父目录即 115 同步根目录 ID
return st.p.RemotePath
}
if cid, ok := st.dirPathToID[dir]; ok && cid != "" {
return cid
}
// 父目录未在缓存中(父目录可能本次未扫描到),降级为用户配置的同步根 cid,
// 由上传端尽力处理(可能失败记日志,不影响下载)。
return st.p.RemotePath
}
return st.remoteUploadPath(rel)
}
// taskExists 检查是否已有同目录、同目标的进行中/已完成任务(避免重复入队)。
func (st *strmSyncState) taskExists(kind, syncPathID, localPath string) bool {
ctx := st.ctx
@@ -688,6 +1198,14 @@ func (st *strmSyncState) taskExists(kind, syncPathID, localPath string) bool {
// pruneLocal 清理本地多余 .strm 与元数据(远端已不存在),可选删除空目录。
func (st *strmSyncState) pruneLocal() error {
// 增量同步保护:本次远端扫描不完整(目录详情解析失败 / 文件父路径降级)时,
// seenVideo/seenMeta 覆盖不全,按"远端不存在"清理会误删刚下载或已存在的本地文件,
// 进而触发"下次增量重新下载"的循环。此时跳过清理,仅做进度落库。
if st.syncType == model.StrmSyncTypeIncremental && st.scanIncomplete.Load() {
st.s.log.Warn("strm 增量同步跳过清理:本次远端扫描不完整,prune 已禁用",
zap.String("path_id", st.p.ID))
return nil
}
localRoot := filepath.Clean(st.p.LocalPath)
var dirs []string
err := filepath.WalkDir(localRoot, func(path string, d os.DirEntry, err error) error {
@@ -748,12 +1266,16 @@ func (st *strmSyncState) pruneLocal() error {
return nil
}
// touchProgress 每处理若干个文件落库一次进度。
// touchProgress 进度计数并限流防抖落库(避免高频写 SQLite 导致锁竞争)。
func (st *strmSyncState) touchProgress() {
st.mu.Lock()
st.rec.Total++
st.processed++
flush := st.processed%100 == 0
now := time.Now()
flush := st.processed%100 == 0 || (st.processed%20 == 0 && now.Sub(st.lastProgressFlush) >= 2*time.Second)
if flush {
st.lastProgressFlush = now
}
st.mu.Unlock()
if flush {
st.flushProgress()
@@ -769,6 +1291,18 @@ func (st *strmSyncState) flushProgress() {
}
}
// updateSyncMessage 实时更新同步阶段提示信息,让前端界面清晰了解当前进度。
func (st *strmSyncState) updateSyncMessage(msg string) {
st.mu.Lock()
st.rec.Message = msg
st.p.LastSyncMessage = msg
rec := *st.rec
p := *st.p
st.mu.Unlock()
_ = st.s.repo.StrmSyncRecord.Update(st.ctx, &rec)
_ = st.s.repo.StrmSyncPath.Update(st.ctx, &p)
}
// ─── 定时同步巡检 ──────────────────────────────────────────────────────────────
func (s *StrmService) cronLoop(ctx context.Context) {
+263 -1
View File
@@ -2,10 +2,13 @@ package service
import (
"context"
"net/http"
"net/http/httptest"
"net/url"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
@@ -18,6 +21,7 @@ import (
"github.com/ShukeBta/MMTL/internal/model"
"github.com/ShukeBta/MMTL/internal/repository"
"github.com/ShukeBta/MMTL/internal/service/cloud"
"github.com/ShukeBta/MMTL/internal/service/cloud115"
)
// testStrmService 构建带内存库的 StrmService。
@@ -35,7 +39,7 @@ func testStrmService(t *testing.T) *StrmService {
t.Cleanup(func() { _ = sqlDB.Close() })
}
if err := db.AutoMigrate(&model.StrmAccount{}, &model.StrmSyncPath{}, &model.StrmSyncRecord{},
&model.StrmDownloadTask{}, &model.StrmUploadTask{}, &model.Setting{}); err != nil {
&model.StrmDownloadTask{}, &model.StrmUploadTask{}, &model.StrmDirCache{}, &model.Setting{}); err != nil {
t.Fatal(err)
}
repos := repository.New(db)
@@ -159,6 +163,56 @@ func TestLocalStrmSync(t *testing.T) {
}
}
// TestStrmFullAndIncrementalSync 测试增量同步与全量同步模式切换及记录
func TestStrmFullAndIncrementalSync(t *testing.T) {
svc := testStrmService(t)
src := t.TempDir()
out := t.TempDir()
writeFile(t, filepath.Join(src, "电影", "星际穿越.mkv"), "fake-video-data")
p := syncPathRecord(t, svc, model.StrmProviderLocal, src, out, true)
// 1. 默认触发增量同步
if err := svc.StartSync(context.Background(), p.ID, model.StrmSyncTypeIncremental); err != nil {
t.Fatal(err)
}
record := waitSyncDone(t, svc, p.ID, 10*time.Second)
if record.Status != model.StrmSyncRecordDone {
t.Fatalf("sync status = %s, message = %s", record.Status, record.Message)
}
if record.SyncType != model.StrmSyncTypeIncremental {
t.Fatalf("expected sync_type = incremental, got %s", record.SyncType)
}
if record.NewStrm != 1 {
t.Fatalf("expected 1 new strm, got %d", record.NewStrm)
}
// 2. 再次执行增量同步,应当跳过
if err := svc.StartSync(context.Background(), p.ID, model.StrmSyncTypeIncremental); err != nil {
t.Fatal(err)
}
record = waitSyncDone(t, svc, p.ID, 10*time.Second)
if record.SyncType != model.StrmSyncTypeIncremental {
t.Fatalf("expected sync_type = incremental, got %s", record.SyncType)
}
if record.Skipped != 1 {
t.Fatalf("expected 1 skipped, got %d", record.Skipped)
}
// 3. 执行全量同步
if err := svc.StartSync(context.Background(), p.ID, model.StrmSyncTypeFull); err != nil {
t.Fatal(err)
}
record = waitSyncDone(t, svc, p.ID, 10*time.Second)
if record.SyncType != model.StrmSyncTypeFull {
t.Fatalf("expected sync_type = full, got %s", record.SyncType)
}
if record.Status != model.StrmSyncRecordDone {
t.Fatalf("full sync failed: status = %s, message = %s", record.Status, record.Message)
}
}
// TestStrmCronMatches cron 表达式匹配。
func TestStrmCronMatches(t *testing.T) {
cases := []struct {
@@ -446,3 +500,211 @@ func TestWalkRemoteConcurrent(t *testing.T) {
t.Errorf("生成的 .strm 数量 = %d,期望 5", strmCount)
}
}
// TestStrmBatchEnqueueAndConcurrentClaim 测试大规模批量入库及多协程并发认领无死锁
func TestStrmBatchEnqueueAndConcurrentClaim(t *testing.T) {
svc := testStrmService(t)
ctx := context.Background()
// 1. 批量插入 200 个下载任务
tasks := make([]*model.StrmDownloadTask, 0, 200)
for i := 0; i < 200; i++ {
tasks = append(tasks, &model.StrmDownloadTask{
SyncPathID: "test-sync-path",
AccountID: "test-acct",
Provider: model.StrmProvider115,
FileName: filepath.Base(string(rune('a'+i%26))) + ".nfo",
LocalPath: filepath.Join(t.TempDir(), string(rune('a'+i%26)), "test.nfo"),
Status: model.StrmTaskPending,
})
}
if err := svc.repo.StrmDownload.CreateInBatches(ctx, tasks, 50); err != nil {
t.Fatalf("CreateInBatches failed: %v", err)
}
// 2. 验证 ActiveLocalPathMap
activeMap, err := svc.repo.StrmDownload.GetActiveLocalPathMap(ctx, "test-sync-path")
if err != nil {
t.Fatalf("GetActiveLocalPathMap failed: %v", err)
}
if len(activeMap) == 0 {
t.Fatal("expected active local path map to have entries")
}
// 3. 模拟 6 个 worker 并发 ClaimPendingDownload
claimedCount := 0
var claimMu sync.Mutex
var wg sync.WaitGroup
for w := 0; w < 6; w++ {
wg.Add(1)
go func() {
defer wg.Done()
for {
batch, err := svc.repo.StrmDownload.ClaimPendingDownload(ctx, 10)
if err != nil {
t.Errorf("concurrent ClaimPendingDownload failed: %v", err)
return
}
if len(batch) == 0 {
return
}
claimMu.Lock()
claimedCount += len(batch)
claimMu.Unlock()
}
}()
}
wg.Wait()
if claimedCount != 200 {
t.Fatalf("expected all 200 tasks claimed, got %d", claimedCount)
}
}
// TestStrmDuplicateFileConflictResolution 测试远端存在多个同名不同大小文件时,本地确定性仲裁,避免增量死循环
func TestStrmDuplicateFileConflictResolution(t *testing.T) {
svc := testStrmService(t)
localDir := t.TempDir()
p := &model.StrmSyncPath{
Base: model.Base{ID: "dup-test-path"},
Provider: model.StrmProvider115,
RemotePath: "root",
LocalPath: localDir,
DownloadMeta: true,
}
st := &strmSyncState{
s: svc,
ctx: context.Background(),
p: p,
cfg: &strmPathConfig{DownloadMeta: true, MetaExt: []string{"nfo"}},
rec: &model.StrmSyncRecord{},
seenMeta: map[string]bool{},
remoteMeta: map[string]int64{},
seenMetaTarget: map[string]cloud.FileEntry{},
seenVideoTarget: map[string]cloud.FileEntry{},
}
// 模拟远端同目录下存在两个同名不同大小的 nfo 文件 (115 历史重复上传)
// entry1: 较早文件 (MTime: 1000, Size: 100)
entry1 := cloud.FileEntry{ID: "f1", Name: "test.nfo", Size: 100, MTime: 1000, PickCode: "p1"}
// entry2: 较新文件 (MTime: 2000, Size: 200)
entry2 := cloud.FileEntry{ID: "f2", Name: "test.nfo", Size: 200, MTime: 2000, PickCode: "p2"}
// 第一次全量处理:两者都在列表中
st.handleMeta(entry1, "test.nfo", ".nfo")
st.handleMeta(entry2, "test.nfo", ".nfo")
st.flushPendingDownloads()
// 验证仲裁结果:最终只产生 1 个下载任务,且使用的是首个匹配项 (Size 100/p1)
tasks, _, err := svc.repo.StrmDownload.List(context.Background(), "", 1, 10)
if err != nil {
t.Fatal(err)
}
if len(tasks) != 1 {
t.Fatalf("expected 1 download task after conflict resolution, got %d", len(tasks))
}
if tasks[0].Size != 100 || tasks[0].RemoteRef != "p1" {
t.Fatalf("expected task with size 100/p1, got size=%d ref=%s", tasks[0].Size, tasks[0].RemoteRef)
}
// 模拟该任务下载落盘完成
writeFile(t, filepath.Join(localDir, "test.nfo"), strings.Repeat("x", 100))
// 第二次增量同步:两者再次依次扫描
st2 := &strmSyncState{
s: svc,
ctx: context.Background(),
p: p,
cfg: &strmPathConfig{DownloadMeta: true, MetaExt: []string{"nfo"}},
rec: &model.StrmSyncRecord{},
seenMeta: map[string]bool{},
remoteMeta: map[string]int64{},
seenMetaTarget: map[string]cloud.FileEntry{},
seenVideoTarget: map[string]cloud.FileEntry{},
}
st2.handleMeta(entry1, "test.nfo", ".nfo")
st2.handleMeta(entry2, "test.nfo", ".nfo")
st2.flushPendingDownloads()
// 验证:不会新增任何下载任务,NewMeta 为 0,增量跳过
if st2.rec.NewMeta != 0 {
t.Fatalf("expected 0 new meta on incremental sync, got %d", st2.rec.NewMeta)
}
}
// TestWalk115FlatAbortsOnDirResolveFailure 回归测试:115 开放平台 token 失效/目录详情
// 解析失败时,同步必须中止而不是带着塌缩的 rel 继续处理,否则会导致本地大量元数据
// 被误判为"云端不存在"而重复下载/上传,甚至误删本地文件(用户反馈"云盘没动却重下重传")。
func TestWalk115FlatAbortsOnDirResolveFailure(t *testing.T) {
svc := testStrmService(t)
localDir := t.TempDir()
acct := &model.StrmAccount{
Name: "fake115",
Provider: "cloud115",
Config: "{}",
Enabled: true,
}
if err := svc.repo.StrmAccount.Create(context.Background(), acct); err != nil {
t.Fatal(err)
}
p := &model.StrmSyncPath{
Base: model.Base{ID: "abort-path"},
AccountID: acct.ID,
Provider: model.StrmProvider115,
RemotePath: "0",
LocalPath: localDir,
}
// 115 mock:文件列表返回一个视频(父目录 999 不在缓存,需要 get_info),
// get_info 恒返回 access_token 格式错误(40140123)→ 目录树解析失败。
api := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/open/ufile/files":
w.Write([]byte(`{"state":true,"count":1,"data":[{"fid":"100","pid":"999","fc":1,"fn":"movie.mkv","pc":"pc1","upt":1700000000,"fs":1024}]}`))
case "/open/folder/get_info":
w.Write([]byte(`{"state":false,"code":40140123,"message":"access_token 格式错误"}`))
default:
t.Errorf("unexpected path %s", r.URL.Path)
}
}))
defer api.Close()
oldPro := cloud115.ProAPIBase
cloud115.ProAPIBase = api.URL
defer func() { cloud115.ProAPIBase = oldPro }()
oc := cloud115.NewOpenClient("app", "at", "rt")
st := &strmSyncState{
s: svc,
ctx: context.Background(),
p: p,
provider: cloud.NewOpenAPI115("app", "at", "rt"),
cfg: &strmPathConfig{VideoExt: []string{"mkv"}, MetaExt: []string{"nfo"}, AddPath: 1, DownloadMeta: false},
rec: &model.StrmSyncRecord{},
syncType: model.StrmSyncTypeFull,
dirCache: sync.Map{},
seenVideo: map[string]bool{},
seenMeta: map[string]bool{},
remoteMeta: map[string]int64{},
}
err := st.walk115Flat(oc)
if err == nil {
t.Fatal("expected walk115Flat to abort on dir-resolve failure, got nil error")
}
// 中止后不允许产生任何部分写入(本地不允许生成 .strm 文件)。
var strmCount int
_ = filepath.WalkDir(localDir, func(path string, d os.DirEntry, err error) error {
if err == nil && !d.IsDir() && strings.HasSuffix(d.Name(), ".strm") {
strmCount++
}
return nil
})
if strmCount != 0 {
t.Fatalf("expected no .strm written after abort, got %d", strmCount)
}
}
+16 -4
View File
@@ -57,7 +57,7 @@ func (t *TMDbProvider) searchMovieCandidates(ctx context.Context, query string,
apiKey := t.resolveAPIKey(ctx)
if apiKey == "" {
return nil, nil
return nil, errors.New("TMDb API Key 未配置,请先在「系统设置 → API配置」中填写")
}
base := t.resolveBaseURL(ctx)
@@ -65,7 +65,7 @@ func (t *TMDbProvider) searchMovieCandidates(ctx context.Context, query string,
q.Set("api_key", apiKey)
q.Set("query", query)
q.Set("language", language)
q.Set("include_adult", "false")
q.Set("include_adult", "true")
if year > 0 {
q.Set("year", fmt.Sprintf("%d", year))
}
@@ -79,6 +79,12 @@ func (t *TMDbProvider) searchMovieCandidates(ctx context.Context, query string,
if err := t.getJSON(ctx, u, &p); err != nil {
return nil, err
}
if len(p.Results) == 0 && year > 0 {
// 移除年份限制重试一次,避免年份微小差异(如 2019 vs 2020)导致无搜索结果
q.Del("year")
u = base + "/search/movie?" + q.Encode()
_ = t.getJSON(ctx, u, &p)
}
if len(p.Results) == 0 {
return nil, nil
}
@@ -137,7 +143,7 @@ func (t *TMDbProvider) searchTVCandidates(ctx context.Context, query string, yea
apiKey := t.resolveAPIKey(ctx)
if apiKey == "" {
return nil, nil
return nil, errors.New("TMDb API Key 未配置,请先在「系统设置 → API配置」中填写")
}
base := t.resolveBaseURL(ctx)
@@ -145,7 +151,7 @@ func (t *TMDbProvider) searchTVCandidates(ctx context.Context, query string, yea
q.Set("api_key", apiKey)
q.Set("query", query)
q.Set("language", language)
q.Set("include_adult", "false")
q.Set("include_adult", "true")
if year > 0 {
q.Set("first_air_date_year", fmt.Sprintf("%d", year))
}
@@ -159,6 +165,12 @@ func (t *TMDbProvider) searchTVCandidates(ctx context.Context, query string, yea
if err := t.getJSON(ctx, u, &p); err != nil {
return nil, err
}
if len(p.Results) == 0 && year > 0 {
// 移除年份限制重试一次,避免年份微小差异导致无搜索结果
q.Del("first_air_date_year")
u = base + "/search/tv?" + q.Encode()
_ = t.getJSON(ctx, u, &p)
}
if len(p.Results) == 0 {
return nil, nil
}
+13 -5
View File
@@ -123,12 +123,20 @@ func LibraryVisibleForUser(ctx context.Context, repo *repository.Container, lib
return false
}
if repo != nil && repo.DB != nil {
var count int64
var totalCount int64
_ = repo.DB.WithContext(ctx).Model(&model.Media{}).
Where("library_id = ? AND nsfw = ?", lib.ID, true).
Count(&count).Error
if count > 0 {
return false
Where("library_id = ?", lib.ID).
Count(&totalCount).Error
if totalCount > 0 {
var nsfwCount int64
_ = repo.DB.WithContext(ctx).Model(&model.Media{}).
Where("library_id = ? AND nsfw = ?", lib.ID, true).
Count(&nsfwCount).Error
// 仅当整库媒体全部为成人内容(纯成人库)时才隐藏整库;
// 含有普通内容的混合媒体库保持库本身可见,具体 NSFW 条目在媒体列表内过滤。
if nsfwCount == totalCount {
return false
}
}
}
return true
+23
View File
@@ -0,0 +1,23 @@
// Package webui embeds the built React SPA so a single binary can serve the
// UI without requiring a separate web/dist directory on disk. The embed
// happens at compile time, so `web/dist` must exist before `go build` runs
// (the CI pipeline builds it via `npm run build` first).
package webui
import (
"embed"
"io/fs"
)
//go:embed all:dist
var distFS embed.FS
// DistFS returns the embedded SPA build artifacts rooted at the directory
// containing index.html (i.e. web/dist).
func DistFS() fs.FS {
sub, err := fs.Sub(distFS, "dist")
if err != nil {
panic(err)
}
return sub
}
+69 -18
View File
@@ -1,6 +1,45 @@
import { api } from './client'
import type { AccessLog, Setting, User } from '../types'
export interface DatabaseStatus {
type: 'sqlite' | 'postgres'
dsn?: string
db_path?: string
open_conns: number
in_use: number
idle: number
max_open_conns: number
table_counts?: Record<string, number>
}
export interface PostgresTestResult {
success: boolean
latency_ms?: number
version?: string
message?: string
error?: string
}
export interface DatabaseMigrationResult {
success: boolean
total_rows: number
table_rows?: Record<string, number>
duration_ms: number
message?: string
error?: string
}
export interface DatabaseConnectionPayload {
type?: string
dsn?: string
host?: string
port?: number
user?: string
password?: string
dbname?: string
sslmode?: string
}
export interface SystemUpdateStatus {
image: string
current_version?: string
@@ -55,21 +94,33 @@ export const adminAPI = {
systemUpdateApply: () => api.post<SystemUpdateStatus>('/admin/system/update/apply').then((r) => r.data),
testAdultScraper: (payload: {
engine?: string
server_url?: string
token?: string
javdb_url?: string
javbus_url?: string
cookie?: string
}) =>
api
.post<{
success: boolean
latency_ms?: number
providers?: string[]
message?: string
error?: string
}>('/admin/adult/test-scraper', payload)
.then((r) => r.data),
}
testAdultScraper: (payload: {
engine?: string
server_url?: string
token?: string
javdb_url?: string
javbus_url?: string
cookie?: string
}) =>
api
.post<{
success: boolean
latency_ms?: number
providers?: string[]
message?: string
error?: string
}>('/admin/adult/test-scraper', payload)
.then((r) => r.data),
getDatabaseStatus: () =>
api.get<DatabaseStatus>('/admin/database/status').then((r) => r.data),
testDatabaseConnection: (payload: DatabaseConnectionPayload) =>
api.post<PostgresTestResult>('/admin/database/test', payload).then((r) => r.data),
migrateDatabase: (payload: DatabaseConnectionPayload) =>
api.post<DatabaseMigrationResult>('/admin/database/migrate', payload).then((r) => r.data),
saveDatabaseConfig: (payload: DatabaseConnectionPayload) =>
api.post<{ message: string; type: string }>('/admin/database/save-config', payload).then((r) => r.data),
}
+13 -2
View File
@@ -102,8 +102,19 @@ export const libraryAPI = {
createWithRoots: (name: string, type: string, roots: LibraryRootInput[], coverURL = '') =>
api.post<Library>('/libraries', { name, type, roots, cover_url: coverURL }).then((r) => r.data),
update: (id: string, payload: { cover_url: string }) =>
api.patch<Library>(`/libraries/${id}`, payload).then((r) => r.data),
createPerSubfolder: (parentPath: string, type: string, coverURL = '') =>
api.post<{ libraries: Library[] }>('/libraries', { path: parentPath, type, cover_url: coverURL, create_per_subfolder: true }).then((r) => r.data),
update: (
id: string,
payload: {
cover_url?: string
sort_order?: number | null
carousel_enabled?: boolean | null
},
) => api.patch<Library>(`/libraries/${id}`, payload).then((r) => r.data),
reorder: (ids: string[]) => api.put('/libraries/reorder', { ids }).then((r) => r.data),
remove: (id: string) => api.delete(`/libraries/${id}`).then((r) => r.data),
+57
View File
@@ -0,0 +1,57 @@
import { api } from './client'
import type { ScrapeQueueSnapshot } from '../types/scraper'
export interface EnqueueScrapeOptions {
episode_images?: boolean
episode_artwork?: boolean
refresh_matched?: boolean
include_matched?: boolean
}
export const scraperAPI = {
queue: (status?: string, page = 1, pageSize = 50) =>
api
.get<ScrapeQueueSnapshot>('/admin/scraper/queue', {
params: { status, page, page_size: pageSize },
})
.then((r) => r.data),
cancelTask: (id: string) =>
api.post(`/admin/scraper/queue/${id}/cancel`).then((r) => r.data),
retryTask: (id: string) =>
api.post(`/admin/scraper/queue/${id}/retry`).then((r) => r.data),
deleteTask: (id: string) =>
api.delete(`/admin/scraper/queue/${id}`).then((r) => r.data),
batchAction: (action: 'delete' | 'retry' | 'cancel', ids: string[]) =>
api
.post<{ affected: number; action: string }>('/admin/scraper/queue/batch', { action, ids })
.then((r) => r.data),
clearDone: () =>
api.post<{ deleted: number }>('/admin/scraper/queue/clear-done').then((r) => r.data),
clearFinished: () =>
api.post<{ deleted: number }>('/admin/scraper/queue/clear-finished').then((r) => r.data),
clearCanceled: () =>
api.post<{ deleted: number }>('/admin/scraper/queue/clear-canceled').then((r) => r.data),
retryFailed: () =>
api.post<{ retried: number }>('/admin/scraper/queue/retry-failed').then((r) => r.data),
cancelPending: () =>
api.post<{ canceled: number }>('/admin/scraper/queue/cancel-pending').then((r) => r.data),
enqueueLibrary: (libraryId: string, options?: EnqueueScrapeOptions) =>
api
.post<{ enqueued: number }>(`/admin/scraper/queue/enqueue-library/${libraryId}`, options ?? {})
.then((r) => r.data),
enqueueAll: (options?: EnqueueScrapeOptions) =>
api
.post<{ enqueued: number }>('/admin/scraper/queue/enqueue-all', options ?? {})
.then((r) => r.data),
}
+30 -1
View File
@@ -107,7 +107,8 @@ export const strmAPI = {
deletePath: (id: string) => api.delete(`/admin/strm/paths/${id}`).then((r) => r.data),
startSync: (id: string) => api.post(`/admin/strm/paths/${id}/sync`).then((r) => r.data),
startSync: (id: string, mode: 'incremental' | 'full' = 'incremental') =>
api.post(`/admin/strm/paths/${id}/sync`, null, { params: { mode } }).then((r) => r.data),
cancelSync: (id: string) => api.post(`/admin/strm/paths/${id}/cancel`).then((r) => r.data),
@@ -116,6 +117,13 @@ export const strmAPI = {
.get<StrmSyncRecord[]>('/admin/strm/records', { params: pathId ? { path_id: pathId } : {} })
.then((r) => r.data),
deleteRecord: (id: string) => api.delete(`/admin/strm/records/${id}`).then((r) => r.data),
clearRecords: (pathId?: string) =>
api
.delete<{ deleted: number }>('/admin/strm/records', { params: pathId ? { path_id: pathId } : {} })
.then((r) => r.data),
// ── 本地目录浏览(同步目录选择器) ────────────────────────
listLocalDirs: (path?: string) =>
api
@@ -136,12 +144,21 @@ export const strmAPI = {
retryDownload: (id: string) =>
api.post(`/admin/strm/downloads/${id}/retry`).then((r) => r.data),
deleteDownload: (id: string) =>
api.delete(`/admin/strm/downloads/${id}`).then((r) => r.data),
batchActionDownloads: (action: 'delete' | 'retry' | 'cancel', ids: string[]) =>
api.post<{ affected: number; action: string }>('/admin/strm/downloads/batch', { action, ids }).then((r) => r.data),
clearDoneDownloads: () =>
api.post<{ deleted: number }>('/admin/strm/downloads/clear-done').then((r) => r.data),
clearFinishedDownloads: () =>
api.post<{ deleted: number }>('/admin/strm/downloads/clear-finished').then((r) => r.data),
clearCanceledDownloads: () =>
api.post<{ deleted: number }>('/admin/strm/downloads/clear-canceled').then((r) => r.data),
retryFailedDownloads: () =>
api.post<{ retried: number }>('/admin/strm/downloads/retry-failed').then((r) => r.data),
@@ -160,4 +177,16 @@ export const strmAPI = {
retryUpload: (id: string) =>
api.post(`/admin/strm/uploads/${id}/retry`).then((r) => r.data),
deleteUpload: (id: string) =>
api.delete(`/admin/strm/uploads/${id}`).then((r) => r.data),
batchActionUploads: (action: 'delete' | 'retry' | 'cancel', ids: string[]) =>
api.post<{ affected: number; action: string }>('/admin/strm/uploads/batch', { action, ids }).then((r) => r.data),
cancelPendingUploads: () =>
api.post<{ canceled: number }>('/admin/strm/uploads/cancel-pending').then((r) => r.data),
clearCanceledUploads: () =>
api.post<{ deleted: number }>('/admin/strm/uploads/clear-canceled').then((r) => r.data),
}
+4
View File
@@ -33,6 +33,9 @@ const StrmDownloadQueuePage = lazy(() =>
const StrmUploadQueuePage = lazy(() =>
import('./pages/StrmQueuePage').then((m) => ({ default: m.StrmUploadQueuePage })),
)
const ScraperQueuePage = lazy(() =>
import('./pages/ScraperQueuePage').then((m) => ({ default: m.ScraperQueuePage })),
)
export type AppRoute = {
path?: string
@@ -62,5 +65,6 @@ export const appRoutes: AppRoute[] = [
{ path: 'strm', element: <StrmManagePage />, adminOnly: true },
{ path: 'strm/downloads', element: <StrmDownloadQueuePage />, adminOnly: true },
{ path: 'strm/uploads', element: <StrmUploadQueuePage />, adminOnly: true },
{ path: 'scraper/queue', element: <ScraperQueuePage />, adminOnly: true },
{ path: 'admin', element: <AdminPage />, adminOnly: true },
]
+2
View File
@@ -5,6 +5,7 @@ import {
FolderOpen,
Library,
Settings,
Sparkles,
Upload,
User,
Users,
@@ -22,6 +23,7 @@ export type LayoutNavItem = {
export const LAYOUT_NAV_ITEMS: LayoutNavItem[] = [
{ to: '/profile', label: '个人资料', icon: User },
{ to: '/libraries?from=admin', label: '媒体库', icon: Library },
{ to: '/scraper/queue', label: '刮削队列', icon: Sparkles, adminOnly: true },
{ to: '/admin', label: '用户管理', icon: Users, adminOnly: true },
{ to: '/files', label: '文件管理', icon: FolderOpen, adminOnly: true },
{ to: '/settings', label: '系统设置', icon: Settings, adminOnly: true },
+4
View File
@@ -13,9 +13,11 @@ export function AdminLibraryPanel() {
type={createForm.type}
coverURL={createForm.coverURL}
roots={createForm.roots}
createPerSubfolder={createForm.createPerSubfolder}
onNameChange={createForm.setName}
onTypeChange={createForm.setType}
onCoverURLChange={createForm.setCoverURL}
onCreatePerSubfolderChange={createForm.setCreatePerSubfolder}
onRootChange={createForm.updateRoot}
onAddRoot={createForm.addRoot}
onRemoveRoot={createForm.removeRoot}
@@ -33,6 +35,8 @@ export function AdminLibraryPanel() {
onRemoveLibrary={libraryActions.removeLibrary}
onAddLibraryRoot={libraryActions.addLibraryRoot}
onEditLibraryCover={libraryActions.editLibraryCover}
onToggleCarousel={libraryActions.toggleCarouselLibrary}
onReorder={libraryActions.reorderLibraries}
/>
<LibraryStorageStats />
</div>
+26 -6
View File
@@ -9,9 +9,11 @@ type CreateFormProps = {
type: string
coverURL: string
roots: RootDraft[]
createPerSubfolder: boolean
onNameChange: (value: string) => void
onTypeChange: (value: string) => void
onCoverURLChange: (value: string) => void
onCreatePerSubfolderChange: (value: boolean) => void
onRootChange: (index: number, patch: Partial<RootDraft>) => void
onAddRoot: () => void
onRemoveRoot: (index: number) => void
@@ -23,9 +25,11 @@ export function AdminLibraryCreateForm({
type,
coverURL,
roots,
createPerSubfolder,
onNameChange,
onTypeChange,
onCoverURLChange,
onCreatePerSubfolderChange,
onRootChange,
onAddRoot,
onRemoveRoot,
@@ -50,9 +54,9 @@ export function AdminLibraryCreateForm({
<>
<form onSubmit={onSubmit} className="glass-panel grid gap-3 md:grid-cols-4">
<input
required
required={!createPerSubfolder}
className="input-base"
placeholder="名称"
placeholder={createPerSubfolder ? '父级媒体库名(批量模式忽略)' : '名称'}
value={name}
onChange={(e) => onNameChange(e.target.value)}
/>
@@ -81,15 +85,31 @@ export function AdminLibraryCreateForm({
onRemove={onRemoveRoot}
/>
))}
<button type="button" className="inline-flex items-center gap-2 rounded-lg border px-3 py-2 text-sm" onClick={onAddRoot}>
<Plus size={16} /> 添加路径
</button>
{!createPerSubfolder && (
<button type="button" className="inline-flex items-center gap-2 rounded-lg border px-3 py-2 text-sm" onClick={onAddRoot}>
<Plus size={16} /> 添加路径
</button>
)}
</div>
<p className="md:col-span-4 -mt-2 text-xs text-sand-500">
支持直接点选或手动输入;名称和类型与现有媒体库一致时,会自动把这里填写的路径追加到该媒体库。
</p>
<label className="md:col-span-4 flex items-center gap-2 text-sm text-ink-100">
<input
type="checkbox"
className="h-4 w-4 accent-brand-400"
checked={createPerSubfolder}
onChange={(e) => onCreatePerSubfolderChange(e.target.checked)}
/>
<span>按目录下每个子文件夹各建一个媒体库(媒体库名取子文件夹名)</span>
</label>
{createPerSubfolder && (
<p className="md:col-span-4 -mt-2 text-xs text-sand-500">
批处理模式:仅取上方第一个路径作为父级目录,会为其中每个子文件夹分别创建媒体库,可自选类型用于整体推断。
</p>
)}
<button type="submit" className="neon-button md:col-span-4">
新建 / 追加路径
{createPerSubfolder ? '按目录批量创建' : '新建 / 追加路径'}
</button>
</form>
+188 -20
View File
@@ -1,5 +1,6 @@
import { useState, type MouseEvent, type ReactNode } from 'react'
import { Folder, Image, MoreVertical, Plus, Power, PowerOff, RefreshCw, Save, Trash2 } from 'lucide-react'
import { useEffect, useRef, useState, type DragEvent, type MouseEvent, type ReactNode } from 'react'
import { createPortal } from 'react-dom'
import { Folder, GripVertical, Image, MoreVertical, Plus, Power, PowerOff, RefreshCw, Save, Trash2 } from 'lucide-react'
import { LocalDirBrowserDialog } from '../components/LocalDirBrowserDialog'
import type { Library, LibraryRoot } from '../types'
@@ -18,11 +19,15 @@ type LibraryTableProps = {
onRemoveLibrary: (library: Library) => void
onAddLibraryRoot: (library: Library, path?: string, name?: string) => void
onEditLibraryCover: (library: Library) => void
onToggleCarousel: (library: Library) => void
onReorder: (orderedLibs: Library[]) => void
}
export function AdminLibraryTable({ libs, ...actions }: LibraryTableProps) {
const [browsingRoot, setBrowsingRoot] = useState<{ libraryID: string; root: LibraryRoot; initialPath?: string } | null>(null)
const [addingRootLib, setAddingRootLib] = useState<Library | null>(null)
const [draggingId, setDraggingId] = useState<string | null>(null)
const dragOverId = useRef<string | null>(null)
const handleSelectRootPath = (selectedPath: string) => {
if (browsingRoot) {
@@ -40,15 +45,39 @@ export function AdminLibraryTable({ libs, ...actions }: LibraryTableProps) {
}
}
const handleReorder = (fromId: string, overId: string) => {
if (fromId === overId) return
const copy = [...libs]
const fromIndex = copy.findIndex((l) => l.id === fromId)
const overIndex = copy.findIndex((l) => l.id === overId)
if (fromIndex < 0 || overIndex < 0) return
const [moved] = copy.splice(fromIndex, 1)
copy.splice(overIndex, 0, moved)
actions.onReorder(copy)
}
const handleDrop = (e: DragEvent, overId: string) => {
e.preventDefault()
dragOverId.current = null
setDraggingId(null)
if (draggingId && overId !== draggingId) {
handleReorder(draggingId, overId)
}
}
return (
<>
<div className="glass-panel overflow-x-auto !p-3">
<table className="w-full min-w-[900px] text-left text-sm">
<table className="w-full min-w-[960px] text-left text-sm">
<thead className="text-xs uppercase tracking-wider text-sand-500">
<tr>
<th className="w-8 text-center" title="拖动排序">
<GripVertical size={13} className="mx-auto text-gray-300" />
</th>
<th className="w-28 py-2">名称</th>
<th>路径</th>
<th className="w-20">类型</th>
<th className="w-24">轮播</th>
<th className="w-12 text-right">操作</th>
</tr>
</thead>
@@ -59,6 +88,23 @@ export function AdminLibraryTable({ libs, ...actions }: LibraryTableProps) {
library={library}
onBrowseRoot={(root) => setBrowsingRoot({ libraryID: library.id, root, initialPath: root.path })}
onOpenAddRoot={() => setAddingRootLib(library)}
dragging={draggingId === library.id}
dragOver={dragOverId.current === library.id}
onDragStart={(e) => {
e.dataTransfer.effectAllowed = 'move'
dragOverId.current = null
setDraggingId(library.id)
}}
onDragOver={(e) => {
e.preventDefault()
e.dataTransfer.dropEffect = 'move'
if (dragOverId.current !== library.id) dragOverId.current = library.id
}}
onDragEnd={() => {
dragOverId.current = null
setDraggingId(null)
}}
onDrop={(e) => handleDrop(e, library.id)}
{...actions}
/>
))}
@@ -90,11 +136,37 @@ type LibraryTableRowProps = Omit<LibraryTableProps, 'libs'> & {
library: Library
onBrowseRoot: (root: LibraryRoot) => void
onOpenAddRoot: () => void
dragging?: boolean
dragOver?: boolean
onDragStart?: (e: DragEvent) => void
onDragOver?: (e: DragEvent) => void
onDragEnd?: () => void
onDrop?: (e: DragEvent) => void
}
function LibraryTableRow({ library, ...actions }: LibraryTableRowProps) {
function LibraryTableRow({ library, dragging, dragOver, onDragStart, onDragOver, onDragEnd, onDrop, ...actions }: LibraryTableRowProps) {
return (
<tr className="border-t border-gray-200">
<tr
draggable={false}
onDragStart={onDragStart}
onDragOver={onDragOver}
onDragEnd={onDragEnd}
onDrop={onDrop}
className={`border-t border-gray-200 transition-colors ${
dragging ? 'bg-primary-400/10 opacity-60' : dragOver ? 'bg-primary-400/5' : ''
}`}
>
<td className="py-2 text-center">
<button
type="button"
draggable
onDragStart={onDragStart}
className="inline-flex cursor-grab items-center justify-center rounded p-1 text-gray-400 transition hover:bg-gray-100 hover:text-brand-500 active:cursor-grabbing"
title="拖动以调整媒体库显示顺序"
>
<GripVertical size={16} />
</button>
</td>
<td className="py-2 pr-3 font-medium text-ink-600">
<div className="flex items-center gap-2">
{library.cover_url && <img src={library.cover_url} alt="" className="h-10 w-8 rounded object-cover" />}
@@ -105,6 +177,9 @@ function LibraryTableRow({ library, ...actions }: LibraryTableRowProps) {
<LibraryRootsCell library={library} {...actions} />
</td>
<td className="px-3 text-ink-100">{library.type}</td>
<td className="py-2 text-ink-100">
<CarouselToggle library={library} onToggleCarousel={actions.onToggleCarousel} />
</td>
<td className="py-2 text-right">
<LibraryActionsCell library={library} {...actions} />
</td>
@@ -112,6 +187,25 @@ function LibraryTableRow({ library, ...actions }: LibraryTableRowProps) {
)
}
function CarouselToggle({ library, onToggleCarousel }: { library: Library; onToggleCarousel: (library: Library) => void }) {
const on = Boolean(library.carousel_enabled)
return (
<button
type="button"
onClick={() => onToggleCarousel(library)}
className={`inline-flex items-center gap-1.5 rounded-lg border px-2.5 py-1 text-xs font-semibold transition ${
on
? 'border-brand-500/50 bg-brand-500/10 text-brand-500'
: 'border-gray-300 bg-white text-ink-50 hover:border-gray-400'
}`}
title={on ? '已开启首页海报轮播(点击关闭)' : '未开启首页海报轮播(点击开启)'}
>
<span className={`h-3.5 w-3.5 rounded-full ${on ? 'bg-brand-500' : 'bg-gray-300'}`} />
{on ? '参与轮播' : '未参与'}
</button>
)
}
function LibraryRootsCell({ library, ...actions }: LibraryTableRowProps) {
const roots = library.roots?.length ? library.roots : [fallbackLibraryRoot(library)]
return (
@@ -255,18 +349,90 @@ function LibraryActionsCell({ library, onScanLibrary, onRemoveLibrary, onOpenAdd
}
function ActionMenu({ label, children }: { label: string; children: ReactNode }) {
const [isOpen, setIsOpen] = useState(false)
const [coords, setCoords] = useState<{ top?: number; bottom?: number; right: number } | null>(null)
const triggerRef = useRef<HTMLButtonElement>(null)
const menuRef = useRef<HTMLDivElement>(null)
const toggleMenu = (e: MouseEvent) => {
e.stopPropagation()
if (isOpen) {
setIsOpen(false)
return
}
if (triggerRef.current) {
const rect = triggerRef.current.getBoundingClientRect()
const spaceBelow = window.innerHeight - rect.bottom
const estimatedHeight = 180
const openUpward = spaceBelow < estimatedHeight && rect.top > estimatedHeight
setCoords({
top: openUpward ? undefined : rect.bottom + 4,
bottom: openUpward ? window.innerHeight - rect.top + 4 : undefined,
right: window.innerWidth - rect.right,
})
setIsOpen(true)
}
}
useEffect(() => {
if (!isOpen) return
const handleClickOutside = (e: globalThis.MouseEvent) => {
if (
menuRef.current &&
!menuRef.current.contains(e.target as Node) &&
triggerRef.current &&
!triggerRef.current.contains(e.target as Node)
) {
setIsOpen(false)
}
}
const handleScrollOrResize = () => setIsOpen(false)
window.addEventListener('mousedown', handleClickOutside)
window.addEventListener('scroll', handleScrollOrResize, true)
window.addEventListener('resize', handleScrollOrResize)
return () => {
window.removeEventListener('mousedown', handleClickOutside)
window.removeEventListener('scroll', handleScrollOrResize, true)
window.removeEventListener('resize', handleScrollOrResize)
}
}, [isOpen])
return (
<details className="group relative inline-flex justify-end">
<summary
className="flex h-8 w-8 cursor-pointer list-none items-center justify-center rounded-lg border border-gray-200 bg-white text-ink-50 transition hover:border-primary-400/50 hover:text-brand-500 [&::-webkit-details-marker]:hidden"
<>
<button
ref={triggerRef}
type="button"
onClick={toggleMenu}
className={`inline-flex h-8 w-8 cursor-pointer items-center justify-center rounded-lg border transition ${
isOpen
? 'border-brand-500 bg-brand-500/10 text-brand-500'
: 'border-gray-200 bg-white text-ink-50 hover:border-primary-400/50 hover:text-brand-500'
}`}
title={label}
>
<MoreVertical size={16} />
</summary>
<div className="absolute right-0 top-9 z-30 min-w-28 rounded-lg border border-gray-200 bg-white p-1 shadow-lg">
{children}
</div>
</details>
</button>
{isOpen &&
coords &&
createPortal(
<div
ref={menuRef}
style={{
position: 'fixed',
top: coords.top !== undefined ? `${coords.top}px` : undefined,
bottom: coords.bottom !== undefined ? `${coords.bottom}px` : undefined,
right: `${coords.right}px`,
zIndex: 99999,
}}
className="min-w-32 rounded-xl border border-gray-200/90 bg-white p-1.5 shadow-2xl backdrop-blur"
onClick={() => setIsOpen(false)}
>
{children}
</div>,
document.body,
)}
</>
)
}
@@ -283,17 +449,19 @@ function MenuButton({
onClick: () => void
children: ReactNode
}) {
const handleClick = (event: MouseEvent<HTMLButtonElement>) => {
event.currentTarget.closest('details')?.removeAttribute('open')
onClick()
}
return (
<button
className={`flex w-full items-center gap-2 rounded-md px-2.5 py-2 text-left text-xs transition ${
danger ? 'text-red-500 hover:bg-red-50' : 'text-ink-100 hover:bg-gray-50 hover:text-brand-500'
type="button"
className={`flex w-full items-center gap-2 rounded-lg px-3 py-2 text-left text-xs font-medium transition ${
danger
? 'text-red-500 hover:bg-red-50'
: 'text-ink-100 hover:bg-gray-100 hover:text-brand-500'
}`}
title={label}
onClick={handleClick}
onClick={(e) => {
e.stopPropagation()
onClick()
}}
>
{icon}
<span>{children}</span>
+449
View File
@@ -0,0 +1,449 @@
import { useEffect, useState } from 'react'
import toast from 'react-hot-toast'
import {
Activity,
ArrowRightLeft,
CheckCircle2,
Database,
HardDrive,
HelpCircle,
Loader2,
RefreshCw,
Save,
Server,
ShieldCheck,
Zap,
} from 'lucide-react'
import {
adminAPI,
type DatabaseConnectionPayload,
type DatabaseMigrationResult,
type DatabaseStatus,
type PostgresTestResult,
} from '../api/admin'
import { confirmAction } from '../components/confirmAction'
export function DatabaseSettingsPanel() {
const [status, setStatus] = useState<DatabaseStatus | null>(null)
const [loading, setLoading] = useState(true)
const [mode, setMode] = useState<'form' | 'dsn'>('form')
// 表单状态
const [formData, setFormData] = useState<DatabaseConnectionPayload>({
host: '127.0.0.1',
port: 5432,
user: 'postgres',
password: '',
dbname: 'mmtl',
sslmode: 'disable',
dsn: '',
})
// 测试与操作状态
const [testing, setTesting] = useState(false)
const [testResult, setTestResult] = useState<PostgresTestResult | null>(null)
const [migrating, setMigrating] = useState(false)
const [migrationResult, setMigrationResult] = useState<DatabaseMigrationResult | null>(null)
const [saving, setSaving] = useState(false)
const refreshStatus = () => {
setLoading(true)
return adminAPI
.getDatabaseStatus()
.then(setStatus)
.catch((err) => toast.error('获取数据库状态失败: ' + (err.message || '网络错误')))
.finally(() => setLoading(false))
}
useEffect(() => {
refreshStatus().catch(() => undefined)
}, [])
const getPayload = (): DatabaseConnectionPayload => {
if (mode === 'dsn') {
return { type: 'postgres', dsn: formData.dsn?.trim() || '' }
}
return {
type: 'postgres',
host: formData.host?.trim() || '',
port: Number(formData.port) || 5432,
user: formData.user?.trim() || '',
password: formData.password || '',
dbname: formData.dbname?.trim() || 'mmtl',
sslmode: formData.sslmode || 'disable',
}
}
const handleTestConnection = async () => {
const payload = getPayload()
if (mode === 'form' && (!payload.host || !payload.user)) {
toast.error('请填写 PostgreSQL 主机和用户名')
return
}
if (mode === 'dsn' && !payload.dsn) {
toast.error('请填写 PostgreSQL DSN')
return
}
setTesting(true)
setTestResult(null)
try {
const res = await adminAPI.testDatabaseConnection(payload)
setTestResult(res)
if (res.success) {
toast.success(`连接成功!延迟: ${res.latency_ms}ms`)
} else {
toast.error(res.error || '连接失败')
}
} catch (err: any) {
const errorMsg = err.response?.data?.error || err.message || '测试连接异常'
setTestResult({ success: false, error: errorMsg })
toast.error(errorMsg)
} finally {
setTesting(false)
}
}
const handleMigrate = async () => {
const payload = getPayload()
if (status?.type === 'postgres') {
const ok = await confirmAction({
title: '覆盖/同步确认',
message: '当前已经处于 PostgreSQL 模式,继续迁移将覆盖/合并目标库的数据,确定继续吗?',
})
if (!ok) return
} else {
const ok = await confirmAction({
title: '开始数据库迁移',
message: '即将把当前 SQLite 数据库中的所有媒体、用户、播放记录、设置等全量迁移到目标 PostgreSQL 数据库。确定开始吗?',
})
if (!ok) return
}
setMigrating(true)
setMigrationResult(null)
try {
const res = await adminAPI.migrateDatabase(payload)
setMigrationResult(res)
if (res.success) {
toast.success(`数据迁移完成!共迁移 ${res.total_rows} 条记录`)
} else {
toast.error(res.error || '数据迁移失败')
}
} catch (err: any) {
const errorMsg = err.response?.data?.error || err.message || '迁移发生错误'
setMigrationResult({ success: false, total_rows: 0, duration_ms: 0, error: errorMsg })
toast.error(errorMsg)
} finally {
setMigrating(false)
}
}
const handleSaveAndSwitch = async () => {
const payload = getPayload()
const ok = await confirmAction({
title: '切换数据库',
message: '保存后系统配置将更新为使用 PostgreSQL。需要重启 MMTL 服务使新数据库生效。确定保存吗?',
})
if (!ok) return
setSaving(true)
try {
const res = await adminAPI.saveDatabaseConfig(payload)
toast.success(res.message || '数据库配置已保存,请重启服务生效')
refreshStatus()
} catch (err: any) {
toast.error(err.response?.data?.error || err.message || '保存配置失败')
} finally {
setSaving(false)
}
}
return (
<div className="space-y-6">
{/* 头部标题 */}
<div className="flex items-center justify-between">
<div className="flex items-center gap-3">
<Database className="h-6 w-6 text-brand-500" />
<div>
<h2 className="font-display text-lg font-semibold text-ink-600">数据库设置与迁移</h2>
<p className="text-xs text-ink-50">
管理系统底层数据库,支持在 SQLite(本地嵌入式)与 PostgreSQL(高性能关系库)之间平滑切换与数据迁移
</p>
</div>
</div>
<button
onClick={refreshStatus}
disabled={loading}
className="flex items-center gap-1.5 rounded-lg border border-gray-200 bg-sand-200/50 px-3 py-1.5 text-xs text-ink-100 hover:bg-sand-200 disabled:opacity-50"
>
<RefreshCw size={14} className={loading ? 'animate-spin' : ''} />
刷新状态
</button>
</div>
{/* 当前数据库状态卡片 */}
<div className="glass-panel p-5 space-y-4">
<div className="flex items-center justify-between border-b border-gray-200 pb-3">
<div className="flex items-center gap-2">
<Server size={18} className="text-brand-500" />
<span className="font-medium text-sm text-ink-600">当前运行引擎</span>
</div>
<div className="flex items-center gap-2">
<span
className={
'inline-flex items-center gap-1 rounded-full px-2.5 py-0.5 text-xs font-semibold ' +
(status?.type === 'postgres'
? 'bg-blue-500/10 text-blue-400 border border-blue-500/20'
: 'bg-emerald-500/10 text-emerald-400 border border-emerald-500/20')
}
>
<Zap size={12} />
{status?.type === 'postgres' ? 'PostgreSQL' : 'SQLite (WAL 优化)'}
</span>
</div>
</div>
<div className="grid grid-cols-1 md:grid-cols-3 gap-4 text-xs">
<div className="space-y-1 rounded-lg bg-sand-200/30 p-3">
<p className="text-sand-500">存储位置 / 连接</p>
<p className="font-mono text-ink-600 break-all">
{status?.type === 'postgres'
? status.dsn || '配置的 PostgreSQL 实例'
: status?.db_path || './data/mmtl.db'}
</p>
</div>
<div className="space-y-1 rounded-lg bg-sand-200/30 p-3">
<p className="text-sand-500">连接池活跃 / 最大</p>
<p className="font-mono text-ink-600">
活跃: {status?.in_use ?? 0} · 空闲: {status?.idle ?? 0} · 上限:{' '}
{status?.max_open_conns ?? 16}
</p>
</div>
<div className="space-y-1 rounded-lg bg-sand-200/30 p-3">
<p className="text-sand-500">核心表记录概览</p>
<p className="text-ink-600">
媒体: {status?.table_counts?.media ?? 0} · 用户: {status?.table_counts?.users ?? 0} ·
播放记录: {status?.table_counts?.playback_histories ?? 0}
</p>
</div>
</div>
</div>
{/* 配置 PostgreSQL */}
<div className="glass-panel p-5 space-y-5">
<div className="flex items-center justify-between border-b border-gray-200 pb-3">
<div className="flex items-center gap-2">
<HardDrive size={18} className="text-brand-500" />
<span className="font-medium text-sm text-ink-600">配置目标 PostgreSQL</span>
</div>
<div className="flex rounded-lg bg-sand-200/40 p-0.5 text-xs">
<button
onClick={() => setMode('form')}
className={
'rounded-md px-3 py-1 transition ' +
(mode === 'form'
? 'bg-brand-500 text-white font-medium shadow-sm'
: 'text-sand-500 hover:text-ink-600')
}
>
分段表单
</button>
<button
onClick={() => setMode('dsn')}
className={
'rounded-md px-3 py-1 transition ' +
(mode === 'dsn'
? 'bg-brand-500 text-white font-medium shadow-sm'
: 'text-sand-500 hover:text-ink-600')
}
>
完整 DSN
</button>
</div>
</div>
{mode === 'form' ? (
<div className="grid grid-cols-1 md:grid-cols-2 gap-4 text-sm">
<div className="space-y-1.5">
<label className="text-xs font-medium text-sand-500">主机地址 (Host)</label>
<input
type="text"
value={formData.host || ''}
onChange={(e) => setFormData({ ...formData, host: e.target.value })}
placeholder="例如 127.0.0.1 或 postgres"
className="w-full rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none"
/>
</div>
<div className="space-y-1.5">
<label className="text-xs font-medium text-sand-500">端口 (Port)</label>
<input
type="number"
value={formData.port || 5432}
onChange={(e) => setFormData({ ...formData, port: Number(e.target.value) })}
placeholder="5432"
className="w-full rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none"
/>
</div>
<div className="space-y-1.5">
<label className="text-xs font-medium text-sand-500">数据库名 (Database)</label>
<input
type="text"
value={formData.dbname || ''}
onChange={(e) => setFormData({ ...formData, dbname: e.target.value })}
placeholder="mmtl"
className="w-full rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none"
/>
</div>
<div className="space-y-1.5">
<label className="text-xs font-medium text-sand-500">用户名 (User)</label>
<input
type="text"
value={formData.user || ''}
onChange={(e) => setFormData({ ...formData, user: e.target.value })}
placeholder="postgres"
className="w-full rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none"
/>
</div>
<div className="space-y-1.5">
<label className="text-xs font-medium text-sand-500">密码 (Password)</label>
<input
type="password"
value={formData.password || ''}
onChange={(e) => setFormData({ ...formData, password: e.target.value })}
placeholder="••••••••"
className="w-full rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none"
/>
</div>
<div className="space-y-1.5">
<label className="text-xs font-medium text-sand-500">SSL 模式 (SSL Mode)</label>
<select
value={formData.sslmode || 'disable'}
onChange={(e) => setFormData({ ...formData, sslmode: e.target.value })}
className="w-full rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 focus:border-brand-500 focus:outline-none"
>
<option value="disable">disable (关闭 SSL)</option>
<option value="require">require (强制 SSL)</option>
<option value="verify-ca">verify-ca (验证 CA)</option>
<option value="verify-full">verify-full (严格验证证书与主机名)</option>
</select>
</div>
</div>
) : (
<div className="space-y-1.5 text-sm">
<label className="text-xs font-medium text-sand-500">
PostgreSQL DSN 字符串 (URL 格式)
</label>
<input
type="text"
value={formData.dsn || ''}
onChange={(e) => setFormData({ ...formData, dsn: e.target.value })}
placeholder="postgres://user:password@127.0.0.1:5432/mmtl?sslmode=disable"
className="w-full font-mono text-xs rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none"
/>
</div>
)}
{/* 测试结果卡片 */}
{testResult && (
<div
className={
'flex items-start gap-2.5 rounded-lg p-3.5 text-xs ' +
(testResult.success
? 'bg-emerald-500/10 border border-emerald-500/20 text-emerald-400'
: 'bg-rose-500/10 border border-rose-500/20 text-rose-400')
}
>
{testResult.success ? (
<CheckCircle2 size={16} className="mt-0.5 shrink-0" />
) : (
<HelpCircle size={16} className="mt-0.5 shrink-0" />
)}
<div className="space-y-0.5">
<p className="font-semibold">
{testResult.success ? `连接测试通过 (${testResult.latency_ms} ms)` : '连接测试未通过'}
</p>
{testResult.version && <p className="text-ink-100">{testResult.version}</p>}
{testResult.error && <p className="text-rose-300 font-mono">{testResult.error}</p>}
</div>
</div>
)}
{/* 迁移结果卡片 */}
{migrationResult && (
<div
className={
'rounded-lg p-3.5 text-xs space-y-2 ' +
(migrationResult.success
? 'bg-blue-500/10 border border-blue-500/20 text-blue-400'
: 'bg-rose-500/10 border border-rose-500/20 text-rose-400')
}
>
<div className="flex items-center gap-2 font-semibold">
<ShieldCheck size={16} />
<span>{migrationResult.message || (migrationResult.success ? '迁移完成' : '迁移失败')}</span>
{migrationResult.success && (
<span className="text-sand-500 text-[11px]">
(耗时: {migrationResult.duration_ms} ms)
</span>
)}
</div>
{migrationResult.table_rows && Object.keys(migrationResult.table_rows).length > 0 && (
<div className="grid grid-cols-2 sm:grid-cols-3 gap-2 pt-1 font-mono text-[11px] text-ink-100">
{Object.entries(migrationResult.table_rows).map(([tbl, count]) => (
<div key={tbl} className="rounded bg-sand-200/40 px-2 py-1 flex justify-between">
<span>{tbl}:</span>
<span className="font-bold text-brand-400">{count} 条</span>
</div>
))}
</div>
)}
{migrationResult.error && (
<p className="text-rose-300 font-mono">{migrationResult.error}</p>
)}
</div>
)}
{/* 操作按钮区 */}
<div className="flex flex-wrap items-center justify-between gap-3 pt-2 border-t border-gray-200">
<button
type="button"
onClick={handleTestConnection}
disabled={testing || migrating || saving}
className="flex items-center gap-1.5 rounded-lg border border-gray-200 bg-sand-200/60 px-4 py-2 text-xs font-medium text-ink-600 hover:bg-sand-200 disabled:opacity-50"
>
{testing ? <Loader2 size={14} className="animate-spin" /> : <Activity size={14} />}
测试连接
</button>
<div className="flex items-center gap-2">
<button
type="button"
onClick={handleMigrate}
disabled={testing || migrating || saving}
className="flex items-center gap-1.5 rounded-lg border border-primary-500/30 bg-primary-500/10 px-4 py-2 text-xs font-medium text-brand-400 hover:bg-primary-500/20 disabled:opacity-50"
>
{migrating ? (
<Loader2 size={14} className="animate-spin" />
) : (
<ArrowRightLeft size={14} />
)}
一键数据迁移到 PostgreSQL
</button>
<button
type="button"
onClick={handleSaveAndSwitch}
disabled={testing || migrating || saving}
className="neon-button text-xs disabled:opacity-50"
>
{saving ? <Loader2 size={14} className="animate-spin" /> : <Save size={14} />}
保存并切换
</button>
</div>
</div>
</div>
</div>
)
}
+9 -29
View File
@@ -13,21 +13,12 @@ import {
HomeLoadingState,
} from './HomePageSections'
const CAROUSEL_STORAGE_KEY = 'mmtl.home.carousel_libraries'
const hasArtwork = (media?: Media | null) => !!(media?.poster_url || media?.backdrop_url)
const asArray = <T,>(value: unknown): T[] => (Array.isArray(value) ? (value as T[]) : [])
export function HomePage() {
const [libraries, setLibraries] = useState<Library[]>([])
const [libraryData, setLibraryData] = useState<Record<string, { cards: SeriesCard[]; items: Media[]; total: number }>>({})
const [selectedLibraryIds, setSelectedLibraryIds] = useState<string[]>(() => {
try {
const saved = localStorage.getItem(CAROUSEL_STORAGE_KEY)
return saved ? JSON.parse(saved) : []
} catch {
return []
}
})
const [history, setHistory] = useState<HistoryItem[]>([])
const [loading, setLoading] = useState(true)
@@ -46,18 +37,6 @@ export function HomePage() {
setLibraries(libs)
setHistory(hist.filter((h) => h && !h.completed && !!h.media))
// Set default selected libraries if none saved yet
setSelectedLibraryIds((current) => {
if (current.length > 0) return current
const allIds = libs.map((l) => l.id)
try {
localStorage.setItem(CAROUSEL_STORAGE_KEY, JSON.stringify(allIds))
} catch {
// ignore
}
return allIds
})
// Fetch media items for all libraries in parallel
const isSeriesType = (type?: string) => type === 'tv' || type === 'anime' || type === 'variety'
const results = await Promise.allSettled(
@@ -132,16 +111,18 @@ export function HomePage() {
return counts
}, [libraries, libraryData])
// Compute items to show in the Hero Carousel
// Compute items to show in the Hero Carousel. Only libraries flagged
// carousel_enabled contribute; a library's items array is empty for
// series-type libs (loaded via /series), so gate on cards instead.
const carouselItems = useMemo(() => {
const candidateMedia: Media[] = []
const effectiveSelectedIds =
selectedLibraryIds.length > 0 ? selectedLibraryIds : libraries.map((l) => l.id)
const effectiveSelectedIds = libraries
.filter((l) => l.carousel_enabled === true)
.map((l) => l.id)
for (const libId of effectiveSelectedIds) {
const data = libraryData[libId]
if (data && data.items.length > 0) {
// Pick representative items from series cards or raw items
if (data && data.cards.length > 0) {
for (const card of data.cards) {
if (hasArtwork(card.rep)) {
candidateMedia.push(card.rep)
@@ -150,9 +131,8 @@ export function HomePage() {
}
}
// Sort by artwork score / rating or shuffle / interleave
// Fallback to all loaded items with artwork
if (candidateMedia.length === 0) {
// Fallback to all loaded items with artwork
for (const lib of libraries) {
const data = libraryData[lib.id]
if (data) {
@@ -164,7 +144,7 @@ export function HomePage() {
}
return candidateMedia.slice(0, 10)
}, [selectedLibraryIds, libraries, libraryData])
}, [libraries, libraryData])
const empty =
!loading &&
+36 -30
View File
@@ -1,7 +1,8 @@
import { useEffect, useMemo, useState } from 'react'
import { useCallback, useEffect, useMemo, useState } from 'react'
import { libraryAPI } from '../api/library'
import { toolsAPI } from '../api/tools'
import { openManageLibrariesDialog } from '../components/manageLibrariesDialog'
import {
LibrariesContent,
LibrariesEmptyState,
@@ -16,6 +17,32 @@ export function LibrariesPage() {
const [repairEpisodeArtwork, setRepairEpisodeArtwork] = useState(false)
const [repairMsg, setRepairMsg] = useState('')
const loadLibraries = useCallback(async () => {
setLoading(true)
try {
const libs = await libraryAPI.list()
const rows = await Promise.all(libs.map(async (library) => {
try {
if (isSeriesLibraryType(library.type)) {
const [seriesPage, mediaPage] = await Promise.all([
libraryAPI.listSeries(library.id, 1, 10),
libraryAPI.listMedia(library.id, 1, 1, { groupVersions: false }),
])
return { library, items: [], total: mediaPage.total, cards: seriesPage.items ?? [] } satisfies LibraryPreview
}
const page = await libraryAPI.listMedia(library.id, 1, 160, { groupVersions: false })
const cards = latestLibraryCards(page.items)
return { library, items: page.items, total: page.total, cards } satisfies LibraryPreview
} catch {
return { library, items: [], total: 0, cards: [] } satisfies LibraryPreview
}
}))
setPreviews(rows)
} finally {
setLoading(false)
}
}, [])
async function handleRepairRescrape() {
if (repairing) return
setRepairing(true)
@@ -30,36 +57,14 @@ export function LibrariesPage() {
}
}
const handleManageLibraries = async () => {
await openManageLibrariesDialog()
await loadLibraries()
}
useEffect(() => {
let cancelled = false
async function load() {
setLoading(true)
try {
const libs = await libraryAPI.list()
const rows = await Promise.all(libs.map(async (library) => {
try {
if (isSeriesLibraryType(library.type)) {
const [seriesPage, mediaPage] = await Promise.all([
libraryAPI.listSeries(library.id, 1, 10),
libraryAPI.listMedia(library.id, 1, 1, { groupVersions: false }),
])
return { library, items: [], total: mediaPage.total, cards: seriesPage.items ?? [] } satisfies LibraryPreview
}
const page = await libraryAPI.listMedia(library.id, 1, 160, { groupVersions: false })
const cards = latestLibraryCards(page.items)
return { library, items: page.items, total: page.total, cards } satisfies LibraryPreview
} catch {
return { library, items: [], total: 0, cards: [] } satisfies LibraryPreview
}
}))
if (!cancelled) setPreviews(rows)
} finally {
if (!cancelled) setLoading(false)
}
}
load()
return () => { cancelled = true }
}, [])
loadLibraries().catch(() => undefined)
}, [loadLibraries])
const total = useMemo(() => previews.reduce((sum, preview) => sum + preview.total, 0), [previews])
@@ -77,6 +82,7 @@ export function LibrariesPage() {
repairing={repairing}
onRepairEpisodeArtworkChange={setRepairEpisodeArtwork}
onRepairRescrape={handleRepairRescrape}
onManageLibraries={handleManageLibraries}
/>
{previews.length === 0 ? (
+8 -3
View File
@@ -1,12 +1,11 @@
import type { ReactNode } from 'react'
import { Link } from 'react-router-dom'
import { motion } from 'framer-motion'
import { ArrowRight, Film, FolderOpen, Library as LibraryIcon, Music, PlayCircle, RefreshCw, Tv } from 'lucide-react'
import { ArrowRight, Film, FolderOpen, Library as LibraryIcon, Music, PlayCircle, RefreshCw, Sparkles, Tv } from 'lucide-react'
import { imageURL } from '../api/client'
import { EpisodeArtworkToggle } from '../components/EpisodeArtworkToggle'
import { MediaCard } from '../components/MediaCard'
import { openManageLibrariesDialog } from '../components/manageLibrariesDialog'
import { seriesCardLink } from '../utils/groupSeries'
import { libraryDisplayPath } from './libraryDisplayModel'
import { libraryArtworkItems, type LibraryPreview } from './librariesPageModel'
@@ -37,6 +36,7 @@ export function LibrariesHeader({
repairing,
onRepairEpisodeArtworkChange,
onRepairRescrape,
onManageLibraries,
}: {
previewCount: number
total: number
@@ -45,6 +45,7 @@ export function LibrariesHeader({
repairing: boolean
onRepairEpisodeArtworkChange: (value: boolean) => void
onRepairRescrape: () => void
onManageLibraries: () => void
}) {
return (
<div className="flex flex-wrap items-end justify-between gap-4">
@@ -72,7 +73,11 @@ export function LibrariesHeader({
<RefreshCw size={14} className={repairing ? 'animate-spin' : ''} />
{repairing ? '正在启动…' : '全库修复+重刮'}
</button>
<button type="button" onClick={() => openManageLibrariesDialog()} className="btn-outline">
<Link to="/scraper/queue" className="btn-outline inline-flex items-center gap-1.5" title="查看正在进行的刮削任务与进度">
<Sparkles size={14} className="text-brand-500" />
<span>刮削队列</span>
</Link>
<button type="button" onClick={onManageLibraries} className="btn-outline">
管理媒体库
</button>
</div>
+17 -15
View File
@@ -113,21 +113,23 @@ export function LibraryPage() {
return (
<div className="space-y-6">
<LibraryPageHeader
library={library}
itemCount={isSeries ? seriesCards.length : total}
loadingAllText={loadingAllText}
scanProgress={scanProgress}
isAdmin={role === 'admin'}
scrapeEpisodeArtwork={scrapeEpisodeArtwork}
scanning={scanning}
scraping={scraping}
repairing={repairing}
onScrapeEpisodeArtworkChange={setScrapeEpisodeArtwork}
onScan={handleScan}
onScrape={() => setScrapeDialogOpen(true)}
onRepairRescrape={handleRepairRescrape}
/>
{!selectedSeries && (
<LibraryPageHeader
library={library}
itemCount={isSeries ? seriesCards.length : total}
loadingAllText={loadingAllText}
scanProgress={scanProgress}
isAdmin={role === 'admin'}
scrapeEpisodeArtwork={scrapeEpisodeArtwork}
scanning={scanning}
scraping={scraping}
repairing={repairing}
onScrapeEpisodeArtworkChange={setScrapeEpisodeArtwork}
onScan={handleScan}
onScrape={() => setScrapeDialogOpen(true)}
onRepairRescrape={handleRepairRescrape}
/>
)}
<LibraryMediaSections
isSeries={isSeries}
+4 -4
View File
@@ -55,10 +55,10 @@ export function LibraryPageHeader({
className="h-10"
/>
<button onClick={onScan} disabled={scanning} className="btn-outline">
{scanning ? '扫描中…' : '立即扫描'}
{scanning ? '扫描中…' : '扫描媒体库'}
</button>
<button onClick={onScrape} disabled={scraping} className="btn-outline">
{scraping ? '刮削中…' : '刮削元数据'}
<button onClick={onScrape} disabled={scraping} className="btn-outline" title="对整个媒体库执行刮削元数据">
{scraping ? '刮削中…' : '整库刮削元数据'}
</button>
<button
onClick={onRepairRescrape}
@@ -66,7 +66,7 @@ export function LibraryPageHeader({
className="btn-outline"
title="回填本库占位符外部 ID 并重刮,修正空 ID / 拆集问题"
>
{repairing ? '修复中…' : '修复+重刮本库'}
{repairing ? '修复中…' : '修复+重刮整库'}
</button>
</div>
)}
-305
View File
@@ -1,305 +0,0 @@
import { useEffect, useState } from 'react'
import {
Check,
CheckSquare,
Film,
FolderOpen,
HeartHandshake,
Layers,
Loader2,
Music,
Save,
SlidersHorizontal,
Square,
Tv,
} from 'lucide-react'
import toast from 'react-hot-toast'
import { adminAPI } from '../api/admin'
import { imageURL } from '../api/client'
import { libraryAPI } from '../api/library'
import type { Library, Setting } from '../types'
import { groupSeries, type SeriesCard } from '../utils/groupSeries'
import { getLibraryArtworks } from './librariesPageModel'
const CAROUSEL_STORAGE_KEY = 'mmtl.home.carousel_libraries'
const SETTING_KEY_CAROUSEL = 'home.carousel_libraries'
const TYPE_ICONS: Record<string, React.ReactNode> = {
movie: <Film size={20} />,
movies: <Film size={20} />,
tv: <Tv size={20} />,
series: <Tv size={20} />,
anime: <Layers size={20} />,
shows: <Tv size={20} />,
variety: <Tv size={20} />,
music: <Music size={20} />,
adult: <HeartHandshake size={20} />,
}
const TYPE_LABELS: Record<string, string> = {
movie: '电影',
movies: '电影',
tv: '剧集',
series: '剧集',
anime: '动漫',
shows: '综艺',
variety: '综艺',
music: '音乐',
adult: 'Adult',
}
export function LibrarySettingsPanel() {
const [libraries, setLibraries] = useState<Library[]>([])
const [libraryCards, setLibraryCards] = useState<Record<string, SeriesCard[]>>({})
const [selectedIds, setSelectedIds] = useState<string[]>([])
const [loading, setLoading] = useState(true)
const [saving, setSaving] = useState(false)
const [dirty, setDirty] = useState(false)
useEffect(() => {
async function load() {
setLoading(true)
try {
const [libs, settings] = await Promise.all([
libraryAPI.list({ includeHidden: true }).catch(() => [] as Library[]),
adminAPI.listSettings().catch(() => [] as Setting[]),
])
const libList = Array.isArray(libs) ? libs : []
setLibraries(libList)
// 异步拉取各个媒体库的前几个条目用于封面展示
Promise.allSettled(
libList.map(async (lib) => {
const page = await libraryAPI.listMedia(lib.id, 1, 10)
const items = Array.isArray(page?.items) ? page.items : []
return { id: lib.id, cards: groupSeries(items) }
}),
).then((results) => {
const map: Record<string, SeriesCard[]> = {}
for (const r of results) {
if (r.status === 'fulfilled' && r.value) {
map[r.value.id] = r.value.cards
}
}
setLibraryCards(map)
})
// 优先读取系统配置,其次读取 localStorage,默认全选
const settingItem = Array.isArray(settings)
? settings.find((s) => s.key === SETTING_KEY_CAROUSEL)
: undefined
let initialIds: string[] | null = null
if (settingItem?.value) {
try {
initialIds = JSON.parse(settingItem.value)
} catch {
// ignore
}
}
if (!initialIds) {
try {
const saved = localStorage.getItem(CAROUSEL_STORAGE_KEY)
if (saved) initialIds = JSON.parse(saved)
} catch {
// ignore
}
}
if (Array.isArray(initialIds)) {
setSelectedIds(initialIds)
} else {
setSelectedIds(libList.map((l) => l.id))
}
} finally {
setLoading(false)
}
}
load()
}, [])
const toggleLibrary = (id: string) => {
setSelectedIds((prev) => {
const next = prev.includes(id) ? prev.filter((item) => item !== id) : [...prev, id]
setDirty(true)
return next
})
}
const selectAll = () => {
setSelectedIds(libraries.map((l) => l.id))
setDirty(true)
}
const deselectAll = () => {
setSelectedIds([])
setDirty(true)
}
const handleSave = async () => {
setSaving(true)
try {
const jsonValue = JSON.stringify(selectedIds)
await adminAPI.updateSetting(SETTING_KEY_CAROUSEL, jsonValue)
try {
localStorage.setItem(CAROUSEL_STORAGE_KEY, jsonValue)
} catch {
// ignore
}
toast.success('海报轮播设置已保存')
setDirty(false)
} catch (err: unknown) {
const msg =
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
'保存设置失败'
toast.error(msg)
} finally {
setSaving(false)
}
}
if (loading) {
return (
<div className="flex justify-center py-12 text-ink-50">
<Loader2 className="animate-spin" />
</div>
)
}
return (
<div className="glass-panel space-y-6">
{/* 头部说明与快捷操作 */}
<div className="flex flex-col gap-4 sm:flex-row sm:items-center sm:justify-between border-b border-gray-200/80 pb-4">
<div className="flex items-start gap-3">
<div className="rounded-xl border border-primary-400/40 bg-primary-400/10 p-2 text-brand-500 mt-0.5">
<SlidersHorizontal size={20} />
</div>
<div>
<h3 className="font-display text-lg font-bold text-ink-600">首页海报轮播设置</h3>
<p className="text-xs text-sand-500 mt-0.5">
选择参与首页顶部大图海报轮播推荐的媒体库。勾选的媒体库内容将轮流展示在系统首页顶部。
</p>
</div>
</div>
<div className="flex items-center gap-2 self-end sm:self-auto">
<button
type="button"
onClick={selectAll}
className="flex items-center gap-1.5 rounded-lg border border-gray-200 bg-white px-2.5 py-1.5 text-xs font-semibold text-ink-100 transition hover:border-primary-400/50 hover:text-brand-500"
>
<CheckSquare size={14} />
<span>全选</span>
</button>
<button
type="button"
onClick={deselectAll}
className="flex items-center gap-1.5 rounded-lg border border-gray-200 bg-white px-2.5 py-1.5 text-xs font-semibold text-ink-100 transition hover:border-primary-400/50 hover:text-brand-500"
>
<Square size={14} />
<span>清空</span>
</button>
</div>
</div>
{/* 媒体库列表卡片 */}
{libraries.length === 0 ? (
<div className="py-8 text-center text-xs text-sand-500">
暂无可用媒体库,请先添加媒体库后再配置海报轮播。
</div>
) : (
<div className="grid grid-cols-1 gap-3 sm:grid-cols-2 lg:grid-cols-3">
{libraries.map((lib) => {
const isSelected = selectedIds.includes(lib.id)
const cards = libraryCards[lib.id] || []
const artwork = getLibraryArtworks(lib, cards)
return (
<div
key={lib.id}
onClick={() => toggleLibrary(lib.id)}
className={`flex cursor-pointer items-center justify-between rounded-2xl border p-4 transition-all duration-200 select-none ${
isSelected
? 'border-brand-500/60 bg-primary-400/10 shadow-sm'
: 'border-gray-200 bg-white/70 hover:border-gray-300'
}`}
>
<div className="flex items-center gap-3 min-w-0">
<div
className={`grid h-12 w-16 shrink-0 gap-0.5 overflow-hidden rounded-xl bg-gray-100 shadow-inner transition-colors ${
artwork.length > 1 ? 'grid-cols-2' : 'grid-cols-1'
} ${
isSelected ? 'ring-2 ring-brand-500/30' : ''
}`}
>
{artwork.length > 0 ? (
artwork.map(({ src, version }, index) => (
<img
key={`${src}-${index}`}
src={imageURL(src, version)}
alt=""
className="h-full w-full object-cover"
referrerPolicy="no-referrer"
onError={(e) => {
e.currentTarget.style.display = 'none'
}}
/>
))
) : (
<div className="flex h-full w-full items-center justify-center text-ink-50">
{TYPE_ICONS[lib.type] || <FolderOpen size={20} />}
</div>
)}
</div>
<div className="min-w-0">
<div className="flex items-center gap-1.5">
<span className="truncate font-display text-sm font-bold text-ink-600">
{lib.name}
</span>
<span className="shrink-0 rounded px-1.5 py-0.5 text-[10px] font-bold border border-gray-200 bg-gray-50 text-ink-50">
{TYPE_LABELS[lib.type] || '自定义'}
</span>
</div>
<p className="text-xs text-sand-500 mt-0.5">
{isSelected ? '已启用轮播' : '未参与轮播'}
</p>
</div>
</div>
<div
className={`flex h-6 w-6 shrink-0 items-center justify-center rounded-lg border transition-colors ${
isSelected
? 'border-brand-500 bg-brand-500 text-white'
: 'border-gray-300 bg-white text-transparent'
}`}
>
<Check size={14} strokeWidth={3} />
</div>
</div>
)
})}
</div>
)}
{/* 底部保存按钮 */}
<div className="flex items-center justify-between pt-2 border-t border-gray-200/80">
<span className="text-xs text-sand-500">
已选择 {selectedIds.length} / {libraries.length} 个媒体库参与轮播
</span>
<button
type="button"
onClick={handleSave}
disabled={saving || !dirty}
className="neon-button disabled:opacity-50"
>
{saving ? <Loader2 size={16} className="animate-spin" /> : <Save size={16} />}
保存设置
</button>
</div>
</div>
)
}

Some files were not shown because too many files have changed in this diff Show More