diff --git a/app/application/recommendation_service.py b/app/application/recommendation_service.py index 256de39..350a52b 100644 --- a/app/application/recommendation_service.py +++ b/app/application/recommendation_service.py @@ -82,8 +82,8 @@ class RecommendationService: if self._recommender.is_available(): ids = await self._recommender.similar_artist_ids(artist_id, limit=limit) if ids is not None: - found = [a for i in ids if (a := await self._artists.get_by_id(i))] - return REASON_ML, found + by_id = {a.id: a for a in await self._artists.get_many(ids)} + return REASON_ML, [by_id[i] for i in ids if i in by_id] found = await self._artists.list_similar(artist_id=artist_id, limit=limit) return REASON_SIMILAR, found @@ -183,5 +183,7 @@ class RecommendationService: return None, REASON_DISCOVER async def _hydrate_tracks(self, ids: list[uuid.UUID]) -> list[Track]: - """Resolve ids → tracks preserving order, skipping any that vanished.""" - return [t for i in ids if (t := await self._tracks.get_by_id(i))] + """Resolve ids → tracks preserving order, skipping any that vanished. + One batched query rather than N per-id round-trips.""" + by_id = {t.id: t for t in await self._tracks.get_many(ids)} + return [by_id[i] for i in ids if i in by_id] diff --git a/app/domain/ports.py b/app/domain/ports.py index 786de76..8b80e71 100644 --- a/app/domain/ports.py +++ b/app/domain/ports.py @@ -142,6 +142,11 @@ class ArtistRepository(Protocol): class TrackRepository(Protocol): async def get_by_id(self, track_id: uuid.UUID) -> Track | None: ... + async def get_many(self, ids: list[uuid.UUID]) -> list[Track]: + """Resolve multiple ids in one query (unordered) — batches the per-id + lookups radio/similar would otherwise fan out into N round-trips.""" + ... + async def get_by_source(self, source: str, source_id: str) -> Track | None: ... async def add( self, diff --git a/app/infrastructure/db/repositories/track_repository.py b/app/infrastructure/db/repositories/track_repository.py index 9b5cd59..0c250b4 100644 --- a/app/infrastructure/db/repositories/track_repository.py +++ b/app/infrastructure/db/repositories/track_repository.py @@ -46,6 +46,16 @@ class SqlAlchemyTrackRepository: row = await self._session.get(TrackModel, track_id) return _to_entity(row) if row is not None else None + async def get_many(self, ids: list[uuid.UUID]) -> list[Track]: + if not ids: + return [] + rows = ( + (await self._session.execute(select(TrackModel).where(TrackModel.id.in_(ids)))) + .scalars() + .all() + ) + return [_to_entity(r) for r in rows] + async def get_by_source(self, source: str, source_id: str) -> Track | None: row = ( await self._session.execute(