fix: reduce emby load and stabilize strm scan

This commit is contained in:
ShukeBta
2026-06-11 00:18:51 +08:00
parent 859d6e9d24
commit 542e85a067
22 changed files with 1040 additions and 91 deletions
+1
View File
@@ -22,6 +22,7 @@ services:
image: ghcr.io/shukebta/mediastation-go:${MEDIASTATION_IMAGE_TAG:-latest}
container_name: mediastation-go
restart: unless-stopped
init: true
pull_policy: missing
ports:
+22 -1
View File
@@ -75,7 +75,28 @@ func AutoMigrate(db *gorm.DB) error {
if err := db.AutoMigrate(model.AllModels()...); err != nil {
return err
}
return enforceTelegramBindingOneToOne(db)
if err := enforceTelegramBindingOneToOne(db); err != nil {
return err
}
return ensurePerformanceIndexes(db)
}
func ensurePerformanceIndexes(db *gorm.DB) error {
statements := []string{
`CREATE INDEX IF NOT EXISTS idx_media_library_created_active ON media(library_id, created_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_media_library_episode_active ON media(library_id, season_num, episode_num, created_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_media_series_active ON media(series_id, season_num, episode_num) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_favorites_user_media_active ON favorites(user_id, media_id) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_playback_histories_user_media_active ON playback_histories(user_id, media_id, watched_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_playback_histories_resume_active ON playback_histories(user_id, completed, watched_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_play_profiles_user_created_active ON play_profiles(user_id, created_at DESC) WHERE deleted_at IS NULL`,
}
for _, stmt := range statements {
if err := db.Exec(stmt).Error; err != nil {
return err
}
}
return nil
}
func enforceTelegramBindingOneToOne(db *gorm.DB) error {
+28
View File
@@ -45,3 +45,31 @@ func TestEnforceTelegramBindingOneToOneCleansDuplicatesAndAddsIndex(t *testing.T
t.Fatal("expected unique index to reject another active binding for the same user")
}
}
func TestEnsurePerformanceIndexesCreatesHotPathIndexes(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.Media{}, &model.Favorite{}, &model.PlaybackHistory{}, &model.PlayProfile{}); err != nil {
t.Fatal(err)
}
if err := ensurePerformanceIndexes(db); err != nil {
t.Fatal(err)
}
for _, name := range []string{
"idx_media_library_created_active",
"idx_media_library_episode_active",
"idx_favorites_user_media_active",
"idx_playback_histories_user_media_active",
"idx_play_profiles_user_created_active",
} {
var count int
if err := db.Raw(`SELECT COUNT(1) FROM sqlite_master WHERE type = 'index' AND name = ?`, name).Scan(&count).Error; err != nil {
t.Fatal(err)
}
if count != 1 {
t.Fatalf("index %s count = %d, want 1", name, count)
}
}
}
+3
View File
@@ -269,6 +269,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 == "cloud.auto_sync_enabled" && !service.ParseBoolSetting(req.Value, false) && svc.Scan != nil {
_ = svc.Scan.CancelAllCloudScans()
}
c.Status(http.StatusNoContent)
}
}
+26 -2
View File
@@ -7,6 +7,7 @@ package handler
import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
@@ -596,6 +597,17 @@ func embyResumeItemsHandler(svc *service.Container) gin.HandlerFunc {
}
}
func embyItemsCountsHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"MovieCount": 0,
"SeriesCount": 0,
"EpisodeCount": 0,
"ItemCount": 0,
})
}
}
// ─── Images ──────────────────────────────────────────────────────────────────
// embyItemImageHandler 把 /Items/{id}/Images/Primary 等请求直接输出为图片。
@@ -603,14 +615,18 @@ func embyResumeItemsHandler(svc *service.Container) gin.HandlerFunc {
// /api/img 会变成 401,所以这里复用 ImageProxy 但不再走 /api 路由。
func embyItemImageHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
ctx, cancel := context.WithTimeout(c.Request.Context(), 8*time.Second)
defer cancel()
req := c.Request.WithContext(ctx)
id := c.Param("id")
imgType := strings.ToLower(c.Param("type"))
raw, err := svc.Emby.ImageURL(c.Request.Context(), id, imgType)
raw, err := svc.Emby.ImageURL(ctx, id, imgType)
if err != nil || raw == "" {
c.Status(http.StatusNotFound)
return
}
if typ, ref, ok := parseCloudPlayImageURL(raw); ok {
c.Request = req
serveCloudResolvedLink(svc, c, typ, ref)
return
}
@@ -618,7 +634,7 @@ func embyItemImageHandler(svc *service.Container) gin.HandlerFunc {
c.Status(http.StatusNotFound)
return
}
if err := svc.ImageProxy.Serve(c.Request.Context(), c.Writer, c.Request, raw); err != nil {
if err := svc.ImageProxy.Serve(ctx, c.Writer, req, raw); err != nil {
c.Status(http.StatusNotFound)
}
}
@@ -1059,6 +1075,8 @@ func registerEmbyRoutes(r *gin.Engine, jwtSecret string, svc *service.Container)
auth.GET("/Items", embyItemsHandler(svc))
auth.GET("/Users/:userId/Items", embyItemsHandler(svc))
auth.GET("/Items/Counts", embyItemsCountsHandler(svc))
auth.GET("/Users/:userId/Items/Counts", embyItemsCountsHandler(svc))
auth.GET("/Items/:id", embyItemByIDHandler(svc))
auth.GET("/Users/:userId/Items/:id", embyUserItemByIDHandler(svc))
auth.GET("/Shows/:id/Seasons", embyShowSeasonsHandler(svc))
@@ -1079,7 +1097,9 @@ func registerEmbyRoutes(r *gin.Engine, jwtSecret string, svc *service.Container)
auth.GET("/Videos/:id/stream.:container", embyVideoStreamHandler(svc))
auth.HEAD("/Videos/:id/stream.:container", embyVideoStreamHandler(svc))
auth.GET("/Videos/:id/original", embyVideoStreamHandler(svc))
auth.HEAD("/Videos/:id/original", embyVideoStreamHandler(svc))
auth.GET("/Videos/:id/original.:container", embyVideoStreamHandler(svc))
auth.HEAD("/Videos/:id/original.:container", embyVideoStreamHandler(svc))
auth.GET("/Videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.HEAD("/Videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.GET("/Videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
@@ -1121,6 +1141,8 @@ func registerLowercaseEmbyAuthRoutes(auth *gin.RouterGroup, svc *service.Contain
auth.GET("/items", embyItemsHandler(svc))
auth.GET("/users/:userId/items", embyItemsHandler(svc))
auth.GET("/items/counts", embyItemsCountsHandler(svc))
auth.GET("/users/:userId/items/counts", embyItemsCountsHandler(svc))
auth.GET("/items/:id", embyItemByIDHandler(svc))
auth.GET("/users/:userId/items/:id", embyUserItemByIDHandler(svc))
auth.GET("/shows/:id/seasons", embyShowSeasonsHandler(svc))
@@ -1141,7 +1163,9 @@ func registerLowercaseEmbyAuthRoutes(auth *gin.RouterGroup, svc *service.Contain
auth.GET("/videos/:id/stream.:container", embyVideoStreamHandler(svc))
auth.HEAD("/videos/:id/stream.:container", embyVideoStreamHandler(svc))
auth.GET("/videos/:id/original", embyVideoStreamHandler(svc))
auth.HEAD("/videos/:id/original", embyVideoStreamHandler(svc))
auth.GET("/videos/:id/original.:container", embyVideoStreamHandler(svc))
auth.HEAD("/videos/:id/original.:container", embyVideoStreamHandler(svc))
auth.GET("/videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.HEAD("/videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.GET("/videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
+103
View File
@@ -276,6 +276,50 @@ func TestEmbyVirtualFoldersRouteReturnsJSON(t *testing.T) {
}
}
func TestEmbyItemsCountsRouteReturnsJSON(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(&model.User{}); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
const secret = "test-secret"
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{Repo: repos})
for _, path := range []string{"/Items/Counts", "/Users/user-1/Items/Counts", "/items/counts"} {
req := httptest.NewRequest(http.MethodGet, path, nil)
req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("%s status=%d body=%s", path, w.Code, w.Body.String())
}
var body map[string]any
if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
t.Fatalf("%s decode response: %v", path, err)
}
if _, ok := body["MovieCount"]; !ok {
t.Fatalf("%s missing MovieCount: %#v", path, body)
}
}
}
func TestEmbyItemImageServesWithoutAPIAuth(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
@@ -531,6 +575,65 @@ func TestEmbyLowercaseVideoStreamRouteServesMedia(t *testing.T) {
}
}
func TestEmbyLowercaseOriginalHeadRouteServesHeaders(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
dir := t.TempDir()
mediaPath := filepath.Join(dir, "sample.mp4")
if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
t.Fatalf("write media: %v", err)
}
lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "media-1"},
LibraryID: lib.ID,
Title: "Lowercase Original",
Path: mediaPath,
Container: "mp4",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
const secret = "test-secret"
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
})
req := httptest.NewRequest(http.MethodHead, "/videos/media-1/original.mp4?api_key="+signedTestToken(t, secret), nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
if w.Body.Len() != 0 {
t.Fatalf("HEAD response should not include body, got %q", w.Body.String())
}
}
func TestEmbyLowercaseVideoHLSRouteDoesNot404WhenDirectOnly(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+327 -44
View File
@@ -22,6 +22,7 @@ import (
"sort"
"strconv"
"strings"
"sync"
"time"
"go.uber.org/zap"
@@ -51,6 +52,14 @@ type EmbyService struct {
cfg *config.Config
log *zap.Logger
repo *repository.Container
virtualMu sync.RWMutex
virtualSeries map[string]embySeriesCacheEntry
virtualSeasons map[string]embySeasonCacheEntry
virtualArtwork map[string]embyArtworkCacheEntry
visibilityMu sync.RWMutex
visibilityCache map[string]embyVisibilityCacheEntry
}
// NewEmbyService is the constructor.
@@ -187,10 +196,10 @@ func (e *EmbyService) Views(ctx context.Context, userID string) (map[string]any,
return nil, err
}
libs = FilterShadowedCloudLibraries(libs)
visibility := UserDefaultMediaVisibility(ctx, e.repo, userID)
visibility := e.mediaVisibility(ctx, userID)
items := make([]map[string]any, 0, len(libs))
for _, l := range libs {
if !LibraryVisibleForUser(ctx, e.repo, l, visibility) {
if !e.libraryVisibleFromCachedVisibility(l, visibility) {
continue
}
items = append(items, e.libraryAsView(&l))
@@ -248,6 +257,9 @@ type ItemsParams struct {
const (
embyVirtualSeriesPrefix = "msgo-series-"
embyVirtualSeasonPrefix = "msgo-season-"
embyVirtualCacheTTL = 10 * time.Minute
embyVisibilityCacheTTL = 30 * time.Second
embySeriesGroupingLimit = 5000
)
var (
@@ -281,6 +293,27 @@ type embySeasonGroup struct {
Episodes []model.Media
}
type embySeriesCacheEntry struct {
group embySeriesGroup
expiresAt time.Time
}
type embySeasonCacheEntry struct {
season embySeasonGroup
expiresAt time.Time
}
type embyArtworkCacheEntry struct {
primary string
backdrop string
expiresAt time.Time
}
type embyVisibilityCacheEntry struct {
visibility MediaVisibility
expiresAt time.Time
}
// 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.
@@ -513,19 +546,7 @@ func (e *EmbyService) LatestItems(ctx context.Context, userID, parentID string,
q = e.applyUserMediaVisibility(ctx, q, userID)
if parentID != "" {
if episodic, err := e.libraryIsEpisodic(ctx, parentID); err == nil && episodic {
resp, err := e.seriesItemsForLibrary(ctx, parentID, ItemsParams{
UserID: userID,
ParentID: parentID,
Limit: limit,
StartIndex: 0,
SortBy: "datecreated",
SortOrder: "Descending",
})
if err != nil {
return nil, err
}
items, _ := resp["Items"].([]map[string]any)
return items, nil
return e.latestSeriesItemsForLibrary(ctx, userID, parentID, limit)
}
q = q.Where("library_id = ?", parentID)
}
@@ -557,6 +578,36 @@ func (e *EmbyService) LatestItems(ctx context.Context, userID, parentID string,
return out, nil
}
func (e *EmbyService) latestSeriesItemsForLibrary(ctx context.Context, userID, libraryID string, limit int) ([]map[string]any, error) {
if limit <= 0 || limit > 100 {
limit = 20
}
rowLimit := limit * 40
if rowLimit < 200 {
rowLimit = 200
}
if rowLimit > embySeriesGroupingLimit {
rowLimit = embySeriesGroupingLimit
}
q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
Where("library_id = ? AND (season_num > 0 OR episode_num > 0)", libraryID)
q = e.applyUserMediaVisibility(ctx, q, userID)
var rows []model.Media
if err := q.Order("created_at desc").Limit(rowLimit).Find(&rows).Error; err != nil {
return nil, err
}
groups := e.seriesGroupsFromMedia(rows)
sortSeriesGroups(groups, ItemsParams{SortBy: "datecreated", SortOrder: "Descending"})
if len(groups) > limit {
groups = groups[:limit]
}
items := make([]map[string]any, 0, len(groups))
for _, group := range groups {
items = append(items, e.seriesPayload(group))
}
return items, nil
}
// ResumeItems 列出有未完成播放进度的媒体。
func (e *EmbyService) ResumeItems(ctx context.Context, userID string, limit int) (map[string]any, error) {
if limit <= 0 || limit > 100 {
@@ -690,8 +741,15 @@ func (e *EmbyService) seriesItemsForLibrary(ctx context.Context, libraryID strin
if p.SearchTerm != "" {
q = q.Where("title LIKE ? OR original_name LIKE ?", "%"+p.SearchTerm+"%", "%"+p.SearchTerm+"%")
}
rowLimit := p.StartIndex + maxInt(p.Limit*40, 1000)
if rowLimit < p.Limit {
rowLimit = p.Limit
}
if rowLimit > embySeriesGroupingLimit {
rowLimit = embySeriesGroupingLimit
}
var rows []model.Media
if err := q.Order("created_at desc").Find(&rows).Error; err != nil {
if err := q.Order("created_at desc").Limit(rowLimit).Find(&rows).Error; err != nil {
return nil, err
}
groups := e.seriesGroupsFromMedia(rows)
@@ -723,21 +781,142 @@ func (e *EmbyService) libraryIsEpisodic(ctx context.Context, libraryID string) (
return count > 0, err
}
func (e *EmbyService) rememberSeriesGroup(group embySeriesGroup) {
if e == nil || strings.TrimSpace(group.ID) == "" {
return
}
expiresAt := time.Now().Add(embyVirtualCacheTTL)
e.virtualMu.Lock()
defer e.virtualMu.Unlock()
if e.virtualSeries == nil {
e.virtualSeries = make(map[string]embySeriesCacheEntry)
}
if e.virtualSeasons == nil {
e.virtualSeasons = make(map[string]embySeasonCacheEntry)
}
if e.virtualArtwork == nil {
e.virtualArtwork = make(map[string]embyArtworkCacheEntry)
}
if len(e.virtualSeries) > 2000 || len(e.virtualSeasons) > 5000 || len(e.virtualArtwork) > 7000 {
e.virtualSeries = make(map[string]embySeriesCacheEntry)
e.virtualSeasons = make(map[string]embySeasonCacheEntry)
e.virtualArtwork = make(map[string]embyArtworkCacheEntry)
}
e.virtualSeries[group.ID] = embySeriesCacheEntry{group: group, expiresAt: expiresAt}
e.virtualArtwork[group.ID] = embyArtworkCacheEntry{primary: group.PosterURL, backdrop: group.BackdropURL, expiresAt: expiresAt}
e.virtualArtwork[group.ID+"-bd"] = embyArtworkCacheEntry{primary: group.PosterURL, backdrop: group.BackdropURL, expiresAt: expiresAt}
for _, season := range e.seasonsForSeries(group) {
e.virtualSeasons[season.ID] = embySeasonCacheEntry{season: season, expiresAt: expiresAt}
e.virtualArtwork[season.ID] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
e.virtualArtwork[season.ID+"-bd"] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
}
}
func (e *EmbyService) rememberSeasonGroup(season embySeasonGroup) {
if e == nil || strings.TrimSpace(season.ID) == "" {
return
}
expiresAt := time.Now().Add(embyVirtualCacheTTL)
e.virtualMu.Lock()
defer e.virtualMu.Unlock()
if e.virtualSeasons == nil {
e.virtualSeasons = make(map[string]embySeasonCacheEntry)
}
if e.virtualArtwork == nil {
e.virtualArtwork = make(map[string]embyArtworkCacheEntry)
}
e.virtualSeasons[season.ID] = embySeasonCacheEntry{season: season, expiresAt: expiresAt}
e.virtualArtwork[season.ID] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
e.virtualArtwork[season.ID+"-bd"] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
}
func (e *EmbyService) cachedSeriesGroup(id string) (embySeriesGroup, bool) {
if e == nil || strings.TrimSpace(id) == "" {
return embySeriesGroup{}, false
}
now := time.Now()
e.virtualMu.RLock()
entry, ok := e.virtualSeries[id]
e.virtualMu.RUnlock()
if !ok || now.After(entry.expiresAt) {
if ok {
e.virtualMu.Lock()
delete(e.virtualSeries, id)
e.virtualMu.Unlock()
}
return embySeriesGroup{}, false
}
return entry.group, true
}
func (e *EmbyService) cachedSeasonGroup(id string) (embySeasonGroup, bool) {
if e == nil || strings.TrimSpace(id) == "" {
return embySeasonGroup{}, false
}
now := time.Now()
e.virtualMu.RLock()
entry, ok := e.virtualSeasons[id]
e.virtualMu.RUnlock()
if !ok || now.After(entry.expiresAt) {
if ok {
e.virtualMu.Lock()
delete(e.virtualSeasons, id)
e.virtualMu.Unlock()
}
return embySeasonGroup{}, false
}
return entry.season, true
}
func (e *EmbyService) cachedArtworkURL(id, imageType string) (string, bool) {
if e == nil || strings.TrimSpace(id) == "" {
return "", false
}
now := time.Now()
e.virtualMu.RLock()
entry, ok := e.virtualArtwork[id]
e.virtualMu.RUnlock()
if !ok || now.After(entry.expiresAt) {
if ok {
e.virtualMu.Lock()
delete(e.virtualArtwork, id)
e.virtualMu.Unlock()
}
return "", false
}
switch strings.ToLower(imageType) {
case "backdrop", "art":
if entry.backdrop != "" {
return entry.backdrop, true
}
}
if entry.primary != "" {
return entry.primary, true
}
return entry.backdrop, entry.backdrop != ""
}
func (e *EmbyService) findSeriesGroup(ctx context.Context, id, userID string) (embySeriesGroup, bool, error) {
if strings.TrimSpace(id) == "" {
return embySeriesGroup{}, false, nil
}
if strings.HasPrefix(id, embyVirtualSeriesPrefix) {
if group, ok := e.cachedSeriesGroup(id); ok {
return group, true, nil
}
}
var rows []model.Media
q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("season_num > 0 OR episode_num > 0")
q = e.applyUserMediaVisibility(ctx, q, userID)
if !strings.HasPrefix(id, embyVirtualSeriesPrefix) {
q = q.Where("series_id = ?", id)
}
if err := q.Order("season_num asc, episode_num asc, created_at asc").Find(&rows).Error; err != nil {
if err := q.Order("season_num asc, episode_num asc, created_at asc").Limit(embySeriesGroupingLimit).Find(&rows).Error; err != nil {
return embySeriesGroup{}, false, err
}
for _, group := range e.seriesGroupsFromMedia(rows) {
if group.ID == id {
e.rememberSeriesGroup(group)
return group, true, nil
}
}
@@ -767,18 +946,23 @@ func (e *EmbyService) findSeasonGroup(ctx context.Context, id, userID string) (e
if strings.TrimSpace(id) == "" || !strings.HasPrefix(id, embyVirtualSeasonPrefix) {
return embySeasonGroup{}, false, nil
}
if season, ok := e.cachedSeasonGroup(id); ok {
return season, true, nil
}
var rows []model.Media
q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
Where("season_num > 0 OR episode_num > 0")
q = e.applyUserMediaVisibility(ctx, q, userID)
if err := q.
Order("season_num asc, episode_num asc, created_at asc").
Limit(embySeriesGroupingLimit).
Find(&rows).Error; err != nil {
return embySeasonGroup{}, false, err
}
for _, series := range e.seriesGroupsFromMedia(rows) {
for _, season := range e.seasonsForSeries(series) {
if season.ID == id {
e.rememberSeriesGroup(series)
return season, true, nil
}
}
@@ -875,6 +1059,7 @@ func (e *EmbyService) seasonsForSeries(series embySeriesGroup) []embySeasonGroup
}
func (e *EmbyService) seriesPayload(group embySeriesGroup) map[string]any {
e.rememberSeriesGroup(group)
imageTags := map[string]string{}
backdropTags := []string{}
if group.PosterURL != "" {
@@ -908,6 +1093,7 @@ func (e *EmbyService) seriesPayload(group embySeriesGroup) map[string]any {
}
func (e *EmbyService) seasonPayload(season embySeasonGroup) map[string]any {
e.rememberSeasonGroup(season)
imageTags := map[string]string{}
backdropTags := []string{}
if season.Series.PosterURL != "" {
@@ -949,18 +1135,16 @@ func (e *EmbyService) ImageURL(ctx context.Context, id, imageType string) (strin
return backdrop
}
if strings.HasPrefix(id, embyVirtualSeasonPrefix) {
if season, ok, err := e.findSeasonGroup(ctx, id, ""); err != nil {
return "", err
} else if ok {
return pick(season.Series.PosterURL, season.Series.BackdropURL), nil
if raw, ok := e.cachedArtworkURL(id, imageType); ok {
return raw, nil
}
return "", nil
}
if strings.HasPrefix(id, embyVirtualSeriesPrefix) {
if series, ok, err := e.findSeriesGroup(ctx, id, ""); err != nil {
return "", err
} else if ok {
return pick(series.PosterURL, series.BackdropURL), nil
if raw, ok := e.cachedArtworkURL(id, imageType); ok {
return raw, nil
}
return "", nil
}
m, err := e.repo.Media.FindByID(ctx, id)
if err == nil && m != nil {
@@ -1100,10 +1284,10 @@ func emptyUserData() map[string]any {
}
func (e *EmbyService) applyUserMediaVisibility(ctx context.Context, q *gorm.DB, userID string) *gorm.DB {
visibility := UserDefaultMediaVisibility(ctx, e.repo, userID)
visibility := e.mediaVisibility(ctx, userID)
if !visibility.IncludeNSFW {
q = q.Where("nsfw = ?", false)
if hidden := e.hiddenLibraryIDs(ctx, visibility); len(hidden) > 0 {
if hidden := visibility.HiddenLibraryIDs; len(hidden) > 0 {
q = q.Where("library_id NOT IN ?", hidden)
}
}
@@ -1114,7 +1298,7 @@ func (e *EmbyService) applyUserMediaVisibility(ctx context.Context, q *gorm.DB,
}
func (e *EmbyService) filterMediaRowsForUser(ctx context.Context, rows []model.Media, userID string) []model.Media {
visibility := UserDefaultMediaVisibility(ctx, e.repo, userID)
visibility := e.mediaVisibility(ctx, userID)
if visibility.IncludeNSFW && len(visibility.AllowedLibraryIDs) == 0 {
return rows
}
@@ -1123,7 +1307,7 @@ func (e *EmbyService) filterMediaRowsForUser(ctx context.Context, rows []model.M
allowed[id] = true
}
hiddenLibraries := map[string]bool{}
for _, id := range e.hiddenLibraryIDs(ctx, visibility) {
for _, id := range visibility.HiddenLibraryIDs {
hiddenLibraries[id] = true
}
out := rows[:0]
@@ -1142,6 +1326,75 @@ func (e *EmbyService) filterMediaRowsForUser(ctx context.Context, rows []model.M
return out
}
func (e *EmbyService) mediaVisibility(ctx context.Context, userID string) MediaVisibility {
if e == nil {
return MediaVisibility{IncludeNSFW: true}
}
key := strings.TrimSpace(userID)
now := time.Now()
e.visibilityMu.RLock()
entry, ok := e.visibilityCache[key]
e.visibilityMu.RUnlock()
if ok && now.Before(entry.expiresAt) {
return cloneMediaVisibility(entry.visibility)
}
visibility := UserDefaultMediaVisibility(ctx, e.repo, userID)
if !visibility.IncludeNSFW {
visibility.HiddenLibraryIDs = e.hiddenLibraryIDs(ctx, visibility)
}
visibility = cloneMediaVisibility(visibility)
e.visibilityMu.Lock()
if e.visibilityCache == nil {
e.visibilityCache = make(map[string]embyVisibilityCacheEntry)
}
if len(e.visibilityCache) > 1000 {
e.visibilityCache = make(map[string]embyVisibilityCacheEntry)
}
e.visibilityCache[key] = embyVisibilityCacheEntry{
visibility: cloneMediaVisibility(visibility),
expiresAt: now.Add(embyVisibilityCacheTTL),
}
e.visibilityMu.Unlock()
return visibility
}
func cloneMediaVisibility(visibility MediaVisibility) MediaVisibility {
if visibility.AllowedLibraryIDs != nil {
visibility.AllowedLibraryIDs = append([]string(nil), visibility.AllowedLibraryIDs...)
}
if visibility.HiddenLibraryIDs != nil {
visibility.HiddenLibraryIDs = append([]string(nil), visibility.HiddenLibraryIDs...)
}
return visibility
}
func (e *EmbyService) libraryVisibleFromCachedVisibility(lib model.Library, visibility MediaVisibility) bool {
if len(visibility.AllowedLibraryIDs) > 0 {
allowed := false
for _, id := range visibility.AllowedLibraryIDs {
if id == lib.ID {
allowed = true
break
}
}
if !allowed {
return false
}
}
if visibility.IncludeNSFW {
return true
}
for _, id := range visibility.HiddenLibraryIDs {
if id == lib.ID {
return false
}
}
return true
}
func (e *EmbyService) hiddenLibraryIDs(ctx context.Context, visibility MediaVisibility) []string {
if visibility.IncludeNSFW {
return nil
@@ -1228,24 +1481,34 @@ func (e *EmbyService) playableMedia(ctx context.Context, id, userID string) (*mo
// 直链给搜索接口)。/PlaybackInfo 走 false 路径,URL 指向 Emby 兼容
// /Videos/{id}/stream(客户端会继续携带 X-Emby-Token 或 append api_key)。
func (e *EmbyService) mediaSource(m *model.Media, asEmbedded, directOnly bool) map[string]any {
container := strings.Trim(strings.ToLower(m.Container), ". ")
if container == "" {
container = strings.TrimPrefix(strings.ToLower(filepath.Ext(m.Path)), ".")
}
if container == "" && strings.TrimSpace(m.STRMURL) != "" {
container = "strm"
}
src := map[string]any{
"Id": m.ID,
"Name": m.Title,
"Path": m.Path,
"Container": m.Container,
"Size": m.SizeBytes,
"Protocol": "Http",
"Type": "Default",
"IsRemote": false,
"SupportsTranscoding": !directOnly,
"SupportsDirectStream": true,
"SupportsDirectPlay": true,
"SupportsProbing": true,
"RunTimeTicks": int64(m.DurationSec) * 10_000_000,
"MediaStreams": e.mediaStreams(m),
"Id": m.ID,
"Name": m.Title,
"Path": m.Path,
"Container": container,
"Size": m.SizeBytes,
"Protocol": "Http",
"Type": "Default",
"IsRemote": false,
"RequiresOpening": false,
"RequiresClosing": false,
"ReadAtNativeFramerate": false,
"SupportsTranscoding": !directOnly,
"SupportsDirectStream": true,
"SupportsDirectPlay": true,
"SupportsProbing": true,
"RunTimeTicks": int64(m.DurationSec) * 10_000_000,
"MediaStreams": e.mediaStreams(m),
}
if !asEmbedded {
src["DirectStreamUrl"] = "/Videos/" + m.ID + "/stream"
src["DirectStreamUrl"] = embyDirectStreamURL(m.ID, container)
// 直连解码模式下不下发 TranscodingUrl,迫使客户端本地解码直连,
// 宿主机不参与转码。
if !directOnly {
@@ -1264,6 +1527,15 @@ func (e *EmbyService) mediaSource(m *model.Media, asEmbedded, directOnly bool) m
return src
}
func embyDirectStreamURL(mediaID, container string) string {
mediaID = strings.TrimSpace(mediaID)
container = strings.Trim(strings.ToLower(container), ". ")
if container == "" || container == "strm" {
return "/Videos/" + mediaID + "/stream"
}
return "/Videos/" + mediaID + "/stream." + container
}
func (e *EmbyService) mediaStreams(m *model.Media) []map[string]any {
streams := []map[string]any{}
if m.VideoCodec != "" || m.Width > 0 {
@@ -1290,6 +1562,17 @@ func (e *EmbyService) mediaStreams(m *model.Media) []map[string]any {
"IsExternal": false,
})
}
if len(streams) == 0 {
streams = append(streams, map[string]any{
"Codec": "unknown",
"Type": "Video",
"Index": 0,
"IsDefault": true,
"IsForced": false,
"IsExternal": false,
"DisplayTitle": "Video",
})
}
return streams
}
+51 -2
View File
@@ -1,6 +1,7 @@
package service
import (
"context"
"testing"
"github.com/glebarez/sqlite"
@@ -95,11 +96,55 @@ func TestEmbyItemsExposeSeriesSeasonEpisodeHierarchy(t *testing.T) {
if sources[0]["Id"] != "ep-1" {
t.Fatalf("series playback should fall back to first episode: %#v", sources)
}
if sources[0]["DirectStreamUrl"] != "/Videos/ep-1/stream" {
if sources[0]["DirectStreamUrl"] != "/Videos/ep-1/stream.mkv" {
t.Fatalf("playback should use Emby-compatible stream URL: %#v", sources[0])
}
}
func TestEmbyVirtualSeriesArtworkUsesListCache(t *testing.T) {
svc := newTestEmbyService(t)
lib := model.Library{Name: "番剧", Path: `/media/anime`, Type: "anime", Enabled: true}
if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
media := model.Media{
Base: model.Base{ID: "ep-1"},
LibraryID: lib.ID,
Title: "剑来",
Path: `/media/anime/剑来/Season 01/剑来 - S01E01.mkv`,
PosterURL: `/poster.jpg`,
BackdropURL: `/backdrop.jpg`,
SeasonNum: 1,
EpisodeNum: 1,
}
if err := svc.repo.DB.Create(&media).Error; err != nil {
t.Fatalf("create media: %v", err)
}
root, err := svc.Items(t.Context(), ItemsParams{ParentID: lib.ID, Limit: 50})
if err != nil {
t.Fatalf("library items: %v", err)
}
items := root["Items"].([]map[string]any)
seriesID := items[0]["Id"].(string)
cancelled, cancel := context.WithCancel(t.Context())
cancel()
poster, err := svc.ImageURL(cancelled, seriesID, "Primary")
if err != nil {
t.Fatalf("image url from cache: %v", err)
}
if poster != "/poster.jpg" {
t.Fatalf("poster = %q, want cached poster", poster)
}
backdrop, err := svc.ImageURL(cancelled, seriesID, "Backdrop")
if err != nil {
t.Fatalf("backdrop url from cache: %v", err)
}
if backdrop != "/backdrop.jpg" {
t.Fatalf("backdrop = %q, want cached backdrop", backdrop)
}
}
func TestEmbyCloudAnimeUsesSeriesNameFromChineseSeasonFolder(t *testing.T) {
svc := newTestEmbyService(t)
lib := model.Library{Name: "OpenList · 国漫", Path: `cloud://openlist/国漫`, Type: "anime", Enabled: true}
@@ -283,7 +328,7 @@ func TestEmbyPlaybackInfoRespectsDirectPlayOnly(t *testing.T) {
if _, ok := src["TranscodingUrl"]; ok {
t.Fatalf("expected no TranscodingUrl in direct-only mode: %#v", src)
}
if src["SupportsDirectPlay"] != true || src["DirectStreamUrl"] != "/Videos/m-1/stream" {
if src["SupportsDirectPlay"] != true || src["DirectStreamUrl"] != "/Videos/m-1/stream.mkv" {
t.Fatalf("direct-only must still allow direct play: %#v", src)
}
}
@@ -319,6 +364,10 @@ func TestEmbyPlaybackInfoKeepsSTRMBehindStreamEndpoint(t *testing.T) {
if src["Path"] != "/api/cloud/play/quark?ref=f1" {
t.Fatalf("path should expose the strm target for diagnostics: %#v", src)
}
streams := src["MediaStreams"].([]map[string]any)
if len(streams) == 0 || streams[0]["Type"] != "Video" {
t.Fatalf("strm media should expose a fallback video stream for Android clients: %#v", src)
}
}
func newTestEmbyService(t *testing.T) *EmbyService {
+84 -11
View File
@@ -58,6 +58,7 @@ type ScannerService struct {
cloudScanMu sync.Mutex
cloudScans map[string]*cloudScanEntry
cloudSlots chan struct{}
}
// NewScannerService is the constructor.
@@ -73,6 +74,7 @@ func NewScannerService(
cfg: cfg, log: log, repo: repo, hub: hub,
probe: probe, scraper: scraper,
cloudScans: make(map[string]*cloudScanEntry),
cloudSlots: make(chan struct{}, 1),
}
}
@@ -232,6 +234,35 @@ func (s *ScannerService) updateCloudScanProgress(libraryID, stage string, dirs,
entry.status.FilesPerSecond = filesPerSecond
}
func (s *ScannerService) acquireCloudScanSlot(ctx context.Context, libraryID string) (func(), error) {
if s == nil {
return func() {}, nil
}
s.cloudScanMu.Lock()
if s.cloudSlots == nil {
s.cloudSlots = make(chan struct{}, 1)
}
slots := s.cloudSlots
if entry := s.cloudScans[libraryID]; entry != nil {
entry.status.Stage = "queued"
entry.status.UpdatedAt = time.Now()
}
s.cloudScanMu.Unlock()
select {
case slots <- struct{}{}:
s.cloudScanMu.Lock()
if entry := s.cloudScans[libraryID]; entry != nil && entry.status.State == "running" {
entry.status.Stage = "listing"
entry.status.UpdatedAt = time.Now()
}
s.cloudScanMu.Unlock()
return func() { <-slots }, nil
case <-ctx.Done():
return nil, ctx.Err()
}
}
// CloudScanStatuses returns the current or most recent status per cloud library.
func (s *ScannerService) CloudScanStatuses() []CloudScanStatus {
if s == nil {
@@ -421,6 +452,15 @@ func (s *ScannerService) scanLibrary(ctx context.Context, libraryID string, auto
}
return nil, err
}
release, err := s.acquireCloudScanSlot(scanCtx, lib.ID)
if err != nil {
res := &ScanResult{LibraryID: lib.ID}
if finish != nil {
finish(res, err)
}
return res, err
}
defer release()
res, err := s.scanCloudLibrary(scanCtx, lib, mount, autoScrape)
if finish != nil {
finish(res, err)
@@ -639,6 +679,11 @@ func (s *ScannerService) scanCloudLibrary(ctx context.Context, lib *model.Librar
if err := walkCloud(rootDir, rootDisplayDir, nil); err != nil {
return res, err
}
existingPaths, err := s.existingCloudMediaPaths(ctx, lib.ID)
if err != nil {
s.log.Warn("load existing cloud media paths failed", zap.String("library_id", lib.ID), zap.Error(err))
existingPaths = nil
}
for _, candidate := range candidates {
select {
case <-ctx.Done():
@@ -646,7 +691,7 @@ func (s *ScannerService) scanCloudLibrary(ctx context.Context, lib *model.Librar
default:
}
seen[candidate.path] = struct{}{}
s.ingestCloudFile(ctx, lib, typ, candidate.ref, candidate.path, candidate.name, candidate.size, candidate.localMeta, res)
s.ingestCloudFile(ctx, lib, typ, candidate.ref, candidate.path, candidate.name, candidate.size, candidate.localMeta, existingPaths, res)
publishProgress("importing", res.Visited == 1 || res.Visited%100 == 0)
}
removed, err := s.pruneMissingCloudMedia(ctx, lib.ID, seen)
@@ -678,6 +723,26 @@ func (s *ScannerService) scanCloudLibrary(ctx context.Context, lib *model.Librar
return res, nil
}
func (s *ScannerService) existingCloudMediaPaths(ctx context.Context, libraryID string) (map[string]struct{}, error) {
var rows []struct {
Path string
}
if err := s.repo.DB.WithContext(ctx).
Model(&model.Media{}).
Select("path").
Where("library_id = ? AND path LIKE ?", libraryID, "cloud://%").
Find(&rows).Error; err != nil {
return nil, err
}
out := make(map[string]struct{}, len(rows))
for _, row := range rows {
if row.Path != "" {
out[row.Path] = struct{}{}
}
}
return out, nil
}
func (s *ScannerService) shadowedCloudLibrary(ctx context.Context, lib *model.Library) *CloudMountConflict {
libs, err := s.repo.Library.List(ctx)
if err != nil {
@@ -687,7 +752,7 @@ func (s *ScannerService) shadowedCloudLibrary(ctx context.Context, lib *model.Li
return CloudLibraryShadowed(libs, *lib)
}
func (s *ScannerService) ingestCloudFile(ctx context.Context, lib *model.Library, typ, ref, path, name string, size int64, localMeta *LocalMetadata, res *ScanResult) {
func (s *ScannerService) ingestCloudFile(ctx context.Context, lib *model.Library, typ, ref, path, name string, size int64, localMeta *LocalMetadata, existingPaths map[string]struct{}, res *ScanResult) {
res.Visited++
ext := strings.ToLower(filepath.Ext(name))
title, year := CleanQuery(name)
@@ -706,7 +771,13 @@ func (s *ScannerService) ingestCloudFile(ctx context.Context, lib *model.Library
}
}
}
isNewMedia := !s.mediaPathExists(ctx, path)
isNewMedia := false
if existingPaths != nil {
_, exists := existingPaths[path]
isNewMedia = !exists
} else {
isNewMedia = !s.mediaPathExists(ctx, path)
}
m := &model.Media{
LibraryID: lib.ID,
Title: title,
@@ -739,14 +810,16 @@ func (s *ScannerService) ingestCloudFile(ctx context.Context, lib *model.Library
} else {
res.Updated++
}
s.hub.Publish("scan", map[string]any{
"library_id": lib.ID,
"path": path,
"visited": res.Visited,
"added": res.Added,
"updated": res.Updated,
"cloud": true,
})
if s.hub != nil && (res.Visited == 1 || res.Visited%100 == 0) {
s.hub.Publish("scan", map[string]any{
"library_id": lib.ID,
"path": path,
"visited": res.Visited,
"added": res.Added,
"updated": res.Updated,
"cloud": true,
})
}
}
func cloudSeriesTitleFromMediaPath(mediaPath string) (string, int) {
+8
View File
@@ -92,6 +92,14 @@ func TestScanCloudLibraryImportsRecursivePlayableMedia(t *testing.T) {
t.Fatalf("root media path/strm wrong: path=%q strm=%q", rows[0].Path, rows[0].STRMURL)
}
res, err = scanner.ScanLibrary(t.Context(), lib.ID)
if err != nil {
t.Fatalf("rescan same cloud: %v", err)
}
if res.Added != 0 || res.Updated != 2 {
t.Fatalf("same cloud rescan should update existing rows only, got %#v", res)
}
empty = true
res, err = scanner.ScanLibrary(t.Context(), lib.ID)
if err != nil {
+26 -10
View File
@@ -174,24 +174,40 @@ func (s *SchedulerService) RunNow(ctx context.Context, name string) error {
}
func (s *SchedulerService) loop(ctx context.Context, j *scheduledJob) {
t := time.NewTicker(j.interval)
defer t.Stop()
// Run once shortly after startup so the initial state is fresh.
first := time.NewTimer(15 * time.Second)
defer first.Stop()
s.loopWithInitialDelay(ctx, j, 15*time.Second)
}
func (s *SchedulerService) loopWithInitialDelay(ctx context.Context, j *scheduledJob, initialDelay time.Duration) {
delay := initialDelay
for {
if delay < 0 {
delay = 0
}
timer := time.NewTimer(delay)
select {
case <-ctx.Done():
if !timer.Stop() {
select {
case <-timer.C:
default:
}
}
return
case <-s.stopCh:
if !timer.Stop() {
select {
case <-timer.C:
default:
}
}
return
case <-first.C:
case <-t.C:
case <-timer.C:
}
if err := s.runOnce(ctx, j); err != nil {
s.log.Warn("scheduled job failed",
zap.String("name", j.name), zap.Error(err))
}
delay = j.interval
}
}
@@ -343,13 +359,13 @@ func (s *SchedulerService) jobSyncCloudLibraries(ctx context.Context) error {
func (s *SchedulerService) autoCloudSyncEnabled(ctx context.Context) bool {
if s.repo == nil || s.repo.Setting == nil {
return true
return false
}
v, err := s.repo.Setting.Get(ctx, "cloud.auto_sync_enabled")
if err != nil {
return true
return false
}
return parseBoolSetting(v, true)
return parseBoolSetting(v, false)
}
func (s *SchedulerService) cloudSyncInterval(ctx context.Context) time.Duration {
+86
View File
@@ -1,10 +1,12 @@
package service
import (
"context"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"sync/atomic"
"testing"
"time"
@@ -146,6 +148,9 @@ func TestSchedulerCloudSyncImportsMountedCloudLibrary(t *testing.T) {
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), "cloud.auto_sync_enabled", "true"); err != nil {
t.Fatal(err)
}
scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
scanner.SetStorageConfig(storage)
scheduler := NewSchedulerService(log, repos, scanner, nil, nil, storage, NewHub(log), "")
@@ -161,3 +166,84 @@ func TestSchedulerCloudSyncImportsMountedCloudLibrary(t *testing.T) {
t.Fatalf("strm url = %q", media.STRMURL)
}
}
func TestSchedulerCloudSyncDisabledByDefault(t *testing.T) {
var requests atomic.Int32
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
requests.Add(1)
_, _ = w.Write([]byte(`{"status":200,"code":0,"data":{"list":[
{"fid":"f1","file_name":"Cloud.Movie.2026.mkv","dir":false,"size":1024}
]}}`))
}))
defer upstream.Close()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}); err != nil {
t.Fatal(err)
}
repos := repository.New(db)
log := zap.NewNop()
storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
if _, err := storage.Save(t.Context(), StorageInput{
Type: "quark",
Config: map[string]any{
"cookie": "kps=test",
"base": upstream.URL,
},
}); err != nil {
t.Fatal(err)
}
lib := model.Library{Name: "夸克网盘", Path: "cloud://quark/0", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
scanner.SetStorageConfig(storage)
scheduler := NewSchedulerService(log, repos, scanner, nil, nil, storage, NewHub(log), "")
if err := scheduler.jobSyncCloudLibraries(t.Context()); err != nil {
t.Fatalf("disabled cloud sync should be a no-op: %v", err)
}
if got := requests.Load(); got != 0 {
t.Fatalf("cloud sync made %d upstream requests while disabled by default", got)
}
if got := countMedia(t, repos); got != 0 {
t.Fatalf("media count = %d, want 0 while cloud sync disabled by default", got)
}
}
func TestSchedulerLoopWaitsIntervalAfterSlowRun(t *testing.T) {
scheduler := NewSchedulerService(zap.NewNop(), nil, nil, nil, nil, nil, nil, "")
ctx, cancel := context.WithCancel(t.Context())
defer cancel()
var runs atomic.Int32
job := &scheduledJob{
name: "slow",
interval: 25 * time.Millisecond,
run: func(ctx context.Context) error {
runs.Add(1)
time.Sleep(50 * time.Millisecond)
return nil
},
}
done := make(chan struct{})
go func() {
scheduler.loopWithInitialDelay(ctx, job, time.Millisecond)
close(done)
}()
time.Sleep(120 * time.Millisecond)
cancel()
select {
case <-done:
case <-time.After(250 * time.Millisecond):
t.Fatal("scheduler loop did not stop")
}
if got := runs.Load(); got > 2 {
t.Fatalf("slow job ran %d times; scheduler should not catch up missed ticks", got)
}
}
+79 -3
View File
@@ -33,6 +33,8 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
const STRMEnabledSettingKey = "strm.enabled"
// StreamService serves media files with proper Range support so browsers can
// seek into the stream.
type StreamService struct {
@@ -77,6 +79,10 @@ func (s *StreamService) directPlayOnly(ctx context.Context) bool {
// own relative API endpoints — never to an absolute external direct link —
// so the JWT is never leaked off-site (e.g. to the cloud CDN).
func withAuthToken(target string, r *http.Request) string {
return withAuthTokenForInternalRedirect(target, r, "")
}
func withAuthTokenForInternalRedirect(target string, r *http.Request, publicBase string) string {
if r == nil {
return target
}
@@ -84,7 +90,13 @@ func withAuthToken(target string, r *http.Request) string {
return target
}
u, err := url.Parse(target)
if err != nil || u.IsAbs() {
if err != nil {
return target
}
if u.IsAbs() && !isInternalAPIURL(u, r, publicBase) {
return target
}
if !strings.HasPrefix(strings.ToLower(u.Path), "/api/") {
return target
}
tok := requestToken(r)
@@ -99,6 +111,58 @@ func withAuthToken(target string, r *http.Request) string {
return u.String()
}
func absoluteInternalRedirect(target string, r *http.Request) string {
if r == nil || target == "" || strings.HasPrefix(target, "//") {
return target
}
u, err := url.Parse(target)
if err != nil || u.IsAbs() || !strings.HasPrefix(target, "/") {
return target
}
scheme := strings.TrimSpace(r.Header.Get("X-Forwarded-Proto"))
if scheme == "" {
if r.TLS != nil {
scheme = "https"
} else {
scheme = "http"
}
}
host := strings.TrimSpace(r.Header.Get("X-Forwarded-Host"))
if host == "" {
host = r.Host
}
if host == "" {
return target
}
u.Scheme = scheme
u.Host = host
return u.String()
}
func isInternalAPIURL(u *url.URL, r *http.Request, publicBase string) bool {
if u == nil || !strings.HasPrefix(strings.ToLower(u.Path), "/api/") {
return false
}
targetHost := strings.ToLower(strings.TrimSpace(u.Host))
if targetHost == "" {
return true
}
if r != nil {
if host := strings.ToLower(strings.TrimSpace(r.Host)); host != "" && targetHost == host {
return true
}
if host := strings.ToLower(strings.TrimSpace(r.Header.Get("X-Forwarded-Host"))); host != "" && targetHost == host {
return true
}
}
if publicBase != "" {
if base, err := url.Parse(publicBase); err == nil && strings.EqualFold(strings.TrimSpace(base.Host), targetHost) {
return true
}
}
return false
}
// requestToken extracts the bearer JWT from the incoming request the same way
// the auth middleware does (Authorization header, Emby token headers, or the
// token / api_key query params used by <video>.src).
@@ -119,6 +183,17 @@ func requestToken(r *http.Request) string {
return ""
}
func STRMPlaybackEnabled(ctx context.Context, repo *repository.Container) bool {
if repo == nil || repo.Setting == nil {
return true
}
v, err := repo.Setting.Get(ctx, STRMEnabledSettingKey)
if err != nil || strings.TrimSpace(v) == "" {
return true
}
return parseBoolSetting(v, true)
}
// ServeFile streams the file backing the given media ID using
// http.ServeContent so HEAD / Range / If-Modified-Since are handled for free.
//
@@ -133,8 +208,9 @@ func (s *StreamService) ServeFile(w http.ResponseWriter, r *http.Request, mediaI
if m == nil {
return ErrMediaNotFound
}
if strings.TrimSpace(m.STRMURL) != "" {
http.Redirect(w, r, withAuthToken(m.STRMURL, r), http.StatusFound)
if strings.TrimSpace(m.STRMURL) != "" && STRMPlaybackEnabled(r.Context(), s.repo) {
target := withAuthTokenForInternalRedirect(m.STRMURL, r, PublicServerURL(r.Context(), s.repo, s.cfg))
http.Redirect(w, r, absoluteInternalRedirect(target, r), http.StatusFound)
return nil
}
f, err := os.Open(m.Path)
+87
View File
@@ -2,9 +2,18 @@ package service
import (
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestWithAuthTokenPropagatesToInternalRedirect(t *testing.T) {
@@ -36,6 +45,84 @@ func TestWithAuthTokenNeverLeaksToAbsoluteURL(t *testing.T) {
}
}
func TestWithAuthTokenPropagatesToSameOriginAbsoluteInternalURL(t *testing.T) {
r := httptest.NewRequest(http.MethodGet, "http://media.example/Videos/m-1/stream?api_key=jwt123", nil)
got := withAuthTokenForInternalRedirect("http://media.example/api/cloud/play/openlist?ref=abc", r, "http://media.example")
u, err := url.Parse(got)
if err != nil {
t.Fatalf("parse: %v", err)
}
if u.Query().Get("token") != "jwt123" || u.Query().Get("ref") != "abc" {
t.Fatalf("same-origin internal URL should keep ref and receive token: %q", got)
}
}
func TestServeFileRedirectsInternalSTRMAsAbsoluteURLWithToken(t *testing.T) {
repos := newStreamTestRepo(t)
if err := repos.DB.Create(&model.Media{
Base: model.Base{ID: "cloud-1"},
Title: "Cloud",
Path: "cloud://openlist/Movie.mkv",
STRMURL: "/api/cloud/play/openlist?ref=movie",
}).Error; err != nil {
t.Fatal(err)
}
svc := NewStreamService(&config.Config{}, zap.NewNop(), repos, nil)
req := httptest.NewRequest(http.MethodGet, "http://nas.local:18080/api/stream/cloud-1?api_key=jwt123", nil)
w := httptest.NewRecorder()
if err := svc.ServeFile(w, req, "cloud-1"); err != nil {
t.Fatal(err)
}
if w.Code != http.StatusFound {
t.Fatalf("status = %d, want 302", w.Code)
}
loc := w.Header().Get("Location")
if !strings.HasPrefix(loc, "http://nas.local:18080/api/cloud/play/openlist?") ||
!strings.Contains(loc, "ref=movie") ||
!strings.Contains(loc, "token=jwt123") {
t.Fatalf("redirect Location should be absolute and tokenized, got %q", loc)
}
}
func TestServeFileHonorsSTRMPlaybackDisabled(t *testing.T) {
repos := newStreamTestRepo(t)
if err := repos.Setting.Set(t.Context(), STRMEnabledSettingKey, "false"); err != nil {
t.Fatal(err)
}
if err := repos.DB.Create(&model.Media{
Base: model.Base{ID: "cloud-1"},
Title: "Cloud",
Path: "cloud://openlist/Movie.mkv",
STRMURL: "/api/cloud/play/openlist?ref=movie",
}).Error; err != nil {
t.Fatal(err)
}
svc := NewStreamService(&config.Config{}, zap.NewNop(), repos, nil)
req := httptest.NewRequest(http.MethodGet, "http://nas.local:18080/api/stream/cloud-1?api_key=jwt123", nil)
w := httptest.NewRecorder()
err := svc.ServeFile(w, req, "cloud-1")
if err != ErrMediaNotFound {
t.Fatalf("disabled STRM should not redirect cloud media, err=%v status=%d location=%q", err, w.Code, w.Header().Get("Location"))
}
if loc := w.Header().Get("Location"); loc != "" {
t.Fatalf("disabled STRM leaked redirect Location %q", loc)
}
}
func newStreamTestRepo(t *testing.T) *repository.Container {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.Media{}, &model.Setting{}); err != nil {
t.Fatal(err)
}
return repository.New(db)
}
func TestRequestTokenFromBearerHeader(t *testing.T) {
h := http.Header{}
h.Set("Authorization", "Bearer hdrtok")
+3
View File
@@ -10,6 +10,9 @@ export const api = axios.create({
timeout: 30000,
})
export const LONG_REQUEST_TIMEOUT = 120_000
export const BATCH_REQUEST_TIMEOUT = 300_000
// Flag to prevent multiple simultaneous refresh attempts
let isRefreshing = false
let refreshSubscribers: Array<{
+4 -3
View File
@@ -1,4 +1,4 @@
import { api } from './client'
import { api, BATCH_REQUEST_TIMEOUT, LONG_REQUEST_TIMEOUT } from './client'
import type { Library, Media, ScanResult } from '../types'
export interface MediaPage {
@@ -22,15 +22,16 @@ export const libraryAPI = {
remove: (id: string) => api.delete(`/libraries/${id}`).then((r) => r.data),
scan: (id: string) =>
api.post<ScanResult>(`/libraries/${id}/scan`).then((r) => r.data),
api.post<ScanResult>(`/libraries/${id}/scan`, null, { timeout: BATCH_REQUEST_TIMEOUT }).then((r) => r.data),
scrape: (id: string) =>
api.post(`/libraries/${id}/scrape`).then((r) => r.data),
api.post(`/libraries/${id}/scrape`, null, { timeout: BATCH_REQUEST_TIMEOUT }).then((r) => r.data),
listMedia: (id: string, page = 1, pageSize = 50) =>
api
.get<MediaPage>(`/libraries/${id}/media`, {
params: { page, page_size: pageSize },
timeout: LONG_REQUEST_TIMEOUT,
})
.then((r) => r.data),
}
+10 -3
View File
@@ -1,4 +1,4 @@
import { api } from './client'
import { api, BATCH_REQUEST_TIMEOUT, LONG_REQUEST_TIMEOUT } from './client'
export type StorageType = 'alist' | 'openlist' | 's3' | 'webdav' | 'cloud115' | 'quark' | 'clouddrive2'
@@ -85,6 +85,8 @@ export const storageAPI = {
.post<{ ok: boolean; error?: string }>(`/admin/storage/${type}/test`, {
type,
config,
}, {
timeout: LONG_REQUEST_TIMEOUT,
})
.then((r) => r.data),
@@ -99,7 +101,9 @@ export const storageAPI = {
},
) =>
api
.post<{ result: CloudUploadResult; error?: string }>(`/admin/storage/${type}/upload-local`, input)
.post<{ result: CloudUploadResult; error?: string }>(`/admin/storage/${type}/upload-local`, input, {
timeout: BATCH_REQUEST_TIMEOUT,
})
.then((r) => r.data),
scanAllCloud: () =>
@@ -126,6 +130,7 @@ export const cloudAPI = {
api
.get<{ items: CloudEntry[]; error?: string }>(`/admin/cloud/${type}/list`, {
params: { dir },
timeout: LONG_REQUEST_TIMEOUT,
})
.then((r) => r.data),
@@ -136,7 +141,9 @@ export const cloudAPI = {
mount: (type: StorageType, dir = '', name = '', media_type = 'movie', dir_path = '') =>
api
.post(`/admin/cloud/${type}/mount`, { dir, dir_path, name, media_type })
.post(`/admin/cloud/${type}/mount`, { dir, dir_path, name, media_type }, {
timeout: LONG_REQUEST_TIMEOUT,
})
.then((r) => r.data),
qrStart: (type: StorageType) =>
+4 -2
View File
@@ -1,4 +1,4 @@
import { api } from './client'
import { api, BATCH_REQUEST_TIMEOUT } from './client'
export type GenerateSTRMInput = {
library_id: string
@@ -33,5 +33,7 @@ export const strmAPI = {
importURL: (libraryID: string, title: string, url: string) =>
api.post('/strm/import', { library_id: libraryID, title, url }).then((r) => r.data),
generate: (input: GenerateSTRMInput) =>
api.post<GenerateSTRMResult>('/strm/generate', input).then((r) => r.data),
api
.post<GenerateSTRMResult>('/strm/generate', input, { timeout: BATCH_REQUEST_TIMEOUT })
.then((r) => r.data),
}
+14 -3
View File
@@ -2,9 +2,12 @@ import { useEffect, useRef } from 'react'
import { useAuthStore } from '../stores/auth'
const MAX_RECONNECT_ATTEMPTS = 5
// useWebSocket opens a single connection to /api/ws and dispatches every
// message to the supplied handler. Auto-reconnects with a 3 s back-off
// while the auth token is present.
// message to the supplied handler. Auto-reconnects with back-off while the
// auth token is present, but stops after repeated failures so an expired token
// cannot create an endless /api/ws 401 loop.
export function useWebSocket(onEvent: (topic: string, payload: unknown) => void) {
const ref = useRef<WebSocket | null>(null)
const token = useAuthStore((s) => s.token)
@@ -13,12 +16,18 @@ export function useWebSocket(onEvent: (topic: string, payload: unknown) => void)
if (!token) return
let closed = false
let timer: number | undefined
let reconnectAttempts = 0
const open = () => {
if (closed) return
if (reconnectAttempts >= MAX_RECONNECT_ATTEMPTS) return
const proto = window.location.protocol === 'https:' ? 'wss:' : 'ws:'
const url = `${proto}//${window.location.host}/api/ws?token=${encodeURIComponent(token)}`
const ws = new WebSocket(url)
ref.current = ws
ws.onopen = () => {
reconnectAttempts = 0
}
ws.onmessage = (ev) => {
try {
const msg = JSON.parse(ev.data)
@@ -31,7 +40,9 @@ export function useWebSocket(onEvent: (topic: string, payload: unknown) => void)
}
ws.onclose = () => {
if (closed) return
timer = window.setTimeout(open, 3_000)
reconnectAttempts += 1
const delay = Math.min(3_000 * reconnectAttempts, 30_000)
timer = window.setTimeout(open, delay)
}
}
+11 -1
View File
@@ -13,6 +13,8 @@ import { useAuthStore } from '../stores/auth'
import { getSeriesKey, groupSeries, isEpisodeLike, seriesTitle, type SeriesCard } from '../utils/groupSeries'
import { useWebSocket } from '../hooks/useWebSocket'
const MAX_LIBRARY_ITEMS_IN_BROWSER = 3_000
export function LibraryPage() {
const { id = '' } = useParams()
const [searchParams, setSearchParams] = useSearchParams()
@@ -77,9 +79,10 @@ export function LibraryPage() {
setLoading(true)
setItems([])
const loadAll = async () => {
const pageSize = 2000
const pageSize = 500
let page = 1
let collected: Media[] = []
let warnedLargeLibrary = false
try {
for (;;) {
const d = await libraryAPI.listMedia(id, page, pageSize)
@@ -88,6 +91,13 @@ export function LibraryPage() {
setItems(collected)
setTotal(d.total)
if (collected.length >= d.total || d.items.length < pageSize) break
if (collected.length >= MAX_LIBRARY_ITEMS_IN_BROWSER) {
if (!warnedLargeLibrary) {
warnedLargeLibrary = true
toast(`媒体库条目较多,已先加载前 ${MAX_LIBRARY_ITEMS_IN_BROWSER} 条,避免浏览器卡死。请使用搜索或更细的媒体库目录浏览。`)
}
break
}
page += 1
}
} finally {
+2 -2
View File
@@ -304,8 +304,8 @@ const GROUPS: SettingGroup[] = [
key: 'cloud.auto_sync_enabled',
label: '自动同步网盘媒体库',
type: 'toggle',
hint: '开启后后台会按间隔刷新已挂载的 cloud:// 媒体库,自动生成或更新 302/STRM 播放入口;不会下载网盘文件到本地。',
defaultValue: 'true',
hint: '默认关闭,避免 NAS 反复递归读取大型网盘目录。需要定时刷新时再开启;手动扫描仍可在外部存储页面执行。',
defaultValue: 'false',
},
{
key: 'cloud.sync_interval_seconds',
+61 -4
View File
@@ -1,5 +1,5 @@
import { FormEvent, useEffect, useState } from 'react'
import { Link as LinkIcon, Loader2, Plus, Search, Trash2, Wand2 } from 'lucide-react'
import { Link as LinkIcon, Loader2, Plus, Save, Search, Trash2, Wand2 } from 'lucide-react'
import toast from 'react-hot-toast'
import { adminAPI } from '../api/admin'
@@ -21,6 +21,9 @@ export function StrmPage() {
const [generateLibraryID, setGenerateLibraryID] = useState('')
const [baseURL, setBaseURL] = useState('')
const [outputDir, setOutputDir] = useState('')
const [strmEnabled, setStrmEnabled] = useState(true)
const [autoGenerate, setAutoGenerate] = useState(false)
const [savingSettings, setSavingSettings] = useState(false)
const [overwrite, setOverwrite] = useState(false)
const [generating, setGenerating] = useState(false)
const [generateResult, setGenerateResult] = useState<GenerateSTRMResult | null>(null)
@@ -45,6 +48,8 @@ export function StrmPage() {
const settings = Object.fromEntries(rows.map((row) => [row.key, row.value]))
setBaseURL(settings['app.server_url'] || settings['strm.base_url'] || '')
setOutputDir(settings['strm.output_dir'] || '')
setStrmEnabled(settings['strm.enabled'] !== 'false')
setAutoGenerate(settings['strm.auto_generate_enabled'] === 'true')
})
.catch(() => undefined)
}, [])
@@ -69,7 +74,7 @@ export function StrmPage() {
base_url: baseURL.trim().replace(/\/+$/, ''),
output_dir: outputDir.trim(),
overwrite,
enabled: true,
enabled: autoGenerate,
include_local: true,
})
setGenerateResult(result)
@@ -85,6 +90,24 @@ export function StrmPage() {
}
}
const saveSTRMSettings = async () => {
setSavingSettings(true)
try {
await Promise.all([
adminAPI.updateSetting('strm.enabled', String(strmEnabled)),
adminAPI.updateSetting('strm.auto_generate_enabled', String(autoGenerate)),
])
toast.success(strmEnabled ? 'STRM 播放已启用' : 'STRM 播放已关闭')
} catch (err: unknown) {
const msg =
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
'保存 STRM 开关失败'
toast.error(msg)
} finally {
setSavingSettings(false)
}
}
const onImport = async (e: FormEvent) => {
e.preventDefault()
if (!libraryID || !title.trim() || !url.trim()) return
@@ -185,10 +208,44 @@ export function StrmPage() {
只需要填写自己的访问域名,系统会按媒体库内每个媒体批量生成可播放的 .strm 文件。
</p>
</div>
<span className="rounded-full border border-emerald-300/40 bg-emerald-400/10 px-3 py-1 text-xs font-semibold text-emerald-500">
开启后可 STRM 播放
<span className={`rounded-full border px-3 py-1 text-xs font-semibold ${
strmEnabled
? 'border-emerald-300/40 bg-emerald-400/10 text-emerald-500'
: 'border-red-300/40 bg-red-400/10 text-red-500'
}`}>
{strmEnabled ? 'STRM 播放已启用' : 'STRM 播放已关闭'}
</span>
</div>
<div className="grid gap-3 rounded-2xl border border-gray-200 bg-white/70 p-4 md:grid-cols-[1fr_1fr_auto]">
<label className="flex items-start gap-3 text-sm text-ink-100">
<input
type="checkbox"
className="mt-1 h-4 w-4 accent-primary-400"
checked={strmEnabled}
onChange={(e) => setStrmEnabled(e.target.checked)}
/>
<span>
<span className="block font-medium text-ink-600">启用 STRM 播放</span>
<span className="text-xs text-ink-50">关闭后不会跳转 STRM/网盘直链;本地文件仍按本地文件播放。</span>
</span>
</label>
<label className="flex items-start gap-3 text-sm text-ink-100">
<input
type="checkbox"
className="mt-1 h-4 w-4 accent-primary-400"
checked={autoGenerate}
onChange={(e) => setAutoGenerate(e.target.checked)}
/>
<span>
<span className="block font-medium text-ink-600">扫描后自动刷新 STRM 文件</span>
<span className="text-xs text-ink-50">默认关闭,避免扫描大型网盘库时重复写文件。</span>
</span>
</label>
<button type="button" className="neon-button self-center" disabled={savingSettings} onClick={saveSTRMSettings}>
{savingSettings ? <Loader2 size={16} className="animate-spin" /> : <Save size={16} />}
保存开关
</button>
</div>
<form onSubmit={onGenerate} className="grid gap-3 md:grid-cols-4">
<select
required