mirror of
https://github.com/truewhile/MeBox.git
synced 2026-09-28 11:16:37 +08:00
Compare commits
69 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| e6473300a7 | |||
| 994f64f753 | |||
| 90064a5480 | |||
| 6e8eac9887 | |||
| d3051eaffe | |||
| 496a897782 | |||
| 22b7290ee1 | |||
| 14037d5dea | |||
| 152db3fb9f | |||
| 0413d123da | |||
| 5f6bd7b5cd | |||
| 3c25bb5d61 | |||
| 659b91b000 | |||
| fe5b3bd56a | |||
| 82bbb116ae | |||
| 9f5ff7e6f0 | |||
| a00504080a | |||
| fc6e2e6f10 | |||
| 65c5f3e4bf | |||
| 07e340251b | |||
| 9b956b928b | |||
| c3187f6e3f | |||
| 60c815a8b3 | |||
| 41b155ea31 | |||
| 2888ae8bf7 | |||
| 7363064d89 | |||
| 1d53bf2ae1 | |||
| 618165ec31 | |||
| 87c66a9b8c | |||
| 1ea4724261 | |||
| 3d372f039e | |||
| 0384017e98 | |||
| 0332579d5f | |||
| 6aefe18caa | |||
| ef72fc8d83 | |||
| 4764c09572 | |||
| 98ca766a37 | |||
| 13c9035b76 | |||
| ad6d0ba21d | |||
| 431f7f088b | |||
| 3f13ed1113 | |||
| 9d359c40dd | |||
| 0e7dbd6215 | |||
| 7fd8de91cb | |||
| 5f323eb2ce | |||
| c0ac8bf11a | |||
| b676733af7 | |||
| 7a2027a3a7 | |||
| 9ffb74adce | |||
| 13faff7078 | |||
| a8a4e88d86 | |||
| 8fa5db88ff | |||
| 3c325f81c8 | |||
| 585434010c | |||
| eb1a705cae | |||
| 10770b2b77 | |||
| b61c51e064 | |||
| 019ecbec7b | |||
| 6a96c5640e | |||
| 4025e92cb4 | |||
| 2355419ef9 | |||
| 86214ea796 | |||
| 774f2d4695 | |||
| efa64051ea | |||
| 5095347ace | |||
| ad9260f8fe | |||
| 4ae2502096 | |||
| c05b5259ef | |||
| da1fb02c9d |
@@ -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
|
||||
|
||||
@@ -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 ./...
|
||||
|
||||
|
||||
@@ -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)。
|
||||
|
||||
+9
-23
@@ -12,10 +12,7 @@ package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"strings"
|
||||
@@ -108,30 +105,19 @@ func main() {
|
||||
|
||||
router := buildRouter(cfg, logger, services)
|
||||
|
||||
srv := &http.Server{
|
||||
Addr: fmt.Sprintf(":%d", cfg.App.Port),
|
||||
Handler: router,
|
||||
ReadHeaderTimeout: 15 * time.Second,
|
||||
serverMgr := newServerManager(cfg, logger, router)
|
||||
services.ReloadHTTPServer = serverMgr.Reload
|
||||
if err := serverMgr.Start(); err != nil {
|
||||
logger.Fatal("listen failed", zap.Error(err))
|
||||
}
|
||||
|
||||
ln, err := net.Listen("tcp", srv.Addr)
|
||||
if err != nil {
|
||||
logger.Fatal("listen failed", zap.String("addr", srv.Addr), zap.Error(err))
|
||||
}
|
||||
localIP := getLocalIP()
|
||||
logger.Info("server is ready",
|
||||
zap.String("local", fmt.Sprintf("http://%s:%d", localIP, cfg.App.Port)),
|
||||
zap.String("listen", srv.Addr),
|
||||
)
|
||||
go func() {
|
||||
if err := srv.Serve(ln); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||
logger.Fatal("listen failed", zap.Error(err))
|
||||
scheme := "http"
|
||||
if cfg.App.HTTPSEnabled {
|
||||
scheme = "https"
|
||||
}
|
||||
}()
|
||||
go func() {
|
||||
if publicIP := getPublicIP(3 * time.Second); publicIP != "" {
|
||||
logger.Info("server public endpoint",
|
||||
zap.String("public", fmt.Sprintf("http://%s:%d", publicIP, cfg.App.Port)),
|
||||
zap.String("public", fmt.Sprintf("%s://%s:%d", scheme, publicIP, cfg.App.Port)),
|
||||
)
|
||||
}
|
||||
}()
|
||||
@@ -145,7 +131,7 @@ func main() {
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
||||
defer cancel()
|
||||
if err := srv.Shutdown(ctx); err != nil {
|
||||
if err := serverMgr.Shutdown(ctx); err != nil {
|
||||
logger.Error("graceful shutdown failed", zap.Error(err))
|
||||
}
|
||||
services.Close()
|
||||
|
||||
@@ -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
@@ -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) {
|
||||
|
||||
@@ -0,0 +1,267 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/config"
|
||||
"github.com/ShukeBta/MMTL/internal/service"
|
||||
)
|
||||
|
||||
// tlsPair 记录当前正在服务的证书,用于判断是否需要重新绑定监听。
|
||||
type tlsPair struct {
|
||||
cert tls.Certificate
|
||||
certPEM string
|
||||
keyPEM string
|
||||
// version 是解析后的证书/私钥指纹;内容或磁盘文件变化都会导致其改变,
|
||||
// 据此决定是否需要重新绑定监听。
|
||||
version string
|
||||
}
|
||||
|
||||
// serverManager 负责 MMTL 的 HTTP/HTTPS 监听。HTTPS 设置保存后调用 Reload,
|
||||
// 在同一个端口上把明文 HTTP 与 TLS 监听热切换,无需重启进程:
|
||||
//
|
||||
// - 关闭旧监听释放端口(同一进程内 Windows 不允许重复绑定同一端口);
|
||||
// - 按最新配置重新绑定并立即对外服务;
|
||||
// - 旧服务器随后优雅退出,正在进行的播放/请求不会被立刻掐断。
|
||||
//
|
||||
// 任何校验失败都会中止切换并保留旧监听,保证用户不会被锁在服务外面。
|
||||
type serverManager struct {
|
||||
cfg *config.Config
|
||||
log *zap.Logger
|
||||
handler http.Handler
|
||||
addr string
|
||||
|
||||
mu sync.Mutex
|
||||
srv *http.Server
|
||||
ln net.Listener
|
||||
pair *tlsPair
|
||||
stopCh chan struct{}
|
||||
autoReloadStarted bool
|
||||
}
|
||||
|
||||
func newServerManager(cfg *config.Config, log *zap.Logger, handler http.Handler) *serverManager {
|
||||
return &serverManager{
|
||||
cfg: cfg,
|
||||
log: log,
|
||||
handler: handler,
|
||||
addr: fmt.Sprintf(":%d", cfg.App.Port),
|
||||
stopCh: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
// Start 启动监听。即使 HTTPS 配置损坏也退回明文 HTTP 继续启动,避免服务冷启动失败。
|
||||
func (m *serverManager) Start() error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
pair, err := m.desiredPair()
|
||||
if err != nil {
|
||||
m.log.Error("invalid HTTPS config at startup, serving plain HTTP instead", zap.Error(err))
|
||||
pair = nil
|
||||
}
|
||||
if err := m.bind(pair); err != nil {
|
||||
return err
|
||||
}
|
||||
m.logServerReady()
|
||||
m.maybeStartAutoReloadLocked()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Reload 依据最新配置热切换监听。返回的错误会带给调用它的设置接口;若新监听
|
||||
// 绑定失败会自动回滚到旧配置继续服务。
|
||||
func (m *serverManager) Reload() error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
pair, err := m.desiredPair()
|
||||
if err != nil {
|
||||
m.log.Error("server reload aborted", zap.Error(err))
|
||||
return err
|
||||
}
|
||||
if m.pairEquals(pair) {
|
||||
return nil
|
||||
}
|
||||
|
||||
oldSrv, oldLn, oldPair := m.srv, m.ln, m.pair
|
||||
if oldLn != nil {
|
||||
_ = oldLn.Close() // 释放端口后再绑定新监听
|
||||
}
|
||||
m.srv, m.ln, m.pair = nil, nil, nil
|
||||
|
||||
firstErr := m.bind(pair)
|
||||
if firstErr != nil {
|
||||
m.log.Error("bind new listener failed, rolling back to previous", zap.Error(firstErr))
|
||||
if rbErr := m.bind(oldPair); rbErr != nil {
|
||||
return fmt.Errorf("reload failed: %v; rollback failed: %v", firstErr, rbErr)
|
||||
}
|
||||
}
|
||||
// 新监听已就绪,让旧服务器在新连接切换到新监听后优雅退出。
|
||||
m.drain(oldSrv)
|
||||
m.logServerReady()
|
||||
m.maybeStartAutoReloadLocked()
|
||||
return firstErr
|
||||
}
|
||||
|
||||
// Shutdown 优雅停止当前服务器(用于进程退出)。
|
||||
func (m *serverManager) Shutdown(ctx context.Context) error {
|
||||
select {
|
||||
case <-m.stopCh:
|
||||
default:
|
||||
close(m.stopCh)
|
||||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.srv == nil {
|
||||
return nil
|
||||
}
|
||||
return m.srv.Shutdown(ctx)
|
||||
}
|
||||
|
||||
// desiredPair 根据当前配置计算目标监听形态:nil 表示明文 HTTP,非 nil 表示 TLS。
|
||||
// 证书/私钥按"路径优先、内容兜底"解析,并校验是否匹配。
|
||||
func (m *serverManager) desiredPair() (*tlsPair, error) {
|
||||
if m.cfg == nil || !m.cfg.App.HTTPSEnabled {
|
||||
return nil, nil
|
||||
}
|
||||
certPEM, err := service.ResolveSSLMaterial(m.cfg.App.SSLCert, m.cfg.App.SSLCertPath, "证书")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
keyPEM, err := service.ResolveSSLMaterial(m.cfg.App.SSLKey, m.cfg.App.SSLKeyPath, "私钥")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := service.ValidateSSLKeyPair(certPEM, keyPEM); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cert, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSL 证书/私钥无效:%v", err)
|
||||
}
|
||||
return &tlsPair{
|
||||
cert: cert,
|
||||
certPEM: certPEM,
|
||||
keyPEM: keyPEM,
|
||||
version: certPEM + "\x00" + keyPEM,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// maybeStartAutoReloadLocked 在证书/私钥通过文件路径配置时,幂等地启动后台轮询,
|
||||
// 便于运行中切换到路径方式(或换证)后无需重启也能热更新。调用方需持有 m.mu。
|
||||
func (m *serverManager) maybeStartAutoReloadLocked() {
|
||||
if m.autoReloadStarted {
|
||||
return
|
||||
}
|
||||
if !m.pathBased() {
|
||||
return
|
||||
}
|
||||
m.autoReloadStarted = true
|
||||
m.startAutoReload()
|
||||
}
|
||||
|
||||
// pathBased 是否至少有一侧证书/私钥通过文件路径配置。
|
||||
func (m *serverManager) pathBased() bool {
|
||||
return strings.TrimSpace(m.cfg.App.SSLCertPath) != "" || strings.TrimSpace(m.cfg.App.SSLKeyPath) != ""
|
||||
}
|
||||
|
||||
// startAutoReload 后台轮询文件变更并自动热更新,方便换证。
|
||||
func (m *serverManager) startAutoReload() {
|
||||
ticker := time.NewTicker(30 * time.Second)
|
||||
go func() {
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-m.stopCh:
|
||||
return
|
||||
case <-ticker.C:
|
||||
if !m.pathBased() {
|
||||
continue // 路径已清空(改回内容配置),不再轮询
|
||||
}
|
||||
if err := m.Reload(); err != nil {
|
||||
m.log.Warn("periodic https reload failed", zap.Error(err))
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// pairEquals 判断目标配置与当前监听是否一致,一致则无需重新绑定。
|
||||
func (m *serverManager) pairEquals(pair *tlsPair) bool {
|
||||
if pair == nil && m.pair == nil {
|
||||
return true
|
||||
}
|
||||
if pair == nil || m.pair == nil {
|
||||
return false
|
||||
}
|
||||
return pair.version == m.pair.version
|
||||
}
|
||||
|
||||
// bind 创建并按需启用 TLS 的监听,异步开始服务。
|
||||
func (m *serverManager) bind(pair *tlsPair) error {
|
||||
ln, err := net.Listen("tcp", m.addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("listen %s: %w", m.addr, err)
|
||||
}
|
||||
srv := &http.Server{
|
||||
Handler: m.handler,
|
||||
ReadHeaderTimeout: 15 * time.Second,
|
||||
}
|
||||
if pair != nil {
|
||||
ln = tls.NewListener(ln, &tls.Config{
|
||||
Certificates: []tls.Certificate{pair.cert},
|
||||
MinVersion: tls.VersionTLS12,
|
||||
})
|
||||
}
|
||||
m.srv, m.ln, m.pair = srv, ln, pair
|
||||
go m.serve(srv, ln)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *serverManager) serve(s *http.Server, ln net.Listener) {
|
||||
if err := s.Serve(ln); err != nil &&
|
||||
!errors.Is(err, http.ErrServerClosed) && !errors.Is(err, net.ErrClosed) {
|
||||
m.log.Fatal("listen failed", zap.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
// drain 让旧服务器在后台优雅退出(等待进行中的连接完成或在超时后强制关闭)。
|
||||
func (m *serverManager) drain(s *http.Server) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
go func(s *http.Server) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
||||
defer cancel()
|
||||
if err := s.Shutdown(ctx); err != nil && !errors.Is(err, context.DeadlineExceeded) {
|
||||
m.log.Warn("drain old server failed", zap.Error(err))
|
||||
}
|
||||
}(s)
|
||||
}
|
||||
|
||||
func (m *serverManager) logServerReady() {
|
||||
scheme := "http"
|
||||
if m.pair != nil {
|
||||
scheme = "https"
|
||||
}
|
||||
localIP := getLocalIP()
|
||||
m.log.Info("server is ready",
|
||||
zap.String("scheme", scheme),
|
||||
zap.String("local", fmt.Sprintf("%s://%s:%d", scheme, localIP, m.cfg.App.Port)),
|
||||
zap.String("listen", m.addr),
|
||||
)
|
||||
if m.pair != nil {
|
||||
m.log.Info("HTTPS is enabled; plain HTTP is no longer served on this port",
|
||||
zap.String("addr", m.addr),
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/pem"
|
||||
"math/big"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/config"
|
||||
)
|
||||
|
||||
func makeTestPairPEM(t *testing.T) (certPEM, keyPEM string) {
|
||||
t.Helper()
|
||||
priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tpl := &x509.Certificate{
|
||||
SerialNumber: big.NewInt(1),
|
||||
Subject: pkix.Name{CommonName: "localhost"},
|
||||
NotBefore: time.Now().Add(-time.Hour),
|
||||
NotAfter: time.Now().Add(24 * time.Hour),
|
||||
DNSNames: []string{"localhost"},
|
||||
}
|
||||
der, err := x509.CreateCertificate(rand.Reader, tpl, tpl, &priv.PublicKey, priv)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
keyDER, err := x509.MarshalECPrivateKey(priv)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
certPEM = strings.TrimSpace(string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})))
|
||||
keyPEM = strings.TrimSpace(string(pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER})))
|
||||
return certPEM, keyPEM
|
||||
}
|
||||
|
||||
func newTestServerManager(t *testing.T) *serverManager {
|
||||
t.Helper()
|
||||
cfg := &config.Config{}
|
||||
cfg.App.Port = 18081
|
||||
return newServerManager(cfg, zap.NewNop(), http.NewServeMux())
|
||||
}
|
||||
|
||||
func TestDesiredPairModes(t *testing.T) {
|
||||
m := newTestServerManager(t)
|
||||
|
||||
if p, err := m.desiredPair(); err != nil || p != nil {
|
||||
t.Fatalf("disabled should be nil pair, got p=%v err=%v", p, err)
|
||||
}
|
||||
|
||||
certPEM, keyPEM := makeTestPairPEM(t)
|
||||
m.cfg.App.HTTPSEnabled = true
|
||||
m.cfg.App.SSLCert, m.cfg.App.SSLKey = certPEM, keyPEM
|
||||
p, err := m.desiredPair()
|
||||
if err != nil || p == nil || p.version == "" {
|
||||
t.Fatalf("content pair failed: p=%v err=%v", p, err)
|
||||
}
|
||||
|
||||
dir := t.TempDir()
|
||||
certPath, keyPath := filepath.Join(dir, "cert.pem"), filepath.Join(dir, "key.pem")
|
||||
if err := os.WriteFile(certPath, []byte(certPEM), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(keyPath, []byte(keyPEM), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
m.cfg.App.SSLCert, m.cfg.App.SSLKey = "", ""
|
||||
m.cfg.App.SSLCertPath, m.cfg.App.SSLKeyPath = certPath, keyPath
|
||||
p2, err := m.desiredPair()
|
||||
if err != nil || p2 == nil {
|
||||
t.Fatalf("path pair failed: %v", err)
|
||||
}
|
||||
|
||||
m.cfg.App.SSLKeyPath = filepath.Join(dir, "nope.pem")
|
||||
if _, err := m.desiredPair(); err == nil {
|
||||
t.Fatal("expected error when key file missing")
|
||||
}
|
||||
m.cfg.App.SSLKeyPath = keyPath
|
||||
|
||||
// 替换文件(换一套新的有效证书)后版本号应变化,触发热更新。
|
||||
newCert, newKey := makeTestPairPEM(t)
|
||||
if err := os.WriteFile(certPath, []byte(newCert), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(keyPath, []byte(newKey), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
p3, err := m.desiredPair()
|
||||
if err != nil {
|
||||
t.Fatalf("replace: %v", err)
|
||||
}
|
||||
if p3.version == p2.version {
|
||||
t.Fatal("version should change after files replaced")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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="
|
||||
)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -49,6 +49,17 @@ type AppConfig struct {
|
||||
Env string `mapstructure:"env"`
|
||||
DataDir string `mapstructure:"data_dir"`
|
||||
WebDir string `mapstructure:"web_dir"`
|
||||
// HTTPSEnabled 是否仅通过 HTTPS 提供访问。启用时必须同时配置
|
||||
// SSLCert / SSLKey(或 SSLCertPath / SSLKeyPath),保存后服务会热切换到 HTTPS。
|
||||
HTTPSEnabled bool `mapstructure:"https_enabled"`
|
||||
// SSLCert 是 PEM 编码的 SSL 证书内容。
|
||||
SSLCert string `mapstructure:"ssl_cert"`
|
||||
// SSLKey 是 PEM 编码的 SSL 私钥内容。
|
||||
SSLKey string `mapstructure:"ssl_key"`
|
||||
// SSLCertPath 是 SSL 证书文件路径;非空时优先于 SSLCert 从文件读取。
|
||||
SSLCertPath string `mapstructure:"ssl_cert_path"`
|
||||
// SSLKeyPath 是 SSL 私钥文件路径;非空时优先于 SSLKey 从文件读取。
|
||||
SSLKeyPath string `mapstructure:"ssl_key_path"`
|
||||
FFmpegPath string `mapstructure:"ffmpeg_path"`
|
||||
FFprobePath string `mapstructure:"ffprobe_path"`
|
||||
// FFprobeMaxConcurrent limits concurrent ffprobe/ffmpeg metadata probes.
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"])
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
"github.com/ShukeBta/MMTL/internal/service"
|
||||
@@ -50,6 +51,10 @@ func updateSettingHandler(svc *service.Container) gin.HandlerFunc {
|
||||
_ = svc.Repo.DB.WithContext(c.Request.Context()).Model(&model.User{}).Where("hide_adult = ?", false).Update("hide_adult", true).Error
|
||||
}
|
||||
service.ApplyRuntimeSetting(svc.Cfg, req.Key, req.Value)
|
||||
if err := applyHTTPSetting(svc, req.Key, req.Value); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if svc.FFprobe != nil && (req.Key == "ffprobe.max_concurrent" || req.Key == "app.ffprobe_max_concurrent") {
|
||||
svc.FFprobe.SetMaxConcurrent(svc.Cfg.App.FFprobeMaxConcurrent)
|
||||
}
|
||||
@@ -63,6 +68,83 @@ func updateSettingHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// applyHTTPSetting 校验 HTTPS 相关设置,并在可行时热重载监听。
|
||||
// 必须在 ApplyRuntimeSetting 之后调用,这样 svc.Cfg 已反映刚保存的值。
|
||||
//
|
||||
// 规则:
|
||||
// - https.enabled=true 时强制要求证书与私钥都已配置(内容或路径均可)且匹配,
|
||||
// 否则返回错误("如果启用就必须配置 SSL 证书和密钥");
|
||||
// - 证书/私钥(内容或路径)单独保存时只校验格式;若 HTTPS 已开启且新的整体
|
||||
// 配置可解析匹配才触发重载,避免"只存了新证书、私钥还没保存"时用旧私钥带
|
||||
// 新证书对外提供服务。
|
||||
func applyHTTPSetting(svc *service.Container, key, value string) error {
|
||||
skipReload := func(reason string) {
|
||||
if svc.Log != nil {
|
||||
svc.Log.Warn("https setting saved but not applied yet", zap.String("key", key), zap.String("reason", reason))
|
||||
}
|
||||
}
|
||||
switch key {
|
||||
case "https.enabled":
|
||||
if svc.Cfg.App.HTTPSEnabled {
|
||||
if _, err := service.ResolveSSLKeyPair(svc.Cfg.App.SSLCert, svc.Cfg.App.SSLCertPath, svc.Cfg.App.SSLKey, svc.Cfg.App.SSLKeyPath); err != nil {
|
||||
return fmt.Errorf("启用 HTTPS 失败:%v", err)
|
||||
}
|
||||
}
|
||||
case "https.cert", "https.cert_path", "https.key", "https.key_path":
|
||||
if err := validateSSLMaterialSource(key, value); err != nil {
|
||||
return err
|
||||
}
|
||||
if !svc.Cfg.App.HTTPSEnabled {
|
||||
return nil
|
||||
}
|
||||
if !httpsPairReady(svc) {
|
||||
skipReload("证书与私钥尚未匹配,等待另一半保存后生效")
|
||||
return nil
|
||||
}
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
if svc.ReloadHTTPServer != nil {
|
||||
return svc.ReloadHTTPServer()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateSSLMaterialSource 校验刚保存的证书/私钥来源(内容或路径)本身格式合法。
|
||||
func validateSSLMaterialSource(key, value string) error {
|
||||
switch key {
|
||||
case "https.cert":
|
||||
return service.ValidateSSLCert(value)
|
||||
case "https.cert_path":
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return nil // 清空路径也允许,启用时由整体校验把关
|
||||
}
|
||||
pemStr, err := service.ResolveSSLMaterial("", value, "证书")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return service.ValidateSSLCert(pemStr)
|
||||
case "https.key":
|
||||
return service.ValidateSSLKey(value)
|
||||
case "https.key_path":
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return nil
|
||||
}
|
||||
pemStr, err := service.ResolveSSLMaterial("", value, "私钥")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return service.ValidateSSLKey(pemStr)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// httpsPairReady 判断基于当前配置解析出的证书/私钥是否完整且匹配。
|
||||
func httpsPairReady(svc *service.Container) bool {
|
||||
_, err := service.ResolveSSLKeyPair(svc.Cfg.App.SSLCert, svc.Cfg.App.SSLCertPath, svc.Cfg.App.SSLKey, svc.Cfg.App.SSLKeyPath)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
type testAdultScraperReq struct {
|
||||
Engine string `json:"engine"`
|
||||
ServerURL string `json:"server_url"`
|
||||
|
||||
+73
-14
@@ -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")
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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})
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 是逻辑媒体库下的一条真实物理/挂载路径。
|
||||
|
||||
@@ -54,7 +54,9 @@ func AllModels() []interface{} {
|
||||
&StrmAccount{},
|
||||
&StrmSyncPath{},
|
||||
&StrmSyncRecord{},
|
||||
&StrmDownloadTask{},
|
||||
&StrmUploadTask{},
|
||||
}
|
||||
&StrmDownloadTask{},
|
||||
&StrmUploadTask{},
|
||||
&StrmDirCache{},
|
||||
&ScrapeTask{},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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"` // 相对根目录的路径
|
||||
}
|
||||
|
||||
|
||||
@@ -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 更新测试结果。
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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).
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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:"-"` // 原始响应体(外层附加字段用)
|
||||
}
|
||||
@@ -280,7 +281,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,57 @@ 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, "")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -92,6 +92,10 @@ func TestDanmakuFetchHashMatchLayer(t *testing.T) {
|
||||
require.Equal(t, "xml", res.SourceType)
|
||||
require.Contains(t, res.Raw, "弹幕Hash命中")
|
||||
require.Empty(t, res.Candidates)
|
||||
require.Equal(t, "测试动画", res.AnimeTitle)
|
||||
require.Equal(t, "第1话", res.EpisodeTitle)
|
||||
require.Equal(t, int64(25484), res.EpisodeID)
|
||||
require.Equal(t, "hash", res.MatchMode)
|
||||
|
||||
// match 请求体:文件名去扩展名并 URL 转义(官方接口要求,实测验证)、
|
||||
// hash、大小、matchMode 齐全。
|
||||
@@ -274,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, "手动搜索弹幕")
|
||||
}
|
||||
@@ -67,11 +67,17 @@ type DanmakuRenderConfig struct {
|
||||
//
|
||||
// Candidates is non-nil when multiple anime matched the search and the player
|
||||
// must ask the user which one to use (disambiguation); Raw is empty then.
|
||||
// AnimeTitle, EpisodeTitle, EpisodeID and MatchMode provide matched danmaku
|
||||
// metadata so the player UI can display which episode was loaded.
|
||||
type DanmakuFetchResult struct {
|
||||
DanmakuRenderConfig
|
||||
SourceType string `json:"source_type"`
|
||||
Raw string `json:"raw,omitempty"`
|
||||
Candidates []DanmakuAnime `json:"candidates,omitempty"`
|
||||
SourceType string `json:"source_type"`
|
||||
Raw string `json:"raw,omitempty"`
|
||||
Candidates []DanmakuAnime `json:"candidates,omitempty"`
|
||||
AnimeTitle string `json:"anime_title,omitempty"`
|
||||
EpisodeTitle string `json:"episode_title,omitempty"`
|
||||
EpisodeID int64 `json:"episode_id,omitempty"`
|
||||
MatchMode string `json:"match_mode,omitempty"`
|
||||
}
|
||||
|
||||
// DanmakuAnime is one search hit (an anime) with its episode list, mirroring
|
||||
@@ -184,74 +190,90 @@ func (s *DanmakuService) Fetch(ctx context.Context, mediaID, keyword, episodeID
|
||||
configured := strings.TrimRight(strings.TrimSpace(res.Source), "/")
|
||||
official := danmakuOfficialBase
|
||||
|
||||
// 手动指定弹幕库:跳过识别,直接拉取该库(自定义源失败回退官方)。
|
||||
if target := strings.TrimSpace(episodeID); target != "" {
|
||||
raw, st, err := s.fetchCommentWithFallback(ctx, configured, official, target)
|
||||
if err != nil {
|
||||
s.log.Warn("danmaku comment fetch failed", zap.String("media_id", mediaID), zap.String("episode_id", target), zap.Error(err))
|
||||
return res, err
|
||||
}
|
||||
res.Raw, res.SourceType = raw, st
|
||||
return res, nil
|
||||
}
|
||||
|
||||
term, media, err := s.searchTerms(ctx, mediaID)
|
||||
if err != nil {
|
||||
return res, err
|
||||
}
|
||||
manualKeyword := strings.TrimSpace(keyword) != ""
|
||||
if kw := strings.TrimSpace(keyword); kw != "" {
|
||||
term.name = kw
|
||||
}
|
||||
if strings.TrimSpace(term.name) == "" {
|
||||
return res, nil
|
||||
}
|
||||
|
||||
target := ""
|
||||
|
||||
// 1) hash 识别:始终走官方 /api/v2/match。
|
||||
if media != nil && media.Path != "" {
|
||||
if hash, ok := s.mediaHash(ctx, media); ok {
|
||||
fileSize := media.SizeBytes
|
||||
if strings.EqualFold(filepath.Ext(media.Path), ".strm") {
|
||||
fileSize = 0 // strm 行的 SizeBytes 是文本大小,不是视频大小
|
||||
}
|
||||
matches, err := s.matchOfficial(ctx, danmakuMatchFileName(media.Path), hash, fileSize, media.DurationSec)
|
||||
// 手动指定弹幕库:跳过识别,直接拉取该库(自定义源失败回退官方)。
|
||||
if target := strings.TrimSpace(episodeID); target != "" {
|
||||
raw, st, err := s.fetchCommentWithFallback(ctx, configured, official, target)
|
||||
if err != nil {
|
||||
s.log.Warn("danmaku hash match failed", zap.String("media_id", mediaID), zap.Error(err))
|
||||
} else if len(matches) > 0 {
|
||||
target = fmt.Sprintf("%d", matches[0].EpisodeID)
|
||||
s.log.Warn("danmaku comment fetch failed", zap.String("media_id", mediaID), zap.String("episode_id", target), zap.Error(err))
|
||||
return res, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 2) 按播放的文件名 + 集数搜索(keyword 手动覆盖时跳过,直接走第 3 层)。
|
||||
if target == "" && !manualKeyword && media != nil && media.Path != "" {
|
||||
if fileName := danmakuMatchFileName(media.Path); fileName != "" && fileName != term.name {
|
||||
if candidates, err := s.searchCandidatesWithFallback(ctx, configured, official, fileName, term.episode); err == nil &&
|
||||
len(candidates) == 1 && len(candidates[0].Episodes) > 0 {
|
||||
target = fmt.Sprintf("%d", candidates[0].Episodes[0].EpisodeID)
|
||||
res.Raw, res.SourceType = raw, st
|
||||
if id, parseErr := strconv.ParseInt(target, 10, 64); parseErr == nil {
|
||||
res.EpisodeID = id
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 3) 现有自动识别:标题层级(original_name → title → 文件名)+ 集数,
|
||||
// 多结果返回候选列表交给播放器(歧义处理)。
|
||||
if target == "" {
|
||||
candidates, err := s.searchCandidatesWithFallback(ctx, configured, official, term.name, term.episode)
|
||||
if err != nil {
|
||||
s.log.Warn("danmaku search failed", zap.String("media_id", mediaID), zap.String("name", term.name), zap.String("episode", term.episode), zap.Error(err))
|
||||
return res, err
|
||||
}
|
||||
if len(candidates) != 1 {
|
||||
res.Candidates = candidates
|
||||
res.MatchMode = "manual"
|
||||
return res, nil
|
||||
}
|
||||
if len(candidates[0].Episodes) == 0 {
|
||||
return res, errors.New("no danmaku library found for this video")
|
||||
|
||||
term, media, err := s.searchTerms(ctx, mediaID)
|
||||
if err != nil {
|
||||
return res, err
|
||||
}
|
||||
manualKeyword := strings.TrimSpace(keyword) != ""
|
||||
if kw := strings.TrimSpace(keyword); kw != "" {
|
||||
term.name = kw
|
||||
}
|
||||
if strings.TrimSpace(term.name) == "" {
|
||||
return res, nil
|
||||
}
|
||||
|
||||
target := ""
|
||||
|
||||
// 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") {
|
||||
fileSize = 0 // strm 行的 SizeBytes 是文本大小,不是视频大小
|
||||
}
|
||||
matches, err := s.matchOfficial(ctx, danmakuMatchFileName(media.Path), hash, fileSize, media.DurationSec)
|
||||
if err != nil {
|
||||
s.log.Warn("danmaku hash match failed", zap.String("media_id", mediaID), zap.Error(err))
|
||||
} else if len(matches) > 0 {
|
||||
target = fmt.Sprintf("%d", matches[0].EpisodeID)
|
||||
res.AnimeTitle = matches[0].AnimeTitle
|
||||
res.EpisodeTitle = matches[0].EpisodeTitle
|
||||
res.EpisodeID = matches[0].EpisodeID
|
||||
res.MatchMode = "hash"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 2) 按播放的文件名 + 集数搜索(keyword 手动覆盖时跳过,直接走第 3 层)。
|
||||
if target == "" && !manualKeyword && media != nil && media.Path != "" {
|
||||
if fileName := danmakuMatchFileName(media.Path); fileName != "" && fileName != term.name {
|
||||
if candidates, err := s.searchCandidatesWithFallback(ctx, configured, official, fileName, term.episode); err == nil &&
|
||||
len(candidates) == 1 && len(candidates[0].Episodes) > 0 {
|
||||
target = fmt.Sprintf("%d", candidates[0].Episodes[0].EpisodeID)
|
||||
res.AnimeTitle = candidates[0].AnimeTitle
|
||||
res.EpisodeTitle = candidates[0].Episodes[0].EpisodeTitle
|
||||
res.EpisodeID = candidates[0].Episodes[0].EpisodeID
|
||||
res.MatchMode = "filename"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 3) 现有自动识别:标题层级(original_name → title → 文件名)+ 集数,
|
||||
// 多结果返回候选列表交给播放器(歧义处理)。
|
||||
if target == "" {
|
||||
candidates, err := s.searchCandidatesWithFallback(ctx, configured, official, term.name, term.episode)
|
||||
if err != nil {
|
||||
s.log.Warn("danmaku search failed", zap.String("media_id", mediaID), zap.String("name", term.name), zap.String("episode", term.episode), zap.Error(err))
|
||||
return res, err
|
||||
}
|
||||
if len(candidates) != 1 {
|
||||
res.Candidates = candidates
|
||||
return res, nil
|
||||
}
|
||||
if len(candidates[0].Episodes) == 0 {
|
||||
return res, errors.New("no danmaku library found for this video")
|
||||
}
|
||||
target = fmt.Sprintf("%d", candidates[0].Episodes[0].EpisodeID)
|
||||
res.AnimeTitle = candidates[0].AnimeTitle
|
||||
res.EpisodeTitle = candidates[0].Episodes[0].EpisodeTitle
|
||||
res.EpisodeID = candidates[0].Episodes[0].EpisodeID
|
||||
res.MatchMode = "search"
|
||||
}
|
||||
target = fmt.Sprintf("%d", candidates[0].Episodes[0].EpisodeID)
|
||||
}
|
||||
|
||||
raw, st, err := s.fetchCommentWithFallback(ctx, configured, official, target)
|
||||
if err != nil {
|
||||
|
||||
@@ -142,6 +142,10 @@ func TestDanmakuFetchWithDandanplaySource(t *testing.T) {
|
||||
require.Equal(t, "xml", res.SourceType)
|
||||
require.Contains(t, res.Raw, "弹幕A")
|
||||
require.Contains(t, res.Raw, `p="0.5,1,16777215,user1"`)
|
||||
require.Equal(t, "测试动画", res.AnimeTitle)
|
||||
require.Equal(t, "第1话", res.EpisodeTitle)
|
||||
require.Equal(t, int64(25484), res.EpisodeID)
|
||||
require.Equal(t, "search", res.MatchMode)
|
||||
}
|
||||
|
||||
func TestDanmakuFetchUsesOriginalNameForSearch(t *testing.T) {
|
||||
@@ -262,6 +266,8 @@ func TestDanmakuFetchWithExplicitEpisodeID(t *testing.T) {
|
||||
require.True(t, res.Enabled)
|
||||
require.Contains(t, res.Raw, "显式指定弹幕")
|
||||
require.Empty(t, res.Candidates)
|
||||
require.Equal(t, int64(99999), res.EpisodeID)
|
||||
require.Equal(t, "manual", res.MatchMode)
|
||||
}
|
||||
func TestDetectDanmakuSourceType(t *testing.T) {
|
||||
cases := []struct {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/ed25519"
|
||||
"crypto/rsa"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ResolveSSLMaterial 解析一份 SSL 材料(证书或私钥)的 PEM 内容:
|
||||
// 优先读取 path 指向的文件,其次使用内容;两者都为空时返回错误。
|
||||
// what 用于错误提示("证书" / "私钥")。
|
||||
func ResolveSSLMaterial(content, path, what string) (string, error) {
|
||||
p := strings.TrimSpace(path)
|
||||
if p != "" {
|
||||
if info, err := os.Stat(p); err != nil {
|
||||
return "", fmt.Errorf("SSL %s文件不可访问:%s(%v)", what, p, err)
|
||||
} else if info.IsDir() {
|
||||
return "", fmt.Errorf("SSL %s路径指向的是目录,请填写文件路径:%s", what, p)
|
||||
}
|
||||
b, err := os.ReadFile(p)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("读取 SSL %s文件失败:%s(%v)", what, p, err)
|
||||
}
|
||||
return strings.TrimSpace(string(b)), nil
|
||||
}
|
||||
c := strings.TrimSpace(content)
|
||||
if c == "" {
|
||||
return "", fmt.Errorf("SSL %s未配置:请填写内容或文件路径", what)
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
|
||||
// ResolveSSLKeyPair 解析证书与私钥(各自支持 内容或路径),校验格式与匹配后
|
||||
// 返回可用的 tls.Certificate。任何一步失败都会给出明确错误。
|
||||
func ResolveSSLKeyPair(certContent, certPath, keyContent, keyPath string) (*tls.Certificate, error) {
|
||||
certPEM, err := ResolveSSLMaterial(certContent, certPath, "证书")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
keyPEM, err := ResolveSSLMaterial(keyContent, keyPath, "私钥")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cert, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSL 证书/私钥无效:%v", err)
|
||||
}
|
||||
if !sslKeyMatchesCert(cert) {
|
||||
return nil, errors.New("SSL 证书与私钥不匹配")
|
||||
}
|
||||
return &cert, nil
|
||||
}
|
||||
|
||||
// PathFingerprint 返回文件路径的内容指纹(路径 + 大小 + 修改时间),用于检测
|
||||
// 文件是否被替换过;文件不存在时返回 ("", false)。
|
||||
func PathFingerprint(p string) (string, bool) {
|
||||
p = strings.TrimSpace(p)
|
||||
if p == "" {
|
||||
return "", false
|
||||
}
|
||||
info, err := os.Stat(p)
|
||||
if err != nil {
|
||||
return "", false
|
||||
}
|
||||
return fmt.Sprintf("path:%s|size:%d|mtime:%d", filepath.Clean(p), info.Size(), info.ModTime().UnixNano()), true
|
||||
}
|
||||
|
||||
// ValidateSSLCert 校验 s 是一个可解析的 PEM 编码 X.509 证书。
|
||||
func ValidateSSLCert(s string) error {
|
||||
block, _ := pem.Decode([]byte(strings.TrimSpace(s)))
|
||||
if block == nil {
|
||||
return errors.New("SSL 证书格式无效:未找到 PEM 数据")
|
||||
}
|
||||
if block.Type != "CERTIFICATE" {
|
||||
return fmt.Errorf("SSL 证书格式无效:期望 CERTIFICATE,实际为 %s", block.Type)
|
||||
}
|
||||
if _, err := x509.ParseCertificate(block.Bytes); err != nil {
|
||||
return fmt.Errorf("SSL 证书解析失败:%v", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateSSLKey 校验 s 是一个可解析的 PEM 编码私钥。
|
||||
func ValidateSSLKey(s string) error {
|
||||
block, _ := pem.Decode([]byte(strings.TrimSpace(s)))
|
||||
if block == nil {
|
||||
return errors.New("SSL 私钥格式无效:未找到 PEM 数据")
|
||||
}
|
||||
if _, err := parsePrivateKeyBlock(block); err != nil {
|
||||
return fmt.Errorf("SSL 私钥解析失败:%v", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateSSLKeyPair 校验证书与私钥都存在、可解析且相互匹配。
|
||||
func ValidateSSLKeyPair(certPEM, keyPEM string) error {
|
||||
cert, err := tls.X509KeyPair([]byte(strings.TrimSpace(certPEM)), []byte(strings.TrimSpace(keyPEM)))
|
||||
if err != nil {
|
||||
return fmt.Errorf("SSL 证书/私钥无效:%v", err)
|
||||
}
|
||||
if !sslKeyMatchesCert(cert) {
|
||||
return errors.New("SSL 证书与私钥不匹配")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// sslKeyMatchesCert 通过公钥是否一致来判断私钥确实对应证书。
|
||||
func sslKeyMatchesCert(cert tls.Certificate) bool {
|
||||
if len(cert.Certificate) == 0 || cert.PrivateKey == nil {
|
||||
return false
|
||||
}
|
||||
leaf, err := x509.ParseCertificate(cert.Certificate[0])
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
privPub := publicKeyOf(cert.PrivateKey)
|
||||
if privPub == nil {
|
||||
return false
|
||||
}
|
||||
eq, ok := leaf.PublicKey.(interface {
|
||||
Equal(x crypto.PublicKey) bool
|
||||
})
|
||||
return ok && eq.Equal(privPub)
|
||||
}
|
||||
|
||||
// publicKeyOf 从各类私钥中提取对应的公钥。
|
||||
func publicKeyOf(priv crypto.PrivateKey) crypto.PublicKey {
|
||||
switch k := priv.(type) {
|
||||
case *rsa.PrivateKey:
|
||||
return &k.PublicKey
|
||||
case *ecdsa.PrivateKey:
|
||||
return &k.PublicKey
|
||||
case ed25519.PrivateKey:
|
||||
return k.Public()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// parsePrivateKeyBlock 支持 PKCS#8 / PKCS#1 RSA / EC 三种常见私钥格式。
|
||||
func parsePrivateKeyBlock(block *pem.Block) (crypto.PrivateKey, error) {
|
||||
if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
if key, err := x509.ParseECPrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
return nil, errors.New("无法解析私钥(支持 PKCS#8 / PKCS#1 RSA / EC)")
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/pem"
|
||||
"math/big"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func makeTestKeyPair(t *testing.T) (certPEM, keyPEM string) {
|
||||
t.Helper()
|
||||
priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tpl := &x509.Certificate{
|
||||
SerialNumber: big.NewInt(1),
|
||||
Subject: pkix.Name{CommonName: "localhost"},
|
||||
NotBefore: time.Now().Add(-time.Hour),
|
||||
NotAfter: time.Now().Add(24 * time.Hour),
|
||||
DNSNames: []string{"localhost"},
|
||||
}
|
||||
der, err := x509.CreateCertificate(rand.Reader, tpl, tpl, &priv.PublicKey, priv)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
keyDER, err := x509.MarshalECPrivateKey(priv)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
certPEM = strings.TrimSpace(string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})))
|
||||
keyPEM = strings.TrimSpace(string(pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER})))
|
||||
return certPEM, keyPEM
|
||||
}
|
||||
|
||||
func TestResolveSSLMaterial(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
certPEM, _ := makeTestKeyPair(t)
|
||||
path := filepath.Join(dir, "cert.pem")
|
||||
if err := os.WriteFile(path, []byte(certPEM), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
content string
|
||||
path string
|
||||
want string
|
||||
err bool
|
||||
}{
|
||||
{name: "content only", content: certPEM, want: certPEM},
|
||||
{name: "path only", path: path, want: certPEM},
|
||||
{name: "path wins over content", content: "bogus", path: path, want: certPEM},
|
||||
{name: "both empty", err: true},
|
||||
{name: "missing file", path: filepath.Join(dir, "missing.pem"), err: true},
|
||||
{name: "path is dir", path: dir, err: true},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := ResolveSSLMaterial(tc.content, tc.path, "证书")
|
||||
if tc.err {
|
||||
if err == nil {
|
||||
t.Fatalf("expected error, got %q", got)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if got != tc.want {
|
||||
t.Fatalf("got %q want %q", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveSSLKeyPair(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
certPEM, keyPEM := makeTestKeyPair(t)
|
||||
certPath := filepath.Join(dir, "cert.pem")
|
||||
keyPath := filepath.Join(dir, "key.pem")
|
||||
if err := os.WriteFile(certPath, []byte(certPEM), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(keyPath, []byte(keyPEM), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if _, err := ResolveSSLKeyPair(certPEM, "", keyPEM, ""); err != nil {
|
||||
t.Fatalf("content pair: %v", err)
|
||||
}
|
||||
if _, err := ResolveSSLKeyPair("", certPath, "", keyPath); err != nil {
|
||||
t.Fatalf("path pair: %v", err)
|
||||
}
|
||||
if _, err := ResolveSSLKeyPair(certPEM, "", "", keyPath); err != nil {
|
||||
t.Fatalf("mixed pair: %v", err)
|
||||
}
|
||||
|
||||
otherCert, _ := makeTestKeyPair(t)
|
||||
if _, err := ResolveSSLKeyPair(otherCert, "", keyPEM, ""); err == nil {
|
||||
t.Fatal("expected mismatch error")
|
||||
}
|
||||
|
||||
if got, ok := PathFingerprint(certPath); !ok || got == "" {
|
||||
t.Fatalf("PathFingerprint failed: got=%q ok=%v", got, ok)
|
||||
}
|
||||
if _, ok := PathFingerprint(filepath.Join(dir, "missing.pem")); ok {
|
||||
t.Fatal("PathFingerprint should report missing file")
|
||||
}
|
||||
if _, ok := PathFingerprint(" "); ok {
|
||||
t.Fatal("empty PathFingerprint should not be ok")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -100,6 +100,16 @@ func ApplyRuntimeSetting(cfg *config.Config, key, value string) {
|
||||
}
|
||||
case "transcode.video_bitrate", "transcoder.video_bitrate":
|
||||
cfg.Transcoder.VideoBitrate = value
|
||||
case "https.enabled":
|
||||
cfg.App.HTTPSEnabled = parseBoolSetting(value, false)
|
||||
case "https.cert":
|
||||
cfg.App.SSLCert = value
|
||||
case "https.key":
|
||||
cfg.App.SSLKey = value
|
||||
case "https.cert_path":
|
||||
cfg.App.SSLCertPath = strings.TrimSpace(value)
|
||||
case "https.key_path":
|
||||
cfg.App.SSLKeyPath = strings.TrimSpace(value)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -58,12 +58,17 @@ 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
|
||||
// 会调用它把 HTTP/HTTPS 监听热切换到最新配置;nil 表示未注入(测试环境)。
|
||||
ReloadHTTPServer func() error
|
||||
}
|
||||
|
||||
// New 构建服务容器。
|
||||
@@ -98,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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 路径。
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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),
|
||||
}
|
||||
|
||||
@@ -28,6 +28,19 @@ export interface DanmakuFetchResult {
|
||||
area: string
|
||||
raw?: string
|
||||
candidates?: DanmakuAnime[]
|
||||
anime_title?: string
|
||||
episode_title?: string
|
||||
episode_id?: number
|
||||
match_mode?: 'hash' | 'filename' | 'search' | 'manual' | string
|
||||
}
|
||||
|
||||
export interface DanmakuLoadedInfo {
|
||||
animeTitle?: string
|
||||
episodeTitle?: string
|
||||
episodeId?: number | string
|
||||
matchMode?: 'hash' | 'filename' | 'search' | 'manual' | string
|
||||
totalCount: number
|
||||
sourceType?: 'auto' | 'xml' | 'json'
|
||||
}
|
||||
|
||||
export type DanmakuFetchOptions = {
|
||||
|
||||
+13
-2
@@ -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),
|
||||
|
||||
|
||||
@@ -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
@@ -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),
|
||||
}
|
||||
@@ -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 },
|
||||
]
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import { useEffect, useRef } from 'react'
|
||||
import { create, type Manager, type ManagerPlugin } from 'danmu'
|
||||
|
||||
import { danmakuAPI, type DanmakuAnime } from '../api/danmaku'
|
||||
import { danmakuAPI, type DanmakuAnime, type DanmakuLoadedInfo } from '../api/danmaku'
|
||||
import type { Media } from '../types'
|
||||
import { parseDanmaku, type Comment } from '../utils/parseDanmaku'
|
||||
|
||||
@@ -28,8 +28,8 @@ type DanmakuStageProps = {
|
||||
search?: string | null
|
||||
/** Explicit danmaku library chosen by the user; null = auto-resolve. */
|
||||
episodeId?: number | string | null
|
||||
/** Called after each fetch attempt (success or error) finishes. */
|
||||
onLoaded?: () => void
|
||||
/** Called after each fetch attempt (success or error) finishes with metadata. */
|
||||
onLoaded?: (info: DanmakuLoadedInfo | null) => void
|
||||
/** Called when multiple anime matched and the user must pick one. */
|
||||
onCandidates?: (candidates: DanmakuAnime[]) => void
|
||||
}
|
||||
@@ -109,6 +109,14 @@ export function DanmakuStage({
|
||||
// 弹幕层不拦截播放器控制栏的点击。
|
||||
holder.style.pointerEvents = 'none'
|
||||
|
||||
// 监听 holder 尺寸变化(全屏/退出全屏/窗口缩放),实时重置弹道与容器边界
|
||||
const ro = new ResizeObserver(() => {
|
||||
if (!disposed && managerRef.current) {
|
||||
managerRef.current.format()
|
||||
}
|
||||
})
|
||||
ro.observe(holder)
|
||||
|
||||
const applyLiveSettings = () => {
|
||||
const { opacity: liveOpacity, area: liveArea } = liveRef.current
|
||||
manager.setOpacity(liveOpacity)
|
||||
@@ -119,6 +127,10 @@ export function DanmakuStage({
|
||||
applyLiveSettings()
|
||||
|
||||
const loadDanmaku = async () => {
|
||||
let loadedInfo: DanmakuLoadedInfo | null = null
|
||||
comments = []
|
||||
nextIndex = 0
|
||||
manager.clear()
|
||||
try {
|
||||
const res = await danmakuAPI.fetch(media.id, {
|
||||
kw: search ?? undefined,
|
||||
@@ -136,14 +148,28 @@ export function DanmakuStage({
|
||||
.filter((c) => Number.isFinite(c.time) && c.time >= 0)
|
||||
.sort((a, b) => a.time - b.time)
|
||||
nextIndex = 0
|
||||
loadedInfo = {
|
||||
animeTitle: res.anime_title,
|
||||
episodeTitle: res.episode_title,
|
||||
episodeId: res.episode_id ?? (episodeId ? Number(episodeId) || episodeId : undefined),
|
||||
matchMode: res.match_mode,
|
||||
totalCount: comments.length,
|
||||
sourceType: res.source_type,
|
||||
}
|
||||
} else {
|
||||
comments = []
|
||||
loadedInfo = {
|
||||
totalCount: 0,
|
||||
}
|
||||
}
|
||||
} catch {
|
||||
// 拉取失败时静默关闭弹幕,不打断播放。
|
||||
comments = []
|
||||
loadedInfo = {
|
||||
totalCount: 0,
|
||||
}
|
||||
} finally {
|
||||
if (!disposed) onLoaded?.()
|
||||
if (!disposed) onLoaded?.(loadedInfo)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -215,6 +241,7 @@ export function DanmakuStage({
|
||||
|
||||
return () => {
|
||||
disposed = true
|
||||
ro.disconnect()
|
||||
cancelAnimationFrame(raf)
|
||||
video.removeEventListener('play', onPlay)
|
||||
video.removeEventListener('playing', onPlay)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user