mirror of
https://github.com/truewhile/MeBox.git
synced 2026-09-28 11:16:37 +08:00
Compare commits
57 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| db52792ada | |||
| a8065680bb | |||
| 3304f09b1d | |||
| ba621cfc4a | |||
| deca7735a7 | |||
| b503fdee7a | |||
| 3fe37e050b | |||
| d3233a62c0 | |||
| 165eee7b36 | |||
| 78526afc9c | |||
| f534e0607a | |||
| db64a6c093 | |||
| 5c9e7fcaa6 | |||
| af67f4cd6e | |||
| 73de139d1f | |||
| e0cc481b96 | |||
| 7ac75ec69a | |||
| 0e73e43cbd | |||
| bd41ab3fb4 | |||
| 73b95a8e38 | |||
| 5a5555d28a | |||
| 4dc9cbe2e7 | |||
| 6af0e5fdcf | |||
| f0055b76fe | |||
| 0eb3f104f4 | |||
| c171b38155 | |||
| 89d7a6cbb2 | |||
| 85918d1196 | |||
| fa9307e2e1 | |||
| 5ab18c2725 | |||
| a4a2bde1a4 | |||
| 28113f5fdc | |||
| 42b8805e94 | |||
| 51a64f41b9 | |||
| a7bd9a942c | |||
| 97491a6175 | |||
| ea5cb3a130 | |||
| 921010926b | |||
| 7425c3d57b | |||
| 1b7d4eef46 | |||
| c30dab56a3 | |||
| e365250440 | |||
| 47d10e1f58 | |||
| e6473300a7 | |||
| 994f64f753 | |||
| 90064a5480 | |||
| 6e8eac9887 | |||
| d3051eaffe | |||
| 496a897782 | |||
| 22b7290ee1 | |||
| 14037d5dea | |||
| 152db3fb9f | |||
| 0413d123da | |||
| 5f6bd7b5cd | |||
| 3c25bb5d61 | |||
| 659b91b000 | |||
| fe5b3bd56a |
@@ -4,9 +4,7 @@ name: AuTo Docker Image
|
||||
on:
|
||||
push:
|
||||
branches: [main]
|
||||
pull_request:
|
||||
branches: [main]
|
||||
|
||||
|
||||
# 保留手动触发作为备选
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
@@ -177,9 +175,32 @@ jobs:
|
||||
path: web/dist
|
||||
retention-days: 1
|
||||
|
||||
build-binaries:
|
||||
needs: [version-and-publish, build-frontend]
|
||||
# 先创建(幂等)空的 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:
|
||||
@@ -228,41 +249,18 @@ jobs:
|
||||
else
|
||||
tar -czf "mmtl_${{ matrix.goos }}_${{ matrix.goarch }}.tar.gz" -C package mmtl
|
||||
fi
|
||||
- name: Upload package
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: mmtl-package-${{ matrix.goos }}-${{ matrix.goarch }}
|
||||
path: |
|
||||
mmtl_${{ matrix.goos }}_${{ matrix.goarch }}.zip
|
||||
mmtl_${{ matrix.goos }}_${{ matrix.goarch }}.tar.gz
|
||||
if-no-files-found: ignore
|
||||
retention-days: 1
|
||||
|
||||
publish-release:
|
||||
needs: [version-and-publish, build-binaries]
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- name: Download all artifacts
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
path: release-assets
|
||||
merge-multiple: true
|
||||
- name: Create / update GitHub Release
|
||||
- name: Upload to GitHub Release
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
RELEASE_TAG: ${{ needs.version-and-publish.outputs.tag }}
|
||||
run: |
|
||||
set -eux
|
||||
echo "tag=$RELEASE_TAG"
|
||||
# 若 release 不存在则创建(tag 已由 version-and-publish 推送)
|
||||
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
|
||||
# 上传所有平台产物(已存在的同名 asset 会直接覆盖)
|
||||
for f in release-assets/mmtl_*.zip release-assets/mmtl_*.tar.gz; do
|
||||
[ -e "$f" ] && gh release upload "$RELEASE_TAG" "$f" --clobber || true
|
||||
done
|
||||
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
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
# Beta 分支自动构建流水线
|
||||
#
|
||||
# 触发:push 到 beta 分支 / PR 到 beta / 手动触发。
|
||||
# 产出:
|
||||
# 1. 前端 + 后端编译验证(go vet / go test / go build)
|
||||
# 2. 多平台可执行二进制 artifact(linux/amd64、linux/arm64、windows/amd64)
|
||||
# 3. ghcr.io/{owner}/mmtl:beta 多架构 Docker 镜像(linux/amd64 + linux/arm64)
|
||||
#
|
||||
# 与 main 分支的发布流(Auto-docker-publish.yml)隔离:beta 不做版本递增、
|
||||
# 不打 release tag,只构建带 -beta 标识的产物供测试。
|
||||
|
||||
name: Beta Build
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [beta]
|
||||
pull_request:
|
||||
branches: [beta]
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
env:
|
||||
BETA_VERSION_PREFIX: beta
|
||||
|
||||
jobs:
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# 1) 编译验证 + 多平台二进制产物
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
test-and-build:
|
||||
name: Test & build artifacts
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Resolve beta version
|
||||
id: version
|
||||
run: |
|
||||
BASE_VERSION=$(cat VERSION 2>/dev/null || echo "0.0.0")
|
||||
SHA_SHORT=${GITHUB_SHA:0:7}
|
||||
echo "full_version=${BASE_VERSION}-beta.${SHA_SHORT}" >> "$GITHUB_OUTPUT"
|
||||
|
||||
# The binary embeds the SPA (web/dist) via go:embed, so 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
|
||||
|
||||
- uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version: '1.25'
|
||||
cache: true
|
||||
|
||||
- name: go vet
|
||||
run: go vet ./...
|
||||
|
||||
- name: go test
|
||||
run: go test ./...
|
||||
|
||||
- name: go build (host)
|
||||
run: go build ./...
|
||||
|
||||
# 多平台可执行文件(嵌入刚构建的 web/dist)
|
||||
- name: Build linux/amd64
|
||||
run: CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -trimpath -ldflags="-s -w -X main.version=${{ steps.version.outputs.full_version }}" -o dist/mmtl-beta-linux-amd64 ./cmd/server
|
||||
- name: Build linux/arm64
|
||||
run: CGO_ENABLED=0 GOOS=linux GOARCH=arm64 go build -trimpath -ldflags="-s -w -X main.version=${{ steps.version.outputs.full_version }}" -o dist/mmtl-beta-linux-arm64 ./cmd/server
|
||||
- name: Build windows/amd64
|
||||
run: CGO_ENABLED=0 GOOS=windows GOARCH=amd64 go build -trimpath -ldflags="-s -w -X main.version=${{ steps.version.outputs.full_version }}" -o dist/mmtl-beta-windows-amd64.exe ./cmd/server
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: mmtl-beta-binaries
|
||||
path: dist/*
|
||||
if-no-files-found: error
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# 2) Beta Docker 镜像(ghcr.io/{owner}/mmtl:beta)
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
docker-beta:
|
||||
name: Build & push beta Docker image
|
||||
needs: test-and-build
|
||||
runs-on: ubuntu-latest
|
||||
# PR 事件不推送镜像,仅 push beta / 手动触发时推送
|
||||
if: github.event_name != 'pull_request'
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Resolve beta version
|
||||
id: version
|
||||
run: |
|
||||
BASE_VERSION=$(cat VERSION 2>/dev/null || echo "0.0.0")
|
||||
SHA_SHORT=${GITHUB_SHA:0:7}
|
||||
echo "full_version=${BASE_VERSION}-beta.${SHA_SHORT}" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- uses: docker/setup-qemu-action@v3
|
||||
- uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Log in to GHCR
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.actor }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Build & push
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
platforms: linux/amd64,linux/arm64
|
||||
push: true
|
||||
provenance: false
|
||||
sbom: false
|
||||
tags: ghcr.io/${{ github.repository_owner }}/mmtl:beta
|
||||
labels: |
|
||||
org.opencontainers.image.revision=${{ github.sha }}
|
||||
org.opencontainers.image.source=${{ github.repository }}
|
||||
build-args: |
|
||||
VERSION=${{ steps.version.outputs.full_version }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
@@ -173,6 +173,11 @@ func isFrontendLibraryRoute(path string) bool {
|
||||
if strings.Contains(id, "/") {
|
||||
return false
|
||||
}
|
||||
// 远程 Emby 挂载库的伪装 ID(embyremote~account~remote)也是前端库路由,
|
||||
// 需要交给 SPA 而非当作 Emby API 路径 404。
|
||||
if strings.HasPrefix(id, "embyremote~") {
|
||||
return true
|
||||
}
|
||||
if len(id) != 36 {
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -68,6 +68,7 @@ require (
|
||||
github.com/tklauser/numcpus v0.6.1 // indirect
|
||||
github.com/twitchyliquid64/golang-asm v0.15.1 // indirect
|
||||
github.com/ugorji/go/codec v1.2.11 // indirect
|
||||
github.com/ulikunitz/xz v0.5.12 // indirect
|
||||
github.com/yusufpapurcu/wmi v1.2.4 // indirect
|
||||
go.uber.org/multierr v1.10.0 // indirect
|
||||
golang.org/x/arch v0.3.0 // indirect
|
||||
|
||||
@@ -152,6 +152,8 @@ github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS
|
||||
github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08=
|
||||
github.com/ugorji/go/codec v1.2.11 h1:BMaWp1Bb6fHwEtbplGBGJ498wD+LKlNSl25MjdZY4dU=
|
||||
github.com/ugorji/go/codec v1.2.11/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZgYf6w6lg=
|
||||
github.com/ulikunitz/xz v0.5.12 h1:37Nm15o69RwBkXM0J6A5OlE67RZTfzUxTj8fB3dfcsc=
|
||||
github.com/ulikunitz/xz v0.5.12/go.mod h1:nbz6k7qbPmH4IRqmfOplQw/tblSgqTqBwxkY0oWt/14=
|
||||
github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0=
|
||||
github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
|
||||
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 830 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 1.3 MiB |
@@ -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="
|
||||
)
|
||||
@@ -29,7 +29,7 @@ func setDefaults(v *viper.Viper) {
|
||||
v.SetDefault("database.dsn", "")
|
||||
v.SetDefault("database.wal_mode", true)
|
||||
v.SetDefault("database.busy_timeout", 5000)
|
||||
v.SetDefault("database.cache_size", -20000)
|
||||
v.SetDefault("database.cache_size", -40000)
|
||||
v.SetDefault("database.max_open_conns", defaultDatabaseMaxOpenConns)
|
||||
v.SetDefault("database.max_idle_conns", defaultDatabaseMaxIdleConns)
|
||||
|
||||
@@ -43,6 +43,7 @@ func setDefaults(v *viper.Viper) {
|
||||
v.SetDefault("logging.max_backups", 10)
|
||||
|
||||
v.SetDefault("cache.cache_dir", "./cache")
|
||||
v.SetDefault("cache.images_max_size_mb", 500)
|
||||
v.SetDefault("cache.cleanup_interval_min", 60)
|
||||
v.SetDefault("cache.redis_url", "")
|
||||
v.SetDefault("cache.redis_prefix", "mmtl")
|
||||
|
||||
@@ -44,6 +44,9 @@ func (c *Config) normalize() error {
|
||||
if c.Cache.CacheDir == "" {
|
||||
c.Cache.CacheDir = filepath.Join(c.App.DataDir, "cache")
|
||||
}
|
||||
if c.Cache.ImagesMaxSizeMB < 0 {
|
||||
c.Cache.ImagesMaxSizeMB = 0
|
||||
}
|
||||
if c.Cache.RedisPrefix == "" {
|
||||
c.Cache.RedisPrefix = "mmtl"
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -44,11 +44,11 @@ type TranscoderConfig struct {
|
||||
|
||||
// AppConfig 保存运行时应用参数。
|
||||
type AppConfig struct {
|
||||
Port int `mapstructure:"port"`
|
||||
Debug bool `mapstructure:"debug"`
|
||||
Env string `mapstructure:"env"`
|
||||
DataDir string `mapstructure:"data_dir"`
|
||||
WebDir string `mapstructure:"web_dir"`
|
||||
Port int `mapstructure:"port"`
|
||||
Debug bool `mapstructure:"debug"`
|
||||
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"`
|
||||
@@ -59,7 +59,7 @@ type AppConfig struct {
|
||||
// SSLCertPath 是 SSL 证书文件路径;非空时优先于 SSLCert 从文件读取。
|
||||
SSLCertPath string `mapstructure:"ssl_cert_path"`
|
||||
// SSLKeyPath 是 SSL 私钥文件路径;非空时优先于 SSLKey 从文件读取。
|
||||
SSLKeyPath string `mapstructure:"ssl_key_path"`
|
||||
SSLKeyPath string `mapstructure:"ssl_key_path"`
|
||||
FFmpegPath string `mapstructure:"ffmpeg_path"`
|
||||
FFprobePath string `mapstructure:"ffprobe_path"`
|
||||
// FFprobeMaxConcurrent limits concurrent ffprobe/ffmpeg metadata probes.
|
||||
@@ -116,6 +116,7 @@ type LoggingConfig struct {
|
||||
// CacheConfig 控制磁盘转码/刮削缓存。
|
||||
type CacheConfig struct {
|
||||
CacheDir string `mapstructure:"cache_dir"`
|
||||
ImagesMaxSizeMB int `mapstructure:"images_max_size_mb"`
|
||||
MaxDiskUsageMB int `mapstructure:"max_disk_usage_mb"`
|
||||
TTLHours int `mapstructure:"ttl_hours"`
|
||||
AutoCleanup bool `mapstructure:"auto_cleanup"`
|
||||
|
||||
@@ -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,7 +408,7 @@ func TestOpenSQLiteMigrationSourceUsesFallbackSourcePath(t *testing.T) {
|
||||
_ = sqlDB2.Close()
|
||||
}
|
||||
}()
|
||||
copied, err := copyModelTables(src2, dst, 2)
|
||||
_, copied, err := copyModelTables(src2, dst, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
@@ -20,12 +20,23 @@ func AutoMigrate(db *gorm.DB) error {
|
||||
if err := ensureLibraryRootsCompatibility(db); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ensureEmbyMountsCompatibility(db); err != nil {
|
||||
return err
|
||||
}
|
||||
if isSQLite(db) {
|
||||
return ensureMediaSearchIndex(db)
|
||||
if err := ensureMediaSearchIndex(db); err != nil {
|
||||
return err
|
||||
}
|
||||
return ensureSQLiteQueryOptimizer(db)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensureSQLiteQueryOptimizer(db *gorm.DB) error {
|
||||
// Refresh planner statistics so indexes on large media tables are used.
|
||||
return db.Exec("ANALYZE").Error
|
||||
}
|
||||
|
||||
func ensurePostgresColumnCompatibility(db *gorm.DB) error {
|
||||
if !isPostgres(db) {
|
||||
return nil
|
||||
@@ -77,3 +88,25 @@ func ensurePerformanceIndexes(db *gorm.DB) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensureEmbyMountsCompatibility(db *gorm.DB) error {
|
||||
if !db.Migrator().HasTable(&model.EmbyMount{}) {
|
||||
return nil
|
||||
}
|
||||
if !db.Migrator().HasColumn(&model.EmbyMount{}, "sort_order") {
|
||||
if err := db.Migrator().AddColumn(&model.EmbyMount{}, "sort_order"); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
// 针对已有数据:如果存在多个 sort_order=0/NULL 的记录,按创建时间顺序赋予稳定递增的序号
|
||||
var zeroCount int64
|
||||
if err := db.Model(&model.EmbyMount{}).Where("sort_order = 0 OR sort_order IS NULL").Count(&zeroCount).Error; err == nil && zeroCount > 1 {
|
||||
var mounts []model.EmbyMount
|
||||
if err := db.Order("created_at asc, id asc").Find(&mounts).Error; err == nil {
|
||||
for i, m := range mounts {
|
||||
_ = db.Exec("UPDATE emby_mounts SET sort_order = ? WHERE id = ?", i, m.ID).Error
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
)
|
||||
|
||||
func TestEnsureEmbyMountsCompatibility(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Create a table without sort_order simulating an older schema
|
||||
if err := db.Exec(`CREATE TABLE emby_mounts (
|
||||
id varchar(36) PRIMARY KEY,
|
||||
created_at datetime,
|
||||
updated_at datetime,
|
||||
deleted_at datetime,
|
||||
account_id text,
|
||||
remote_view_id text,
|
||||
remote_view_name text,
|
||||
collection_type text,
|
||||
name text,
|
||||
proxy_play numeric DEFAULT false,
|
||||
enabled numeric DEFAULT true
|
||||
)`).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Insert older rows
|
||||
now := time.Now()
|
||||
_ = db.Exec("INSERT INTO emby_mounts (id, name, created_at) VALUES (?, ?, ?)", "m1", "Mount 1", now.Add(-2*time.Hour)).Error
|
||||
_ = db.Exec("INSERT INTO emby_mounts (id, name, created_at) VALUES (?, ?, ?)", "m2", "Mount 2", now.Add(-1*time.Hour)).Error
|
||||
|
||||
// Run compatibility migration
|
||||
if err := ensureEmbyMountsCompatibility(db); err != nil {
|
||||
t.Fatalf("ensureEmbyMountsCompatibility failed: %v", err)
|
||||
}
|
||||
|
||||
// Verify column sort_order exists and values are initialized sequentially
|
||||
if !db.Migrator().HasColumn(&model.EmbyMount{}, "sort_order") {
|
||||
t.Fatal("expected sort_order column to be added")
|
||||
}
|
||||
|
||||
var m1, m2 model.EmbyMount
|
||||
if err := db.Where("id = ?", "m1").First(&m1).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.Where("id = ?", "m2").First(&m2).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if m1.SortOrder != 0 || m2.SortOrder != 1 {
|
||||
t.Fatalf("unexpected sort orders: m1=%d, m2=%d", m1.SortOrder, m2.SortOrder)
|
||||
}
|
||||
}
|
||||
@@ -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,10 @@ 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(536870912)"
|
||||
if cfg.Database.WALMode {
|
||||
dsn += "&_pragma=wal_autocheckpoint(1000)"
|
||||
}
|
||||
return dsn
|
||||
}
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@ package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
@@ -27,6 +28,9 @@ func listUsersHandler(svc *service.Container) gin.HandlerFunc {
|
||||
if svc.Sessions != nil {
|
||||
svc.Sessions.ApplyToUsers(c.Request.Context(), users)
|
||||
}
|
||||
for i := range users {
|
||||
users[i].PopulateComputedFields()
|
||||
}
|
||||
c.JSON(http.StatusOK, users)
|
||||
}
|
||||
}
|
||||
@@ -198,6 +202,64 @@ func updateUserStatusHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
updated.PopulateComputedFields()
|
||||
c.JSON(http.StatusOK, updated)
|
||||
}
|
||||
}
|
||||
|
||||
type adminUpdateUserLibrariesReq struct {
|
||||
AllowedLibraryIDs *[]string `json:"allowed_library_ids"`
|
||||
}
|
||||
|
||||
func updateUserLibrariesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req adminUpdateUserLibrariesReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
userID := c.Param("id")
|
||||
user, err := svc.Repo.User.FindByID(c.Request.Context(), userID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if user == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "user not found"})
|
||||
return
|
||||
}
|
||||
|
||||
var rawJSON string
|
||||
if req.AllowedLibraryIDs != nil && len(*req.AllowedLibraryIDs) > 0 {
|
||||
var cleanIDs []string
|
||||
for _, id := range *req.AllowedLibraryIDs {
|
||||
trimmed := strings.TrimSpace(id)
|
||||
if trimmed != "" {
|
||||
cleanIDs = append(cleanIDs, trimmed)
|
||||
}
|
||||
}
|
||||
if len(cleanIDs) > 0 {
|
||||
data, err := json.Marshal(cleanIDs)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
rawJSON = string(data)
|
||||
}
|
||||
}
|
||||
|
||||
updates := map[string]any{"allowed_library_ids": rawJSON}
|
||||
if err := svc.Repo.User.UpdateFields(c.Request.Context(), userID, updates); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
updated, err := svc.Repo.User.FindByID(c.Request.Context(), userID)
|
||||
if err != nil || updated == nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to reload user"})
|
||||
return
|
||||
}
|
||||
updated.PopulateComputedFields()
|
||||
c.JSON(http.StatusOK, updated)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -64,6 +64,9 @@ func updateSettingHandler(svc *service.Container) gin.HandlerFunc {
|
||||
if req.Key == "transcode.hw_enabled" || req.Key == "transcode.hw_accel" || req.Key == "transcoder.hardware_accel" || req.Key == "transcoder.encoder" {
|
||||
svc.Transcoder.StopAll()
|
||||
}
|
||||
if req.Key == "cache.images_max_size_mb" && svc.Scheduler != nil {
|
||||
_ = svc.Scheduler.RunNowAsync(c.Request.Context(), "image_cache_cleanup")
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package handler
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -47,3 +48,88 @@ func TestDeleteUserRefusesRecentRealtimeSession(t *testing.T) {
|
||||
t.Fatal("recent realtime user should not be deleted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateUserLibraries(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(model.AllModels()...); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
user := model.User{Base: model.Base{ID: "u1"}, Username: "alice", PasswordHash: "x", Role: "user", IsActive: true}
|
||||
lib1 := model.Library{Base: model.Base{ID: "lib-1"}, Name: "电影", Type: "movie", Path: "/movie"}
|
||||
lib2 := model.Library{Base: model.Base{ID: "lib-2"}, Name: "剧集", Type: "tv", Path: "/tv"}
|
||||
lib3 := model.Library{Base: model.Base{ID: "lib-3"}, Name: "动漫", Type: "anime", Path: "/anime"}
|
||||
if err := repos.DB.Create(&user).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := repos.DB.Create(&[]model.Library{lib1, lib2, lib3}).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
svc := &service.Container{Repo: repos}
|
||||
router := gin.New()
|
||||
router.PATCH("/admin/users/:id/libraries", updateUserLibrariesHandler(svc))
|
||||
|
||||
// 1. 设置限制为 lib-1 和 lib-2
|
||||
body := `{"allowed_library_ids":["lib-1","lib-2"]}`
|
||||
req := httptest.NewRequest(http.MethodPatch, "/admin/users/u1/libraries", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body = %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
found, err := repos.User.FindByID(t.Context(), "u1")
|
||||
if err != nil || found == nil {
|
||||
t.Fatal("user not found")
|
||||
}
|
||||
allowed := found.DecodeAllowedLibraryIDs()
|
||||
if len(allowed) != 2 || allowed[0] != "lib-1" || allowed[1] != "lib-2" {
|
||||
t.Fatalf("expected [lib-1, lib-2], got %v", allowed)
|
||||
}
|
||||
|
||||
// 验证可见性
|
||||
vis := service.UserDefaultMediaVisibility(t.Context(), repos, "u1")
|
||||
if len(vis.AllowedLibraryIDs) != 2 {
|
||||
t.Fatalf("expected 2 allowed libraries, got %v", vis.AllowedLibraryIDs)
|
||||
}
|
||||
if !service.LibraryVisibleForUser(t.Context(), repos, lib1, vis) {
|
||||
t.Fatal("lib1 should be visible")
|
||||
}
|
||||
if !service.LibraryVisibleForUser(t.Context(), repos, lib2, vis) {
|
||||
t.Fatal("lib2 should be visible")
|
||||
}
|
||||
if service.LibraryVisibleForUser(t.Context(), repos, lib3, vis) {
|
||||
t.Fatal("lib3 should not be visible")
|
||||
}
|
||||
|
||||
// 2. 清空限制,恢复全部可见
|
||||
bodyEmpty := `{"allowed_library_ids":[]}`
|
||||
reqEmpty := httptest.NewRequest(http.MethodPatch, "/admin/users/u1/libraries", strings.NewReader(bodyEmpty))
|
||||
reqEmpty.Header.Set("Content-Type", "application/json")
|
||||
wEmpty := httptest.NewRecorder()
|
||||
router.ServeHTTP(wEmpty, reqEmpty)
|
||||
|
||||
if wEmpty.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body = %s", wEmpty.Code, wEmpty.Body.String())
|
||||
}
|
||||
|
||||
foundReset, _ := repos.User.FindByID(t.Context(), "u1")
|
||||
if len(foundReset.DecodeAllowedLibraryIDs()) != 0 {
|
||||
t.Fatalf("expected nil or empty, got %v", foundReset.DecodeAllowedLibraryIDs())
|
||||
}
|
||||
|
||||
visReset := service.UserDefaultMediaVisibility(t.Context(), repos, "u1")
|
||||
if len(visReset.AllowedLibraryIDs) != 0 {
|
||||
t.Fatalf("expected no library restrictions, got %v", visReset.AllowedLibraryIDs)
|
||||
}
|
||||
if !service.LibraryVisibleForUser(t.Context(), repos, lib3, visReset) {
|
||||
t.Fatal("lib3 should now be visible")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -89,6 +89,7 @@ func meHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
u.PopulateComputedFields()
|
||||
c.JSON(http.StatusOK, u)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,148 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/config"
|
||||
"github.com/ShukeBta/MMTL/internal/middleware"
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
"github.com/ShukeBta/MMTL/internal/repository"
|
||||
"github.com/ShukeBta/MMTL/internal/service"
|
||||
)
|
||||
|
||||
func TestMountedEmbyPlayingProgressAndResumePipeline(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(model.AllModels()...); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
|
||||
repos := repository.New(db)
|
||||
user := &model.User{
|
||||
Base: model.Base{ID: "user-1"},
|
||||
Username: "test_viewer",
|
||||
PasswordHash: "x",
|
||||
Role: "user",
|
||||
Tier: "free",
|
||||
IsActive: true,
|
||||
}
|
||||
if err := repos.User.Create(t.Context(), user); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
|
||||
cfg := &config.Config{}
|
||||
logger := zap.NewNop()
|
||||
svc := &service.Container{
|
||||
Repo: repos,
|
||||
Emby: service.NewEmbyService(cfg, logger, repos),
|
||||
Sessions: service.NewSessionTrackerService(logger),
|
||||
Playback: service.NewPlaybackService(logger, repos),
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
// 注册带认证的路由,模拟已登录用户
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Set(middleware.CtxUserID, user.ID)
|
||||
c.Next()
|
||||
})
|
||||
router.POST("/Sessions/Playing/Progress", embyPlayingProgressHandler(svc))
|
||||
router.GET("/Items", embyItemsHandler(svc))
|
||||
router.GET("/Users/:userId/Items/Resume", embyResumeItemsHandler(svc))
|
||||
router.GET("/Sessions", embySessionsHandler(svc))
|
||||
|
||||
remoteMediaID := service.EncodeEmbyRemoteID("mount-1", "remote-item-123")
|
||||
|
||||
// 1. 测试上报进度:客户端使用小写 query 参数 itemId / positionTicks
|
||||
progressReq := httptest.NewRequest(
|
||||
http.MethodPost,
|
||||
"/Sessions/Playing/Progress?itemId="+remoteMediaID+"&positionTicks=300000000&runTimeTicks=1000000000",
|
||||
nil,
|
||||
)
|
||||
wProgress := httptest.NewRecorder()
|
||||
router.ServeHTTP(wProgress, progressReq)
|
||||
if wProgress.Code != http.StatusNoContent {
|
||||
t.Fatalf("progress status = %d, body = %s", wProgress.Code, wProgress.Body.String())
|
||||
}
|
||||
|
||||
// 验证已持久化到 PlaybackHistory
|
||||
var hist model.PlaybackHistory
|
||||
if err := db.Where("user_id = ? AND media_id = ?", user.ID, remoteMediaID).First(&hist).Error; err != nil {
|
||||
t.Fatalf("playback history not saved: %v", err)
|
||||
}
|
||||
if hist.PositionMs != 30000 {
|
||||
t.Fatalf("expected position_ms = 30000, got %d", hist.PositionMs)
|
||||
}
|
||||
|
||||
// 2. 测试 Filters=IsResumable 能够包含该远程条目
|
||||
resumableReq := httptest.NewRequest(
|
||||
http.MethodGet,
|
||||
"/Items?Filters=IsResumable",
|
||||
nil,
|
||||
)
|
||||
wResumable := httptest.NewRecorder()
|
||||
router.ServeHTTP(wResumable, resumableReq)
|
||||
if wResumable.Code != http.StatusOK {
|
||||
t.Fatalf("items resumable status = %d, body = %s", wResumable.Code, wResumable.Body.String())
|
||||
}
|
||||
var resumableEnvelope map[string]any
|
||||
if err := json.Unmarshal(wResumable.Body.Bytes(), &resumableEnvelope); err != nil {
|
||||
t.Fatalf("decode resumable: %v", err)
|
||||
}
|
||||
// 因为没有配置真实的远程客户端连接,该远程条目在当前离线测试中不会 panic 崩溃,并且正常响应 Envelope
|
||||
if resumableEnvelope["TotalRecordCount"] == nil {
|
||||
t.Fatalf("missing TotalRecordCount in resumable envelope")
|
||||
}
|
||||
|
||||
// 3. 测试 /Users/:userId/Items/Resume 别名路由
|
||||
resumeAliasReq := httptest.NewRequest(
|
||||
http.MethodGet,
|
||||
"/Users/"+user.ID+"/Items/Resume",
|
||||
nil,
|
||||
)
|
||||
wResumeAlias := httptest.NewRecorder()
|
||||
router.ServeHTTP(wResumeAlias, resumeAliasReq)
|
||||
if wResumeAlias.Code != http.StatusOK {
|
||||
t.Fatalf("resume alias status = %d, body = %s", wResumeAlias.Code, wResumeAlias.Body.String())
|
||||
}
|
||||
|
||||
// 4. 测试 /Sessions 返回 NowPlayingItem
|
||||
sessionsReq := httptest.NewRequest(http.MethodGet, "/Sessions", nil)
|
||||
wSessions := httptest.NewRecorder()
|
||||
router.ServeHTTP(wSessions, sessionsReq)
|
||||
if wSessions.Code != http.StatusOK {
|
||||
t.Fatalf("sessions status = %d, body = %s", wSessions.Code, wSessions.Body.String())
|
||||
}
|
||||
var sessionsList []map[string]any
|
||||
if err := json.Unmarshal(wSessions.Body.Bytes(), &sessionsList); err != nil {
|
||||
t.Fatalf("decode sessions: %v", err)
|
||||
}
|
||||
if len(sessionsList) == 0 {
|
||||
t.Fatalf("expected at least 1 session")
|
||||
}
|
||||
nowPlaying, ok := sessionsList[0]["NowPlayingItem"].(map[string]any)
|
||||
if !ok || nowPlaying["Id"] != remoteMediaID {
|
||||
t.Fatalf("expected NowPlayingItem with id %q, got %#v", remoteMediaID, sessionsList[0]["NowPlayingItem"])
|
||||
}
|
||||
}
|
||||
|
||||
func signMockToken(secret, userID string) string {
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
|
||||
"sub": userID,
|
||||
"exp": time.Now().Add(time.Hour).Unix(),
|
||||
})
|
||||
s, _ := token.SignedString([]byte(secret))
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,222 @@
|
||||
// Emby 挂载管理 HTTP 层:远程 Emby 服务器(账号)下的媒体库挂载 CRUD,
|
||||
// 以及账号远程媒体库(View)列表预览。
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
"github.com/ShukeBta/MMTL/internal/service"
|
||||
)
|
||||
|
||||
// embyMountView 挂载的对外 JSON(附带账号信息)。
|
||||
type embyMountView struct {
|
||||
model.EmbyMount
|
||||
AccountName string `json:"account_name"`
|
||||
}
|
||||
|
||||
// embyMountInput 创建挂载的请求体(单个或批量)。
|
||||
type embyMountInput struct {
|
||||
AccountID string `json:"account_id" binding:"required"`
|
||||
Views []embyViewInput `json:"views" binding:"required,min=1"`
|
||||
}
|
||||
|
||||
type embyViewInput struct {
|
||||
RemoteViewID string `json:"remote_view_id" binding:"required"`
|
||||
RemoteViewName string `json:"remote_view_name"`
|
||||
CollectionType string `json:"collection_type"`
|
||||
Name string `json:"name"`
|
||||
ProxyPlay bool `json:"proxy_play"`
|
||||
}
|
||||
|
||||
func embyMountViews(mounts []model.EmbyMount, accounts map[string]string) []embyMountView {
|
||||
out := make([]embyMountView, 0, len(mounts))
|
||||
for _, m := range mounts {
|
||||
out = append(out, embyMountView{EmbyMount: m, AccountName: accounts[m.AccountID]})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// embyAccountViewsHandler 列出账号上的远程媒体库(View),供挂载选择。
|
||||
func embyAccountViewsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
acct := svc.EmbyRemote.AccountByID(c.Request.Context(), c.Param("id"))
|
||||
if acct == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "账号不存在或已禁用"})
|
||||
return
|
||||
}
|
||||
views, err := svc.EmbyRemote.RemoteViews(c.Request.Context(), acct)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
type viewEntry struct {
|
||||
RemoteViewID string `json:"remote_view_id"`
|
||||
RemoteViewName string `json:"remote_view_name"`
|
||||
CollectionType string `json:"collection_type"`
|
||||
ChildCount int `json:"child_count"`
|
||||
AlreadyMounted bool `json:"already_mounted"`
|
||||
}
|
||||
mounted := map[string]bool{}
|
||||
if mounts, err := svc.EmbyRemote.ListMountsByAccount(c.Request.Context(), acct.ID); err == nil {
|
||||
for _, m := range mounts {
|
||||
mounted[m.RemoteViewID] = true
|
||||
}
|
||||
}
|
||||
out := make([]viewEntry, 0, len(views))
|
||||
for _, v := range views {
|
||||
viewID := service.RemoteItemIDString(v)
|
||||
if strings.TrimSpace(viewID) == "" {
|
||||
continue
|
||||
}
|
||||
out = append(out, viewEntry{
|
||||
RemoteViewID: viewID,
|
||||
RemoteViewName: service.RemoteItemNameString(v),
|
||||
CollectionType: service.RemoteItemCollectionType(v),
|
||||
ChildCount: service.RemoteItemChildCount(v),
|
||||
AlreadyMounted: mounted[viewID],
|
||||
})
|
||||
}
|
||||
c.JSON(http.StatusOK, out)
|
||||
}
|
||||
}
|
||||
|
||||
// listEmbyMountsHandler 列出全部挂载。
|
||||
func listEmbyMountsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
mounts, err := svc.EmbyRemote.ListMounts(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
names := map[string]string{}
|
||||
if accounts, err := svc.EmbyRemote.ListAccounts(c.Request.Context()); err == nil {
|
||||
for _, a := range accounts {
|
||||
names[a.ID] = a.Name
|
||||
}
|
||||
}
|
||||
out := embyMountViews(mounts, names)
|
||||
if out == nil {
|
||||
out = []embyMountView{}
|
||||
}
|
||||
c.JSON(http.StatusOK, out)
|
||||
}
|
||||
}
|
||||
|
||||
// createEmbyMountsHandler 批量创建挂载(同一账号下的多个远程媒体库)。
|
||||
func createEmbyMountsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req embyMountInput
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
mounts := make([]*model.EmbyMount, 0, len(req.Views))
|
||||
for _, v := range req.Views {
|
||||
mounts = append(mounts, &model.EmbyMount{
|
||||
AccountID: req.AccountID,
|
||||
RemoteViewID: v.RemoteViewID,
|
||||
RemoteViewName: v.RemoteViewName,
|
||||
CollectionType: v.CollectionType,
|
||||
Name: v.Name,
|
||||
ProxyPlay: v.ProxyPlay,
|
||||
Enabled: true,
|
||||
})
|
||||
}
|
||||
if _, err := svc.EmbyRemote.CreateMounts(c.Request.Context(), mounts); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true, "created": len(mounts)})
|
||||
}
|
||||
}
|
||||
|
||||
// fullMountEmbyAccountHandler 全量挂载:把账号所有远程媒体库一次挂载进来。
|
||||
func fullMountEmbyAccountHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
acct := svc.EmbyRemote.AccountByID(c.Request.Context(), c.Param("id"))
|
||||
if acct == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "账号不存在或已禁用"})
|
||||
return
|
||||
}
|
||||
proxy := c.Query("proxy") == "1" || c.Query("proxy") == "true"
|
||||
n, err := svc.EmbyRemote.FullMountAccount(c.Request.Context(), acct, proxy)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true, "created": n})
|
||||
}
|
||||
}
|
||||
|
||||
// updateEmbyMountHandler 更新挂载(显示名 / 代理开关 / 启用)。
|
||||
func updateEmbyMountHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req struct {
|
||||
Name *string `json:"name"`
|
||||
ProxyPlay *bool `json:"proxy_play"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
mount, err := svc.EmbyRemote.MountByID(c.Request.Context(), c.Param("id"))
|
||||
if err != nil || mount == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "挂载不存在"})
|
||||
return
|
||||
}
|
||||
if req.Name != nil {
|
||||
mount.Name = *req.Name
|
||||
}
|
||||
if req.ProxyPlay != nil {
|
||||
mount.ProxyPlay = *req.ProxyPlay
|
||||
}
|
||||
if req.Enabled != nil {
|
||||
mount.Enabled = *req.Enabled
|
||||
}
|
||||
if _, err := svc.EmbyRemote.UpdateMount(c.Request.Context(), mount.ID, mount); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, mount)
|
||||
}
|
||||
}
|
||||
|
||||
// deleteEmbyMountHandler 删除挂载。
|
||||
func deleteEmbyMountHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.EmbyRemote.DeleteMount(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 reorderEmbyMountsReq struct {
|
||||
IDs []string `json:"ids" binding:"required"`
|
||||
}
|
||||
|
||||
// reorderEmbyMountsHandler 批量重排挂载媒体库顺序。
|
||||
func reorderEmbyMountsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req reorderEmbyMountsReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if svc.EmbyRemote == nil {
|
||||
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "emby remote service not available"})
|
||||
return
|
||||
}
|
||||
if err := svc.EmbyRemote.ReorderMounts(c.Request.Context(), req.IDs); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/database"
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
"github.com/ShukeBta/MMTL/internal/repository"
|
||||
"github.com/ShukeBta/MMTL/internal/service"
|
||||
)
|
||||
|
||||
func TestReorderEmbyMountsHandler(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := database.AutoMigrate(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
ctx := t.Context()
|
||||
|
||||
m1 := &model.EmbyMount{AccountID: "acct-1", RemoteViewID: "view-1", Name: "Mount 1"}
|
||||
m2 := &model.EmbyMount{AccountID: "acct-1", RemoteViewID: "view-2", Name: "Mount 2"}
|
||||
_ = repos.EmbyMount.Create(ctx, m1)
|
||||
_ = repos.EmbyMount.Create(ctx, m2)
|
||||
|
||||
svc := &service.Container{
|
||||
Repo: repos,
|
||||
EmbyRemote: service.NewEmbyRemoteService(nil, zap.NewNop(), repos, nil),
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.PUT("/admin/emby/mounts/reorder", reorderEmbyMountsHandler(svc))
|
||||
|
||||
body, _ := json.Marshal(map[string]any{
|
||||
"ids": []string{m2.ID, m1.ID},
|
||||
})
|
||||
req := httptest.NewRequest(http.MethodPut, "/admin/emby/mounts/reorder", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200, got %d: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
list, err := repos.EmbyMount.List(ctx)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(list) != 2 || list[0].ID != m2.ID || list[1].ID != m1.ID {
|
||||
t.Fatalf("expected order [m2, m1], got [m%s, m%s]", list[0].ID, list[1].ID)
|
||||
}
|
||||
}
|
||||
@@ -34,10 +34,19 @@ func embyPlaybackInfoHandler(svc *service.Container) gin.HandlerFunc {
|
||||
// embySubtitleStreamHandler serves an external subtitle track advertised in a
|
||||
// MediaSource's MediaStreams via its Emby index
|
||||
// (/Videos/:id/Subtitles/:index/Stream). The index maps to a discovered
|
||||
// sideloaded subtitle file next to the video (SRT/ASS/SSA/VTT, local or
|
||||
// cloud://), following the same layout appended by mediaStreams.
|
||||
// sideloaded subtitle track next to the video (SRT/ASS/SSA/VTT, local or
|
||||
// cloud://), following the same layout appended by mediaStreams. 远程 Emby
|
||||
// 条目的字幕直接反向代理远程。
|
||||
func embySubtitleStreamHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
encodedID := c.Param("id")
|
||||
if accountID, remoteID, ok := service.DecodeEmbyRemoteID(encodedID); ok {
|
||||
if err := svc.Emby.ProxyRemoteSubtitle(c.Request.Context(), c.Writer, c.Request, accountID, remoteID, c.Param("index")); err != nil {
|
||||
embyError(c, http.StatusNotFound, "subtitle not found")
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
uid := c.Param("userId")
|
||||
if uid == "" {
|
||||
uid = embyUserID(c)
|
||||
@@ -213,12 +222,26 @@ func embyAppendAPIKey(raw, token string) string {
|
||||
return u.String()
|
||||
}
|
||||
|
||||
// embyVideoStreamHandler 是 GET /Videos/{id}/stream 的入口,
|
||||
// 直接代理到我们的 /api/stream/{id}(同一个 ServeFile)。
|
||||
// embyVideoStreamHandler 是 GET /Videos/{id}/stream 的入口。
|
||||
// 远程 Emby 条目(embyremote~ 前缀)走反向代理;本地条目直接代理到
|
||||
// /api/stream/{id}(同一个 ServeFile)。
|
||||
func embyVideoStreamHandler(svc *service.Container, cloudMode string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
encodedID := c.Param("id")
|
||||
if accountID, remoteID, ok := service.DecodeEmbyRemoteID(encodedID); ok {
|
||||
if err := svc.Emby.ProxyRemoteVideoStream(c.Request.Context(), c.Writer, c.Request, accountID, remoteID); err != nil {
|
||||
if errors.Is(err, service.ErrEmbyRemoteNotFound) {
|
||||
c.Status(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
if !c.Writer.Written() {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
uid := embyUserID(c)
|
||||
item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
|
||||
item, err := svc.Emby.Item(c.Request.Context(), encodedID, uid)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -296,6 +319,11 @@ func embyShouldRedirectVideoStreamToSTRM(c *gin.Context, svc *service.Container,
|
||||
|
||||
func embyVideoHLSPlaylistHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// 远程 Emby 条目不做本地转码(播放地址已由 PlaybackInfo 指向远程/代理直连)。
|
||||
if service.IsEmbyRemoteID(c.Param("id")) {
|
||||
c.Status(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
uid := embyUserID(c)
|
||||
item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
|
||||
if err != nil || item == nil || svc.Stream == nil {
|
||||
@@ -319,6 +347,10 @@ func embyVideoHLSPlaylistHandler(svc *service.Container) gin.HandlerFunc {
|
||||
|
||||
func embyVideoHLSSegmentHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if service.IsEmbyRemoteID(c.Param("id")) {
|
||||
c.Status(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
uid := embyUserID(c)
|
||||
item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
|
||||
if err != nil || item == nil || svc.Stream == nil {
|
||||
|
||||
@@ -12,8 +12,13 @@ import (
|
||||
|
||||
type embyPlayingReq struct {
|
||||
ItemId string `json:"ItemId"`
|
||||
ItemIDLower string `json:"itemId"`
|
||||
ID string `json:"Id"`
|
||||
IDLower string `json:"id"`
|
||||
PositionTicks int64 `json:"PositionTicks"`
|
||||
PositionLower int64 `json:"positionTicks"`
|
||||
RunTimeTicks int64 `json:"RunTimeTicks"`
|
||||
RunTimeLower int64 `json:"runTimeTicks"`
|
||||
}
|
||||
|
||||
func embyPlayingProgressHandler(svc *service.Container) gin.HandlerFunc {
|
||||
@@ -25,16 +30,25 @@ func embyPlayingProgressHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
var req embyPlayingReq
|
||||
_ = c.ShouldBindJSON(&req)
|
||||
if req.ItemId == "" {
|
||||
req.ItemId = c.Query("ItemId")
|
||||
itemID := embyFirstNonEmptyString(req.ItemId, req.ItemIDLower, req.ID, req.IDLower)
|
||||
if itemID == "" {
|
||||
itemID = embyFirstNonEmptyString(firstQueryValue(c, "ItemId", "itemId", "Id", "id"))
|
||||
}
|
||||
if req.PositionTicks == 0 {
|
||||
req.PositionTicks, _ = strconv.ParseInt(c.Query("PositionTicks"), 10, 64)
|
||||
pos := req.PositionTicks
|
||||
if pos == 0 {
|
||||
pos = req.PositionLower
|
||||
}
|
||||
if req.RunTimeTicks == 0 {
|
||||
req.RunTimeTicks, _ = strconv.ParseInt(c.Query("RunTimeTicks"), 10, 64)
|
||||
if pos == 0 {
|
||||
pos, _ = strconv.ParseInt(firstQueryValue(c, "PositionTicks", "positionTicks"), 10, 64)
|
||||
}
|
||||
if req.ItemId == "" {
|
||||
runTime := req.RunTimeTicks
|
||||
if runTime == 0 {
|
||||
runTime = req.RunTimeLower
|
||||
}
|
||||
if runTime == 0 {
|
||||
runTime, _ = strconv.ParseInt(firstQueryValue(c, "RunTimeTicks", "runTimeTicks"), 10, 64)
|
||||
}
|
||||
if itemID == "" {
|
||||
c.Status(http.StatusOK)
|
||||
return
|
||||
}
|
||||
@@ -43,7 +57,10 @@ func embyPlayingProgressHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.Status(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
_ = svc.Emby.RecordProgress(c.Request.Context(), uid, req.ItemId, req.PositionTicks, req.RunTimeTicks)
|
||||
if err := svc.Emby.RecordProgress(c.Request.Context(), uid, itemID, pos, runTime); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
stopped := strings.Contains(strings.ToLower(c.FullPath()+" "+c.Request.URL.Path), "stopped")
|
||||
if svc.Sessions != nil {
|
||||
svc.Sessions.RecordPlayback(c.Request.Context(), uid, "",
|
||||
@@ -51,9 +68,9 @@ func embyPlayingProgressHandler(svc *service.Container) gin.HandlerFunc {
|
||||
clientInfo.DeviceName,
|
||||
clientInfo.Client,
|
||||
c.ClientIP(),
|
||||
req.ItemId,
|
||||
req.PositionTicks,
|
||||
req.RunTimeTicks,
|
||||
itemID,
|
||||
pos,
|
||||
runTime,
|
||||
stopped)
|
||||
}
|
||||
if svc.Device != nil && !stopped {
|
||||
|
||||
@@ -161,6 +161,8 @@ func registerEmbyAuthenticatedItemRoutes(auth *gin.RouterGroup, svc *service.Con
|
||||
auth.GET("/Users/:userId/Items/Counts", embyItemsCountsHandler(svc))
|
||||
auth.GET("/Items/Latest", embyLatestItemsHandler(svc))
|
||||
auth.GET("/Items/Resume", embyResumeItemsHandler(svc))
|
||||
auth.GET("/Users/:userId/Items/Resume", embyResumeItemsHandler(svc))
|
||||
auth.GET("/UserItems/Resume", embyResumeItemsHandler(svc))
|
||||
auth.GET("/Items/:id", embyItemByIDHandler(svc))
|
||||
auth.GET("/Users/:userId/Items/:id", embyUserItemByIDHandler(svc))
|
||||
auth.GET("/Shows/:id/Seasons", embyShowSeasonsHandler(svc))
|
||||
|
||||
@@ -32,6 +32,8 @@ func registerLowercaseEmbyItemRoutes(auth *gin.RouterGroup, svc *service.Contain
|
||||
auth.GET("/users/:userId/items/counts", embyItemsCountsHandler(svc))
|
||||
auth.GET("/items/latest", embyLatestItemsHandler(svc))
|
||||
auth.GET("/items/resume", embyResumeItemsHandler(svc))
|
||||
auth.GET("/users/:userId/items/resume", embyResumeItemsHandler(svc))
|
||||
auth.GET("/useritems/resume", embyResumeItemsHandler(svc))
|
||||
auth.GET("/items/:id", embyItemByIDHandler(svc))
|
||||
auth.GET("/users/:userId/items/:id", embyUserItemByIDHandler(svc))
|
||||
auth.GET("/shows/:id/seasons", embyShowSeasonsHandler(svc))
|
||||
|
||||
@@ -41,7 +41,17 @@ func embySessionsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
"SupportsRemoteControl": true,
|
||||
}
|
||||
if itemID != "" && sess.IsPlaying {
|
||||
row["NowPlayingItem"] = gin.H{"Id": itemID}
|
||||
nowPlaying := gin.H{"Id": itemID}
|
||||
if svc.Emby != nil {
|
||||
if item, _ := svc.Emby.Item(c.Request.Context(), itemID, sess.UserID); item != nil {
|
||||
for _, key := range []string{"Name", "Type", "RunTimeTicks", "PrimaryImageItemId", "ImageTags", "SeriesName", "SeasonName", "IndexNumber", "ParentIndexNumber"} {
|
||||
if val, ok := item[key]; ok && val != nil {
|
||||
nowPlaying[key] = val
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
row["NowPlayingItem"] = nowPlaying
|
||||
}
|
||||
out = append(out, row)
|
||||
}
|
||||
|
||||
+266
-49
@@ -7,6 +7,7 @@ import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
@@ -16,18 +17,41 @@ import (
|
||||
)
|
||||
|
||||
type createLibraryReq struct {
|
||||
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"`
|
||||
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"`
|
||||
}
|
||||
|
||||
// webLibraryPayload 是 /api/libraries 返回的库条目:本地库与远程 Emby 挂载库
|
||||
// 统一结构(远程库附加 is_remote_emby / remote_source 只读标记)。
|
||||
type webLibraryPayload struct {
|
||||
model.Library
|
||||
IsRemoteEmby bool `json:"is_remote_emby,omitempty"`
|
||||
RemoteSource string `json:"remote_source,omitempty"`
|
||||
Total int64 `json:"total,omitempty"`
|
||||
Cards []service.SeriesCard `json:"cards,omitempty"`
|
||||
}
|
||||
|
||||
// remoteLibraryItemTypes 远程库内容拉取时按 CollectionType 过滤直属条目,
|
||||
// 避免电影库里的合集文件夹(Folder) 漏出为电影卡片。
|
||||
func remoteLibraryItemTypes(collectionType string) string {
|
||||
switch collectionType {
|
||||
case "movies":
|
||||
return "Movie"
|
||||
case "tvshows":
|
||||
return "Series"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
libs, err := svc.Media.ListLibraries(c.Request.Context())
|
||||
ctx := c.Request.Context()
|
||||
libs, err := svc.Media.ListLibraries(ctx)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -35,23 +59,102 @@ func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
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)
|
||||
libs = service.FilterDisplayCloudLibraries(ctx, svc.Repo, libs)
|
||||
visibility := mediaVisibilityForRequest(c, svc)
|
||||
filtered := libs[:0]
|
||||
for _, lib := range libs {
|
||||
if service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, lib, visibility) {
|
||||
if service.LibraryVisibleForUser(ctx, svc.Repo, lib, visibility) {
|
||||
filtered = append(filtered, lib)
|
||||
}
|
||||
}
|
||||
libs = filtered
|
||||
}
|
||||
c.JSON(http.StatusOK, libs)
|
||||
withPreview := c.Query("with_preview") == "1" || c.Query("with_preview") == "true"
|
||||
limit := 10
|
||||
if withPreview {
|
||||
limit, _ = strconv.Atoi(c.DefaultQuery("preview_limit", c.DefaultQuery("limit", "10")))
|
||||
if limit <= 0 {
|
||||
limit = 10
|
||||
} else if limit > 100 {
|
||||
limit = 100
|
||||
}
|
||||
}
|
||||
out := make([]webLibraryPayload, 0, len(libs)+8)
|
||||
if withPreview {
|
||||
previews, err := svc.Media.ListLibrariesWithPreview(ctx, libs, mediaVisibilityForRequest(c, svc), limit)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
for _, p := range previews {
|
||||
out = append(out, webLibraryPayload{Library: p.Library, Total: p.Total, Cards: p.Cards})
|
||||
}
|
||||
} else {
|
||||
for _, l := range libs {
|
||||
out = append(out, webLibraryPayload{Library: l})
|
||||
}
|
||||
}
|
||||
// 远程 Emby 挂载库追加在本地库之后。
|
||||
if svc.EmbyRemote != nil {
|
||||
if views, err := svc.EmbyRemote.RemoteLibraries(ctx); err == nil {
|
||||
remotePayloads := make([]webLibraryPayload, len(views))
|
||||
for i, v := range views {
|
||||
remotePayloads[i] = webLibraryPayload{Library: v.Library, IsRemoteEmby: true, RemoteSource: v.AccountName}
|
||||
}
|
||||
if withPreview && len(views) > 0 {
|
||||
const maxRemotePreviewWorkers = 6
|
||||
sem := make(chan struct{}, maxRemotePreviewWorkers)
|
||||
var wg sync.WaitGroup
|
||||
for i, v := range views {
|
||||
i, v := i, v
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
select {
|
||||
case sem <- struct{}{}:
|
||||
defer func() { <-sem }()
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
acct := svc.EmbyRemote.AccountByID(ctx, v.AccountID)
|
||||
if acct == nil {
|
||||
return
|
||||
}
|
||||
tmpMount := &model.EmbyMount{Base: model.Base{ID: v.MountID}}
|
||||
itemTypes := remoteLibraryItemTypes(v.CollectionType)
|
||||
if _, total, err := svc.EmbyRemote.RemoteLibraryMedia(ctx, tmpMount, acct, v.RemoteID, itemTypes, 0, 1); err == nil {
|
||||
remotePayloads[i].Total = total
|
||||
}
|
||||
if cards, err := svc.EmbyRemote.RemoteLatestCards(ctx, tmpMount, acct, v.RemoteID, limit); err == nil {
|
||||
remotePayloads[i].Cards = cards
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
}
|
||||
out = append(out, remotePayloads...)
|
||||
}
|
||||
}
|
||||
c.JSON(http.StatusOK, out)
|
||||
}
|
||||
}
|
||||
|
||||
func getLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
lib, err := svc.Repo.Library.FindByID(c.Request.Context(), c.Param("id"))
|
||||
ctx := c.Request.Context()
|
||||
id := c.Param("id")
|
||||
// 远程 Emby 挂载库详情。
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(id) {
|
||||
mountID, remoteID, _ := service.DecodeEmbyRemoteID(id)
|
||||
view, err := svc.EmbyRemote.RemoteLibraryByID(ctx, mountID, remoteID)
|
||||
if err != nil || view == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, webLibraryPayload{Library: view.Library, IsRemoteEmby: true, RemoteSource: view.AccountName})
|
||||
return
|
||||
}
|
||||
lib, err := svc.Repo.Library.FindByID(ctx, id)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -63,14 +166,14 @@ func getLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
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)) {
|
||||
libs := service.FilterDisplayCloudLibraries(ctx, svc.Repo, []model.Library{*lib})
|
||||
if len(libs) == 0 || !service.LibraryVisibleForUser(ctx, svc.Repo, libs[0], mediaVisibilityForRequest(c, svc)) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, libs[0])
|
||||
c.JSON(http.StatusOK, webLibraryPayload{Library: libs[0]})
|
||||
} else {
|
||||
c.JSON(http.StatusOK, lib)
|
||||
c.JSON(http.StatusOK, webLibraryPayload{Library: *lib})
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -89,38 +192,38 @@ func createLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
if len(roots) == 0 && strings.TrimSpace(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
|
||||
roots = append(roots, service.LibraryRootInput{Path: req.Path})
|
||||
}
|
||||
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()) }()
|
||||
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
|
||||
}
|
||||
for _, root := range lib.Roots {
|
||||
if root.Enabled {
|
||||
queueLibraryRootScan(svc, lib.ID, root.ID)
|
||||
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
|
||||
}
|
||||
c.JSON(http.StatusCreated, gin.H{"libraries": created})
|
||||
return
|
||||
}
|
||||
l, err := svc.Media.CreateLibraryWithRootsAndCover(c.Request.Context(), req.Name, req.Type, req.CoverURL, roots)
|
||||
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
|
||||
@@ -147,7 +250,9 @@ roots = append(roots, service.LibraryRootInput{Path: req.Path})
|
||||
}
|
||||
|
||||
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 {
|
||||
@@ -157,9 +262,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 {
|
||||
@@ -170,6 +283,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")
|
||||
@@ -187,8 +319,37 @@ func deleteLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
func listMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
ctx := c.Request.Context()
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
size, _ := strconv.Atoi(c.DefaultQuery("page_size", "50"))
|
||||
// 远程 Emby 库:转发远程直属条目并映射为本地 Media 结构(分页由远程承接)。
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(id) {
|
||||
mountID, remoteID, _ := service.DecodeEmbyRemoteID(id)
|
||||
mount, acct, _ := svc.EmbyRemote.ResolveMount(ctx, mountID)
|
||||
if mount == nil || acct == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
itemTypes := ""
|
||||
if view, err := svc.EmbyRemote.RemoteLibraryByID(ctx, mountID, remoteID); err == nil && view != nil {
|
||||
itemTypes = remoteLibraryItemTypes(view.CollectionType)
|
||||
}
|
||||
items, total, err := svc.EmbyRemote.RemoteLibraryMedia(ctx, mount, acct, remoteID, itemTypes, (page-1)*size, size)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if items == nil {
|
||||
items = []model.Media{}
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"items": items,
|
||||
"total": total,
|
||||
"page": page,
|
||||
"page_size": size,
|
||||
})
|
||||
return
|
||||
}
|
||||
groupVersions := c.DefaultQuery("group_versions", "1") != "0"
|
||||
if !groupVersions {
|
||||
items, total, err := svc.Media.ListMediaVisible(c.Request.Context(), id, page, size, mediaVisibilityForRequest(c, svc))
|
||||
@@ -226,7 +387,33 @@ func listMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
|
||||
func getMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
m, err := svc.Media.GetMedia(c.Request.Context(), c.Param("id"))
|
||||
ctx := c.Request.Context()
|
||||
id := c.Param("id")
|
||||
// 远程 Emby 条目:拉远程详情并映射为本地 Media 结构。
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(id) {
|
||||
mountID, remoteID, _ := service.DecodeEmbyRemoteID(id)
|
||||
mount, acct, _ := svc.EmbyRemote.ResolveMount(ctx, mountID)
|
||||
if mount == nil || acct == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
m, err := svc.EmbyRemote.RemoteMediaDetail(ctx, mount, acct, remoteID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if m == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
if !mediaVisibleForRequest(c, svc, m) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, m)
|
||||
return
|
||||
}
|
||||
m, err := svc.Media.GetMedia(ctx, id)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -320,7 +507,37 @@ func searchMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
|
||||
func streamHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
m, err := svc.Media.GetMedia(c.Request.Context(), c.Param("id"))
|
||||
ctx := c.Request.Context()
|
||||
id := c.Param("id")
|
||||
// 远程 Emby 条目:按挂载代理配置分流——代理走 MMTL 反代,否则 302 直连。
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(id) {
|
||||
if !enforceScopedPlaybackToken(c, id) {
|
||||
return
|
||||
}
|
||||
mountID, remoteID, _ := service.DecodeEmbyRemoteID(id)
|
||||
mount, acct, _ := svc.EmbyRemote.ResolveMount(ctx, mountID)
|
||||
if mount == nil || acct == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
if mount.ProxyPlay {
|
||||
if err := svc.Emby.ProxyRemoteVideoStream(ctx, c.Writer, c.Request, mountID, remoteID); err != nil {
|
||||
if !c.Writer.Written() {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
target, err := svc.EmbyRemote.WebStreamURL(ctx, acct, remoteID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
setRedirectNoStoreHeaders(c)
|
||||
c.Redirect(http.StatusFound, target)
|
||||
return
|
||||
}
|
||||
m, err := svc.Media.GetMedia(ctx, id)
|
||||
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
|
||||
@@ -19,21 +19,40 @@ import (
|
||||
"github.com/ShukeBta/MMTL/internal/service"
|
||||
)
|
||||
|
||||
func findMediaForPlaybackEndpoint(c *gin.Context, svc *service.Container, id string) (*model.Media, error) {
|
||||
ctx := c.Request.Context()
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(id) {
|
||||
mountID, remoteID, _ := service.DecodeEmbyRemoteID(id)
|
||||
mount, acct, _ := svc.EmbyRemote.ResolveMount(ctx, mountID)
|
||||
if mount == nil || acct == nil {
|
||||
return nil, nil
|
||||
}
|
||||
return svc.EmbyRemote.RemoteMediaDetail(ctx, mount, acct, remoteID)
|
||||
}
|
||||
return svc.Repo.Media.FindByID(ctx, id)
|
||||
}
|
||||
|
||||
// playbackInfoHandler returns the media row + a `stream_url` the React
|
||||
// player can hit. Mirrors the Python project's surface.
|
||||
func playbackInfoHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
|
||||
id := c.Param("id")
|
||||
m, err := findMediaForPlaybackEndpoint(c, svc, id)
|
||||
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "media not found"})
|
||||
return
|
||||
}
|
||||
token := externalPlaybackToken(c, svc, m.ID, m.DurationSec)
|
||||
profileQuery := externalProfileQuery(c)
|
||||
hlsURL := "/api/hls/" + m.ID + "/index.m3u8?token=" + url.QueryEscape(token) + profileQuery
|
||||
if service.IsEmbyRemoteID(m.ID) || service.IsStrmMediaRow(m) {
|
||||
// Emby 远程挂载与 STRM 媒体一样,默认直连播放,不提供转码地址
|
||||
hlsURL = ""
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"media": m,
|
||||
"stream_url": "/api/stream/" + m.ID + "?token=" + url.QueryEscape(token) + profileQuery,
|
||||
"hls_url": "/api/hls/" + m.ID + "/index.m3u8?token=" + url.QueryEscape(token) + profileQuery,
|
||||
"hls_url": hlsURL,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -68,7 +87,8 @@ func playbackProgressHandler(svc *service.Container) gin.HandlerFunc {
|
||||
// produce the per-player launch URL.
|
||||
func externalPlayersHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
|
||||
id := c.Param("id")
|
||||
m, err := findMediaForPlaybackEndpoint(c, svc, id)
|
||||
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "media not found"})
|
||||
return
|
||||
@@ -93,7 +113,8 @@ func externalPlayersHandler(svc *service.Container) gin.HandlerFunc {
|
||||
// token query string the external player needs.
|
||||
func externalURLHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
|
||||
id := c.Param("id")
|
||||
m, err := findMediaForPlaybackEndpoint(c, svc, id)
|
||||
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "media not found"})
|
||||
return
|
||||
|
||||
@@ -404,8 +404,77 @@ func newPlaybackScopeTestRouter(t *testing.T) (*gin.Engine, *service.Container,
|
||||
router := gin.New()
|
||||
api := router.Group("/api")
|
||||
api.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret))
|
||||
api.GET("/playback/:id/info", playbackInfoHandler(svc))
|
||||
api.GET("/playback/:id/external-url", externalURLHandler(svc))
|
||||
api.GET("/playback/:id/external-players", externalPlayersHandler(svc))
|
||||
api.GET("/stream/:id", streamHandler(svc))
|
||||
api.GET("/hls/:id/index.m3u8", hlsPlaylistHandler(svc))
|
||||
api.GET("/media/:id/subtitles", listSubtitlesHandler(svc))
|
||||
return router, svc, cfg.Secrets.JWTSecret
|
||||
}
|
||||
|
||||
func TestPlaybackInfoForSTRMMediaDisablesHLS(t *testing.T) {
|
||||
router, _, secret := newPlaybackScopeTestRouter(t)
|
||||
loginToken := signedTestToken(t, secret)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://nas.local/api/playback/media-1/info", nil)
|
||||
req.Header.Set("Authorization", "Bearer "+loginToken)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
var payload struct {
|
||||
StreamURL string `json:"stream_url"`
|
||||
HlsURL string `json:"hls_url"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
|
||||
t.Fatalf("decode: %v", err)
|
||||
}
|
||||
if payload.StreamURL == "" {
|
||||
t.Fatalf("expected non-empty stream_url")
|
||||
}
|
||||
if payload.HlsURL != "" {
|
||||
t.Fatalf("expected empty hls_url for STRM media, got %q", payload.HlsURL)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHLSPlaylistForRemoteEmbyMediaDisabled(t *testing.T) {
|
||||
router, svc, secret := newPlaybackScopeTestRouter(t)
|
||||
svc.EmbyRemote = &service.EmbyRemoteService{}
|
||||
loginToken := signedTestToken(t, secret)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://nas.local/api/hls/embyremote~acct1~item1/index.m3u8", nil)
|
||||
req.Header.Set("Authorization", "Bearer "+loginToken)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusConflict {
|
||||
t.Fatalf("status = %d, want %d (409 StatusConflict)", w.Code, http.StatusConflict)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListSubtitlesForRemoteEmbyMediaReturnsEmptyTracks(t *testing.T) {
|
||||
router, svc, secret := newPlaybackScopeTestRouter(t)
|
||||
svc.EmbyRemote = &service.EmbyRemoteService{}
|
||||
loginToken := signedTestToken(t, secret)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://nas.local/api/media/embyremote~acct1~item1/subtitles", nil)
|
||||
req.Header.Set("Authorization", "Bearer "+loginToken)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200 OK", w.Code)
|
||||
}
|
||||
var payload struct {
|
||||
Tracks []any `json:"tracks"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
|
||||
t.Fatalf("decode: %v", err)
|
||||
}
|
||||
if payload.Tracks == nil || len(payload.Tracks) != 0 {
|
||||
t.Fatalf("expected empty tracks array, got %v", payload.Tracks)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -20,9 +20,41 @@ 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)
|
||||
|
||||
// FFmpeg/FFprobe 工具:状态查询 + 一键下载安装(自动匹配当前平台)。
|
||||
admin.GET("/tools/ffmpeg/status", ffToolsStatusHandler(svc))
|
||||
admin.POST("/tools/ffmpeg/install", ffToolsInstallHandler(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) {
|
||||
// Emby 挂载管理:远程 Emby 媒体库挂载(账号复用 strm/accounts)
|
||||
admin.GET("/emby/accounts/:id/views", embyAccountViewsHandler(svc))
|
||||
admin.POST("/emby/accounts/:id/full-mount", fullMountEmbyAccountHandler(svc))
|
||||
admin.GET("/emby/mounts", listEmbyMountsHandler(svc))
|
||||
admin.POST("/emby/mounts", createEmbyMountsHandler(svc))
|
||||
admin.PUT("/emby/mounts/reorder", reorderEmbyMountsHandler(svc))
|
||||
admin.PUT("/emby/mounts/:id", updateEmbyMountHandler(svc))
|
||||
admin.DELETE("/emby/mounts/:id", deleteEmbyMountHandler(svc))
|
||||
|
||||
admin.GET("/strm/accounts", listStrmAccountsHandler(svc))
|
||||
admin.POST("/strm/accounts", createStrmAccountHandler(svc))
|
||||
admin.PUT("/strm/accounts/:id", updateStrmAccountHandler(svc))
|
||||
@@ -50,6 +82,8 @@ func registerAdminStrmRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
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))
|
||||
@@ -58,8 +92,13 @@ func registerAdminStrmRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.GET("/strm/uploads", uploadQueueHandler(svc))
|
||||
admin.POST("/strm/uploads/:id/cancel", cancelStrmUploadHandler(svc))
|
||||
admin.POST("/strm/uploads/:id/retry", retryStrmUploadHandler(svc))
|
||||
admin.POST("/strm/uploads/cancel-pending", cancelPendingUploadsHandler(svc))
|
||||
admin.DELETE("/strm/uploads/:id", deleteStrmUploadHandler(svc))
|
||||
admin.POST("/strm/uploads/batch", batchActionUploadsHandler(svc))
|
||||
admin.POST("/strm/uploads/clear-done", clearDoneUploadsHandler(svc))
|
||||
admin.POST("/strm/uploads/clear-finished", clearFinishedUploadsHandler(svc))
|
||||
admin.POST("/strm/uploads/clear-canceled", clearCanceledUploadsHandler(svc))
|
||||
admin.POST("/strm/uploads/retry-failed", retryAllFailedUploadsHandler(svc))
|
||||
admin.POST("/strm/uploads/cancel-pending", cancelPendingUploadsHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminUserRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
@@ -69,6 +108,7 @@ func registerAdminUserRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.PATCH("/users/:id/password", resetUserPasswordHandler(svc))
|
||||
admin.PATCH("/users/:id/status", updateUserStatusHandler(svc))
|
||||
admin.PATCH("/users/:id/role", adminUpdateRoleHandler(svc))
|
||||
admin.PATCH("/users/:id/libraries", updateUserLibrariesHandler(svc))
|
||||
admin.DELETE("/users/:id", deleteUserHandler(svc))
|
||||
admin.GET("/settings", listSettingsHandler(svc))
|
||||
admin.PUT("/settings", updateSettingHandler(svc))
|
||||
@@ -116,3 +156,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))
|
||||
@@ -38,6 +39,7 @@ func registerAuthedLibraryRoutes(authed *gin.RouterGroup, svc *service.Container
|
||||
|
||||
func registerAuthedMediaRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.GET("/media/:id", getMediaHandler(svc))
|
||||
authed.GET("/media/:id/episodes", listMediaEpisodesHandler(svc))
|
||||
authed.GET("/media", searchMediaHandler(svc))
|
||||
authed.PATCH("/media/:id/metadata", middleware.AdminRequired(), updateMediaMetadataHandler(svc))
|
||||
authed.POST("/media/:id/scrape", middleware.AdminRequired(), scrapeOneHandler(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})
|
||||
}
|
||||
}
|
||||
+100
-2
@@ -65,7 +65,49 @@ func listSeasonsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
func listLibrarySeriesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
libID := c.Param("id")
|
||||
if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), libID); err == nil && lib != nil {
|
||||
ctx := c.Request.Context()
|
||||
// 远程剧集库:远程 Series 映射为系列卡片。
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(libID) {
|
||||
mountID, remoteID, _ := service.DecodeEmbyRemoteID(libID)
|
||||
mount, acct, _ := svc.EmbyRemote.ResolveMount(ctx, mountID)
|
||||
if mount == nil || acct == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
cards, err := svc.EmbyRemote.RemoteSeriesCards(ctx, mount, acct, remoteID)
|
||||
if err != nil {
|
||||
writeInternalOrCanceled(c, err)
|
||||
return
|
||||
}
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
size, _ := strconv.Atoi(c.DefaultQuery("page_size", "500"))
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if size <= 0 || size > 1000 {
|
||||
size = 500
|
||||
}
|
||||
start := (page - 1) * size
|
||||
if start > len(cards) {
|
||||
start = len(cards)
|
||||
}
|
||||
end := start + size
|
||||
if end > len(cards) {
|
||||
end = len(cards)
|
||||
}
|
||||
pageItems := cards[start:end]
|
||||
if pageItems == nil {
|
||||
pageItems = []service.SeriesCard{}
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"items": pageItems,
|
||||
"total": len(cards),
|
||||
"page": page,
|
||||
"page_size": size,
|
||||
})
|
||||
return
|
||||
}
|
||||
if lib, err := svc.Repo.Library.FindByID(ctx, libID); err == nil && lib != nil {
|
||||
if !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, *lib, mediaVisibilityForRequest(c, svc)) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
@@ -114,7 +156,27 @@ func listLibrarySeriesEpisodesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "key is required"})
|
||||
return
|
||||
}
|
||||
if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), libID); err == nil && lib != nil {
|
||||
ctx := c.Request.Context()
|
||||
// 远程系列 key(伪装系列 ID):转发远程该系列全部剧集。
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(key) {
|
||||
mountID, remoteSeriesID, _ := service.DecodeEmbyRemoteID(key)
|
||||
mount, acct, _ := svc.EmbyRemote.ResolveMount(ctx, mountID)
|
||||
if mount == nil || acct == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
items, err := svc.EmbyRemote.RemoteEpisodes(ctx, mount, acct, remoteSeriesID)
|
||||
if err != nil {
|
||||
writeInternalOrCanceled(c, err)
|
||||
return
|
||||
}
|
||||
if items == nil {
|
||||
items = []model.Media{}
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": items, "total": len(items)})
|
||||
return
|
||||
}
|
||||
if lib, err := svc.Repo.Library.FindByID(ctx, libID); err == nil && lib != nil {
|
||||
if !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, *lib, mediaVisibilityForRequest(c, svc)) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
@@ -128,3 +190,39 @@ func listLibrarySeriesEpisodesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusOK, gin.H{"items": items, "total": len(items)})
|
||||
}
|
||||
}
|
||||
|
||||
func listMediaEpisodesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "id is required"})
|
||||
return
|
||||
}
|
||||
ctx := c.Request.Context()
|
||||
// 远程条目:单集→同系列集列表;系列/季/文件夹→子集;电影→自身单条。
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(id) {
|
||||
mountID, remoteID, _ := service.DecodeEmbyRemoteID(id)
|
||||
mount, acct, _ := svc.EmbyRemote.ResolveMount(ctx, mountID)
|
||||
if mount == nil || acct == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
items, err := svc.EmbyRemote.RemoteEpisodes(ctx, mount, acct, remoteID)
|
||||
if err != nil {
|
||||
writeInternalOrCanceled(c, err)
|
||||
return
|
||||
}
|
||||
if items == nil {
|
||||
items = []model.Media{}
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": items, "total": len(items)})
|
||||
return
|
||||
}
|
||||
items, err := svc.Media.ListMediaEpisodes(ctx, id, mediaVisibilityForRequest(c, svc))
|
||||
if err != nil {
|
||||
writeInternalOrCanceled(c, err)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": items, "total": len(items)})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
@@ -14,7 +13,13 @@ import (
|
||||
|
||||
func hlsPlaylistHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
m, err := svc.Media.GetMedia(c.Request.Context(), c.Param("id"))
|
||||
id := c.Param("id")
|
||||
// 远程 Emby 挂载媒体与 STRM 一样,默认直连播放,不进行转码。
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(id) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": "transcode disabled"})
|
||||
return
|
||||
}
|
||||
m, err := svc.Media.GetMedia(c.Request.Context(), id)
|
||||
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
@@ -44,7 +49,12 @@ func hlsPlaylistHandler(svc *service.Container) gin.HandlerFunc {
|
||||
|
||||
func hlsSegmentHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
m, err := svc.Media.GetMedia(c.Request.Context(), c.Param("id"))
|
||||
id := c.Param("id")
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(id) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": "transcode disabled"})
|
||||
return
|
||||
}
|
||||
m, err := svc.Media.GetMedia(c.Request.Context(), id)
|
||||
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
@@ -135,28 +145,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 +164,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})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+105
-7
@@ -28,17 +28,25 @@ type strmAccountView struct {
|
||||
model.StrmAccount
|
||||
HasCredential bool `json:"has_credential"`
|
||||
ProviderLabel string `json:"provider_label"`
|
||||
// ProxyPlay 仅远程 Emby 挂载账号返回:播放流量是否经过 MMTL 代理(编辑回显用)。
|
||||
ProxyPlay *bool `json:"proxy_play,omitempty"`
|
||||
}
|
||||
|
||||
func strmAccountViews(accounts []model.StrmAccount) []strmAccountView {
|
||||
func strmAccountViews(svc *service.Container, accounts []model.StrmAccount) []strmAccountView {
|
||||
out := make([]strmAccountView, 0, len(accounts))
|
||||
for i := range accounts {
|
||||
a := accounts[i]
|
||||
out = append(out, strmAccountView{
|
||||
view := strmAccountView{
|
||||
StrmAccount: a,
|
||||
HasCredential: service.HasStrmAccountCredential(&a),
|
||||
ProviderLabel: providerLabelOf(a.Provider),
|
||||
})
|
||||
}
|
||||
if a.Provider == model.StrmProviderEmbyRemote && svc != nil && svc.EmbyRemote != nil {
|
||||
if proxyPlay, err := svc.EmbyRemote.ProxyPlayOf(&a); err == nil {
|
||||
view.ProxyPlay = &proxyPlay
|
||||
}
|
||||
}
|
||||
out = append(out, view)
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -58,7 +66,7 @@ func listStrmAccountsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, strmAccountViews(accounts))
|
||||
c.JSON(http.StatusOK, strmAccountViews(svc, accounts))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -74,7 +82,7 @@ func createStrmAccountHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
views := strmAccountViews([]model.StrmAccount{*acct})
|
||||
views := strmAccountViews(svc, []model.StrmAccount{*acct})
|
||||
c.JSON(http.StatusOK, views[0])
|
||||
}
|
||||
}
|
||||
@@ -92,7 +100,7 @@ func updateStrmAccountHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
views := strmAccountViews([]model.StrmAccount{*acct})
|
||||
views := strmAccountViews(svc, []model.StrmAccount{*acct})
|
||||
c.JSON(http.StatusOK, views[0])
|
||||
}
|
||||
}
|
||||
@@ -114,7 +122,7 @@ func testStrmAccountHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "网盘账号不存在"})
|
||||
return
|
||||
}
|
||||
views := strmAccountViews([]model.StrmAccount{*acct})
|
||||
views := strmAccountViews(svc, []model.StrmAccount{*acct})
|
||||
c.JSON(http.StatusOK, views[0])
|
||||
}
|
||||
}
|
||||
@@ -393,6 +401,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 {
|
||||
@@ -439,6 +504,28 @@ func clearCanceledUploadsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
func clearDoneUploadsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
n, err := svc.Strm.ClearDoneUploadTasks(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 clearFinishedUploadsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
n, err := svc.Strm.ClearFinishedUploadTasks(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())
|
||||
@@ -450,6 +537,17 @@ func retryAllFailedDownloadsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
func retryAllFailedUploadsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
n, err := svc.Strm.RetryAllFailedUploadTasks(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 cancelPendingDownloadsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
n, err := svc.Strm.CancelPendingDownloadTasks(c.Request.Context())
|
||||
|
||||
@@ -50,16 +50,19 @@ func TestStrmAdminRoutesAreRegistered(t *testing.T) {
|
||||
"GET /api/admin/strm/downloads",
|
||||
"POST /api/admin/strm/downloads/:id/cancel",
|
||||
"POST /api/admin/strm/downloads/:id/retry",
|
||||
"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",
|
||||
"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/clear-done",
|
||||
"POST /api/admin/strm/uploads/clear-finished",
|
||||
"POST /api/admin/strm/uploads/clear-canceled",
|
||||
"POST /api/admin/strm/uploads/retry-failed",
|
||||
"POST /api/admin/strm/uploads/cancel-pending",
|
||||
"GET /api/strm/play/:provider/:file",
|
||||
} {
|
||||
if !routes[want] {
|
||||
t.Fatalf("%s route is not registered", want)
|
||||
|
||||
@@ -11,7 +11,12 @@ import (
|
||||
|
||||
func listSubtitlesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
tracks, err := svc.Subtitle.Discover(c.Request.Context(), c.Param("id"))
|
||||
id := c.Param("id")
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(id) {
|
||||
c.JSON(http.StatusOK, gin.H{"tracks": []service.SubtitleTrack{}})
|
||||
return
|
||||
}
|
||||
tracks, err := svc.Subtitle.Discover(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -31,7 +36,7 @@ func serveSubtitleHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
c.Header("Content-Type", "text/vtt; charset=utf-8")
|
||||
c.Header("Cache-Control", "public, max-age=3600")
|
||||
c.Header("Cache-Control", "no-cache, no-store, must-revalidate")
|
||||
if err := svc.Subtitle.Serve(c.Request.Context(), c.Param("id"), path, c.Writer); err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": err.Error()})
|
||||
return
|
||||
|
||||
@@ -1,144 +0,0 @@
|
||||
// Package handler — system tools detection.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"os/exec"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/config"
|
||||
"github.com/ShukeBta/MMTL/internal/service"
|
||||
)
|
||||
|
||||
// SystemHandler handles system-related endpoints.
|
||||
type SystemHandler struct {
|
||||
cfg *config.Config
|
||||
log *zap.Logger
|
||||
svc *service.Container
|
||||
}
|
||||
|
||||
// NewSystemHandler is the constructor.
|
||||
func NewSystemHandler(cfg *config.Config, log *zap.Logger, svc *service.Container) *SystemHandler {
|
||||
return &SystemHandler{cfg: cfg, log: log, svc: svc}
|
||||
}
|
||||
|
||||
// ToolStatus represents the detection status of a system tool.
|
||||
type ToolStatus struct {
|
||||
Name string `json:"name"`
|
||||
DisplayName string `json:"display_name"`
|
||||
ConfigKey string `json:"config_key"`
|
||||
Path string `json:"path,omitempty"`
|
||||
Detected bool `json:"detected"`
|
||||
Version string `json:"version,omitempty"`
|
||||
}
|
||||
|
||||
// GetToolsStatus returns the status of system tools.
|
||||
func (h *SystemHandler) GetToolsStatus(c *gin.Context) {
|
||||
tools := []ToolStatus{
|
||||
{Name: "ffprobe", DisplayName: "FFprobe", ConfigKey: "app.ffprobe_path"},
|
||||
{Name: "ffmpeg", DisplayName: "FFmpeg", ConfigKey: "app.ffmpeg_path"},
|
||||
}
|
||||
|
||||
for i := range tools {
|
||||
// Check configured path first
|
||||
var configuredPath string
|
||||
switch tools[i].ConfigKey {
|
||||
case "app.ffprobe_path":
|
||||
configuredPath = h.cfg.App.FFprobePath
|
||||
if configuredPath == "" {
|
||||
configuredPath = "ffprobe"
|
||||
}
|
||||
case "app.ffmpeg_path":
|
||||
configuredPath = h.cfg.App.FFmpegPath
|
||||
if configuredPath == "" {
|
||||
configuredPath = "ffmpeg"
|
||||
}
|
||||
}
|
||||
|
||||
// Try to find the tool
|
||||
path, err := exec.LookPath(configuredPath)
|
||||
if err == nil {
|
||||
tools[i].Detected = true
|
||||
tools[i].Path = path
|
||||
// Try to get version
|
||||
tools[i].Version = getToolVersion(path)
|
||||
}
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"tools": tools,
|
||||
})
|
||||
}
|
||||
|
||||
// getToolVersion attempts to get the version of a tool.
|
||||
func getToolVersion(path string) string {
|
||||
out, err := exec.Command(path, "-version").Output()
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
// Extract first line as version info
|
||||
lines := strings.Split(string(out), "\n")
|
||||
if len(lines) > 0 {
|
||||
return strings.TrimSpace(lines[0])
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// InstallTools attempts to auto-install system tools (ffmpeg/ffprobe)
|
||||
func (h *SystemHandler) InstallTools(c *gin.Context) {
|
||||
h.log.Info("Received tools auto-install request")
|
||||
|
||||
// Call service layer to auto-install
|
||||
ffprobePath, ffmpegPath := service.AutoInstallFFmpeg(h.log, h.cfg)
|
||||
|
||||
result := gin.H{
|
||||
"installed": ffprobePath != "" || ffmpegPath != "",
|
||||
}
|
||||
|
||||
if ffprobePath != "" {
|
||||
result["ffprobe_path"] = ffprobePath
|
||||
result["ffprobe_installed"] = true
|
||||
}
|
||||
if ffmpegPath != "" {
|
||||
result["ffmpeg_path"] = ffmpegPath
|
||||
result["ffmpeg_installed"] = true
|
||||
}
|
||||
|
||||
// Re-detect tool status
|
||||
tools := []ToolStatus{
|
||||
{Name: "ffprobe", DisplayName: "FFprobe", ConfigKey: "app.ffprobe_path"},
|
||||
{Name: "ffmpeg", DisplayName: "FFmpeg", ConfigKey: "app.ffmpeg_path"},
|
||||
}
|
||||
|
||||
for i := range tools {
|
||||
var configuredPath string
|
||||
switch tools[i].ConfigKey {
|
||||
case "app.ffprobe_path":
|
||||
configuredPath = h.cfg.App.FFprobePath
|
||||
if configuredPath == "" {
|
||||
configuredPath = "ffprobe"
|
||||
}
|
||||
case "app.ffmpeg_path":
|
||||
configuredPath = h.cfg.App.FFmpegPath
|
||||
if configuredPath == "" {
|
||||
configuredPath = "ffmpeg"
|
||||
}
|
||||
}
|
||||
|
||||
path, err := exec.LookPath(configuredPath)
|
||||
if err == nil {
|
||||
tools[i].Detected = true
|
||||
tools[i].Path = path
|
||||
tools[i].Version = getToolVersion(path)
|
||||
}
|
||||
}
|
||||
|
||||
result["tools"] = tools
|
||||
|
||||
h.log.Info("Tool installation completed", zap.Any("result", result))
|
||||
c.JSON(http.StatusOK, result)
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
// Package handler — FFmpeg/FFprobe 工具安装端点。
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/service"
|
||||
)
|
||||
|
||||
// ffToolsStatusHandler 返回 ffmpeg/ffprobe 当前安装状态
|
||||
// (GET /api/admin/tools/ffmpeg/status)。
|
||||
func ffToolsStatusHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if svc == nil || svc.FFTools == nil {
|
||||
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "FFmpeg 工具服务不可用"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, svc.FFTools.Status(c.Request.Context()))
|
||||
}
|
||||
}
|
||||
|
||||
// ffToolsInstallHandler 触发后台下载安装(POST /api/admin/tools/ffmpeg/install)。
|
||||
// 自动匹配当前运行环境(OS+架构),安装到 data/tools/ffmpeg/ 并把路径写入设置。
|
||||
func ffToolsInstallHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if svc == nil || svc.FFTools == nil {
|
||||
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "FFmpeg 工具服务不可用"})
|
||||
return
|
||||
}
|
||||
if err := svc.FFTools.StartInstall(c.Request.Context()); err != nil {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, svc.FFTools.Status(c.Request.Context()))
|
||||
}
|
||||
}
|
||||
@@ -32,7 +32,14 @@ func mediaVisibilityForRequest(c *gin.Context, svc *service.Container) service.M
|
||||
return visibility
|
||||
}
|
||||
visibility.IncludeNSFW = adultEnabled && profile.AllowAdult && !userHidesAdult
|
||||
visibility.AllowedLibraryIDs = profileAllowedLibraryIDs(*profile)
|
||||
profileAllowed := profileAllowedLibraryIDs(*profile)
|
||||
if len(profileAllowed) > 0 {
|
||||
if len(visibility.AllowedLibraryIDs) > 0 {
|
||||
visibility.AllowedLibraryIDs = service.IntersectStrings(visibility.AllowedLibraryIDs, profileAllowed)
|
||||
} else {
|
||||
visibility.AllowedLibraryIDs = profileAllowed
|
||||
}
|
||||
}
|
||||
if !visibility.IncludeNSFW {
|
||||
visibility.HiddenLibraryIDs = service.AdultLibraryIDs(c.Request.Context(), svc.Repo)
|
||||
} else {
|
||||
|
||||
@@ -125,6 +125,18 @@ func historyContinueHandler(svc *service.Container) gin.HandlerFunc {
|
||||
for _, r := range rows {
|
||||
m, ok := mIdx[r.MediaID]
|
||||
if !ok {
|
||||
if svc.EmbyRemote != nil && service.IsEmbyRemoteID(r.MediaID) {
|
||||
mountID, remoteID, _ := service.DecodeEmbyRemoteID(r.MediaID)
|
||||
if mount, acct, _ := svc.EmbyRemote.ResolveMount(c.Request.Context(), mountID); mount != nil && acct != nil {
|
||||
if rm, err := svc.EmbyRemote.RemoteMediaDetail(c.Request.Context(), mount, acct, remoteID); err == nil && rm != nil {
|
||||
out = append(out, gin.H{
|
||||
"history": r,
|
||||
"media": *rm,
|
||||
})
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
out = append(out, gin.H{
|
||||
@@ -163,7 +175,7 @@ func historyDeleteHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "status must be completed or incomplete"})
|
||||
return
|
||||
}
|
||||
res := q.Unscoped().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
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
// Emby 媒体库挂载模型。
|
||||
//
|
||||
// 远程 Emby 账号(StrmAccount.Provider = emby_remote)只是一个服务器连接;
|
||||
// 「挂载」才决定把该服务器的哪个媒体库(View)暴露到本项目的媒体库中。
|
||||
// 这样同一个 Emby 服务器可以按库选择挂载,且每个挂载独立控制是否由 MMTL
|
||||
// 代理播放流量。
|
||||
package model
|
||||
|
||||
// EmbyMount 是远程 Emby 服务器上一个媒体库(View)的挂载配置。
|
||||
type EmbyMount struct {
|
||||
Base
|
||||
AccountID string `gorm:"size:36;index" json:"account_id"` // StrmAccount.ID(provider=emby_remote)
|
||||
RemoteViewID string `gorm:"size:128" json:"remote_view_id"` // 远程 Emby 的 View Id
|
||||
RemoteViewName string `gorm:"size:255" json:"remote_view_name"` // 远程媒体库原名(展示冗余)
|
||||
CollectionType string `gorm:"size:32" json:"collection_type"` // movies / tvshows / music ...
|
||||
Name string `gorm:"size:255" json:"name,omitempty"` // 覆盖显示名(可选,默认「账号 · 库名」)
|
||||
SortOrder int `gorm:"default:0;index" json:"sort_order"` // 手动排序用,越小越靠前
|
||||
ProxyPlay bool `gorm:"default:false" json:"proxy_play"` // 该挂载播放流量是否经 MMTL 反向代理
|
||||
Enabled bool `gorm:"default:true" json:"enabled"` // 是否在媒体库中展示
|
||||
}
|
||||
@@ -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,8 +54,10 @@ func AllModels() []interface{} {
|
||||
&StrmAccount{},
|
||||
&StrmSyncPath{},
|
||||
&StrmSyncRecord{},
|
||||
&StrmDownloadTask{},
|
||||
&StrmUploadTask{},
|
||||
&StrmDirCache{},
|
||||
}
|
||||
&StrmDownloadTask{},
|
||||
&StrmUploadTask{},
|
||||
&StrmDirCache{},
|
||||
&ScrapeTask{},
|
||||
&EmbyMount{},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
+12
-12
@@ -13,6 +13,7 @@ const (
|
||||
StrmProviderCloudDrive = "clouddrive2" // CloudDrive2(WebDAV 桥接)
|
||||
StrmProviderOpenList = "openlist" // OpenList / AList 兼容桥接
|
||||
StrmProviderLocal = "local" // 本地目录(无账号)
|
||||
StrmProviderEmbyRemote = "emby_remote" // 远程 Emby 服务器(API 网关聚合挂载,不走 STRM 同步)
|
||||
)
|
||||
|
||||
// StrmAccount 是一个网盘账号(STRM 同步数据源凭据)。
|
||||
@@ -37,17 +38,17 @@ type StrmSyncPath struct {
|
||||
RemotePath string `gorm:"size:1024" json:"remote_path"` // 远端目录:115=目录ID,OpenList/CD2=路径,local=源目录
|
||||
LocalPath string `gorm:"size:1024" json:"local_path"` // STRM/元数据本地输出目录
|
||||
// STRM 链接配置(空值继承全局 strm.* 设置)
|
||||
StrmBaseURL string `gorm:"size:512" json:"strm_base_url"` // 覆盖 strm.base_url
|
||||
VideoExt string `gorm:"size:512" json:"video_ext"` // 逗号分隔,覆盖 strm.video_ext
|
||||
MetaExt string `gorm:"size:512" json:"meta_ext"` // 逗号分隔,覆盖 strm.meta_ext
|
||||
ExcludeName string `gorm:"size:512" json:"exclude_name"` // 逗号分隔,文件名包含即跳过
|
||||
MinVideoSizeMB int64 `json:"min_video_size_mb"` // 小于该大小(MB)的视频不生成 STRM
|
||||
AddPath int `json:"add_path"` // STRM 链接 path 参数:1=完整远端路径 2=仅文件名 3=不带
|
||||
DownloadMeta bool `gorm:"default:true" json:"download_meta"` // 同步时下载元数据文件(nfo/图片/字幕)
|
||||
UploadMeta bool `json:"upload_meta"` // 同步时把本地元数据上传到远端
|
||||
DeleteDir bool `json:"delete_dir"` // 清理多余文件时删除空目录
|
||||
Cron string `gorm:"size:128" json:"cron"` // 5 段 cron 表达式(可选)
|
||||
EnableCron bool `json:"enable_cron"` // 是否按 Cron 定时同步
|
||||
StrmBaseURL string `gorm:"size:512" json:"strm_base_url"` // 覆盖 strm.base_url
|
||||
VideoExt string `gorm:"size:512" json:"video_ext"` // 逗号分隔,覆盖 strm.video_ext
|
||||
MetaExt string `gorm:"size:512" json:"meta_ext"` // 逗号分隔,覆盖 strm.meta_ext
|
||||
ExcludeName string `gorm:"size:512" json:"exclude_name"` // 逗号分隔,文件名包含即跳过
|
||||
MinVideoSizeMB int64 `json:"min_video_size_mb"` // 小于该大小(MB)的视频不生成 STRM
|
||||
AddPath int `json:"add_path"` // STRM 链接 path 参数:1=完整远端路径 2=仅文件名 3=不带
|
||||
DownloadMeta bool `gorm:"default:true" json:"download_meta"` // 同步时下载元数据文件(nfo/图片/字幕)
|
||||
UploadMeta bool `json:"upload_meta"` // 同步时把本地元数据上传到远端
|
||||
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"`
|
||||
@@ -139,4 +140,3 @@ type StrmDirCache struct {
|
||||
DirID string `gorm:"size:128;index:idx_strm_dir_cache,priority:2" json:"dir_id"`
|
||||
Path string `gorm:"size:1024" json:"path"` // 相对根目录的路径
|
||||
}
|
||||
|
||||
|
||||
+36
-1
@@ -1,6 +1,10 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// User 是本地账户。第一个注册的管理员(或种子管理员)获得 "admin" 角色;
|
||||
// 其他所有用户默认为 "user"。
|
||||
@@ -17,6 +21,10 @@ type User struct {
|
||||
ForcePasswordReset bool `gorm:"default:false" json:"force_password_reset"`
|
||||
IsActive bool `gorm:"default:true" json:"is_active"`
|
||||
LastLoginAt *time.Time `json:"last_login_at,omitempty"`
|
||||
// AllowedLibraryIDs 存储管理员为该用户指定的受限可访问媒体库 ID 列表(JSON 字符串)。
|
||||
// 为空时代表不限制(全库可访问)。
|
||||
AllowedLibraryIDs string `gorm:"type:text" json:"-"`
|
||||
AllowedLibraryList []string `gorm:"-" json:"allowed_library_ids,omitempty"`
|
||||
// ExpiredAt is the account expiry time. Nil means the account never
|
||||
// expires. When set and in the past, the account is treated as expired
|
||||
// (login blocked) until an admin or a redemption code renews it.
|
||||
@@ -31,3 +39,30 @@ type User struct {
|
||||
RealtimeOnline bool `gorm:"-" json:"realtime_online,omitempty"`
|
||||
RealtimeDeviceCount int `gorm:"-" json:"realtime_device_count,omitempty"`
|
||||
}
|
||||
|
||||
// DecodeAllowedLibraryIDs 解析 AllowedLibraryIDs 字段。
|
||||
func (u *User) DecodeAllowedLibraryIDs() []string {
|
||||
if u == nil || strings.TrimSpace(u.AllowedLibraryIDs) == "" {
|
||||
return nil
|
||||
}
|
||||
var ids []string
|
||||
if err := json.Unmarshal([]byte(u.AllowedLibraryIDs), &ids); err != nil {
|
||||
return nil
|
||||
}
|
||||
var out []string
|
||||
for _, id := range ids {
|
||||
trimmed := strings.TrimSpace(id)
|
||||
if trimmed != "" {
|
||||
out = append(out, trimmed)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// PopulateComputedFields 填充非 DB 虚拟计算字段(如 AllowedLibraryList)。
|
||||
func (u *User) PopulateComputedFields() {
|
||||
if u == nil {
|
||||
return
|
||||
}
|
||||
u.AllowedLibraryList = u.DecodeAllowedLibraryIDs()
|
||||
}
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
)
|
||||
|
||||
// EmbyMountRepository 持久化远程 Emby 媒体库挂载。
|
||||
type EmbyMountRepository struct{ db *gorm.DB }
|
||||
|
||||
func (r *EmbyMountRepository) Create(ctx context.Context, m *model.EmbyMount) error {
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if m != nil && m.SortOrder == 0 {
|
||||
var maxSort int
|
||||
_ = tx.Model(&model.EmbyMount{}).Select("COALESCE(MAX(sort_order), -1)").Scan(&maxSort)
|
||||
m.SortOrder = maxSort + 1
|
||||
}
|
||||
return tx.Create(m).Error
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
func (r *EmbyMountRepository) CreateInBatches(ctx context.Context, mounts []*model.EmbyMount, batchSize int) error {
|
||||
if len(mounts) == 0 {
|
||||
return nil
|
||||
}
|
||||
if batchSize <= 0 {
|
||||
batchSize = 50
|
||||
}
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var maxSort int
|
||||
_ = tx.Model(&model.EmbyMount{}).Select("COALESCE(MAX(sort_order), -1)").Scan(&maxSort)
|
||||
for _, m := range mounts {
|
||||
if m != nil && m.SortOrder == 0 {
|
||||
maxSort++
|
||||
m.SortOrder = maxSort
|
||||
}
|
||||
}
|
||||
return tx.CreateInBatches(mounts, batchSize).Error
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
func (r *EmbyMountRepository) FindByID(ctx context.Context, id string) (*model.EmbyMount, error) {
|
||||
var m model.EmbyMount
|
||||
err := r.db.WithContext(ctx).Where("id = ?", id).First(&m).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &m, nil
|
||||
}
|
||||
|
||||
func (r *EmbyMountRepository) List(ctx context.Context) ([]model.EmbyMount, error) {
|
||||
var rows []model.EmbyMount
|
||||
err := r.db.WithContext(ctx).Order("sort_order asc, created_at asc").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
func (r *EmbyMountRepository) ListByAccountID(ctx context.Context, accountID string) ([]model.EmbyMount, error) {
|
||||
var rows []model.EmbyMount
|
||||
err := r.db.WithContext(ctx).Where("account_id = ?", accountID).Order("sort_order asc, created_at asc").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
func (r *EmbyMountRepository) SetSortOrder(ctx context.Context, ids []string) error {
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
for i, id := range ids {
|
||||
if err := tx.Model(&model.EmbyMount{}).Where("id = ?", id).
|
||||
Update("sort_order", i).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
func (r *EmbyMountRepository) CountByAccountID(ctx context.Context, accountID string) (int64, error) {
|
||||
var count int64
|
||||
err := r.db.WithContext(ctx).Model(&model.EmbyMount{}).Where("account_id = ?", accountID).Count(&count).Error
|
||||
return count, err
|
||||
}
|
||||
|
||||
func (r *EmbyMountRepository) Update(ctx context.Context, m *model.EmbyMount) error {
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Model(&model.EmbyMount{}).Where("id = ?", m.ID).Updates(map[string]any{
|
||||
"name": m.Name,
|
||||
"proxy_play": m.ProxyPlay,
|
||||
"enabled": m.Enabled,
|
||||
"remote_view_id": m.RemoteViewID,
|
||||
"remote_view_name": m.RemoteViewName,
|
||||
"collection_type": m.CollectionType,
|
||||
"updated_at": time.Now(),
|
||||
}).Error
|
||||
})
|
||||
}
|
||||
|
||||
func (r *EmbyMountRepository) Delete(ctx context.Context, id string) error {
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.EmbyMount{}).Error
|
||||
})
|
||||
}
|
||||
|
||||
// DeleteByAccountID 删除账号下全部挂载(删除账号时级联清理)。
|
||||
func (r *EmbyMountRepository) DeleteByAccountID(ctx context.Context, accountID string) (int64, error) {
|
||||
var count int64
|
||||
err := withSQLiteBusyRetry(ctx, func() error {
|
||||
res := r.db.WithContext(ctx).Where("account_id = ?", accountID).Delete(&model.EmbyMount{})
|
||||
count = res.RowsAffected
|
||||
return res.Error
|
||||
})
|
||||
return count, err
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/database"
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
)
|
||||
|
||||
func TestEmbyMountSortOrderAndReorder(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := database.AutoMigrate(db); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := New(db)
|
||||
ctx := t.Context()
|
||||
|
||||
// 1. Create mounts and verify auto-assigned sort_order
|
||||
m1 := &model.EmbyMount{AccountID: "acct-1", RemoteViewID: "view-1", Name: "Mount 1"}
|
||||
m2 := &model.EmbyMount{AccountID: "acct-1", RemoteViewID: "view-2", Name: "Mount 2"}
|
||||
m3 := &model.EmbyMount{AccountID: "acct-1", RemoteViewID: "view-3", Name: "Mount 3"}
|
||||
|
||||
if err := repos.EmbyMount.Create(ctx, m1); err != nil {
|
||||
t.Fatalf("create m1: %v", err)
|
||||
}
|
||||
if err := repos.EmbyMount.Create(ctx, m2); err != nil {
|
||||
t.Fatalf("create m2: %v", err)
|
||||
}
|
||||
if err := repos.EmbyMount.Create(ctx, m3); err != nil {
|
||||
t.Fatalf("create m3: %v", err)
|
||||
}
|
||||
|
||||
if m1.SortOrder >= m2.SortOrder || m2.SortOrder >= m3.SortOrder {
|
||||
t.Fatalf("expected ascending sort order on create: m1=%d, m2=%d, m3=%d",
|
||||
m1.SortOrder, m2.SortOrder, m3.SortOrder)
|
||||
}
|
||||
|
||||
// 2. Query list and verify initial order
|
||||
list, err := repos.EmbyMount.List(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("list mounts: %v", err)
|
||||
}
|
||||
if len(list) != 3 || list[0].ID != m1.ID || list[1].ID != m2.ID || list[2].ID != m3.ID {
|
||||
t.Fatalf("unexpected list order: %+v", list)
|
||||
}
|
||||
|
||||
// 3. Reorder: m3, m1, m2
|
||||
if err := repos.EmbyMount.SetSortOrder(ctx, []string{m3.ID, m1.ID, m2.ID}); err != nil {
|
||||
t.Fatalf("SetSortOrder failed: %v", err)
|
||||
}
|
||||
|
||||
// 4. Query list again and verify updated order
|
||||
reordered, err := repos.EmbyMount.List(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("list mounts after reorder: %v", err)
|
||||
}
|
||||
if len(reordered) != 3 {
|
||||
t.Fatalf("expected 3 mounts, got %d", len(reordered))
|
||||
}
|
||||
if reordered[0].ID != m3.ID || reordered[1].ID != m1.ID || reordered[2].ID != m2.ID {
|
||||
t.Fatalf("expected order [m3, m1, m2], got: %s, %s, %s",
|
||||
reordered[0].ID, reordered[1].ID, reordered[2].ID)
|
||||
}
|
||||
if reordered[0].SortOrder != 0 || reordered[1].SortOrder != 1 || reordered[2].SortOrder != 2 {
|
||||
t.Fatalf("unexpected sort orders: %d, %d, %d",
|
||||
reordered[0].SortOrder, reordered[1].SortOrder, reordered[2].SortOrder)
|
||||
}
|
||||
}
|
||||
@@ -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).Unscoped().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
|
||||
|
||||
@@ -3,6 +3,8 @@ package repository
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"gorm.io/gorm"
|
||||
@@ -87,6 +89,18 @@ func (r *MediaRepository) ListByLibraryFiltered(ctx context.Context, libraryID s
|
||||
}
|
||||
|
||||
func (r *MediaRepository) ListByLibrariesFiltered(ctx context.Context, libraryIDs []string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
|
||||
items, total, err := r.listByLibrariesFiltered(ctx, libraryIDs, offset, limit, filter, true)
|
||||
return items, total, err
|
||||
}
|
||||
|
||||
// ListByLibrariesFilteredNoCount skips the COUNT query when the caller already
|
||||
// knows totals or only needs a bounded slice (e.g. home-page previews).
|
||||
func (r *MediaRepository) ListByLibrariesFilteredNoCount(ctx context.Context, libraryIDs []string, offset, limit int, filter MediaQueryFilter) ([]model.Media, error) {
|
||||
items, _, err := r.listByLibrariesFiltered(ctx, libraryIDs, offset, limit, filter, false)
|
||||
return items, err
|
||||
}
|
||||
|
||||
func (r *MediaRepository) listByLibrariesFiltered(ctx context.Context, libraryIDs []string, offset, limit int, filter MediaQueryFilter, withCount bool) ([]model.Media, int64, error) {
|
||||
var items []model.Media
|
||||
var total int64
|
||||
if len(libraryIDs) == 0 {
|
||||
@@ -99,8 +113,10 @@ func (r *MediaRepository) ListByLibrariesFiltered(ctx context.Context, libraryID
|
||||
q = q.Where("library_id IN ?", libraryIDs)
|
||||
}
|
||||
q = applyMediaQueryFilter(q, filter)
|
||||
if err := q.Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
if withCount {
|
||||
if err := q.Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
}
|
||||
// 多级排序消除"随机"观感:
|
||||
// 1. release_date desc — 精确上映/首播日期新→旧
|
||||
@@ -114,6 +130,104 @@ func (r *MediaRepository) ListByLibrariesFiltered(ctx context.Context, libraryID
|
||||
return items, total, err
|
||||
}
|
||||
|
||||
type rankedMediaRow struct {
|
||||
model.Media
|
||||
MmtlRN int `gorm:"column:mmtl_rn"`
|
||||
}
|
||||
|
||||
// ListRecentByLibraries returns up to perLibrary recent items for each library
|
||||
// in a single query using a window function (avoids N+1 on home preview).
|
||||
func (r *MediaRepository) ListRecentByLibraries(ctx context.Context, libraryIDs []string, perLibrary int, filter MediaQueryFilter) (map[string][]model.Media, error) {
|
||||
out := make(map[string][]model.Media, len(libraryIDs))
|
||||
if len(libraryIDs) == 0 || perLibrary <= 0 {
|
||||
return out, nil
|
||||
}
|
||||
|
||||
var libClause string
|
||||
var args []interface{}
|
||||
if len(libraryIDs) == 1 {
|
||||
libClause = "library_id = ?"
|
||||
args = append(args, libraryIDs[0])
|
||||
} else {
|
||||
libClause = "library_id IN ?"
|
||||
args = append(args, libraryIDs)
|
||||
}
|
||||
where := "deleted_at IS NULL AND " + libClause
|
||||
if filterSQL, filterArgs := mediaQueryFilterSQL(filter); filterSQL != "" {
|
||||
where += " AND " + filterSQL
|
||||
args = append(args, filterArgs...)
|
||||
}
|
||||
args = append(args, perLibrary)
|
||||
|
||||
sql := fmt.Sprintf(`
|
||||
SELECT * FROM (
|
||||
SELECT *, ROW_NUMBER() OVER (
|
||||
PARTITION BY library_id
|
||||
ORDER BY release_date DESC, year DESC, updated_at DESC, created_at DESC, id DESC
|
||||
) AS mmtl_rn
|
||||
FROM media
|
||||
WHERE %s
|
||||
) ranked
|
||||
WHERE mmtl_rn <= ?
|
||||
`, where)
|
||||
|
||||
var rows []rankedMediaRow
|
||||
if err := r.db.WithContext(ctx).Raw(sql, args...).Scan(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, row := range rows {
|
||||
out[row.LibraryID] = append(out[row.LibraryID], row.Media)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func mediaQueryFilterSQL(filter MediaQueryFilter) (string, []interface{}) {
|
||||
var parts []string
|
||||
var args []interface{}
|
||||
if !filter.IncludeNSFW {
|
||||
parts = append(parts, "nsfw = ?")
|
||||
args = append(args, false)
|
||||
}
|
||||
if len(filter.HiddenLibraryIDs) > 0 {
|
||||
parts = append(parts, "library_id NOT IN ?")
|
||||
args = append(args, filter.HiddenLibraryIDs)
|
||||
}
|
||||
if len(filter.AllowedLibraryIDs) > 0 {
|
||||
parts = append(parts, "library_id IN ?")
|
||||
args = append(args, filter.AllowedLibraryIDs)
|
||||
}
|
||||
return strings.Join(parts, " AND "), args
|
||||
}
|
||||
|
||||
type libraryCountRow struct {
|
||||
LibraryID string `gorm:"column:library_id"`
|
||||
Total int64 `gorm:"column:total"`
|
||||
}
|
||||
|
||||
// CountByLibraries returns a map of library_id -> total media count for the given library IDs.
|
||||
func (r *MediaRepository) CountByLibraries(ctx context.Context, libraryIDs []string, filter MediaQueryFilter) (map[string]int64, error) {
|
||||
out := make(map[string]int64, len(libraryIDs))
|
||||
if len(libraryIDs) == 0 {
|
||||
return out, nil
|
||||
}
|
||||
var rows []libraryCountRow
|
||||
q := r.db.WithContext(ctx).Model(&model.Media{}).
|
||||
Select("library_id, count(*) as total")
|
||||
if len(libraryIDs) == 1 {
|
||||
q = q.Where("library_id = ?", libraryIDs[0])
|
||||
} else {
|
||||
q = q.Where("library_id IN ?", libraryIDs)
|
||||
}
|
||||
q = applyMediaQueryFilter(q, filter)
|
||||
if err := q.Group("library_id").Scan(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, row := range rows {
|
||||
out[row.LibraryID] = row.Total
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// DeleteByLibrary purges all media tied to a library.
|
||||
func (r *MediaRepository) DeleteByLibrary(ctx context.Context, libraryID string) error {
|
||||
// FTS 行由 media 表上的触发器同步清理(物理删除触发 FTS 清理)。
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/database"
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
)
|
||||
|
||||
func TestListRecentByLibraries(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := database.AutoMigrate(db); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := New(db)
|
||||
|
||||
lib1 := model.Library{Name: "电影", Path: "/media/movies", Type: "movie", Enabled: true}
|
||||
if err := repos.Library.Create(t.Context(), &lib1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
lib2 := model.Library{Name: "动漫", Path: "/media/anime", Type: "anime", Enabled: true}
|
||||
if err := repos.Library.Create(t.Context(), &lib2); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
now := time.Date(2026, 7, 2, 12, 0, 0, 0, time.UTC)
|
||||
var rows []model.Media
|
||||
for i := 1; i <= 5; i++ {
|
||||
rows = append(rows, model.Media{
|
||||
Base: model.Base{ID: fmt.Sprintf("movie-%02d", i), CreatedAt: now.Add(time.Duration(i) * time.Hour)},
|
||||
LibraryID: lib1.ID,
|
||||
Title: fmt.Sprintf("电影%d", i),
|
||||
Path: fmt.Sprintf("/media/movies/电影%d/movie%d.mp4", i, i),
|
||||
})
|
||||
}
|
||||
for i := 1; i <= 8; i++ {
|
||||
rows = append(rows, model.Media{
|
||||
Base: model.Base{ID: fmt.Sprintf("anime-ep-%02d", i), CreatedAt: now.Add(time.Duration(i) * time.Minute)},
|
||||
LibraryID: lib2.ID,
|
||||
Title: fmt.Sprintf("某动漫 第%d集", i),
|
||||
Path: fmt.Sprintf("/media/anime/某动漫/Season 01/某动漫.S01E%02d.mp4", i),
|
||||
SeasonNum: 1,
|
||||
EpisodeNum: i,
|
||||
})
|
||||
}
|
||||
if err := repos.DB.Create(&rows).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
filter := MediaQueryFilter{IncludeNSFW: true}
|
||||
got, err := repos.Media.ListRecentByLibraries(t.Context(), []string{lib1.ID, lib2.ID}, 3, filter)
|
||||
if err != nil {
|
||||
t.Fatalf("ListRecentByLibraries failed: %v", err)
|
||||
}
|
||||
if len(got[lib1.ID]) != 3 {
|
||||
t.Fatalf("lib1 recent count = %d, want 3", len(got[lib1.ID]))
|
||||
}
|
||||
if len(got[lib2.ID]) != 3 {
|
||||
t.Fatalf("lib2 recent count = %d, want 3", len(got[lib2.ID]))
|
||||
}
|
||||
if got[lib1.ID][0].ID != "movie-05" {
|
||||
t.Fatalf("lib1 newest = %q, want movie-05", got[lib1.ID][0].ID)
|
||||
}
|
||||
}
|
||||
@@ -32,6 +32,8 @@ type Container struct {
|
||||
StrmDownload *StrmDownloadTaskRepository
|
||||
StrmUpload *StrmUploadTaskRepository
|
||||
StrmDirCache *StrmDirCacheRepository
|
||||
ScrapeTask *ScrapeTaskRepository
|
||||
EmbyMount *EmbyMountRepository
|
||||
}
|
||||
|
||||
// New 将每个 repository 连接到单个 *gorm.DB。
|
||||
@@ -60,5 +62,7 @@ func New(db *gorm.DB) *Container {
|
||||
StrmDownload: &StrmDownloadTaskRepository{db: db},
|
||||
StrmUpload: &StrmUploadTaskRepository{db: db},
|
||||
StrmDirCache: &StrmDirCacheRepository{db: db},
|
||||
ScrapeTask: &ScrapeTaskRepository{db: db},
|
||||
EmbyMount: &EmbyMountRepository{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
|
||||
}
|
||||
@@ -308,6 +308,66 @@ func (r *StrmDownloadTaskRepository) Delete(ctx context.Context, id string) erro
|
||||
})
|
||||
}
|
||||
|
||||
// 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) {
|
||||
var count int64
|
||||
@@ -555,6 +615,89 @@ func (r *StrmUploadTaskRepository) Delete(ctx context.Context, id string) 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
|
||||
}
|
||||
|
||||
// ClearDone 清空全部已完成上传任务。
|
||||
func (r *StrmUploadTaskRepository) ClearDone(ctx context.Context) (int64, error) {
|
||||
var count int64
|
||||
err := withSQLiteBusyRetry(ctx, func() error {
|
||||
res := r.db.WithContext(ctx).Unscoped().Where("status = ?", model.StrmTaskDone).Delete(&model.StrmUploadTask{})
|
||||
count = res.RowsAffected
|
||||
return res.Error
|
||||
})
|
||||
return count, err
|
||||
}
|
||||
|
||||
// ClearFinished 清空全部已完成与失败上传任务(包括已完成、失败及取消)。
|
||||
func (r *StrmUploadTaskRepository) 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.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}).
|
||||
Delete(&model.StrmUploadTask{})
|
||||
count = res.RowsAffected
|
||||
return res.Error
|
||||
})
|
||||
return count, err
|
||||
}
|
||||
|
||||
// ClearCanceled 清空全部已取消上传任务。
|
||||
func (r *StrmUploadTaskRepository) ClearCanceled(ctx context.Context) (int64, error) {
|
||||
var count int64
|
||||
@@ -566,6 +709,27 @@ func (r *StrmUploadTaskRepository) ClearCanceled(ctx context.Context) (int64, er
|
||||
return count, err
|
||||
}
|
||||
|
||||
// RetryAllFailed 把所有失败任务重置回待处理,清空错误与重试计数。
|
||||
func (r *StrmUploadTaskRepository) RetryAllFailed(ctx context.Context) (int64, error) {
|
||||
var count int64
|
||||
err := withSQLiteBusyRetry(ctx, func() error {
|
||||
res := r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).
|
||||
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 批量取消所有排队中和进行中的任务。
|
||||
func (r *StrmUploadTaskRepository) CancelPending(ctx context.Context) (int64, error) {
|
||||
now := time.Now()
|
||||
@@ -657,5 +821,3 @@ func (r *StrmDirCacheRepository) DeleteBySyncPathID(ctx context.Context, syncPat
|
||||
return r.db.WithContext(ctx).Unscoped().Where("sync_path_id = ?", syncPathID).Delete(&model.StrmDirCache{}).Error
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -4,8 +4,11 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
@@ -42,3 +45,96 @@ func walkAndPrune(root string, cutoff time.Time) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// PruneImageCacheResult holds stats from an image cache prune operation.
|
||||
type PruneImageCacheResult struct {
|
||||
TotalFilesBefore int
|
||||
TotalBytesBefore int64
|
||||
DeletedFiles int
|
||||
FreedBytes int64
|
||||
RemainingBytes int64
|
||||
}
|
||||
|
||||
type imageCacheFileEntry struct {
|
||||
path string
|
||||
size int64
|
||||
modTime time.Time
|
||||
}
|
||||
|
||||
// PruneImageCache scans imagesDir for cached image files. If the total disk usage
|
||||
// exceeds maxSizeBytes, it removes files starting from the oldest (by ModTime)
|
||||
// until disk usage falls to or below targetSizeBytes (80% of maxSizeBytes).
|
||||
//
|
||||
// In-flight temporary files (*.tmp) are skipped to avoid corrupting concurrent writes.
|
||||
// Empty subdirectories left behind are best-effort removed.
|
||||
func PruneImageCache(imagesDir string, maxSizeBytes int64) (PruneImageCacheResult, error) {
|
||||
var result PruneImageCacheResult
|
||||
if imagesDir == "" || maxSizeBytes <= 0 {
|
||||
return result, nil
|
||||
}
|
||||
if _, err := os.Stat(imagesDir); err != nil {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
var (
|
||||
dirs []string
|
||||
entries []imageCacheFileEntry
|
||||
)
|
||||
|
||||
_ = filepath.Walk(imagesDir, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
if info.IsDir() {
|
||||
if path != imagesDir {
|
||||
dirs = append(dirs, path)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
// Skip temporary files created during image download.
|
||||
name := info.Name()
|
||||
if strings.HasSuffix(name, ".tmp") || strings.HasPrefix(name, "img-") && strings.Contains(name, ".tmp") {
|
||||
return nil
|
||||
}
|
||||
size := info.Size()
|
||||
result.TotalFilesBefore++
|
||||
result.TotalBytesBefore += size
|
||||
entries = append(entries, imageCacheFileEntry{
|
||||
path: path,
|
||||
size: size,
|
||||
modTime: info.ModTime(),
|
||||
})
|
||||
return nil
|
||||
})
|
||||
|
||||
result.RemainingBytes = result.TotalBytesBefore
|
||||
if result.TotalBytesBefore <= maxSizeBytes {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// High/Low watermark: prune down to 80% of max size to leave headroom
|
||||
// and prevent disk thrashing on consecutive writes.
|
||||
targetSizeBytes := maxSizeBytes * 80 / 100
|
||||
|
||||
sort.Slice(entries, func(i, j int) bool {
|
||||
return entries[i].modTime.Before(entries[j].modTime)
|
||||
})
|
||||
|
||||
for _, entry := range entries {
|
||||
if result.RemainingBytes <= targetSizeBytes {
|
||||
break
|
||||
}
|
||||
if err := os.Remove(entry.path); err == nil || errors.Is(err, os.ErrNotExist) {
|
||||
result.DeletedFiles++
|
||||
result.FreedBytes += entry.size
|
||||
result.RemainingBytes -= entry.size
|
||||
}
|
||||
}
|
||||
|
||||
// Clean up emptied subdirectories from deepest to shallowest.
|
||||
for i := len(dirs) - 1; i >= 0; i-- {
|
||||
_ = os.Remove(dirs[i])
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,162 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
func TestPruneImageCache_UnderLimit(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
|
||||
file1 := filepath.Join(dir, "img1")
|
||||
file2 := filepath.Join(dir, "img2")
|
||||
if err := os.WriteFile(file1, make([]byte, 100), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(file2, make([]byte, 200), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Max limit is 500 bytes, total is 300 bytes -> no prune
|
||||
res, err := PruneImageCache(dir, 500)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if res.DeletedFiles != 0 {
|
||||
t.Fatalf("expected 0 deleted files, got %d", res.DeletedFiles)
|
||||
}
|
||||
if res.TotalFilesBefore != 2 || res.TotalBytesBefore != 300 || res.RemainingBytes != 300 {
|
||||
t.Fatalf("unexpected stats: %+v", res)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPruneImageCache_OverLimitLRU(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
|
||||
now := time.Now()
|
||||
// Create 4 files of 100 bytes each, with distinct mtime
|
||||
fOldest := filepath.Join(dir, "oldest")
|
||||
fMidOld := filepath.Join(dir, "mid_old")
|
||||
fMidNew := filepath.Join(dir, "mid_new")
|
||||
fNewest := filepath.Join(dir, "newest")
|
||||
|
||||
for _, f := range []string{fOldest, fMidOld, fMidNew, fNewest} {
|
||||
if err := os.WriteFile(f, make([]byte, 100), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
_ = os.Chtimes(fOldest, now.Add(-4*time.Hour), now.Add(-4*time.Hour))
|
||||
_ = os.Chtimes(fMidOld, now.Add(-3*time.Hour), now.Add(-3*time.Hour))
|
||||
_ = os.Chtimes(fMidNew, now.Add(-2*time.Hour), now.Add(-2*time.Hour))
|
||||
_ = os.Chtimes(fNewest, now.Add(-1*time.Hour), now.Add(-1*time.Hour))
|
||||
|
||||
// Total = 400 bytes. Max limit = 300 bytes.
|
||||
// Target = 300 * 80 / 100 = 240 bytes.
|
||||
// Deleting oldest (100) brings total to 300 (> 240).
|
||||
// Deleting mid_old (100) brings total to 200 (<= 240).
|
||||
// Total deleted = 2 files (200 bytes), remaining = 200 bytes.
|
||||
res, err := PruneImageCache(dir, 300)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if res.DeletedFiles != 2 {
|
||||
t.Fatalf("expected 2 deleted files, got %d", res.DeletedFiles)
|
||||
}
|
||||
if res.FreedBytes != 200 {
|
||||
t.Fatalf("expected 200 freed bytes, got %d", res.FreedBytes)
|
||||
}
|
||||
if res.RemainingBytes != 200 {
|
||||
t.Fatalf("expected 200 remaining bytes, got %d", res.RemainingBytes)
|
||||
}
|
||||
|
||||
// Verify oldest and mid_old were deleted, mid_new and newest still exist
|
||||
if _, err := os.Stat(fOldest); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected oldest file to be deleted, got err=%v", err)
|
||||
}
|
||||
if _, err := os.Stat(fMidOld); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected mid_old file to be deleted, got err=%v", err)
|
||||
}
|
||||
if _, err := os.Stat(fMidNew); err != nil {
|
||||
t.Fatalf("expected mid_new file to exist, got err=%v", err)
|
||||
}
|
||||
if _, err := os.Stat(fNewest); err != nil {
|
||||
t.Fatalf("expected newest file to exist, got err=%v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPruneImageCache_SkipsTmpFiles(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
|
||||
fTmp := filepath.Join(dir, "img-123.tmp")
|
||||
fImg := filepath.Join(dir, "cached_img")
|
||||
|
||||
if err := os.WriteFile(fTmp, make([]byte, 500), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(fImg, make([]byte, 100), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Limit is 200 bytes. fTmp (500) is ignored, only fImg (100) is counted <= 200.
|
||||
res, err := PruneImageCache(dir, 200)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if res.DeletedFiles != 0 {
|
||||
t.Fatalf("expected 0 deleted files, got %d", res.DeletedFiles)
|
||||
}
|
||||
if _, err := os.Stat(fTmp); err != nil {
|
||||
t.Fatalf("expected tmp file to remain untouched, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPruneImageCache_ZeroOrNegativeLimit(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
f := filepath.Join(dir, "img")
|
||||
if err := os.WriteFile(f, make([]byte, 100), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
res, err := PruneImageCache(dir, 0)
|
||||
if err != nil || res.DeletedFiles != 0 {
|
||||
t.Fatalf("expected no-op for 0 limit, got %+v, err=%v", res, err)
|
||||
}
|
||||
|
||||
res, err = PruneImageCache(dir, -10)
|
||||
if err != nil || res.DeletedFiles != 0 {
|
||||
t.Fatalf("expected no-op for negative limit, got %+v, err=%v", res, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchedulerJobCleanImageCache(t *testing.T) {
|
||||
cacheRoot := t.TempDir()
|
||||
imagesDir := filepath.Join(cacheRoot, "images")
|
||||
if err := os.MkdirAll(imagesDir, 0o750); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
f := filepath.Join(imagesDir, "old_poster")
|
||||
if err := os.WriteFile(f, make([]byte, 2*1024*1024), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
scheduler := NewSchedulerService(zap.NewNop(), nil, nil, nil, nil, nil, cacheRoot)
|
||||
// Set limit to 1MB; our file is 2MB -> should be pruned
|
||||
scheduler.SetImagesMaxSizeMBProvider(func() int {
|
||||
return 1
|
||||
})
|
||||
|
||||
if err := scheduler.jobCleanImageCache(context.Background()); err != nil {
|
||||
t.Fatalf("jobCleanImageCache failed: %v", err)
|
||||
}
|
||||
|
||||
if _, err := os.Stat(f); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected file to be pruned, got err=%v", err)
|
||||
}
|
||||
}
|
||||
@@ -30,6 +30,7 @@ const (
|
||||
Type115 = "cloud115" // 115 网盘
|
||||
TypeCloudDrive2 = "clouddrive2" // CloudDrive2 桥接网盘
|
||||
TypeOpenList = "openlist" // OpenList / AList-compatible bridge
|
||||
TypeEmbyRemote = "emby_remote" // 远程 Emby 服务器(API 网关挂载)
|
||||
)
|
||||
|
||||
// ErrUnsupported is returned for an unknown provider type.
|
||||
@@ -37,11 +38,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"`
|
||||
MTime int64 `json:"mtime,omitempty"`
|
||||
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"`
|
||||
}
|
||||
@@ -101,6 +102,8 @@ func New(typ string, cfg map[string]any, client *http.Client) (Provider, error)
|
||||
return newCloudDrive2(cfg, client), nil
|
||||
case TypeOpenList:
|
||||
return newOpenList(cfg, client), nil
|
||||
case TypeEmbyRemote:
|
||||
return newEmby(cfg, client), nil
|
||||
default:
|
||||
return nil, ErrUnsupported
|
||||
}
|
||||
@@ -108,7 +111,7 @@ func New(typ string, cfg map[string]any, client *http.Client) (Provider, error)
|
||||
|
||||
// IsCloudType reports whether typ is a cloud-disk provider.
|
||||
func IsCloudType(typ string) bool {
|
||||
return typ == Type115 || typ == TypeCloudDrive2 || typ == TypeOpenList
|
||||
return typ == Type115 || typ == TypeCloudDrive2 || typ == TypeOpenList || typ == TypeEmbyRemote
|
||||
}
|
||||
|
||||
// str coerces a config value to a trimmed string.
|
||||
|
||||
@@ -0,0 +1,246 @@
|
||||
// Emby remote provider: exposes a remote Emby server through the same
|
||||
// Provider interface used by cloud disks, so account CRUD / connectivity
|
||||
// test / directory browser work unchanged. This is a thin adapter — the
|
||||
// federated Emby API aggregation (Views / Items / PlaybackInfo / streaming
|
||||
// proxy) lives in service.EmbyRemoteService and does not go through the
|
||||
// cloud-disk sync machinery.
|
||||
package cloud
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Emby 远程挂载类型(service 层聚合走 EmbyRemoteService,不走 STRM 同步)。
|
||||
|
||||
// embyProvider implements Provider against a remote Emby server using an
|
||||
// api_key (token) for authentication. DirectLink.Resolve returns the remote
|
||||
// stream URL; whether MMTL reverse-proxies the bytes is decided by the
|
||||
// emby.proxy_play account config (defaults to off).
|
||||
type embyProvider struct {
|
||||
base string // e.g. http://host:8096(自动补 /emby 前缀)
|
||||
username string
|
||||
password string
|
||||
token string // api_key
|
||||
userID string // 远程用户 Id
|
||||
proxyPlay bool
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
type embyUserPayload struct {
|
||||
Id string `json:"Id"`
|
||||
}
|
||||
|
||||
type embyLoginResponse struct {
|
||||
AccessToken string `json:"AccessToken"`
|
||||
User embyUserPayload `json:"User"`
|
||||
}
|
||||
|
||||
type embyPingResponse struct {
|
||||
ServerName string `json:"ServerName"`
|
||||
}
|
||||
|
||||
// newEmby builds the provider from the account config map.
|
||||
func newEmby(cfg map[string]any, client *http.Client) Provider {
|
||||
p := &embyProvider{
|
||||
base: strings.TrimRight(str(cfg["url"]), "/"),
|
||||
username: str(cfg["username"]),
|
||||
password: str(cfg["password"]),
|
||||
token: firstNonEmpty(str(cfg["api_key"]), str(cfg["token"])),
|
||||
userID: str(cfg["remote_user_id"]),
|
||||
proxyPlay: boolish(cfg["proxy_play"]),
|
||||
client: client,
|
||||
}
|
||||
if p.client == nil {
|
||||
p.client = &http.Client{Transport: &embyUATransport{base: http.DefaultTransport}}
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
// embyUATransport 给远程 Emby 请求注入浏览器 UA(防 Cloudflare 风控拦截)。
|
||||
type embyUATransport struct {
|
||||
base http.RoundTripper
|
||||
}
|
||||
|
||||
func (t *embyUATransport) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
if strings.TrimSpace(req.Header.Get("User-Agent")) == "" {
|
||||
req.Header.Set("User-Agent", defaultUA)
|
||||
}
|
||||
return t.base.RoundTrip(req)
|
||||
}
|
||||
|
||||
// embyBase normalizes the address so requests go to /emby/... endpoints.
|
||||
func (p *embyProvider) embyBase() string {
|
||||
base := strings.TrimRight(p.base, "/")
|
||||
if !strings.Contains(base, "/emby") {
|
||||
base += "/emby"
|
||||
}
|
||||
return base
|
||||
}
|
||||
|
||||
// externalBase 不追加 /emby(内嵌媒体资源 URL 使用 /emby 会更贴近习惯,此处
|
||||
// 与 embyBase 保持一致:所有端点统一以 /emby 开头)。
|
||||
func (p *embyProvider) apiBase() string { return p.embyBase() }
|
||||
|
||||
func (p *embyProvider) Type() string { return TypeEmbyRemote }
|
||||
|
||||
// Ping 验证地址连通性与凭据(/System/Info)。
|
||||
func (p *embyProvider) Ping(ctx context.Context) error {
|
||||
if p.base == "" {
|
||||
return errors.New("缺少 Emby 地址")
|
||||
}
|
||||
token, err := p.ensureToken(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return p.doJSON(ctx, http.MethodGet, "/System/Info", nil, token, &embyPingResponse{})
|
||||
}
|
||||
|
||||
// doJSON 向远程 Emby 发起带 api_key 的请求并解析 JSON 响应。
|
||||
func (p *embyProvider) doJSON(ctx context.Context, method, path string, body io.Reader, token string, out any) error {
|
||||
endpoint := p.apiBase() + path
|
||||
if token != "" {
|
||||
sep := "?"
|
||||
if strings.Contains(endpoint, "?") {
|
||||
sep = "&"
|
||||
}
|
||||
endpoint += sep + "api_key=" + url.QueryEscape(token)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, method, endpoint, body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("X-Emby-Token", token)
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
resp, err := p.client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode >= 300 {
|
||||
if resp.StatusCode == http.StatusUnauthorized {
|
||||
return ErrEmbyUnauthorized
|
||||
}
|
||||
data, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
||||
return fmt.Errorf("emby 请求失败(%d): %s", resp.StatusCode, strings.TrimSpace(string(data)))
|
||||
}
|
||||
if out == nil {
|
||||
return nil
|
||||
}
|
||||
return json.NewDecoder(resp.Body).Decode(out)
|
||||
}
|
||||
|
||||
// ErrEmbyUnauthorized 表示远程凭据失效(触发重新认证/打回测试)。
|
||||
var ErrEmbyUnauthorized = errors.New("emby 认证失败或凭据已失效")
|
||||
|
||||
// ensureToken 返回可用 api_key:已有则直接用,否则尝试账号密码认证。
|
||||
func (p *embyProvider) ensureToken(ctx context.Context) (string, error) {
|
||||
if strings.TrimSpace(p.token) != "" {
|
||||
return p.token, nil
|
||||
}
|
||||
if strings.TrimSpace(p.username) == "" {
|
||||
return "", errors.New("缺少 Emby 凭据(token 或 用户名/密码)")
|
||||
}
|
||||
payload := map[string]string{"Username": p.username, "Pw": p.password}
|
||||
data, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.apiBase()+"/Users/AuthenticateByName", strings.NewReader(string(data)))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("X-Emby-Authorization", `MediaBrowser Client="MMTL", Device="MMTL-Federated", DeviceId="mmtl-federated", Version="1.0"`)
|
||||
resp, err := p.client.Do(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode >= 300 {
|
||||
return "", fmt.Errorf("emby 登录失败(%d)", resp.StatusCode)
|
||||
}
|
||||
var login embyLoginResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&login); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if strings.TrimSpace(login.AccessToken) == "" {
|
||||
return "", errors.New("emby 登录成功但未返回 AccessToken")
|
||||
}
|
||||
p.token = login.AccessToken
|
||||
if login.User.Id != "" {
|
||||
p.userID = login.User.Id
|
||||
}
|
||||
return p.token, nil
|
||||
}
|
||||
|
||||
// embyItemSummary 目录浏览所需的最小 Emby 条目字段。
|
||||
type embyItemSummary struct {
|
||||
Id string `json:"Id"`
|
||||
Name string `json:"Name"`
|
||||
Type string `json:"Type"`
|
||||
IsFolder bool `json:"IsFolder"`
|
||||
ChildCount int `json:"ChildCount"`
|
||||
RunTimeTicks int64 `json:"RunTimeTicks"`
|
||||
}
|
||||
|
||||
type embyItemListResponse struct {
|
||||
Items []embyItemSummary `json:"Items"`
|
||||
}
|
||||
|
||||
// List 把远程媒体库(View)展开为目录树:dirID 为空=媒体库列表;否则返回该
|
||||
// 目录(Movie/Series/Season/Folder)下的条目。用于账号「浏览目录」调试入口。
|
||||
func (p *embyProvider) List(ctx context.Context, dirID string) ([]FileEntry, error) {
|
||||
token, err := p.ensureToken(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
userID := p.userID
|
||||
if userID == "" {
|
||||
userID = "0" // 某些 Emby 允许用 0 代表管理员
|
||||
}
|
||||
path := "/Users/" + url.PathEscape(userID) + "/Items"
|
||||
if dirID != "" {
|
||||
path += "?ParentId=" + url.QueryEscape(dirID)
|
||||
} else {
|
||||
path += "?IncludeItemTypes=CollectionFolder"
|
||||
}
|
||||
var out embyItemListResponse
|
||||
if err := p.doJSON(ctx, http.MethodGet, path, nil, token, &out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
entries := make([]FileEntry, 0, len(out.Items))
|
||||
for _, it := range out.Items {
|
||||
size := int64(0)
|
||||
if it.RunTimeTicks > 0 {
|
||||
size = it.RunTimeTicks / 10_000_000 // 秒
|
||||
}
|
||||
entries = append(entries, FileEntry{
|
||||
ID: it.Id,
|
||||
Name: it.Name,
|
||||
IsDir: it.IsFolder || it.Type != "Movie",
|
||||
Size: size,
|
||||
})
|
||||
}
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
// Resolve 返回远程 Emby 直链。Proxy=true 时由调用方(StrmService.ProxyDirect)
|
||||
// 反向代理流量;false 时 302 到直链。默认不代理(播放字节不经过 MMTL)。
|
||||
func (p *embyProvider) Resolve(ctx context.Context, fileRef string) (*DirectLink, error) {
|
||||
token, err := p.ensureToken(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u := p.apiBase() + "/Videos/" + url.PathEscape(fileRef) + "/stream"
|
||||
u += "?api_key=" + url.QueryEscape(token) + "&Static=true&MediaSourceId=" + url.QueryEscape(fileRef)
|
||||
return &DirectLink{URL: u, Headers: map[string]string{"X-Emby-Token": token}, Proxy: p.proxyPlay}, nil
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
package cloud
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// fakeEmbyServer 记录请求,按路径返回远程 Emby 风格响应。
|
||||
func fakeEmbyServer(t *testing.T) *httptest.Server {
|
||||
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch {
|
||||
case r.Method == http.MethodPost && r.URL.Path == "/emby/Users/AuthenticateByName":
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"AccessToken":"remote-token","User":{"Id":"user-9"}}`))
|
||||
case r.URL.Path == "/emby/System/Info":
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"ServerName":"RemoteEmby"}`))
|
||||
case r.URL.Path == "/emby/Users/user-9/Items" && r.URL.Query().Get("ParentId") == "":
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"Items":[{"Id":"view-1","Name":"Movies","Type":"CollectionFolder","IsFolder":true}]}`))
|
||||
case r.URL.Path == "/emby/Users/user-9/Items":
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"Items":[{"Id":"movie-1","Name":"Avatar","Type":"Movie","IsFolder":false}]}`))
|
||||
case strings.Contains(r.URL.Path, "/emby/Videos/movie-1/stream"):
|
||||
w.Header().Set("Content-Type", "video/mp4")
|
||||
_, _ = w.Write([]byte("fake-video-bytes"))
|
||||
default:
|
||||
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
|
||||
}
|
||||
}))
|
||||
}
|
||||
|
||||
func TestEmbyProviderPingAuthenticatesAndGetsToken(t *testing.T) {
|
||||
srv := fakeEmbyServer(t)
|
||||
defer srv.Close()
|
||||
|
||||
p, err := New(TypeEmbyRemote, map[string]any{
|
||||
"url": srv.URL,
|
||||
"username": "alice",
|
||||
"password": "secret",
|
||||
}, srv.Client())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := p.Ping(context.Background()); err != nil {
|
||||
t.Fatalf("ping: %v", err)
|
||||
}
|
||||
// 认证成功后 token 被记住,第二次 Ping 不应再走登录。
|
||||
if err := p.Ping(context.Background()); err != nil {
|
||||
t.Fatalf("ping 2: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbyProviderListViewsAndChildren(t *testing.T) {
|
||||
srv := fakeEmbyServer(t)
|
||||
defer srv.Close()
|
||||
|
||||
p, err := New(TypeEmbyRemote, map[string]any{
|
||||
"url": srv.URL,
|
||||
"api_key": "fixed-token",
|
||||
"remote_user_id": "user-9",
|
||||
}, srv.Client())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
root, err := p.List(context.Background(), "")
|
||||
if err != nil {
|
||||
t.Fatalf("list root: %v", err)
|
||||
}
|
||||
if len(root) != 1 || root[0].Name != "Movies" || !root[0].IsDir {
|
||||
t.Fatalf("root listing = %+v", root)
|
||||
}
|
||||
children, err := p.List(context.Background(), "view-1")
|
||||
if err != nil {
|
||||
t.Fatalf("list children: %v", err)
|
||||
}
|
||||
if len(children) != 1 || children[0].Name != "Avatar" || children[0].ID != "movie-1" {
|
||||
t.Fatalf("children = %+v", children)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbyProviderResolveDirectURLByDefault(t *testing.T) {
|
||||
srv := fakeEmbyServer(t)
|
||||
defer srv.Close()
|
||||
|
||||
p, err := New(TypeEmbyRemote, map[string]any{
|
||||
"url": srv.URL,
|
||||
"api_key": "fixed-token",
|
||||
"remote_user_id": "user-9",
|
||||
}, srv.Client())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
link, err := p.Resolve(context.Background(), "movie-1")
|
||||
if err != nil {
|
||||
t.Fatalf("resolve: %v", err)
|
||||
}
|
||||
if !strings.Contains(link.URL, "/emby/Videos/movie-1/stream") {
|
||||
t.Fatalf("url = %q", link.URL)
|
||||
}
|
||||
if !strings.Contains(link.URL, "api_key=fixed-token") {
|
||||
t.Fatalf("url missing api_key: %q", link.URL)
|
||||
}
|
||||
// 默认不代理播放流量。
|
||||
if link.Proxy {
|
||||
t.Fatal("emby remote must not proxy by default")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbyProviderResolveProxyWhenConfigured(t *testing.T) {
|
||||
srv := fakeEmbyServer(t)
|
||||
defer srv.Close()
|
||||
|
||||
p, err := New(TypeEmbyRemote, map[string]any{
|
||||
"url": srv.URL,
|
||||
"api_key": "fixed-token",
|
||||
"remote_user_id": "user-9",
|
||||
"proxy_play": "true",
|
||||
}, srv.Client())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
link, err := p.Resolve(context.Background(), "movie-1")
|
||||
if err != nil {
|
||||
t.Fatalf("resolve: %v", err)
|
||||
}
|
||||
if !link.Proxy {
|
||||
t.Fatal("proxy_play=true must mark link as proxied")
|
||||
}
|
||||
if link.URL == "" {
|
||||
t.Fatal("proxy link must still carry the remote URL")
|
||||
}
|
||||
}
|
||||
@@ -185,7 +185,7 @@ func (c *OpenClient) doJSON(ctx context.Context, method, rawURL string, form map
|
||||
// refresh_token 刷新后重试一次。刷新失败或重试后仍失败才返回,
|
||||
// 避免长时间同步因 token 过期而整体失败。
|
||||
if isTokenCode(base.Code) {
|
||||
if access && c.tryRefreshTokenLocked() {
|
||||
if access && c.tryRefreshTokenLocked(ctx) {
|
||||
continue
|
||||
}
|
||||
if access {
|
||||
@@ -259,19 +259,43 @@ func (c *OpenClient) doAuthJSONWithUA(ctx context.Context, method, rawURL string
|
||||
}
|
||||
|
||||
// tryRefreshTokenLocked 并发安全地刷新 access_token;成功返回 true(调用方
|
||||
// 应使用内存中的新 token 重试原请求)。refresh_token 已失效时也会清空内存 token。
|
||||
func (c *OpenClient) tryRefreshTokenLocked() bool {
|
||||
// 应使用内存中的新 token 重试原请求)。
|
||||
//
|
||||
// 对"refresh_token 本身已失效/被吊销"(IsRefreshTokenDead,如 40140114/116/119/120)
|
||||
// 这类不可恢复的错误直接放弃并清空内存 token(提示需重新授权)。
|
||||
// 对其它失败(网络瞬时抖动、刷新接口可重试错误码等)做指数退避重试几次再放弃,
|
||||
// 避免同步长任务中途 token 到期时恰好撞上一个短暂的刷新失败就整体失败。
|
||||
func (c *OpenClient) tryRefreshTokenLocked(ctx context.Context) bool {
|
||||
c.tokenMu.Lock()
|
||||
defer c.tokenMu.Unlock()
|
||||
token, err := c.RefreshToken(c.RefreshTokenStr)
|
||||
if err != nil {
|
||||
for attempt := 0; attempt < refreshAttempts; attempt++ {
|
||||
token, err := c.RefreshToken(c.RefreshTokenStr)
|
||||
if err == nil {
|
||||
c.SetAuthToken(token.AccessToken, token.RefreshToken)
|
||||
return true
|
||||
}
|
||||
if IsRefreshTokenDead(err) {
|
||||
c.SetAuthToken("", "")
|
||||
return false
|
||||
}
|
||||
// 可恢复失败:退避后重试。ctx 取消时立即放弃。
|
||||
if attempt < refreshAttempts-1 {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return false
|
||||
case <-time.After(refreshBackoff(attempt)):
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
c.SetAuthToken(token.AccessToken, token.RefreshToken)
|
||||
return true
|
||||
return false
|
||||
}
|
||||
|
||||
// refreshAttempts 是刷新 access_token 失败时的最大尝试次数(含首次)。
|
||||
const refreshAttempts = 3
|
||||
|
||||
// refreshBackoff 返回第 attempt 次(从 0 计)刷新失败后的退避时长(指数退避)。
|
||||
func refreshBackoff(attempt int) time.Duration {
|
||||
return time.Duration(200*(1<<attempt)) * time.Millisecond // 200ms, 400ms
|
||||
}
|
||||
|
||||
// IsThrottleCode 判断是否为限流错误码。
|
||||
@@ -281,7 +305,7 @@ func IsThrottleCode(code int) bool {
|
||||
|
||||
func isTokenCode(code int) bool {
|
||||
switch code {
|
||||
case AccessTokenAuthFail, AccessAuthInvalid, AccessTokenExpiryCode, RefreshTokenInvalid:
|
||||
case AccessTokenAuthFail, AccessAuthInvalid, AccessTokenExpiryCode, AccessTokenFormatInvalid, RefreshTokenInvalid:
|
||||
return true
|
||||
}
|
||||
return false
|
||||
|
||||
@@ -429,3 +429,46 @@ func TestRemoteFileDetailRelativePath(t *testing.T) {
|
||||
t.Errorf("d3.RelativePath = %q, want %q", got, "")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFsListRefreshContinue 验证 access_token 在请求中途过期(40140126)时:
|
||||
// 自动用 refresh_token 刷新得到新 token,然后对原请求重试成功(同步得以继续)。
|
||||
func TestFsListRefreshContinue(t *testing.T) {
|
||||
var filesCalls int
|
||||
var refreshCalls int
|
||||
mockAPI(t, func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.URL.Path {
|
||||
case "/open/refreshToken":
|
||||
refreshCalls++
|
||||
w.Write([]byte(`{"state":true,"data":{"access_token":"at2","refresh_token":"rt2","expires_in":7200}}`))
|
||||
case "/open/ufile/files":
|
||||
filesCalls++
|
||||
switch filesCalls {
|
||||
case 1:
|
||||
// 第一次用旧 access_token,返回过期错误,应触发刷新
|
||||
w.Write([]byte(`{"state":false,"code":40140126,"message":"access_token 校验失败"}`))
|
||||
default:
|
||||
// 刷新后续请求应使用新 access_token
|
||||
if got := r.Header.Get("Authorization"); got != "Bearer at2" {
|
||||
t.Errorf("retried request auth = %q, want Bearer at2", got)
|
||||
}
|
||||
w.Write([]byte(`{"state":true,"path":[],"data":[{"fid":"200","fc":"1","fn":"a.mkv","fs":123,"pc":"pickA"}]}`))
|
||||
}
|
||||
default:
|
||||
t.Errorf("unexpected path %s", r.URL.Path)
|
||||
}
|
||||
})
|
||||
c := NewOpenClient("100195125", "at1", "rt1")
|
||||
files, _, err := c.GetFsList(context.Background(), "0", 0, 100)
|
||||
if err != nil {
|
||||
t.Fatalf("expected sync to continue after refresh, got error: %v", err)
|
||||
}
|
||||
if filesCalls != 2 {
|
||||
t.Fatalf("want 2 files calls (original + retried), got %d", filesCalls)
|
||||
}
|
||||
if refreshCalls == 0 {
|
||||
t.Fatal("expected refresh_token to be used once")
|
||||
}
|
||||
if len(files) != 1 {
|
||||
t.Fatalf("want 1 file, got %d", len(files))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -5,6 +5,8 @@ package cloud115
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
@@ -181,6 +183,24 @@ func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input O
|
||||
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),
|
||||
@@ -188,8 +208,8 @@ func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input O
|
||||
CompleteMultipartUpload: &oss.CompleteMultipartUpload{
|
||||
Parts: completeParts,
|
||||
},
|
||||
Callback: oss.Ptr(input.Callback),
|
||||
CallbackVar: oss.Ptr(input.CallbackVar),
|
||||
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)
|
||||
|
||||
@@ -40,7 +40,9 @@ func FileSHA1Partial(path string, start, end int64) (string, error) {
|
||||
}
|
||||
length := end - start + 1
|
||||
h := sha1.New()
|
||||
if _, err := io.CopyN(h, f, length); err != nil {
|
||||
// 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
|
||||
|
||||
@@ -1,9 +1,14 @@
|
||||
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) {
|
||||
@@ -39,6 +44,26 @@ func TestFileSHA1Partial(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// 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 {
|
||||
@@ -92,3 +117,81 @@ func TestBaseNameOf(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -62,4 +62,4 @@ func xorDecode(hexStr string) string {
|
||||
func dandanplaySignature(appID, appSecret string, ts int64, path string) string {
|
||||
sum := sha256.Sum256([]byte(appID + strconv.FormatInt(ts, 10) + path + appSecret))
|
||||
return base64.StdEncoding.EncodeToString(sum[:])
|
||||
}
|
||||
}
|
||||
|
||||
@@ -81,4 +81,4 @@ func TestDanmakuCredentialsSelection(t *testing.T) {
|
||||
require.False(t, ok)
|
||||
require.Empty(t, id)
|
||||
require.Empty(t, key)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -278,69 +278,171 @@ 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")
|
||||
// 视频即便能命中 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)
|
||||
// 官方服务同时提供 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)
|
||||
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, "自动识别弹幕")
|
||||
// 1) 默认自动识别:命中 Hash 识别
|
||||
resAuto, err := svc.Fetch(ctx, "mManual", "", "")
|
||||
require.NoError(t, err)
|
||||
require.True(t, matchCalled)
|
||||
require.Equal(t, "hash", resAuto.MatchMode)
|
||||
require.Equal(t, int64(25484), resAuto.EpisodeID)
|
||||
require.Contains(t, resAuto.Raw, "自动识别弹幕")
|
||||
|
||||
// 2) 用户传入手动搜索关键词:跳过 Hash 识别,命中搜索结果动画B
|
||||
resManual, err := svc.Fetch(ctx, "mManual", "手动搜索动画B", "")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "search", resManual.MatchMode)
|
||||
require.Equal(t, int64(99999), resManual.EpisodeID)
|
||||
require.Equal(t, "手动搜索动画B", resManual.AnimeTitle)
|
||||
require.Contains(t, resManual.Raw, "手动搜索弹幕")
|
||||
}
|
||||
// 2) 用户传入手动搜索关键词:跳过 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, "手动搜索弹幕")
|
||||
}
|
||||
|
||||
// Emby 远程挂载条目:通过伪装 ID 解析出流直链,通过 Range 提取 16MB 前缀计算 hash 并匹配弹幕。
|
||||
func TestDanmakuFetchEmbyRemoteHashViaDirectLink(t *testing.T) {
|
||||
content := bytes.Repeat([]byte("emby-remote-video-bytes-9876543210"), 300)
|
||||
sum := md5.Sum(content)
|
||||
wantHash := hex.EncodeToString(sum[:])
|
||||
|
||||
var gotRange string
|
||||
var rangeHits int
|
||||
rangeSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
rangeHits++
|
||||
gotRange = r.Header.Get("Range")
|
||||
w.Header().Set("Content-Type", "application/octet-stream")
|
||||
_, _ = w.Write(content)
|
||||
}))
|
||||
t.Cleanup(rangeSrv.Close)
|
||||
|
||||
var seen string
|
||||
official := danmakuOfficialServer(t,
|
||||
`{"success":true,"isMatched":true,"matches":[{"episodeId":25484,"animeId":2001,"animeTitle":"芙莉莲","episodeTitle":"第1话"}]}`,
|
||||
`<?xml version="1.0"?><i><d p="1.2,1,16777215,user1">Emby远程弹幕命中</d></i>`,
|
||||
&seen)
|
||||
overrideDanmakuOfficialBase(t, official.URL)
|
||||
|
||||
remoteMediaID := EncodeEmbyRemoteID("mount-123", "remote-item-456")
|
||||
svc := newDanmakuTestService(t)
|
||||
svc.SetRemoteMediaResolver(func(_ context.Context, encodedID string) (*model.Media, string, error) {
|
||||
require.Equal(t, remoteMediaID, encodedID)
|
||||
return &model.Media{
|
||||
Base: model.Base{ID: remoteMediaID},
|
||||
Title: "葬送的芙莉莲",
|
||||
EpisodeTitle: "第1话",
|
||||
EpisodeNum: 1,
|
||||
Path: "/mnt/emby/anime/Frieren/S01E01.mkv",
|
||||
SizeBytes: int64(len(content)),
|
||||
DurationSec: 1400,
|
||||
}, rangeSrv.URL, nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
res, err := svc.Fetch(ctx, remoteMediaID, "", "")
|
||||
require.NoError(t, err)
|
||||
require.True(t, res.Enabled)
|
||||
require.Equal(t, "hash", res.MatchMode)
|
||||
require.Equal(t, int64(25484), res.EpisodeID)
|
||||
require.Equal(t, "芙莉莲", res.AnimeTitle)
|
||||
require.Contains(t, res.Raw, "Emby远程弹幕命中")
|
||||
require.Contains(t, gotRange, "bytes=0-")
|
||||
require.Contains(t, seen, `"fileHash":"`+wantHash+`"`)
|
||||
require.Contains(t, seen, `"fileName":"`+url.QueryEscape("S01E01")+`"`)
|
||||
require.Equal(t, 1, rangeHits)
|
||||
|
||||
// 第二次拉取验证 hashCache 命中,不重复请求 rangeSrv
|
||||
res2, err := svc.Fetch(ctx, remoteMediaID, "", "")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "hash", res2.MatchMode)
|
||||
require.Equal(t, 1, rangeHits)
|
||||
}
|
||||
|
||||
// Emby 远程直链拉取失败时(如网络异常),能平滑降级走番剧原名/标题关键词搜索。
|
||||
func TestDanmakuFetchEmbyRemoteStreamFailedFallsBackToSearch(t *testing.T) {
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/api/v2/search/episodes", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
// 文件名搜索 ep01 时无结果,模拟文件名未匹配
|
||||
if r.URL.Query().Get("anime") == "ep01" {
|
||||
fmt.Fprint(w, `{"hasMore":false,"animes":[]}`)
|
||||
return
|
||||
}
|
||||
// 降级到番剧名搜索命中
|
||||
fmt.Fprint(w, `{"hasMore":false,"animes":[{"animeId":3001,"animeTitle":"降级搜索番剧","episodes":[{"episodeId":7799,"episodeTitle":"第1话"}]}]}`)
|
||||
})
|
||||
mux.HandleFunc("/api/v2/comment/7799", 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.8,1,16777215,user1">降级搜索弹幕</d></i>`)
|
||||
})
|
||||
official := httptest.NewServer(mux)
|
||||
t.Cleanup(official.Close)
|
||||
overrideDanmakuOfficialBase(t, official.URL)
|
||||
|
||||
remoteMediaID := EncodeEmbyRemoteID("mount-123", "remote-item-789")
|
||||
svc := newDanmakuTestService(t)
|
||||
// 返回一个不存在的流服务地址模拟 Range 拉取失败
|
||||
svc.SetRemoteMediaResolver(func(_ context.Context, encodedID string) (*model.Media, string, error) {
|
||||
return &model.Media{
|
||||
Base: model.Base{ID: remoteMediaID},
|
||||
Title: "降级搜索番剧",
|
||||
EpisodeNum: 1,
|
||||
Path: "/mnt/emby/anime/fallback/ep01.mkv",
|
||||
DurationSec: 1200,
|
||||
}, "http://127.0.0.1:1/invalid-stream", nil
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
res, err := svc.Fetch(ctx, remoteMediaID, "", "")
|
||||
require.NoError(t, err)
|
||||
require.True(t, res.Enabled)
|
||||
require.Equal(t, "search", res.MatchMode)
|
||||
require.Equal(t, int64(7799), res.EpisodeID)
|
||||
require.Equal(t, "降级搜索番剧", res.AnimeTitle)
|
||||
require.Contains(t, res.Raw, "降级搜索弹幕")
|
||||
}
|
||||
|
||||
@@ -94,6 +94,10 @@ type DanmakuEpisode struct {
|
||||
EpisodeTitle string `json:"episodeTitle"`
|
||||
}
|
||||
|
||||
// DanmakuRemoteMediaResolver resolves an Emby remote pseudo-ID (e.g. embyremote~mount~id)
|
||||
// into a memory model.Media and a direct stream URL.
|
||||
type DanmakuRemoteMediaResolver func(ctx context.Context, encodedID string) (*model.Media, string, error)
|
||||
|
||||
// DanmakuService fetches danmaku for a media item through the dandanplay
|
||||
// protocol: match by 16MB-prefix hash, then search for an episode id by the
|
||||
// video's name, then fetch the comment library XML. The React player parses
|
||||
@@ -108,6 +112,10 @@ type DanmakuService struct {
|
||||
// StrmService.ResolvePlay; nil means strm sources are skipped.
|
||||
strmResolve func(ctx context.Context, provider string, q url.Values) (*StrmPlayResult, error)
|
||||
|
||||
// remoteResolve resolves an Emby remote pseudo-ID into *model.Media and
|
||||
// direct stream URL for range hashing.
|
||||
remoteResolve DanmakuRemoteMediaResolver
|
||||
|
||||
hashCacheMu sync.Mutex
|
||||
hashCache map[string]string // stamp → 16MB-prefix MD5
|
||||
}
|
||||
@@ -136,6 +144,14 @@ func (s *DanmakuService) SetStrmResolver(resolve func(ctx context.Context, provi
|
||||
}
|
||||
}
|
||||
|
||||
// SetRemoteMediaResolver wires the resolver used to fetch metadata and direct
|
||||
// stream URLs for Emby remote mounted media.
|
||||
func (s *DanmakuService) SetRemoteMediaResolver(resolve DanmakuRemoteMediaResolver) {
|
||||
if s != nil {
|
||||
s.remoteResolve = resolve
|
||||
}
|
||||
}
|
||||
|
||||
// Config reads danmaku settings from the runtime settings table.
|
||||
func (s *DanmakuService) Config(ctx context.Context) DanmakuRenderConfig {
|
||||
cfg := DanmakuRenderConfig{
|
||||
@@ -190,90 +206,94 @@ 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
|
||||
if id, parseErr := strconv.ParseInt(target, 10, 64); parseErr == nil {
|
||||
res.EpisodeID = id
|
||||
}
|
||||
res.MatchMode = "manual"
|
||||
return res, nil
|
||||
}
|
||||
|
||||
term, media, err := s.searchTerms(ctx, mediaID)
|
||||
// 手动指定弹幕库:跳过识别,直接拉取该库(自定义源失败回退官方)。
|
||||
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
|
||||
}
|
||||
manualKeyword := strings.TrimSpace(keyword) != ""
|
||||
if kw := strings.TrimSpace(keyword); kw != "" {
|
||||
term.name = kw
|
||||
res.Raw, res.SourceType = raw, st
|
||||
if id, parseErr := strconv.ParseInt(target, 10, 64); parseErr == nil {
|
||||
res.EpisodeID = id
|
||||
}
|
||||
if strings.TrimSpace(term.name) == "" {
|
||||
res.MatchMode = "manual"
|
||||
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(keyword 手动覆盖时跳过,直接走第 3 层)。
|
||||
if target == "" && !manualKeyword && media != nil && (media.Path != "" || IsEmbyRemoteID(media.ID)) {
|
||||
if hash, ok := s.mediaHash(ctx, media); ok {
|
||||
fileSize := media.SizeBytes
|
||||
if media.Path != "" && strings.EqualFold(filepath.Ext(media.Path), ".strm") {
|
||||
fileSize = 0 // strm 行的 SizeBytes 是文本大小,不是视频大小
|
||||
}
|
||||
matchName := danmakuMatchFileName(media.Path)
|
||||
if matchName == "" {
|
||||
matchName = term.name
|
||||
}
|
||||
matches, err := s.matchOfficial(ctx, matchName, 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
|
||||
}
|
||||
|
||||
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"
|
||||
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"
|
||||
}
|
||||
|
||||
raw, st, err := s.fetchCommentWithFallback(ctx, configured, official, target)
|
||||
if err != nil {
|
||||
@@ -311,6 +331,29 @@ type danmakuSearchTerms struct {
|
||||
// (movies / unknown) is left empty so the search does not filter by episode.
|
||||
func (s *DanmakuService) searchTerms(ctx context.Context, mediaID string) (danmakuSearchTerms, *model.Media, error) {
|
||||
var term danmakuSearchTerms
|
||||
if IsEmbyRemoteID(mediaID) {
|
||||
if s == nil || s.remoteResolve == nil {
|
||||
return term, nil, errors.New("remote emby resolver unavailable")
|
||||
}
|
||||
m, _, err := s.remoteResolve(ctx, mediaID)
|
||||
if err != nil || m == nil {
|
||||
if err != nil {
|
||||
return term, nil, err
|
||||
}
|
||||
return term, nil, errors.New("media not found")
|
||||
}
|
||||
if name := strings.TrimSpace(m.OriginalName); name != "" {
|
||||
term.name = name
|
||||
} else if name := strings.TrimSpace(m.Title); name != "" {
|
||||
term.name = name
|
||||
} else {
|
||||
term.name = danmakuMatchFileName(m.Path)
|
||||
}
|
||||
if m.EpisodeNum > 0 {
|
||||
term.episode = strconv.Itoa(m.EpisodeNum)
|
||||
}
|
||||
return term, m, nil
|
||||
}
|
||||
if s == nil || s.repo == nil || s.repo.Media == nil {
|
||||
return term, nil, errors.New("media repository unavailable")
|
||||
}
|
||||
@@ -484,7 +527,14 @@ func (s *DanmakuService) hashCachePut(stamp, hash string) {
|
||||
// ("xxx.mkv.strm") — so a second strip removes a real video extension only
|
||||
// (filepath.Ext would misread names like "xxx.第01话" as having an extension).
|
||||
func danmakuMatchFileName(path string) string {
|
||||
base := filepath.Base(path)
|
||||
if path == "" {
|
||||
return ""
|
||||
}
|
||||
clean := strings.ReplaceAll(path, "\\", "/")
|
||||
if idx := strings.LastIndex(clean, "/"); idx >= 0 {
|
||||
clean = clean[idx+1:]
|
||||
}
|
||||
base := filepath.Base(clean)
|
||||
if ext := filepath.Ext(base); ext != "" {
|
||||
base = strings.TrimSuffix(base, ext)
|
||||
}
|
||||
@@ -493,14 +543,20 @@ func danmakuMatchFileName(path string) string {
|
||||
base = strings.TrimSuffix(base, filepath.Ext(base))
|
||||
}
|
||||
}
|
||||
return base
|
||||
return strings.TrimSpace(base)
|
||||
}
|
||||
|
||||
// mediaHash returns the dandanplay match hash (MD5 of the first 16MB of the
|
||||
// video). Local videos are hashed straight from disk; .strm indirections are
|
||||
// resolved (local path / direct link) and only the 16MB prefix is downloaded.
|
||||
// video). Local videos are hashed straight from disk; .strm indirections and
|
||||
// remote Emby streams are range-fetched and only the 16MB prefix is downloaded.
|
||||
func (s *DanmakuService) mediaHash(ctx context.Context, media *model.Media) (string, bool) {
|
||||
if media == nil || media.Path == "" {
|
||||
if media == nil {
|
||||
return "", false
|
||||
}
|
||||
if IsEmbyRemoteID(media.ID) {
|
||||
return s.hashEmbyRemote(ctx, media)
|
||||
}
|
||||
if media.Path == "" {
|
||||
return "", false
|
||||
}
|
||||
if strings.EqualFold(filepath.Ext(media.Path), ".strm") {
|
||||
@@ -517,6 +573,43 @@ func (s *DanmakuService) mediaHash(ctx context.Context, media *model.Media) (str
|
||||
return s.hashLocalFile(media.Path)
|
||||
}
|
||||
|
||||
// hashEmbyRemote computes the 16MB-prefix MD5 of a remote Emby stream via HTTP Range.
|
||||
func (s *DanmakuService) hashEmbyRemote(ctx context.Context, media *model.Media) (string, bool) {
|
||||
if media == nil || media.ID == "" {
|
||||
return "", false
|
||||
}
|
||||
if h, ok := s.hashCacheGet("e|" + media.ID); ok {
|
||||
return h, true
|
||||
}
|
||||
if s.remoteResolve == nil {
|
||||
return "", false
|
||||
}
|
||||
_, streamURL, err := s.remoteResolve(ctx, media.ID)
|
||||
if err != nil || strings.TrimSpace(streamURL) == "" {
|
||||
if err != nil {
|
||||
s.log.Warn("danmaku emby stream url resolve failed, hash layer skipped",
|
||||
zap.String("media_id", media.ID), zap.Error(err))
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
body, err := s.openRangeBody(ctx, streamURL, nil)
|
||||
if err != nil || body == nil {
|
||||
if err != nil {
|
||||
s.log.Warn("danmaku emby range fetch failed, hash layer skipped",
|
||||
zap.String("media_id", media.ID), zap.Error(err))
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
defer body.Close()
|
||||
h := md5.New()
|
||||
if _, err := io.Copy(h, io.LimitReader(body, danmakuHashPrefixBytes)); err != nil {
|
||||
return "", false
|
||||
}
|
||||
hash := hex.EncodeToString(h.Sum(nil))
|
||||
s.hashCachePut("e|"+media.ID, hash)
|
||||
return hash, true
|
||||
}
|
||||
|
||||
// hashLocalFile computes the MD5 of the first 16MB of a local video, cached
|
||||
// by path+size+mtime so repeated danmaku loads skip the disk read.
|
||||
func (s *DanmakuService) hashLocalFile(path string) (string, bool) {
|
||||
@@ -694,7 +787,7 @@ func (s *DanmakuService) matchOfficial(ctx context.Context, fileName, fileHash s
|
||||
return nil, fmt.Errorf("danmaku match returned HTTP %d", resp.StatusCode)
|
||||
}
|
||||
var out struct {
|
||||
Success bool `json:"success"`
|
||||
Success bool `json:"success"`
|
||||
Matches []danmakuMatch `json:"matches"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &out); err != nil {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -11,6 +11,15 @@ import (
|
||||
|
||||
// ImageURL returns artwork for a media/series/season item id.
|
||||
func (e *EmbyService) ImageURL(ctx context.Context, id, imageType string) (string, error) {
|
||||
// 远程 Emby 条目:直接返回远程图片绝对地址,由 ImageProxy 拉取透传。
|
||||
if e.remote != nil && IsEmbyRemoteID(id) {
|
||||
mountID, remoteID, _ := DecodeEmbyRemoteID(id)
|
||||
mount, acct, _ := e.remote.ResolveMount(ctx, mountID)
|
||||
if mount == nil || acct == nil {
|
||||
return "", nil
|
||||
}
|
||||
return e.remote.RemoteImageURL(ctx, acct, remoteID, imageType)
|
||||
}
|
||||
pick := func(primary, backdrop string) string {
|
||||
switch strings.ToLower(imageType) {
|
||||
case "backdrop", "art":
|
||||
|
||||
@@ -48,6 +48,7 @@ type EmbyService struct {
|
||||
repo *repository.Container
|
||||
cache *RuntimeCacheService
|
||||
subtitle *SubtitleService
|
||||
remote *EmbyRemoteService // 远程 Emby 联邦聚合(可为 nil:未启用)
|
||||
|
||||
virtualMu sync.RWMutex
|
||||
virtualSeries map[string]embySeriesCacheEntry
|
||||
@@ -66,6 +67,14 @@ func NewEmbyService(cfg *config.Config, log *zap.Logger, repo *repository.Contai
|
||||
return &EmbyService{cfg: cfg, log: log, repo: repo}
|
||||
}
|
||||
|
||||
// SetEmbyRemote 注入远程 Emby 联邦聚合服务(nil 表示未启用)。
|
||||
func (e *EmbyService) SetEmbyRemote(remote *EmbyRemoteService) *EmbyService {
|
||||
if e != nil {
|
||||
e.remote = remote
|
||||
}
|
||||
return e
|
||||
}
|
||||
|
||||
func (e *EmbyService) SetRuntimeCache(cache *RuntimeCacheService) *EmbyService {
|
||||
if e != nil {
|
||||
e.cache = cache
|
||||
@@ -123,7 +132,8 @@ type embyVisibilityCacheEntry struct {
|
||||
|
||||
// Items paginates media in Emby's hierarchy. Episodic libraries are exposed as
|
||||
// Series -> Season -> Episode so Infuse/Vidhub/SenPlayer stop treating every
|
||||
// episode as a separate movie card.
|
||||
// episode as a separate movie card. 带 embyremote~ 前缀的 ParentID / 搜索自动
|
||||
// 路由到远程 Emby(联邦聚合,远程数据不落库)。
|
||||
func (e *EmbyService) Items(ctx context.Context, p ItemsParams) (map[string]any, error) {
|
||||
if p.Limit <= 0 || p.Limit > 500 {
|
||||
p.Limit = 50
|
||||
@@ -135,6 +145,33 @@ func (e *EmbyService) Items(ctx context.Context, p ItemsParams) (map[string]any,
|
||||
return emptyItemsEnvelope(p.StartIndex), nil
|
||||
}
|
||||
|
||||
if e.remote != nil {
|
||||
// 远程目录浏览:ParentId 带远程前缀 → 完整转发给远程 Emby 承接分页。
|
||||
if IsEmbyRemoteID(p.ParentID) {
|
||||
mountID, _, _ := DecodeEmbyRemoteID(p.ParentID)
|
||||
mount, acct, _ := e.remote.ResolveMount(ctx, mountID)
|
||||
if mount == nil || acct == nil {
|
||||
return emptyItemsEnvelope(p.StartIndex), nil
|
||||
}
|
||||
out, err := e.remote.RemoteItems(ctx, mount, acct, p)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := e.mergeRemoteUserData(ctx, p.UserID, out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
// 全局搜索:无 ParentId 且带搜索词 → 聚合本地 + 全部远程。
|
||||
if p.ParentID == "" && p.SearchTerm != "" {
|
||||
return e.aggregatedSearch(ctx, p)
|
||||
}
|
||||
}
|
||||
|
||||
if containsEmbyFilter(p.Filters, "IsResumable") {
|
||||
return e.resumableItems(ctx, p)
|
||||
}
|
||||
|
||||
if len(p.IDs) > 0 {
|
||||
items := make([]map[string]any, 0, len(p.IDs))
|
||||
for _, id := range p.IDs {
|
||||
@@ -212,3 +249,88 @@ func (e *EmbyService) Items(ctx context.Context, p ItemsParams) (map[string]any,
|
||||
}
|
||||
return e.mediaItems(ctx, p)
|
||||
}
|
||||
|
||||
// aggregatedSearch 把本地媒体库与全部启用的远程 Emby 的搜索结果合并为一个
|
||||
// 分页载荷。本地结果保持原有分页语义,远程各自取一页(Limit 同款)后按
|
||||
// SortBy 做稳定排序切片。
|
||||
func (e *EmbyService) aggregatedSearch(ctx context.Context, p ItemsParams) (map[string]any, error) {
|
||||
local, err := e.mediaItems(ctx, p)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
type remoteResult struct {
|
||||
items []any
|
||||
}
|
||||
mounts, aerr := e.remote.ListMounts(ctx)
|
||||
results := make([]remoteResult, 0, len(mounts))
|
||||
if aerr == nil {
|
||||
for i := range mounts {
|
||||
m := mounts[i]
|
||||
if !m.Enabled {
|
||||
continue
|
||||
}
|
||||
acct := e.remote.AccountByID(ctx, m.AccountID)
|
||||
if acct == nil {
|
||||
continue
|
||||
}
|
||||
// 按挂载逐个搜索:搜索结果归属明确(伪装 ID 正确),也天然只搜已
|
||||
// 挂载的媒体库。
|
||||
searchParams := p
|
||||
searchParams.ParentID = "" // RemoteSearchMount 内部设 ParentId
|
||||
remote, rerr := e.remote.RemoteSearchMount(ctx, &m, acct, p)
|
||||
if rerr != nil {
|
||||
if e.log != nil {
|
||||
e.log.Warn("remote emby search failed",
|
||||
zap.String("account", acct.Name), zap.Error(rerr))
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := e.mergeRemoteUserData(ctx, p.UserID, remote); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if raw, ok := remote["Items"].([]any); ok {
|
||||
results = append(results, remoteResult{items: raw})
|
||||
} else if rawMap, ok := remote["Items"].([]map[string]any); ok {
|
||||
converted := make([]any, 0, len(rawMap))
|
||||
for _, m := range rawMap {
|
||||
converted = append(converted, any(m))
|
||||
}
|
||||
results = append(results, remoteResult{items: converted})
|
||||
}
|
||||
}
|
||||
}
|
||||
items := make([]any, 0, len(localItemsAsAny(local))+len(results)*p.Limit)
|
||||
items = append(items, localItemsAsAny(local)...)
|
||||
for _, res := range results {
|
||||
items = append(items, res.items...)
|
||||
}
|
||||
return sliceSearchItems(items, p), nil
|
||||
}
|
||||
|
||||
func localItemsAsAny(envelope map[string]any) []any {
|
||||
if envelope == nil {
|
||||
return nil
|
||||
}
|
||||
if raw, ok := envelope["Items"].([]any); ok {
|
||||
return raw
|
||||
}
|
||||
if raw, ok := envelope["Items"].([]map[string]any); ok {
|
||||
converted := make([]any, 0, len(raw))
|
||||
for _, m := range raw {
|
||||
converted = append(converted, any(m))
|
||||
}
|
||||
return converted
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// sliceSearchItems 对合并结果按请求排序做简单归类后分页。远程返回已按远程
|
||||
// 排序规则排好,这里保持稳定顺序,只做首/尾切片,避免过度重排造成分页跳动。
|
||||
func sliceSearchItems(items []any, p ItemsParams) map[string]any {
|
||||
total := len(items)
|
||||
if p.StartIndex >= total {
|
||||
return map[string]any{"Items": []any{}, "TotalRecordCount": total, "StartIndex": p.StartIndex}
|
||||
}
|
||||
end := minInt(p.StartIndex+p.Limit, total)
|
||||
return map[string]any{"Items": items[p.StartIndex:end], "TotalRecordCount": total, "StartIndex": p.StartIndex}
|
||||
}
|
||||
|
||||
@@ -11,6 +11,25 @@ import (
|
||||
|
||||
// Item 单条目详情。
|
||||
func (e *EmbyService) Item(ctx context.Context, mediaID, userID string) (map[string]any, error) {
|
||||
if e == nil {
|
||||
return nil, nil
|
||||
}
|
||||
// 远程 Emby 条目:不查本地库,直接向远程转发(保持远程最新元数据)。
|
||||
if e.remote != nil && IsEmbyRemoteID(mediaID) {
|
||||
mountID, remoteID, _ := DecodeEmbyRemoteID(mediaID)
|
||||
mount, acct, _ := e.remote.ResolveMount(ctx, mountID)
|
||||
if mount == nil || acct == nil {
|
||||
return nil, nil
|
||||
}
|
||||
out, err := e.remote.RemoteItem(ctx, mount, acct, remoteID)
|
||||
if err != nil || out == nil {
|
||||
return out, err
|
||||
}
|
||||
if err := e.mergeRemoteUserData(ctx, userID, out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
if lib, err := e.repo.Library.FindByID(ctx, mediaID); err != nil {
|
||||
return nil, err
|
||||
} else if lib != nil {
|
||||
@@ -71,11 +90,26 @@ func (e *EmbyService) Item(ctx context.Context, mediaID, userID string) (map[str
|
||||
return e.itemPayload(ctx, m, fav, pos), nil
|
||||
}
|
||||
|
||||
// LatestItems 最近添加,全库或指定库。
|
||||
// LatestItems 最近添加,全库或指定库。远程媒体库(parentID 带前缀)直接透传远程。
|
||||
func (e *EmbyService) LatestItems(ctx context.Context, userID, parentID string, limit int) ([]map[string]any, error) {
|
||||
if limit <= 0 || limit > 100 {
|
||||
limit = 20
|
||||
}
|
||||
if e.remote != nil && IsEmbyRemoteID(parentID) {
|
||||
mountID, remoteParent, _ := DecodeEmbyRemoteID(parentID)
|
||||
mount, acct, _ := e.remote.ResolveMount(ctx, mountID)
|
||||
if mount == nil || acct == nil {
|
||||
return nil, nil
|
||||
}
|
||||
out, err := e.remote.RemoteLatest(ctx, mount, acct, remoteParent, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := e.mergeRemoteUserData(ctx, userID, out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
cacheKey := e.embyLatestCacheKey(userID, parentID, limit)
|
||||
var cached embyLatestCacheValue
|
||||
if e.cache != nil && e.cache.GetJSON(ctx, cacheKey, &cached) {
|
||||
@@ -143,41 +177,88 @@ func (e *EmbyService) latestSeriesItemsForLibrary(ctx context.Context, userID, l
|
||||
|
||||
// ResumeItems 列出有未完成播放进度的媒体。
|
||||
func (e *EmbyService) ResumeItems(ctx context.Context, userID string, limit int) (map[string]any, error) {
|
||||
if limit <= 0 || limit > 100 {
|
||||
limit = 20
|
||||
return e.resumableItems(ctx, ItemsParams{UserID: userID, Limit: limit})
|
||||
}
|
||||
|
||||
// resumableItems 返回未完成播放进度的媒体(包含本地媒体与挂载的远程媒体),支持分页。
|
||||
func (e *EmbyService) resumableItems(ctx context.Context, p ItemsParams) (map[string]any, error) {
|
||||
if p.Limit <= 0 || p.Limit > 100 {
|
||||
p.Limit = 50
|
||||
}
|
||||
if p.StartIndex < 0 {
|
||||
p.StartIndex = 0
|
||||
}
|
||||
if strings.TrimSpace(p.UserID) == "" {
|
||||
return map[string]any{"Items": []any{}, "TotalRecordCount": int64(0), "StartIndex": p.StartIndex}, nil
|
||||
}
|
||||
|
||||
var hist []model.PlaybackHistory
|
||||
if err := e.repo.DB.WithContext(ctx).
|
||||
Where("user_id = ? AND completed = ? AND position_ms > 0", userID, false).
|
||||
Order("watched_at desc").Limit(limit).Find(&hist).Error; err != nil {
|
||||
Where("user_id = ? AND completed = ? AND position_ms > 0", p.UserID, false).
|
||||
Order("watched_at desc").Find(&hist).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(hist) == 0 {
|
||||
return map[string]any{"Items": []any{}, "TotalRecordCount": 0}, nil
|
||||
return map[string]any{"Items": []any{}, "TotalRecordCount": int64(0), "StartIndex": p.StartIndex}, nil
|
||||
}
|
||||
ids := make([]string, 0, len(hist))
|
||||
posByID := map[string]int64{}
|
||||
|
||||
localIDs := make([]string, 0, len(hist))
|
||||
for _, h := range hist {
|
||||
ids = append(ids, h.MediaID)
|
||||
posByID[h.MediaID] = h.PositionMs
|
||||
}
|
||||
var medias []model.Media
|
||||
q := e.repo.DB.WithContext(ctx).Where("id IN ?", ids)
|
||||
q = e.applyUserMediaVisibility(ctx, q, userID)
|
||||
if err := q.Find(&medias).Error; err != nil {
|
||||
return nil, err
|
||||
if !IsEmbyRemoteID(h.MediaID) {
|
||||
localIDs = append(localIDs, h.MediaID)
|
||||
}
|
||||
}
|
||||
byID := map[string]*model.Media{}
|
||||
for i := range medias {
|
||||
byID[medias[i].ID] = &medias[i]
|
||||
if len(localIDs) > 0 {
|
||||
var medias []model.Media
|
||||
q := e.repo.DB.WithContext(ctx).Where("id IN ?", localIDs)
|
||||
q = e.applyUserMediaVisibility(ctx, q, p.UserID)
|
||||
if err := q.Find(&medias).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i := range medias {
|
||||
byID[medias[i].ID] = &medias[i]
|
||||
}
|
||||
}
|
||||
|
||||
items := make([]map[string]any, 0, len(hist))
|
||||
for _, h := range hist {
|
||||
if m, ok := byID[h.MediaID]; ok {
|
||||
items = append(items, e.itemPayload(ctx, m, false, posByID[h.MediaID]))
|
||||
if p.ParentID != "" && m.LibraryID != p.ParentID && m.SeriesID != p.ParentID {
|
||||
continue
|
||||
}
|
||||
items = append(items, e.itemPayload(ctx, m, false, h.PositionMs))
|
||||
continue
|
||||
}
|
||||
if e.remote == nil || !IsEmbyRemoteID(h.MediaID) {
|
||||
continue
|
||||
}
|
||||
mountID, remoteID, _ := DecodeEmbyRemoteID(h.MediaID)
|
||||
mount, acct, err := e.remote.ResolveMount(ctx, mountID)
|
||||
if err != nil || mount == nil || acct == nil {
|
||||
continue
|
||||
}
|
||||
item, err := e.remote.RemoteItem(ctx, mount, acct, remoteID)
|
||||
if err != nil || item == nil {
|
||||
continue
|
||||
}
|
||||
if p.ParentID != "" {
|
||||
parentID, _ := item["ParentId"].(string)
|
||||
seriesID, _ := item["SeriesId"].(string)
|
||||
if parentID != p.ParentID && seriesID != p.ParentID && mountID != p.ParentID {
|
||||
continue
|
||||
}
|
||||
}
|
||||
item["UserData"] = mergedRemoteUserData(item["UserData"], &h)
|
||||
items = append(items, item)
|
||||
}
|
||||
return map[string]any{"Items": items, "TotalRecordCount": len(items)}, nil
|
||||
|
||||
total := int64(len(items))
|
||||
if p.StartIndex >= len(items) {
|
||||
return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil
|
||||
}
|
||||
end := minInt(p.StartIndex+p.Limit, len(items))
|
||||
return map[string]any{"Items": items[p.StartIndex:end], "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil
|
||||
}
|
||||
|
||||
func (e *EmbyService) itemPayload(ctx context.Context, m *model.Media, fav bool, posMs int64) map[string]any {
|
||||
|
||||
@@ -199,7 +199,9 @@ func (e *EmbyService) appendSubtitleStreams(ctx context.Context, streams []map[s
|
||||
if e == nil || e.subtitle == nil || m == nil {
|
||||
return streams
|
||||
}
|
||||
tracks, err := e.subtitle.Discover(ctx, m.ID)
|
||||
// Emby 字幕只列表外挂字幕文件:云盘/strm 媒体的容器内嵌字幕不做服务端
|
||||
// 提取,客户端直连播放直链时自行解析。
|
||||
tracks, err := e.subtitle.DiscoverExternalOnly(ctx, m.ID)
|
||||
if err != nil || len(tracks) == 0 {
|
||||
return streams
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
@@ -14,7 +15,28 @@ import (
|
||||
)
|
||||
|
||||
// PlaybackInfo returns a PlaybackInfoResponse usable by Emby clients.
|
||||
// 远程 Emby 条目直接转发远程 PlaybackInfo,并按账号 proxy_play 配置决定
|
||||
// 播放地址指向远程(直连)还是 MMTL 本地代理端点。
|
||||
func (e *EmbyService) PlaybackInfo(ctx context.Context, mediaID, userID string) (map[string]any, error) {
|
||||
if e.remote != nil && IsEmbyRemoteID(mediaID) {
|
||||
mountID, remoteID, _ := DecodeEmbyRemoteID(mediaID)
|
||||
mount, acct, err := e.remote.ResolveMount(ctx, mountID)
|
||||
if err != nil {
|
||||
return nil, ErrEmbyRemoteNotFound
|
||||
}
|
||||
out, err := e.remote.RemotePlaybackInfo(ctx, mount, acct, remoteID, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if out == nil {
|
||||
return nil, ErrEmbyRemoteNotFound
|
||||
}
|
||||
if err := e.mergeRemoteUserData(ctx, userID, out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out["PlaySessionId"] = fmt.Sprintf("remote-%s-%d", mountID, time.Now().Unix())
|
||||
return out, nil
|
||||
}
|
||||
m, err := e.playableMedia(ctx, mediaID, userID)
|
||||
if err != nil || m == nil {
|
||||
return nil, err
|
||||
@@ -25,6 +47,65 @@ func (e *EmbyService) PlaybackInfo(ctx context.Context, mediaID, userID string)
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ErrEmbyRemoteNotFound 表示伪装 ID 对应的远程挂载账号不存在/已禁用。
|
||||
var ErrEmbyRemoteNotFound = fmt.Errorf("remote emby account not found")
|
||||
|
||||
// RemoteAccountByID 供 handler 层解码伪装 ID 后获取远程账号。
|
||||
func (e *EmbyService) RemoteAccountByID(ctx context.Context, accountID string) *model.StrmAccount {
|
||||
if e == nil || e.remote == nil {
|
||||
return nil
|
||||
}
|
||||
return e.remote.AccountByID(ctx, accountID)
|
||||
}
|
||||
|
||||
// ProxyRemoteVideoStream 反向代理远程 Emby 视频流(保留 Range)。
|
||||
func (e *EmbyService) ProxyRemoteVideoStream(ctx context.Context, w http.ResponseWriter, r *http.Request, mountID, remoteID string) error {
|
||||
if e == nil || e.remote == nil {
|
||||
return ErrEmbyRemoteNotFound
|
||||
}
|
||||
_, acct, err := e.remote.ResolveMount(ctx, mountID)
|
||||
if err != nil {
|
||||
return ErrEmbyRemoteNotFound
|
||||
}
|
||||
return e.remote.ProxyVideoStream(ctx, w, r, acct, remoteID)
|
||||
}
|
||||
|
||||
// ProxyRemoteSubtitle 反向代理远程 Emby 字幕流。
|
||||
func (e *EmbyService) ProxyRemoteSubtitle(ctx context.Context, w http.ResponseWriter, r *http.Request, mountID, remoteID, index string) error {
|
||||
if e == nil || e.remote == nil {
|
||||
return ErrSubtitleNotFound
|
||||
}
|
||||
_, acct, err := e.remote.ResolveMount(ctx, mountID)
|
||||
if err != nil {
|
||||
return ErrSubtitleNotFound
|
||||
}
|
||||
return e.remote.ProxySubtitle(ctx, w, r, acct, remoteID, index)
|
||||
}
|
||||
|
||||
// ProxyRemoteSetPlayed 把已看/未看状态透传到远程 Emby。
|
||||
func (e *EmbyService) ProxyRemoteSetPlayed(ctx context.Context, mountID, remoteID string, played bool) error {
|
||||
if e == nil || e.remote == nil {
|
||||
return ErrEmbyRemoteNotFound
|
||||
}
|
||||
_, acct, err := e.remote.ResolveMount(ctx, mountID)
|
||||
if err != nil {
|
||||
return ErrEmbyRemoteNotFound
|
||||
}
|
||||
return e.remote.ProxySetPlayed(ctx, acct, remoteID, played)
|
||||
}
|
||||
|
||||
// ProxyRemoteSetFavorite 把收藏/取消收藏状态透传到远程 Emby。
|
||||
func (e *EmbyService) ProxyRemoteSetFavorite(ctx context.Context, mountID, remoteID string, favorite bool) error {
|
||||
if e == nil || e.remote == nil {
|
||||
return ErrEmbyRemoteNotFound
|
||||
}
|
||||
_, acct, err := e.remote.ResolveMount(ctx, mountID)
|
||||
if err != nil {
|
||||
return ErrEmbyRemoteNotFound
|
||||
}
|
||||
return e.remote.ProxySetFavorite(ctx, acct, remoteID, favorite)
|
||||
}
|
||||
|
||||
// ServeSubtitleStream resolves the Emby /Videos/:id/Subtitles/:index/Stream
|
||||
// request to one of the media's sideloaded external subtitle tracks and writes
|
||||
// the original (unconverted) subtitle bytes to w — matching the source Codec
|
||||
@@ -32,10 +113,18 @@ func (e *EmbyService) PlaybackInfo(ctx context.Context, mediaID, userID string)
|
||||
// so the DeliveryUrl advertised in MediaStreams lines up exactly with the
|
||||
// served track: subtitles start at 1 when no audio stream is present, otherwise
|
||||
// at 2 (after Video 0 + Audio 1).
|
||||
//
|
||||
// 只服务外挂字幕文件(DiscoverExternalOnly):云盘/strm 媒体的容器内嵌字幕
|
||||
// 不做服务端提取,客户端直连播放直链时自行解析。
|
||||
func (e *EmbyService) ServeSubtitleStream(ctx context.Context, w io.Writer, mediaID, indexStr string, userID string) error {
|
||||
if e == nil || e.subtitle == nil {
|
||||
return ErrSubtitleUnavailable
|
||||
}
|
||||
if e.remote != nil && IsEmbyRemoteID(mediaID) {
|
||||
// 远程字幕由反向代理透传(需要 http.ResponseWriter 能力),handler 层
|
||||
// 已对远程 ID 走 ProxyRemoteSubtitle,这里不重复处理。
|
||||
return ErrSubtitleNotFound
|
||||
}
|
||||
m, err := e.playableMedia(ctx, mediaID, userID)
|
||||
if err != nil || m == nil {
|
||||
return ErrSubtitleNotFound
|
||||
@@ -44,7 +133,7 @@ func (e *EmbyService) ServeSubtitleStream(ctx context.Context, w io.Writer, medi
|
||||
if err != nil || index < 1 {
|
||||
return ErrSubtitleNotFound
|
||||
}
|
||||
tracks, err := e.subtitle.Discover(ctx, m.ID)
|
||||
tracks, err := e.subtitle.DiscoverExternalOnly(ctx, m.ID)
|
||||
if err != nil || len(tracks) == 0 {
|
||||
return ErrSubtitleNotFound
|
||||
}
|
||||
@@ -74,7 +163,7 @@ func (e *EmbyService) SubtitleStreamCodec(ctx context.Context, mediaID, indexStr
|
||||
if err != nil || index < 1 {
|
||||
return ""
|
||||
}
|
||||
tracks, err := e.subtitle.Discover(ctx, m.ID)
|
||||
tracks, err := e.subtitle.DiscoverExternalOnly(ctx, m.ID)
|
||||
if err != nil || len(tracks) == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
@@ -0,0 +1,921 @@
|
||||
// EmbyRemoteService 是「远程 Emby 联邦聚合」核心:把挂载的远程 Emby 服务器
|
||||
// 作为外部媒体源,通过 MMTL 的 Emby 兼容 API 透出。
|
||||
//
|
||||
// 设计要点:
|
||||
// - 远程媒体的元数据完全不落库:每次请求实时向远程 Emby 拉取;
|
||||
// - 条目 ID 用 embyremote~{accountID}~{remoteID} 伪装(见 emby_remote_ids.go),
|
||||
// 客户端拿伪装 ID 回来时按账号路由回远程;
|
||||
// - 播放分流由账号级 proxy_play 配置决定:
|
||||
// 不代理(默认)= MediaSource 下发热门远程绝对 URL,播放字节完全不经过 MMTL;
|
||||
// 代理 = 下发 MMTL 本地 /Videos/{encoded} 端点,由 ProxyVideoStream 反向拉流。
|
||||
//
|
||||
// 配置复用 STRM 账号体系(StrmAccount.Provider = emby_remote),CRUD/加密/连通
|
||||
// 测试全部走既有 /admin/strm/accounts 接口,不需要新增数据表。
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha1"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/config"
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
"github.com/ShukeBta/MMTL/internal/repository"
|
||||
)
|
||||
|
||||
// embyRemoteHTTPTimeout 远程 Emby 常规 API 请求超时(流式代理不在此列)。
|
||||
const embyRemoteHTTPTimeout = 15 * time.Second
|
||||
|
||||
// embyRemoteUA 桌面浏览器 UA:远程 Emby 前方若有 Cloudflare/WAF 会拦截
|
||||
// Go-http-client 等非浏览器 UA(403 error code: 1010),必须伪装浏览器。
|
||||
const embyRemoteUA = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/124.0 Safari/537.36"
|
||||
|
||||
// embyRemoteTransport 统一给远程请求注入浏览器 UA。
|
||||
type embyRemoteTransport struct {
|
||||
base http.RoundTripper
|
||||
}
|
||||
|
||||
func (t *embyRemoteTransport) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
if strings.TrimSpace(req.Header.Get("User-Agent")) == "" {
|
||||
req.Header.Set("User-Agent", embyRemoteUA)
|
||||
}
|
||||
return t.base.RoundTrip(req)
|
||||
}
|
||||
|
||||
// EmbyRemoteConfig 是一个远程 Emby 账号的解密配置。
|
||||
type EmbyRemoteConfig struct {
|
||||
BaseURL string // http://host:8096(无需 /emby 后缀)
|
||||
Username string
|
||||
Password string
|
||||
Token string // api_key(手动填写或自动认证获得)
|
||||
RemoteUserID string // 远程用户 Id(自动认证后回填)
|
||||
ProxyPlay bool // true=播放流量经 MMTL 反向代理;false=客户端直连远程
|
||||
}
|
||||
|
||||
// EmbyRemoteService 提供对远程 Emby 服务器的读写封装。
|
||||
type EmbyRemoteService struct {
|
||||
cfg *config.Config
|
||||
log *zap.Logger
|
||||
repo *repository.Container
|
||||
crypto *CryptoService
|
||||
http *http.Client
|
||||
cache *RuntimeCacheService
|
||||
}
|
||||
|
||||
// NewEmbyRemoteService 构造远程 Emby 聚合服务。
|
||||
func NewEmbyRemoteService(cfg *config.Config, log *zap.Logger, repo *repository.Container, crypto *CryptoService) *EmbyRemoteService {
|
||||
return &EmbyRemoteService{
|
||||
cfg: cfg,
|
||||
log: log,
|
||||
repo: repo,
|
||||
crypto: crypto,
|
||||
http: &http.Client{
|
||||
Timeout: embyRemoteHTTPTimeout,
|
||||
Transport: &embyRemoteTransport{base: http.DefaultTransport},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) SetRuntimeCache(cache *RuntimeCacheService) *EmbyRemoteService {
|
||||
if r != nil {
|
||||
r.cache = cache
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) remoteMediaCacheTTL() time.Duration {
|
||||
seconds := 15
|
||||
if r != nil && r.cfg != nil && r.cfg.Cache.MediaTTLSeconds > 0 {
|
||||
seconds = r.cfg.Cache.MediaTTLSeconds
|
||||
}
|
||||
return time.Duration(seconds) * time.Second
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) remoteCacheKey(parts ...string) string {
|
||||
sum := sha1.Sum([]byte(strings.Join(parts, "|")))
|
||||
return "media:embyremote:" + hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) invalidateRemoteMediaCache(ctx context.Context) {
|
||||
if r != nil && r.cache != nil {
|
||||
r.cache.DeletePrefix(ctx, "media:embyremote:")
|
||||
}
|
||||
}
|
||||
|
||||
// ListAccounts 返回全部启用的远程 Emby 挂载账号。
|
||||
func (r *EmbyRemoteService) ListAccounts(ctx context.Context) ([]model.StrmAccount, error) {
|
||||
accounts, err := r.repo.StrmAccount.List(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]model.StrmAccount, 0, len(accounts))
|
||||
for i := range accounts {
|
||||
if accounts[i].Provider == model.StrmProviderEmbyRemote {
|
||||
out = append(out, accounts[i])
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// AccountByID 按 ID 查找远程 Emby 挂载账号(不存在或类型不符返回 nil)。
|
||||
func (r *EmbyRemoteService) AccountByID(ctx context.Context, id string) *model.StrmAccount {
|
||||
if strings.TrimSpace(id) == "" {
|
||||
return nil
|
||||
}
|
||||
acct, err := r.repo.StrmAccount.FindByID(ctx, id)
|
||||
if err != nil || acct == nil {
|
||||
return nil
|
||||
}
|
||||
if acct.Provider != model.StrmProviderEmbyRemote || !acct.Enabled {
|
||||
return nil
|
||||
}
|
||||
return acct
|
||||
}
|
||||
|
||||
// ─── 媒体库挂载管理 ─────────────────────────────────────────────────────────────
|
||||
|
||||
// ListMounts 返回全部挂载。
|
||||
func (r *EmbyRemoteService) ListMounts(ctx context.Context) ([]model.EmbyMount, error) {
|
||||
return r.repo.EmbyMount.List(ctx)
|
||||
}
|
||||
|
||||
// ListMountsByAccount 返回指定账号的挂载。
|
||||
func (r *EmbyRemoteService) ListMountsByAccount(ctx context.Context, accountID string) ([]model.EmbyMount, error) {
|
||||
return r.repo.EmbyMount.ListByAccountID(ctx, accountID)
|
||||
}
|
||||
|
||||
// MountByID 按 ID 查挂载。
|
||||
func (r *EmbyRemoteService) MountByID(ctx context.Context, id string) (*model.EmbyMount, error) {
|
||||
return r.repo.EmbyMount.FindByID(ctx, id)
|
||||
}
|
||||
|
||||
// CreateMount 创建一个挂载(校验账号类型与远程 View 编号)。
|
||||
func (r *EmbyRemoteService) CreateMount(ctx context.Context, m *model.EmbyMount) (*model.EmbyMount, error) {
|
||||
if strings.TrimSpace(m.AccountID) == "" || strings.TrimSpace(m.RemoteViewID) == "" {
|
||||
return nil, errors.New("缺少账号或远程媒体库")
|
||||
}
|
||||
if r.AccountByID(ctx, m.AccountID) == nil {
|
||||
return nil, errors.New("远程 Emby 账号不存在或已禁用")
|
||||
}
|
||||
if err := r.repo.EmbyMount.Create(ctx, m); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r.invalidateRemoteMediaCache(ctx)
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// CreateMounts 批量创建挂载(幂等:已存在的远程库自动跳过)。
|
||||
func (r *EmbyRemoteService) CreateMounts(ctx context.Context, mounts []*model.EmbyMount) (int, error) {
|
||||
if len(mounts) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
existing, err := r.repo.EmbyMount.ListByAccountID(ctx, mounts[0].AccountID)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
have := make(map[string]bool, len(existing))
|
||||
for _, e := range existing {
|
||||
have[e.RemoteViewID] = true
|
||||
}
|
||||
fresh := make([]*model.EmbyMount, 0, len(mounts))
|
||||
for _, m := range mounts {
|
||||
if m == nil || have[m.RemoteViewID] {
|
||||
continue
|
||||
}
|
||||
fresh = append(fresh, m)
|
||||
}
|
||||
if len(fresh) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
if err := r.repo.EmbyMount.CreateInBatches(ctx, fresh, 50); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
r.invalidateRemoteMediaCache(ctx)
|
||||
return len(fresh), nil
|
||||
}
|
||||
|
||||
// UpdateMount 更新挂载(名称 / 代理 / 启用)。
|
||||
func (r *EmbyRemoteService) UpdateMount(ctx context.Context, id string, m *model.EmbyMount) (*model.EmbyMount, error) {
|
||||
existing, err := r.repo.EmbyMount.FindByID(ctx, id)
|
||||
if err != nil || existing == nil {
|
||||
return nil, errNotFoundOr(err, "挂载不存在")
|
||||
}
|
||||
existing.Name = strings.TrimSpace(m.Name)
|
||||
existing.ProxyPlay = m.ProxyPlay
|
||||
existing.Enabled = m.Enabled
|
||||
if err := r.repo.EmbyMount.Update(ctx, existing); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r.invalidateRemoteMediaCache(ctx)
|
||||
return existing, nil
|
||||
}
|
||||
|
||||
// DeleteMount 删除挂载。
|
||||
func (r *EmbyRemoteService) DeleteMount(ctx context.Context, id string) error {
|
||||
err := r.repo.EmbyMount.Delete(ctx, id)
|
||||
if err == nil {
|
||||
r.invalidateRemoteMediaCache(ctx)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// ReorderMounts 批量重排挂载媒体库顺序。
|
||||
func (r *EmbyRemoteService) ReorderMounts(ctx context.Context, ids []string) error {
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
if err := r.repo.EmbyMount.SetSortOrder(ctx, ids); err != nil {
|
||||
return err
|
||||
}
|
||||
r.invalidateRemoteMediaCache(ctx)
|
||||
return nil
|
||||
}
|
||||
|
||||
// FullMountAccount 把账号的全部远程媒体库(View)挂载进来(幂等,已存在跳过)。
|
||||
func (r *EmbyRemoteService) FullMountAccount(ctx context.Context, acct *model.StrmAccount, proxyPlayDefault bool) (int, error) {
|
||||
views, err := r.RemoteViews(ctx, acct)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
mounts := make([]*model.EmbyMount, 0, len(views))
|
||||
for _, v := range views {
|
||||
viewID := remoteItemString(v, "Id")
|
||||
if viewID == "" {
|
||||
continue
|
||||
}
|
||||
mounts = append(mounts, &model.EmbyMount{
|
||||
AccountID: acct.ID,
|
||||
RemoteViewID: viewID,
|
||||
RemoteViewName: remoteItemString(v, "Name"),
|
||||
CollectionType: remoteItemString(v, "CollectionType"),
|
||||
ProxyPlay: proxyPlayDefault,
|
||||
Enabled: true,
|
||||
})
|
||||
}
|
||||
return r.CreateMounts(ctx, mounts)
|
||||
}
|
||||
|
||||
// ResolveMount 按伪装 ID 的第一段(挂载 ID)解析挂载与其所属账号。
|
||||
// 远程条目/媒体库的伪装 ID 格式:embyremote~{mountID}~{remoteID}。
|
||||
func (r *EmbyRemoteService) ResolveMount(ctx context.Context, mountID string) (*model.EmbyMount, *model.StrmAccount, error) {
|
||||
mount, err := r.repo.EmbyMount.FindByID(ctx, mountID)
|
||||
if err != nil || mount == nil || !mount.Enabled {
|
||||
return nil, nil, errors.New("挂载不存在或已禁用")
|
||||
}
|
||||
acct := r.AccountByID(ctx, mount.AccountID)
|
||||
if acct == nil {
|
||||
return nil, nil, errors.New("远程 Emby 账号不存在或已禁用")
|
||||
}
|
||||
return mount, acct, nil
|
||||
}
|
||||
|
||||
// AutoSeedMounts 兼容迁移:已有 emby_remote 账号但没有任何挂载时,自动把
|
||||
// 其全部媒体库挂载进来(代理沿用账号旧配置),保证旧部署升级后媒体库不消失。
|
||||
// 幂等:每个账号只在挂载数为 0 时执行一次。
|
||||
func (r *EmbyRemoteService) AutoSeedMounts(ctx context.Context) {
|
||||
accounts, err := r.ListAccounts(ctx)
|
||||
if err != nil || len(accounts) == 0 {
|
||||
return
|
||||
}
|
||||
for i := range accounts {
|
||||
acct := &accounts[i]
|
||||
count, err := r.repo.EmbyMount.CountByAccountID(ctx, acct.ID)
|
||||
if err != nil || count > 0 {
|
||||
continue
|
||||
}
|
||||
cfg, cfgErr := r.configOf(acct)
|
||||
if cfgErr != nil {
|
||||
continue
|
||||
}
|
||||
n, seedErr := r.FullMountAccount(ctx, acct, cfg.ProxyPlay)
|
||||
if seedErr != nil {
|
||||
if r.log != nil {
|
||||
r.log.Warn("auto-seed emby mounts failed",
|
||||
zap.String("account", acct.Name), zap.Error(seedErr))
|
||||
}
|
||||
} else if n > 0 {
|
||||
if r.log != nil {
|
||||
r.log.Info("auto-seeded emby mounts",
|
||||
zap.String("account", acct.Name), zap.Int("mounts", n))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// configOf 解密账号配置。
|
||||
func (r *EmbyRemoteService) configOf(acct *model.StrmAccount) (*EmbyRemoteConfig, error) {
|
||||
raw := map[string]string{}
|
||||
if acct != nil && strings.TrimSpace(acct.Config) != "" {
|
||||
if err := json.Unmarshal([]byte(acct.Config), &raw); err != nil {
|
||||
return nil, fmt.Errorf("decode emby account config: %w", err)
|
||||
}
|
||||
}
|
||||
cfg := &EmbyRemoteConfig{
|
||||
BaseURL: strings.TrimRight(strings.TrimSpace(raw["url"]), "/"),
|
||||
Username: strings.TrimSpace(raw["username"]),
|
||||
Password: r.crypto.Decrypt(raw["password"]),
|
||||
Token: firstNonEmptyStr(r.crypto.Decrypt(raw["api_key"]), r.crypto.Decrypt(raw["token"])),
|
||||
RemoteUserID: strings.TrimSpace(raw["remote_user_id"]),
|
||||
ProxyPlay: parseBoolSetting(raw["proxy_play"], false),
|
||||
}
|
||||
if cfg.BaseURL == "" {
|
||||
return nil, errors.New("缺少 Emby 地址")
|
||||
}
|
||||
if !strings.HasPrefix(cfg.BaseURL, "http://") && !strings.HasPrefix(cfg.BaseURL, "https://") {
|
||||
return nil, errors.New("Emby 地址必须以 http:// 或 https:// 开头")
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func firstNonEmptyStr(values ...string) string {
|
||||
for _, v := range values {
|
||||
if strings.TrimSpace(v) != "" {
|
||||
return v
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// embyBase 把地址规范为不带尾部斜杠的 /emby 根。
|
||||
func (r *EmbyRemoteService) embyBase(cfg *EmbyRemoteConfig) string {
|
||||
base := strings.TrimRight(cfg.BaseURL, "/")
|
||||
if !strings.HasSuffix(base, "/emby") {
|
||||
base += "/emby"
|
||||
}
|
||||
return base
|
||||
}
|
||||
|
||||
// ensureToken 返回可用的 api_key:已有则直接用;否则用用户名/密码认证并回写
|
||||
// 数据库(自动获得的 token 与 remote_user_id 会加密保存在账号配置里)。
|
||||
func (r *EmbyRemoteService) ensureToken(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig) error {
|
||||
if strings.TrimSpace(cfg.Token) != "" {
|
||||
return nil
|
||||
}
|
||||
if strings.TrimSpace(cfg.Username) == "" || strings.TrimSpace(cfg.Password) == "" {
|
||||
return errors.New("缺少 Emby 凭据:请填写 api_key 或 用户名/密码")
|
||||
}
|
||||
body, _ := json.Marshal(map[string]string{"Username": cfg.Username, "Pw": cfg.Password})
|
||||
endpoint := r.embyBase(cfg) + "/Users/AuthenticateByName"
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, strings.NewReader(string(body)))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("X-Emby-Authorization", `MediaBrowser Client="MMTL", Device="MMTL-Federated", DeviceId="mmtl-federated", Version="1.0"`)
|
||||
resp, err := r.http.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("连接远程 Emby 失败: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode >= 300 {
|
||||
data, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
||||
return fmt.Errorf("远程 Emby 登录失败(%d): %s", resp.StatusCode, strings.TrimSpace(string(data)))
|
||||
}
|
||||
var login struct {
|
||||
AccessToken string `json:"AccessToken"`
|
||||
User struct {
|
||||
Id string `json:"Id"`
|
||||
} `json:"User"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&login); err != nil {
|
||||
return err
|
||||
}
|
||||
if strings.TrimSpace(login.AccessToken) == "" {
|
||||
return errors.New("远程 Emby 未返回 AccessToken")
|
||||
}
|
||||
cfg.Token = login.AccessToken
|
||||
if login.User.Id != "" {
|
||||
cfg.RemoteUserID = login.User.Id
|
||||
}
|
||||
return r.persistToken(ctx, acct, cfg)
|
||||
}
|
||||
|
||||
// persistToken 把认证得到的 token / user id 加密写回账号配置(下次请求免登录)。
|
||||
func (r *EmbyRemoteService) persistToken(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig) error {
|
||||
if acct == nil {
|
||||
return nil
|
||||
}
|
||||
raw := map[string]string{}
|
||||
if strings.TrimSpace(acct.Config) != "" {
|
||||
_ = json.Unmarshal([]byte(acct.Config), &raw)
|
||||
}
|
||||
raw["api_key"] = r.crypto.Encrypt(cfg.Token)
|
||||
raw["remote_user_id"] = cfg.RemoteUserID
|
||||
if strings.TrimSpace(raw["username"]) == "" {
|
||||
raw["username"] = cfg.Username
|
||||
}
|
||||
data, err := json.Marshal(raw)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
acct.Config = string(data)
|
||||
return r.repo.StrmAccount.Update(ctx, acct)
|
||||
}
|
||||
|
||||
// doGet 向远程 Emby 发起带 api_key 的 GET,把响应 JSON 解码到 out。
|
||||
// 401 时自动重认证一次再重试(凭据过期场景)。
|
||||
func (r *EmbyRemoteService) doGet(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig, path string, q url.Values, out any) error {
|
||||
for attempt := 0; attempt < 2; attempt++ {
|
||||
if err := r.ensureToken(ctx, acct, cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
endpoint := r.embyBase(cfg) + path
|
||||
if q != nil {
|
||||
endpoint += "?" + q.Encode()
|
||||
} else {
|
||||
endpoint += "?api_key=" + url.QueryEscape(cfg.Token)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("X-Emby-Token", cfg.Token)
|
||||
resp, err := r.http.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("请求远程 Emby 失败: %w", err)
|
||||
}
|
||||
data, readErr := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
|
||||
resp.Body.Close()
|
||||
if readErr != nil {
|
||||
return readErr
|
||||
}
|
||||
if resp.StatusCode == http.StatusUnauthorized && attempt == 0 {
|
||||
// token 失效:清空后重认证重试一次。
|
||||
cfg.Token = ""
|
||||
if acct != nil {
|
||||
raw := map[string]string{}
|
||||
_ = json.Unmarshal([]byte(acct.Config), &raw)
|
||||
delete(raw, "api_key")
|
||||
enc, _ := json.Marshal(raw)
|
||||
acct.Config = string(enc)
|
||||
_ = r.repo.StrmAccount.Update(ctx, acct)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if resp.StatusCode >= 300 {
|
||||
return fmt.Errorf("远程 Emby 请求失败(%d): %s", resp.StatusCode, strings.TrimSpace(string(data)))
|
||||
}
|
||||
if out == nil {
|
||||
return nil
|
||||
}
|
||||
return json.Unmarshal(data, out)
|
||||
}
|
||||
return errors.New("远程 Emby 认证重试失败")
|
||||
}
|
||||
|
||||
// TestConnection 连通性测试:确保地址可达且凭据有效;成功时回写自动认证信息。
|
||||
func (r *EmbyRemoteService) TestConnection(ctx context.Context, acct *model.StrmAccount) error {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := r.ensureToken(ctx, acct, cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
var out json.RawMessage
|
||||
return r.doGet(ctx, acct, cfg, "/System/Info", nil, &out)
|
||||
}
|
||||
|
||||
// ProxyPlayOf 返回账号是否配置了播放代理(供账号列表/编辑回显)。
|
||||
func (r *EmbyRemoteService) ProxyPlayOf(acct *model.StrmAccount) (bool, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return cfg.ProxyPlay, nil
|
||||
}
|
||||
|
||||
// ─── 元数据 / 目录聚合 ─────────────────────────────────────────────────────────
|
||||
|
||||
// RemoteViews 拉取远程媒体库(View)列表,返回远程原始 view map(未重写)。
|
||||
func (r *EmbyRemoteService) RemoteViews(ctx context.Context, acct *model.StrmAccount) ([]map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cacheKey := r.remoteCacheKey("views", acct.ID, r.remoteUserID(cfg))
|
||||
var cached []map[string]any
|
||||
if r.cache != nil && r.cache.GetJSON(ctx, cacheKey, &cached) {
|
||||
return cached, nil
|
||||
}
|
||||
q := url.Values{"api_key": {cfg.Token}}
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
}
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Views", q, &body); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if body.Items == nil {
|
||||
body.Items = []map[string]any{}
|
||||
}
|
||||
if r.cache != nil {
|
||||
r.cache.SetJSON(ctx, cacheKey, body.Items, r.remoteMediaCacheTTL())
|
||||
}
|
||||
return body.Items, nil
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) remoteUserID(cfg *EmbyRemoteConfig) string {
|
||||
if strings.TrimSpace(cfg.RemoteUserID) != "" {
|
||||
return cfg.RemoteUserID
|
||||
}
|
||||
return "0" // 未认证出的兜底:部分 Emby 接受 0 代表管理员
|
||||
}
|
||||
|
||||
// RemoteItems 向远程 Emby 转发 /Items 浏览/搜索请求,返回重写后的响应载荷。
|
||||
// p 的分页/排序/过滤参数原样转发,分页语义完全由远程承接。
|
||||
func (r *EmbyRemoteService) RemoteItems(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, p ItemsParams) (map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_, remoteParent, _ := DecodeEmbyRemoteID(p.ParentID)
|
||||
q := url.Values{}
|
||||
if remoteParent != "" {
|
||||
q.Set("ParentId", remoteParent)
|
||||
}
|
||||
q.Set("UserId", r.remoteUserID(cfg))
|
||||
q.Set("Limit", strconv.Itoa(p.Limit))
|
||||
q.Set("StartIndex", strconv.Itoa(p.StartIndex))
|
||||
if p.SearchTerm != "" {
|
||||
q.Set("SearchTerm", p.SearchTerm)
|
||||
}
|
||||
if p.Recursive {
|
||||
q.Set("Recursive", "true")
|
||||
}
|
||||
if p.SortBy != "" {
|
||||
q.Set("SortBy", p.SortBy)
|
||||
}
|
||||
if p.SortOrder != "" {
|
||||
q.Set("SortOrder", p.SortOrder)
|
||||
}
|
||||
if len(p.IncludeItemTypes) > 0 {
|
||||
q.Set("IncludeItemTypes", strings.Join(p.IncludeItemTypes, ","))
|
||||
}
|
||||
if len(p.Filters) > 0 {
|
||||
q.Set("Filters", strings.Join(p.Filters, ","))
|
||||
}
|
||||
path := "/Users/" + url.PathEscape(r.remoteUserID(cfg)) + "/Items"
|
||||
var out map[string]any
|
||||
if err := r.doGet(ctx, acct, cfg, path, q, &out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if out == nil {
|
||||
out = map[string]any{"Items": []any{}, "TotalRecordCount": 0, "StartIndex": p.StartIndex}
|
||||
}
|
||||
RewriteEmbyRemoteIDs(out, mount.ID)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// RemoteSearchMount 对单个挂载的媒体库执行全局搜索(ParentId=挂载的远程库,
|
||||
// Recursive 返回库内全部命中),结果归属明确可直接伪装。
|
||||
func (r *EmbyRemoteService) RemoteSearchMount(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, p ItemsParams) (map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
q := url.Values{}
|
||||
q.Set("ParentId", mount.RemoteViewID)
|
||||
q.Set("Recursive", "true")
|
||||
q.Set("SearchTerm", p.SearchTerm)
|
||||
q.Set("UserId", r.remoteUserID(cfg))
|
||||
q.Set("Limit", strconv.Itoa(p.Limit))
|
||||
q.Set("StartIndex", strconv.Itoa(p.StartIndex))
|
||||
if p.SortBy != "" {
|
||||
q.Set("SortBy", p.SortBy)
|
||||
}
|
||||
if p.SortOrder != "" {
|
||||
q.Set("SortOrder", p.SortOrder)
|
||||
}
|
||||
if len(p.IncludeItemTypes) > 0 {
|
||||
q.Set("IncludeItemTypes", strings.Join(p.IncludeItemTypes, ","))
|
||||
}
|
||||
var out map[string]any
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if out == nil {
|
||||
out = map[string]any{"Items": []any{}, "TotalRecordCount": 0}
|
||||
}
|
||||
RewriteEmbyRemoteIDs(out, mount.ID)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// RemoteItem 拉取远程单条目详情(含响应的重写)。
|
||||
func (r *EmbyRemoteService) RemoteItem(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) (map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
path := "/Users/" + url.PathEscape(r.remoteUserID(cfg)) + "/Items/" + url.PathEscape(remoteID)
|
||||
var out map[string]any
|
||||
if err := r.doGet(ctx, acct, cfg, path, nil, &out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
RewriteEmbyRemoteIDs(out, mount.ID)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// RemoteLatest 拉取远程「最近添加」(用于 /Items/Latest 聚合)。
|
||||
func (r *EmbyRemoteService) RemoteLatest(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, parentID string, limit int) ([]map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
q := url.Values{"Limit": {strconv.Itoa(limit)}}
|
||||
if parentID != "" {
|
||||
q.Set("ParentId", parentID)
|
||||
}
|
||||
path := "/Users/" + url.PathEscape(r.remoteUserID(cfg)) + "/Items/Latest"
|
||||
var out []map[string]any
|
||||
if err := r.doGet(ctx, acct, cfg, path, q, &out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
RewriteEmbyRemoteIDs(out, mount.ID)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// RemotePlaybackInfo 拉取远程 PlaybackInfo,并按挂载的 proxy_play 配置重写
|
||||
// 播放 URL:不代理=指向远程绝对地址(播放字节不过 MMTL);代理=指向 MMTL
|
||||
// 本地 /Videos/{encodedID} 端点(由 ProxyVideoStream 反代)。
|
||||
func (r *EmbyRemoteService) RemotePlaybackInfo(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID, userID string) (map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
q := url.Values{"UserId": {r.remoteUserID(cfg)}}
|
||||
path := "/Items/" + url.PathEscape(remoteID) + "/PlaybackInfo"
|
||||
var out map[string]any
|
||||
if err := r.doGet(ctx, acct, cfg, path, q, &out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
RewriteEmbyRemoteIDs(out, mount.ID)
|
||||
r.rewritePlayURLs(out, mount, cfg, remoteID)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// rewritePlayURLs 按挂载代理模式重写载荷内 MediaSources 的播放地址。
|
||||
// 远程 Emby 的 PlaybackInfo 通常不返回 DirectStreamUrl(客户端靠它拼
|
||||
// /Videos/{Id}/stream),因此这里总是强制构造播放地址,完全由 MMTL 掌控
|
||||
// 直连(远程绝对 URL)或代理(本地 /Videos/{encoded})的最终去向。
|
||||
func (r *EmbyRemoteService) rewritePlayURLs(value any, mount *model.EmbyMount, cfg *EmbyRemoteConfig, remoteID string) {
|
||||
encoded := EncodeEmbyRemoteID(mount.ID, remoteID)
|
||||
sources := collectMediaSources(value)
|
||||
if sources == nil {
|
||||
return
|
||||
}
|
||||
for _, src := range sources {
|
||||
mediaSourceID, _ := src["Id"].(string)
|
||||
var streamPath, subtitlePlayURL string
|
||||
if mount.ProxyPlay {
|
||||
streamPath = "/Videos/" + url.PathEscape(encoded) + "/stream"
|
||||
subtitlePlayURL = "/Videos/" + url.PathEscape(encoded)
|
||||
} else {
|
||||
base := r.embyBase(cfg)
|
||||
streamPath = base + "/Videos/" + url.PathEscape(remoteID) + "/stream?api_key=" + url.QueryEscape(cfg.Token) + "&Static=true"
|
||||
if mediaSourceID != "" {
|
||||
streamPath += "&MediaSourceId=" + url.QueryEscape(mediaSourceID)
|
||||
}
|
||||
subtitlePlayURL = base + "/Videos/" + url.PathEscape(remoteID)
|
||||
}
|
||||
// 直连/代理地址总是下发(PlaybackInfo 语义:客户端直接请求该 URL)。
|
||||
src["DirectStreamUrl"] = streamPath
|
||||
if _, exists := src["TranscodingUrl"]; exists {
|
||||
src["TranscodingUrl"] = streamPath
|
||||
}
|
||||
rewriteSubtitleDeliveryURLs(src, subtitlePlayURL, cfg)
|
||||
}
|
||||
}
|
||||
|
||||
// collectMediaSources 从载荷中取出所有 MediaSources(顶层或嵌套 Items 内)。
|
||||
func collectMediaSources(value any) []map[string]any {
|
||||
var out []map[string]any
|
||||
switch typed := value.(type) {
|
||||
case map[string]any:
|
||||
if sources, ok := typed["MediaSources"].([]map[string]any); ok {
|
||||
out = append(out, sources...)
|
||||
} else if sources, ok := typed["MediaSources"].([]any); ok {
|
||||
for _, s := range sources {
|
||||
if m, isMap := s.(map[string]any); isMap {
|
||||
out = append(out, m)
|
||||
}
|
||||
}
|
||||
}
|
||||
if items, ok := typed["Items"]; ok {
|
||||
out = append(out, collectMediaSources(items)...)
|
||||
}
|
||||
case []any:
|
||||
for _, item := range typed {
|
||||
out = append(out, collectMediaSources(item)...)
|
||||
}
|
||||
case []map[string]any:
|
||||
for _, item := range typed {
|
||||
out = append(out, collectMediaSources(item)...)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
var embySubtitleDeliveryRE = regexp.MustCompile(`/Subtitles/(\d+)/Stream(\.[A-Za-z0-9]+)?`)
|
||||
|
||||
// rewriteSubtitleDeliveryURLs 把 MediaSource 内字幕轨道的 DeliveryUrl 改写到
|
||||
// subtitlePlayURL 前缀(客户端请求本地代理端点 / 远程绝对地址)。
|
||||
func rewriteSubtitleDeliveryURLs(src map[string]any, playURL string, cfg *EmbyRemoteConfig) {
|
||||
if src == nil {
|
||||
return
|
||||
}
|
||||
streams, ok := src["MediaStreams"].([]any)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
for _, s := range streams {
|
||||
stream, isMap := s.(map[string]any)
|
||||
if !isMap || stream["Type"] != "Subtitle" {
|
||||
continue
|
||||
}
|
||||
raw, _ := stream["DeliveryUrl"].(string)
|
||||
if raw == "" {
|
||||
continue
|
||||
}
|
||||
idx := "1"
|
||||
if m := embySubtitleDeliveryRE.FindStringSubmatch(raw); len(m) >= 2 {
|
||||
idx = m[1]
|
||||
}
|
||||
ext := ""
|
||||
if m := embySubtitleDeliveryRE.FindStringSubmatch(raw); len(m) >= 3 {
|
||||
ext = m[2]
|
||||
}
|
||||
base := strings.TrimRight(playURL, "/")
|
||||
delivery := base + "/Subtitles/" + idx + "/Stream" + ext
|
||||
if !cfg.ProxyPlay && strings.TrimSpace(cfg.Token) != "" {
|
||||
delivery += "?api_key=" + url.QueryEscape(cfg.Token)
|
||||
}
|
||||
stream["DeliveryUrl"] = delivery
|
||||
}
|
||||
}
|
||||
|
||||
// RemoteImageURL 构造远程图片绝对地址(由既有 ImageProxy 拉取透传)。
|
||||
func (r *EmbyRemoteService) RemoteImageURL(ctx context.Context, acct *model.StrmAccount, remoteID, imageType string) (string, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return r.embyBase(cfg) + "/Items/" + url.PathEscape(remoteID) + "/Images/" + url.PathEscape(strings.ToLower(imageType)) +
|
||||
"?api_key=" + url.QueryEscape(cfg.Token), nil
|
||||
}
|
||||
|
||||
// ─── 播放代理 ─────────────────────────────────────────────────────────────────
|
||||
|
||||
// ProxyVideoStream 反向代理远程 Emby 视频流(保留 Range 以支持拖动)。
|
||||
func (r *EmbyRemoteService) ProxyVideoStream(ctx context.Context, w http.ResponseWriter, req *http.Request, acct *model.StrmAccount, remoteID string) error {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := r.ensureToken(ctx, acct, cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
endpoint := r.embyBase(cfg) + "/Videos/" + url.PathEscape(remoteID) + "/stream"
|
||||
q := url.Values{}
|
||||
if mediaSourceID := strings.TrimSpace(req.URL.Query().Get("MediaSourceId")); mediaSourceID != "" {
|
||||
q.Set("MediaSourceId", mediaSourceID)
|
||||
}
|
||||
// 代理是纯 byte 中继:始终要求远程原文件直连(Static=true 阻止远程触发
|
||||
// ffmpeg 转码调度——远程转码可能未配置/故障,导致整个代理 500)。
|
||||
q.Set("Static", "true")
|
||||
q.Set("api_key", cfg.Token)
|
||||
if encoded := q.Encode(); encoded != "" {
|
||||
endpoint += "?" + encoded
|
||||
}
|
||||
upstream, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
upstream.Header.Set("X-Emby-Token", cfg.Token)
|
||||
if rangeHeader := req.Header.Get("Range"); rangeHeader != "" {
|
||||
upstream.Header.Set("Range", rangeHeader)
|
||||
}
|
||||
resp, err := r.http.Do(upstream)
|
||||
if err != nil {
|
||||
return fmt.Errorf("连接远程 Emby 视频流失败: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode >= 400 {
|
||||
data, _ := io.ReadAll(io.LimitReader(resp.Body, 256))
|
||||
return fmt.Errorf("远程 Emby 视频流失败(%d): %s", resp.StatusCode, strings.TrimSpace(string(data)))
|
||||
}
|
||||
for _, header := range []string{"Content-Type", "Content-Length", "Content-Range", "Accept-Ranges", "ETag", "Cache-Control"} {
|
||||
if value := resp.Header.Get(header); value != "" {
|
||||
w.Header().Set(header, value)
|
||||
}
|
||||
}
|
||||
if resp.StatusCode == http.StatusPartialContent || resp.StatusCode == http.StatusOK {
|
||||
w.WriteHeader(resp.StatusCode)
|
||||
} else {
|
||||
w.WriteHeader(resp.StatusCode)
|
||||
}
|
||||
if resp.StatusCode == http.StatusPartialContent || resp.StatusCode == http.StatusOK {
|
||||
_, _ = io.Copy(w, resp.Body)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ProxySubtitle 反向代理远程 Emby 字幕流。
|
||||
func (r *EmbyRemoteService) ProxySubtitle(ctx context.Context, w http.ResponseWriter, req *http.Request, acct *model.StrmAccount, remoteID, index string) error {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := r.ensureToken(ctx, acct, cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
endpoint := r.embyBase(cfg) + "/Videos/" + url.PathEscape(remoteID) + "/Subtitles/" + url.PathEscape(index) + "/Stream"
|
||||
endpoint += "?api_key=" + url.QueryEscape(cfg.Token)
|
||||
upstream, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
upstream.Header.Set("X-Emby-Token", cfg.Token)
|
||||
resp, err := r.http.Do(upstream)
|
||||
if err != nil {
|
||||
return fmt.Errorf("连接远程 Emby 字幕流失败: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode >= 400 {
|
||||
return fmt.Errorf("远程 Emby 字幕流失败(%d)", resp.StatusCode)
|
||||
}
|
||||
if value := resp.Header.Get("Content-Type"); value != "" {
|
||||
w.Header().Set("Content-Type", value)
|
||||
}
|
||||
w.WriteHeader(resp.StatusCode)
|
||||
if resp.StatusCode == http.StatusOK {
|
||||
_, _ = io.Copy(w, resp.Body)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ─── 播放状态透传 ──────────────────────────────────────────────────────────────
|
||||
|
||||
// ProxySetPlayed 把「已看/未看」状态透传到远程 Emby(MMTL 本地不落库)。
|
||||
func (r *EmbyRemoteService) ProxySetPlayed(ctx context.Context, acct *model.StrmAccount, remoteID string, played bool) error {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := r.ensureToken(ctx, acct, cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
method := http.MethodPost
|
||||
path := "/Users/" + url.PathEscape(r.remoteUserID(cfg)) + "/PlayedItems/" + url.PathEscape(remoteID)
|
||||
if !played {
|
||||
method = http.MethodDelete
|
||||
}
|
||||
return r.doMutate(ctx, acct, cfg, method, path)
|
||||
}
|
||||
|
||||
// ProxySetFavorite 把「收藏/取消收藏」状态透传到远程 Emby。
|
||||
func (r *EmbyRemoteService) ProxySetFavorite(ctx context.Context, acct *model.StrmAccount, remoteID string, favorite bool) error {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := r.ensureToken(ctx, acct, cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
method := http.MethodPost
|
||||
path := "/Users/" + url.PathEscape(r.remoteUserID(cfg)) + "/FavoriteItems/" + url.PathEscape(remoteID)
|
||||
if !favorite {
|
||||
method = http.MethodDelete
|
||||
}
|
||||
return r.doMutate(ctx, acct, cfg, method, path)
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) doMutate(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig, method, path string) error {
|
||||
endpoint := r.embyBase(cfg) + path + "?api_key=" + url.QueryEscape(cfg.Token)
|
||||
req, err := http.NewRequestWithContext(ctx, method, endpoint, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("X-Emby-Token", cfg.Token)
|
||||
resp, err := r.http.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("请求远程 Emby 失败: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode >= 300 {
|
||||
data, _ := io.ReadAll(io.LimitReader(resp.Body, 256))
|
||||
return fmt.Errorf("远程 Emby 状态同步失败(%d): %s", resp.StatusCode, strings.TrimSpace(string(data)))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
// 远程 Emby 聚合的 ID 伪装。
|
||||
//
|
||||
// MMTL 作为 Emby 联邦网关把多个远程 Emby 服务器的媒体库透明聚合到自身的
|
||||
// Emby API 之下,远程条目完全不落库。为了把本地 ID 与多个远程服务器的 ID
|
||||
// 隔离开,远程条目在返回给客户端之前统一被改写为:
|
||||
//
|
||||
// embyremote~{accountID}~{remoteID}
|
||||
//
|
||||
// 客户端后续对图片 / 详情 / 播放 / 播放状态 的请求都会携带这个伪装 ID,
|
||||
// 服务端据此解码出对应账号与原始 ID,直接向远程 Emby 转发。
|
||||
package service
|
||||
|
||||
import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
// EmbyRemoteIDPrefix 远程条目伪装 ID 的前缀(本地 UUID 与 Emby ID 不会出现 "~")。
|
||||
const EmbyRemoteIDPrefix = "embyremote~"
|
||||
|
||||
// IsEmbyRemoteID 报告 id 是否是伪装过的远程 Emby 条目 ID。
|
||||
func IsEmbyRemoteID(id string) bool {
|
||||
return strings.HasPrefix(id, EmbyRemoteIDPrefix)
|
||||
}
|
||||
|
||||
// EncodeEmbyRemoteID 把 (账号 ID, 远程条目 ID) 伪装为对外暴露的 ID。
|
||||
func EncodeEmbyRemoteID(accountID, remoteID string) string {
|
||||
return EmbyRemoteIDPrefix + accountID + "~" + remoteID
|
||||
}
|
||||
|
||||
// DecodeEmbyRemoteID 拆分伪装 ID 为 (账号 ID, 远程原始 ID)。不是伪装 ID 时返回
|
||||
// ok=false。远程 ID 本身允许包含 "~"(使用 SplitN 只切第一刀)。
|
||||
func DecodeEmbyRemoteID(id string) (accountID, remoteID string, ok bool) {
|
||||
if !IsEmbyRemoteID(id) {
|
||||
return "", "", false
|
||||
}
|
||||
rest := strings.TrimPrefix(id, EmbyRemoteIDPrefix)
|
||||
parts := strings.SplitN(rest, "~", 2)
|
||||
if len(parts) != 2 || strings.TrimSpace(parts[0]) == "" || strings.TrimSpace(parts[1]) == "" {
|
||||
return "", "", false
|
||||
}
|
||||
return parts[0], parts[1], true
|
||||
}
|
||||
|
||||
// embyRemoteStringIDs 是条目 JSON 中需要伪装(编码)成远程 ID 的字符串字段。
|
||||
// 图片 / 详情 / 播放请求都会以这些字段的值作为 ID 回指 MMTL。
|
||||
var embyRemoteStringIDs = []string{
|
||||
"Id",
|
||||
"ParentId",
|
||||
"SeriesId",
|
||||
"SeasonId",
|
||||
"PrimaryImageItemId",
|
||||
"DisplayPreferencesId",
|
||||
}
|
||||
|
||||
// RewriteEmbyRemoteIDs 在内存中把远程 Emby 返回的载荷里的所有条目 ID 替换为
|
||||
// 伪装 ID(防止与本地、多远程冲突),嵌套 Items / Map 数组递归处理。
|
||||
//
|
||||
// MediaSources 里的 Id / MediaSourceId 保持不变:客户端只把它们作为查询
|
||||
// 参数原样带回,转发时直接送回远程即可。播放 URL 的重写由服务层
|
||||
// (rewriteEmbyRemotePlayURLs)按直连/代理模式处理。
|
||||
func RewriteEmbyRemoteIDs(value any, accountID string) {
|
||||
switch typed := value.(type) {
|
||||
case map[string]any:
|
||||
rewriteEmbyRemoteIDsMap(typed, accountID)
|
||||
case []any:
|
||||
for _, item := range typed {
|
||||
RewriteEmbyRemoteIDs(item, accountID)
|
||||
}
|
||||
case []map[string]any:
|
||||
for _, item := range typed {
|
||||
rewriteEmbyRemoteIDsMap(item, accountID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func rewriteEmbyRemoteIDsMap(m map[string]any, accountID string) {
|
||||
if m == nil {
|
||||
return
|
||||
}
|
||||
for _, key := range embyRemoteStringIDs {
|
||||
if raw, ok := m[key].(string); ok && raw != "" {
|
||||
m[key] = EncodeEmbyRemoteID(accountID, raw)
|
||||
}
|
||||
}
|
||||
if tags, ok := m["ImageTags"].(map[string]any); ok {
|
||||
for k, v := range tags {
|
||||
if s, isStr := v.(string); isStr && s != "" {
|
||||
tags[k] = EncodeEmbyRemoteID(accountID, s)
|
||||
}
|
||||
}
|
||||
}
|
||||
if tags, ok := m["ImageTags"].(map[string]string); ok {
|
||||
for k, v := range tags {
|
||||
if v != "" {
|
||||
tags[k] = EncodeEmbyRemoteID(accountID, v)
|
||||
}
|
||||
}
|
||||
}
|
||||
if tags, ok := m["BackdropImageTags"].([]any); ok {
|
||||
for i := range tags {
|
||||
if s, isStr := tags[i].(string); isStr && s != "" {
|
||||
tags[i] = EncodeEmbyRemoteID(accountID, s)
|
||||
}
|
||||
}
|
||||
}
|
||||
if items, ok := m["Items"]; ok {
|
||||
RewriteEmbyRemoteIDs(items, accountID)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,102 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestEmbyRemoteIDEncodeDecode(t *testing.T) {
|
||||
encoded := EncodeEmbyRemoteID("acct-1", "item-123")
|
||||
want := "embyremote~acct-1~item-123"
|
||||
if encoded != want {
|
||||
t.Fatalf("encoded = %q, want %q", encoded, want)
|
||||
}
|
||||
if !IsEmbyRemoteID(encoded) {
|
||||
t.Fatalf("IsEmbyRemoteID(%q) = false", encoded)
|
||||
}
|
||||
acctID, remoteID, ok := DecodeEmbyRemoteID(encoded)
|
||||
if !ok || acctID != "acct-1" || remoteID != "item-123" {
|
||||
t.Fatalf("decode = (%q, %q, %v)", acctID, remoteID, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeEmbyRemoteIDAllowsTildeInRemoteID(t *testing.T) {
|
||||
// 远程 ID 本身允许包含 "~":只切第一刀。
|
||||
encoded := EncodeEmbyRemoteID("acct-1", "a~b~c")
|
||||
acctID, remoteID, ok := DecodeEmbyRemoteID(encoded)
|
||||
if !ok || acctID != "acct-1" || remoteID != "a~b~c" {
|
||||
t.Fatalf("decode = (%q, %q, %v)", acctID, remoteID, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeEmbyRemoteIDRejectsLocalUUIDs(t *testing.T) {
|
||||
if _, _, ok := DecodeEmbyRemoteID("550e8400-e29b-41d4-a716-446655440000"); ok {
|
||||
t.Fatal("local UUID must not decode as remote id")
|
||||
}
|
||||
if _, _, ok := DecodeEmbyRemoteID("embyremote~only-acct"); ok {
|
||||
t.Fatal("malformed remote id must not decode")
|
||||
}
|
||||
if _, _, ok := DecodeEmbyRemoteID("embyremote~~"); ok {
|
||||
t.Fatal("empty parts must not decode")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteEmbyRemoteIDs(t *testing.T) {
|
||||
payload := map[string]any{
|
||||
"Id": "item-1",
|
||||
"ParentId": "folder-1",
|
||||
"SeriesId": "series-1",
|
||||
"SeasonId": "season-1",
|
||||
"PrimaryImageItemId": "item-1",
|
||||
"DisplayPreferencesId": "folder-1",
|
||||
"ImageTags": map[string]any{
|
||||
"Primary": "item-1",
|
||||
},
|
||||
"BackdropImageTags": []any{"item-1-bd"},
|
||||
"Items": []any{
|
||||
map[string]any{"Id": "item-2", "ParentId": "folder-2"},
|
||||
},
|
||||
// MediaSource 的 Id 保持原样(客户端仅作为 MediaSourceId 查询参数)。
|
||||
"MediaSources": []any{
|
||||
map[string]any{
|
||||
"Id": "ms-9",
|
||||
"DirectStreamUrl": "/Videos/item-1/stream",
|
||||
"MediaStreams": []any{
|
||||
map[string]any{"Type": "Subtitle", "DeliveryUrl": "/Videos/item-1/Subtitles/2/Stream.srt"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
RewriteEmbyRemoteIDs(payload, "acct-1")
|
||||
|
||||
if got := payload["Id"]; got != "embyremote~acct-1~item-1" {
|
||||
t.Fatalf("Id = %v", got)
|
||||
}
|
||||
if got := payload["ParentId"]; got != "embyremote~acct-1~folder-1" {
|
||||
t.Fatalf("ParentId = %v", got)
|
||||
}
|
||||
if got := payload["SeriesId"]; got != "embyremote~acct-1~series-1" {
|
||||
t.Fatalf("SeriesId = %v", got)
|
||||
}
|
||||
if got := payload["SeasonId"]; got != "embyremote~acct-1~season-1" {
|
||||
t.Fatalf("SeasonId = %v", got)
|
||||
}
|
||||
if got := payload["ImageTags"].(map[string]any)["Primary"]; got != "embyremote~acct-1~item-1" {
|
||||
t.Fatalf("ImageTags.Primary = %v", got)
|
||||
}
|
||||
if got := payload["BackdropImageTags"].([]any)[0]; got != "embyremote~acct-1~item-1-bd" {
|
||||
t.Fatalf("BackdropImageTags[0] = %v", got)
|
||||
}
|
||||
nested := payload["Items"].([]any)[0].(map[string]any)
|
||||
if nested["Id"] != "embyremote~acct-1~item-2" {
|
||||
t.Fatalf("nested Id = %v", nested["Id"])
|
||||
}
|
||||
|
||||
// MediaSource.Id 与 URL 不被 ID 重写器触碰(URL 由代理模式函数改写)。
|
||||
ms := payload["MediaSources"].([]any)[0].(map[string]any)
|
||||
if ms["Id"] != "ms-9" {
|
||||
t.Fatalf("MediaSource.Id must stay raw, got %v", ms["Id"])
|
||||
}
|
||||
if ms["DirectStreamUrl"] != "/Videos/item-1/stream" {
|
||||
t.Fatalf("DirectStreamUrl must stay raw, got %v", ms["DirectStreamUrl"])
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
)
|
||||
|
||||
// rewriteSubtitleDeliveryURLs 只应改动字幕轨道的 DeliveryUrl,其余媒体流不动。
|
||||
func TestRewriteSubtitleDeliveryURLsProxyMode(t *testing.T) {
|
||||
src := map[string]any{
|
||||
"MediaStreams": []any{
|
||||
map[string]any{"Type": "Video", "DeliveryUrl": "/Videos/x/stream"},
|
||||
map[string]any{"Type": "Audio", "DeliveryUrl": "/Videos/x/stream"},
|
||||
map[string]any{"Type": "Subtitle", "DeliveryUrl": "/Videos/item-1/ms-9/Subtitles/2/Stream.srt"},
|
||||
},
|
||||
}
|
||||
rewriteSubtitleDeliveryURLs(src, "/Videos/embyremote~acct-1~item-1", &EmbyRemoteConfig{})
|
||||
streams := src["MediaStreams"].([]any)
|
||||
if got := streams[0].(map[string]any)["DeliveryUrl"]; got != "/Videos/x/stream" {
|
||||
t.Fatalf("video DeliveryUrl must stay, got %v", got)
|
||||
}
|
||||
want := "/Videos/embyremote~acct-1~item-1/Subtitles/2/Stream.srt"
|
||||
if got := streams[2].(map[string]any)["DeliveryUrl"]; got != want {
|
||||
t.Fatalf("subtitle DeliveryUrl = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteSubtitleDeliveryURLsDirectMode(t *testing.T) {
|
||||
src := map[string]any{
|
||||
"MediaStreams": []any{
|
||||
map[string]any{"Type": "Subtitle", "DeliveryUrl": "/Videos/item-1/ms-9/Subtitles/1/Stream.ass"},
|
||||
},
|
||||
}
|
||||
cfg := &EmbyRemoteConfig{Token: "tok123"}
|
||||
rewriteSubtitleDeliveryURLs(src, "http://remote:8096/emby/Videos/item-1", cfg)
|
||||
streams := src["MediaStreams"].([]any)
|
||||
want := "http://remote:8096/emby/Videos/item-1/Subtitles/1/Stream.ass?api_key=tok123"
|
||||
if got := streams[0].(map[string]any)["DeliveryUrl"]; got != want {
|
||||
t.Fatalf("subtitle DeliveryUrl = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteSubtitleDeliveryURLsFallsBackIndexOne(t *testing.T) {
|
||||
src := map[string]any{
|
||||
"MediaStreams": []any{
|
||||
map[string]any{"Type": "Subtitle", "DeliveryUrl": "custom/url"},
|
||||
},
|
||||
}
|
||||
rewriteSubtitleDeliveryURLs(src, "/Videos/embyremote~acct-1~item-1", &EmbyRemoteConfig{})
|
||||
streams := src["MediaStreams"].([]any)
|
||||
want := "/Videos/embyremote~acct-1~item-1/Subtitles/1/Stream"
|
||||
if got := streams[0].(map[string]any)["DeliveryUrl"]; got != want {
|
||||
t.Fatalf("subtitle DeliveryUrl = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapRemoteItemToMediaExtractsCodecsAndContainer(t *testing.T) {
|
||||
r := &EmbyRemoteService{}
|
||||
item := map[string]any{
|
||||
"Id": "item-100",
|
||||
"Name": "Test Movie",
|
||||
"Container": "mkv",
|
||||
"MediaStreams": []any{
|
||||
map[string]any{
|
||||
"Type": "Video",
|
||||
"Codec": "h264",
|
||||
"Width": 1920,
|
||||
"Height": 1080,
|
||||
},
|
||||
map[string]any{
|
||||
"Type": "Audio",
|
||||
"Codec": "aac",
|
||||
},
|
||||
},
|
||||
"MediaSources": []any{
|
||||
map[string]any{
|
||||
"Container": "mkv",
|
||||
"Size": int64(104857600),
|
||||
},
|
||||
},
|
||||
}
|
||||
media := r.MapRemoteItemToMedia(t.Context(), nil, &model.StrmAccount{Base: model.Base{ID: "acct-1"}}, &EmbyRemoteConfig{}, item)
|
||||
if media.Container != "mkv" {
|
||||
t.Fatalf("media.Container = %v, want mkv", media.Container)
|
||||
}
|
||||
if media.VideoCodec != "h264" {
|
||||
t.Fatalf("media.VideoCodec = %v, want h264", media.VideoCodec)
|
||||
}
|
||||
if media.AudioCodec != "aac" {
|
||||
t.Fatalf("media.AudioCodec = %v, want aac", media.AudioCodec)
|
||||
}
|
||||
if media.Width != 1920 || media.Height != 1080 {
|
||||
t.Fatalf("resolution = %dx%d, want 1920x1080", media.Width, media.Height)
|
||||
}
|
||||
if media.SizeBytes != 104857600 {
|
||||
t.Fatalf("size = %d, want 104857600", media.SizeBytes)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,735 @@
|
||||
// 网页端远程 Emby 库映射。
|
||||
//
|
||||
// 网页端(React UI)的媒体库/媒体浏览走项目自有 REST API(/api/libraries、
|
||||
// /api/libraries/:id/media、/api/media/:id 等),数据结构为 model.Library /
|
||||
// model.Media / SeriesCard。远程 Emby 挂载的数据不落库,因此这里把远程
|
||||
// Emby 的 JSON item 映射为与本地完全一致的结构,让网页端无感知地浏览
|
||||
// 远程库;播放统一走 /api/stream/{伪装ID}(302 到远程 Emby 原地址)。
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/url"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
)
|
||||
|
||||
// RemoteLibraryView 是一个网页端可见的远程媒体库(对应一个挂载的远程 View)。
|
||||
type RemoteLibraryView struct {
|
||||
Library model.Library
|
||||
MountID string
|
||||
AccountID string
|
||||
RemoteID string
|
||||
CollectionType string
|
||||
AccountName string
|
||||
}
|
||||
|
||||
// RemoteLibraries 把所有启用挂载的远程媒体库映射为网页媒体库列表
|
||||
// (只有显式挂载的库才出现在本项目媒体库中)。
|
||||
func (r *EmbyRemoteService) RemoteLibraries(ctx context.Context) ([]RemoteLibraryView, error) {
|
||||
mounts, err := r.ListMounts(ctx)
|
||||
if err != nil || len(mounts) == 0 {
|
||||
return nil, err
|
||||
}
|
||||
type accountData struct {
|
||||
acct *model.StrmAccount
|
||||
cfg *EmbyRemoteConfig
|
||||
viewByName map[string]map[string]any
|
||||
}
|
||||
acctData := map[string]*accountData{}
|
||||
for i := range mounts {
|
||||
m := &mounts[i]
|
||||
if !m.Enabled {
|
||||
continue
|
||||
}
|
||||
if _, ok := acctData[m.AccountID]; ok {
|
||||
continue
|
||||
}
|
||||
acct := r.AccountByID(ctx, m.AccountID)
|
||||
if acct == nil {
|
||||
acctData[m.AccountID] = nil
|
||||
continue
|
||||
}
|
||||
cfg, cfgErr := r.configOf(acct)
|
||||
if cfgErr != nil {
|
||||
acctData[m.AccountID] = nil
|
||||
continue
|
||||
}
|
||||
views, viewErr := r.RemoteViews(ctx, acct)
|
||||
if viewErr != nil {
|
||||
if r.log != nil {
|
||||
r.log.Warn("web remote emby views failed",
|
||||
zap.String("account", acct.Name), zap.Error(viewErr))
|
||||
}
|
||||
acctData[m.AccountID] = nil
|
||||
continue
|
||||
}
|
||||
viewByName := map[string]map[string]any{}
|
||||
for _, v := range views {
|
||||
viewByName[remoteItemString(v, "Id")] = v
|
||||
}
|
||||
acctData[m.AccountID] = &accountData{
|
||||
acct: acct,
|
||||
cfg: cfg,
|
||||
viewByName: viewByName,
|
||||
}
|
||||
}
|
||||
|
||||
out := make([]RemoteLibraryView, 0, len(mounts))
|
||||
for i := range mounts {
|
||||
m := &mounts[i]
|
||||
if !m.Enabled {
|
||||
continue
|
||||
}
|
||||
data := acctData[m.AccountID]
|
||||
if data == nil || data.viewByName == nil {
|
||||
continue
|
||||
}
|
||||
v, ok := data.viewByName[m.RemoteViewID]
|
||||
if !ok {
|
||||
continue // 远程已删除该媒体库
|
||||
}
|
||||
lib := r.mapRemoteMountToLibrary(m, data.acct, data.cfg, v)
|
||||
if lib == nil {
|
||||
continue
|
||||
}
|
||||
out = append(out, RemoteLibraryView{
|
||||
Library: *lib,
|
||||
MountID: m.ID,
|
||||
AccountID: data.acct.ID,
|
||||
RemoteID: m.RemoteViewID,
|
||||
CollectionType: m.CollectionType,
|
||||
AccountName: data.acct.Name,
|
||||
})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// RemoteLibraryByID 按伪装 ID 查远程库视图(详情接口用)。
|
||||
func (r *EmbyRemoteService) RemoteLibraryByID(ctx context.Context, mountID, remoteViewID string) (*RemoteLibraryView, error) {
|
||||
views, err := r.RemoteLibraries(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, v := range views {
|
||||
if v.MountID == mountID && v.RemoteID == remoteViewID {
|
||||
cp := v
|
||||
return &cp, nil
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// mapRemoteMountToLibrary 把挂载信息 + 远程 View item 映射为网页库结构。
|
||||
func (r *EmbyRemoteService) mapRemoteMountToLibrary(mount *model.EmbyMount, acct *model.StrmAccount, cfg *EmbyRemoteConfig, item map[string]any) *model.Library {
|
||||
if mount == nil {
|
||||
return nil
|
||||
}
|
||||
name := strings.TrimSpace(mount.Name)
|
||||
if name == "" {
|
||||
name = strings.TrimSpace(remoteItemString(item, "Name"))
|
||||
}
|
||||
if name == "" {
|
||||
name = acct.Name
|
||||
} else if !strings.Contains(name, acct.Name) {
|
||||
name = acct.Name + " · " + name
|
||||
}
|
||||
libType := "movie"
|
||||
switch mount.CollectionType {
|
||||
case "tvshows":
|
||||
libType = "tv"
|
||||
case "music":
|
||||
libType = "music"
|
||||
}
|
||||
lib := &model.Library{
|
||||
Base: model.Base{ID: EncodeEmbyRemoteID(mount.ID, mount.RemoteViewID)},
|
||||
Name: name,
|
||||
Type: libType,
|
||||
Enabled: true,
|
||||
SortOrder: 1000 + mount.SortOrder, // 远程库排在本地库之后,且保持挂载库排序
|
||||
}
|
||||
// 远程媒体库封面只有真实存在图片标签才下发。
|
||||
if remoteItemHasImageTag(item, "Primary") {
|
||||
lib.CoverURL = r.remoteItemImageURL(cfg, mount.RemoteViewID, "Primary")
|
||||
}
|
||||
return lib
|
||||
}
|
||||
|
||||
// MapRemoteItemToMedia 把远程 Emby item JSON 映射为本地 Media 结构。
|
||||
// poster/backdrop 只有在远程确实存在图片标签时才填 URL(避免对无图条目
|
||||
// 发出必失败的图片请求导致前端破图);剧集回退到系列海报(SeriesPrimaryImage)。
|
||||
func (r *EmbyRemoteService) MapRemoteItemToMedia(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, cfg *EmbyRemoteConfig, item map[string]any) model.Media {
|
||||
encodeScope := acct.ID
|
||||
if mount != nil {
|
||||
encodeScope = mount.ID
|
||||
}
|
||||
remoteID := remoteItemString(item, "Id")
|
||||
// 条目可能已被 RewriteEmbyRemoteIDs 伪装(图片/嵌套 ID 需要原始远程 ID)。
|
||||
if _, rid, ok := DecodeEmbyRemoteID(remoteID); ok {
|
||||
remoteID = rid
|
||||
}
|
||||
seriesID := remoteItemString(item, "SeriesId")
|
||||
if _, rid, ok := DecodeEmbyRemoteID(seriesID); ok {
|
||||
seriesID = rid
|
||||
}
|
||||
rating := remoteItemFloat(item, "CommunityRating")
|
||||
if rating == 0 {
|
||||
rating = remoteItemFloat(item, "CriticRating")
|
||||
}
|
||||
media := model.Media{
|
||||
Base: model.Base{ID: EncodeEmbyRemoteID(encodeScope, remoteID)},
|
||||
Title: remoteItemString(item, "Name"),
|
||||
OriginalName: remoteItemString(item, "OriginalTitle"),
|
||||
Overview: remoteItemString(item, "Overview"),
|
||||
Year: remoteItemInt(item, "ProductionYear"),
|
||||
Rating: float32(rating),
|
||||
Path: remoteItemString(item, "Path"),
|
||||
Genres: remoteItemGenres(item),
|
||||
ScrapeStatus: "done",
|
||||
}
|
||||
if date, ok := parseEmbyRemoteDate(remoteItemString(item, "DateCreated")); ok {
|
||||
media.CreatedAt = date
|
||||
media.UpdatedAt = date
|
||||
}
|
||||
// 只有远程明确存在图片标签才下发图片 URL。
|
||||
if remoteItemHasImageTag(item, "Primary") {
|
||||
media.PosterURL = r.remoteItemImageURL(cfg, remoteID, "Primary")
|
||||
}
|
||||
if remoteItemHasImageTag(item, "Backdrop") || len(remoteBackdropTags(item)) > 0 {
|
||||
media.BackdropURL = r.remoteItemImageURL(cfg, remoteID, "Backdrop")
|
||||
}
|
||||
if ticks := remoteItemInt64(item, "RunTimeTicks"); ticks > 0 {
|
||||
media.DurationSec = int(ticks / 10_000_000)
|
||||
}
|
||||
if date, ok := parseEmbyRemoteDate(remoteItemString(item, "PremiereDate")); ok {
|
||||
media.ReleaseDate = date.Format("2006-01-02")
|
||||
if media.Year == 0 {
|
||||
media.Year = date.Year()
|
||||
}
|
||||
} else if date, ok := embyPremiereDate(remoteItemString(item, "PremiereDate")); ok {
|
||||
media.ReleaseDate = date.Format("2006-01-02")
|
||||
if media.Year == 0 {
|
||||
media.Year = date.Year()
|
||||
}
|
||||
}
|
||||
if providerIDs, ok := item["ProviderIds"].(map[string]any); ok {
|
||||
if v := anyString(providerIDs["Tmdb"]); v != "" {
|
||||
media.TMDbID, _ = strconv.Atoi(v)
|
||||
}
|
||||
if v := anyString(providerIDs["Imdb"]); v != "" {
|
||||
media.TheTVDBID = v
|
||||
}
|
||||
if v := anyString(providerIDs["Douban"]); v != "" {
|
||||
media.DoubanID = v
|
||||
}
|
||||
}
|
||||
media.Container = remoteItemString(item, "Container")
|
||||
media.Width = remoteItemInt(item, "Width")
|
||||
media.Height = remoteItemInt(item, "Height")
|
||||
|
||||
extractStreamInfo := func(streams []any) {
|
||||
for _, s := range streams {
|
||||
sm, ok := s.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
typ := remoteItemString(sm, "Type")
|
||||
if strings.EqualFold(typ, "Video") {
|
||||
if media.VideoCodec == "" {
|
||||
media.VideoCodec = remoteItemString(sm, "Codec")
|
||||
}
|
||||
if media.Width == 0 {
|
||||
media.Width = remoteItemInt(sm, "Width")
|
||||
}
|
||||
if media.Height == 0 {
|
||||
media.Height = remoteItemInt(sm, "Height")
|
||||
}
|
||||
} else if strings.EqualFold(typ, "Audio") {
|
||||
if media.AudioCodec == "" {
|
||||
media.AudioCodec = remoteItemString(sm, "Codec")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if streams, ok := item["MediaStreams"].([]any); ok {
|
||||
extractStreamInfo(streams)
|
||||
} else if streams, ok := item["MediaStreams"].([]map[string]any); ok {
|
||||
anyStreams := make([]any, len(streams))
|
||||
for i, v := range streams {
|
||||
anyStreams[i] = v
|
||||
}
|
||||
extractStreamInfo(anyStreams)
|
||||
}
|
||||
|
||||
if sources, ok := item["MediaSources"].([]any); ok && len(sources) > 0 {
|
||||
if sourceMap, ok := sources[0].(map[string]any); ok {
|
||||
if media.Container == "" {
|
||||
media.Container = remoteItemString(sourceMap, "Container")
|
||||
}
|
||||
if media.SizeBytes == 0 {
|
||||
media.SizeBytes = remoteItemInt64(sourceMap, "Size")
|
||||
}
|
||||
if streams, ok := sourceMap["MediaStreams"].([]any); ok && (media.VideoCodec == "" || media.AudioCodec == "") {
|
||||
extractStreamInfo(streams)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
switch remoteItemString(item, "Type") {
|
||||
case "Episode":
|
||||
media.SeasonNum = remoteItemInt(item, "ParentIndexNumber")
|
||||
media.EpisodeNum = remoteItemInt(item, "IndexNumber")
|
||||
media.EpisodeTitle = remoteItemString(item, "Name")
|
||||
if seriesName := remoteItemString(item, "SeriesName"); seriesName != "" {
|
||||
media.Title = seriesName
|
||||
}
|
||||
// 单集通常无独立海报:若远程返回 SeriesPrimaryImageTag(需要
|
||||
// Fields=SeriesPrimaryImage)且系列有图,则回退到系列海报。
|
||||
if media.PosterURL == "" && seriesID != "" &&
|
||||
strings.TrimSpace(remoteItemString(item, "SeriesPrimaryImageTag")) != "" {
|
||||
media.PosterURL = r.remoteItemImageURL(cfg, seriesID, "Primary")
|
||||
}
|
||||
default: // Movie / Series / Season / Folder
|
||||
media.SeasonNum = 0
|
||||
media.EpisodeNum = 0
|
||||
}
|
||||
return media
|
||||
}
|
||||
|
||||
// RemoteLibraryMedia 拉远程库直属条目(电影库=Movie,剧集库=Series),映射分页。
|
||||
func (r *EmbyRemoteService) RemoteLibraryMedia(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteViewID string, itemTypes string, offset, limit int) ([]model.Media, int64, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if itemTypes == "" {
|
||||
itemTypes = "Movie,Series" // 未知类型时两者都取(前端自行按 episode-like 分组)
|
||||
}
|
||||
cacheKey := r.remoteCacheKey("library-media", acct.ID, mount.ID, remoteViewID, itemTypes, strconv.Itoa(offset), strconv.Itoa(limit))
|
||||
var cached struct {
|
||||
Items []model.Media `json:"items"`
|
||||
TotalRecordCount int64 `json:"total_record_count"`
|
||||
}
|
||||
if r.cache != nil && r.cache.GetJSON(ctx, cacheKey, &cached) {
|
||||
return cached.Items, cached.TotalRecordCount, nil
|
||||
}
|
||||
q := url.Values{}
|
||||
q.Set("ParentId", remoteViewID)
|
||||
q.Set("IncludeItemTypes", itemTypes)
|
||||
q.Set("Recursive", "false")
|
||||
q.Set("StartIndex", strconv.Itoa(offset))
|
||||
q.Set("Limit", strconv.Itoa(limit))
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,SeriesPrimaryImage,MediaStreams,MediaSources,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating")
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
TotalRecordCount int64 `json:"TotalRecordCount"`
|
||||
}
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &body); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
items := make([]model.Media, 0, len(body.Items))
|
||||
for _, it := range body.Items {
|
||||
RewriteEmbyRemoteIDs(it, mount.ID) // 嵌套/关联 ID 一并伪装
|
||||
items = append(items, r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it))
|
||||
}
|
||||
if r.cache != nil {
|
||||
r.cache.SetJSON(ctx, cacheKey, struct {
|
||||
Items []model.Media `json:"items"`
|
||||
TotalRecordCount int64 `json:"total_record_count"`
|
||||
}{Items: items, TotalRecordCount: body.TotalRecordCount}, r.remoteMediaCacheTTL())
|
||||
}
|
||||
return items, body.TotalRecordCount, nil
|
||||
}
|
||||
|
||||
// RemoteMediaDetail 拉远程单条目映射为 Media(网页详情页)。
|
||||
func (r *EmbyRemoteService) RemoteMediaDetail(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) (*model.Media, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
path := "/Users/" + url.PathEscape(r.remoteUserID(cfg)) + "/Items/" + url.PathEscape(remoteID)
|
||||
path += "?Fields=Overview,Genres,ProviderIds,People,Studios,Path,MediaStreams,MediaSources,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating"
|
||||
var out map[string]any
|
||||
if err := r.doGet(ctx, acct, cfg, path, nil, &out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
RewriteEmbyRemoteIDs(out, mount.ID)
|
||||
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, out)
|
||||
return &m, nil
|
||||
}
|
||||
|
||||
// RemoteEpisodes 拉远程条目下的集列表(Series/Season/Folder→子集;Episode→同系列;
|
||||
// Movie→自身单条),按季/集排序,与本地 ListMediaEpisodes 行为一致。
|
||||
func (r *EmbyRemoteService) RemoteEpisodes(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) ([]model.Media, error) {
|
||||
detail, err := r.RemoteMediaDetail(ctx, mount, acct, remoteID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 用远程详情载荷精判类型(Episode→同系列;Series/Season/Folder→子集;Movie→单条)。
|
||||
itemType := r.remoteItemType(ctx, acct, remoteID)
|
||||
if itemType == "" {
|
||||
itemType = remoteItemTypeOf(detail)
|
||||
}
|
||||
var parentID string
|
||||
switch itemType {
|
||||
case "Episode":
|
||||
parentID = r.remoteItemSeriesID(ctx, acct, remoteID)
|
||||
if parentID == "" {
|
||||
parentID = remoteID
|
||||
}
|
||||
case "Season", "Folder", "Series":
|
||||
parentID = remoteID
|
||||
default: // Movie
|
||||
return []model.Media{*detail}, nil
|
||||
}
|
||||
rows, _, err := r.remoteEpisodesOf(ctx, mount, acct, parentID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sort.SliceStable(rows, func(i, j int) bool {
|
||||
if rows[i].SeasonNum != rows[j].SeasonNum {
|
||||
return rows[i].SeasonNum < rows[j].SeasonNum
|
||||
}
|
||||
if rows[i].EpisodeNum != rows[j].EpisodeNum {
|
||||
return rows[i].EpisodeNum < rows[j].EpisodeNum
|
||||
}
|
||||
return rows[i].Title < rows[j].Title
|
||||
})
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) remoteEpisodesOf(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, parentID string) ([]model.Media, int64, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
q := url.Values{}
|
||||
q.Set("ParentId", parentID)
|
||||
q.Set("IncludeItemTypes", "Episode")
|
||||
q.Set("Recursive", "true")
|
||||
q.Set("StartIndex", "0")
|
||||
q.Set("Limit", "500")
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,SeriesPrimaryImage,MediaStreams,MediaSources,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating")
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
TotalRecordCount int64 `json:"TotalRecordCount"`
|
||||
}
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &body); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
items := make([]model.Media, 0, len(body.Items))
|
||||
for _, it := range body.Items {
|
||||
RewriteEmbyRemoteIDs(it, mount.ID)
|
||||
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it)
|
||||
items = append(items, m)
|
||||
}
|
||||
return items, body.TotalRecordCount, nil
|
||||
}
|
||||
|
||||
// RemoteSeriesCards 远程剧集库的系列卡片(ChildCount 作为集数)。
|
||||
func (r *EmbyRemoteService) RemoteSeriesCards(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteViewID string) ([]SeriesCard, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cacheKey := r.remoteCacheKey("series-cards", acct.ID, mount.ID, remoteViewID)
|
||||
var cached []SeriesCard
|
||||
if r.cache != nil && r.cache.GetJSON(ctx, cacheKey, &cached) {
|
||||
return cached, nil
|
||||
}
|
||||
q := url.Values{}
|
||||
q.Set("ParentId", remoteViewID)
|
||||
q.Set("IncludeItemTypes", "Series")
|
||||
q.Set("Recursive", "false")
|
||||
q.Set("StartIndex", "0")
|
||||
q.Set("Limit", "1000")
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,RecursiveItemCount,SeriesPrimaryImage,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating")
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
}
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &body); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cards := make([]SeriesCard, 0, len(body.Items))
|
||||
for _, it := range body.Items {
|
||||
RewriteEmbyRemoteIDs(it, mount.ID)
|
||||
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it)
|
||||
// 集数优先用递归条目数(ChildCount 只算直属 Season 文件夹数)。
|
||||
count := remoteItemInt(it, "RecursiveItemCount")
|
||||
if count == 0 {
|
||||
count = remoteItemInt(it, "ChildCount")
|
||||
}
|
||||
if count == 0 {
|
||||
count = 1
|
||||
}
|
||||
cards = append(cards, SeriesCard{Key: m.ID, Rep: m, LinkMedia: m, Count: count})
|
||||
}
|
||||
if r.cache != nil {
|
||||
r.cache.SetJSON(ctx, cacheKey, cards, r.remoteMediaCacheTTL())
|
||||
}
|
||||
return cards, nil
|
||||
}
|
||||
|
||||
// RemoteLatestCards 远程库最新条目(首页预览卡片),映射 SeriesCard。
|
||||
func (r *EmbyRemoteService) RemoteLatestCards(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteViewID string, limit int) ([]SeriesCard, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cacheKey := r.remoteCacheKey("latest-cards", acct.ID, mount.ID, remoteViewID, strconv.Itoa(limit))
|
||||
var cached []SeriesCard
|
||||
if r.cache != nil && r.cache.GetJSON(ctx, cacheKey, &cached) {
|
||||
return cached, nil
|
||||
}
|
||||
items, err := r.RemoteLatest(ctx, mount, acct, remoteViewID, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cards := make([]SeriesCard, 0, len(items))
|
||||
for _, it := range items {
|
||||
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it)
|
||||
cards = append(cards, SeriesCard{Key: m.ID, Rep: m, LinkMedia: m, Count: 0})
|
||||
}
|
||||
if r.cache != nil {
|
||||
r.cache.SetJSON(ctx, cacheKey, cards, r.remoteMediaCacheTTL())
|
||||
}
|
||||
return cards, nil
|
||||
}
|
||||
|
||||
// WebStreamURL 远程条目的网页播放地址(302 直连远程 Emby 流端点)。
|
||||
func (r *EmbyRemoteService) WebStreamURL(ctx context.Context, acct *model.StrmAccount, remoteID string) (string, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := r.ensureToken(ctx, acct, cfg); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return r.embyBase(cfg) + "/Videos/" + url.PathEscape(remoteID) +
|
||||
"/stream?api_key=" + url.QueryEscape(cfg.Token) + "&Static=true", nil
|
||||
}
|
||||
|
||||
// remoteItemType 轻量查询远程条目 Type(避免依赖映射载荷)。
|
||||
func (r *EmbyRemoteService) remoteItemType(ctx context.Context, acct *model.StrmAccount, remoteID string) string {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
var out map[string]any
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items/"+url.PathEscape(remoteID), nil, &out); err != nil {
|
||||
return ""
|
||||
}
|
||||
return remoteItemString(out, "Type")
|
||||
}
|
||||
|
||||
// remoteItemSeriesID 轻量查询 Episode 的 SeriesId。
|
||||
func (r *EmbyRemoteService) remoteItemSeriesID(ctx context.Context, acct *model.StrmAccount, remoteID string) string {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
var out map[string]any
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items/"+url.PathEscape(remoteID), nil, &out); err != nil {
|
||||
return ""
|
||||
}
|
||||
return remoteItemString(out, "SeriesId")
|
||||
}
|
||||
|
||||
// ─── 远程 item JSON 取值辅助 ────────────────────────────────────────────────
|
||||
|
||||
func anyString(v any) string {
|
||||
if s, ok := v.(string); ok {
|
||||
return s
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func remoteItemString(item map[string]any, key string) string {
|
||||
if item == nil {
|
||||
return ""
|
||||
}
|
||||
if s, ok := item[key].(string); ok {
|
||||
return s
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func remoteItemInt(item map[string]any, key string) int {
|
||||
if item == nil {
|
||||
return 0
|
||||
}
|
||||
switch v := item[key].(type) {
|
||||
case float64:
|
||||
return int(v)
|
||||
case int:
|
||||
return v
|
||||
case int64:
|
||||
return int(v)
|
||||
case string:
|
||||
n, _ := strconv.Atoi(v)
|
||||
return n
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func remoteItemInt64(item map[string]any, key string) int64 {
|
||||
if item == nil {
|
||||
return 0
|
||||
}
|
||||
switch v := item[key].(type) {
|
||||
case float64:
|
||||
return int64(v)
|
||||
case int64:
|
||||
return v
|
||||
case int:
|
||||
return int64(v)
|
||||
case string:
|
||||
n, _ := strconv.ParseInt(v, 10, 64)
|
||||
return n
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func remoteItemFloat(item map[string]any, key string) float64 {
|
||||
if item == nil {
|
||||
return 0
|
||||
}
|
||||
switch v := item[key].(type) {
|
||||
case float64:
|
||||
return v
|
||||
case int:
|
||||
return float64(v)
|
||||
case string:
|
||||
f, _ := strconv.ParseFloat(v, 64)
|
||||
return f
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// remoteItemGenres 合并 GenreItems / Genres 数组为逗号分隔字符串(前端 parseCSV 消费)。
|
||||
func remoteItemGenres(item map[string]any) string {
|
||||
seen := map[string]bool{}
|
||||
var parts []string
|
||||
collect := func(arr any) {
|
||||
list, ok := arr.([]any)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
for _, it := range list {
|
||||
var name string
|
||||
if m, isMap := it.(map[string]any); isMap {
|
||||
name = remoteItemString(m, "Name")
|
||||
} else if s, isStr := it.(string); isStr {
|
||||
name = s
|
||||
}
|
||||
if name != "" && !seen[name] {
|
||||
seen[name] = true
|
||||
parts = append(parts, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
collect(item["GenreItems"])
|
||||
collect(item["Genres"])
|
||||
return strings.Join(parts, ",")
|
||||
}
|
||||
|
||||
// remoteItemImageURL 构造远程条目图片绝对地址(带 api_key;前端经 /api/img 代理)。
|
||||
func (r *EmbyRemoteService) remoteItemImageURL(cfg *EmbyRemoteConfig, remoteID, imageType string) string {
|
||||
if remoteID == "" {
|
||||
return ""
|
||||
}
|
||||
imageType = strings.ToLower(imageType)
|
||||
if imageType == "" {
|
||||
imageType = "primary"
|
||||
}
|
||||
return r.embyBase(cfg) + "/Items/" + url.PathEscape(remoteID) + "/Images/" + url.PathEscape(imageType) +
|
||||
"?api_key=" + url.QueryEscape(cfg.Token)
|
||||
}
|
||||
|
||||
// remoteItemHasImageTag 远程 item 是否带某类型图片标签(Emby 的 ImageTags map)。
|
||||
func remoteItemHasImageTag(item map[string]any, typ string) bool {
|
||||
if item == nil {
|
||||
return false
|
||||
}
|
||||
switch tags := item["ImageTags"].(type) {
|
||||
case map[string]any:
|
||||
_, ok := tags[typ]
|
||||
return ok
|
||||
case map[string]string:
|
||||
_, ok := tags[typ]
|
||||
return ok
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// remoteBackdropTags 远程 item 的 BackdropImageTags 数组。
|
||||
func remoteBackdropTags(item map[string]any) []any {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
switch tags := item["BackdropImageTags"].(type) {
|
||||
case []any:
|
||||
return tags
|
||||
case []string:
|
||||
out := make([]any, 0, len(tags))
|
||||
for _, s := range tags {
|
||||
out = append(out, s)
|
||||
}
|
||||
return out
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// remoteItemTypeOf 从映射后的 Media 推断远程类型(无详情载荷时兜底)。
|
||||
func remoteItemTypeOf(m *model.Media) string {
|
||||
if m == nil {
|
||||
return "Movie"
|
||||
}
|
||||
if m.EpisodeNum > 0 || m.SeasonNum > 0 {
|
||||
return "Episode"
|
||||
}
|
||||
return "Movie"
|
||||
}
|
||||
|
||||
// ─── 供 handler 层使用的远程 View 条目取值(导出薄封装) ──────────────────────
|
||||
|
||||
// RemoteItemIDString 提取远程 View 条目的 Id。
|
||||
func RemoteItemIDString(item map[string]any) string { return remoteItemString(item, "Id") }
|
||||
|
||||
// RemoteItemNameString 提取远程 View 条目的 Name。
|
||||
func RemoteItemNameString(item map[string]any) string { return remoteItemString(item, "Name") }
|
||||
|
||||
// RemoteItemCollectionType 提取远程 View 条目的 CollectionType。
|
||||
func RemoteItemCollectionType(item map[string]any) string {
|
||||
return remoteItemString(item, "CollectionType")
|
||||
}
|
||||
|
||||
// RemoteItemChildCount 提取远程 View 条目的 ChildCount。
|
||||
func RemoteItemChildCount(item map[string]any) int { return remoteItemInt(item, "ChildCount") }
|
||||
|
||||
func parseEmbyRemoteDate(s string) (time.Time, bool) {
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" {
|
||||
return time.Time{}, false
|
||||
}
|
||||
for _, layout := range []string{
|
||||
time.RFC3339Nano,
|
||||
time.RFC3339,
|
||||
"2006-01-02T15:04:05.9999999Z",
|
||||
"2006-01-02T15:04:05.9999999",
|
||||
"2006-01-02T15:04:05",
|
||||
"2006-01-02",
|
||||
} {
|
||||
if t, err := time.Parse(layout, s); err == nil {
|
||||
return t, true
|
||||
}
|
||||
}
|
||||
return time.Time{}, false
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
)
|
||||
|
||||
func TestMapRemoteItemToMediaSortingFields(t *testing.T) {
|
||||
svc := &EmbyRemoteService{}
|
||||
mount := &model.EmbyMount{Base: model.Base{ID: "mount-1"}}
|
||||
acct := &model.StrmAccount{Base: model.Base{ID: "acct-1"}}
|
||||
cfg := &EmbyRemoteConfig{BaseURL: "http://localhost:8096"}
|
||||
|
||||
item := map[string]any{
|
||||
"Id": "item-1",
|
||||
"Name": "测试电影",
|
||||
"OriginalTitle": "Test Movie",
|
||||
"ProductionYear": 2023,
|
||||
"CommunityRating": 8.5,
|
||||
"PremiereDate": "2023-05-12T00:00:00.0000000Z",
|
||||
"DateCreated": "2024-01-15T08:30:00.0000000Z",
|
||||
}
|
||||
|
||||
media := svc.MapRemoteItemToMedia(context.Background(), mount, acct, cfg, item)
|
||||
|
||||
if media.ReleaseDate != "2023-05-12" {
|
||||
t.Fatalf("ReleaseDate = %q, want %q", media.ReleaseDate, "2023-05-12")
|
||||
}
|
||||
if media.Year != 2023 {
|
||||
t.Fatalf("Year = %d, want 2023", media.Year)
|
||||
}
|
||||
if media.Rating != 8.5 {
|
||||
t.Fatalf("Rating = %f, want 8.5", media.Rating)
|
||||
}
|
||||
expectedCreated, _ := time.Parse(time.RFC3339, "2024-01-15T08:30:00Z")
|
||||
if !media.CreatedAt.Equal(expectedCreated) {
|
||||
t.Fatalf("CreatedAt = %v, want %v", media.CreatedAt, expectedCreated)
|
||||
}
|
||||
if !media.UpdatedAt.Equal(expectedCreated) {
|
||||
t.Fatalf("UpdatedAt = %v, want %v", media.UpdatedAt, expectedCreated)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapRemoteItemToMediaCriticRatingFallback(t *testing.T) {
|
||||
svc := &EmbyRemoteService{}
|
||||
mount := &model.EmbyMount{Base: model.Base{ID: "mount-1"}}
|
||||
acct := &model.StrmAccount{Base: model.Base{ID: "acct-1"}}
|
||||
cfg := &EmbyRemoteConfig{BaseURL: "http://localhost:8096"}
|
||||
|
||||
item := map[string]any{
|
||||
"Id": "item-2",
|
||||
"Name": "评分测试",
|
||||
"CriticRating": 9.2,
|
||||
"PremiereDate": "2022-10-01",
|
||||
}
|
||||
|
||||
media := svc.MapRemoteItemToMedia(context.Background(), mount, acct, cfg, item)
|
||||
if media.Rating != 9.2 {
|
||||
t.Fatalf("Rating = %f, want 9.2 from CriticRating", media.Rating)
|
||||
}
|
||||
if media.Year != 2022 {
|
||||
t.Fatalf("Year = %d, want 2022 from PremiereDate", media.Year)
|
||||
}
|
||||
}
|
||||
@@ -124,7 +124,8 @@ func (e *EmbyService) userPayload(u *model.User) map[string]any {
|
||||
}
|
||||
}
|
||||
|
||||
// Views 返回 Emby 中"虚拟根目录"——每个 library 一个条目。
|
||||
// Views 返回 Emby 中"虚拟根目录"——每个 library 一个条目,外加所有启用的
|
||||
// 远程 Emby 挂载的媒体库(联邦聚合)。
|
||||
func (e *EmbyService) Views(ctx context.Context, userID string) (map[string]any, error) {
|
||||
libs, err := e.repo.Library.List(ctx)
|
||||
if err != nil {
|
||||
@@ -132,16 +133,92 @@ func (e *EmbyService) Views(ctx context.Context, userID string) (map[string]any,
|
||||
}
|
||||
libs = FilterDisplayCloudLibraries(ctx, e.repo, libs)
|
||||
visibility := e.mediaVisibility(ctx, userID)
|
||||
items := make([]map[string]any, 0, len(libs))
|
||||
items := make([]map[string]any, 0, len(libs)+4)
|
||||
for _, l := range libs {
|
||||
if !e.libraryVisibleFromCachedVisibility(l, visibility) {
|
||||
continue
|
||||
}
|
||||
items = append(items, e.libraryAsView(ctx, &l))
|
||||
}
|
||||
for _, remote := range e.remoteViews(ctx) {
|
||||
items = append(items, remote)
|
||||
}
|
||||
return map[string]any{"Items": items, "TotalRecordCount": len(items), "StartIndex": 0}, nil
|
||||
}
|
||||
|
||||
// remoteViews 返回全部启用挂载的远程媒体库视图(只有显式挂载的库才出现)。
|
||||
func (e *EmbyService) remoteViews(ctx context.Context) []map[string]any {
|
||||
if e == nil || e.remote == nil {
|
||||
return nil
|
||||
}
|
||||
views, err := e.remote.RemoteLibraries(ctx)
|
||||
if err != nil || len(views) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make([]map[string]any, 0, len(views))
|
||||
for _, v := range views {
|
||||
out = append(out, remoteLibraryViewPayload(v))
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// remoteLibraryViewPayload 把挂载库展示信息标准化为 Emby View payload
|
||||
// (ID 用挂载伪装,名称=挂载显示名)。
|
||||
func remoteLibraryViewPayload(v RemoteLibraryView) map[string]any {
|
||||
encoded := EncodeEmbyRemoteID(v.MountID, v.RemoteID)
|
||||
collectionType := v.CollectionType
|
||||
if !isSupportedEmbyCollectionType(collectionType) {
|
||||
collectionType = "mixed"
|
||||
}
|
||||
name := strings.TrimSpace(v.Library.Name)
|
||||
imageTags := map[string]string{}
|
||||
if strings.TrimSpace(v.Library.CoverURL) != "" {
|
||||
imageTags["Primary"] = encoded
|
||||
}
|
||||
return map[string]any{
|
||||
"Id": encoded,
|
||||
"Name": name,
|
||||
"CollectionType": collectionType,
|
||||
"ServerId": embyServerID,
|
||||
"Type": "CollectionFolder",
|
||||
"IsFolder": true,
|
||||
"Path": "",
|
||||
"SortName": strings.ToLower(name),
|
||||
"DateCreated": time.Now().UTC().Format(time.RFC3339),
|
||||
"CanDelete": false,
|
||||
"CanDownload": false,
|
||||
"DisplayPreferencesId": encoded,
|
||||
"PrimaryImageItemId": encoded,
|
||||
"PrimaryImageAspectRatio": 1.7777777777777777,
|
||||
"RecursiveItemCount": 0,
|
||||
"ChildCount": 0,
|
||||
"SpecialFeatureCount": 0,
|
||||
"EnableMediaSourceDisplay": true,
|
||||
"PlayAccess": "Full",
|
||||
"ExternalUrls": []any{},
|
||||
"ProviderIds": map[string]string{},
|
||||
"Genres": []string{},
|
||||
"Tags": []string{},
|
||||
"ImageTags": imageTags,
|
||||
"BackdropImageTags": []string{},
|
||||
"UserData": map[string]any{
|
||||
"PlaybackPositionTicks": 0,
|
||||
"PlayCount": 0,
|
||||
"IsFavorite": false,
|
||||
"Played": false,
|
||||
"UnplayedItemCount": 0,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func isSupportedEmbyCollectionType(t string) bool {
|
||||
switch t {
|
||||
case "movies", "tvshows", "music", "mixed", "homevideos", "boxsets":
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (e *EmbyService) libraryAsView(ctx context.Context, l *model.Library) map[string]any {
|
||||
collectionType := "movies"
|
||||
switch l.Type {
|
||||
|
||||
@@ -12,8 +12,16 @@ import (
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
)
|
||||
|
||||
// SetFavorite 把 mediaID 标为 userID 的收藏。
|
||||
// SetFavorite 把 mediaID 标为 userID 的收藏。远程 Emby 条目直接透传到对应
|
||||
// 服务器(本地不落库)。
|
||||
func (e *EmbyService) SetFavorite(ctx context.Context, userID, mediaID string, favorite bool) error {
|
||||
if e.remote != nil && IsEmbyRemoteID(mediaID) {
|
||||
acctID, remoteID, _ := DecodeEmbyRemoteID(mediaID)
|
||||
if err := e.ProxyRemoteSetFavorite(ctx, acctID, remoteID, favorite); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if favorite {
|
||||
var f model.Favorite
|
||||
err := e.repo.DB.WithContext(ctx).
|
||||
@@ -31,11 +39,23 @@ func (e *EmbyService) SetFavorite(ctx context.Context, userID, mediaID string, f
|
||||
}
|
||||
|
||||
// MarkPlayed 把 mediaID 标为已看(写一个 100% 进度的 history 行)。
|
||||
// 远程 Emby 条目直接透传到对应服务器(本地不落库)。
|
||||
func (e *EmbyService) MarkPlayed(ctx context.Context, userID, mediaID string, played bool) error {
|
||||
if e.remote != nil && IsEmbyRemoteID(mediaID) {
|
||||
acctID, remoteID, _ := DecodeEmbyRemoteID(mediaID)
|
||||
if err := e.ProxyRemoteSetPlayed(ctx, acctID, remoteID, played); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if !played {
|
||||
return e.repo.DB.WithContext(ctx).
|
||||
err := e.repo.DB.WithContext(ctx).
|
||||
Where("user_id = ? AND media_id = ?", userID, mediaID).
|
||||
Delete(&model.PlaybackHistory{}).Error
|
||||
if err == nil {
|
||||
e.invalidateEmbyItemsCache(ctx)
|
||||
}
|
||||
return err
|
||||
}
|
||||
m, err := e.repo.Media.FindByID(ctx, mediaID)
|
||||
if err != nil || m == nil {
|
||||
@@ -45,7 +65,7 @@ func (e *EmbyService) MarkPlayed(ctx context.Context, userID, mediaID string, pl
|
||||
if dur <= 0 {
|
||||
dur = 1
|
||||
}
|
||||
return e.repo.History.Upsert(ctx, &model.PlaybackHistory{
|
||||
err = e.repo.History.Upsert(ctx, &model.PlaybackHistory{
|
||||
UserID: userID,
|
||||
MediaID: mediaID,
|
||||
PositionMs: dur,
|
||||
@@ -53,6 +73,10 @@ func (e *EmbyService) MarkPlayed(ctx context.Context, userID, mediaID string, pl
|
||||
WatchedAt: time.Now(),
|
||||
Completed: true,
|
||||
})
|
||||
if err == nil {
|
||||
e.invalidateEmbyItemsCache(ctx)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// RecordProgress 记录播放进度(来自 Emby 客户端的 /Sessions/Playing/Progress)。
|
||||
@@ -63,10 +87,27 @@ func (e *EmbyService) RecordProgress(ctx context.Context, userID, mediaID string
|
||||
// runtimeTicks 缺失时回退到 media.DurationSec
|
||||
if m, _ := e.repo.Media.FindByID(ctx, mediaID); m != nil {
|
||||
dur = int64(m.DurationSec) * 1000
|
||||
} else if IsEmbyRemoteID(mediaID) {
|
||||
// 远程挂载条目:尝试从既有历史记录或远程详情补齐时长
|
||||
var oldHist model.PlaybackHistory
|
||||
if err := e.repo.DB.WithContext(ctx).Where("user_id = ? AND media_id = ?", userID, mediaID).First(&oldHist).Error; err == nil && oldHist.DurationMs > 0 {
|
||||
dur = oldHist.DurationMs
|
||||
} else if e.remote != nil {
|
||||
mountID, remoteID, _ := DecodeEmbyRemoteID(mediaID)
|
||||
if mount, acct, _ := e.remote.ResolveMount(ctx, mountID); mount != nil && acct != nil {
|
||||
if item, _ := e.remote.RemoteItem(ctx, mount, acct, remoteID); item != nil {
|
||||
if ticks, ok := item["RunTimeTicks"].(float64); ok && ticks > 0 {
|
||||
dur = int64(ticks) / 10_000
|
||||
} else if ticks, ok := item["RunTimeTicks"].(int64); ok && ticks > 0 {
|
||||
dur = ticks / 10_000
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
completed := dur > 0 && pos >= dur*9/10
|
||||
return e.repo.History.Upsert(ctx, &model.PlaybackHistory{
|
||||
err := e.repo.History.Upsert(ctx, &model.PlaybackHistory{
|
||||
UserID: userID,
|
||||
MediaID: mediaID,
|
||||
PositionMs: pos,
|
||||
@@ -74,6 +115,115 @@ func (e *EmbyService) RecordProgress(ctx context.Context, userID, mediaID string
|
||||
WatchedAt: time.Now(),
|
||||
Completed: completed,
|
||||
})
|
||||
if err == nil {
|
||||
e.invalidateEmbyItemsCache(ctx)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// mergeRemoteUserData applies the current MMTL user's locally recorded playback
|
||||
// state to remote Emby payloads. Remote metadata remains authoritative unless the
|
||||
// user has played the item through MMTL.
|
||||
func (e *EmbyService) mergeRemoteUserData(ctx context.Context, userID string, payload any) error {
|
||||
if strings.TrimSpace(userID) == "" || payload == nil {
|
||||
return nil
|
||||
}
|
||||
items := remoteItemMaps(payload)
|
||||
ids := make([]string, 0, len(items))
|
||||
seen := make(map[string]struct{}, len(items))
|
||||
for _, item := range items {
|
||||
id, _ := item["Id"].(string)
|
||||
if !IsEmbyRemoteID(id) {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; !ok {
|
||||
ids = append(ids, id)
|
||||
seen[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
var histories []model.PlaybackHistory
|
||||
if err := e.repo.DB.WithContext(ctx).Where("user_id = ? AND media_id IN ?", userID, ids).Find(&histories).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
byMediaID := make(map[string]*model.PlaybackHistory, len(histories))
|
||||
for i := range histories {
|
||||
byMediaID[histories[i].MediaID] = &histories[i]
|
||||
}
|
||||
for _, item := range items {
|
||||
id, _ := item["Id"].(string)
|
||||
if h := byMediaID[id]; h != nil {
|
||||
item["UserData"] = mergedRemoteUserData(item["UserData"], h)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func remoteItemMaps(payload any) []map[string]any {
|
||||
items := make([]map[string]any, 0)
|
||||
var visit func(any)
|
||||
visit = func(value any) {
|
||||
switch typed := value.(type) {
|
||||
case map[string]any:
|
||||
if _, ok := typed["Id"].(string); ok {
|
||||
items = append(items, typed)
|
||||
}
|
||||
if nested, ok := typed["Items"]; ok {
|
||||
visit(nested)
|
||||
}
|
||||
case []any:
|
||||
for _, value := range typed {
|
||||
visit(value)
|
||||
}
|
||||
case []map[string]any:
|
||||
for _, value := range typed {
|
||||
visit(value)
|
||||
}
|
||||
}
|
||||
}
|
||||
visit(payload)
|
||||
return items
|
||||
}
|
||||
|
||||
func mergedRemoteUserData(raw any, history *model.PlaybackHistory) map[string]any {
|
||||
userData := map[string]any{}
|
||||
if existing, ok := raw.(map[string]any); ok {
|
||||
for key, value := range existing {
|
||||
userData[key] = value
|
||||
}
|
||||
}
|
||||
duration := history.DurationMs
|
||||
position := history.PositionMs
|
||||
percentage := float64(0)
|
||||
if duration > 0 {
|
||||
percentage = float64(position) / float64(duration) * 100
|
||||
}
|
||||
userData["PlaybackPositionTicks"] = position * 10_000
|
||||
userData["Played"] = history.Completed
|
||||
userData["PlayedPercentage"] = percentage
|
||||
if history.Completed {
|
||||
playCount := 0
|
||||
switch value := userData["PlayCount"].(type) {
|
||||
case int:
|
||||
playCount = value
|
||||
case int64:
|
||||
playCount = int(value)
|
||||
case float64:
|
||||
playCount = int(value)
|
||||
}
|
||||
if playCount < 1 {
|
||||
userData["PlayCount"] = 1
|
||||
}
|
||||
}
|
||||
return userData
|
||||
}
|
||||
|
||||
func (e *EmbyService) invalidateEmbyItemsCache(ctx context.Context) {
|
||||
if e.cache != nil {
|
||||
e.cache.DeletePrefix(ctx, "media:emby:")
|
||||
}
|
||||
}
|
||||
|
||||
func splitCSV(s string) []string {
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
)
|
||||
|
||||
func TestMergedRemoteUserData(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
raw any
|
||||
history model.PlaybackHistory
|
||||
position int64
|
||||
played bool
|
||||
percent float64
|
||||
count int
|
||||
preserve any
|
||||
}{
|
||||
{
|
||||
name: "in-progress preserves remote fields",
|
||||
raw: map[string]any{
|
||||
"PlayCount": 2,
|
||||
"Custom": "remote-value",
|
||||
},
|
||||
history: model.PlaybackHistory{PositionMs: 25_000, DurationMs: 100_000},
|
||||
position: 250_000_000,
|
||||
played: false,
|
||||
percent: 25,
|
||||
count: 2,
|
||||
preserve: "remote-value",
|
||||
},
|
||||
{
|
||||
name: "completed ensures a play count",
|
||||
raw: map[string]any{"PlayCount": 0},
|
||||
history: model.PlaybackHistory{PositionMs: 100_000, DurationMs: 100_000, Completed: true},
|
||||
position: 1_000_000_000,
|
||||
played: true,
|
||||
percent: 100,
|
||||
count: 1,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
out := mergedRemoteUserData(tt.raw, &tt.history)
|
||||
if got := out["PlaybackPositionTicks"]; got != tt.position {
|
||||
t.Fatalf("PlaybackPositionTicks = %#v, want %d", got, tt.position)
|
||||
}
|
||||
if got := out["Played"]; got != tt.played {
|
||||
t.Fatalf("Played = %#v, want %t", got, tt.played)
|
||||
}
|
||||
if got := out["PlayedPercentage"]; got != tt.percent {
|
||||
t.Fatalf("PlayedPercentage = %#v, want %v", got, tt.percent)
|
||||
}
|
||||
if got := out["PlayCount"]; got != tt.count {
|
||||
t.Fatalf("PlayCount = %#v, want %d", got, tt.count)
|
||||
}
|
||||
if tt.preserve != nil && out["Custom"] != tt.preserve {
|
||||
t.Fatalf("Custom = %#v, want %#v", out["Custom"], tt.preserve)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoteItemMapsFindsEnvelopeItems(t *testing.T) {
|
||||
remoteID := EncodeEmbyRemoteID("mount-1", "item-1")
|
||||
payload := map[string]any{
|
||||
"Items": []any{
|
||||
map[string]any{"Id": remoteID},
|
||||
map[string]any{"Id": "local-item"},
|
||||
},
|
||||
}
|
||||
items := remoteItemMaps(payload)
|
||||
if len(items) != 2 {
|
||||
t.Fatalf("item count = %d, want 2", len(items))
|
||||
}
|
||||
if items[0]["Id"] != remoteID {
|
||||
t.Fatalf("first item ID = %#v, want %q", items[0]["Id"], remoteID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecordProgressFallbacksToExistingHistoryDuration(t *testing.T) {
|
||||
svc := newTestEmbyService(t)
|
||||
remoteID := EncodeEmbyRemoteID("mount-test", "item-999")
|
||||
user := &model.User{Username: "resume_test_user", Role: "user", Tier: "free", IsActive: true}
|
||||
if err := svc.repo.User.Create(t.Context(), user); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
|
||||
// 先以有 runtimeTicks 写入首次进度
|
||||
if err := svc.RecordProgress(t.Context(), user.ID, remoteID, 10_000_000, 100_000_000); err != nil {
|
||||
t.Fatalf("first record progress: %v", err)
|
||||
}
|
||||
// 再次上报,但某些客户端此时发了 0 runtimeTicks
|
||||
if err := svc.RecordProgress(t.Context(), user.ID, remoteID, 95_000_000, 0); err != nil {
|
||||
t.Fatalf("second record progress: %v", err)
|
||||
}
|
||||
|
||||
var hist model.PlaybackHistory
|
||||
if err := svc.repo.DB.Where("user_id = ? AND media_id = ?", user.ID, remoteID).First(&hist).Error; err != nil {
|
||||
t.Fatalf("find hist: %v", err)
|
||||
}
|
||||
if hist.DurationMs != 10_000 {
|
||||
t.Fatalf("expected duration 10000ms, got %d", hist.DurationMs)
|
||||
}
|
||||
if !hist.Completed {
|
||||
t.Fatalf("expected 95%% progress to be completed")
|
||||
}
|
||||
}
|
||||
@@ -1,196 +1,358 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"archive/tar"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/ulikunitz/xz"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/config"
|
||||
)
|
||||
|
||||
// AutoInstallFFmpeg is only called by the admin tool-install endpoint. The
|
||||
// server must not auto-download or keep ffmpeg/ffprobe running during startup.
|
||||
func AutoInstallFFmpeg(log *zap.Logger, cfg *config.Config) (ffprobePath, ffmpegPath string) {
|
||||
// 1. 优先使用配置 / PATH / 本机常见软件目录中的现有工具。
|
||||
if path, err := resolveLocalExecutable(cfg.App.FFprobePath, "ffprobe"); err == nil {
|
||||
ffprobePath = path
|
||||
cfg.App.FFprobePath = path
|
||||
log.Info("found local ffprobe", zap.String("path", path))
|
||||
}
|
||||
if path, err := resolveLocalExecutable(cfg.App.FFmpegPath, "ffmpeg"); err == nil {
|
||||
ffmpegPath = path
|
||||
cfg.App.FFmpegPath = path
|
||||
log.Info("found local ffmpeg", zap.String("path", path))
|
||||
}
|
||||
if ffprobePath != "" || ffmpegPath != "" {
|
||||
return ffprobePath, ffmpegPath
|
||||
}
|
||||
// ffmpegDownloadTarget 描述某个平台对应的官方构建下载源。
|
||||
type ffmpegDownloadTarget struct {
|
||||
Label string // 展示名,如 "Windows x86_64"
|
||||
Kind string // 压缩包类型:zip / tar.xz
|
||||
Archives []string // 依次尝试的下载地址(主源 + 备用源)
|
||||
}
|
||||
|
||||
// 2. 检查默认安装位置。
|
||||
defaultDir := getDefaultInstallDir()
|
||||
ffprobeDefault := filepath.Join(defaultDir, "bin", "ffprobe.exe")
|
||||
ffmpegDefault := filepath.Join(defaultDir, "bin", "ffmpeg.exe")
|
||||
|
||||
if _, err := os.Stat(ffprobeDefault); err == nil {
|
||||
log.Info("在默认位置找到 ffprobe", zap.String("path", ffprobeDefault))
|
||||
return ffprobeDefault, ffmpegDefault
|
||||
}
|
||||
|
||||
// 3. 尝试自动安装。
|
||||
log.Warn("未找到 ffmpeg/ffprobe,尝试自动安装...")
|
||||
installed, err := tryAutoInstall(log, defaultDir)
|
||||
if err != nil {
|
||||
log.Error("自动安装失败,请手动安装 ffmpeg", zap.Error(err))
|
||||
return "", ""
|
||||
}
|
||||
|
||||
if installed {
|
||||
if _, err := os.Stat(ffprobeDefault); err == nil {
|
||||
log.Info("自动安装成功", zap.String("path", ffprobeDefault))
|
||||
// 更新配置
|
||||
updateConfigPaths(cfg, ffprobeDefault, ffmpegDefault)
|
||||
return ffprobeDefault, ffmpegDefault
|
||||
// ffmpegTargetForPlatform 按当前运行环境(OS+架构)选择下载源。
|
||||
func ffmpegTargetForPlatform() (*ffmpegDownloadTarget, error) {
|
||||
switch runtime.GOOS {
|
||||
case "windows":
|
||||
switch runtime.GOARCH {
|
||||
case "amd64":
|
||||
return &ffmpegDownloadTarget{
|
||||
Label: "Windows x86_64",
|
||||
Kind: "zip",
|
||||
Archives: []string{
|
||||
"https://www.gyan.dev/ffmpeg/builds/ffmpeg-release-essentials.zip",
|
||||
"https://github.com/BtbN/FFmpeg-Builds/releases/download/latest/ffmpeg-master-latest-win64-gpl.zip",
|
||||
},
|
||||
}, nil
|
||||
case "386":
|
||||
return &ffmpegDownloadTarget{
|
||||
Label: "Windows x86",
|
||||
Kind: "zip",
|
||||
Archives: []string{
|
||||
"https://github.com/BtbN/FFmpeg-Builds/releases/download/latest/ffmpeg-master-latest-win32-gpl.zip",
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
case "linux":
|
||||
switch runtime.GOARCH {
|
||||
case "amd64":
|
||||
return &ffmpegDownloadTarget{
|
||||
Label: "Linux x86_64",
|
||||
Kind: "tar.xz",
|
||||
Archives: []string{
|
||||
"https://github.com/BtbN/FFmpeg-Builds/releases/download/latest/ffmpeg-master-latest-linux64-gpl.tar.xz",
|
||||
"https://johnvansickle.com/ffmpeg/releases/ffmpeg-release-amd64-static.tar.xz",
|
||||
},
|
||||
}, nil
|
||||
case "arm64":
|
||||
return &ffmpegDownloadTarget{
|
||||
Label: "Linux ARM64",
|
||||
Kind: "tar.xz",
|
||||
Archives: []string{
|
||||
"https://github.com/BtbN/FFmpeg-Builds/releases/download/latest/ffmpeg-master-latest-linuxarm64-gpl.tar.xz",
|
||||
"https://johnvansickle.com/ffmpeg/releases/ffmpeg-release-arm64-static.tar.xz",
|
||||
},
|
||||
}, nil
|
||||
case "arm":
|
||||
return &ffmpegDownloadTarget{
|
||||
Label: "Linux ARM (32 位)",
|
||||
Kind: "tar.xz",
|
||||
Archives: []string{
|
||||
"https://github.com/BtbN/FFmpeg-Builds/releases/download/latest/ffmpeg-master-latest-linuxarmhf-gpl.tar.xz",
|
||||
},
|
||||
}, nil
|
||||
case "loong64":
|
||||
return &ffmpegDownloadTarget{
|
||||
Label: "Linux LoongArch64",
|
||||
Kind: "tar.xz",
|
||||
Archives: []string{
|
||||
"https://github.com/BtbN/FFmpeg-Builds/releases/download/latest/ffmpeg-master-latest-linuxloongarch64-gpl.tar.xz",
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
case "darwin":
|
||||
switch runtime.GOARCH {
|
||||
case "amd64":
|
||||
return &ffmpegDownloadTarget{
|
||||
Label: "macOS x86_64",
|
||||
Kind: "zip",
|
||||
Archives: []string{
|
||||
"https://github.com/BtbN/FFmpeg-Builds/releases/download/latest/ffmpeg-master-latest-osx64-gpl.zip",
|
||||
},
|
||||
}, nil
|
||||
case "arm64":
|
||||
return &ffmpegDownloadTarget{
|
||||
Label: "macOS Apple Silicon",
|
||||
Kind: "zip",
|
||||
Archives: []string{
|
||||
"https://github.com/BtbN/FFmpeg-Builds/releases/download/latest/ffmpeg-master-latest-osxarm64-gpl.zip",
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
||||
return "", ""
|
||||
return nil, fmt.Errorf("暂不支持自动下载的平台 %s/%s,请手动填写 ffmpeg/ffprobe 路径", runtime.GOOS, runtime.GOARCH)
|
||||
}
|
||||
|
||||
// getDefaultInstallDir 返回默认安装目录
|
||||
func getDefaultInstallDir() string {
|
||||
exePath, err := os.Executable()
|
||||
// installFFmpegTools 按平台下载并安装 ffmpeg/ffprobe 到 data/tools/ffmpeg/,
|
||||
// 返回两个可执行文件的绝对路径。progress 用于回传阶段消息(UI 展示)。
|
||||
func installFFmpegTools(ctx context.Context, log *zap.Logger, cfg *config.Config, progress func(string)) (ffmpegPath, ffprobePath string, err error) {
|
||||
target, err := ffmpegTargetForPlatform()
|
||||
if err != nil {
|
||||
return "./tools/ffmpeg"
|
||||
return "", "", err
|
||||
}
|
||||
exeDir := filepath.Dir(exePath)
|
||||
return filepath.Join(exeDir, "tools", "ffmpeg")
|
||||
}
|
||||
|
||||
// tryAutoInstall 尝试自动下载并安装 ffmpeg
|
||||
func tryAutoInstall(log *zap.Logger, installDir string) (bool, error) {
|
||||
if runtime.GOOS == "windows" {
|
||||
return downloadFFmpegWindows(log, installDir)
|
||||
}
|
||||
|
||||
return false, fmt.Errorf("不支持的操作系统: %s", runtime.GOOS)
|
||||
}
|
||||
|
||||
// downloadFFmpegWindows 下载 Windows 版本的 ffmpeg
|
||||
func downloadFFmpegWindows(log *zap.Logger, installDir string) (bool, error) {
|
||||
log.Info("开始下载 ffmpeg...")
|
||||
|
||||
// 创建安装目录
|
||||
installDir := filepath.Join(cfg.App.DataDir, "tools", "ffmpeg")
|
||||
if err := os.MkdirAll(installDir, 0o750); err != nil {
|
||||
return false, fmt.Errorf("创建安装目录失败: %w", err)
|
||||
return "", "", fmt.Errorf("创建安装目录失败: %w", err)
|
||||
}
|
||||
|
||||
tempDir, err := os.MkdirTemp("", "mmtl-ffmpeg-*")
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("创建临时目录失败: %w", err)
|
||||
return "", "", fmt.Errorf("创建临时目录失败: %w", err)
|
||||
}
|
||||
defer os.RemoveAll(tempDir)
|
||||
|
||||
// 下载 URL (使用 gyani.org 的静态构建)
|
||||
arch := "win64"
|
||||
if !is64Bit() {
|
||||
arch = "win32"
|
||||
progress("下载 " + target.Label + " 版本…")
|
||||
archivePath := filepath.Join(tempDir, "ffmpeg-archive."+target.Kind)
|
||||
if err := downloadFFmpegArchive(ctx, log, target.Archives, archivePath); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
|
||||
// 先尝试从 gyan.dev 下载(更可靠)
|
||||
downloadURL := fmt.Sprintf("https://www.gyan.dev/ffmpeg/builds/ffmpeg-release-essentials.zip")
|
||||
progress("解压…")
|
||||
extractDir := filepath.Join(tempDir, "extract")
|
||||
if err := extractFFmpegArchive(target.Kind, archivePath, extractDir); err != nil {
|
||||
return "", "", fmt.Errorf("解压失败: %w", err)
|
||||
}
|
||||
|
||||
log.Info("下载 ffmpeg", zap.String("url", downloadURL))
|
||||
exeSuffix := ""
|
||||
if runtime.GOOS == "windows" {
|
||||
exeSuffix = ".exe"
|
||||
}
|
||||
srcFFmpeg, srcFFprobe, err := locateFFmpegBinaries(extractDir, exeSuffix)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
|
||||
// 使用 Go 下载
|
||||
zipPath := filepath.Join(tempDir, "ffmpeg.zip")
|
||||
if err := downloadFile(log, downloadURL, zipPath); err != nil {
|
||||
// 尝试备用 URL
|
||||
backupURL := fmt.Sprintf("https://github.com/BtbN/FFmpeg-Builds/releases/download/latest/ffmpeg-master-latest-%s-gpl.zip", arch)
|
||||
log.Info("尝试备用下载地址", zap.String("url", backupURL))
|
||||
if err2 := downloadFile(log, backupURL, zipPath); err2 != nil {
|
||||
return false, fmt.Errorf("下载失败: %v, %v", err, err2)
|
||||
progress("安装到 data 目录…")
|
||||
ffmpegPath = filepath.Join(installDir, "ffmpeg"+exeSuffix)
|
||||
ffprobePath = filepath.Join(installDir, "ffprobe"+exeSuffix)
|
||||
if err := copyFileMode(srcFFmpeg, ffmpegPath); err != nil {
|
||||
return "", "", fmt.Errorf("复制 ffmpeg 失败: %w", err)
|
||||
}
|
||||
if err := copyFileMode(srcFFprobe, ffprobePath); err != nil {
|
||||
_ = os.Remove(ffmpegPath)
|
||||
return "", "", fmt.Errorf("复制 ffprobe 失败: %w", err)
|
||||
}
|
||||
|
||||
// 验证两个工具都能运行(失败则回滚,避免留下坏文件)。
|
||||
for _, bin := range []string{ffmpegPath, ffprobePath} {
|
||||
cmd := exec.CommandContext(ctx, bin, "-version") // #nosec G204 -- bin 是安装目录中刚写入的固定文件名。
|
||||
if out, verr := cmd.Output(); verr != nil {
|
||||
_ = os.Remove(ffmpegPath)
|
||||
_ = os.Remove(ffprobePath)
|
||||
return "", "", fmt.Errorf("安装后 %s 无法运行:%v", filepath.Base(bin), verr)
|
||||
} else if log != nil {
|
||||
log.Info("ffmpeg 工具安装验证通过", zap.String("bin", filepath.Base(bin)),
|
||||
zap.String("version", strings.TrimSpace(strings.SplitN(string(out), "\n", 2)[0])))
|
||||
}
|
||||
}
|
||||
|
||||
// 解压
|
||||
log.Info("解压 ffmpeg...")
|
||||
extractDir := filepath.Join(tempDir, "extract")
|
||||
if err := unzip(log, zipPath, extractDir); err != nil {
|
||||
return false, fmt.Errorf("解压失败: %w", err)
|
||||
}
|
||||
|
||||
packageRoot, err := findFFmpegPackageRoot(extractDir)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if err := copyDirContents(packageRoot, installDir); err != nil {
|
||||
return false, fmt.Errorf("复制 ffmpeg 文件失败: %w", err)
|
||||
}
|
||||
|
||||
ffmpegBin := filepath.Join(installDir, "bin", "ffmpeg.exe")
|
||||
ffprobeBin := filepath.Join(installDir, "bin", "ffprobe.exe")
|
||||
if _, err := os.Stat(ffmpegBin); err != nil {
|
||||
return false, fmt.Errorf("安装后未找到 ffmpeg: %w", err)
|
||||
}
|
||||
if _, err := os.Stat(ffprobeBin); err != nil {
|
||||
return false, fmt.Errorf("安装后未找到 ffprobe: %w", err)
|
||||
}
|
||||
|
||||
log.Info("ffmpeg 安装完成", zap.String("dir", installDir))
|
||||
return true, nil
|
||||
progress("安装完成")
|
||||
return ffmpegPath, ffprobePath, nil
|
||||
}
|
||||
|
||||
// downloadFile 下载文件
|
||||
func downloadFile(log *zap.Logger, url, filepath string) error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
|
||||
defer cancel()
|
||||
// downloadFFmpegArchive 按顺序尝试下载源,全部失败才返回错误。
|
||||
func downloadFFmpegArchive(ctx context.Context, log *zap.Logger, urls []string, dest string) error {
|
||||
var lastErr error
|
||||
for i, u := range urls {
|
||||
if i > 0 && log != nil {
|
||||
log.Warn("ffmpeg 主下载源不可用,切换备用源", zap.String("url", u))
|
||||
}
|
||||
if err := downloadFFmpegFile(ctx, log, u, dest); err != nil {
|
||||
lastErr = err
|
||||
if log != nil {
|
||||
log.Warn("ffmpeg 下载失败", zap.String("url", u), zap.Error(err))
|
||||
}
|
||||
continue
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("所有下载源均失败:%v", lastErr)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", url, nil)
|
||||
// downloadFFmpegFile 下载单个归档文件(最多 10 分钟,限制大小上限)。
|
||||
func downloadFFmpegFile(ctx context.Context, log *zap.Logger, url, dest string) error {
|
||||
ctx, cancel := context.WithTimeout(ctx, 10*time.Minute)
|
||||
defer cancel()
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
req.Header.Set("User-Agent", "MMTL/ffmpeg-installer ("+runtime.GOOS+"/"+runtime.GOARCH+")")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("下载失败,HTTP 状态码: %d", resp.StatusCode)
|
||||
return fmt.Errorf("下载失败,HTTP 状态码 %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
out, err := os.Create(filepath) // #nosec G304 -- filepath is generated by the installer under its temporary download directory.
|
||||
out, err := os.Create(dest) // #nosec G304 -- dest 是安装器在临时目录下生成的文件。
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer out.Close()
|
||||
|
||||
_, err = io.Copy(out, resp.Body)
|
||||
return err
|
||||
n, err := io.Copy(out, io.LimitReader(resp.Body, 500<<20+1)) // 归档上限 500MB
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n > 500<<20 {
|
||||
return fmt.Errorf("归档文件过大(>500MB): %s", url)
|
||||
}
|
||||
if log != nil {
|
||||
log.Info("ffmpeg 归档下载完成", zap.String("url", url), zap.Int64("bytes", n))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// updateConfigPaths 更新配置文件中的路径
|
||||
func updateConfigPaths(cfg *config.Config, ffprobePath, ffmpegPath string) {
|
||||
cfg.App.FFprobePath = ffprobePath
|
||||
cfg.App.FFmpegPath = ffmpegPath
|
||||
|
||||
// 保存到配置文件
|
||||
// 这里需要调用 config 包的保存函数
|
||||
log := zap.L().Named("config")
|
||||
log.Info("已更新 ffmpeg 路径配置",
|
||||
zap.String("ffprobe", ffprobePath),
|
||||
zap.String("ffmpeg", ffmpegPath))
|
||||
// extractFFmpegArchive 按类型解压 zip 或 tar.xz。
|
||||
func extractFFmpegArchive(kind, archivePath, destDir string) error {
|
||||
switch kind {
|
||||
case "zip":
|
||||
return unzip(nil, archivePath, destDir)
|
||||
case "tar.xz":
|
||||
return untarXZ(archivePath, destDir)
|
||||
default:
|
||||
return fmt.Errorf("不支持的归档类型: %s", kind)
|
||||
}
|
||||
}
|
||||
|
||||
// is64Bit 检查是否为 64 位系统
|
||||
func is64Bit() bool {
|
||||
return true // 简化处理,假设为 64 位
|
||||
// untarXZ 解压 .tar.xz 归档(GNU tar + xz 流式解压,纯 Go 无外部依赖),
|
||||
// 路径安全校验与 ZIP 解压一致。
|
||||
func untarXZ(archivePath, destDir string) error {
|
||||
if err := os.MkdirAll(destDir, 0o750); err != nil {
|
||||
return err
|
||||
}
|
||||
destRoot, err := filepath.Abs(destDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
f, err := os.Open(archivePath) // #nosec G304 -- archivePath 是安装器在临时目录下生成的文件。
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
xzReader, err := xz.NewReader(f)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tr := tar.NewReader(xzReader)
|
||||
var totalWritten int64
|
||||
for {
|
||||
hdr, err := tr.Next()
|
||||
if err == io.EOF {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
target, err := safeZipTarget(destRoot, hdr.Name) // 与 ZIP 相同的路径穿越防护
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
switch hdr.Typeflag {
|
||||
case tar.TypeDir:
|
||||
if err := os.MkdirAll(target, 0o750); err != nil {
|
||||
return err
|
||||
}
|
||||
case tar.TypeReg, tar.TypeRegA:
|
||||
if err := os.MkdirAll(filepath.Dir(target), 0o750); err != nil {
|
||||
return err
|
||||
}
|
||||
dst, err := os.OpenFile(target, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, os.FileMode(hdr.Mode).Perm())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
written, copyErr := io.Copy(dst, io.LimitReader(tr, maxFFmpegZipEntryBytes+1))
|
||||
totalWritten += written
|
||||
closeErr := dst.Close()
|
||||
if copyErr != nil {
|
||||
return copyErr
|
||||
}
|
||||
if closeErr != nil {
|
||||
return closeErr
|
||||
}
|
||||
if written > maxFFmpegZipEntryBytes || totalWritten > maxFFmpegZipTotalBytes {
|
||||
return fmt.Errorf("tar 内容过大: %s", hdr.Name)
|
||||
}
|
||||
default:
|
||||
// 符号链接/设备等一律跳过(静态构建不会依赖它们)。
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// locateFFmpegBinaries 在解压目录中查找 ffmpeg/ffprobe 可执行文件(兼容
|
||||
// 不同构建包的目录布局:gyan 的 bin/、BtbN/johnvansickle 的根目录等)。
|
||||
func locateFFmpegBinaries(root, exeSuffix string) (ffmpeg, ffprobe string, err error) {
|
||||
wantFFmpeg := "ffmpeg" + strings.ToLower(exeSuffix)
|
||||
wantFFprobe := "ffprobe" + strings.ToLower(exeSuffix)
|
||||
err = filepath.WalkDir(root, func(path string, d os.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if d.IsDir() {
|
||||
return nil
|
||||
}
|
||||
switch strings.ToLower(d.Name()) {
|
||||
case wantFFmpeg:
|
||||
if ffmpeg == "" {
|
||||
ffmpeg = path
|
||||
}
|
||||
case wantFFprobe:
|
||||
if ffprobe == "" {
|
||||
ffprobe = path
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("扫描解压目录失败: %w", err)
|
||||
}
|
||||
if ffmpeg == "" || ffprobe == "" {
|
||||
return "", "", fmt.Errorf("解压内容中未找到 ffmpeg/ffprobe 可执行文件")
|
||||
}
|
||||
return ffmpeg, ffprobe, nil
|
||||
}
|
||||
|
||||
// copyFileMode 复制文件并赋予可执行权限。
|
||||
func copyFileMode(src, dst string) error {
|
||||
in, err := os.Open(src) // #nosec G304 -- src 来自解压目录遍历结果。
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer in.Close()
|
||||
out, err := os.OpenFile(dst, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o755)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := io.Copy(out, in); err != nil {
|
||||
_ = out.Close()
|
||||
return err
|
||||
}
|
||||
return out.Close()
|
||||
}
|
||||
|
||||
@@ -41,7 +41,9 @@ func unzip(log *zap.Logger, zipPath, destDir string) error {
|
||||
}
|
||||
info := file.FileInfo()
|
||||
if info.Mode()&os.ModeSymlink != 0 {
|
||||
log.Warn("跳过 ZIP 符号链接", zap.String("name", file.Name))
|
||||
if log != nil {
|
||||
log.Warn("跳过 ZIP 符号链接", zap.String("name", file.Name))
|
||||
}
|
||||
continue
|
||||
}
|
||||
if info.IsDir() {
|
||||
@@ -109,94 +111,3 @@ func safeZipTarget(destRoot, name string) (string, error) {
|
||||
}
|
||||
return targetAbs, nil
|
||||
}
|
||||
|
||||
func findFFmpegPackageRoot(root string) (string, error) {
|
||||
var ffmpegPath string
|
||||
var ffprobePath string
|
||||
|
||||
err := filepath.WalkDir(root, func(path string, d os.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if d.IsDir() {
|
||||
return nil
|
||||
}
|
||||
|
||||
switch strings.ToLower(d.Name()) {
|
||||
case "ffmpeg.exe":
|
||||
ffmpegPath = path
|
||||
case "ffprobe.exe":
|
||||
ffprobePath = path
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("扫描解压目录失败: %w", err)
|
||||
}
|
||||
if ffmpegPath == "" || ffprobePath == "" {
|
||||
return "", fmt.Errorf("解压后未找到 ffmpeg/ffprobe 可执行文件")
|
||||
}
|
||||
|
||||
return filepath.Dir(filepath.Dir(ffmpegPath)), nil
|
||||
}
|
||||
|
||||
func copyDirContents(srcDir, dstDir string) error {
|
||||
entries, err := os.ReadDir(srcDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, entry := range entries {
|
||||
srcPath := filepath.Join(srcDir, entry.Name())
|
||||
dstPath := filepath.Join(dstDir, entry.Name())
|
||||
if err := copyTree(srcPath, dstPath); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func copyTree(srcPath, dstPath string) error {
|
||||
info, err := os.Stat(srcPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if info.IsDir() {
|
||||
if err := os.MkdirAll(dstPath, info.Mode()); err != nil {
|
||||
return err
|
||||
}
|
||||
entries, err := os.ReadDir(srcPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, entry := range entries {
|
||||
if err := copyTree(filepath.Join(srcPath, entry.Name()), filepath.Join(dstPath, entry.Name())); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
in, err := os.Open(srcPath) // #nosec G304 -- srcPath is produced by walking the validated extracted ffmpeg package tree.
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer in.Close()
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(dstPath), 0o750); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
out, err := os.Create(dstPath) // #nosec G304 -- dstPath is generated under the configured ffmpeg install directory.
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer out.Close()
|
||||
|
||||
if _, err := io.Copy(out, in); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return out.Close()
|
||||
}
|
||||
|
||||
@@ -1,57 +0,0 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"os"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
)
|
||||
|
||||
// CheckFFmpegStatus 检查 ffmpeg/ffprobe 状态 (供 API 使用)
|
||||
func CheckFFmpegStatus(ffprobePath, ffmpegPath string) map[string]interface{} {
|
||||
status := map[string]interface{}{
|
||||
"ffprobe_installed": false,
|
||||
"ffmpeg_installed": false,
|
||||
"auto_installable": runtime.GOOS == "windows",
|
||||
}
|
||||
|
||||
if ffprobePath != "" {
|
||||
if _, err := os.Stat(ffprobePath); err == nil {
|
||||
status["ffprobe_installed"] = true
|
||||
status["ffprobe_path"] = ffprobePath
|
||||
|
||||
// 获取版本
|
||||
cmd := exec.Command(ffprobePath, "-version")
|
||||
out, err := cmd.Output()
|
||||
if err == nil {
|
||||
// 提取版本信息(第一行)
|
||||
lines := bytes.Split(out, []byte("\n"))
|
||||
if len(lines) > 0 {
|
||||
version := string(bytes.TrimSpace(lines[0]))
|
||||
status["ffprobe_version"] = version
|
||||
status["ffprobe_security"] = EvaluateFFmpegSecurity(version)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if ffmpegPath != "" {
|
||||
if _, err := os.Stat(ffmpegPath); err == nil {
|
||||
status["ffmpeg_installed"] = true
|
||||
status["ffmpeg_path"] = ffmpegPath
|
||||
|
||||
cmd := exec.Command(ffmpegPath, "-version")
|
||||
out, err := cmd.Output()
|
||||
if err == nil {
|
||||
lines := bytes.Split(out, []byte("\n"))
|
||||
if len(lines) > 0 {
|
||||
version := string(bytes.TrimSpace(lines[0]))
|
||||
status["ffmpeg_version"] = version
|
||||
status["ffmpeg_security"] = EvaluateFFmpegSecurity(version)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return status
|
||||
}
|
||||
@@ -0,0 +1,211 @@
|
||||
// Package service — ffmpeg/ffprobe 自动下载安装。
|
||||
//
|
||||
// FFmpegToolsService 负责「一键下载」:点击后按当前运行平台(OS+架构)选择
|
||||
// 官方构建包(Windows: gyan.dev / BtbN;Linux: BtbN / johnvansickle;
|
||||
// macOS: BtbN),下载解压 ffmpeg/ffprobe 到 data 目录(data/tools/ffmpeg/),
|
||||
// 并把绝对路径写入设置(ffmpeg.path / ffprobe.path),系统随即使用安装的
|
||||
// 工具,无需手动填写路径。
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/config"
|
||||
"github.com/ShukeBta/MMTL/internal/repository"
|
||||
)
|
||||
|
||||
// FFmpegToolsService 管理 ffmpeg/ffprobe 的自动下载安装状态(单飞:同一时间
|
||||
// 只允许一个安装任务)。
|
||||
type FFmpegToolsService struct {
|
||||
cfg *config.Config
|
||||
log *zap.Logger
|
||||
repo *repository.Container
|
||||
|
||||
mu sync.Mutex
|
||||
running bool
|
||||
msg string // 最近阶段/结果消息
|
||||
errMsg string // 最近一次失败原因
|
||||
started time.Time // 最近一次安装开始时间
|
||||
done time.Time // 最近一次安装结束时间
|
||||
}
|
||||
|
||||
// NewFFmpegToolsService 构造工具安装服务。
|
||||
func NewFFmpegToolsService(cfg *config.Config, log *zap.Logger, repo *repository.Container) *FFmpegToolsService {
|
||||
return &FFmpegToolsService{cfg: cfg, log: log, repo: repo}
|
||||
}
|
||||
|
||||
// ffmpegInstallDir 返回 data 目录下的安装位置。
|
||||
func (s *FFmpegToolsService) ffmpegInstallDir() string {
|
||||
return filepath.Join(s.cfg.App.DataDir, "tools", "ffmpeg")
|
||||
}
|
||||
|
||||
// installedBinaries 检查安装目录中是否已存在 ffmpeg/ffprobe 可执行文件。
|
||||
func (s *FFmpegToolsService) installedBinaries() (ffmpeg, ffprobe string) {
|
||||
exe := ""
|
||||
if runtime.GOOS == "windows" {
|
||||
exe = ".exe"
|
||||
}
|
||||
ffmpeg = filepath.Join(s.ffmpegInstallDir(), "ffmpeg"+exe)
|
||||
ffprobe = filepath.Join(s.ffmpegInstallDir(), "ffprobe"+exe)
|
||||
if _, err := os.Stat(ffmpeg); err != nil {
|
||||
return "", ""
|
||||
}
|
||||
if _, err := os.Stat(ffprobe); err != nil {
|
||||
return "", ""
|
||||
}
|
||||
return ffmpeg, ffprobe
|
||||
}
|
||||
|
||||
// ffToolVersion 取工具第一行版本信息。
|
||||
func ffToolVersion(path string) string {
|
||||
if path == "" {
|
||||
return ""
|
||||
}
|
||||
out, err := exec.Command(path, "-version").Output() // #nosec G204 -- path 来自配置/安装目录中的已知工具。
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
line := strings.SplitN(strings.TrimSpace(string(out)), "\n", 2)
|
||||
if len(line) == 0 || strings.TrimSpace(line[0]) == "" {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(line[0])
|
||||
}
|
||||
|
||||
// ffToolInfo 是单个工具的安装状态(返回给前端展示)。
|
||||
type ffToolInfo struct {
|
||||
Installed bool `json:"installed"`
|
||||
Path string `json:"path,omitempty"`
|
||||
Version string `json:"version,omitempty"`
|
||||
}
|
||||
|
||||
// Status 返回当前安装状态(供 GET /api/admin/tools/ffmpeg/status 使用)。
|
||||
func (s *FFmpegToolsService) Status(ctx context.Context) map[string]any {
|
||||
s.mu.Lock()
|
||||
running, msg, errMsg, started, done := s.running, s.msg, s.errMsg, s.started, s.done
|
||||
s.mu.Unlock()
|
||||
|
||||
startedAt, doneAt := "", ""
|
||||
if !started.IsZero() {
|
||||
startedAt = started.Format(time.RFC3339)
|
||||
}
|
||||
if !done.IsZero() {
|
||||
doneAt = done.Format(time.RFC3339)
|
||||
}
|
||||
|
||||
out := map[string]any{
|
||||
"installing": running,
|
||||
"message": msg,
|
||||
"error": errMsg,
|
||||
"started_at": startedAt,
|
||||
"finished_at": doneAt,
|
||||
"install_dir": s.ffmpegInstallDir(),
|
||||
}
|
||||
target, targetErr := ffmpegTargetForPlatform()
|
||||
if targetErr != nil {
|
||||
out["target"] = map[string]any{"label": targetErr.Error()}
|
||||
} else {
|
||||
out["target"] = map[string]any{
|
||||
"os": runtime.GOOS,
|
||||
"arch": runtime.GOARCH,
|
||||
"label": target.Label,
|
||||
}
|
||||
}
|
||||
// 报告「系统当前实际会使用」的工具:优先已生效配置(安装完成会把设置指到
|
||||
// data 目录),其次 PATH / 常见目录。
|
||||
ffmpegPath, ferr := resolveLocalExecutable(s.cfg.App.FFmpegPath, "ffmpeg")
|
||||
ffprobePath, perr := resolveLocalExecutable(s.cfg.App.FFprobePath, "ffprobe")
|
||||
out["ffmpeg"] = ffToolInfo{Installed: ferr == nil, Path: ffmpegPath, Version: ffToolVersion(ffmpegPath)}
|
||||
out["ffprobe"] = ffToolInfo{Installed: perr == nil, Path: ffprobePath, Version: ffToolVersion(ffprobePath)}
|
||||
return out
|
||||
}
|
||||
|
||||
// StartInstall 启动后台安装(幂等)。正在安装时返回错误;data 目录已有完整
|
||||
// 工具时直接应用路径设置并返回(无需重新下载)。
|
||||
func (s *FFmpegToolsService) StartInstall(ctx context.Context) error {
|
||||
s.mu.Lock()
|
||||
if s.running {
|
||||
s.mu.Unlock()
|
||||
return errors.New("工具正在安装中,请稍候")
|
||||
}
|
||||
if ffmpeg, ffprobe := s.installedBinaries(); ffmpeg != "" && ffprobe != "" {
|
||||
s.mu.Unlock()
|
||||
s.setMessage("检测到已安装,直接应用配置")
|
||||
if err := s.applyInstalledPaths(ctx, ffmpeg, ffprobe); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
s.running = true
|
||||
s.errMsg = ""
|
||||
s.started = time.Now()
|
||||
s.mu.Unlock()
|
||||
|
||||
s.setMessage("准备下载…")
|
||||
go s.runInstall()
|
||||
return nil
|
||||
}
|
||||
|
||||
// runInstall 在后台执行下载、解压、验证与配置落盘。
|
||||
func (s *FFmpegToolsService) runInstall() {
|
||||
defer func() {
|
||||
s.mu.Lock()
|
||||
s.running = false
|
||||
s.done = time.Now()
|
||||
s.mu.Unlock()
|
||||
}()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
ffmpegPath, ffprobePath, err := installFFmpegTools(ctx, s.log, s.cfg, s.setMessage)
|
||||
if err != nil {
|
||||
s.mu.Lock()
|
||||
s.errMsg = err.Error()
|
||||
s.msg = "安装失败"
|
||||
s.mu.Unlock()
|
||||
s.log.Error("ffmpeg 工具安装失败", zap.Error(err))
|
||||
return
|
||||
}
|
||||
if err := s.applyInstalledPaths(ctx, ffmpegPath, ffprobePath); err != nil {
|
||||
s.mu.Lock()
|
||||
s.errMsg = "安装完成,但写入设置失败:" + err.Error()
|
||||
s.msg = "安装完成,设置写入失败"
|
||||
s.mu.Unlock()
|
||||
s.log.Error("写入 ffmpeg 工具路径设置失败", zap.Error(err))
|
||||
return
|
||||
}
|
||||
s.setMessage("安装完成")
|
||||
s.log.Info("ffmpeg 工具安装完成",
|
||||
zap.String("ffmpeg", ffmpegPath), zap.String("ffprobe", ffprobePath))
|
||||
}
|
||||
|
||||
// applyInstalledPaths 把安装后的路径写入设置表并热应用到运行配置。
|
||||
func (s *FFmpegToolsService) applyInstalledPaths(ctx context.Context, ffmpeg, ffprobe string) error {
|
||||
if s.repo != nil && s.repo.Setting != nil {
|
||||
if err := s.repo.Setting.Set(ctx, "ffmpeg.path", ffmpeg); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.repo.Setting.Set(ctx, "ffprobe.path", ffprobe); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
ApplyRuntimeSetting(s.cfg, "ffmpeg.path", ffmpeg)
|
||||
ApplyRuntimeSetting(s.cfg, "ffprobe.path", ffprobe)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *FFmpegToolsService) setMessage(msg string) {
|
||||
s.mu.Lock()
|
||||
s.msg = msg
|
||||
s.mu.Unlock()
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/config"
|
||||
)
|
||||
|
||||
func TestFFmpegTargetForPlatform(t *testing.T) {
|
||||
target, err := ffmpegTargetForPlatform()
|
||||
switch runtime.GOOS {
|
||||
case "windows", "linux", "darwin":
|
||||
if err != nil {
|
||||
t.Fatalf("supported platform %s/%s should resolve a target: %v", runtime.GOOS, runtime.GOARCH, err)
|
||||
}
|
||||
if target == nil || target.Label == "" || len(target.Archives) == 0 {
|
||||
t.Fatalf("target incomplete: %#v", target)
|
||||
}
|
||||
if target.Kind != "zip" && target.Kind != "tar.xz" {
|
||||
t.Fatalf("unexpected archive kind: %s", target.Kind)
|
||||
}
|
||||
for _, u := range target.Archives {
|
||||
if !strings.HasPrefix(u, "https://") {
|
||||
t.Fatalf("archive url not https: %s", u)
|
||||
}
|
||||
}
|
||||
default:
|
||||
if err == nil {
|
||||
t.Fatalf("unsupported platform %s/%s should fail", runtime.GOOS, runtime.GOARCH)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeZipTargetRejectsTraversal(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
for _, name := range []string{"../evil", "..\\evil", "/etc/passwd", "a/../../evil"} {
|
||||
if _, err := safeZipTarget(root, name); err == nil {
|
||||
t.Fatalf("expected traversal rejection for %q", name)
|
||||
}
|
||||
}
|
||||
if _, err := safeZipTarget(root, "bin/ffmpeg.exe"); err != nil {
|
||||
t.Fatalf("valid relative path should pass: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFFmpegToolsStatusNoPanic(t *testing.T) {
|
||||
svc := NewFFmpegToolsService(&config.Config{}, zap.NewNop(), nil)
|
||||
st := svc.Status(context.Background())
|
||||
for _, key := range []string{"installing", "message", "error", "install_dir", "target", "ffmpeg", "ffprobe"} {
|
||||
if _, ok := st[key]; !ok {
|
||||
t.Fatalf("status missing key %q: %#v", key, st)
|
||||
}
|
||||
}
|
||||
ffmpeg, ok := st["ffmpeg"].(ffToolInfo)
|
||||
if !ok {
|
||||
t.Fatalf("ffmpeg field not ffToolInfo: %T", st["ffmpeg"])
|
||||
}
|
||||
if ffmpeg.Installed {
|
||||
t.Fatalf("empty config should not report installed ffmpeg: %#v", st)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartInstallRejectsConcurrent(t *testing.T) {
|
||||
svc := NewFFmpegToolsService(&config.Config{App: config.AppConfig{DataDir: t.TempDir()}}, zap.NewNop(), nil)
|
||||
// 不真实运行:直接占用 running 标记模拟进行中的安装。
|
||||
svc.mu.Lock()
|
||||
svc.running = true
|
||||
svc.mu.Unlock()
|
||||
if err := svc.StartInstall(context.Background()); err == nil {
|
||||
t.Fatalf("second install while running should be rejected")
|
||||
}
|
||||
svc.mu.Lock()
|
||||
svc.running = false
|
||||
svc.mu.Unlock()
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user