diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 6b1d509..f2ec069 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -7,7 +7,7 @@ repos: - id: trailing-whitespace - repo: https://github.com/astral-sh/ruff-pre-commit # Ruff version - rev: v0.15.12 + rev: v0.16.4 hooks: # Run linter - id: ruff-check diff --git a/pyproject.toml b/pyproject.toml index 438a9cd..38bfaca 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -75,6 +75,10 @@ line-ending = "lf" select = ["E", "F", "B", "I", "N", "UP", "A", "PTH", "W", "RUF", "C4", "PIE", "Q", "FLY"] # "ANN" ignore = ["E501", "F401", "N806"] +# TODO: Re-enable once F401 is introduced globally +# [tool.ruff.lint.per-file-ignores] +# "src/load_test_db.py" = ["F401"] + # [Based]Pyright: Type checker/LSP [tool.pyright] executionEnvironments = [ diff --git a/src/alembic/versions/b793b71fff3e_update_event_table_to_have_groups_.py b/src/alembic/versions/b793b71fff3e_update_event_table_to_have_groups_.py new file mode 100644 index 0000000..2f9183b --- /dev/null +++ b/src/alembic/versions/b793b71fff3e_update_event_table_to_have_groups_.py @@ -0,0 +1,80 @@ +"""update event table to have groups, location, organizer and status + +Revision ID: b793b71fff3e +Revises: 81f44578da6d +Create Date: 2026-08-28 17:51:07.556757 + +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa +from sqlalchemy.dialects import postgresql + +# revision identifiers, used by Alembic. +revision: str = 'b793b71fff3e' +down_revision: Union[str, None] = '81f44578da6d' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('event_info', sa.Column('start_datetime', sa.DateTime(timezone=True), nullable=False)) + op.add_column('event_info', sa.Column('end_datetime', sa.DateTime(timezone=True), nullable=False)) + op.add_column('event_info', sa.Column('group_id', sa.Uuid(), nullable=True)) + op.add_column('event_info', sa.Column('location', sa.Text(), nullable=True)) + op.add_column('event_info', sa.Column('organizer', sa.Text(), nullable=True)) + op.add_column('event_info', sa.Column('status', sa.Enum('cancelled', 'scheduled', name='valid_status', native_enum=False, create_constraint=True), nullable=False)) + op.add_column('event_info', sa.Column('url', sa.Text(), nullable=True)) + op.add_column('event_info', sa.Column('image_id', sa.Integer(), nullable=True)) + op.alter_column('event_info', 'description', + existing_type=sa.TEXT(), + nullable=False) + op.alter_column('event_info', 'name', + existing_type=sa.VARCHAR(length=64), + type_=sa.Text(), + existing_nullable=False) + op.create_index(op.f('ix_event_info_group_id'), 'event_info', ['group_id'], unique=False) + op.create_foreign_key(op.f('fk_event_info_image_id_image_asset'), 'event_info', 'image_asset', ['image_id'], ['image_id']) + op.drop_constraint(op.f('ck_event_info_valid_frequency_value'), 'event_info', type_='check') + op.drop_constraint(op.f('ck_event_info_check_repeat_start_date_before_repeat_end_date'), 'event_info', type_='check') + op.drop_constraint(op.f('ck_event_info_check_start_time_before_end_time'), 'event_info', type_='check') + op.create_check_constraint(op.f('ck_event_info_check_start_datetime_before_end_datetime'), 'event_info', 'start_datetime <= end_datetime') + op.drop_column('event_info', 'end_time') + op.drop_column('event_info', 'frequency') + op.drop_column('event_info', 'repeat_end_date') + op.drop_column('event_info', 'repeat_start_date') + op.drop_column('event_info', 'start_time') + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('event_info', sa.Column('start_time', postgresql.TIMESTAMP(timezone=True), autoincrement=False, nullable=False)) + op.add_column('event_info', sa.Column('repeat_start_date', sa.DATE(), autoincrement=False, nullable=True)) + op.add_column('event_info', sa.Column('repeat_end_date', sa.DATE(), autoincrement=False, nullable=True)) + op.add_column('event_info', sa.Column('frequency', sa.VARCHAR(length=64), server_default=sa.text("'NONE'::character varying"), autoincrement=False, nullable=False)) + op.add_column('event_info', sa.Column('end_time', postgresql.TIMESTAMP(timezone=True), autoincrement=False, nullable=False)) + op.drop_constraint(op.f('ck_event_info_check_start_datetime_before_end_datetime'), 'event_info', type_='check') + op.create_check_constraint(op.f('ck_event_info_check_start_time_before_end_time'), 'event_info', 'start_time < end_time') + op.create_check_constraint(op.f('ck_event_info_valid_frequency_value'), 'event_info', "frequency IN ('NONE', 'DAILY', 'WEEKLY', 'MONTHLY', 'SEMESTERLY', 'YEARLY')") + op.create_check_constraint(op.f('ck_event_info_check_repeat_start_date_before_repeat_end_date'), 'event_info', 'repeat_start_date < repeat_end_date') + op.drop_constraint(op.f('fk_event_info_image_id_image_asset'), 'event_info', type_='foreignkey') + op.drop_index(op.f('ix_event_info_group_id'), table_name='event_info') + op.alter_column('event_info', 'name', + existing_type=sa.Text(), + type_=sa.VARCHAR(length=64), + existing_nullable=False) + op.alter_column('event_info', 'description', + existing_type=sa.TEXT(), + nullable=True) + op.drop_column('event_info', 'image_id') + op.drop_column('event_info', 'url') + op.drop_column('event_info', 'status') + op.drop_column('event_info', 'organizer') + op.drop_column('event_info', 'location') + op.drop_column('event_info', 'group_id') + op.drop_column('event_info', 'end_datetime') + op.drop_column('event_info', 'start_datetime') + # ### end Alembic commands ### diff --git a/src/database.py b/src/database.py index 8185a2f..dec8f64 100644 --- a/src/database.py +++ b/src/database.py @@ -1,6 +1,5 @@ import asyncio import contextlib -import os from collections.abc import AsyncGenerator from typing import Annotated, Any diff --git a/src/dependencies.py b/src/dependencies.py index 8c1323c..e611852 100644 --- a/src/dependencies.py +++ b/src/dependencies.py @@ -1,14 +1,20 @@ from typing import Annotated -from fastapi import Cookie, Depends, HTTPException, status +from fastapi import Cookie, Depends, HTTPException, Path, status import auth import auth.crud import database from auth.constants import COOKIE_SESSION_KEY -from config import settings from utils.permissions import is_user_election_admin, is_user_website_admin +# Dependency to ensure years in paths are valid +# Honestly don't know if this DB will be running past year 3000 +YearPath = Annotated[int, Path(ge=2000, le=3000)] + +# Dependency to ensure months in paths are valid +MonthPath = Annotated[int, Path(ge=1, le=12)] + async def optional_user( db_session: database.DBSession, session_id: Annotated[str | None, Cookie(alias=COOKIE_SESSION_KEY)] = None diff --git a/src/event/constants.py b/src/event/constants.py index 8764031..a7f4efe 100644 --- a/src/event/constants.py +++ b/src/event/constants.py @@ -1,10 +1,6 @@ from enum import StrEnum -class EventFrequencyEnum(StrEnum): - NONE = "NONE" - DAILY = "DAILY" - WEEKLY = "WEEKLY" - MONTHLY = "MONTHLY" - SEMESTERLY = "SEMESTERLY" - YEARLY = "YEARLY" +class EventStatusEnum(StrEnum): + CANCELLED = "cancelled" + SCHEDULED = "scheduled" diff --git a/src/event/crud.py b/src/event/crud.py index ece1124..e46c176 100644 --- a/src/event/crud.py +++ b/src/event/crud.py @@ -1,58 +1,93 @@ from collections.abc import Sequence -from datetime import date, datetime +from datetime import datetime +from uuid import UUID -from sqlalchemy import and_, delete, extract, or_, select +from sqlalchemy import delete, select from sqlalchemy.ext.asyncio import AsyncSession +from constants import TZ_INFO +from event.constants import EventStatusEnum from event.tables import EventDB -async def get_all_events(db_session: AsyncSession) -> Sequence[EventDB]: - events = (await db_session.scalars(select(EventDB))).all() - return events +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 -async def get_events_for_this_year( - db_session: AsyncSession, - year: int, -) -> Sequence[EventDB]: - events = ( - await db_session.scalars( - select(EventDB).where( - or_(extract("year", EventDB.start_time) == year, extract("year", EventDB.end_time) == year) - ) - ) - ).all() - return events + 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]: - events = ( - await db_session.scalars( - select(EventDB).where( - or_( - and_(extract("year", EventDB.start_time) == year, extract("month", EventDB.start_time) == month), - and_(extract("year", EventDB.end_time) == year, extract("month", EventDB.end_time) == month), - ) - ) - ) - ).all() + """ + 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) + ) + + if not include_cancelled: + query = query.where(EventDB.status != EventStatusEnum.CANCELLED) + + events = (await db_session.scalars(query)).all() + return events async def get_event_by_eid(db_session: AsyncSession, eid: int) -> EventDB | None: - return (await db_session.execute(select(EventDB).where(EventDB.eid == eid))).scalar_one_or_none() + return await db_session.get(EventDB, eid) + + +async def get_events_by_group_id(db_session: AsyncSession, group_id: UUID) -> Sequence[EventDB]: + query = select(EventDB).where(EventDB.group_id == group_id).order_by(EventDB.start_datetime, EventDB.end_datetime) + + result = await db_session.execute(query) + + return result.scalars().all() -async def create_event(db_session: AsyncSession, info: EventDB): +def create_event(db_session: AsyncSession, info: EventDB) -> None: db_session.add(info) -async def delete_event(db_session: AsyncSession, eid: int): - result = await db_session.execute(delete(EventDB).where(EventDB.eid == eid)) - # Return the number of rows affected - return result.rowcount +def create_bulk_event(db_session: AsyncSession, events: list[EventDB]) -> None: + db_session.add_all(events) + + +async def delete_event(db_session: AsyncSession, eid: int) -> int | None: + query = delete(EventDB).where(EventDB.eid == eid).returning(EventDB.eid) + return await db_session.scalar(query) + + +async def delete_group_events(db_session: AsyncSession, group_id: UUID) -> Sequence[int]: + query = delete(EventDB).where(EventDB.group_id == group_id).returning(EventDB.eid) + + return (await db_session.scalars(query)).all() diff --git a/src/event/models.py b/src/event/models.py index 40c8977..041d22a 100644 --- a/src/event/models.py +++ b/src/event/models.py @@ -1,54 +1,74 @@ -import datetime +from collections.abc import Sequence +from uuid import UUID -from pydantic import BaseModel, ConfigDict, model_validator +from pydantic import AwareDatetime, BaseModel, ConfigDict, model_validator -from event.constants import EventFrequencyEnum +from event.constants import EventStatusEnum class BaseEvent(BaseModel): name: str - start_time: datetime.datetime - end_time: datetime.datetime - description: str | None = None - frequency: EventFrequencyEnum | None = None - repeat_start_date: datetime.date | None = None - repeat_end_date: datetime.date | None = None + description: str + start_datetime: AwareDatetime + end_datetime: AwareDatetime + location: str | None = None + organizer: str | None = None + status: EventStatusEnum + url: str | None = None + image_id: int | None = None @model_validator(mode="after") def validate_time_range(self) -> "BaseEvent": - if self.start_time >= self.end_time: - raise ValueError("The event start must be before the event end") - - if self.repeat_start_date and self.repeat_end_date: - if self.repeat_start_date > self.repeat_end_date: - raise ValueError("The event repeat start date must be before the end date") - - if (self.repeat_start_date is None) != (self.repeat_end_date is None): - raise ValueError("The event must have both repeat start and repeat end or have neither.") - + if self.start_datetime > self.end_datetime: + raise ValueError("Event start times cannot be greater than end times") return self class Event(BaseEvent): model_config = ConfigDict(from_attributes=True) + eid: int + group_id: UUID | None = None class EventCreate(BaseEvent): pass +class GroupEvent(BaseModel): + group_id: UUID + events: list[Event] + + class EventUpdate(BaseModel): + """ + Partial patch payload for PATCH-style updates. Deliberately does NOT + inherit from BaseEvent. Inherting from BaseEvent would also inherit the validation + which will caude bugs for None values, to avoid that we would need a db call. Hence, + the validation is done in the routing layer. Every field is optional here since a client + only sends the fields they want to change. + + Group IDs cannot be modified once created, you have to create a new group of events. + """ + model_config = ConfigDict(extra="forbid") name: str | None = None - start_time: datetime.datetime | None = None - end_time: datetime.datetime | None = None description: str | None = None - frequency: EventFrequencyEnum | None = None - repeat_start_date: datetime.date | None = None - repeat_end_date: datetime.date | None = None + start_datetime: AwareDatetime | None = None + end_datetime: AwareDatetime | None = None + location: str | None = None + organizer: str | None = None + status: EventStatusEnum | None = None + url: str | None = None + image_id: int | None = None class EventDelete(BaseModel): result: bool eid: int + + +class GroupEventDeleteResponse(BaseModel): + result: bool + group_id: UUID + deleted_eids: Sequence[int] diff --git a/src/event/tables.py b/src/event/tables.py index c176fa7..3644b44 100644 --- a/src/event/tables.py +++ b/src/event/tables.py @@ -1,26 +1,37 @@ -from datetime import date, datetime +from datetime import datetime +from uuid import UUID -from sqlalchemy import CheckConstraint, Date, DateTime, Integer, String, Text, text +from sqlalchemy import CheckConstraint, DateTime, Enum, ForeignKey, Integer, Text, Uuid from sqlalchemy.orm import Mapped, mapped_column from database import Base -from event.constants import EventFrequencyEnum +from event.constants import EventStatusEnum class EventDB(Base): __tablename__ = "event_info" eid: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) - description: Mapped[str] = mapped_column(Text, nullable=True) - name: Mapped[str] = mapped_column(String(64)) - start_time: Mapped[datetime] = mapped_column(DateTime(timezone=True)) - end_time: Mapped[datetime] = mapped_column(DateTime(timezone=True)) - frequency: Mapped[EventFrequencyEnum] = mapped_column(String(64), server_default=text("'NONE'")) - repeat_start_date: Mapped[date] = mapped_column(Date, nullable=True) - repeat_end_date: Mapped[date] = mapped_column(Date, nullable=True) + description: Mapped[str] = mapped_column(Text) + name: Mapped[str] = mapped_column(Text) + start_datetime: Mapped[datetime] = mapped_column(DateTime(timezone=True)) + end_datetime: Mapped[datetime] = mapped_column(DateTime(timezone=True)) + group_id: Mapped[UUID | None] = mapped_column(Uuid, nullable=True, index=True) + location: Mapped[str | None] = mapped_column(Text, nullable=True) + organizer: Mapped[str | None] = mapped_column(Text, nullable=True) + status: Mapped[EventStatusEnum] = mapped_column( + Enum( + EventStatusEnum, + native_enum=False, + create_constraint=True, + validate_strings=True, + values_callable=lambda enum: [status.value for status in enum], + name="valid_status", + ) + ) + url: Mapped[str | None] = mapped_column(Text, nullable=True) + image_id: Mapped[int | None] = mapped_column(Integer, ForeignKey("image_asset.image_id"), nullable=True) __table_args__ = ( - CheckConstraint("start_time < end_time", name="check_start_time_before_end_time"), - CheckConstraint("repeat_start_date < repeat_end_date", name="check_repeat_start_date_before_repeat_end_date"), - CheckConstraint(frequency.in_([e.value for e in EventFrequencyEnum]), name="valid_frequency_value"), + CheckConstraint("start_datetime <= end_datetime", name="check_start_datetime_before_end_datetime"), ) diff --git a/src/event/urls.py b/src/event/urls.py index 9821490..08f78a5 100644 --- a/src/event/urls.py +++ b/src/event/urls.py @@ -1,11 +1,14 @@ -from fastapi import APIRouter, Depends, HTTPException, status +import uuid +from typing import Annotated + +from fastapi import APIRouter, Body, Depends, HTTPException, status from fastapi.encoders import jsonable_encoder from pydantic import ValidationError import database import event.crud -from dependencies import perm_admin -from event.models import Event, EventCreate, EventDelete, EventUpdate +from dependencies import MonthPath, YearPath, perm_admin +from event.models import Event, EventCreate, EventDelete, EventUpdate, GroupEvent, GroupEventDeleteResponse from event.tables import EventDB from utils.shared_models import DetailModel @@ -21,37 +24,22 @@ response_model=list[Event], operation_id="get_all_events", ) -async def get_all_events( - db_session: database.DBSession, -): - events_list = await event.crud.get_all_events(db_session) - - return events_list - - -@router.get( - "/{year}", - description="Get events that start OR end in this year", - response_model=list[Event], - operation_id="get_events_for_this_year", -) -async def get_events_for_this_year( - db_session: database.DBSession, - year: int, -): - events_list = await event.crud.get_events_for_this_year(db_session, year) +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 start OR end in the given year and 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: int, month: int): - events_list = await event.crud.get_events_for_this_year_month(db_session, 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) return events_list @@ -69,7 +57,7 @@ async def get_events_for_this_year_month(db_session: database.DBSession, year: i ) async def create_event(db_session: database.DBSession, body: EventCreate): new_event = EventDB(**body.model_dump()) - await event.crud.create_event( + event.crud.create_event( db_session, new_event, ) @@ -80,6 +68,56 @@ async def create_event(db_session: database.DBSession, body: EventCreate): return new_event +@router.post( + "/group", + description="Creates one or more events under the same group key.", + response_model=GroupEvent, + status_code=status.HTTP_201_CREATED, + operation_id="create_group_event", + dependencies=[Depends(perm_admin)], +) +async def create_group_events(db_session: database.DBSession, body: Annotated[list[EventCreate], Body(min_length=1)]): + g_id = uuid.uuid4() + # Create new EventDBs by injecting group_id + new_event_list = [EventDB(**e.model_dump(exclude={"group_id"}), group_id=g_id) for e in body] + + event.crud.create_bulk_event( + db_session, + new_event_list, + ) + + await db_session.flush() + + response = GroupEvent(group_id=g_id, events=[Event.model_validate(db_event) for db_event in new_event_list]) + + await db_session.commit() + return response + + +@router.post( + "/group/{group_id}", + description="Adds an event to an existing group", + response_model=Event, + status_code=status.HTTP_201_CREATED, + responses={404: {"description": "Group doesn't exist."}}, + operation_id="add_event_to_group", + dependencies=[Depends(perm_admin)], +) +async def add_event_to_group(db_session: database.DBSession, group_id: uuid.UUID, new_event: EventCreate): + + # NOTE: There is a race condition here: + # If the group is deleted after the existence check this will still insert the new event. + group = await event.crud.get_events_by_group_id(db_session, group_id) + if not group: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Group doesn't exist.") + + create_event = EventDB(**new_event.model_dump(), group_id=group_id) + event.crud.create_event(db_session, create_event) + await db_session.commit() + await db_session.refresh(create_event) + return create_event + + @router.patch( "/{eid}", description="Update an Event detail", @@ -93,12 +131,11 @@ async def update_event(db_session: database.DBSession, eid: int, body: EventUpda if db_event is None: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Event doesn't exist.") - db_data = Event.model_validate(db_event).model_dump() + db_data = Event.model_validate(db_event) patch_data = body.model_dump(exclude_unset=True) - merged_data = {**db_data, **patch_data} try: - Event.model_validate(merged_data) + updated = Event.model_validate(db_data.model_dump() | patch_data) except ValidationError as e: raise HTTPException( status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=jsonable_encoder(e.errors()) @@ -108,9 +145,8 @@ async def update_event(db_session: database.DBSession, eid: int, body: EventUpda setattr(db_event, key, value) await db_session.commit() - await db_session.refresh(db_event) - return db_event + return updated @router.delete( @@ -122,10 +158,28 @@ async def update_event(db_session: database.DBSession, eid: int, body: EventUpda dependencies=[Depends(perm_admin)], ) async def delete_event(db_session: database.DBSession, eid: int): - rows_deleted = await event.crud.delete_event(db_session, eid) + deleted_eid = await event.crud.delete_event(db_session, eid) + + if deleted_eid is None: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Event doesn't exist.") + + await db_session.commit() + return EventDelete(result=True, eid=deleted_eid) + + +@router.delete( + "/group/{group_id}", + description="Delete event(s) with the given group_id", + response_model=GroupEventDeleteResponse, + responses={404: {"description": "Event doesn't exist."}}, + operation_id="delete_group_event", + dependencies=[Depends(perm_admin)], +) +async def delete_group_event(db_session: database.DBSession, group_id: uuid.UUID): + deleted_eids = await event.crud.delete_group_events(db_session, group_id) - if rows_deleted == 0: + if not deleted_eids: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Event doesn't exist.") await db_session.commit() - return EventDelete(result=True, eid=eid) + return GroupEventDeleteResponse(result=True, group_id=group_id, deleted_eids=deleted_eids)