Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 23 additions & 46 deletions src/event/crud.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
9 changes: 8 additions & 1 deletion src/event/models.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -29,6 +29,7 @@ class Event(BaseEvent):

eid: int
group_id: UUID | None = None
image_url: str | None = None


class EventCreate(BaseEvent):
Expand Down Expand Up @@ -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.")
37 changes: 15 additions & 22 deletions src/event/urls.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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


Expand Down
123 changes: 123 additions & 0 deletions tests/integration/test_events.py
Original file line number Diff line number Diff line change
@@ -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