From b26c4680981fd87452e7a63cd4983820b603bd37 Mon Sep 17 00:00:00 2001 From: mai Date: Fri, 14 Aug 2026 22:15:52 -0400 Subject: [PATCH 1/3] feat: media and media variants tables --- .../versions/4315ee2f32f5_media_service.py | 59 +++++++++++++++++++ 1 file changed, 59 insertions(+) create mode 100644 backend/alembic/versions/4315ee2f32f5_media_service.py diff --git a/backend/alembic/versions/4315ee2f32f5_media_service.py b/backend/alembic/versions/4315ee2f32f5_media_service.py new file mode 100644 index 0000000..19f4ded --- /dev/null +++ b/backend/alembic/versions/4315ee2f32f5_media_service.py @@ -0,0 +1,59 @@ +"""media service + +Revision ID: 4315ee2f32f5 +Revises: ce8c0c720681 +Create Date: 2026-08-14 21:54:24.487092 +""" + +from collections.abc import Sequence + +from alembic import op + +revision: str = "4315ee2f32f5" +down_revision: str | None = "ce8c0c720681" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + +UPGRADE = """ +CREATE TABLE media ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + owner_id UUID NOT NULL, + s3_key TEXT NOT NULL UNIQUE, + original_filename TEXT NOT NULL, + media_type TEXT NOT NULL, + mime_type TEXT NOT NULL, + size_bytes BIGINT NOT NULL, + status TEXT NOT NULL, + metadata JSONB, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE TABLE media_variants ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + media_id UUID NOT NULL REFERENCES media(id) ON DELETE CASCADE, + variant_type TEXT NOT NULL, + s3_key TEXT NOT NULL UNIQUE, + mime_type TEXT NOT NULL, + size_bytes BIGINT NOT NULL, + metadata JSONB, + status TEXT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + + UNIQUE (media_id, variant_type) +); +""" + +DOWNGRADE = """ +DROP TABLE IF EXISTS media; +DROP TABLE IF EXISTS media_variants; +""" + + +def upgrade() -> None: + op.execute(UPGRADE) + + +def downgrade() -> None: + op.execute(DOWNGRADE) From d0fbd7efbecb3e924f8a199eff28d9244ded79e7 Mon Sep 17 00:00:00 2001 From: mai Date: Fri, 14 Aug 2026 22:51:33 -0400 Subject: [PATCH 2/3] feat: media service api --- .../versions/4315ee2f32f5_media_service.py | 15 +- backend/src/admin/api/dependencies.py | 12 + backend/src/admin/api/router.py | 2 + backend/src/admin/api/v1/media.py | 77 ++++++ backend/src/admin/core/storage.py | 44 +++- backend/src/admin/domain/enums.py | 6 + backend/src/admin/domain/media.py | 26 ++ backend/src/admin/domain/permissions.py | 4 + backend/src/admin/repositories/__init__.py | 2 + backend/src/admin/repositories/audit.py | 25 ++ backend/src/admin/repositories/media.py | 117 +++++++++ backend/src/admin/schemas/media.py | 86 +++++++ backend/src/admin/services/access.py | 38 +-- backend/src/admin/services/guards.py | 28 ++- backend/src/admin/services/media.py | 223 ++++++++++++++++++ 15 files changed, 675 insertions(+), 30 deletions(-) create mode 100644 backend/src/admin/api/v1/media.py create mode 100644 backend/src/admin/domain/media.py create mode 100644 backend/src/admin/repositories/media.py create mode 100644 backend/src/admin/schemas/media.py create mode 100644 backend/src/admin/services/media.py diff --git a/backend/alembic/versions/4315ee2f32f5_media_service.py b/backend/alembic/versions/4315ee2f32f5_media_service.py index 19f4ded..0392667 100644 --- a/backend/alembic/versions/4315ee2f32f5_media_service.py +++ b/backend/alembic/versions/4315ee2f32f5_media_service.py @@ -20,34 +20,41 @@ owner_id UUID NOT NULL, s3_key TEXT NOT NULL UNIQUE, original_filename TEXT NOT NULL, - media_type TEXT NOT NULL, mime_type TEXT NOT NULL, size_bytes BIGINT NOT NULL, + purpose TEXT NOT NULL, + visibility TEXT NOT NULL CHECK (visibility IN ('public', 'private')), status TEXT NOT NULL, metadata JSONB, created_at TIMESTAMPTZ NOT NULL DEFAULT now(), - updated_at TIMESTAMPTZ NOT NULL DEFAULT now() + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + + UNIQUE (id, visibility) ); +CREATE INDEX ix_media_owner ON media (owner_id, created_at DESC); + CREATE TABLE media_variants ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), - media_id UUID NOT NULL REFERENCES media(id) ON DELETE CASCADE, + media_id UUID NOT NULL, variant_type TEXT NOT NULL, s3_key TEXT NOT NULL UNIQUE, mime_type TEXT NOT NULL, size_bytes BIGINT NOT NULL, + visibility TEXT NOT NULL, metadata JSONB, status TEXT NOT NULL, created_at TIMESTAMPTZ NOT NULL DEFAULT now(), updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + FOREIGN KEY (media_id, visibility) REFERENCES media (id, visibility) ON DELETE CASCADE, UNIQUE (media_id, variant_type) ); """ DOWNGRADE = """ -DROP TABLE IF EXISTS media; DROP TABLE IF EXISTS media_variants; +DROP TABLE IF EXISTS media; """ diff --git a/backend/src/admin/api/dependencies.py b/backend/src/admin/api/dependencies.py index 1665166..14a90ba 100644 --- a/backend/src/admin/api/dependencies.py +++ b/backend/src/admin/api/dependencies.py @@ -22,6 +22,7 @@ AccessRequestRepository, AuditRepository, InvitationRepository, + MediaRepository, RoleRepository, UserRepository, ) @@ -30,6 +31,7 @@ from admin.services.access import AccessService from admin.services.access_request import AccessRequestService from admin.services.invitation import InvitationService +from admin.services.media import MediaService from admin.services.member import MemberService bearer_scheme = HTTPBearer(auto_error=True) @@ -91,11 +93,16 @@ def get_audit_repository(connection: Connection) -> AuditRepository: return AuditRepository(connection) +def get_media_repository(connection: Connection) -> MediaRepository: + return MediaRepository(connection) + + Users = Annotated[UserRepository, Depends(get_user_repository)] Roles = Annotated[RoleRepository, Depends(get_role_repository)] Invitations = Annotated[InvitationRepository, Depends(get_invitation_repository)] AccessRequests = Annotated[AccessRequestRepository, Depends(get_access_request_repository)] Audit = Annotated[AuditRepository, Depends(get_audit_repository)] +Media = Annotated[MediaRepository, Depends(get_media_repository)] def get_access_service( @@ -138,10 +145,15 @@ def get_access_request_service( ) +def get_media_service(media: Media, storage: Storage, audit: Audit) -> MediaService: + return MediaService(media=media, storage=storage, audit=audit) + + AccessServiceDep = Annotated[AccessService, Depends(get_access_service)] MemberServiceDep = Annotated[MemberService, Depends(get_member_service)] InvitationServiceDep = Annotated[InvitationService, Depends(get_invitation_service)] AccessRequestServiceDep = Annotated[AccessRequestService, Depends(get_access_request_service)] +MediaServiceDep = Annotated[MediaService, Depends(get_media_service)] async def get_identity( diff --git a/backend/src/admin/api/router.py b/backend/src/admin/api/router.py index 74ae0ad..dbc7a5d 100644 --- a/backend/src/admin/api/router.py +++ b/backend/src/admin/api/router.py @@ -4,6 +4,7 @@ access_requests, health, invitations, + media, members, roles, session, @@ -15,6 +16,7 @@ api_router.include_router(invitations.router) api_router.include_router(access_requests.router) api_router.include_router(roles.router) +api_router.include_router(media.router) root_router = APIRouter() root_router.include_router(health.router) diff --git a/backend/src/admin/api/v1/media.py b/backend/src/admin/api/v1/media.py new file mode 100644 index 0000000..8adecf9 --- /dev/null +++ b/backend/src/admin/api/v1/media.py @@ -0,0 +1,77 @@ +import uuid +from typing import Annotated + +from fastapi import APIRouter, Query, status + +from admin.api.dependencies import CurrentUser, MediaServiceDep +from admin.schemas.media import ( + MediaPresetRead, + MediaRead, + MediaUploadRequest, + MediaUploadTicket, +) + +router = APIRouter(prefix="/media", tags=["media"]) + + +@router.get("/presets", response_model=list[MediaPresetRead]) +async def list_presets(service: MediaServiceDep, _: CurrentUser) -> list[MediaPresetRead]: + return service.presets() + + +@router.post( + "/upload-tickets", + response_model=list[MediaUploadTicket], + status_code=status.HTTP_201_CREATED, +) +async def create_upload_tickets( + payload: MediaUploadRequest, service: MediaServiceDep, context: CurrentUser +) -> list[MediaUploadTicket]: + return await service.create_upload_tickets(actor=context.user, payload=payload) + + +@router.get("", response_model=list[MediaRead]) +async def list_media( + service: MediaServiceDep, + context: CurrentUser, + ids: Annotated[list[uuid.UUID], Query()], +) -> list[MediaRead]: + return await service.get_files( + actor=context.user, actor_permissions=context.permissions, media_ids=ids + ) + + +@router.delete("", status_code=status.HTTP_204_NO_CONTENT) +async def delete_media( + service: MediaServiceDep, + context: CurrentUser, + ids: Annotated[list[uuid.UUID], Query()], +) -> None: + await service.delete_files( + actor=context.user, actor_permissions=context.permissions, media_ids=ids + ) + + +@router.post("/{media_id}/complete", response_model=MediaRead) +async def complete_upload( + media_id: uuid.UUID, service: MediaServiceDep, context: CurrentUser +) -> MediaRead: + return await service.complete_upload(actor=context.user, media_id=media_id) + + +@router.get("/{media_id}", response_model=MediaRead) +async def get_media( + media_id: uuid.UUID, service: MediaServiceDep, context: CurrentUser +) -> MediaRead: + return await service.get_file( + actor=context.user, actor_permissions=context.permissions, media_id=media_id + ) + + +@router.delete("/{media_id}", status_code=status.HTTP_204_NO_CONTENT) +async def delete_single_media( + media_id: uuid.UUID, service: MediaServiceDep, context: CurrentUser +) -> None: + await service.delete_file( + actor=context.user, actor_permissions=context.permissions, media_id=media_id + ) diff --git a/backend/src/admin/core/storage.py b/backend/src/admin/core/storage.py index 46233d0..5caf336 100644 --- a/backend/src/admin/core/storage.py +++ b/backend/src/admin/core/storage.py @@ -1,8 +1,9 @@ import asyncio import uuid from dataclasses import dataclass -from datetime import UTC, datetime +from datetime import UTC, datetime, timedelta from enum import StrEnum +from itertools import batched from typing import Any import boto3 @@ -24,10 +25,10 @@ class MediaVisibility(StrEnum): PRESIGN_CACHE_RATIO = 0.8 PUBLIC_PREFIX = "public" PRIVATE_PREFIX = "private" +S3_DELETE_BATCH_SIZE = 1000 IMMUTABLE_CACHE_CONTROL = "public, max-age=31536000, immutable" PRIVATE_CACHE_CONTROL = "private, max-age=0, no-store" -STRING_ADAPTER = TypeAdapter(str) def cache_control_for(visibility: MediaVisibility) -> str: @@ -50,6 +51,15 @@ class ObjectMetadata: content_type: str +@dataclass(frozen=True, slots=True) +class MediaUrl: + url: str + expires_at: datetime | None = None + + +MEDIA_URL_ADAPTER = TypeAdapter(MediaUrl) + + def build_object_key(*, visibility: MediaVisibility, filename: str) -> str: prefix = PUBLIC_PREFIX if visibility is MediaVisibility.PUBLIC else PRIVATE_PREFIX stamp = datetime.now(UTC) @@ -70,6 +80,10 @@ def __init__(self, config: StorageConfig, cache: Cache) -> None: config=Config(signature_version="s3v4", s3={"addressing_style": "path"}), ) + @property + def max_upload_bytes(self) -> int: + return self._config.max_upload_bytes + def create_upload_ticket( self, *, @@ -103,22 +117,28 @@ def create_upload_ticket( def public_url(self, key: str) -> str: return self._config.public_url_for(key) - async def signed_download_url(self, key: str) -> str: + async def url_for(self, *, key: str, visibility: MediaVisibility) -> MediaUrl: + if visibility is MediaVisibility.PUBLIC: + return MediaUrl(url=self.public_url(key)) + return await self.signed_download_url(key) + + async def signed_download_url(self, key: str) -> MediaUrl: ttl = self._config.download_url_ttl_seconds - async def generate() -> str: - return await asyncio.to_thread( + async def generate() -> MediaUrl: + url = await asyncio.to_thread( self._client.generate_presigned_url, "get_object", Params={"Bucket": self._config.bucket_name, "Key": key}, ExpiresIn=ttl, ) + return MediaUrl(url=url, expires_at=datetime.now(UTC) + timedelta(seconds=ttl)) return await self._cache.fetch( CacheNamespace.MEDIA, f"presign:get:{key}", generate, - adapter=STRING_ADAPTER, + adapter=MEDIA_URL_ADAPTER, ttl=ttl * PRESIGN_CACHE_RATIO, ) @@ -140,3 +160,15 @@ async def delete(self, key: str) -> None: self._client.delete_object, Bucket=self._config.bucket_name, Key=key ) await self._cache.bump(CacheNamespace.MEDIA) + + async def delete_many(self, keys: list[str]) -> None: + if not keys: + return + + for batch in batched(keys, S3_DELETE_BATCH_SIZE): + await asyncio.to_thread( + self._client.delete_objects, + Bucket=self._config.bucket_name, + Delete={"Objects": [{"Key": key} for key in batch], "Quiet": True}, + ) + await self._cache.bump(CacheNamespace.MEDIA) diff --git a/backend/src/admin/domain/enums.py b/backend/src/admin/domain/enums.py index 9227699..2a4bbc6 100644 --- a/backend/src/admin/domain/enums.py +++ b/backend/src/admin/domain/enums.py @@ -28,6 +28,10 @@ class AccessRequestStatus(StrEnum): DENIED = "denied" +class MediaPurpose(StrEnum): + AVATAR = "avatar" + + class AuditAction(StrEnum): USER_PROVISIONED = "user.provisioned" USER_SUSPENDED = "user.suspended" @@ -40,3 +44,5 @@ class AuditAction(StrEnum): ACCESS_REQUEST_CREATED = "access_request.created" ACCESS_REQUEST_APPROVED = "access_request.approved" ACCESS_REQUEST_DENIED = "access_request.denied" + MEDIA_UPLOADED = "media.uploaded" + MEDIA_DELETED = "media.deleted" diff --git a/backend/src/admin/domain/media.py b/backend/src/admin/domain/media.py new file mode 100644 index 0000000..87286ba --- /dev/null +++ b/backend/src/admin/domain/media.py @@ -0,0 +1,26 @@ +from dataclasses import dataclass + +from admin.core.storage import MediaVisibility +from admin.domain.enums import MediaPurpose + + +@dataclass(frozen=True, slots=True) +class MediaPreset: + max_edge: int + max_bytes: int + visibility: MediaVisibility + mime_types: frozenset[str] + + +PRESETS: dict[MediaPurpose, MediaPreset] = { + MediaPurpose.AVATAR: MediaPreset( + max_edge=512, + max_bytes=1_048_576, + visibility=MediaVisibility.PUBLIC, + mime_types=frozenset({"image/jpeg", "image/png", "image/webp"}), + ), +} + + +def preset_for(purpose: MediaPurpose) -> MediaPreset: + return PRESETS[purpose] diff --git a/backend/src/admin/domain/permissions.py b/backend/src/admin/domain/permissions.py index b964614..33aef2e 100644 --- a/backend/src/admin/domain/permissions.py +++ b/backend/src/admin/domain/permissions.py @@ -12,6 +12,8 @@ class Permission(StrEnum): ACCESS_REQUESTS_READ = "core.access_requests.read" ACCESS_REQUESTS_REVIEW = "core.access_requests.review" AUDIT_READ = "core.audit.read" + MEDIA_READ = "core.media.read" + MEDIA_DELETE = "core.media.delete" @property def description(self) -> str: @@ -28,6 +30,8 @@ def description(self) -> str: Permission.ACCESS_REQUESTS_READ: "View pending access requests", Permission.ACCESS_REQUESTS_REVIEW: "Approve or deny access requests", Permission.AUDIT_READ: "Read the audit log", + Permission.MEDIA_READ: "View private files uploaded by other members", + Permission.MEDIA_DELETE: "Delete files uploaded by other members", } diff --git a/backend/src/admin/repositories/__init__.py b/backend/src/admin/repositories/__init__.py index 7fdc534..a908125 100644 --- a/backend/src/admin/repositories/__init__.py +++ b/backend/src/admin/repositories/__init__.py @@ -1,6 +1,7 @@ from admin.repositories.access_request import AccessRequestRepository from admin.repositories.audit import AuditRepository from admin.repositories.invitation import InvitationRepository +from admin.repositories.media import MediaRepository from admin.repositories.role import RoleRepository from admin.repositories.user import UserRepository @@ -8,6 +9,7 @@ "AccessRequestRepository", "AuditRepository", "InvitationRepository", + "MediaRepository", "RoleRepository", "UserRepository", ] diff --git a/backend/src/admin/repositories/audit.py b/backend/src/admin/repositories/audit.py index f6aface..8ab06e9 100644 --- a/backend/src/admin/repositories/audit.py +++ b/backend/src/admin/repositories/audit.py @@ -29,6 +29,31 @@ async def record(self, entry: AuditEntry) -> None: entry.after, ) + async def record_many(self, entries: list[AuditEntry]) -> None: + if not entries: + return + + await self.connection.executemany( + """ + INSERT INTO audit_logs + (actor_id, actor_email, action, resource_type, resource_id, + before, after) + VALUES ($1, $2, $3, $4, $5, $6, $7) + """, + [ + ( + entry.actor_id, + entry.actor_email, + entry.action.value, + entry.resource_type, + entry.resource_id, + entry.before, + entry.after, + ) + for entry in entries + ], + ) + async def list_entries( self, params: CursorParams, diff --git a/backend/src/admin/repositories/media.py b/backend/src/admin/repositories/media.py new file mode 100644 index 0000000..57f0b1c --- /dev/null +++ b/backend/src/admin/repositories/media.py @@ -0,0 +1,117 @@ +import uuid + +from admin.repositories.base import Repository, required_row +from admin.schemas.media import MediaCreate, MediaRecord, MediaUploadingStatus + +MEDIA_SELECT = """ +SELECT m.id, + m.owner_id, + m.s3_key, + m.original_filename, + m.mime_type, + m.size_bytes, + m.purpose, + m.visibility, + m.status, + m.created_at, + m.updated_at +FROM media m +""" + + +class MediaRepository(Repository): + async def insert(self, payload: MediaCreate) -> MediaRecord: + row = required_row( + await self.connection.fetchrow( + """ + INSERT INTO media ( + owner_id, s3_key, original_filename, mime_type, + size_bytes, purpose, visibility, status + ) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8) + RETURNING id, owner_id, s3_key, original_filename, mime_type, + size_bytes, purpose, visibility, status, created_at, updated_at + """, + payload.owner_id, + payload.s3_key, + payload.original_filename, + payload.mime_type, + payload.size_bytes, + payload.purpose.value, + payload.visibility.value, + payload.status.value, + ) + ) + return MediaRecord.from_row(row) + + async def update_status( + self, + media_id: uuid.UUID, + status: MediaUploadingStatus, + *, + size_bytes: int | None = None, + ) -> MediaRecord | None: + row = await self.connection.fetchrow( + """ + UPDATE media + SET status = $2, + size_bytes = COALESCE($3, size_bytes), + updated_at = now() + WHERE id = $1 + RETURNING id, owner_id, s3_key, original_filename, mime_type, + size_bytes, purpose, visibility, status, created_at, updated_at + """, + media_id, + status.value, + size_bytes, + ) + return MediaRecord.from_optional_row(row) + + async def get_by_id(self, media_id: uuid.UUID) -> MediaRecord | None: + row = await self.connection.fetchrow(f"{MEDIA_SELECT} WHERE m.id = $1", media_id) + return MediaRecord.from_optional_row(row) + + async def get_by_ids(self, media_ids: list[uuid.UUID]) -> list[MediaRecord]: + if not media_ids: + return [] + + rows = await self.connection.fetch( + f"{MEDIA_SELECT} WHERE m.id = ANY($1) ORDER BY m.created_at DESC", media_ids + ) + return MediaRecord.from_rows(rows) + + async def delete_by_id(self, media_id: uuid.UUID) -> MediaRecord | None: + row = await self.connection.fetchrow( + """ + DELETE FROM media + WHERE id = $1 + RETURNING id, owner_id, s3_key, original_filename, mime_type, + size_bytes, purpose, visibility, status, created_at, updated_at + """, + media_id, + ) + return MediaRecord.from_optional_row(row) + + async def delete_by_ids(self, media_ids: list[uuid.UUID]) -> list[MediaRecord]: + if not media_ids: + return [] + + rows = await self.connection.fetch( + """ + DELETE FROM media + WHERE id = ANY($1) + RETURNING id, owner_id, s3_key, original_filename, mime_type, + size_bytes, purpose, visibility, status, created_at, updated_at + """, + media_ids, + ) + return MediaRecord.from_rows(rows) + + async def variant_keys_for(self, media_ids: list[uuid.UUID]) -> list[str]: + if not media_ids: + return [] + + rows = await self.connection.fetch( + "SELECT s3_key FROM media_variants WHERE media_id = ANY($1)", media_ids + ) + return [row["s3_key"] for row in rows] diff --git a/backend/src/admin/schemas/media.py b/backend/src/admin/schemas/media.py new file mode 100644 index 0000000..a8ce7fc --- /dev/null +++ b/backend/src/admin/schemas/media.py @@ -0,0 +1,86 @@ +import uuid +from datetime import datetime +from enum import StrEnum + +from pydantic import Field + +from admin.core.storage import MediaVisibility +from admin.domain.enums import MediaPurpose +from admin.schemas.base import ReadDTO, RequestDTO + + +class MediaUploadingStatus(StrEnum): + UPLOADING = "uploading" + COMPLETED = "completed" + FAILED = "failed" + + +class MediaCreate(RequestDTO): + owner_id: uuid.UUID + s3_key: str + original_filename: str + mime_type: str + size_bytes: int + purpose: MediaPurpose + visibility: MediaVisibility + status: MediaUploadingStatus + + +class MediaUpdate(RequestDTO): + id: uuid.UUID + s3_key: str + owner_id: uuid.UUID + status: MediaUploadingStatus + + +class MediaRecord(ReadDTO): + id: uuid.UUID + owner_id: uuid.UUID + s3_key: str + original_filename: str + mime_type: str + size_bytes: int + purpose: MediaPurpose + visibility: MediaVisibility + status: MediaUploadingStatus + created_at: datetime + updated_at: datetime + + +class MediaRead(ReadDTO): + id: uuid.UUID + url: str + url_expires_at: datetime | None = None + original_filename: str + mime_type: str + size_bytes: int + purpose: MediaPurpose + visibility: MediaVisibility + status: MediaUploadingStatus + created_at: datetime + + +class MediaUploadItem(RequestDTO): + filename: str = Field(min_length=1, max_length=255) + mime_type: str + size_bytes: int = Field(gt=0) + + +class MediaUploadRequest(RequestDTO): + purpose: MediaPurpose + files: list[MediaUploadItem] = Field(min_length=1, max_length=20) + + +class MediaPresetRead(ReadDTO): + purpose: MediaPurpose + max_edge: int + max_bytes: int + mime_types: list[str] + + +class MediaUploadTicket(ReadDTO): + media_id: uuid.UUID + url: str + fields: dict[str, str] + s3_key: str + expires_in: int diff --git a/backend/src/admin/services/access.py b/backend/src/admin/services/access.py index f2accd3..1c8e037 100644 --- a/backend/src/admin/services/access.py +++ b/backend/src/admin/services/access.py @@ -99,24 +99,24 @@ async def _accept_invitation(self, identity: Identity, invitation: InvitationRea ) await self._invitations.mark_accepted(invitation.id) - await self._audit.record( - AuditEntry( - actor_id=user.id, - actor_email=user.email, - action=AuditAction.INVITATION_ACCEPTED, - resource_type="invitation", - resource_id=str(invitation.id), - after={"user_id": str(user.id), "role_key": invitation.role.key}, - ) - ) - await self._audit.record( - AuditEntry( - actor_id=user.id, - actor_email=user.email, - action=AuditAction.USER_PROVISIONED, - resource_type="user", - resource_id=str(user.id), - after={"email": user.email, "source": "invitation"}, - ) + await self._audit.record_many( + [ + AuditEntry( + actor_id=user.id, + actor_email=user.email, + action=AuditAction.INVITATION_ACCEPTED, + resource_type="invitation", + resource_id=str(invitation.id), + after={"user_id": str(user.id), "role_key": invitation.role.key}, + ), + AuditEntry( + actor_id=user.id, + actor_email=user.email, + action=AuditAction.USER_PROVISIONED, + resource_type="user", + resource_id=str(user.id), + after={"email": user.email, "source": "invitation"}, + ), + ] ) return user diff --git a/backend/src/admin/services/guards.py b/backend/src/admin/services/guards.py index 5fb7e53..41a4d87 100644 --- a/backend/src/admin/services/guards.py +++ b/backend/src/admin/services/guards.py @@ -1,8 +1,11 @@ -from admin.core.errors import LastOwnerError, PrivilegeEscalationError +from admin.core.errors import LastOwnerError, PermissionDeniedError, PrivilegeEscalationError +from admin.core.storage import MediaVisibility from admin.domain.access import PermissionSet from admin.domain.permissions import Permission, SystemRole from admin.repositories.user import UserRepository +from admin.schemas.media import MediaRecord from admin.schemas.role import RoleRead +from admin.schemas.user import UserRead def ensure_can_delegate(actor: PermissionSet, role: RoleRead) -> None: @@ -20,3 +23,26 @@ async def ensure_not_last_owner(users: UserRepository, role_key: str) -> None: return if await users.count_active_holders_of_role(SystemRole.OWNER.value) <= 1: raise LastOwnerError("the workspace must keep at least one active owner") + + +def ensure_owns(record: MediaRecord, *, actor: UserRead) -> None: + if record.owner_id != actor.id: + raise PermissionDeniedError("you can only complete your own uploads") + + +def ensure_can_read(record: MediaRecord, *, actor: UserRead, permissions: PermissionSet) -> None: + if record.visibility is MediaVisibility.PUBLIC: + return + if record.owner_id == actor.id: + return + if permissions.allows(Permission.MEDIA_READ): + return + raise PermissionDeniedError("you do not have access to this file") + + +def ensure_can_delete(record: MediaRecord, *, actor: UserRead, permissions: PermissionSet) -> None: + if record.owner_id == actor.id: + return + if permissions.allows(Permission.MEDIA_DELETE): + return + raise PermissionDeniedError("you can only delete your own files") diff --git a/backend/src/admin/services/media.py b/backend/src/admin/services/media.py new file mode 100644 index 0000000..81f7e1d --- /dev/null +++ b/backend/src/admin/services/media.py @@ -0,0 +1,223 @@ +import asyncio +import uuid + +from admin.core.errors import NotFoundError, ValidationError +from admin.core.logging import get_logger +from admin.core.storage import MediaUrl, S3Storage, build_object_key +from admin.domain.access import PermissionSet +from admin.domain.enums import AuditAction +from admin.domain.media import PRESETS, preset_for +from admin.repositories.audit import AuditRepository +from admin.repositories.media import MediaRepository +from admin.schemas.audit import AuditEntry +from admin.schemas.media import ( + MediaCreate, + MediaPresetRead, + MediaRead, + MediaRecord, + MediaUploadingStatus, + MediaUploadRequest, + MediaUploadTicket, +) +from admin.schemas.user import UserRead +from admin.services.guards import ensure_can_delete, ensure_can_read, ensure_owns + +logger = get_logger(__name__) + +RESOURCE_TYPE = "media" + + +class MediaService: + def __init__( + self, *, media: MediaRepository, storage: S3Storage, audit: AuditRepository + ) -> None: + self._media = media + self._storage = storage + self._audit = audit + + def presets(self) -> list[MediaPresetRead]: + return [ + MediaPresetRead( + purpose=purpose, + max_edge=preset.max_edge, + max_bytes=preset.max_bytes, + mime_types=sorted(preset.mime_types), + ) + for purpose, preset in PRESETS.items() + ] + + async def create_upload_tickets( + self, *, actor: UserRead, payload: MediaUploadRequest + ) -> list[MediaUploadTicket]: + preset = preset_for(payload.purpose) + max_bytes = min(preset.max_bytes, self._storage.max_upload_bytes) + + for item in payload.files: + if item.mime_type not in preset.mime_types: + raise ValidationError( + f"{item.mime_type} is not accepted for {payload.purpose.value}", + details={"allowed": sorted(preset.mime_types)}, + ) + if item.size_bytes > max_bytes: + raise ValidationError( + f"{item.filename} is larger than the {max_bytes} byte limit", + details={"max_bytes": max_bytes}, + ) + + tickets: list[MediaUploadTicket] = [] + for item in payload.files: + key = build_object_key(visibility=preset.visibility, filename=item.filename) + ticket = self._storage.create_upload_ticket( + key=key, + content_type=item.mime_type, + visibility=preset.visibility, + max_bytes=max_bytes, + ) + record = await self._media.insert( + MediaCreate( + owner_id=actor.id, + s3_key=key, + original_filename=item.filename, + mime_type=item.mime_type, + size_bytes=item.size_bytes, + purpose=payload.purpose, + visibility=preset.visibility, + status=MediaUploadingStatus.UPLOADING, + ) + ) + tickets.append( + MediaUploadTicket( + media_id=record.id, + url=ticket.url, + fields=ticket.fields, + s3_key=ticket.key, + expires_in=ticket.expires_in, + ) + ) + + return tickets + + async def complete_upload(self, *, actor: UserRead, media_id: uuid.UUID) -> MediaRead: + record = await self._require(media_id) + ensure_owns(record, actor=actor) + + metadata = await self._storage.head(record.s3_key) + if metadata is None: + await self._media.update_status(media_id, MediaUploadingStatus.FAILED) + raise ValidationError("the file was never uploaded") + + updated = await self._media.update_status( + media_id, MediaUploadingStatus.COMPLETED, size_bytes=metadata.size_bytes + ) + if updated is None: + raise NotFoundError("media does not exist") + + await self._audit.record( + AuditEntry( + actor_id=actor.id, + actor_email=actor.email, + action=AuditAction.MEDIA_UPLOADED, + resource_type=RESOURCE_TYPE, + resource_id=str(media_id), + after={ + "purpose": updated.purpose.value, + "mime_type": updated.mime_type, + "size_bytes": updated.size_bytes, + }, + ) + ) + + return await self._to_read(updated) + + async def get_file( + self, *, actor: UserRead, actor_permissions: PermissionSet, media_id: uuid.UUID + ) -> MediaRead: + record = await self._require(media_id) + ensure_can_read(record, actor=actor, permissions=actor_permissions) + return await self._to_read(record) + + async def get_files( + self, *, actor: UserRead, actor_permissions: PermissionSet, media_ids: list[uuid.UUID] + ) -> list[MediaRead]: + records = await self._media.get_by_ids(media_ids) + for record in records: + ensure_can_read(record, actor=actor, permissions=actor_permissions) + + urls = await asyncio.gather(*(self._url_for(record) for record in records)) + return [self._build_read(record, url) for record, url in zip(records, urls, strict=True)] + + async def delete_file( + self, *, actor: UserRead, actor_permissions: PermissionSet, media_id: uuid.UUID + ) -> None: + await self.delete_files( + actor=actor, actor_permissions=actor_permissions, media_ids=[media_id] + ) + + async def delete_files( + self, *, actor: UserRead, actor_permissions: PermissionSet, media_ids: list[uuid.UUID] + ) -> None: + if not media_ids: + return + + records = await self._media.get_by_ids(media_ids) + if len(records) != len(set(media_ids)): + raise NotFoundError("media does not exist") + + for record in records: + ensure_can_delete(record, actor=actor, permissions=actor_permissions) + + variant_keys = await self._media.variant_keys_for(media_ids) + deleted = await self._media.delete_by_ids(media_ids) + + await self._audit.record_many( + [ + AuditEntry( + actor_id=actor.id, + actor_email=actor.email, + action=AuditAction.MEDIA_DELETED, + resource_type=RESOURCE_TYPE, + resource_id=str(record.id), + before={ + "owner_id": str(record.owner_id), + "s3_key": record.s3_key, + "purpose": record.purpose.value, + }, + ) + for record in deleted + ] + ) + + keys = [record.s3_key for record in deleted] + variant_keys + try: + await self._storage.delete_many(keys) + except Exception: + logger.exception("media_s3_delete_failed", keys=keys) + + def public_url(self, s3_key: str) -> str: + return self._storage.public_url(s3_key) + + async def _require(self, media_id: uuid.UUID) -> MediaRecord: + record = await self._media.get_by_id(media_id) + if record is None: + raise NotFoundError("media does not exist") + return record + + async def _url_for(self, record: MediaRecord) -> MediaUrl: + return await self._storage.url_for(key=record.s3_key, visibility=record.visibility) + + async def _to_read(self, record: MediaRecord) -> MediaRead: + return self._build_read(record, await self._url_for(record)) + + def _build_read(self, record: MediaRecord, url: MediaUrl) -> MediaRead: + return MediaRead( + id=record.id, + url=url.url, + url_expires_at=url.expires_at, + original_filename=record.original_filename, + mime_type=record.mime_type, + size_bytes=record.size_bytes, + purpose=record.purpose, + visibility=record.visibility, + status=record.status, + created_at=record.created_at, + ) From a8438bc381c531128a349603beb2515767b5b701 Mon Sep 17 00:00:00 2001 From: mai Date: Fri, 14 Aug 2026 23:03:02 -0400 Subject: [PATCH 3/3] make audit logging non-blocking --- backend/src/admin/api/dependencies.py | 11 ++- backend/src/admin/api/router.py | 2 + backend/src/admin/api/v1/audit.py | 29 +++++++ backend/src/admin/core/audit.py | 11 +++ backend/src/admin/repositories/audit.py | 17 ---- backend/src/admin/services/access.py | 40 +++++---- backend/src/admin/services/access_request.py | 10 +-- backend/src/admin/services/invitation.py | 8 +- backend/src/admin/services/media.py | 14 ++-- backend/src/admin/services/member.py | 8 +- backend/tests/test_audit_log.py | 86 ++++++++++++++++++++ 11 files changed, 176 insertions(+), 60 deletions(-) create mode 100644 backend/src/admin/api/v1/audit.py create mode 100644 backend/src/admin/core/audit.py create mode 100644 backend/tests/test_audit_log.py diff --git a/backend/src/admin/api/dependencies.py b/backend/src/admin/api/dependencies.py index 14a90ba..206be07 100644 --- a/backend/src/admin/api/dependencies.py +++ b/backend/src/admin/api/dependencies.py @@ -6,6 +6,7 @@ from fastapi import Depends, Request from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer +from admin.core.audit import AuditLog from admin.core.cache import Cache as CacheProtocol from admin.core.config import Settings, get_settings from admin.core.errors import ( @@ -93,6 +94,13 @@ def get_audit_repository(connection: Connection) -> AuditRepository: return AuditRepository(connection) +async def get_audit_log(connection: Connection) -> AsyncIterator[AuditLog]: + log = AuditLog() + yield log + if log.entries: + await AuditRepository(connection).record_many(log.entries) + + def get_media_repository(connection: Connection) -> MediaRepository: return MediaRepository(connection) @@ -101,7 +109,8 @@ def get_media_repository(connection: Connection) -> MediaRepository: Roles = Annotated[RoleRepository, Depends(get_role_repository)] Invitations = Annotated[InvitationRepository, Depends(get_invitation_repository)] AccessRequests = Annotated[AccessRequestRepository, Depends(get_access_request_repository)] -Audit = Annotated[AuditRepository, Depends(get_audit_repository)] +Audit = Annotated[AuditLog, Depends(get_audit_log)] +AuditEntries = Annotated[AuditRepository, Depends(get_audit_repository)] Media = Annotated[MediaRepository, Depends(get_media_repository)] diff --git a/backend/src/admin/api/router.py b/backend/src/admin/api/router.py index dbc7a5d..f08e97c 100644 --- a/backend/src/admin/api/router.py +++ b/backend/src/admin/api/router.py @@ -2,6 +2,7 @@ from admin.api.v1 import ( access_requests, + audit, health, invitations, media, @@ -17,6 +18,7 @@ api_router.include_router(access_requests.router) api_router.include_router(roles.router) api_router.include_router(media.router) +api_router.include_router(audit.router) root_router = APIRouter() root_router.include_router(health.router) diff --git a/backend/src/admin/api/v1/audit.py b/backend/src/admin/api/v1/audit.py new file mode 100644 index 0000000..2f4c6a8 --- /dev/null +++ b/backend/src/admin/api/v1/audit.py @@ -0,0 +1,29 @@ +import uuid +from typing import Annotated + +from fastapi import APIRouter, Depends + +from admin.api.dependencies import AuditEntries, AuthContext, require +from admin.domain.permissions import Permission +from admin.schemas.audit import AuditLogRead +from admin.schemas.base import CursorPage, CursorParams + +router = APIRouter(prefix="/audit", tags=["audit"]) + + +@router.get("", response_model=CursorPage[AuditLogRead]) +async def list_audit_entries( + entries: AuditEntries, + params: Annotated[CursorParams, Depends()], + _: Annotated[AuthContext, Depends(require(Permission.AUDIT_READ))], + actor_id: uuid.UUID | None = None, + resource_type: str | None = None, + resource_id: str | None = None, +) -> CursorPage[AuditLogRead]: + rows = await entries.list_entries( + params, + actor_id=actor_id, + resource_type=resource_type, + resource_id=resource_id, + ) + return CursorPage[AuditLogRead].build(rows, params, lambda item: [item.created_at, item.id]) diff --git a/backend/src/admin/core/audit.py b/backend/src/admin/core/audit.py new file mode 100644 index 0000000..6b7cb25 --- /dev/null +++ b/backend/src/admin/core/audit.py @@ -0,0 +1,11 @@ +from dataclasses import dataclass, field + +from admin.schemas.audit import AuditEntry + + +@dataclass(slots=True) +class AuditLog: + entries: list[AuditEntry] = field(default_factory=list) + + def add(self, *entries: AuditEntry) -> None: + self.entries.extend(entries) diff --git a/backend/src/admin/repositories/audit.py b/backend/src/admin/repositories/audit.py index 8ab06e9..12c19da 100644 --- a/backend/src/admin/repositories/audit.py +++ b/backend/src/admin/repositories/audit.py @@ -12,23 +12,6 @@ class AuditRepository(Repository): - async def record(self, entry: AuditEntry) -> None: - await self.connection.execute( - """ - INSERT INTO audit_logs - (actor_id, actor_email, action, resource_type, resource_id, - before, after) - VALUES ($1, $2, $3, $4, $5, $6, $7) - """, - entry.actor_id, - entry.actor_email, - entry.action.value, - entry.resource_type, - entry.resource_id, - entry.before, - entry.after, - ) - async def record_many(self, entries: list[AuditEntry]) -> None: if not entries: return diff --git a/backend/src/admin/services/access.py b/backend/src/admin/services/access.py index 1c8e037..aaa5018 100644 --- a/backend/src/admin/services/access.py +++ b/backend/src/admin/services/access.py @@ -1,3 +1,4 @@ +from admin.core.audit import AuditLog from admin.domain.access import PermissionSet, ResolvedAccess from admin.domain.enums import ( AccessRequestStatus, @@ -6,7 +7,6 @@ UserStatus, ) from admin.repositories.access_request import AccessRequestRepository -from admin.repositories.audit import AuditRepository from admin.repositories.invitation import InvitationRepository from admin.repositories.role import RoleRepository from admin.repositories.user import UserRepository @@ -24,7 +24,7 @@ def __init__( invitations: InvitationRepository, access_requests: AccessRequestRepository, roles: RoleRepository, - audit: AuditRepository, + audit: AuditLog, ) -> None: self._users = users self._invitations = invitations @@ -99,24 +99,22 @@ async def _accept_invitation(self, identity: Identity, invitation: InvitationRea ) await self._invitations.mark_accepted(invitation.id) - await self._audit.record_many( - [ - AuditEntry( - actor_id=user.id, - actor_email=user.email, - action=AuditAction.INVITATION_ACCEPTED, - resource_type="invitation", - resource_id=str(invitation.id), - after={"user_id": str(user.id), "role_key": invitation.role.key}, - ), - AuditEntry( - actor_id=user.id, - actor_email=user.email, - action=AuditAction.USER_PROVISIONED, - resource_type="user", - resource_id=str(user.id), - after={"email": user.email, "source": "invitation"}, - ), - ] + self._audit.add( + AuditEntry( + actor_id=user.id, + actor_email=user.email, + action=AuditAction.INVITATION_ACCEPTED, + resource_type="invitation", + resource_id=str(invitation.id), + after={"user_id": str(user.id), "role_key": invitation.role.key}, + ), + AuditEntry( + actor_id=user.id, + actor_email=user.email, + action=AuditAction.USER_PROVISIONED, + resource_type="user", + resource_id=str(user.id), + after={"email": user.email, "source": "invitation"}, + ), ) return user diff --git a/backend/src/admin/services/access_request.py b/backend/src/admin/services/access_request.py index b06d047..a5e2b69 100644 --- a/backend/src/admin/services/access_request.py +++ b/backend/src/admin/services/access_request.py @@ -1,10 +1,10 @@ import uuid +from admin.core.audit import AuditLog from admin.core.errors import ConflictError, NotFoundError, ValidationError from admin.domain.access import PermissionSet from admin.domain.enums import AccessRequestStatus, AuditAction from admin.repositories.access_request import AccessRequestRepository -from admin.repositories.audit import AuditRepository from admin.repositories.role import RoleRepository from admin.repositories.user import UserRepository from admin.schemas.access_request import ( @@ -26,7 +26,7 @@ def __init__( access_requests: AccessRequestRepository, users: UserRepository, roles: RoleRepository, - audit: AuditRepository, + audit: AuditLog, ) -> None: self._access_requests = access_requests self._users = users @@ -50,7 +50,7 @@ async def create( message=payload.message, ) - await self._audit.record( + self._audit.add( AuditEntry( actor_email=identity.email, action=AuditAction.ACCESS_REQUEST_CREATED, @@ -108,7 +108,7 @@ async def approve( if decided is None: raise ValidationError("access request has already been reviewed") - await self._audit.record( + self._audit.add( AuditEntry( actor_id=actor.id, actor_email=actor.email, @@ -132,7 +132,7 @@ async def deny( if decided is None: raise NotFoundError("no pending access request with that id") - await self._audit.record( + self._audit.add( AuditEntry( actor_id=actor.id, actor_email=actor.email, diff --git a/backend/src/admin/services/invitation.py b/backend/src/admin/services/invitation.py index 97fe49b..784069b 100644 --- a/backend/src/admin/services/invitation.py +++ b/backend/src/admin/services/invitation.py @@ -3,10 +3,10 @@ import uuid from datetime import UTC, datetime, timedelta +from admin.core.audit import AuditLog from admin.core.errors import ConflictError, NotFoundError, ValidationError from admin.domain.access import PermissionSet from admin.domain.enums import AuditAction -from admin.repositories.audit import AuditRepository from admin.repositories.invitation import InvitationRepository from admin.repositories.role import RoleRepository from admin.repositories.user import UserRepository @@ -30,7 +30,7 @@ def __init__( invitations: InvitationRepository, roles: RoleRepository, users: UserRepository, - audit: AuditRepository, + audit: AuditLog, default_ttl_hours: int, ) -> None: self._invitations = invitations @@ -72,7 +72,7 @@ async def create( expires_at=expires_at, ) - await self._audit.record( + self._audit.add( AuditEntry( actor_id=actor.id, actor_email=actor.email, @@ -96,7 +96,7 @@ async def revoke(self, *, actor: UserRead, invitation_id: uuid.UUID) -> None: if not await self._invitations.revoke(invitation_id): raise ValidationError("invitation is already accepted or revoked") - await self._audit.record( + self._audit.add( AuditEntry( actor_id=actor.id, actor_email=actor.email, diff --git a/backend/src/admin/services/media.py b/backend/src/admin/services/media.py index 81f7e1d..cf259e0 100644 --- a/backend/src/admin/services/media.py +++ b/backend/src/admin/services/media.py @@ -1,13 +1,13 @@ import asyncio import uuid +from admin.core.audit import AuditLog from admin.core.errors import NotFoundError, ValidationError from admin.core.logging import get_logger from admin.core.storage import MediaUrl, S3Storage, build_object_key from admin.domain.access import PermissionSet from admin.domain.enums import AuditAction from admin.domain.media import PRESETS, preset_for -from admin.repositories.audit import AuditRepository from admin.repositories.media import MediaRepository from admin.schemas.audit import AuditEntry from admin.schemas.media import ( @@ -28,9 +28,7 @@ class MediaService: - def __init__( - self, *, media: MediaRepository, storage: S3Storage, audit: AuditRepository - ) -> None: + def __init__(self, *, media: MediaRepository, storage: S3Storage, audit: AuditLog) -> None: self._media = media self._storage = storage self._audit = audit @@ -112,7 +110,7 @@ async def complete_upload(self, *, actor: UserRead, media_id: uuid.UUID) -> Medi if updated is None: raise NotFoundError("media does not exist") - await self._audit.record( + self._audit.add( AuditEntry( actor_id=actor.id, actor_email=actor.email, @@ -169,8 +167,8 @@ async def delete_files( variant_keys = await self._media.variant_keys_for(media_ids) deleted = await self._media.delete_by_ids(media_ids) - await self._audit.record_many( - [ + self._audit.add( + *( AuditEntry( actor_id=actor.id, actor_email=actor.email, @@ -184,7 +182,7 @@ async def delete_files( }, ) for record in deleted - ] + ) ) keys = [record.s3_key for record in deleted] + variant_keys diff --git a/backend/src/admin/services/member.py b/backend/src/admin/services/member.py index fc7c9bd..4186b23 100644 --- a/backend/src/admin/services/member.py +++ b/backend/src/admin/services/member.py @@ -1,9 +1,9 @@ import uuid +from admin.core.audit import AuditLog from admin.core.errors import NotFoundError from admin.domain.access import PermissionSet from admin.domain.enums import AuditAction -from admin.repositories.audit import AuditRepository from admin.repositories.role import RoleRepository from admin.repositories.user import UserRepository from admin.schemas.audit import AuditEntry @@ -18,7 +18,7 @@ def __init__( *, users: UserRepository, roles: RoleRepository, - audit: AuditRepository, + audit: AuditLog, ) -> None: self._users = users self._roles = roles @@ -55,7 +55,7 @@ async def grant_role( expires_at=payload.expires_at, ) - await self._audit.record( + self._audit.add( AuditEntry( actor_id=actor.id, actor_email=actor.email, @@ -85,7 +85,7 @@ async def revoke_role(self, *, actor: UserRead, assignment_id: uuid.UUID) -> Use await ensure_not_last_owner(self._users, assignment.role.key) await self._roles.revoke(assignment_id) - await self._audit.record( + self._audit.add( AuditEntry( actor_id=actor.id, actor_email=actor.email, diff --git a/backend/tests/test_audit_log.py b/backend/tests/test_audit_log.py new file mode 100644 index 0000000..ec3cf34 --- /dev/null +++ b/backend/tests/test_audit_log.py @@ -0,0 +1,86 @@ +import uuid +from contextlib import asynccontextmanager + +import pytest +from fastapi import APIRouter, FastAPI, Request +from fastapi.testclient import TestClient + +from admin.api.dependencies import Audit, Connection, get_connection +from admin.core.audit import AuditLog +from admin.core.database import create_pool +from admin.domain.enums import AuditAction +from admin.schemas.audit import AuditEntry + + +def entry(resource_id: str) -> AuditEntry: + return AuditEntry( + action=AuditAction.USER_PROVISIONED, + resource_type="test", + resource_id=resource_id, + ) + + +@pytest.fixture +def app(settings) -> FastAPI: + router = APIRouter() + + @router.post("/write/{resource_id}") + async def write(resource_id: str, audit: Audit) -> dict[str, str]: + audit.add(entry(resource_id), entry(resource_id)) + return {"ok": resource_id} + + @router.post("/write-then-fail/{resource_id}") + async def write_then_fail(resource_id: str, audit: Audit) -> dict[str, str]: + audit.add(entry(resource_id)) + raise RuntimeError("boom") + + @router.get("/count/{resource_id}") + async def count(resource_id: str, connection: Connection) -> dict[str, int]: + rows = await connection.fetchval( + "SELECT count(*) FROM audit_logs WHERE resource_type='test' AND resource_id=$1", + resource_id, + ) + return {"count": rows} + + @asynccontextmanager + async def lifespan(instance: FastAPI): + instance.state.pool = await create_pool(settings.database) + yield + await instance.state.pool.close() + + application = FastAPI(lifespan=lifespan) + + async def connection_override(request: Request): + async with request.app.state.pool.acquire() as connection, connection.transaction(): + yield connection + + application.dependency_overrides[get_connection] = connection_override + application.include_router(router) + return application + + +def test_entries_flush_inside_the_request_transaction(app: FastAPI) -> None: + resource_id = str(uuid.uuid4()) + with TestClient(app) as client: + assert client.post(f"/write/{resource_id}").status_code == 200 + assert client.get(f"/count/{resource_id}").json()["count"] == 2 + + +def test_entries_roll_back_when_the_request_fails(app: FastAPI) -> None: + resource_id = str(uuid.uuid4()) + with TestClient(app, raise_server_exceptions=False) as client: + client.post(f"/write-then-fail/{resource_id}") + assert client.get(f"/count/{resource_id}").json()["count"] == 0 + + +def test_nothing_is_written_when_no_entries_are_added(app: FastAPI) -> None: + resource_id = str(uuid.uuid4()) + with TestClient(app) as client: + assert client.get(f"/count/{resource_id}").json()["count"] == 0 + + +def test_add_accepts_multiple_entries() -> None: + log = AuditLog() + log.add(entry("a")) + log.add(entry("b"), entry("c")) + assert [item.resource_id for item in log.entries] == ["a", "b", "c"]