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
+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")