diff --git a/src/event/crud.py b/src/event/crud.py index e46c176..ae5e7dc 100644 --- a/src/event/crud.py +++ b/src/event/crud.py @@ -5,61 +5,38 @@ from sqlalchemy import delete, select from sqlalchemy.ext.asyncio import AsyncSession +from config import settings from constants import TZ_INFO from event.constants import EventStatusEnum +from event.models import Event, GetEventQueryParams from event.tables import EventDB +from image_asset.tables import ImageAssetDB -async def get_all_events(db_session: AsyncSession, include_cancelled: bool) -> Sequence[EventDB]: - """ - Get all non-cancelled events, by default. - - Args: - db_session: database session - - Returns: - A list of events, ordered by descending start_datetime and end_datetime. - """ - query = select(EventDB) if include_cancelled else select(EventDB).where(EventDB.status != EventStatusEnum.CANCELLED) - return (await db_session.scalars(query.order_by(EventDB.start_datetime.desc(), EventDB.end_datetime.desc()))).all() - - -async def get_events_for_this_year_month( - db_session: AsyncSession, - year: int, - month: int, - include_cancelled: bool, -) -> Sequence[EventDB]: - """ - Gets all events that occur during the month, including ones that start before the month and year, - and end after the month and year, assuming America/Vancouver timezone. - - Args: - db_session: database session - year: the year to check - month: the month to check - - Returns: - A list of events that overlap the month and year, ordered by ascending start_datetime and end_datetime. - """ - period_start = datetime(year, month, 1, tzinfo=TZ_INFO) - - if month == 12: - period_end = datetime(year + 1, 1, 1, tzinfo=TZ_INFO) - else: - period_end = datetime(year, month + 1, 1, tzinfo=TZ_INFO) - query = ( - select(EventDB) - .where(EventDB.start_datetime < period_end, EventDB.end_datetime >= period_start) - .order_by(EventDB.start_datetime, EventDB.end_datetime) - ) +async def get_events(db_session: AsyncSession, q: GetEventQueryParams) -> list[Event]: + query = select(EventDB, ImageAssetDB.storage_key).outerjoin(ImageAssetDB, EventDB.image_id == ImageAssetDB.image_id) + + if q.current: + query = query.where(EventDB.end_datetime > datetime.now(tz=TZ_INFO)) - if not include_cancelled: + if not q.include_cancelled: query = query.where(EventDB.status != EventStatusEnum.CANCELLED) - events = (await db_session.scalars(query)).all() + query = ( + query.order_by(EventDB.start_datetime.desc(), EventDB.end_datetime.desc()) + if q.desc + else query.order_by(EventDB.start_datetime.asc(), EventDB.end_datetime.asc()) + ) + + rows = (await db_session.execute(query)).all() + media_base_url = settings.media_base_url.rstrip("/") - return events + return [ + Event.model_validate(event).model_copy( + update={"image_url": f"{media_base_url}/{storage_key}" if storage_key is not None else None} + ) + for event, storage_key in rows + ] async def get_event_by_eid(db_session: AsyncSession, eid: int) -> EventDB | None: diff --git a/src/event/models.py b/src/event/models.py index 041d22a..1afe67b 100644 --- a/src/event/models.py +++ b/src/event/models.py @@ -1,7 +1,7 @@ from collections.abc import Sequence from uuid import UUID -from pydantic import AwareDatetime, BaseModel, ConfigDict, model_validator +from pydantic import AwareDatetime, BaseModel, ConfigDict, Field, model_validator from event.constants import EventStatusEnum @@ -29,6 +29,7 @@ class Event(BaseEvent): eid: int group_id: UUID | None = None + image_url: str | None = None class EventCreate(BaseEvent): @@ -72,3 +73,9 @@ class GroupEventDeleteResponse(BaseModel): result: bool group_id: UUID deleted_eids: Sequence[int] + + +class GetEventQueryParams(BaseModel): + include_cancelled: bool = Field(False, description="Include cancelled events in the response.") + current: bool = Field(False, description="Only get events that haven't ended yet.") + desc: bool = Field(False, description="Sorts by descending start time and end time.") diff --git a/src/event/urls.py b/src/event/urls.py index 08f78a5..03a9e2d 100644 --- a/src/event/urls.py +++ b/src/event/urls.py @@ -1,14 +1,22 @@ import uuid from typing import Annotated -from fastapi import APIRouter, Body, Depends, HTTPException, status +from fastapi import APIRouter, Body, Depends, HTTPException, Query, status from fastapi.encoders import jsonable_encoder -from pydantic import ValidationError +from pydantic import BaseModel, Field, ValidationError import database import event.crud from dependencies import MonthPath, YearPath, perm_admin -from event.models import Event, EventCreate, EventDelete, EventUpdate, GroupEvent, GroupEventDeleteResponse +from event.models import ( + Event, + EventCreate, + EventDelete, + EventUpdate, + GetEventQueryParams, + GroupEvent, + GroupEventDeleteResponse, +) from event.tables import EventDB from utils.shared_models import DetailModel @@ -20,27 +28,12 @@ @router.get( "", - description="Get all events", + description="Get events, with parameters", response_model=list[Event], - operation_id="get_all_events", + operation_id="get_events", ) -async def get_all_events(db_session: database.DBSession, include_cancelled: bool = False): - events_list = await event.crud.get_all_events(db_session, include_cancelled) - - return events_list - - -@router.get( - "/{year}/{month}", - description="Get events that overlap in the year and month.", - response_model=list[Event], - operation_id="get_events_for_this_year_month", -) -async def get_events_for_this_year_month( - db_session: database.DBSession, year: YearPath, month: MonthPath, include_cancelled: bool = False -): - events_list = await event.crud.get_events_for_this_year_month(db_session, year, month, include_cancelled) - +async def get_all_events(db_session: database.DBSession, q: Annotated[GetEventQueryParams, Query()]): + events_list = await event.crud.get_events(db_session, q) return events_list diff --git a/tests/integration/test_events.py b/tests/integration/test_events.py new file mode 100644 index 0000000..19309a2 --- /dev/null +++ b/tests/integration/test_events.py @@ -0,0 +1,123 @@ +from datetime import UTC, datetime, timedelta + +import pytest +from fastapi import status +from httpx import AsyncClient + +from config import settings +from database import DBSession +from event.constants import EventStatusEnum +from event.tables import EventDB +from image_asset.tables import ImageAssetDB + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +async def seed_events(db_session: DBSession) -> None: + now = datetime.now(UTC) + image = ImageAssetDB( + storage_key="images/ongoing-event.png", + original_filename="ongoing-event.png", + ) + db_session.add(image) + await db_session.flush() + + db_session.add_all( + [ + EventDB( + name="Past event", + description="A scheduled event that has ended.", + start_datetime=now - timedelta(days=4), + end_datetime=now - timedelta(days=3), + status=EventStatusEnum.SCHEDULED, + ), + EventDB( + name="Ongoing event", + description="A scheduled event that has started but not ended.", + start_datetime=now - timedelta(days=1), + end_datetime=now + timedelta(days=1), + status=EventStatusEnum.SCHEDULED, + image_id=image.image_id, + ), + EventDB( + name="Future event", + description="A scheduled event that has not started.", + start_datetime=now + timedelta(days=2), + end_datetime=now + timedelta(days=3), + status=EventStatusEnum.SCHEDULED, + ), + EventDB( + name="Cancelled future event", + description="A cancelled event that has not started.", + start_datetime=now + timedelta(days=4), + end_datetime=now + timedelta(days=5), + status=EventStatusEnum.CANCELLED, + ), + ] + ) + await db_session.commit() + + +@pytest.mark.parametrize( + ("params", "expected_names"), + [ + ({}, ["Past event", "Ongoing event", "Future event"]), + ( + {"include_cancelled": "true"}, + ["Past event", "Ongoing event", "Future event", "Cancelled future event"], + ), + ({"current": "true"}, ["Ongoing event", "Future event"]), + ( + {"current": "true", "include_cancelled": "true"}, + ["Ongoing event", "Future event", "Cancelled future event"], + ), + ], + ids=["defaults", "include-cancelled", "current", "current-and-include-cancelled"], +) +async def test__get_events_applies_filters( + db_session: DBSession, + client: AsyncClient, + params: dict[str, str], + expected_names: list[str], +): + await seed_events(db_session) + + response = await client.get("/api/event", params=params) + + assert response.status_code == status.HTTP_200_OK + assert [event["name"] for event in response.json()] == expected_names + + +async def test__get_events_sorts_descending(db_session: DBSession, client: AsyncClient): + await seed_events(db_session) + + response = await client.get("/api/event", params={"desc": "true"}) + + assert response.status_code == status.HTTP_200_OK + assert [event["name"] for event in response.json()] == [ + "Future event", + "Ongoing event", + "Past event", + ] + + +async def test__get_events_includes_image_url_without_dropping_unlinked_events( + db_session: DBSession, + client: AsyncClient, +): + await seed_events(db_session) + + response = await client.get("/api/event") + + assert response.status_code == status.HTTP_200_OK + events_by_name = {event["name"]: event for event in response.json()} + assert events_by_name["Ongoing event"]["image_url"] == ( + f"{settings.media_base_url.rstrip('/')}/images/ongoing-event.png" + ) + assert events_by_name["Past event"]["image_url"] is None + + +async def test__get_events_rejects_invalid_boolean_query_params(client: AsyncClient): + response = await client.get("/api/event", params={"current": "not-a-boolean"}) + + assert response.status_code == status.HTTP_422_UNPROCESSABLE_CONTENT