diff --git a/backend/src/playlist/sql.py b/backend/src/playlist/sql.py index bdaa604..8b22141 100644 --- a/backend/src/playlist/sql.py +++ b/backend/src/playlist/sql.py @@ -4,10 +4,8 @@ from datetime import datetime from typing import Any import env +import models from aiologger import Logger -from models import Genre -from models import Statistics -from models import Track from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import create_async_engine @@ -52,19 +50,19 @@ async def insert_or_update_track(track_data: dict[str, Any]) -> None: async with async_session_maker() as session: try: async with session.begin(): - genre = await session.get(Genre, track_data.get("genre_id")) + genre = await session.get(models.Genre, track_data.get("genre_id")) if not genre: - genre = Genre(name=track_data["genre_name"]) + genre = models.Genre(name=track_data["genre_name"]) session.add(genre) await session.flush() # Ensures 'genre' is persisted and has an 'id' - track = await session.get(Track, track_data.get("id")) + track = await session.get(models.Track, track_data.get("id")) if track: for key, value in track_data.items(): setattr(track, key, value) await logger.info(f"Updated track: {track.title}") else: - track = Track(**track_data) + track = models.Track(**track_data) session.add(track) await logger.info(f"Inserted new track: {track.title}") @@ -82,7 +80,11 @@ async def cleanup_unused_genres() -> None: async with async_session_maker() as session: try: async with session.begin(): - stmt = select(Genre).outerjoin(Track).filter(Track.id is None) + stmt = ( + select(models.Genre) + .outerjoin(models.Track) + .filter(models.Track.id is None) + ) result = await session.execute(stmt) unused_genres = result.scalars().all() @@ -109,7 +111,7 @@ async def update_statistics_timestamps(operation: str) -> None: try: async with session.begin(): stats = await session.execute( - select(Statistics).order_by(Statistics.id) + select(models.Statistics).order_by(models.Statistics.id) ) stats_obj = stats.scalars().first()