diff --git a/docker-compose.yml b/docker-compose.yml index 7f6ef9d..46afb26 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -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: diff --git a/internal/database/database.go b/internal/database/database.go index 8915aa8..75d4421 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -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 { diff --git a/internal/database/database_test.go b/internal/database/database_test.go index 634ee15..96e380b 100644 --- a/internal/database/database_test.go +++ b/internal/database/database_test.go @@ -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) + } + } +} diff --git a/internal/handler/admin.go b/internal/handler/admin.go index f516b8d..dcd6173 100644 --- a/internal/handler/admin.go +++ b/internal/handler/admin.go @@ -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) } } diff --git a/internal/handler/emby.go b/internal/handler/emby.go index c2ac7e2..b00d411 100644 --- a/internal/handler/emby.go +++ b/internal/handler/emby.go @@ -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)) diff --git a/internal/handler/emby_test.go b/internal/handler/emby_test.go index c7c2d04..e141323 100644 --- a/internal/handler/emby_test.go +++ b/internal/handler/emby_test.go @@ -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{}) diff --git a/internal/service/emby_compat.go b/internal/service/emby_compat.go index 502b593..702c5b5 100644 --- a/internal/service/emby_compat.go +++ b/internal/service/emby_compat.go @@ -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 } diff --git a/internal/service/emby_compat_test.go b/internal/service/emby_compat_test.go index f1a0f07..f4b39d0 100644 --- a/internal/service/emby_compat_test.go +++ b/internal/service/emby_compat_test.go @@ -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 { diff --git a/internal/service/scanner.go b/internal/service/scanner.go index a7eb10f..303bf71 100644 --- a/internal/service/scanner.go +++ b/internal/service/scanner.go @@ -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) { diff --git a/internal/service/scanner_cloud_test.go b/internal/service/scanner_cloud_test.go index 97ceff1..02823ed 100644 --- a/internal/service/scanner_cloud_test.go +++ b/internal/service/scanner_cloud_test.go @@ -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 { diff --git a/internal/service/scheduler.go b/internal/service/scheduler.go index ef28e55..1ec404c 100644 --- a/internal/service/scheduler.go +++ b/internal/service/scheduler.go @@ -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 { diff --git a/internal/service/scheduler_test.go b/internal/service/scheduler_test.go index 6c23922..5478e95 100644 --- a/internal/service/scheduler_test.go +++ b/internal/service/scheduler_test.go @@ -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) + } +} diff --git a/internal/service/stream.go b/internal/service/stream.go index 6153a4d..65365e5 100644 --- a/internal/service/stream.go +++ b/internal/service/stream.go @@ -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