Getting the playlist processing logic organized and implemented.
Signed-off-by: Cliff Hill <xlorep@darkhelm.org>
This commit is contained in:
@@ -2,6 +2,7 @@
|
||||
|
||||
from enum import StrEnum
|
||||
from enum import auto
|
||||
from typing import Generator
|
||||
|
||||
|
||||
class Sublist(StrEnum):
|
||||
@@ -20,3 +21,9 @@ class Sublist(StrEnum):
|
||||
LEAST_RECENTLY_PLAYED = auto()
|
||||
LEAST_OFTEN_PLAYED = auto()
|
||||
RANDOM = auto()
|
||||
|
||||
def standard(self) -> Generator["Sublist", None, None]:
|
||||
"""Generate the standard sublist types one at a time."""
|
||||
yield self.LEAST_RECENTLY_ADDED
|
||||
yield self.LEAST_RECENTLY_PLAYED
|
||||
yield self.LEAST_OFTEN_PLAYED
|
||||
@@ -1,8 +1,8 @@
|
||||
"""All database models are defined here."""
|
||||
|
||||
import enum
|
||||
from datetime import datetime
|
||||
|
||||
from playlist import enums
|
||||
from sqlalchemy import Boolean
|
||||
from sqlalchemy import Column
|
||||
from sqlalchemy import DateTime
|
||||
@@ -112,7 +112,7 @@ class PlaylistTrack:
|
||||
holiday_id: Mapped[int] = mapped_column(ForeignKey("holidays.id"), nullable=True)
|
||||
episode_id: Mapped[int] = mapped_column(ForeignKey("episodes.id"), nullable=True)
|
||||
favorite: Mapped[bool] = mapped_column(Boolean, nullable=True)
|
||||
sublist: Mapped[str] = mapped_column(Enum(enum.Sublist), nullable=True)
|
||||
sublist: Mapped[str] = mapped_column(Enum(enums.Sublist), nullable=True)
|
||||
|
||||
playlist: Mapped["Playlist"] = relationship("Playlist")
|
||||
track: Mapped["Track"] = relationship("Track")
|
||||
@@ -171,6 +171,9 @@ class Track:
|
||||
genre: Mapped["Genre"] = relationship("Genre", back_populates="tracks")
|
||||
category_id: Mapped[int] = mapped_column(ForeignKey("categories.id"))
|
||||
category: Mapped["Category"] = relationship("Category", back_populates="tracks")
|
||||
holidays: Mapped[list["Holiday"]] = relationship(
|
||||
"Holiday", secondary=holiday_track_association, back_populates="tracks"
|
||||
)
|
||||
playlist_tracks: Mapped[list["PlaylistTrack"]] = relationship(
|
||||
"PlaylistTrack", back_populates="track"
|
||||
)
|
||||
|
||||
204
backend/src/playlist/playlist_generator.py
Normal file
204
backend/src/playlist/playlist_generator.py
Normal file
@@ -0,0 +1,204 @@
|
||||
"""Contains the functions that are used to generate the playlist."""
|
||||
|
||||
import asyncio
|
||||
|
||||
from aiologger import Logger
|
||||
from playlist import enums
|
||||
from playlist import sql
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
|
||||
logger = Logger.with_default_handlers()
|
||||
|
||||
|
||||
async def generate_sublists(
|
||||
session: AsyncSession,
|
||||
playlist_entries: set[int],
|
||||
sublist_counts: dict[enums.Sublist, int],
|
||||
is_category: bool,
|
||||
type_id: int,
|
||||
is_favorite: bool,
|
||||
base_count: int,
|
||||
max_count: int,
|
||||
) -> None:
|
||||
"""Generate sublists for a given category and type (favorites or general).
|
||||
|
||||
Args:
|
||||
session (AsyncSession): The database session.
|
||||
playlist_entries (set[int]): The set of track IDs in the playlist.
|
||||
sublist_counts (dict[enums.Sublist, int]): The counts of existing tracks in the
|
||||
given sublist in the current playlist.
|
||||
is_category (bool): True if it is a category, false if it is a holiday.
|
||||
type_id (int): The category or holiday ID for which to generate playlists.
|
||||
is_favorite (bool): True for favorites only, False for general tracks.
|
||||
base_count (int): The base number of tracks to be included from regular sublists.
|
||||
max_count (int): The max number of tracks to be included from regular sublists.
|
||||
"""
|
||||
for sublist in enums.Sublist.standard():
|
||||
existing_count = sublist_counts.get(sublist, 0)
|
||||
allowed_count = base_count - existing_count
|
||||
tracks = await sql.fetch_tracks_for_sublist(
|
||||
session,
|
||||
is_category,
|
||||
type_id,
|
||||
is_favorite,
|
||||
sublist,
|
||||
allowed_count,
|
||||
playlist_entries,
|
||||
)
|
||||
for track in tracks:
|
||||
await sql.insert_into_playlist(
|
||||
session, is_category, type_id, track.id, sublist, is_favorite
|
||||
)
|
||||
playlist_entries.add(track.id)
|
||||
|
||||
# Calculate how many tracks to randomly add
|
||||
remaining_count = max_count - len(playlist_entries)
|
||||
if remaining_count > 0:
|
||||
await generate_random_sublist(
|
||||
session,
|
||||
playlist_entries,
|
||||
sublist_counts.get(enums.Sublist.RANDOM, 0),
|
||||
type_id,
|
||||
is_favorite,
|
||||
remaining_count,
|
||||
)
|
||||
|
||||
|
||||
async def generate_random_sublist(
|
||||
session: AsyncSession,
|
||||
playlist_entries: set[int],
|
||||
existing_count: int,
|
||||
is_category: bool,
|
||||
type_id: int,
|
||||
is_favorite: bool,
|
||||
count: int,
|
||||
) -> None:
|
||||
"""Generate random sublist for the remaining slots in the playlist.
|
||||
|
||||
Args:
|
||||
session (AsyncSession): The database session.
|
||||
playlist_entries (set[int]): The set of track IDs in the playlist.
|
||||
existing_count (int): The number of random tracks already existing in the playlist.
|
||||
is_category (bool): True if it is a category, false if it is a holiday.
|
||||
type_id (int): The category or holiday ID for which to generate the random sublist.
|
||||
is_favorite (bool): True for favorites only, False for general tracks.
|
||||
count (int): Number of random tracks to add.
|
||||
existing_entries (set): Set of track IDs already added to avoid duplicates.
|
||||
"""
|
||||
allowed_count = count - existing_count
|
||||
random_tracks = await sql.fetch_tracks_for_sublist(
|
||||
session,
|
||||
type_id,
|
||||
is_favorite,
|
||||
enums.Sublist.RANDOM,
|
||||
count,
|
||||
allowed_count,
|
||||
playlist_entries,
|
||||
)
|
||||
for track in random_tracks:
|
||||
await sql.insert_into_playlist(
|
||||
session, type_id, track.id, enums.Sublist.RANDOM, is_favorite
|
||||
)
|
||||
playlist_entries.add(track.id)
|
||||
if len(playlist_entries) >= count:
|
||||
break
|
||||
|
||||
|
||||
async def process_playlists(
|
||||
playlist_entries: set[int],
|
||||
playlist_counts: dict[tuple[bool, int, bool, enums.Sublist], int],
|
||||
is_category: bool,
|
||||
type_id: int,
|
||||
favorite_base: int,
|
||||
general_base: int,
|
||||
favorite_max: int,
|
||||
general_max: int,
|
||||
) -> None:
|
||||
"""Process playlists for both favorites and general tracks for a given category or holiday.
|
||||
|
||||
Args:
|
||||
playlist_entries (set[int]): The IDs of the tracks currently in the playlist.
|
||||
playlist_counts (dict[tuple[bool, int, bool, enums.Sublist], int]): Counts of the
|
||||
different sublist types within the playlist.
|
||||
is_category (bool): True if it is a category, false if it is a holiday.
|
||||
type_id (int): The category or holiday ID for which to generate the random sublist.
|
||||
favorite_base (int): The maximum number of tracks for a favorite standard sublist.
|
||||
general_base (int): The maximum number of tracks for a general standard sublist.
|
||||
favorite_max (int): The maximum number of tracks for favorite tracks.
|
||||
general_max (int): The maximum number of tracks for general tracks.
|
||||
"""
|
||||
async with sql.async_session_maker() as session:
|
||||
favorite_counts = {
|
||||
key[3]: count
|
||||
for key, count in playlist_counts.items()
|
||||
if key[0] is is_category and key[1] == type_id and key[2] is True
|
||||
}
|
||||
general_counts = {
|
||||
key[3]: count
|
||||
for key, count in playlist_counts.items()
|
||||
if key[0] is is_category and key[1] == type_id and key[2] is False
|
||||
}
|
||||
await generate_sublists(
|
||||
session,
|
||||
playlist_entries,
|
||||
favorite_counts,
|
||||
is_category,
|
||||
type_id,
|
||||
True,
|
||||
favorite_base,
|
||||
favorite_max,
|
||||
)
|
||||
await generate_sublists(
|
||||
session,
|
||||
playlist_entries,
|
||||
general_counts,
|
||||
is_category,
|
||||
type_id,
|
||||
False,
|
||||
general_base,
|
||||
general_max,
|
||||
)
|
||||
|
||||
|
||||
async def generate_playlist_for_all() -> None:
|
||||
"""Generate playlists for all categories & holidays."""
|
||||
(
|
||||
existing_track_ids,
|
||||
existing_playlist_counts,
|
||||
) = await sql.get_existing_playlist_track_info()
|
||||
category_ids = tuple(await sql.gen_all_category_ids())
|
||||
holiday_ids = tuple(sql.gen_active_holiday_ids())
|
||||
stats = await sql.get_statistics()
|
||||
category_favorite_base = stats.max_regular_favorite_base
|
||||
category_general_base = stats.max_regular_general_base
|
||||
category_favorite_max = stats.max_regular_favorite_tracks
|
||||
category_general_max = stats.max_regular_general_tracks
|
||||
holiday_favorite_base = stats.max_holiday_favorite_base
|
||||
holiday_general_base = stats.max_holiday_general_base
|
||||
holiday_favorite_max = stats.max_holiday_favorite_tracks
|
||||
holiday_general_max = stats.max_holiday_general_tracks
|
||||
tasks = [
|
||||
process_playlists(
|
||||
playlist_entries=existing_track_ids,
|
||||
playlist_counts=existing_playlist_counts,
|
||||
is_category=True,
|
||||
type_id=category_id,
|
||||
favorite_base=category_favorite_base,
|
||||
general_base=category_general_base,
|
||||
favorite_max=category_favorite_max,
|
||||
general_max=category_general_max,
|
||||
)
|
||||
for category_id in category_ids
|
||||
] + [
|
||||
process_playlists(
|
||||
is_category=False,
|
||||
type_id=holiday_id,
|
||||
favorite_base=holiday_favorite_base,
|
||||
general_base=holiday_general_base,
|
||||
favorite_max=holiday_favorite_max,
|
||||
general_max=holiday_general_max,
|
||||
)
|
||||
for holiday_id in holiday_ids
|
||||
]
|
||||
await asyncio.gather(*tasks)
|
||||
@@ -1,11 +1,15 @@
|
||||
"""Contains all SQL code here."""
|
||||
|
||||
from collections import Counter
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from typing import AsyncGenerator
|
||||
|
||||
import env
|
||||
import models
|
||||
from aiologger import Logger
|
||||
from playlist import enums
|
||||
from playlist import env
|
||||
from playlist import models
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
@@ -110,9 +114,7 @@ async def update_statistics_timestamps(operation: str) -> None:
|
||||
async with async_session_maker() as session:
|
||||
try:
|
||||
async with session.begin():
|
||||
stats = await session.execute(
|
||||
select(models.Statistics).order_by(models.Statistics.id)
|
||||
)
|
||||
stats = await session.get(models.Statistics, 1)
|
||||
stats_obj = stats.scalars().first()
|
||||
|
||||
if operation == "start":
|
||||
@@ -214,32 +216,6 @@ async def remove_holiday_by_id(session: AsyncSession, holiday_id: int) -> None:
|
||||
await session.rollback()
|
||||
|
||||
|
||||
async def update_average_track_length(new_average: float) -> None:
|
||||
"""Update the average track length in the statistics model with the geometric mean.
|
||||
|
||||
Args:
|
||||
new_average (float): The new average track length to be set.
|
||||
"""
|
||||
async with async_session_maker() as session:
|
||||
try:
|
||||
async with session.begin():
|
||||
# Assuming there's only one statistics record, or you might need to handle this differently
|
||||
stats = await session.get(
|
||||
models.Statistics, 1
|
||||
) # Adjust based on how you access the statistics record
|
||||
if stats:
|
||||
stats.average_track_length = new_average
|
||||
await session.commit()
|
||||
else:
|
||||
# If no statistics entry exists, create one
|
||||
new_stats = models.Statistics(average_track_length=new_average)
|
||||
session.add(new_stats)
|
||||
await session.commit()
|
||||
except Exception as e:
|
||||
print(f"Failed to update the average track length: {e}")
|
||||
await session.rollback()
|
||||
|
||||
|
||||
async def update_statistics(
|
||||
new_average_duration: float,
|
||||
new_podcast_length: float,
|
||||
@@ -283,9 +259,7 @@ async def update_statistics(
|
||||
try:
|
||||
async with session.begin():
|
||||
# Assuming there's only one statistics record, or you might need to handle this differently
|
||||
stats = await session.get(
|
||||
models.Statistics, 1
|
||||
) # Adjust based on how you access the statistics record
|
||||
stats = await session.get(models.Statistics, 1)
|
||||
if stats:
|
||||
stats.average_track_length = new_average_duration
|
||||
stats.total_podcast_length = new_podcast_length
|
||||
@@ -328,3 +302,173 @@ async def update_statistics(
|
||||
except Exception as e:
|
||||
print(f"Failed to update the statistics model: {e}")
|
||||
await session.rollback()
|
||||
|
||||
|
||||
async def fetch_tracks_for_sublist(
|
||||
session: AsyncSession,
|
||||
is_category: bool,
|
||||
type_id: int,
|
||||
is_favorite: bool,
|
||||
sublist: enums.Sublist,
|
||||
limit: int,
|
||||
playlist_entries: set[int],
|
||||
) -> AsyncGenerator[models.Track, None]:
|
||||
"""Fetch tracks for a given sublist type within a category.
|
||||
|
||||
Args:
|
||||
session (AsyncSession): The database session.
|
||||
is_category (bool): True if a category, false if a holiday.
|
||||
type_id (int): Category or holiday ID for track filtering.
|
||||
is_favorite (bool): Flag indicating if only favorite tracks should be fetched.
|
||||
sublist (enums.Sublist): Type of sublist to fetch.
|
||||
limit (int): Number of tracks to fetch.
|
||||
playlist_entries (set[int]): The set of track IDs already in the playlist.
|
||||
|
||||
Returns:
|
||||
AsyncGenerator[models.Track, None]: Generates the tracks fitting the criteria.
|
||||
"""
|
||||
try:
|
||||
query = select(models.Track)
|
||||
if is_category:
|
||||
query = query.where(models.Track.category_id == type_id)
|
||||
else:
|
||||
query = query.join(models.Track.holidays).where(
|
||||
models.Holiday.id == type_id
|
||||
)
|
||||
|
||||
query = query.filter(models.Track.id.not_in(playlist_entries))
|
||||
|
||||
if is_favorite:
|
||||
query = query.filter(models.Track.rating == 5)
|
||||
|
||||
match sublist:
|
||||
case enums.Sublist.LEAST_RECENTLY_PLAYED:
|
||||
query = query.order_by(models.Track.last_played.asc())
|
||||
case enums.Sublist.LEAST_OFTEN_PLAYED:
|
||||
query = query.order_by(models.Track.play_count.asc())
|
||||
case enums.Sublist.LEAST_RECENTLY_ADDED:
|
||||
query = query.order_by(models.Track.date_added.asc())
|
||||
case enums.Sublist.RANDOM:
|
||||
query = query.order_by(func.random())
|
||||
case _:
|
||||
raise ValueError(f"Unknown sublist type: {sublist}")
|
||||
|
||||
query = query.limit(limit)
|
||||
result = await session.execute(query)
|
||||
for track in result.scalars():
|
||||
yield track
|
||||
except SQLAlchemyError as e:
|
||||
await logger.error(f"Database error occurred while fetching tracks: {e}")
|
||||
raise
|
||||
except Exception as e:
|
||||
await logger.error(f"Unexpected error occurred while fetching tracks: {e}")
|
||||
raise
|
||||
|
||||
|
||||
async def insert_into_playlist( # noqa: C901
|
||||
session: AsyncSession,
|
||||
is_category: bool,
|
||||
type_id: int,
|
||||
track_id: int,
|
||||
sublist: enums.Sublist,
|
||||
is_favorite: bool,
|
||||
) -> None:
|
||||
"""Insert a track into the playlist.
|
||||
|
||||
Args:
|
||||
session (AsyncSession): The database session.
|
||||
is_category (bool): True is a category, false is a holiday.
|
||||
type_id (int): Category or holiday ID associated with the playlist.
|
||||
track_id (int): Track ID to be added to the playlist.
|
||||
sublist (enums.Sublist): Sublist type under which the track is added.
|
||||
is_favorite (bool): Indicates if the track is added as a favorite.
|
||||
"""
|
||||
try:
|
||||
if is_category:
|
||||
playlist_track = models.PlaylistTrack(
|
||||
category_id=type_id,
|
||||
track_id=track_id,
|
||||
sublist_type=sublist,
|
||||
is_favorite=is_favorite,
|
||||
)
|
||||
else:
|
||||
playlist_track = models.PlaylistTrack(
|
||||
holiday_id=type_id,
|
||||
track_id=track_id,
|
||||
sublist_type=sublist,
|
||||
is_favorite=is_favorite,
|
||||
)
|
||||
session.add(playlist_track)
|
||||
await session.commit()
|
||||
except SQLAlchemyError as e:
|
||||
await logger.error(
|
||||
f"Database error occurred while inserting into playlist: {e}"
|
||||
)
|
||||
await session.rollback()
|
||||
raise
|
||||
except Exception as e:
|
||||
await logger.error(
|
||||
f"Unexpected error occurred while inserting into playlist: {e}"
|
||||
)
|
||||
await session.rollback()
|
||||
raise
|
||||
|
||||
async def gen_all_category_ids() -> AsyncGenerator[int, None, None]:
|
||||
"""Generate the current category ids from the database.
|
||||
|
||||
Yields:
|
||||
int: The category id.
|
||||
"""
|
||||
async with async_session_maker() as session:
|
||||
query = select(models.Category)
|
||||
result = session.execute(query)
|
||||
for row in result.scalars():
|
||||
yield row.id
|
||||
|
||||
async def gen_active_holiday_ids() -> AsyncGenerator[int, None, None]:
|
||||
"""Generate the current active holiday ids from the database.
|
||||
|
||||
Yields:
|
||||
int: The holiday id.
|
||||
"""
|
||||
async with async_session_maker() as session:
|
||||
query = select(models.Holiday).where(models.Holiday.is_active is True)
|
||||
result = session.execute(query)
|
||||
for row in result.scalars():
|
||||
yield row.id
|
||||
|
||||
async def get_existing_playlist_track_info() -> (
|
||||
tuple[set[int], dict[tuple[bool, int, bool, enums.Sublist], int]]
|
||||
):
|
||||
"""Get the existing playlist track information.
|
||||
|
||||
Returns:
|
||||
tuple[set[int], dict[tuple[bool, int, bool, enums.Sublist], int]]: Two pieces:
|
||||
First is a set of the ids of all of the tracks in the playlist. Second is a
|
||||
dict with the key being the combination of is_category, type_id, is_favorite,
|
||||
and sublist; the value is the count of that unique key combination.
|
||||
"""
|
||||
async with async_session_maker() as session:
|
||||
# TODO: Set up user-specific query here.
|
||||
query = select(models.PlaylistTrack).where(
|
||||
models.PlaylistTrack.episode_id is None
|
||||
)
|
||||
result = session.execute(query)
|
||||
ids = set()
|
||||
data = []
|
||||
for row in result.scalars():
|
||||
if row.episode_id is not None:
|
||||
continue
|
||||
ids.append(row.track_id)
|
||||
is_category = row.category_id is not None
|
||||
type_id = row.category_id if is_category else row.holiday_id
|
||||
is_favorite = row.is_favorite
|
||||
sublist = row.sublist
|
||||
data.append((is_category, type_id, is_favorite, sublist))
|
||||
data_counts = Counter(data)
|
||||
return ids, data_counts
|
||||
|
||||
async def get_statistics() -> models.Statistics:
|
||||
stats = await session.get(models.Statistics, 1)
|
||||
stats_obj = stats.scalars().first()
|
||||
return stats_obj
|
||||
|
||||
Reference in New Issue
Block a user