diff --git a/internal/requests/service.go b/internal/requests/service.go index 4bffc775..1ca639b2 100644 --- a/internal/requests/service.go +++ b/internal/requests/service.go @@ -9,6 +9,7 @@ import ( "github.com/Silo-Server/silo-server/internal/idgen" "github.com/Silo-Server/silo-server/internal/metadata/tmdb" + "golang.org/x/sync/errgroup" ) type TMDBClient interface { @@ -22,6 +23,8 @@ type TMDBExternalIDClient interface { GetExternalIDs(ctx context.Context, mediaType string, id int) (*tmdb.ExternalIDs, error) } +const externalIDHydrationConcurrency = 4 + type SecretResolver interface { Get(ctx context.Context, key string) (string, error) } @@ -733,6 +736,44 @@ func (s *Service) hydratePresenceCandidate(ctx context.Context, mediaType MediaT return candidate } +func (s *Service) hydratePresenceCandidates(ctx context.Context, mediaType MediaType, candidates []PresenceCandidate) []PresenceCandidate { + if len(candidates) == 0 { + return candidates + } + if _, ok := s.tmdb.(TMDBExternalIDClient); !ok { + return candidates + } + + hydrated := append([]PresenceCandidate(nil), candidates...) + if externalIDHydrationConcurrency <= 1 { + for i := range hydrated { + if ctx.Err() != nil { + return hydrated + } + hydrated[i] = s.hydratePresenceCandidate(ctx, mediaType, hydrated[i]) + } + return hydrated + } + + group, groupCtx := errgroup.WithContext(ctx) + group.SetLimit(externalIDHydrationConcurrency) + for i := range hydrated { + if groupCtx.Err() != nil { + break + } + i := i + group.Go(func() error { + if err := groupCtx.Err(); err != nil { + return err + } + hydrated[i] = s.hydratePresenceCandidate(groupCtx, mediaType, hydrated[i]) + return nil + }) + } + _ = group.Wait() + return hydrated +} + func tmdbMediaType(mediaType MediaType) string { if mediaType == MediaTypeSeries { return "tv" @@ -741,12 +782,16 @@ func tmdbMediaType(mediaType MediaType) string { } func (s *Service) lookupAvailable(ctx context.Context, mediaType MediaType, ids []int) (map[int]bool, error) { + if s.presence == nil { + return map[int]bool{}, nil + } candidates := make([]PresenceCandidate, 0, len(ids)) for _, id := range ids { if id > 0 { - candidates = append(candidates, s.hydratePresenceCandidate(ctx, mediaType, PresenceCandidate{TMDBID: id})) + candidates = append(candidates, PresenceCandidate{TMDBID: id}) } } + candidates = s.hydratePresenceCandidates(ctx, mediaType, candidates) matches, err := s.lookupPresence(ctx, mediaType, candidates) if err != nil { return nil, err diff --git a/internal/requests/service_test.go b/internal/requests/service_test.go index a5bc39a8..87c2824d 100644 --- a/internal/requests/service_test.go +++ b/internal/requests/service_test.go @@ -4,6 +4,7 @@ import ( "context" "errors" "strings" + "sync" "testing" "time" @@ -309,6 +310,59 @@ func TestSearchMarksSeriesAvailableByHydratedTVDBID(t *testing.T) { } } +func TestSearchWithNilPresenceDoesNotHydrateExternalIDs(t *testing.T) { + store := newFakeStore() + store.settings.RequestsEnabled = true + tmdbClient := &fakeTMDBClient{page: &tmdb.MediaPage{Results: []tmdb.MediaResult{{ + ID: 201992, + MediaType: "series", + Title: "The Rookie: Feds", + }}}} + service := NewService(store, tmdbClient, nil) + + _, err := service.Search(context.Background(), testViewer(1), "rookie feds", MediaTypeSeries, 1) + if err != nil { + t.Fatalf("Search returned error: %v", err) + } + if len(tmdbClient.externalIDCalls) != 0 { + t.Fatalf("external ID calls = %v, want none", tmdbClient.externalIDCalls) + } +} + +func TestSearchHydratesMultipleResultsBeforePresenceLookup(t *testing.T) { + store := newFakeStore() + store.settings.RequestsEnabled = true + tmdbClient := &fakeTMDBClient{ + page: &tmdb.MediaPage{Results: []tmdb.MediaResult{ + {ID: 201992, MediaType: "series", Title: "The Rookie: Feds"}, + {ID: 1399, MediaType: "series", Title: "Game of Thrones"}, + }}, + externalIDsByID: map[int]*tmdb.ExternalIDs{ + 201992: {TVDBID: 420105, IMDbID: "tt18076310"}, + 1399: {TVDBID: 121361, IMDbID: "tt0944947"}, + }, + } + presence := &fakePresence{} + service := NewService(store, tmdbClient, presence) + + _, err := service.Search(context.Background(), testViewer(1), "series", MediaTypeSeries, 1) + if err != nil { + t.Fatalf("Search returned error: %v", err) + } + if len(presence.got) != 2 { + t.Fatalf("presence candidates = %d, want 2", len(presence.got)) + } + got := map[int]int{} + for _, candidate := range presence.got { + if candidate.TVDBID != nil { + got[candidate.TMDBID] = *candidate.TVDBID + } + } + if got[201992] != 420105 || got[1399] != 121361 { + t.Fatalf("hydrated tvdb ids = %+v", got) + } +} + func TestSearchEnrichmentHidesOtherRequesterID(t *testing.T) { store := newFakeStore() store.settings.RequestsEnabled = true @@ -963,6 +1017,7 @@ func (f *fakePresence) LookupTMDB(_ context.Context, mediaType MediaType, ids [] } type fakeTMDBClient struct { + mu sync.Mutex page *tmdb.MediaPage externalIDs *tmdb.ExternalIDs externalIDsByID map[int]*tmdb.ExternalIDs @@ -993,7 +1048,9 @@ func (f *fakeTMDBClient) DiscoverPage(context.Context, string, tmdb.DiscoverPara } func (f *fakeTMDBClient) GetExternalIDs(_ context.Context, _ string, id int) (*tmdb.ExternalIDs, error) { + f.mu.Lock() f.externalIDCalls = append(f.externalIDCalls, id) + f.mu.Unlock() if f.externalIDsByID != nil { return f.externalIDsByID[id], nil }