stuff
This commit is contained in:
@@ -29,3 +29,7 @@ class PlaylistUpdate(BaseModel):
|
||||
class PlaylistAddTrack(BaseModel):
|
||||
track_id: uuid.UUID
|
||||
position: float | None = None
|
||||
|
||||
|
||||
class PlaylistReorder(BaseModel):
|
||||
track_ids: list[uuid.UUID]
|
||||
|
||||
+34
-3
@@ -13,11 +13,17 @@ from app.api.deps import (
|
||||
TrackRepoDep,
|
||||
)
|
||||
from app.api.schemas.pagination import PagedResponse
|
||||
from app.api.schemas.playlist import PlaylistAddTrack, PlaylistCreate, PlaylistOut, PlaylistUpdate
|
||||
from app.api.schemas.playlist import (
|
||||
PlaylistAddTrack,
|
||||
PlaylistCreate,
|
||||
PlaylistOut,
|
||||
PlaylistReorder,
|
||||
PlaylistUpdate,
|
||||
)
|
||||
from app.api.schemas.track import TrackOut
|
||||
from app.api.v1.tracks import _build_track_out
|
||||
from app.domain.entities.playlist import Playlist
|
||||
from app.domain.errors import NotFoundError, PermissionDeniedError
|
||||
from app.domain.errors import NotFoundError, PermissionDeniedError, ValidationError
|
||||
from app.infrastructure.db.repositories.playlist_repository import SqlAlchemyPlaylistRepository
|
||||
|
||||
router = APIRouter(prefix="/playlists", tags=["playlists"])
|
||||
@@ -182,7 +188,32 @@ async def remove_playlist_track(
|
||||
|
||||
|
||||
@router.put("/{playlist_id}/tracks/reorder")
|
||||
async def reorder_playlist_tracks(playlist_id: uuid.UUID, _: CurrentUser) -> Any: ...
|
||||
async def reorder_playlist_tracks(
|
||||
playlist_id: uuid.UUID,
|
||||
body: PlaylistReorder,
|
||||
playlist_repo: PlaylistRepoDep,
|
||||
user: CurrentUser,
|
||||
) -> PlaylistOut:
|
||||
playlist = await playlist_repo.get_by_id(playlist_id)
|
||||
if playlist is None:
|
||||
raise NotFoundError(f"Playlist {playlist_id} not found.")
|
||||
if playlist.owner_id != user.id:
|
||||
raise PermissionDeniedError("You don't own this playlist.")
|
||||
|
||||
total = await playlist_repo.get_track_total(playlist_id)
|
||||
current_tracks = (
|
||||
await playlist_repo.get_tracks(playlist_id, limit=total, offset=0) if total else []
|
||||
)
|
||||
current_ids = {t.id for t in current_tracks}
|
||||
given_ids = body.track_ids
|
||||
if len(given_ids) != len(set(given_ids)) or set(given_ids) != current_ids:
|
||||
raise ValidationError("track_ids must be a permutation of the playlist's current tracks.")
|
||||
|
||||
await playlist_repo.reorder_tracks(playlist_id, given_ids)
|
||||
updated = await playlist_repo.get_by_id(playlist_id)
|
||||
assert updated is not None
|
||||
items = await _build_playlist_out([updated], playlist_repo)
|
||||
return items[0]
|
||||
|
||||
|
||||
@router.get("/{playlist_id}/cover")
|
||||
|
||||
@@ -278,11 +278,15 @@ class PlaylistRepository(Protocol):
|
||||
self, playlist_id: uuid.UUID, *, limit: int, offset: int
|
||||
) -> list[Track]: ...
|
||||
async def get_track_total(self, playlist_id: uuid.UUID) -> int: ...
|
||||
async def has_track(self, playlist_id: uuid.UUID, track_id: uuid.UUID) -> bool: ...
|
||||
async def add_track(
|
||||
self, playlist_id: uuid.UUID, track_id: uuid.UUID, *, position: float
|
||||
) -> None: ...
|
||||
async def remove_track(self, playlist_id: uuid.UUID, track_id: uuid.UUID) -> None: ...
|
||||
async def max_position(self, playlist_id: uuid.UUID) -> float: ...
|
||||
async def reorder_tracks(
|
||||
self, playlist_id: uuid.UUID, ordered_track_ids: list[uuid.UUID]
|
||||
) -> None: ...
|
||||
# list must come after any method using list[...] in its signature (name shadowing)
|
||||
async def list(self, *, owner_id: uuid.UUID, limit: int, offset: int) -> list[Playlist]: ...
|
||||
|
||||
|
||||
@@ -138,9 +138,22 @@ class SqlAlchemyPlaylistRepository:
|
||||
async def get_track_total(self, playlist_id: uuid.UUID) -> int:
|
||||
return await self.track_count(playlist_id)
|
||||
|
||||
async def has_track(self, playlist_id: uuid.UUID, track_id: uuid.UUID) -> bool:
|
||||
row = (
|
||||
await self._session.execute(
|
||||
select(PlaylistTrackModel.id).where(
|
||||
PlaylistTrackModel.playlist_id == playlist_id,
|
||||
PlaylistTrackModel.track_id == track_id,
|
||||
)
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
return row is not None
|
||||
|
||||
async def add_track(
|
||||
self, playlist_id: uuid.UUID, track_id: uuid.UUID, *, position: float
|
||||
) -> None:
|
||||
if await self.has_track(playlist_id, track_id):
|
||||
return
|
||||
row = PlaylistTrackModel(playlist_id=playlist_id, track_id=track_id, position=position)
|
||||
self._session.add(row)
|
||||
playlist = await self._session.get(PlaylistModel, playlist_id)
|
||||
@@ -148,6 +161,26 @@ class SqlAlchemyPlaylistRepository:
|
||||
playlist.version = playlist.version + 1
|
||||
await self._session.flush()
|
||||
|
||||
async def reorder_tracks(
|
||||
self, playlist_id: uuid.UUID, ordered_track_ids: list[uuid.UUID]
|
||||
) -> None:
|
||||
rows = (
|
||||
(
|
||||
await self._session.execute(
|
||||
select(PlaylistTrackModel).where(PlaylistTrackModel.playlist_id == playlist_id)
|
||||
)
|
||||
)
|
||||
.scalars()
|
||||
.all()
|
||||
)
|
||||
by_track_id = {row.track_id: row for row in rows}
|
||||
for position, track_id in enumerate(ordered_track_ids, start=1):
|
||||
by_track_id[track_id].position = float(position)
|
||||
playlist = await self._session.get(PlaylistModel, playlist_id)
|
||||
if playlist is not None:
|
||||
playlist.version = playlist.version + 1
|
||||
await self._session.flush()
|
||||
|
||||
async def remove_track(self, playlist_id: uuid.UUID, track_id: uuid.UUID) -> None:
|
||||
row = (
|
||||
await self._session.execute(
|
||||
|
||||
Reference in New Issue
Block a user