From f244a3b31a3b693ef30e9e3b4ea618f9783d1115 Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Mon, 24 Aug 2026 23:58:58 +0100 Subject: [PATCH 01/29] feat: gate mobile endpoints behind face enrollment, fix jsonb round-trip - Add require_onboarded_mobile_user dependency: every mobile endpoint except /auth/* and /enroll now returns 403 until the user completes face enrollment (users.face_embedding is set). - Add is_onboarded to GET /user/auth/me and the avatar upload response so clients can branch on a single field instead of parsing 403s. - Fix asyncpg/SQLAlchemy jsonb handling: register a jsonb type codec on connect so dict params bind correctly and jsonb columns decode back to dict (was crashing audit_events writes, and silently mismatched on notifications/staff_notifications reads). - Add a dev-env fixed OTP (DEV_OTP_BYPASS_CODE) for registration so local/mobile testing doesn't require a real inbox or NATS email flow. --- app/core/config.py | 3 +++ app/deps/token_auth.py | 21 ++++++++++++++++++++ app/infra/database.py | 26 ++++++++++++++++++++++++- app/router/mobile/audit.py | 4 ++-- app/router/mobile/auth.py | 2 ++ app/router/mobile/event.py | 6 +++--- app/router/mobile/notifications.py | 6 +++--- app/router/mobile/photo_approval.py | 6 +++--- app/router/mobile/photos.py | 8 ++++---- app/schema/response/mobile/auth.py | 1 + app/service/user_notification.py | 3 +-- app/service/users.py | 30 ++++++++++++++++++----------- 12 files changed, 87 insertions(+), 29 deletions(-) diff --git a/app/core/config.py b/app/core/config.py index dea2f69a..35ec0daf 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -54,6 +54,9 @@ class Settings(BaseSettings): # Rate Limit Settings RATE_LIMIT_LOGIN_MAX_ATTEMPTS: int = 5 RATE_LIMIT_LOGIN_WINDOW_SECONDS: int = 60 + # In dev env, registration OTPs are fixed to this value and the email/NATS + # send is skipped, so mobile devs can verify without a real inbox. + DEV_OTP_BYPASS_CODE: str = "000000" TRUST_PROXY_HEADERS: bool = True # Admin list defaults ADMIN_USERS_DEFAULT_LIMIT: int = 20 diff --git a/app/deps/token_auth.py b/app/deps/token_auth.py index a14de6c4..7c778309 100644 --- a/app/deps/token_auth.py +++ b/app/deps/token_auth.py @@ -114,3 +114,24 @@ async def get_current_mobile_user( email=user.email or "", session_id=session.id, ) + + +async def require_onboarded_mobile_user( + current_user: Annotated[MobileUserSchema, Depends(get_current_mobile_user)], + container: Annotated[Container, Depends(get_container)], +) -> MobileUserSchema: + """Gate for endpoints that require a completed face enrollment. + Auth (login/me/devices/etc.), /enroll, and /event/join stay reachable + via plain get_current_mobile_user so a new user can always finish + onboarding; everything else (photos, notifications, audits, /event/me) + depends on this instead. + """ + user = await container.auth_service.user_querier.get_user_by_id(id=current_user.user_id) + if user is None: + raise HTTPException(status_code=401, detail="User not found") + if user.face_embedding is None: + raise HTTPException( + status_code=403, + detail="Complete face enrollment before accessing this resource", + ) + return current_user diff --git a/app/infra/database.py b/app/infra/database.py index 956662bd..930b3368 100644 --- a/app/infra/database.py +++ b/app/infra/database.py @@ -1,5 +1,8 @@ -from typing import AsyncGenerator +import json +from typing import Any, AsyncGenerator + import sqlalchemy.ext.asyncio +from sqlalchemy import event from app.core.config import settings @@ -13,6 +16,27 @@ ) +@event.listens_for(engine.sync_engine, "connect") +def _register_jsonb_codec(dbapi_connection: Any, connection_record: Any) -> None: + # SQLAlchemy's asyncpg dialect doesn't forward connect_args={"init": ...} + # to asyncpg (it raises TypeError: unexpected keyword argument 'init'), + # so codec registration has to go through the wrapped connection's + # run_async bridge instead. Without this, asyncpg neither accepts a + # Python dict as a jsonb bind param (raises DataError) nor decodes a + # jsonb column back into one (returns raw JSON text) — every jsonb + # column in the schema (audit metadata, notification payloads) needs + # both directions to work. + dbapi_connection.run_async( + lambda conn: conn.set_type_codec( + "jsonb", + encoder=json.dumps, + decoder=json.loads, + schema="pg_catalog", + format="text", + ) + ) + + async def get_db() -> AsyncGenerator[sqlalchemy.ext.asyncio.AsyncConnection, None]: async with engine.begin() as conn: yield conn diff --git a/app/router/mobile/audit.py b/app/router/mobile/audit.py index fe3226a5..14cba0af 100644 --- a/app/router/mobile/audit.py +++ b/app/router/mobile/audit.py @@ -7,7 +7,7 @@ from app.container import Container, get_container from app.core.constant import AuditEventType -from app.deps.token_auth import MobileUserSchema, get_current_mobile_user +from app.deps.token_auth import MobileUserSchema, require_onboarded_mobile_user from app.schema.response.mobile.audit import AuditEventListResponse, AuditEventSchema router = APIRouter(prefix="/audits", tags=["audits"]) @@ -22,7 +22,7 @@ async def list_audits( limit: int = Query(50, ge=1, le=200), offset: int = Query(0, ge=0), container: Container = Depends(get_container), - _: MobileUserSchema = Depends(get_current_mobile_user), + _: MobileUserSchema = Depends(require_onboarded_mobile_user), ) -> AuditEventListResponse: events = await container.audit_service.list_audit_events( event_type=event_type, diff --git a/app/router/mobile/auth.py b/app/router/mobile/auth.py index ce444c00..a7e4dddd 100644 --- a/app/router/mobile/auth.py +++ b/app/router/mobile/auth.py @@ -206,6 +206,7 @@ async def get_me( email=user.email, name=user.display_name, avatar_url="/user/auth/me/avatar/image" if user.avatar_key else None, + is_onboarded=user.face_embedding is not None, ), devices=device_list, sessions=session_schema, @@ -239,6 +240,7 @@ async def upload_avatar( email=user.email, name=user.display_name, avatar_url="/user/auth/me/avatar/image", + is_onboarded=user.face_embedding is not None, ) diff --git a/app/router/mobile/event.py b/app/router/mobile/event.py index adada86a..33548e75 100644 --- a/app/router/mobile/event.py +++ b/app/router/mobile/event.py @@ -3,7 +3,7 @@ from fastapi import APIRouter, Depends from app.container import Container, get_container -from app.deps.token_auth import MobileUserSchema, get_current_mobile_user +from app.deps.token_auth import MobileUserSchema, require_onboarded_mobile_user from app.schema.request.web.event import JoinEventRequest from app.schema.response.web.event import JoinEventResponse, UserEventResponse @@ -13,7 +13,7 @@ async def join_event( req: JoinEventRequest, container: Container = Depends(get_container), - current_user: MobileUserSchema = Depends(get_current_mobile_user), + current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), )-> JoinEventResponse: return await container.event_service.join_event_by_code( user_id=current_user.user_id, @@ -24,6 +24,6 @@ async def join_event( @router.get("/me", response_model=List[UserEventResponse]) async def get_my_joined_events( container: Container = Depends(get_container), - current_user: MobileUserSchema = Depends(get_current_mobile_user), + current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), )-> List[UserEventResponse]: return await container.event_service.get_my_events(current_user.user_id) diff --git a/app/router/mobile/notifications.py b/app/router/mobile/notifications.py index f8d4a569..ca973fcb 100644 --- a/app/router/mobile/notifications.py +++ b/app/router/mobile/notifications.py @@ -1,7 +1,7 @@ from fastapi import APIRouter, Depends from app.container import Container, get_container -from app.deps.token_auth import MobileUserSchema, get_current_mobile_user +from app.deps.token_auth import MobileUserSchema, require_onboarded_mobile_user from app.schema.request.mobile.notifications import MarkUserNotificationsReadRequest from app.schema.response.mobile.notifications import UserNotificationListResponse @@ -12,7 +12,7 @@ @router.get("", response_model=UserNotificationListResponse) async def get_all_notifications( container: Container = Depends(get_container), - current_user: MobileUserSchema = Depends(get_current_mobile_user), + current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), ) -> UserNotificationListResponse: notifications = await container.user_notifications_service.get_all_notifications( user_id=current_user.user_id, @@ -24,7 +24,7 @@ async def get_all_notifications( async def mark_as_read( req: MarkUserNotificationsReadRequest, container: Container = Depends(get_container), - current_user: MobileUserSchema = Depends(get_current_mobile_user), + current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), ) -> UserNotificationListResponse: notifications = await container.user_notifications_service.mark_notifications_as_read( notification_ids=req.notification_ids, diff --git a/app/router/mobile/photo_approval.py b/app/router/mobile/photo_approval.py index 3aead0d3..9f5f97f0 100644 --- a/app/router/mobile/photo_approval.py +++ b/app/router/mobile/photo_approval.py @@ -4,7 +4,7 @@ from fastapi import APIRouter, Depends, Query from app.container import Container, get_container -from app.deps.token_auth import MobileUserSchema, get_current_mobile_user +from app.deps.token_auth import MobileUserSchema, require_onboarded_mobile_user from app.schema.request.mobile.photo_approval import PhotoApprovalRequest router = APIRouter(prefix="/photos") @@ -15,7 +15,7 @@ async def list_my_approvals( status: Literal["pending", "approved", "rejected"] | None = Query(default=None), limit: int = Query(default=50, ge=1, le=100), offset: int = Query(default=0, ge=0), - current_user: MobileUserSchema = Depends(get_current_mobile_user), + current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), container: Container = Depends(get_container), ) -> list[dict[str, object]]: approvals: list[dict[str, object]] = [] @@ -39,7 +39,7 @@ async def list_my_approvals( async def decide_photo_approval( photo_id: UUID, req: PhotoApprovalRequest, - current_user: MobileUserSchema = Depends(get_current_mobile_user), + current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), container: Container = Depends(get_container), ) -> dict[str, str]: photo_status = await container.photo_approval_service.decide( diff --git a/app/router/mobile/photos.py b/app/router/mobile/photos.py index 00113f97..a9f6036c 100644 --- a/app/router/mobile/photos.py +++ b/app/router/mobile/photos.py @@ -5,7 +5,7 @@ from fastapi.responses import Response from app.container import Container, get_container -from app.deps.token_auth import MobileUserSchema, get_current_mobile_user +from app.deps.token_auth import MobileUserSchema, require_onboarded_mobile_user from app.deps.rate_limit import RateLimiter router = APIRouter(prefix="/photos") @@ -17,7 +17,7 @@ async def list_my_photos( sort: Literal["asc", "desc"] = Query(default="desc"), limit: int = Query(default=50, ge=1, le=100), offset: int = Query(default=0, ge=0), - current_user: MobileUserSchema = Depends(get_current_mobile_user), + current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), container: Container = Depends(get_container), ) -> list[dict[str, object]]: photos = await container.user_photo_service.list_photos( @@ -47,7 +47,7 @@ async def list_event_photos( sort: Literal["asc", "desc"] = Query(default="desc"), limit: int = Query(default=50, ge=1, le=100), offset: int = Query(default=0, ge=0), - current_user: MobileUserSchema = Depends(get_current_mobile_user), + current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), container: Container = Depends(get_container), ) -> dict[str, object]: photos = await container.user_photo_service.list_event_photos( @@ -81,7 +81,7 @@ async def list_event_photos( @router.get("/{photo_id}/image") async def get_photo_image( photo_id: UUID, - current_user: MobileUserSchema = Depends(get_current_mobile_user), + current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), container: Container = Depends(get_container), ) -> Response: data, filename, content_type = await container.user_photo_service.get_photo_bytes( diff --git a/app/schema/response/mobile/auth.py b/app/schema/response/mobile/auth.py index 67bf1398..d9949374 100644 --- a/app/schema/response/mobile/auth.py +++ b/app/schema/response/mobile/auth.py @@ -26,6 +26,7 @@ class UserSchema(BaseModel): email: str name: str | None avatar_url: str | None + is_onboarded: bool class MeResponse(BaseModel): user: UserSchema diff --git a/app/service/user_notification.py b/app/service/user_notification.py index f8181bc0..a269ff17 100644 --- a/app/service/user_notification.py +++ b/app/service/user_notification.py @@ -1,4 +1,3 @@ -import json from typing import Any import uuid @@ -42,7 +41,7 @@ async def create_notification( notification_record = await self.notification_querier.create_notification( user_id=user_id, type=type, - payload=json.dumps(payload), + payload=payload, ) if notification_record is None: raise AppException.internal_error("Failed to create user notification") diff --git a/app/service/users.py b/app/service/users.py index 3864757b..5682b826 100644 --- a/app/service/users.py +++ b/app/service/users.py @@ -176,7 +176,6 @@ async def mobile_register( raise AppException.conflict("Email already in use; please login instead") hashed = hash_password(req.password) - otp = "".join(secrets.choice("0123456789") for _ in range(6)) pending_key = f"pending_user:{req.email}" pending_data = { @@ -185,10 +184,16 @@ async def mobile_register( # Save in Redis for 10 minutes (600 seconds) await redis.set(pending_key, json.dumps(pending_data), expire=600) - await redis.set(f"otp:{req.email}", otp, expire=600) - # Send to NATS - await NatsClient.publish("email.send_otp", json.dumps({"email": req.email, "otp": otp}).encode("utf-8")) + if settings.environment == "dev": + otp = settings.DEV_OTP_BYPASS_CODE + await redis.set(f"otp:{req.email}", otp, expire=600) + logger.info("dev OTP bypass active, otp=%s email=%s", otp, req.email) + else: + otp = "".join(secrets.choice("0123456789") for _ in range(6)) + await redis.set(f"otp:{req.email}", otp, expire=600) + # Send to NATS + await NatsClient.publish("email.send_otp", json.dumps({"email": req.email, "otp": otp}).encode("utf-8")) logger.info("register success, OTP sent") return RegisterPendingResponse( @@ -226,13 +231,16 @@ async def mobile_register_resend_otp( if not raw_data: raise AppException.not_found("No pending registration found for this email") - otp = "".join(secrets.choice("0123456789") for _ in range(6)) - - # Regenerate OTP with 10 mins TTL, without touching the pending_user TTL - await redis.set(f"otp:{email}", otp, expire=600) - - # Send to NATS - await NatsClient.publish("email.send_otp", json.dumps({"email": email, "otp": otp}).encode("utf-8")) + if settings.environment == "dev": + otp = settings.DEV_OTP_BYPASS_CODE + await redis.set(f"otp:{email}", otp, expire=600) + logger.info("dev OTP bypass active, otp=%s email=%s", otp, email) + else: + otp = "".join(secrets.choice("0123456789") for _ in range(6)) + # Regenerate OTP with 10 mins TTL, without touching the pending_user TTL + await redis.set(f"otp:{email}", otp, expire=600) + # Send to NATS + await NatsClient.publish("email.send_otp", json.dumps({"email": email, "otp": otp}).encode("utf-8")) logger.info("resend_otp success, new OTP sent to %s", email) return RegisterPendingResponse( From e17085183ed1a0dfbc10f6a33d2f846feed80d7f Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 00:32:20 +0100 Subject: [PATCH 02/29] fix: handle redis unavailability in enrollment lock acquisition --- app/router/mobile/enrollement.py | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/app/router/mobile/enrollement.py b/app/router/mobile/enrollement.py index f8e2ac10..1c7cc4cf 100644 --- a/app/router/mobile/enrollement.py +++ b/app/router/mobile/enrollement.py @@ -139,12 +139,19 @@ async def image_payloads() -> AsyncIterator[FaceImagePayload]: lock_key = _enrollment_lock_key(user.user_id) lock_value = str(uuid.uuid4()) - lock_acquired = await container.redis.set( - lock_key, - lock_value, - expire=ENROLL_IN_PROGRESS_TTL_SECONDS, - nx=True, - ) + try: + lock_acquired = await container.redis.set( + lock_key, + lock_value, + expire=ENROLL_IN_PROGRESS_TTL_SECONDS, + nx=True, + ) + except Exception as exc: + logger.warning( + "enroll: redis unavailable, failing open (no duplicate-submission lock) for user %s: %s", + user.user_id, exc, + ) + lock_acquired = True if not lock_acquired: raise AppException.conflict( "Enrollment already in progress. Please wait for it to finish." From 75bba7344c830f5cd19e8010c8fa834faeec58c8 Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 01:11:03 +0100 Subject: [PATCH 03/29] fix: update secure cookie setting based on environment --- app/router/web/auth.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/app/router/web/auth.py b/app/router/web/auth.py index 7629cf93..24d7a402 100644 --- a/app/router/web/auth.py +++ b/app/router/web/auth.py @@ -2,6 +2,7 @@ from app.container import Container, get_container from fastapi import Response +from app.core.config import settings from app.deps.cookie_auth import get_current_staff_user from app.deps.rate_limit import RateLimiter from app.schema.request.web.auth import WebAuthRequest @@ -26,7 +27,7 @@ async def admin_login( key="access_token", value=authResponse.access_token, httponly=True, - secure=True, + secure=settings.environment != "dev", samesite="strict", max_age=60 * 60 * 24 * 7, ) From 59c904c9b0bccfe77e8e5dc109f4362b4e37e5a4 Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 01:39:40 +0100 Subject: [PATCH 04/29] feat: automate event lifecycle transitions on start/end time - Add events.end_date (nullable timestamptz) via migration. - Add ActivateDueEvents / ArchiveEndedEvents queries: draft -> scheduled when event_date has passed, scheduled -> archived (+ archived_at) when end_date has passed. - Add a polling worker (app/worker/event_lifecycle) that runs both transitions every EVENT_LIFECYCLE_POLL_INTERVAL_SECONDS (default 60s), wired into make run-workers. - Thread end_date through EventCreate/EventResponse/UserEventResponse and CreateEvent/GetUserEvents. Events with no end_date set are never auto-archived and stay in scheduled until archived manually via the existing endpoint. --- app/core/config.py | 1 + app/schema/request/web/event.py | 1 + app/schema/response/web/event.py | 2 + app/service/event.py | 1 + app/worker/event_lifecycle/__init__.py | 0 app/worker/event_lifecycle/main.py | 36 +++++++++++ db/generated/event_participant.py | 15 +++-- db/generated/events.py | 60 +++++++++++++++---- db/generated/models.py | 1 + db/queries/event_participant.sql | 9 +-- db/queries/events.sql | 24 ++++++-- makefile | 1 + .../sql/down/add_end_date_to_events.sql | 1 + migrations/sql/up/add_end_date_to_events.sql | 1 + .../afb0a93aa21e_add_end_date_to_events.py | 27 +++++++++ 15 files changed, 155 insertions(+), 25 deletions(-) create mode 100644 app/worker/event_lifecycle/__init__.py create mode 100644 app/worker/event_lifecycle/main.py create mode 100644 migrations/sql/down/add_end_date_to_events.sql create mode 100644 migrations/sql/up/add_end_date_to_events.sql create mode 100644 migrations/versions/afb0a93aa21e_add_end_date_to_events.py diff --git a/app/core/config.py b/app/core/config.py index 35ec0daf..c9270f03 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -34,6 +34,7 @@ class Settings(BaseSettings): POSTGRES_PORT: int = 5432 PHOTO_APPROVAL_TIMEOUT_DAYS: int = 7 + EVENT_LIFECYCLE_POLL_INTERVAL_SECONDS: int = 60 # Mobile auth/session defaults MOBILE_SESSION_LIMIT: int = 3 diff --git a/app/schema/request/web/event.py b/app/schema/request/web/event.py index d03b46ff..c7c05e3d 100644 --- a/app/schema/request/web/event.py +++ b/app/schema/request/web/event.py @@ -5,6 +5,7 @@ class EventCreate(BaseModel): name: str event_date: datetime + end_date: Optional[datetime] = None status: Optional[str] = "draft" class JoinEventRequest(BaseModel): diff --git a/app/schema/response/web/event.py b/app/schema/response/web/event.py index 4334fc44..01a64266 100644 --- a/app/schema/response/web/event.py +++ b/app/schema/response/web/event.py @@ -10,6 +10,7 @@ class EventResponse(BaseModel): name: str event_code: str event_date: datetime + end_date: Optional[datetime] = None status: str created_by: uuid.UUID created_at: datetime @@ -30,6 +31,7 @@ class UserEventResponse(BaseModel): id: uuid.UUID name: str event_date: datetime + end_date: Optional[datetime] = None status: str joined_at: datetime diff --git a/app/service/event.py b/app/service/event.py index e697ddd8..fb6165e9 100644 --- a/app/service/event.py +++ b/app/service/event.py @@ -33,6 +33,7 @@ async def create_event(self, req: EventCreate, creator_id: uuid.UUID) -> EventRe name=req.name, event_code=code_created, event_date=req.event_date, + end_date=req.end_date, status=req.status or "draft", created_by=creator_id ) diff --git a/app/worker/event_lifecycle/__init__.py b/app/worker/event_lifecycle/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/app/worker/event_lifecycle/main.py b/app/worker/event_lifecycle/main.py new file mode 100644 index 00000000..7ab11b8d --- /dev/null +++ b/app/worker/event_lifecycle/main.py @@ -0,0 +1,36 @@ +import asyncio + +from app.core.config import settings +from app.core.logger import logger +from app.infra.database import engine +from db.generated import events as event_queries + + +async def run_lifecycle_pass() -> None: + async with engine.begin() as conn: + querier = event_queries.AsyncQuerier(conn) + + activated = [event_id async for event_id in querier.activate_due_events()] + if activated: + logger.info("event_lifecycle: activated %d event(s): %s", len(activated), activated) + + archived = [event_id async for event_id in querier.archive_ended_events()] + if archived: + logger.info("event_lifecycle: archived %d event(s): %s", len(archived), archived) + + +async def main() -> None: + logger.info( + "Event lifecycle worker starting, poll_interval=%ds", + settings.EVENT_LIFECYCLE_POLL_INTERVAL_SECONDS, + ) + while True: + try: + await run_lifecycle_pass() + except Exception: + logger.exception("event_lifecycle: pass failed") + await asyncio.sleep(settings.EVENT_LIFECYCLE_POLL_INTERVAL_SECONDS) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/db/generated/event_participant.py b/db/generated/event_participant.py index 65129c77..2ad09031 100644 --- a/db/generated/event_participant.py +++ b/db/generated/event_participant.py @@ -40,10 +40,11 @@ class GetEventParticipantsRow: GET_USER_EVENTS = """-- name: get_user_events \\:many -SELECT - e.id, - e.name, - e.event_date, +SELECT + e.id, + e.name, + e.event_date, + e.end_date, e.status, ep.joined_at FROM events e @@ -58,6 +59,7 @@ class GetUserEventsRow: id: uuid.UUID name: str event_date: datetime.datetime + end_date: Optional[datetime.datetime] status: Any joined_at: datetime.datetime @@ -164,8 +166,9 @@ async def get_user_events(self, *, user_id: uuid.UUID) -> AsyncIterator[GetUserE id=row[0], name=row[1], event_date=row[2], - status=row[3], - joined_at=row[4], + end_date=row[3], + status=row[4], + joined_at=row[5], ) async def is_user_in_event(self, *, event_id: uuid.UUID, user_id: uuid.UUID) -> Optional[bool]: diff --git a/db/generated/events.py b/db/generated/events.py index 0395bfda..1cd52749 100644 --- a/db/generated/events.py +++ b/db/generated/events.py @@ -13,10 +13,30 @@ from db.generated import models +ACTIVATE_DUE_EVENTS = """-- name: activate_due_events \\:many +UPDATE events +SET status = 'scheduled'\\:\\:event_status +WHERE status = 'draft'\\:\\:event_status + AND event_date <= NOW() +RETURNING id +""" + + +ARCHIVE_ENDED_EVENTS = """-- name: archive_ended_events \\:many +UPDATE events +SET status = 'archived'\\:\\:event_status, + archived_at = NOW() +WHERE status = 'scheduled'\\:\\:event_status + AND end_date IS NOT NULL + AND end_date <= NOW() +RETURNING id +""" + + CREATE_EVENT = """-- name: create_event \\:one -INSERT INTO events (name, event_code, event_date, status, created_by) -VALUES (:p1, :p2, :p3, :p4, :p5) -RETURNING id, name, event_code, event_date, status, created_by, created_at, archived_at +INSERT INTO events (name, event_code, event_date, end_date, status, created_by) +VALUES (:p1, :p2, :p3, :p4, :p5, :p6) +RETURNING id, name, event_code, event_date, status, created_by, created_at, archived_at, end_date """ @@ -25,37 +45,38 @@ class CreateEventParams: name: str event_code: str event_date: datetime.datetime + end_date: Optional[datetime.datetime] status: Any created_by: uuid.UUID DELETE_EVENT = """-- name: delete_event \\:exec -DELETE FROM events +DELETE FROM events WHERE id = :p1 """ GET_EVENT_BY_CODE = """-- name: get_event_by_code \\:one -SELECT id, name, event_code, event_date, status, created_by, created_at, archived_at FROM events +SELECT id, name, event_code, event_date, status, created_by, created_at, archived_at, end_date FROM events WHERE event_code = :p1 """ GET_EVENT_BY_ID = """-- name: get_event_by_id \\:one -SELECT id, name, event_code, event_date, status, created_by, created_at, archived_at FROM events +SELECT id, name, event_code, event_date, status, created_by, created_at, archived_at, end_date FROM events WHERE id = :p1 """ GET_EVENTS_BY_NAME = """-- name: get_events_by_name \\:many -SELECT id, name, event_code, event_date, status, created_by, created_at, archived_at FROM events +SELECT id, name, event_code, event_date, status, created_by, created_at, archived_at, end_date FROM events WHERE name ILIKE '%' || :p1 || '%' ORDER BY event_date DESC """ LIST_EVENTS = """-- name: list_events \\:many -SELECT id, name, event_code, event_date, status, created_by, created_at, archived_at FROM events +SELECT id, name, event_code, event_date, status, created_by, created_at, archived_at, end_date FROM events WHERE -- Filter by Status (Optional) (:p3\\:\\:event_status IS NULL OR status = :p3) @@ -95,7 +116,7 @@ class ListEventsParams: SET status = :p2, archived_at = CASE WHEN :p2 = 'archived'\\:\\:event_status THEN NOW() ELSE archived_at END WHERE id = :p1 -RETURNING id, name, event_code, event_date, status, created_by, created_at, archived_at +RETURNING id, name, event_code, event_date, status, created_by, created_at, archived_at, end_date """ @@ -103,13 +124,24 @@ class AsyncQuerier: def __init__(self, conn: sqlalchemy.ext.asyncio.AsyncConnection): self._conn = conn + async def activate_due_events(self) -> AsyncIterator[uuid.UUID]: + result = await self._conn.stream(sqlalchemy.text(ACTIVATE_DUE_EVENTS)) + async for row in result: + yield row[0] + + async def archive_ended_events(self) -> AsyncIterator[uuid.UUID]: + result = await self._conn.stream(sqlalchemy.text(ARCHIVE_ENDED_EVENTS)) + async for row in result: + yield row[0] + async def create_event(self, arg: CreateEventParams) -> Optional[models.Event]: row = (await self._conn.execute(sqlalchemy.text(CREATE_EVENT), { "p1": arg.name, "p2": arg.event_code, "p3": arg.event_date, - "p4": arg.status, - "p5": arg.created_by, + "p4": arg.end_date, + "p5": arg.status, + "p6": arg.created_by, })).first() if row is None: return None @@ -122,6 +154,7 @@ async def create_event(self, arg: CreateEventParams) -> Optional[models.Event]: created_by=row[5], created_at=row[6], archived_at=row[7], + end_date=row[8], ) async def delete_event(self, *, id: uuid.UUID) -> None: @@ -140,6 +173,7 @@ async def get_event_by_code(self, *, event_code: str) -> Optional[models.Event]: created_by=row[5], created_at=row[6], archived_at=row[7], + end_date=row[8], ) async def get_event_by_id(self, *, id: uuid.UUID) -> Optional[models.Event]: @@ -155,6 +189,7 @@ async def get_event_by_id(self, *, id: uuid.UUID) -> Optional[models.Event]: created_by=row[5], created_at=row[6], archived_at=row[7], + end_date=row[8], ) async def get_events_by_name(self, *, dollar_1: Optional[str]) -> AsyncIterator[models.Event]: @@ -169,6 +204,7 @@ async def get_events_by_name(self, *, dollar_1: Optional[str]) -> AsyncIterator[ created_by=row[5], created_at=row[6], archived_at=row[7], + end_date=row[8], ) async def list_events(self, arg: ListEventsParams) -> AsyncIterator[models.Event]: @@ -191,6 +227,7 @@ async def list_events(self, arg: ListEventsParams) -> AsyncIterator[models.Event created_by=row[5], created_at=row[6], archived_at=row[7], + end_date=row[8], ) async def update_event_status(self, *, id: uuid.UUID, status: Any) -> Optional[models.Event]: @@ -206,4 +243,5 @@ async def update_event_status(self, *, id: uuid.UUID, status: Any) -> Optional[m created_by=row[5], created_at=row[6], archived_at=row[7], + end_date=row[8], ) diff --git a/db/generated/models.py b/db/generated/models.py index 21bac799..86617af5 100644 --- a/db/generated/models.py +++ b/db/generated/models.py @@ -74,6 +74,7 @@ class Event: created_by: uuid.UUID created_at: datetime.datetime archived_at: Optional[datetime.datetime] + end_date: Optional[datetime.datetime] @dataclasses.dataclass() diff --git a/db/queries/event_participant.sql b/db/queries/event_participant.sql index 9203c83d..42de3e66 100644 --- a/db/queries/event_participant.sql +++ b/db/queries/event_participant.sql @@ -6,10 +6,11 @@ RETURNING *; -- name: GetUserEvents :many -- Retrieves all events a specific user has successfully joined -SELECT - e.id, - e.name, - e.event_date, +SELECT + e.id, + e.name, + e.event_date, + e.end_date, e.status, ep.joined_at FROM events e diff --git a/db/queries/events.sql b/db/queries/events.sql index e7a7fd84..5e1fde41 100644 --- a/db/queries/events.sql +++ b/db/queries/events.sql @@ -1,6 +1,6 @@ -- name: CreateEvent :one -INSERT INTO events (name, event_code, event_date, status, created_by) -VALUES ($1, $2, $3, $4, $5) +INSERT INTO events (name, event_code, event_date, end_date, status, created_by) +VALUES ($1, $2, $3, $4, $5, $6) RETURNING *; -- name: GetEventById :one @@ -47,5 +47,21 @@ WHERE id = $1 RETURNING *; -- name: DeleteEvent :exec -DELETE FROM events -WHERE id = $1; \ No newline at end of file +DELETE FROM events +WHERE id = $1; + +-- name: ActivateDueEvents :many +UPDATE events +SET status = 'scheduled'::event_status +WHERE status = 'draft'::event_status + AND event_date <= NOW() +RETURNING id; + +-- name: ArchiveEndedEvents :many +UPDATE events +SET status = 'archived'::event_status, + archived_at = NOW() +WHERE status = 'scheduled'::event_status + AND end_date IS NOT NULL + AND end_date <= NOW() +RETURNING id; \ No newline at end of file diff --git a/makefile b/makefile index 241764ae..36690ed8 100644 --- a/makefile +++ b/makefile @@ -65,6 +65,7 @@ run-workers: uv run python -m app.worker.photo_worker.main & \ uv run python -m app.worker.storage_cleaner.main & \ uv run python -m app.worker.email_worker.main & \ + uv run python -m app.worker.event_lifecycle.main & \ wait lint: diff --git a/migrations/sql/down/add_end_date_to_events.sql b/migrations/sql/down/add_end_date_to_events.sql new file mode 100644 index 00000000..708425cc --- /dev/null +++ b/migrations/sql/down/add_end_date_to_events.sql @@ -0,0 +1 @@ +ALTER TABLE events DROP COLUMN end_date; diff --git a/migrations/sql/up/add_end_date_to_events.sql b/migrations/sql/up/add_end_date_to_events.sql new file mode 100644 index 00000000..5f2f8e72 --- /dev/null +++ b/migrations/sql/up/add_end_date_to_events.sql @@ -0,0 +1 @@ +ALTER TABLE events ADD COLUMN end_date timestamptz; diff --git a/migrations/versions/afb0a93aa21e_add_end_date_to_events.py b/migrations/versions/afb0a93aa21e_add_end_date_to_events.py new file mode 100644 index 00000000..9d7fdcd6 --- /dev/null +++ b/migrations/versions/afb0a93aa21e_add_end_date_to_events.py @@ -0,0 +1,27 @@ +"""add_end_date_to_events + +Revision ID: afb0a93aa21e +Revises: 9ec59aeb192b +Create Date: 2026-08-25 01:34:59.327934 + +""" +from typing import Sequence, Union + +from migrations.helper import run_sql_up, run_sql_down + + +# revision identifiers, used by Alembic. +revision: str = 'afb0a93aa21e' +down_revision: Union[str, Sequence[str], None] = '9ec59aeb192b' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + """Upgrade schema.""" + run_sql_up("add_end_date_to_events") + + +def downgrade() -> None: + """Downgrade schema.""" + run_sql_down("add_end_date_to_events") From 22ff376011e864d37b51221a2b67a17726219d83 Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 01:59:14 +0100 Subject: [PATCH 05/29] feat: expose face_count on photo list endpoints Adds a per-photo face count (subquery over photo_faces) to ListUserPhotos and ListEventPhotosForUser, and surfaces it as face_count on GET /photos and GET /photos/event/{event_id}. Lets clients distinguish solo photos (one face) from group photos (multiple faces) without an extra request per photo. --- app/router/mobile/photos.py | 2 ++ app/service/user_photo.py | 16 +++++++++----- db/generated/photos.py | 44 ++++++++++++++++++++++++++++++++----- db/queries/photos.sql | 6 +++-- 4 files changed, 54 insertions(+), 14 deletions(-) diff --git a/app/router/mobile/photos.py b/app/router/mobile/photos.py index a9f6036c..5d0f357e 100644 --- a/app/router/mobile/photos.py +++ b/app/router/mobile/photos.py @@ -36,6 +36,7 @@ async def list_my_photos( "taken_at": p.taken_at.isoformat() if p.taken_at else None, "day_number": p.day_number, "created_at": p.created_at.isoformat(), + "face_count": p.face_count, } for p in photos ] @@ -72,6 +73,7 @@ async def list_event_photos( "taken_at": p.taken_at.isoformat() if p.taken_at else None, "day_number": p.day_number, "created_at": p.created_at.isoformat(), + "face_count": p.face_count, } for p in photos ], diff --git a/app/service/user_photo.py b/app/service/user_photo.py index f10a6303..58204960 100644 --- a/app/service/user_photo.py +++ b/app/service/user_photo.py @@ -10,8 +10,12 @@ from db.generated import photo_approvals as photo_approval_queries from db.generated import photo_faces as photo_face_queries from db.generated import photos as photo_queries -from db.generated.models import Photo -from db.generated.photos import ListEventPhotosForUserParams, ListUserPhotosParams +from db.generated.photos import ( + ListEventPhotosForUserParams, + ListEventPhotosForUserRow, + ListUserPhotosParams, + ListUserPhotosRow, +) class UserPhotoService: @@ -36,8 +40,8 @@ async def list_photos( sort: str = "desc", limit: int = 50, offset: int = 0, - ) -> list[Photo]: - photos: list[Photo] = [] + ) -> list[ListUserPhotosRow]: + photos: list[ListUserPhotosRow] = [] async for photo in self._photo_querier.list_user_photos( ListUserPhotosParams( user_id=user_id, @@ -58,8 +62,8 @@ async def list_event_photos( sort: str = "desc", limit: int = 50, offset: int = 0, - ) -> list[Photo]: - photos: list[Photo] = [] + ) -> list[ListEventPhotosForUserRow]: + photos: list[ListEventPhotosForUserRow] = [] async for photo in self._photo_querier.list_event_photos_for_user( ListEventPhotosForUserParams( user_id=user_id, diff --git a/db/generated/photos.py b/db/generated/photos.py index f757a9fb..599a46d1 100644 --- a/db/generated/photos.py +++ b/db/generated/photos.py @@ -62,7 +62,8 @@ class CreatePhotoParams: LIST_EVENT_PHOTOS_FOR_USER = """-- name: list_event_photos_for_user \\:many -SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at +SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, + (SELECT COUNT(*) FROM photo_faces pf2 WHERE pf2.photo_id = p.id)\\:\\:int AS face_count FROM photos p WHERE p.event_id = :p2 AND p.status = 'approved' @@ -94,8 +95,23 @@ class ListEventPhotosForUserParams: offset: int +@dataclasses.dataclass() +class ListEventPhotosForUserRow: + id: uuid.UUID + event_id: uuid.UUID + uploaded_by: Optional[uuid.UUID] + storage_key: str + taken_at: Optional[datetime.datetime] + day_number: Optional[int] + visibility: str + status: Any + created_at: datetime.datetime + face_count: int + + LIST_USER_PHOTOS = """-- name: list_user_photos \\:many -SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at +SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, + (SELECT COUNT(*) FROM photo_faces pf2 WHERE pf2.photo_id = p.id)\\:\\:int AS face_count FROM photos p WHERE ( EXISTS ( @@ -125,6 +141,20 @@ class ListUserPhotosParams: offset: int +@dataclasses.dataclass() +class ListUserPhotosRow: + id: uuid.UUID + event_id: uuid.UUID + uploaded_by: Optional[uuid.UUID] + storage_key: str + taken_at: Optional[datetime.datetime] + day_number: Optional[int] + visibility: str + status: Any + created_at: datetime.datetime + face_count: int + + UPDATE_PHOTO_STATUS = """-- name: update_photo_status \\:one UPDATE photos SET status = :p2 @@ -195,7 +225,7 @@ async def get_photo_by_id(self, *, id: uuid.UUID) -> Optional[models.Photo]: created_at=row[8], ) - async def list_event_photos_for_user(self, arg: ListEventPhotosForUserParams) -> AsyncIterator[models.Photo]: + async def list_event_photos_for_user(self, arg: ListEventPhotosForUserParams) -> AsyncIterator[ListEventPhotosForUserRow]: result = await self._conn.stream(sqlalchemy.text(LIST_EVENT_PHOTOS_FOR_USER), { "p1": arg.user_id, "p2": arg.event_id, @@ -204,7 +234,7 @@ async def list_event_photos_for_user(self, arg: ListEventPhotosForUserParams) -> "p5": arg.offset, }) async for row in result: - yield models.Photo( + yield ListEventPhotosForUserRow( id=row[0], event_id=row[1], uploaded_by=row[2], @@ -214,9 +244,10 @@ async def list_event_photos_for_user(self, arg: ListEventPhotosForUserParams) -> visibility=row[6], status=row[7], created_at=row[8], + face_count=row[9], ) - async def list_user_photos(self, arg: ListUserPhotosParams) -> AsyncIterator[models.Photo]: + async def list_user_photos(self, arg: ListUserPhotosParams) -> AsyncIterator[ListUserPhotosRow]: result = await self._conn.stream(sqlalchemy.text(LIST_USER_PHOTOS), { "p1": arg.user_id, "p2": arg.column_2, @@ -225,7 +256,7 @@ async def list_user_photos(self, arg: ListUserPhotosParams) -> AsyncIterator[mod "p5": arg.offset, }) async for row in result: - yield models.Photo( + yield ListUserPhotosRow( id=row[0], event_id=row[1], uploaded_by=row[2], @@ -235,6 +266,7 @@ async def list_user_photos(self, arg: ListUserPhotosParams) -> AsyncIterator[mod visibility=row[6], status=row[7], created_at=row[8], + face_count=row[9], ) async def update_photo_status(self, *, id: uuid.UUID, status: Any) -> Optional[models.Photo]: diff --git a/db/queries/photos.sql b/db/queries/photos.sql index 993c397a..3838e3c2 100644 --- a/db/queries/photos.sql +++ b/db/queries/photos.sql @@ -26,7 +26,8 @@ WHERE id = $1 RETURNING *; -- name: ListUserPhotos :many -SELECT p.* +SELECT p.*, + (SELECT COUNT(*) FROM photo_faces pf2 WHERE pf2.photo_id = p.id)::int AS face_count FROM photos p WHERE ( EXISTS ( @@ -46,7 +47,8 @@ ORDER BY LIMIT $4 OFFSET $5; -- name: ListEventPhotosForUser :many -SELECT p.* +SELECT p.*, + (SELECT COUNT(*) FROM photo_faces pf2 WHERE pf2.photo_id = p.id)::int AS face_count FROM photos p WHERE p.event_id = $2 AND p.status = 'approved' From 53fc8aa8d14e757fcd9f4b628eba8e4c86b15a8c Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 02:35:56 +0100 Subject: [PATCH 06/29] feat: add schema support for direct (non-Drive) bulk uploads --- .../sql/down/add_direct_upload_support.sql | 11 ++++++++ .../sql/up/add_direct_upload_support.sql | 11 ++++++++ .../5425a051d68c_add_direct_upload_support.py | 25 +++++++++++++++++++ 3 files changed, 47 insertions(+) create mode 100644 migrations/sql/down/add_direct_upload_support.sql create mode 100644 migrations/sql/up/add_direct_upload_support.sql create mode 100644 migrations/versions/5425a051d68c_add_direct_upload_support.py diff --git a/migrations/sql/down/add_direct_upload_support.sql b/migrations/sql/down/add_direct_upload_support.sql new file mode 100644 index 00000000..d9837913 --- /dev/null +++ b/migrations/sql/down/add_direct_upload_support.sql @@ -0,0 +1,11 @@ +ALTER TABLE upload_request_photos + ALTER COLUMN drive_file_id SET NOT NULL, + DROP COLUMN transfer_status, + DROP COLUMN source; + +ALTER TABLE upload_requests + DROP COLUMN source; + +ALTER TABLE upload_request_groups + ALTER COLUMN folder_id SET NOT NULL, + DROP COLUMN source; diff --git a/migrations/sql/up/add_direct_upload_support.sql b/migrations/sql/up/add_direct_upload_support.sql new file mode 100644 index 00000000..0319c12f --- /dev/null +++ b/migrations/sql/up/add_direct_upload_support.sql @@ -0,0 +1,11 @@ +ALTER TABLE upload_request_groups + ADD COLUMN source character varying(16) DEFAULT 'drive'::character varying NOT NULL, + ALTER COLUMN folder_id DROP NOT NULL; + +ALTER TABLE upload_requests + ADD COLUMN source character varying(16) DEFAULT 'drive'::character varying NOT NULL; + +ALTER TABLE upload_request_photos + ADD COLUMN source character varying(16) DEFAULT 'drive'::character varying NOT NULL, + ADD COLUMN transfer_status character varying(16) DEFAULT 'uploaded'::character varying NOT NULL, + ALTER COLUMN drive_file_id DROP NOT NULL; diff --git a/migrations/versions/5425a051d68c_add_direct_upload_support.py b/migrations/versions/5425a051d68c_add_direct_upload_support.py new file mode 100644 index 00000000..0e8807ce --- /dev/null +++ b/migrations/versions/5425a051d68c_add_direct_upload_support.py @@ -0,0 +1,25 @@ +"""add_direct_upload_support + +Revision ID: 5425a051d68c +Revises: afb0a93aa21e +Create Date: 2026-08-25 02:35:19.168255 + +""" +from typing import Sequence, Union + +from migrations.helper import run_sql_up, run_sql_down + + +# revision identifiers, used by Alembic. +revision: str = '5425a051d68c' +down_revision: Union[str, Sequence[str], None] = 'afb0a93aa21e' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + run_sql_up("add_direct_upload_support") + + +def downgrade() -> None: + run_sql_down("add_direct_upload_support") From 3e881b85847514923b882acca0bee6d09b112cac Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 02:36:52 +0100 Subject: [PATCH 07/29] feat: add SQLC queries for direct upload registration and transfer tracking --- db/generated/models.py | 8 +- db/generated/upload_request_groups.py | 80 ++++++++-- db/generated/upload_request_photos.py | 217 +++++++++++++++++++++++++- db/generated/upload_requests.py | 34 ++-- db/queries/upload_request_groups.sql | 13 +- db/queries/upload_request_photos.sql | 51 ++++++ db/queries/upload_requests.sql | 5 +- 7 files changed, 369 insertions(+), 39 deletions(-) diff --git a/db/generated/models.py b/db/generated/models.py index 86617af5..2d6bcf62 100644 --- a/db/generated/models.py +++ b/db/generated/models.py @@ -208,13 +208,14 @@ class UploadRequest: photo_count: int rejection_reason: Optional[str] group_id: Optional[uuid.UUID] + source: str @dataclasses.dataclass() class UploadRequestGroup: id: uuid.UUID event_id: uuid.UUID - folder_id: str + folder_id: Optional[str] requested_by: uuid.UUID approved_by: Optional[uuid.UUID] status: Any @@ -227,13 +228,14 @@ class UploadRequestGroup: processed_photo_count: int failed_photo_count: int error_message: Optional[str] + source: str @dataclasses.dataclass() class UploadRequestPhoto: id: uuid.UUID upload_request_id: uuid.UUID - drive_file_id: str + drive_file_id: Optional[str] file_name: str mime_type: str size_bytes: int @@ -244,6 +246,8 @@ class UploadRequestPhoto: visibility: str status: str created_at: datetime.datetime + source: str + transfer_status: str @dataclasses.dataclass() diff --git a/db/generated/upload_request_groups.py b/db/generated/upload_request_groups.py index 039b1f05..ceb91cf5 100644 --- a/db/generated/upload_request_groups.py +++ b/db/generated/upload_request_groups.py @@ -20,7 +20,7 @@ rejection_reason = NULL WHERE id = :p1 AND status = 'pending' -RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source """ @@ -33,7 +33,7 @@ failed_photo_count = :p5, error_message = NULL WHERE id = :p1 -RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source """ @@ -52,21 +52,25 @@ class CompleteUploadRequestGroupProcessingParams: folder_id, requested_by, total_photo_count, - batch_count + batch_count, + source, + processing_status ) VALUES ( - :p1, :p2, :p3, :p4, :p5 + :p1, :p2, :p3, :p4, :p5, :p6, :p7 ) -RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source """ @dataclasses.dataclass() class CreateUploadRequestGroupParams: event_id: uuid.UUID - folder_id: str + folder_id: Optional[str] requested_by: uuid.UUID total_photo_count: int batch_count: int + source: str + processing_status: str DELETE_UPLOAD_REQUEST_GROUP = """-- name: delete_upload_request_group \\:exec @@ -84,7 +88,7 @@ class CreateUploadRequestGroupParams: failed_photo_count = :p5, error_message = :p6 WHERE id = :p1 -RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source """ @@ -99,21 +103,30 @@ class FailUploadRequestGroupProcessingParams: GET_UPLOAD_REQUEST_GROUP_BY_ID = """-- name: get_upload_request_group_by_id \\:one -SELECT id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +SELECT id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source FROM upload_request_groups WHERE id = :p1 """ +INCREMENT_UPLOAD_REQUEST_GROUP_COUNTS = """-- name: increment_upload_request_group_counts \\:one +UPDATE upload_request_groups +SET total_photo_count = total_photo_count + :p2, + batch_count = batch_count + 1 +WHERE id = :p1 +RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source +""" + + LIST_UPLOAD_REQUEST_GROUPS = """-- name: list_upload_request_groups \\:many -SELECT id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +SELECT id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source FROM upload_request_groups ORDER BY created_at DESC """ LIST_UPLOAD_REQUEST_GROUPS_BY_REQUESTER = """-- name: list_upload_request_groups_by_requester \\:many -SELECT id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +SELECT id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source FROM upload_request_groups WHERE requested_by = :p1 ORDER BY created_at DESC @@ -121,7 +134,7 @@ class FailUploadRequestGroupProcessingParams: LIST_UPLOAD_REQUEST_GROUPS_BY_REQUESTER_AND_STATUS = """-- name: list_upload_request_groups_by_requester_and_status \\:many -SELECT id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +SELECT id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source FROM upload_request_groups WHERE requested_by = :p1 AND status = :p2 @@ -130,7 +143,7 @@ class FailUploadRequestGroupProcessingParams: LIST_UPLOAD_REQUEST_GROUPS_BY_STATUS = """-- name: list_upload_request_groups_by_status \\:many -SELECT id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +SELECT id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source FROM upload_request_groups WHERE status = :p1 ORDER BY created_at DESC @@ -145,7 +158,7 @@ class FailUploadRequestGroupProcessingParams: rejection_reason = :p3 WHERE id = :p1 AND status = 'pending' -RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source """ @@ -155,7 +168,7 @@ class FailUploadRequestGroupProcessingParams: error_message = NULL WHERE id = :p1 AND processing_status = 'pending' -RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source """ @@ -166,7 +179,7 @@ class FailUploadRequestGroupProcessingParams: processed_photo_count = :p4, failed_photo_count = :p5 WHERE id = :p1 -RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message +RETURNING id, event_id, folder_id, requested_by, approved_by, status, total_photo_count, batch_count, created_at, approved_at, rejection_reason, processing_status, processed_photo_count, failed_photo_count, error_message, source """ @@ -203,6 +216,7 @@ async def approve_upload_request_group(self, *, id: uuid.UUID, approved_by: Opti processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], ) async def complete_upload_request_group_processing(self, arg: CompleteUploadRequestGroupProcessingParams) -> Optional[models.UploadRequestGroup]: @@ -231,6 +245,7 @@ async def complete_upload_request_group_processing(self, arg: CompleteUploadRequ processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], ) async def create_upload_request_group(self, arg: CreateUploadRequestGroupParams) -> Optional[models.UploadRequestGroup]: @@ -240,6 +255,8 @@ async def create_upload_request_group(self, arg: CreateUploadRequestGroupParams) "p3": arg.requested_by, "p4": arg.total_photo_count, "p5": arg.batch_count, + "p6": arg.source, + "p7": arg.processing_status, })).first() if row is None: return None @@ -259,6 +276,7 @@ async def create_upload_request_group(self, arg: CreateUploadRequestGroupParams) processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], ) async def delete_upload_request_group(self, *, id: uuid.UUID) -> None: @@ -291,6 +309,7 @@ async def fail_upload_request_group_processing(self, arg: FailUploadRequestGroup processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], ) async def get_upload_request_group_by_id(self, *, id: uuid.UUID) -> Optional[models.UploadRequestGroup]: @@ -313,6 +332,30 @@ async def get_upload_request_group_by_id(self, *, id: uuid.UUID) -> Optional[mod processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], + ) + + async def increment_upload_request_group_counts(self, *, id: uuid.UUID, total_photo_count: int) -> Optional[models.UploadRequestGroup]: + row = (await self._conn.execute(sqlalchemy.text(INCREMENT_UPLOAD_REQUEST_GROUP_COUNTS), {"p1": id, "p2": total_photo_count})).first() + if row is None: + return None + return models.UploadRequestGroup( + id=row[0], + event_id=row[1], + folder_id=row[2], + requested_by=row[3], + approved_by=row[4], + status=row[5], + total_photo_count=row[6], + batch_count=row[7], + created_at=row[8], + approved_at=row[9], + rejection_reason=row[10], + processing_status=row[11], + processed_photo_count=row[12], + failed_photo_count=row[13], + error_message=row[14], + source=row[15], ) async def list_upload_request_groups(self) -> AsyncIterator[models.UploadRequestGroup]: @@ -334,6 +377,7 @@ async def list_upload_request_groups(self) -> AsyncIterator[models.UploadRequest processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], ) async def list_upload_request_groups_by_requester(self, *, requested_by: uuid.UUID) -> AsyncIterator[models.UploadRequestGroup]: @@ -355,6 +399,7 @@ async def list_upload_request_groups_by_requester(self, *, requested_by: uuid.UU processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], ) async def list_upload_request_groups_by_requester_and_status(self, *, requested_by: uuid.UUID, status: Any) -> AsyncIterator[models.UploadRequestGroup]: @@ -376,6 +421,7 @@ async def list_upload_request_groups_by_requester_and_status(self, *, requested_ processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], ) async def list_upload_request_groups_by_status(self, *, status: Any) -> AsyncIterator[models.UploadRequestGroup]: @@ -397,6 +443,7 @@ async def list_upload_request_groups_by_status(self, *, status: Any) -> AsyncIte processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], ) async def reject_upload_request_group(self, *, id: uuid.UUID, approved_by: Optional[uuid.UUID], rejection_reason: Optional[str]) -> Optional[models.UploadRequestGroup]: @@ -419,6 +466,7 @@ async def reject_upload_request_group(self, *, id: uuid.UUID, approved_by: Optio processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], ) async def start_upload_request_group_processing(self, *, id: uuid.UUID) -> Optional[models.UploadRequestGroup]: @@ -441,6 +489,7 @@ async def start_upload_request_group_processing(self, *, id: uuid.UUID) -> Optio processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], ) async def update_upload_request_group_import_progress(self, arg: UpdateUploadRequestGroupImportProgressParams) -> Optional[models.UploadRequestGroup]: @@ -469,4 +518,5 @@ async def update_upload_request_group_import_progress(self, arg: UpdateUploadReq processed_photo_count=row[12], failed_photo_count=row[13], error_message=row[14], + source=row[15], ) diff --git a/db/generated/upload_request_photos.py b/db/generated/upload_request_photos.py index 1cd3ebb4..ba33dee1 100644 --- a/db/generated/upload_request_photos.py +++ b/db/generated/upload_request_photos.py @@ -13,6 +13,50 @@ from db.generated import models +CONFIRM_UPLOAD_REQUEST_PHOTO_TRANSFER = """-- name: confirm_upload_request_photo_transfer \\:one +UPDATE upload_request_photos +SET transfer_status = 'uploaded', + size_bytes = :p2, + mime_type = :p3 +WHERE id = :p1 + AND transfer_status IN ('pending_upload', 'failed') +RETURNING id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at, source, transfer_status +""" + + +CREATE_DIRECT_UPLOAD_REQUEST_PHOTO = """-- name: create_direct_upload_request_photo \\:one +INSERT INTO upload_request_photos ( + upload_request_id, + drive_file_id, + file_name, + mime_type, + size_bytes, + staging_storage_key, + taken_at, + day_number, + visibility, + status, + source, + transfer_status +) VALUES ( + :p1, NULL, :p2, :p3, :p4, :p5, :p6, :p7, :p8, 'staged', 'direct', 'pending_upload' +) +RETURNING id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at, source, transfer_status +""" + + +@dataclasses.dataclass() +class CreateDirectUploadRequestPhotoParams: + upload_request_id: uuid.UUID + file_name: str + mime_type: str + size_bytes: int + staging_storage_key: str + taken_at: Optional[datetime.datetime] + day_number: Optional[int] + visibility: str + + CREATE_UPLOAD_REQUEST_PHOTO = """-- name: create_upload_request_photo \\:one INSERT INTO upload_request_photos ( upload_request_id, @@ -28,14 +72,14 @@ ) VALUES ( :p1, :p2, :p3, :p4, :p5, :p6, :p7, :p8, :p9, :p10 ) -RETURNING id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at +RETURNING id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at, source, transfer_status """ @dataclasses.dataclass() class CreateUploadRequestPhotoParams: upload_request_id: uuid.UUID - drive_file_id: str + drive_file_id: Optional[str] file_name: str mime_type: str size_bytes: int @@ -52,15 +96,35 @@ class CreateUploadRequestPhotoParams: """ +FAIL_UPLOAD_REQUEST_PHOTO_TRANSFER = """-- name: fail_upload_request_photo_transfer \\:one +UPDATE upload_request_photos +SET transfer_status = 'failed' +WHERE id = :p1 + AND transfer_status IN ('pending_upload', 'failed') +RETURNING id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at, source, transfer_status +""" + + GET_UPLOAD_REQUEST_PHOTO_BY_ID = """-- name: get_upload_request_photo_by_id \\:one -SELECT id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at +SELECT id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at, source, transfer_status FROM upload_request_photos WHERE id = :p1 """ +LIST_STALE_PENDING_TRANSFER_PHOTOS = """-- name: list_stale_pending_transfer_photos \\:many +SELECT id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at, source, transfer_status +FROM upload_request_photos +WHERE source = 'direct' + AND transfer_status = 'pending_upload' + AND created_at <= NOW() - (:p1 || ' minutes')\\:\\:interval +ORDER BY created_at ASC +LIMIT 500 +""" + + LIST_UPLOAD_REQUEST_PHOTOS_BY_UPLOAD_REQUEST_ID = """-- name: list_upload_request_photos_by_upload_request_id \\:many -SELECT id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at +SELECT id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at, source, transfer_status FROM upload_request_photos WHERE upload_request_id = :p1 ORDER BY created_at ASC @@ -68,19 +132,28 @@ class CreateUploadRequestPhotoParams: LIST_UPLOAD_REQUEST_PHOTOS_BY_UPLOAD_REQUEST_IDS = """-- name: list_upload_request_photos_by_upload_request_ids \\:many -SELECT id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at +SELECT id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at, source, transfer_status FROM upload_request_photos WHERE upload_request_id = ANY(:p1\\:\\:uuid[]) ORDER BY created_at ASC """ +RESET_UPLOAD_REQUEST_PHOTO_TRANSFER_TO_PENDING = """-- name: reset_upload_request_photo_transfer_to_pending \\:one +UPDATE upload_request_photos +SET transfer_status = 'pending_upload' +WHERE id = :p1 + AND transfer_status IN ('pending_upload', 'failed') +RETURNING id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at, source, transfer_status +""" + + UPDATE_UPLOAD_REQUEST_PHOTO_APPROVAL = """-- name: update_upload_request_photo_approval \\:one UPDATE upload_request_photos SET status = :p2, final_storage_key = :p3 WHERE id = :p1 -RETURNING id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at +RETURNING id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at, source, transfer_status """ @@ -88,7 +161,7 @@ class CreateUploadRequestPhotoParams: UPDATE upload_request_photos SET status = :p2 WHERE upload_request_id = :p1 -RETURNING id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at +RETURNING id, upload_request_id, drive_file_id, file_name, mime_type, size_bytes, staging_storage_key, final_storage_key, taken_at, day_number, visibility, status, created_at, source, transfer_status """ @@ -96,6 +169,59 @@ class AsyncQuerier: def __init__(self, conn: sqlalchemy.ext.asyncio.AsyncConnection): self._conn = conn + async def confirm_upload_request_photo_transfer(self, *, id: uuid.UUID, size_bytes: int, mime_type: str) -> Optional[models.UploadRequestPhoto]: + row = (await self._conn.execute(sqlalchemy.text(CONFIRM_UPLOAD_REQUEST_PHOTO_TRANSFER), {"p1": id, "p2": size_bytes, "p3": mime_type})).first() + if row is None: + return None + return models.UploadRequestPhoto( + id=row[0], + upload_request_id=row[1], + drive_file_id=row[2], + file_name=row[3], + mime_type=row[4], + size_bytes=row[5], + staging_storage_key=row[6], + final_storage_key=row[7], + taken_at=row[8], + day_number=row[9], + visibility=row[10], + status=row[11], + created_at=row[12], + source=row[13], + transfer_status=row[14], + ) + + async def create_direct_upload_request_photo(self, arg: CreateDirectUploadRequestPhotoParams) -> Optional[models.UploadRequestPhoto]: + row = (await self._conn.execute(sqlalchemy.text(CREATE_DIRECT_UPLOAD_REQUEST_PHOTO), { + "p1": arg.upload_request_id, + "p2": arg.file_name, + "p3": arg.mime_type, + "p4": arg.size_bytes, + "p5": arg.staging_storage_key, + "p6": arg.taken_at, + "p7": arg.day_number, + "p8": arg.visibility, + })).first() + if row is None: + return None + return models.UploadRequestPhoto( + id=row[0], + upload_request_id=row[1], + drive_file_id=row[2], + file_name=row[3], + mime_type=row[4], + size_bytes=row[5], + staging_storage_key=row[6], + final_storage_key=row[7], + taken_at=row[8], + day_number=row[9], + visibility=row[10], + status=row[11], + created_at=row[12], + source=row[13], + transfer_status=row[14], + ) + async def create_upload_request_photo(self, arg: CreateUploadRequestPhotoParams) -> Optional[models.UploadRequestPhoto]: row = (await self._conn.execute(sqlalchemy.text(CREATE_UPLOAD_REQUEST_PHOTO), { "p1": arg.upload_request_id, @@ -125,11 +251,35 @@ async def create_upload_request_photo(self, arg: CreateUploadRequestPhotoParams) visibility=row[10], status=row[11], created_at=row[12], + source=row[13], + transfer_status=row[14], ) async def delete_upload_request_photos_by_upload_request_id(self, *, upload_request_id: uuid.UUID) -> None: await self._conn.execute(sqlalchemy.text(DELETE_UPLOAD_REQUEST_PHOTOS_BY_UPLOAD_REQUEST_ID), {"p1": upload_request_id}) + async def fail_upload_request_photo_transfer(self, *, id: uuid.UUID) -> Optional[models.UploadRequestPhoto]: + row = (await self._conn.execute(sqlalchemy.text(FAIL_UPLOAD_REQUEST_PHOTO_TRANSFER), {"p1": id})).first() + if row is None: + return None + return models.UploadRequestPhoto( + id=row[0], + upload_request_id=row[1], + drive_file_id=row[2], + file_name=row[3], + mime_type=row[4], + size_bytes=row[5], + staging_storage_key=row[6], + final_storage_key=row[7], + taken_at=row[8], + day_number=row[9], + visibility=row[10], + status=row[11], + created_at=row[12], + source=row[13], + transfer_status=row[14], + ) + async def get_upload_request_photo_by_id(self, *, id: uuid.UUID) -> Optional[models.UploadRequestPhoto]: row = (await self._conn.execute(sqlalchemy.text(GET_UPLOAD_REQUEST_PHOTO_BY_ID), {"p1": id})).first() if row is None: @@ -148,8 +298,31 @@ async def get_upload_request_photo_by_id(self, *, id: uuid.UUID) -> Optional[mod visibility=row[10], status=row[11], created_at=row[12], + source=row[13], + transfer_status=row[14], ) + async def list_stale_pending_transfer_photos(self, *, dollar_1: Optional[str]) -> AsyncIterator[models.UploadRequestPhoto]: + result = await self._conn.stream(sqlalchemy.text(LIST_STALE_PENDING_TRANSFER_PHOTOS), {"p1": dollar_1}) + async for row in result: + yield models.UploadRequestPhoto( + id=row[0], + upload_request_id=row[1], + drive_file_id=row[2], + file_name=row[3], + mime_type=row[4], + size_bytes=row[5], + staging_storage_key=row[6], + final_storage_key=row[7], + taken_at=row[8], + day_number=row[9], + visibility=row[10], + status=row[11], + created_at=row[12], + source=row[13], + transfer_status=row[14], + ) + async def list_upload_request_photos_by_upload_request_id(self, *, upload_request_id: uuid.UUID) -> AsyncIterator[models.UploadRequestPhoto]: result = await self._conn.stream(sqlalchemy.text(LIST_UPLOAD_REQUEST_PHOTOS_BY_UPLOAD_REQUEST_ID), {"p1": upload_request_id}) async for row in result: @@ -167,6 +340,8 @@ async def list_upload_request_photos_by_upload_request_id(self, *, upload_reques visibility=row[10], status=row[11], created_at=row[12], + source=row[13], + transfer_status=row[14], ) async def list_upload_request_photos_by_upload_request_ids(self, *, dollar_1: List[uuid.UUID]) -> AsyncIterator[models.UploadRequestPhoto]: @@ -186,8 +361,32 @@ async def list_upload_request_photos_by_upload_request_ids(self, *, dollar_1: Li visibility=row[10], status=row[11], created_at=row[12], + source=row[13], + transfer_status=row[14], ) + async def reset_upload_request_photo_transfer_to_pending(self, *, id: uuid.UUID) -> Optional[models.UploadRequestPhoto]: + row = (await self._conn.execute(sqlalchemy.text(RESET_UPLOAD_REQUEST_PHOTO_TRANSFER_TO_PENDING), {"p1": id})).first() + if row is None: + return None + return models.UploadRequestPhoto( + id=row[0], + upload_request_id=row[1], + drive_file_id=row[2], + file_name=row[3], + mime_type=row[4], + size_bytes=row[5], + staging_storage_key=row[6], + final_storage_key=row[7], + taken_at=row[8], + day_number=row[9], + visibility=row[10], + status=row[11], + created_at=row[12], + source=row[13], + transfer_status=row[14], + ) + async def update_upload_request_photo_approval(self, *, id: uuid.UUID, status: str, final_storage_key: Optional[str]) -> Optional[models.UploadRequestPhoto]: row = (await self._conn.execute(sqlalchemy.text(UPDATE_UPLOAD_REQUEST_PHOTO_APPROVAL), {"p1": id, "p2": status, "p3": final_storage_key})).first() if row is None: @@ -206,6 +405,8 @@ async def update_upload_request_photo_approval(self, *, id: uuid.UUID, status: s visibility=row[10], status=row[11], created_at=row[12], + source=row[13], + transfer_status=row[14], ) async def update_upload_request_photo_status_by_upload_request_id(self, *, upload_request_id: uuid.UUID, status: str) -> AsyncIterator[models.UploadRequestPhoto]: @@ -225,4 +426,6 @@ async def update_upload_request_photo_status_by_upload_request_id(self, *, uploa visibility=row[10], status=row[11], created_at=row[12], + source=row[13], + transfer_status=row[14], ) diff --git a/db/generated/upload_requests.py b/db/generated/upload_requests.py index b0da8bb0..de8af04e 100644 --- a/db/generated/upload_requests.py +++ b/db/generated/upload_requests.py @@ -20,7 +20,7 @@ rejection_reason = NULL WHERE id = :p1 AND status = 'pending' -RETURNING id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id +RETURNING id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id, source """ @@ -30,11 +30,12 @@ group_id, drive_file_id, requested_by, - photo_count + photo_count, + source ) VALUES ( - :p1, :p2, :p3, :p4, :p5 + :p1, :p2, :p3, :p4, :p5, :p6 ) -RETURNING id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id +RETURNING id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id, source """ @@ -45,6 +46,7 @@ class CreateUploadRequestParams: drive_file_id: Optional[str] requested_by: uuid.UUID photo_count: int + source: str DELETE_UPLOAD_REQUEST = """-- name: delete_upload_request \\:exec @@ -54,21 +56,21 @@ class CreateUploadRequestParams: GET_UPLOAD_REQUEST_BY_ID = """-- name: get_upload_request_by_id \\:one -SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id +SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id, source FROM upload_requests WHERE id = :p1 """ LIST_UPLOAD_REQUESTS = """-- name: list_upload_requests \\:many -SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id +SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id, source FROM upload_requests ORDER BY created_at DESC """ LIST_UPLOAD_REQUESTS_BY_GROUP_ID = """-- name: list_upload_requests_by_group_id \\:many -SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id +SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id, source FROM upload_requests WHERE group_id = :p1 ORDER BY created_at ASC @@ -76,7 +78,7 @@ class CreateUploadRequestParams: LIST_UPLOAD_REQUESTS_BY_REQUESTER = """-- name: list_upload_requests_by_requester \\:many -SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id +SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id, source FROM upload_requests WHERE requested_by = :p1 ORDER BY created_at DESC @@ -84,7 +86,7 @@ class CreateUploadRequestParams: LIST_UPLOAD_REQUESTS_BY_REQUESTER_AND_STATUS = """-- name: list_upload_requests_by_requester_and_status \\:many -SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id +SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id, source FROM upload_requests WHERE requested_by = :p1 AND status = :p2 @@ -93,7 +95,7 @@ class CreateUploadRequestParams: LIST_UPLOAD_REQUESTS_BY_STATUS = """-- name: list_upload_requests_by_status \\:many -SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id +SELECT id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id, source FROM upload_requests WHERE status = :p1 ORDER BY created_at DESC @@ -108,7 +110,7 @@ class CreateUploadRequestParams: rejection_reason = :p3 WHERE id = :p1 AND status = 'pending' -RETURNING id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id +RETURNING id, event_id, drive_file_id, requested_by, approved_by, status, created_at, approved_at, photo_count, rejection_reason, group_id, source """ @@ -132,6 +134,7 @@ async def approve_upload_request(self, *, id: uuid.UUID, approved_by: Optional[u photo_count=row[8], rejection_reason=row[9], group_id=row[10], + source=row[11], ) async def create_upload_request(self, arg: CreateUploadRequestParams) -> Optional[models.UploadRequest]: @@ -141,6 +144,7 @@ async def create_upload_request(self, arg: CreateUploadRequestParams) -> Optiona "p3": arg.drive_file_id, "p4": arg.requested_by, "p5": arg.photo_count, + "p6": arg.source, })).first() if row is None: return None @@ -156,6 +160,7 @@ async def create_upload_request(self, arg: CreateUploadRequestParams) -> Optiona photo_count=row[8], rejection_reason=row[9], group_id=row[10], + source=row[11], ) async def delete_upload_request(self, *, id: uuid.UUID) -> None: @@ -177,6 +182,7 @@ async def get_upload_request_by_id(self, *, id: uuid.UUID) -> Optional[models.Up photo_count=row[8], rejection_reason=row[9], group_id=row[10], + source=row[11], ) async def list_upload_requests(self) -> AsyncIterator[models.UploadRequest]: @@ -194,6 +200,7 @@ async def list_upload_requests(self) -> AsyncIterator[models.UploadRequest]: photo_count=row[8], rejection_reason=row[9], group_id=row[10], + source=row[11], ) async def list_upload_requests_by_group_id(self, *, group_id: Optional[uuid.UUID]) -> AsyncIterator[models.UploadRequest]: @@ -211,6 +218,7 @@ async def list_upload_requests_by_group_id(self, *, group_id: Optional[uuid.UUID photo_count=row[8], rejection_reason=row[9], group_id=row[10], + source=row[11], ) async def list_upload_requests_by_requester(self, *, requested_by: uuid.UUID) -> AsyncIterator[models.UploadRequest]: @@ -228,6 +236,7 @@ async def list_upload_requests_by_requester(self, *, requested_by: uuid.UUID) -> photo_count=row[8], rejection_reason=row[9], group_id=row[10], + source=row[11], ) async def list_upload_requests_by_requester_and_status(self, *, requested_by: uuid.UUID, status: Any) -> AsyncIterator[models.UploadRequest]: @@ -245,6 +254,7 @@ async def list_upload_requests_by_requester_and_status(self, *, requested_by: uu photo_count=row[8], rejection_reason=row[9], group_id=row[10], + source=row[11], ) async def list_upload_requests_by_status(self, *, status: Any) -> AsyncIterator[models.UploadRequest]: @@ -262,6 +272,7 @@ async def list_upload_requests_by_status(self, *, status: Any) -> AsyncIterator[ photo_count=row[8], rejection_reason=row[9], group_id=row[10], + source=row[11], ) async def reject_upload_request(self, *, id: uuid.UUID, approved_by: Optional[uuid.UUID], rejection_reason: Optional[str]) -> Optional[models.UploadRequest]: @@ -280,4 +291,5 @@ async def reject_upload_request(self, *, id: uuid.UUID, approved_by: Optional[uu photo_count=row[8], rejection_reason=row[9], group_id=row[10], + source=row[11], ) diff --git a/db/queries/upload_request_groups.sql b/db/queries/upload_request_groups.sql index 7c800f14..1dfe97ba 100644 --- a/db/queries/upload_request_groups.sql +++ b/db/queries/upload_request_groups.sql @@ -4,12 +4,21 @@ INSERT INTO upload_request_groups ( folder_id, requested_by, total_photo_count, - batch_count + batch_count, + source, + processing_status ) VALUES ( - $1, $2, $3, $4, $5 + $1, $2, $3, $4, $5, $6, $7 ) RETURNING *; +-- name: IncrementUploadRequestGroupCounts :one +UPDATE upload_request_groups +SET total_photo_count = total_photo_count + $2, + batch_count = batch_count + 1 +WHERE id = $1 +RETURNING *; + -- name: GetUploadRequestGroupById :one SELECT * FROM upload_request_groups diff --git a/db/queries/upload_request_photos.sql b/db/queries/upload_request_photos.sql index f78ab850..880f1afe 100644 --- a/db/queries/upload_request_photos.sql +++ b/db/queries/upload_request_photos.sql @@ -48,3 +48,54 @@ RETURNING *; -- name: DeleteUploadRequestPhotosByUploadRequestId :exec DELETE FROM upload_request_photos WHERE upload_request_id = $1; + +-- name: CreateDirectUploadRequestPhoto :one +INSERT INTO upload_request_photos ( + upload_request_id, + drive_file_id, + file_name, + mime_type, + size_bytes, + staging_storage_key, + taken_at, + day_number, + visibility, + status, + source, + transfer_status +) VALUES ( + $1, NULL, $2, $3, $4, $5, $6, $7, $8, 'staged', 'direct', 'pending_upload' +) +RETURNING *; + +-- name: ConfirmUploadRequestPhotoTransfer :one +UPDATE upload_request_photos +SET transfer_status = 'uploaded', + size_bytes = $2, + mime_type = $3 +WHERE id = $1 + AND transfer_status IN ('pending_upload', 'failed') +RETURNING *; + +-- name: FailUploadRequestPhotoTransfer :one +UPDATE upload_request_photos +SET transfer_status = 'failed' +WHERE id = $1 + AND transfer_status IN ('pending_upload', 'failed') +RETURNING *; + +-- name: ResetUploadRequestPhotoTransferToPending :one +UPDATE upload_request_photos +SET transfer_status = 'pending_upload' +WHERE id = $1 + AND transfer_status IN ('pending_upload', 'failed') +RETURNING *; + +-- name: ListStalePendingTransferPhotos :many +SELECT * +FROM upload_request_photos +WHERE source = 'direct' + AND transfer_status = 'pending_upload' + AND created_at <= NOW() - ($1 || ' minutes')::interval +ORDER BY created_at ASC +LIMIT 500; diff --git a/db/queries/upload_requests.sql b/db/queries/upload_requests.sql index 31eb3739..2fd571bc 100644 --- a/db/queries/upload_requests.sql +++ b/db/queries/upload_requests.sql @@ -4,9 +4,10 @@ INSERT INTO upload_requests ( group_id, drive_file_id, requested_by, - photo_count + photo_count, + source ) VALUES ( - $1, $2, $3, $4, $5 + $1, $2, $3, $4, $5, $6 ) RETURNING *; From 0e7a9470f202367cae5cb2b567986d024c41f12a Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 02:41:27 +0100 Subject: [PATCH 08/29] feat: add presigned PUT/stat support; fix Drive-flow params for new source/transfer_status columns --- app/infra/minio.py | 27 +++++++++++++++++ app/service/staged_upload_storage.py | 21 +++++++++++++- app/service/upload_requests.py | 3 ++ tests/unit/test_minio.py | 43 ++++++++++++++++++++++++++++ tests/unit/test_upload_requests.py | 17 ++++++----- 5 files changed, 103 insertions(+), 8 deletions(-) diff --git a/app/infra/minio.py b/app/infra/minio.py index e6249dae..9a666c5c 100644 --- a/app/infra/minio.py +++ b/app/infra/minio.py @@ -2,6 +2,8 @@ import random import string import uuid +from dataclasses import dataclass +from datetime import timedelta from fastapi import UploadFile from miniopy_async.commonconfig import CopySource from miniopy_async.error import S3Error @@ -36,6 +38,12 @@ async def init_minio_client( if not await Bucket.client.bucket_exists(bucket_name): await Bucket.client.make_bucket(bucket_name) +@dataclass(frozen=True) +class ObjectStat: + size: int + content_type: str + + class Bucket: bucket_name: str file_prefix: str @@ -130,6 +138,25 @@ async def copy(self, *, source_object_name: str, target_object_name: str) -> str ) return target_object_name + async def presigned_put_url(self, object_name: str, *, expires_seconds: int) -> str: + return await self.client.presigned_put_object( + bucket_name=self.bucket_name, + object_name=self._object_path(object_name), + expires=timedelta(seconds=expires_seconds), + ) + + async def stat(self, object_name: str) -> ObjectStat | None: + try: + result = await self.client.stat_object( + bucket_name=self.bucket_name, + object_name=self._object_path(object_name), + ) + except S3Error as e: + if e.code == "NoSuchKey": + return None + raise + return ObjectStat(size=result.size or 0, content_type=result.content_type or DEFAULT_CONTENT_TYPE) + image_ext_content_type_map = { "apng": ["image/apng"], "avif": ["image/avif"], diff --git a/app/service/staged_upload_storage.py b/app/service/staged_upload_storage.py index 813fa2bf..e8f9d97a 100644 --- a/app/service/staged_upload_storage.py +++ b/app/service/staged_upload_storage.py @@ -5,7 +5,7 @@ import uuid from app.core.exceptions import AppException -from app.infra.minio import Bucket, IMAGES_BUCKET_NAME +from app.infra.minio import Bucket, IMAGES_BUCKET_NAME, ObjectStat @dataclass(frozen=True) @@ -100,3 +100,22 @@ async def delete_storage_key(self, storage_key: str) -> None: async def get_preview(self, storage_key: str) -> PreviewObject: data, file_name, content_type = await self.bucket.get(storage_key) return PreviewObject(data=data, file_name=file_name, content_type=content_type) + + async def create_presigned_staging_upload( + self, + *, + upload_request_id: uuid.UUID, + photo_id: uuid.UUID, + file_name: str, + expires_seconds: int, + ) -> tuple[str, str]: + storage_key = self.build_staging_key( + upload_request_id=upload_request_id, + photo_id=photo_id, + file_name=file_name, + ) + url = await self.bucket.presigned_put_url(storage_key, expires_seconds=expires_seconds) + return storage_key, url + + async def stat_staging_object(self, storage_key: str) -> ObjectStat | None: + return await self.bucket.stat(storage_key) diff --git a/app/service/upload_requests.py b/app/service/upload_requests.py index 0ab404d6..1d19eef1 100644 --- a/app/service/upload_requests.py +++ b/app/service/upload_requests.py @@ -282,6 +282,7 @@ async def _create_request_with_access_token( drive_file_id=None, requested_by=requested_by.id, photo_count=len(photos), + source="drive", ) ) except IntegrityError as exc: @@ -572,6 +573,8 @@ async def create_group_from_folder( requested_by=requested_by.id, total_photo_count=0, batch_count=0, + source="drive", + processing_status="pending", ) ) except IntegrityError as exc: diff --git a/tests/unit/test_minio.py b/tests/unit/test_minio.py index 66115611..1587c27f 100644 --- a/tests/unit/test_minio.py +++ b/tests/unit/test_minio.py @@ -8,6 +8,7 @@ from app.infra.minio import ( Bucket, ImageBucket, + ObjectStat, WaSimBucket, init_minio_client, ) @@ -200,3 +201,45 @@ async def test_wa_sim_bucket_auto_name(mock_minio_client, mock_upload_file): # WaSimBucket generates 16 digit string assert len(object_name) == 16 assert object_name.isdigit() + + +@pytest.mark.asyncio +async def test_presigned_put_url_calls_client_with_expiry(mock_minio_client): + Bucket.client = mock_minio_client + mock_minio_client.presigned_put_object = AsyncMock(return_value="https://minio.local/signed") + bucket = Bucket("test_bucket", "") + + url = await bucket.presigned_put_url("staging/foo.jpg", expires_seconds=1800) + + assert url == "https://minio.local/signed" + mock_minio_client.presigned_put_object.assert_awaited_once() + kwargs = mock_minio_client.presigned_put_object.call_args[1] + assert kwargs["object_name"] == "staging/foo.jpg" + assert kwargs["expires"].total_seconds() == 1800 + + +@pytest.mark.asyncio +async def test_stat_returns_object_stat_when_present(mock_minio_client): + Bucket.client = mock_minio_client + stat_result = MagicMock(size=12345, content_type="image/jpeg") + mock_minio_client.stat_object = AsyncMock(return_value=stat_result) + bucket = Bucket("test_bucket", "") + + result = await bucket.stat("staging/foo.jpg") + + assert result == ObjectStat(size=12345, content_type="image/jpeg") + + +@pytest.mark.asyncio +async def test_stat_returns_none_when_object_missing(mock_minio_client): + Bucket.client = mock_minio_client + error = S3Error( + code="NoSuchKey", message="not found", resource="", request_id="", + host_id="", response=MagicMock(), + ) + mock_minio_client.stat_object = AsyncMock(side_effect=error) + bucket = Bucket("test_bucket", "") + + result = await bucket.stat("staging/missing.jpg") + + assert result is None diff --git a/tests/unit/test_upload_requests.py b/tests/unit/test_upload_requests.py index 691c30c8..dd31c391 100644 --- a/tests/unit/test_upload_requests.py +++ b/tests/unit/test_upload_requests.py @@ -122,6 +122,7 @@ async def test_create_request_success( rejection_reason=None, created_at=datetime.now(timezone.utc), approved_at=None, + source="drive", ) mock_upload_request_photo_querier.create_upload_request_photo.return_value = UploadRequestPhoto( @@ -138,6 +139,8 @@ async def test_create_request_success( visibility="public", status="staged", created_at=datetime.now(timezone.utc), + source="drive", + transfer_status="uploaded", ) photos = [ @@ -191,7 +194,7 @@ async def test_create_request_duplicate_conflict( request_id = uuid.uuid4() mock_upload_request_querier.create_upload_request.return_value = UploadRequest( - id=request_id, event_id=event_id, group_id=None, drive_file_id=None, requested_by=mock_staff_user.id, photo_count=1, status="pending", approved_by=None, rejection_reason=None, created_at=datetime.now(timezone.utc), approved_at=None + id=request_id, event_id=event_id, group_id=None, drive_file_id=None, requested_by=mock_staff_user.id, photo_count=1, status="pending", approved_by=None, rejection_reason=None, created_at=datetime.now(timezone.utc), approved_at=None, source="drive" ) # Simulate DB Conflict (Duplicate) on photo insert @@ -225,7 +228,7 @@ async def test_create_group_from_folder( group_id = uuid.uuid4() mock_upload_request_group_querier.create_upload_request_group.return_value = UploadRequestGroup( - id=group_id, event_id=event_id, folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="pending", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None + id=group_id, event_id=event_id, folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="pending", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None, source="drive" ) with patch("app.service.upload_requests.NatsClient.publish") as mock_publish: @@ -249,7 +252,7 @@ async def test_process_group_import_no_images( ): group_id = uuid.uuid4() mock_upload_request_group_querier.start_upload_request_group_processing.return_value = UploadRequestGroup( - id=group_id, event_id=uuid.uuid4(), folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="processing", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None + id=group_id, event_id=uuid.uuid4(), folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="processing", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None, source="drive" ) mock_staff_drive_service.staff_user_querier.get_staff_user_by_id.return_value = mock_staff_user @@ -257,7 +260,7 @@ async def test_process_group_import_no_images( with patch("app.service.upload_requests.GoogleDriveClient.list_folder_files", return_value=[]): async def mock_get_group(*args, **kwargs): yield UploadRequestGroup( - id=group_id, event_id=uuid.uuid4(), folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="processing", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None + id=group_id, event_id=uuid.uuid4(), folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="processing", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None, source="drive" ) mock_upload_request_querier.list_upload_requests_by_group_id = mock_get_group @@ -267,7 +270,7 @@ async def mock_list_photos_by_ids(*args, **kwargs): upload_requests_service.upload_request_photo_querier.list_upload_request_photos_by_upload_request_ids = mock_list_photos_by_ids mock_upload_request_group_querier.get_upload_request_group_by_id.return_value = UploadRequestGroup( - id=group_id, event_id=uuid.uuid4(), folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="processing", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None + id=group_id, event_id=uuid.uuid4(), folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="processing", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None, source="drive" ) await upload_requests_service.process_group_import( @@ -294,12 +297,12 @@ async def test_approve_request_without_side_effects( event_id = uuid.uuid4() mock_upload_request_querier.get_upload_request_by_id.return_value = UploadRequest( - id=request_id, event_id=event_id, group_id=None, drive_file_id=None, requested_by=uuid.uuid4(), photo_count=1, status="pending", approved_by=None, rejection_reason=None, created_at=datetime.now(timezone.utc), approved_at=None + id=request_id, event_id=event_id, group_id=None, drive_file_id=None, requested_by=uuid.uuid4(), photo_count=1, status="pending", approved_by=None, rejection_reason=None, created_at=datetime.now(timezone.utc), approved_at=None, source="drive" ) async def mock_list_photos(*args, **kwargs): yield UploadRequestPhoto( - id=photo_id, upload_request_id=request_id, drive_file_id="drive_1", file_name="p.jpg", mime_type="image/jpeg", size_bytes=100, staging_storage_key="stage_key", final_storage_key=None, taken_at=None, day_number=None, visibility="public", status="staged", created_at=datetime.now(timezone.utc) + id=photo_id, upload_request_id=request_id, drive_file_id="drive_1", file_name="p.jpg", mime_type="image/jpeg", size_bytes=100, staging_storage_key="stage_key", final_storage_key=None, taken_at=None, day_number=None, visibility="public", status="staged", created_at=datetime.now(timezone.utc), source="drive", transfer_status="uploaded" ) mock_upload_request_photo_querier.list_upload_request_photos_by_upload_request_id = mock_list_photos From 3f0677916899e303deeb867335c6a43447b1c07f Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 02:41:43 +0100 Subject: [PATCH 09/29] feat: add config settings for direct upload --- app/core/config.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/app/core/config.py b/app/core/config.py index c9270f03..5a7fe12f 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -35,6 +35,10 @@ class Settings(BaseSettings): PHOTO_APPROVAL_TIMEOUT_DAYS: int = 7 EVENT_LIFECYCLE_POLL_INTERVAL_SECONDS: int = 60 + DIRECT_UPLOAD_PRESIGN_EXPIRES_SECONDS: int = 1800 + DIRECT_UPLOAD_STALE_PENDING_MINUTES: int = 45 + DIRECT_UPLOAD_RECONCILE_POLL_INTERVAL_SECONDS: int = 300 + DIRECT_UPLOAD_MAX_BATCH_SIZE: int = 200 # Mobile auth/session defaults MOBILE_SESSION_LIMIT: int = 3 From 43e54a5025523907aeabbaa93286facf6b1de01d Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 02:43:08 +0100 Subject: [PATCH 10/29] feat: add direct upload batch registration, confirm, and fail to UploadRequestsService --- app/schema/internal/uploads.py | 10 ++ app/service/upload_requests.py | 138 ++++++++++++++- tests/unit/test_direct_uploads.py | 284 ++++++++++++++++++++++++++++++ 3 files changed, 431 insertions(+), 1 deletion(-) create mode 100644 tests/unit/test_direct_uploads.py diff --git a/app/schema/internal/uploads.py b/app/schema/internal/uploads.py index c8b91da5..58f82491 100644 --- a/app/schema/internal/uploads.py +++ b/app/schema/internal/uploads.py @@ -8,3 +8,13 @@ class UploadPhotoInput: taken_at: datetime | None day_number: int | None visibility: str + + +@dataclass(frozen=True) +class DirectFileInput: + file_name: str + mime_type: str + size_bytes: int + taken_at: datetime | None + day_number: int | None + visibility: str diff --git a/app/service/upload_requests.py b/app/service/upload_requests.py index 1d19eef1..c8578364 100644 --- a/app/service/upload_requests.py +++ b/app/service/upload_requests.py @@ -8,6 +8,7 @@ from sqlalchemy.exc import IntegrityError +from app.core.config import settings from app.core.constant import AuditEventType from app.core.exceptions import AppException from app.core.logger import logger @@ -17,7 +18,7 @@ GoogleDriveFileMetadata, ) from app.infra.nats import NatsClient, NatsSubjects -from app.schema.internal.uploads import UploadPhotoInput +from app.schema.internal.uploads import DirectFileInput, UploadPhotoInput from app.service.audit import AuditService from app.service.staged_upload_storage import PreviewObject, StagedUploadStorageService from app.service.staff_drive import StaffDriveService @@ -596,6 +597,141 @@ async def create_group_from_folder( ) return UploadRequestGroupDetails(group=upload_group, requests=[]) + async def create_direct_group( + self, + *, + event_id: uuid.UUID, + requested_by: StaffUser, + ) -> UploadRequestGroup: + try: + upload_group = await self.upload_request_group_querier.create_upload_request_group( + upload_request_group_queries.CreateUploadRequestGroupParams( + event_id=event_id, + folder_id=None, + requested_by=requested_by.id, + total_photo_count=0, + batch_count=0, + source="direct", + processing_status="completed", + ) + ) + except IntegrityError as exc: + self._raise_integrity_error(exc) + if upload_group is None: + raise AppException.internal_error("Failed to create upload group") + return upload_group + + async def register_direct_batch( + self, + *, + group_id: uuid.UUID, + files: Sequence[DirectFileInput], + requested_by: StaffUser, + ) -> list[tuple[UploadRequestPhoto, str]]: + if not files: + raise AppException.bad_request("At least one file is required") + if len(files) > settings.DIRECT_UPLOAD_MAX_BATCH_SIZE: + raise AppException.bad_request( + f"A batch can contain at most {settings.DIRECT_UPLOAD_MAX_BATCH_SIZE} files" + ) + for file in files: + if file.mime_type not in self._allowed_mime_types: + raise AppException.image_format_error(f"Unsupported image format: {file.mime_type}") + if file.size_bytes <= 0 or file.size_bytes > self._max_photo_size_bytes: + raise AppException.bad_request(f"{file.file_name} exceeds maximum allowed size") + + group = await self.upload_request_group_querier.get_upload_request_group_by_id(id=group_id) + if group is None: + raise AppException.not_found("Upload group not found") + self._ensure_group_access(current_staff_user=requested_by, upload_group=group) + + upload_request = await self.upload_request_querier.create_upload_request( + upload_request_queries.CreateUploadRequestParams( + event_id=group.event_id, + group_id=group_id, + drive_file_id=None, + requested_by=requested_by.id, + photo_count=len(files), + source="direct", + ) + ) + if upload_request is None: + raise AppException.internal_error("Failed to create upload request") + + results: list[tuple[UploadRequestPhoto, str]] = [] + for file in files: + photo_id = uuid.uuid4() + storage_key, presigned_url = await self.staged_upload_storage.create_presigned_staging_upload( + upload_request_id=upload_request.id, + photo_id=photo_id, + file_name=file.file_name, + expires_seconds=settings.DIRECT_UPLOAD_PRESIGN_EXPIRES_SECONDS, + ) + created_photo = await self.upload_request_photo_querier.create_direct_upload_request_photo( + upload_request_photo_queries.CreateDirectUploadRequestPhotoParams( + upload_request_id=upload_request.id, + file_name=file.file_name, + mime_type=file.mime_type, + size_bytes=file.size_bytes, + staging_storage_key=storage_key, + taken_at=file.taken_at, + day_number=file.day_number, + visibility=file.visibility, + ) + ) + if created_photo is None: + raise AppException.internal_error("Failed to register upload photo") + results.append((created_photo, presigned_url)) + + await self.upload_request_group_querier.increment_upload_request_group_counts( + id=group_id, total_photo_count=len(files), + ) + + return results + + async def confirm_direct_upload( + self, + *, + photo_id: uuid.UUID, + requested_by: StaffUser, + ) -> UploadRequestPhoto: + photo = await self.upload_request_photo_querier.get_upload_request_photo_by_id(id=photo_id) + if photo is None: + raise AppException.not_found("Upload photo not found") + + stat = await self.staged_upload_storage.stat_staging_object(photo.staging_storage_key) + if stat is None: + failed = await self.upload_request_photo_querier.fail_upload_request_photo_transfer(id=photo_id) + if failed is None: + raise AppException.internal_error("Failed to mark upload as failed") + raise AppException.bad_request( + "Upload did not complete — file not found in storage. Retry the upload." + ) + + confirmed = await self.upload_request_photo_querier.confirm_upload_request_photo_transfer( + id=photo_id, + size_bytes=stat.size, + mime_type=stat.content_type, + ) + if confirmed is None: + raise AppException.internal_error("Failed to confirm upload") + return confirmed + + async def fail_direct_upload( + self, + *, + photo_id: uuid.UUID, + requested_by: StaffUser, + ) -> UploadRequestPhoto: + photo = await self.upload_request_photo_querier.get_upload_request_photo_by_id(id=photo_id) + if photo is None: + raise AppException.not_found("Upload photo not found") + + failed = await self.upload_request_photo_querier.fail_upload_request_photo_transfer(id=photo_id) + if failed is None: + raise AppException.internal_error("Failed to mark upload as failed") + return failed + async def process_group_import( self, *, diff --git a/tests/unit/test_direct_uploads.py b/tests/unit/test_direct_uploads.py new file mode 100644 index 00000000..5846e841 --- /dev/null +++ b/tests/unit/test_direct_uploads.py @@ -0,0 +1,284 @@ +import uuid +from datetime import datetime, timezone +from unittest.mock import AsyncMock + +import pytest + +from app.schema.internal.uploads import DirectFileInput +from app.service.staged_upload_storage import StoredObject +from app.service.upload_requests import UploadRequestsService +from db.generated.models import ( + StaffUser, + UploadRequest, + UploadRequestGroup, + UploadRequestPhoto, +) + + +@pytest.fixture +def mock_upload_request_group_querier(): + return AsyncMock() + +@pytest.fixture +def mock_upload_request_querier(): + return AsyncMock() + +@pytest.fixture +def mock_upload_request_photo_querier(): + return AsyncMock() + +@pytest.fixture +def mock_photo_querier(): + return AsyncMock() + +@pytest.fixture +def mock_staged_upload_storage(): + mock = AsyncMock() + mock.store_staging_object.return_value = StoredObject(storage_key="test_storage_key", content_type="image/jpeg", file_name="photo.jpg") + mock.create_presigned_staging_upload.return_value = ("staging/upload-requests/req1/photo1.jpg", "https://minio.local/signed-url") + return mock + +@pytest.fixture +def mock_staff_drive_service(): + mock = AsyncMock() + mock.get_access_token_for_staff_user.return_value = "fake_access_token" + mock.staff_user_querier = AsyncMock() + return mock + +@pytest.fixture +def mock_staff_notifications_service(): + return AsyncMock() + +@pytest.fixture +def mock_audit_service(): + return AsyncMock() + +@pytest.fixture +def upload_requests_service( + mock_upload_request_group_querier, + mock_upload_request_querier, + mock_upload_request_photo_querier, + mock_photo_querier, + mock_staged_upload_storage, + mock_staff_drive_service, + mock_staff_notifications_service, + mock_audit_service, +): + return UploadRequestsService( + upload_request_group_querier=mock_upload_request_group_querier, + upload_request_querier=mock_upload_request_querier, + upload_request_photo_querier=mock_upload_request_photo_querier, + photo_querier=mock_photo_querier, + staged_upload_storage=mock_staged_upload_storage, + staff_drive_service=mock_staff_drive_service, + staff_notifications_service=mock_staff_notifications_service, + audit_service=mock_audit_service, + ) + +@pytest.fixture +def mock_staff_user(): + return StaffUser( + id=uuid.uuid4(), + email="test@multai.com", + password="hash", + role="multi", + created_at=datetime.now(timezone.utc), + updated_at=datetime.now(timezone.utc), + ) + + +def _make_group(group_id, event_id, requested_by_id, **overrides): + defaults = dict( + id=group_id, event_id=event_id, folder_id=None, requested_by=requested_by_id, + approved_by=None, status="pending", total_photo_count=0, batch_count=0, + created_at=datetime.now(timezone.utc), approved_at=None, rejection_reason=None, + processing_status="completed", processed_photo_count=0, failed_photo_count=0, + error_message=None, source="direct", + ) + defaults.update(overrides) + return UploadRequestGroup(**defaults) + + +def _make_request(request_id, event_id, requested_by_id, group_id, **overrides): + defaults = dict( + id=request_id, event_id=event_id, drive_file_id=None, requested_by=requested_by_id, + approved_by=None, status="pending", created_at=datetime.now(timezone.utc), + approved_at=None, photo_count=1, rejection_reason=None, group_id=group_id, + source="direct", + ) + defaults.update(overrides) + return UploadRequest(**defaults) + + +def _make_photo(photo_id, request_id, **overrides): + defaults = dict( + id=photo_id, upload_request_id=request_id, drive_file_id=None, file_name="a.jpg", + mime_type="image/jpeg", size_bytes=1000, staging_storage_key="staging/x.jpg", + final_storage_key=None, taken_at=None, day_number=None, visibility="private", + status="staged", created_at=datetime.now(timezone.utc), + source="direct", transfer_status="pending_upload", + ) + defaults.update(overrides) + return UploadRequestPhoto(**defaults) + + +@pytest.mark.asyncio +async def test_create_direct_group_sets_source_and_completed_processing( + upload_requests_service, + mock_upload_request_group_querier, + mock_staff_user, +): + event_id = uuid.uuid4() + group_id = uuid.uuid4() + mock_upload_request_group_querier.create_upload_request_group.return_value = _make_group( + group_id, event_id, mock_staff_user.id, + ) + + group = await upload_requests_service.create_direct_group( + event_id=event_id, requested_by=mock_staff_user, + ) + + assert group.id == group_id + call_args = mock_upload_request_group_querier.create_upload_request_group.call_args + params = call_args.args[0] if call_args.args else call_args.kwargs["arg"] + assert params.folder_id is None + assert params.source == "direct" + assert params.processing_status == "completed" + + +@pytest.mark.asyncio +async def test_register_direct_batch_creates_pending_photos_and_returns_urls( + upload_requests_service, + mock_upload_request_querier, + mock_upload_request_photo_querier, + mock_upload_request_group_querier, + mock_staged_upload_storage, + mock_staff_user, +): + group_id = uuid.uuid4() + event_id = uuid.uuid4() + request_id = uuid.uuid4() + photo_id = uuid.uuid4() + + mock_upload_request_group_querier.get_upload_request_group_by_id.return_value = _make_group( + group_id, event_id, mock_staff_user.id, + ) + mock_upload_request_querier.create_upload_request.return_value = _make_request( + request_id, event_id, mock_staff_user.id, group_id, + ) + mock_upload_request_photo_querier.create_direct_upload_request_photo.return_value = _make_photo( + photo_id, request_id, staging_storage_key="staging/upload-requests/req1/photo1.jpg", + ) + + results = await upload_requests_service.register_direct_batch( + group_id=group_id, + files=[DirectFileInput( + file_name="a.jpg", mime_type="image/jpeg", size_bytes=1000, + taken_at=None, day_number=None, visibility="private", + )], + requested_by=mock_staff_user, + ) + + assert len(results) == 1 + photo, url = results[0] + assert photo.id == photo_id + assert url == "https://minio.local/signed-url" + mock_staged_upload_storage.create_presigned_staging_upload.assert_awaited() + mock_upload_request_group_querier.increment_upload_request_group_counts.assert_awaited_once_with( + id=group_id, total_photo_count=1, + ) + + +@pytest.mark.asyncio +async def test_register_direct_batch_rejects_oversized_batch( + upload_requests_service, + mock_upload_request_group_querier, + mock_staff_user, +): + group_id = uuid.uuid4() + from app.core.exceptions import AppException + + files = [ + DirectFileInput(file_name=f"{i}.jpg", mime_type="image/jpeg", size_bytes=1000, taken_at=None, day_number=None, visibility="private") + for i in range(201) + ] + + with pytest.raises(Exception): + await upload_requests_service.register_direct_batch( + group_id=group_id, files=files, requested_by=mock_staff_user, + ) + mock_upload_request_group_querier.get_upload_request_group_by_id.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_confirm_direct_upload_success_updates_transfer_status( + upload_requests_service, + mock_upload_request_photo_querier, + mock_staged_upload_storage, + mock_staff_user, +): + from app.infra.minio import ObjectStat + + photo_id = uuid.uuid4() + request_id = uuid.uuid4() + existing_photo = _make_photo(photo_id, request_id, transfer_status="pending_upload") + mock_upload_request_photo_querier.get_upload_request_photo_by_id.return_value = existing_photo + mock_staged_upload_storage.stat_staging_object.return_value = ObjectStat(size=1000, content_type="image/jpeg") + mock_upload_request_photo_querier.confirm_upload_request_photo_transfer.return_value = _make_photo( + photo_id, request_id, transfer_status="uploaded", + ) + + result = await upload_requests_service.confirm_direct_upload( + photo_id=photo_id, requested_by=mock_staff_user, + ) + + assert result.id == photo_id + mock_upload_request_photo_querier.confirm_upload_request_photo_transfer.assert_awaited_once_with( + id=photo_id, size_bytes=1000, mime_type="image/jpeg", + ) + + +@pytest.mark.asyncio +async def test_confirm_direct_upload_marks_failed_when_object_missing( + upload_requests_service, + mock_upload_request_photo_querier, + mock_staged_upload_storage, + mock_staff_user, +): + photo_id = uuid.uuid4() + request_id = uuid.uuid4() + existing_photo = _make_photo(photo_id, request_id, transfer_status="pending_upload") + mock_upload_request_photo_querier.get_upload_request_photo_by_id.return_value = existing_photo + mock_staged_upload_storage.stat_staging_object.return_value = None + mock_upload_request_photo_querier.fail_upload_request_photo_transfer.return_value = _make_photo( + photo_id, request_id, transfer_status="failed", + ) + + with pytest.raises(Exception): + await upload_requests_service.confirm_direct_upload( + photo_id=photo_id, requested_by=mock_staff_user, + ) + + mock_upload_request_photo_querier.fail_upload_request_photo_transfer.assert_awaited_once_with(id=photo_id) + + +@pytest.mark.asyncio +async def test_fail_direct_upload_marks_transfer_failed( + upload_requests_service, + mock_upload_request_photo_querier, + mock_staff_user, +): + photo_id = uuid.uuid4() + request_id = uuid.uuid4() + existing_photo = _make_photo(photo_id, request_id, transfer_status="pending_upload") + mock_upload_request_photo_querier.get_upload_request_photo_by_id.return_value = existing_photo + mock_upload_request_photo_querier.fail_upload_request_photo_transfer.return_value = _make_photo( + photo_id, request_id, transfer_status="failed", + ) + + result = await upload_requests_service.fail_direct_upload( + photo_id=photo_id, requested_by=mock_staff_user, + ) + + assert result.id == photo_id + mock_upload_request_photo_querier.fail_upload_request_photo_transfer.assert_awaited_once_with(id=photo_id) From 7ac3169dd3327c4d8589ecc2dde709fff2f42a30 Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 02:44:04 +0100 Subject: [PATCH 11/29] feat: add resume and pre-approval transfer guard for direct uploads --- app/service/upload_requests.py | 46 +++++++++++++++++++ tests/unit/test_direct_uploads.py | 74 +++++++++++++++++++++++++++++++ 2 files changed, 120 insertions(+) diff --git a/app/service/upload_requests.py b/app/service/upload_requests.py index c8578364..f64e346b 100644 --- a/app/service/upload_requests.py +++ b/app/service/upload_requests.py @@ -338,6 +338,16 @@ async def _approve_request_without_side_effects( if not staged_photos: raise AppException.bad_request("No staged photos found for this upload request") + not_transferred = [ + p for p in staged_photos + if getattr(p, "transfer_status", "uploaded") != "uploaded" + ] + if not_transferred: + raise AppException.bad_request( + f"{len(not_transferred)} photo(s) have not finished uploading. " + "Resume or remove them before approving." + ) + finalized_storage_keys: list[str] = [] created_photos: list[Photo] = [] try: @@ -732,6 +742,42 @@ async def fail_direct_upload( raise AppException.internal_error("Failed to mark upload as failed") return failed + async def resume_direct_group( + self, + *, + group_id: uuid.UUID, + requested_by: StaffUser, + ) -> list[tuple[UploadRequestPhoto, str]]: + group = await self.upload_request_group_querier.get_upload_request_group_by_id(id=group_id) + if group is None: + raise AppException.not_found("Upload group not found") + self._ensure_group_access(current_staff_user=requested_by, upload_group=group) + + request_ids: list[uuid.UUID] = [] + async for req in self.upload_request_querier.list_upload_requests_by_group_id(group_id=group_id): + request_ids.append(req.id) + + results: list[tuple[UploadRequestPhoto, str]] = [] + async for photo in self.upload_request_photo_querier.list_upload_request_photos_by_upload_request_ids( + dollar_1=request_ids + ): + if getattr(photo, "transfer_status", "uploaded") not in ("pending_upload", "failed"): + continue + storage_key, presigned_url = await self.staged_upload_storage.create_presigned_staging_upload( + upload_request_id=photo.upload_request_id, + photo_id=photo.id, + file_name=photo.file_name, + expires_seconds=settings.DIRECT_UPLOAD_PRESIGN_EXPIRES_SECONDS, + ) + reset_photo = await self.upload_request_photo_querier.reset_upload_request_photo_transfer_to_pending( + id=photo.id + ) + if reset_photo is None: + continue + results.append((reset_photo, presigned_url)) + + return results + async def process_group_import( self, *, diff --git a/tests/unit/test_direct_uploads.py b/tests/unit/test_direct_uploads.py index 5846e841..ca73237f 100644 --- a/tests/unit/test_direct_uploads.py +++ b/tests/unit/test_direct_uploads.py @@ -262,6 +262,80 @@ async def test_confirm_direct_upload_marks_failed_when_object_missing( mock_upload_request_photo_querier.fail_upload_request_photo_transfer.assert_awaited_once_with(id=photo_id) +@pytest.mark.asyncio +async def test_approve_request_blocked_when_photo_not_fully_uploaded( + upload_requests_service, + mock_upload_request_querier, + mock_upload_request_photo_querier, + mock_staff_user, +): + request_id = uuid.uuid4() + photo_id = uuid.uuid4() + + mock_upload_request_querier.get_upload_request_by_id.return_value = _make_request( + request_id, uuid.uuid4(), mock_staff_user.id, None, + ) + not_uploaded_photo = _make_photo(photo_id, request_id, transfer_status="pending_upload") + + async def _photos_iter(upload_request_id): + yield not_uploaded_photo + + mock_upload_request_photo_querier.list_upload_request_photos_by_upload_request_id = _photos_iter + + with pytest.raises(Exception) as exc_info: + await upload_requests_service.approve_request( + request_id=request_id, approved_by=mock_staff_user, + ) + assert "have not finished uploading" in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_resume_direct_group_reissues_urls_for_pending_and_failed_only( + upload_requests_service, + mock_upload_request_group_querier, + mock_upload_request_querier, + mock_upload_request_photo_querier, + mock_staged_upload_storage, + mock_staff_user, +): + group_id = uuid.uuid4() + event_id = uuid.uuid4() + request_id = uuid.uuid4() + failed_photo_id = uuid.uuid4() + uploaded_photo_id = uuid.uuid4() + + mock_upload_request_group_querier.get_upload_request_group_by_id.return_value = _make_group( + group_id, event_id, mock_staff_user.id, total_photo_count=2, batch_count=1, failed_photo_count=1, + ) + + async def _requests_iter(group_id): + yield _make_request(request_id, event_id, mock_staff_user.id, group_id, photo_count=2) + + mock_upload_request_querier.list_upload_requests_by_group_id = _requests_iter + + failed_photo = _make_photo(failed_photo_id, request_id, file_name="fail.jpg", staging_storage_key="staging/fail.jpg", transfer_status="failed") + uploaded_photo = _make_photo(uploaded_photo_id, request_id, file_name="ok.jpg", staging_storage_key="staging/ok.jpg", transfer_status="uploaded") + + async def _photos_iter(dollar_1): + for p in [failed_photo, uploaded_photo]: + yield p + + mock_upload_request_photo_querier.list_upload_request_photos_by_upload_request_ids = _photos_iter + mock_staged_upload_storage.create_presigned_staging_upload.return_value = ("staging/fail.jpg", "https://minio.local/resumed") + mock_upload_request_photo_querier.reset_upload_request_photo_transfer_to_pending.return_value = _make_photo( + failed_photo_id, request_id, file_name="fail.jpg", transfer_status="pending_upload", + ) + + results = await upload_requests_service.resume_direct_group( + group_id=group_id, requested_by=mock_staff_user, + ) + + assert len(results) == 1 + photo, url = results[0] + assert photo.id == failed_photo_id + assert url == "https://minio.local/resumed" + + @pytest.mark.asyncio async def test_fail_direct_upload_marks_transfer_failed( upload_requests_service, From 745fd69c8d2c3c1ace825545011b9e98056b0b7e Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 02:44:55 +0100 Subject: [PATCH 12/29] feat: add request/response schemas for direct upload endpoints --- app/schema/request/staff/uploads_direct.py | 22 +++++++++++++++++++++ app/schema/response/staff/upload_groups.py | 6 ++++-- app/schema/response/staff/uploads.py | 5 ++++- app/schema/response/staff/uploads_direct.py | 18 +++++++++++++++++ 4 files changed, 48 insertions(+), 3 deletions(-) create mode 100644 app/schema/request/staff/uploads_direct.py create mode 100644 app/schema/response/staff/uploads_direct.py diff --git a/app/schema/request/staff/uploads_direct.py b/app/schema/request/staff/uploads_direct.py new file mode 100644 index 00000000..068c1b69 --- /dev/null +++ b/app/schema/request/staff/uploads_direct.py @@ -0,0 +1,22 @@ +from datetime import datetime +from typing import Optional +from uuid import UUID + +from pydantic import BaseModel + + +class DirectFileInputRequest(BaseModel): + file_name: str + mime_type: str + size_bytes: int + taken_at: Optional[datetime] = None + day_number: Optional[int] = None + visibility: str = "private" + + +class CreateDirectGroupRequest(BaseModel): + event_id: UUID + + +class RegisterDirectBatchRequest(BaseModel): + files: list[DirectFileInputRequest] diff --git a/app/schema/response/staff/upload_groups.py b/app/schema/response/staff/upload_groups.py index 7a264c2d..f691cc9a 100644 --- a/app/schema/response/staff/upload_groups.py +++ b/app/schema/response/staff/upload_groups.py @@ -18,10 +18,11 @@ class UploadRequestGroupSchema(BaseModel): id: UUID event_id: UUID - folder_id: str + folder_id: str | None requested_by: UUID approved_by: UUID | None status: str + source: str processing_status: str total_photo_count: int batch_count: int @@ -72,8 +73,9 @@ class UploadRequestGroupSummarySchema(BaseModel): model_config = ConfigDict(from_attributes=True) id: UUID event_id: UUID - folder_id: str + folder_id: str | None status: str + source: str processing_status: str total_photo_count: int batch_count: int diff --git a/app/schema/response/staff/uploads.py b/app/schema/response/staff/uploads.py index 863414cd..60a13544 100644 --- a/app/schema/response/staff/uploads.py +++ b/app/schema/response/staff/uploads.py @@ -11,7 +11,7 @@ class UploadRequestPhotoSchema(BaseModel): model_config = ConfigDict(from_attributes=True) id: UUID - drive_file_id: str + drive_file_id: str | None file_name: str mime_type: str size_bytes: int @@ -19,6 +19,8 @@ class UploadRequestPhotoSchema(BaseModel): day_number: int | None visibility: str status: str + source: str + transfer_status: str created_at: datetime @@ -32,6 +34,7 @@ class UploadRequestSchema(BaseModel): requested_by: UUID approved_by: UUID | None status: str + source: str photo_count: int created_at: datetime approved_at: datetime | None diff --git a/app/schema/response/staff/uploads_direct.py b/app/schema/response/staff/uploads_direct.py new file mode 100644 index 00000000..fba3d4c4 --- /dev/null +++ b/app/schema/response/staff/uploads_direct.py @@ -0,0 +1,18 @@ +from uuid import UUID + +from pydantic import BaseModel + + +class DirectUploadFileResponse(BaseModel): + photo_id: UUID + file_name: str + upload_url: str + + +class RegisterDirectBatchResponse(BaseModel): + group_id: UUID + items: list[DirectUploadFileResponse] + + +class ResumeDirectGroupResponse(BaseModel): + items: list[DirectUploadFileResponse] From b57af9ecbf66a3dc126c6cdad49aa7d38071f3b9 Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 02:46:19 +0100 Subject: [PATCH 13/29] feat: add staff router endpoints for direct upload --- app/router/staff/__init__.py | 2 + app/router/staff/uploads_direct.py | 119 +++++++++++++++++++++++++++++ 2 files changed, 121 insertions(+) create mode 100644 app/router/staff/uploads_direct.py diff --git a/app/router/staff/__init__.py b/app/router/staff/__init__.py index 35cbe0ec..11023d7b 100644 --- a/app/router/staff/__init__.py +++ b/app/router/staff/__init__.py @@ -3,8 +3,10 @@ from app.router.staff.drive import router as staff_drive_router from app.router.staff.notifications import router as staff_notifications_router from app.router.staff.uploads import router as staff_uploads_router +from app.router.staff.uploads_direct import router as staff_uploads_direct_router router = APIRouter(prefix="/staff", tags=["staff"]) router.include_router(staff_drive_router) router.include_router(staff_notifications_router) router.include_router(staff_uploads_router) +router.include_router(staff_uploads_direct_router) diff --git a/app/router/staff/uploads_direct.py b/app/router/staff/uploads_direct.py new file mode 100644 index 00000000..0e9dbd13 --- /dev/null +++ b/app/router/staff/uploads_direct.py @@ -0,0 +1,119 @@ +from uuid import UUID + +from fastapi import APIRouter, Depends + +from app.container import Container, get_container +from app.deps.cookie_auth import get_current_staff_user +from app.schema.internal.uploads import DirectFileInput +from app.schema.request.staff.uploads_direct import ( + CreateDirectGroupRequest, + RegisterDirectBatchRequest, +) +from app.schema.response.staff.upload_groups import UploadRequestGroupSchema +from app.schema.response.staff.uploads_direct import ( + DirectUploadFileResponse, + RegisterDirectBatchResponse, + ResumeDirectGroupResponse, +) +from db.generated.models import StaffUser + +router = APIRouter(prefix="/uploads/direct") +# this endpoint are for staff to upload images directly to the system, bypassing the mobile app. This is useful for bulk uploads or for users who cannot use the mobile app. + +@router.post("/groups", response_model=UploadRequestGroupSchema) +async def create_direct_group( + req: CreateDirectGroupRequest, + current_staff_user: StaffUser = Depends(get_current_staff_user), + container: Container = Depends(get_container), +) -> UploadRequestGroupSchema: + group = await container.upload_requests_service.create_direct_group( + event_id=req.event_id, requested_by=current_staff_user, + ) + details = await container.upload_requests_service.get_group_details( + group_id=group.id, current_staff_user=current_staff_user, + ) + return UploadRequestGroupSchema.from_details(details) + + +@router.post("/groups/{group_id}/batches", response_model=RegisterDirectBatchResponse) +async def register_direct_batch( + group_id: UUID, + req: RegisterDirectBatchRequest, + current_staff_user: StaffUser = Depends(get_current_staff_user), + container: Container = Depends(get_container), +) -> RegisterDirectBatchResponse: + results = await container.upload_requests_service.register_direct_batch( + group_id=group_id, + files=[ + DirectFileInput( + file_name=f.file_name, + mime_type=f.mime_type, + size_bytes=f.size_bytes, + taken_at=f.taken_at, + day_number=f.day_number, + visibility=f.visibility, + ) + for f in req.files + ], + requested_by=current_staff_user, + ) + return RegisterDirectBatchResponse( + group_id=group_id, + items=[ + DirectUploadFileResponse(photo_id=photo.id, file_name=photo.file_name, upload_url=url) + for photo, url in results + ], + ) + + +@router.post("/photos/{photo_id}/confirm", response_model=DirectUploadFileResponse) +async def confirm_direct_upload( + photo_id: UUID, + current_staff_user: StaffUser = Depends(get_current_staff_user), + container: Container = Depends(get_container), +) -> DirectUploadFileResponse: + photo = await container.upload_requests_service.confirm_direct_upload( + photo_id=photo_id, requested_by=current_staff_user, + ) + return DirectUploadFileResponse(photo_id=photo.id, file_name=photo.file_name, upload_url="") + + +@router.post("/photos/{photo_id}/fail", response_model=DirectUploadFileResponse) +async def fail_direct_upload( + photo_id: UUID, + current_staff_user: StaffUser = Depends(get_current_staff_user), + container: Container = Depends(get_container), +) -> DirectUploadFileResponse: + photo = await container.upload_requests_service.fail_direct_upload( + photo_id=photo_id, requested_by=current_staff_user, + ) + return DirectUploadFileResponse(photo_id=photo.id, file_name=photo.file_name, upload_url="") + + +@router.post("/groups/{group_id}/resume", response_model=ResumeDirectGroupResponse) +async def resume_direct_group( + group_id: UUID, + current_staff_user: StaffUser = Depends(get_current_staff_user), + container: Container = Depends(get_container), +) -> ResumeDirectGroupResponse: + results = await container.upload_requests_service.resume_direct_group( + group_id=group_id, requested_by=current_staff_user, + ) + return ResumeDirectGroupResponse( + items=[ + DirectUploadFileResponse(photo_id=photo.id, file_name=photo.file_name, upload_url=url) + for photo, url in results + ] + ) + + +@router.get("/groups/{group_id}", response_model=UploadRequestGroupSchema) +async def get_direct_group_status( + group_id: UUID, + current_staff_user: StaffUser = Depends(get_current_staff_user), + container: Container = Depends(get_container), +) -> UploadRequestGroupSchema: + details = await container.upload_requests_service.get_group_details( + group_id=group_id, current_staff_user=current_staff_user, + ) + return UploadRequestGroupSchema.from_details(details) From 1291defc62b9f678ffe62048ac1b1939880cbfa5 Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 02:47:25 +0100 Subject: [PATCH 14/29] feat: add reconciliation worker for stale direct-upload transfers --- app/worker/upload_reconciler/__init__.py | 0 app/worker/upload_reconciler/main.py | 60 ++++++++++++++++++++++++ makefile | 1 + 3 files changed, 61 insertions(+) create mode 100644 app/worker/upload_reconciler/__init__.py create mode 100644 app/worker/upload_reconciler/main.py diff --git a/app/worker/upload_reconciler/__init__.py b/app/worker/upload_reconciler/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/app/worker/upload_reconciler/main.py b/app/worker/upload_reconciler/main.py new file mode 100644 index 00000000..915d7911 --- /dev/null +++ b/app/worker/upload_reconciler/main.py @@ -0,0 +1,60 @@ +import asyncio + +from app.core.config import settings +from app.core.logger import logger +from app.infra.database import engine +from app.service.staged_upload_storage import StagedUploadStorageService +from db.generated import upload_request_photos as upload_request_photo_queries + +storage_service = StagedUploadStorageService() + + +async def run_reconcile_pass() -> None: + async with engine.begin() as conn: + querier = upload_request_photo_queries.AsyncQuerier(conn) + + stale_photos = [ + photo + async for photo in querier.list_stale_pending_transfer_photos( + dollar_1=str(settings.DIRECT_UPLOAD_STALE_PENDING_MINUTES) + ) + ] + + if not stale_photos: + return + + confirmed = 0 + failed = 0 + for photo in stale_photos: + stat = await storage_service.stat_staging_object(photo.staging_storage_key) + if stat is not None: + await querier.confirm_upload_request_photo_transfer( + id=photo.id, size_bytes=stat.size, mime_type=stat.content_type, + ) + confirmed += 1 + else: + await querier.fail_upload_request_photo_transfer(id=photo.id) + failed += 1 + + logger.info( + "upload_reconciler: reconciled %d stale photo(s) — %d confirmed, %d failed", + len(stale_photos), confirmed, failed, + ) + + +async def main() -> None: + logger.info( + "Upload reconciler worker starting, poll_interval=%ds, stale_after=%dmin", + settings.DIRECT_UPLOAD_RECONCILE_POLL_INTERVAL_SECONDS, + settings.DIRECT_UPLOAD_STALE_PENDING_MINUTES, + ) + while True: + try: + await run_reconcile_pass() + except Exception: + logger.exception("upload_reconciler: pass failed") + await asyncio.sleep(settings.DIRECT_UPLOAD_RECONCILE_POLL_INTERVAL_SECONDS) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/makefile b/makefile index 36690ed8..52e37c54 100644 --- a/makefile +++ b/makefile @@ -66,6 +66,7 @@ run-workers: uv run python -m app.worker.storage_cleaner.main & \ uv run python -m app.worker.email_worker.main & \ uv run python -m app.worker.event_lifecycle.main & \ + uv run python -m app.worker.upload_reconciler.main & \ wait lint: From 7154235fb5220e9859a4871b95563a93376cfcfd Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 02:47:37 +0100 Subject: [PATCH 15/29] docs: clarify direct upload router comment --- app/router/staff/uploads_direct.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/app/router/staff/uploads_direct.py b/app/router/staff/uploads_direct.py index 0e9dbd13..267cd9bd 100644 --- a/app/router/staff/uploads_direct.py +++ b/app/router/staff/uploads_direct.py @@ -18,7 +18,7 @@ from db.generated.models import StaffUser router = APIRouter(prefix="/uploads/direct") -# this endpoint are for staff to upload images directly to the system, bypassing the mobile app. This is useful for bulk uploads or for users who cannot use the mobile app. +# this endpoint are for staff to upload images directly to the system and very large files and they can resume and restart and retry . @router.post("/groups", response_model=UploadRequestGroupSchema) async def create_direct_group( From 070d7e3c6e7ba6ed40cbcfdbd8c68de1e6a89ddc Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 03:00:16 +0100 Subject: [PATCH 16/29] feat: add Google Drive write capability for syncing approved direct uploads --- app/core/config.py | 13 +++- app/infra/google_drive.py | 61 +++++++++++++++++++ app/infra/nats.py | 1 + app/service/staff_drive.py | 19 ++++++ db/generated/models.py | 2 + db/generated/photos.py | 61 ++++++++++++++++--- db/queries/photos.sql | 7 +++ .../down/add_drive_sync_fields_to_photos.sql | 3 + .../up/add_drive_sync_fields_to_photos.sql | 3 + ...06b53a9_add_drive_sync_fields_to_photos.py | 25 ++++++++ 10 files changed, 185 insertions(+), 10 deletions(-) create mode 100644 migrations/sql/down/add_drive_sync_fields_to_photos.sql create mode 100644 migrations/sql/up/add_drive_sync_fields_to_photos.sql create mode 100644 migrations/versions/af58506b53a9_add_drive_sync_fields_to_photos.py diff --git a/app/core/config.py b/app/core/config.py index 5a7fe12f..9d01ffa0 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -83,9 +83,20 @@ class Settings(BaseSettings): GOOGLE_CLIENT_ID: str = "" GOOGLE_CLIENT_SECRET: str = "" GOOGLE_REDIRECT_URI: str = "" + # drive.readonly alone can't write; drive.file alone can only see files + # the app itself created, which would break browsing/importing existing + # Drive folders. Both scopes together preserve the existing read/import + # flow and add write access for syncing approved direct uploads back to + # Drive. Existing staff connections keep their old readonly-only grant + # until they disconnect and reconnect through the consent screen. GOOGLE_OAUTH_SCOPES: str = ( - "https://www.googleapis.com/auth/drive.readonly openid email profile" + "https://www.googleapis.com/auth/drive.readonly " + "https://www.googleapis.com/auth/drive.file openid email profile" ) + # Folder ID (from the Drive URL) that approved direct-upload photos get + # synced into. Empty means uploads land in the connected account's Drive + # root instead of a specific folder. + GOOGLE_CLUB_DRIVE_FOLDER_ID: str = "" FACE_ENCRYPTION_KEY: str FIREBASE_CREDENTIALS_PATH: str diff --git a/app/infra/google_drive.py b/app/infra/google_drive.py index f25c4ef6..05633bf7 100644 --- a/app/infra/google_drive.py +++ b/app/infra/google_drive.py @@ -189,6 +189,67 @@ async def get_file_metadata( size_bytes=size_bytes, ) + @staticmethod + async def upload_file( + *, + access_token: str, + file_name: str, + content_type: str, + data: bytes, + folder_id: str | None, + ) -> GoogleDriveFileMetadata: + boundary = "multai-drive-upload-boundary" + metadata: dict[str, object] = {"name": file_name} + if folder_id: + metadata["parents"] = [folder_id] + + body = ( + f"--{boundary}\r\n" + "Content-Type: application/json; charset=UTF-8\r\n\r\n" + f"{json.dumps(metadata)}\r\n" + f"--{boundary}\r\n" + f"Content-Type: {content_type}\r\n\r\n" + ).encode("utf-8") + data + f"\r\n--{boundary}--".encode("utf-8") + + def _request() -> dict[str, object]: + url = ( + "https://www.googleapis.com/upload/drive/v3/files" + "?uploadType=multipart&supportsAllDrives=true&fields=id,name,mimeType,size" + ) + request = urllib.request.Request( + url, + data=body, + headers={ + "Authorization": f"Bearer {access_token}", + "Content-Type": f"multipart/related; boundary={boundary}", + }, + method="POST", + ) + try: + with urllib.request.urlopen(request, timeout=60) as response: + return json.loads(response.read().decode("utf-8")) + except urllib.error.HTTPError as exc: + details = exc.read().decode("utf-8", errors="ignore") + raise AppException.bad_request( + f"Google Drive file upload failed: {details or exc.reason}" + ) from exc + except urllib.error.URLError as exc: + raise AppException.internal_error("Unable to reach Google APIs") from exc + + result = await asyncio.to_thread(_request) + size_raw = result.get("size", "0") + try: + size_bytes = int(size_raw) if isinstance(size_raw, (str, int)) else len(data) + except (TypeError, ValueError): + size_bytes = len(data) + + return GoogleDriveFileMetadata( + id=GoogleDriveClient._require_str(result, "id"), + name=GoogleDriveClient._require_str(result, "name"), + mime_type=GoogleDriveClient._require_str(result, "mimeType"), + size_bytes=size_bytes, + ) + @staticmethod async def download_file( *, diff --git a/app/infra/nats.py b/app/infra/nats.py index e17bd500..8b07e831 100644 --- a/app/infra/nats.py +++ b/app/infra/nats.py @@ -35,6 +35,7 @@ class NatsSubjects(Enum): STAFF_UPLOAD_REQUEST_APPROVED = "staff.upload_request.approved" STAFF_UPLOAD_REQUEST_REJECTED = "staff.upload_request.rejected" PHOTO_PROCESS = "photo.process" + PHOTO_DRIVE_SYNC_REQUESTED = "photo.drive_sync.requested" class NatsClient: diff --git a/app/service/staff_drive.py b/app/service/staff_drive.py index f3b2ba00..b80a98be 100644 --- a/app/service/staff_drive.py +++ b/app/service/staff_drive.py @@ -209,6 +209,25 @@ async def get_system_access_token(self) -> str: connection = await self._refresh_connection_access_token(connection) return self.decrypt(connection.access_token) + async def upload_to_system_drive( + self, + *, + file_name: str, + content_type: str, + data: bytes, + ) -> str: + """Upload bytes to the system/club Drive using the most recently + connected active staff Drive connection. Returns the Drive file id.""" + access_token = await self.get_system_access_token() + metadata = await GoogleDriveClient.upload_file( + access_token=access_token, + file_name=file_name, + content_type=content_type, + data=data, + folder_id=settings.GOOGLE_CLUB_DRIVE_FOLDER_ID or None, + ) + return metadata.id + async def disconnect(self, staff_user_id: uuid.UUID) -> None: connection = await self.get_status(staff_user_id) if connection is None: diff --git a/db/generated/models.py b/db/generated/models.py index 2d6bcf62..7d935995 100644 --- a/db/generated/models.py +++ b/db/generated/models.py @@ -115,6 +115,8 @@ class Photo: visibility: str status: Any created_at: datetime.datetime + drive_file_id: Optional[str] + drive_synced_at: Optional[datetime.datetime] @dataclasses.dataclass() diff --git a/db/generated/photos.py b/db/generated/photos.py index 599a46d1..a978143c 100644 --- a/db/generated/photos.py +++ b/db/generated/photos.py @@ -35,7 +35,7 @@ ) VALUES ( :p1, :p2, :p3, :p4, :p5 ) -RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at +RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at """ @@ -57,12 +57,12 @@ class CreatePhotoParams: GET_PHOTO_BY_ID = """-- name: get_photo_by_id \\:one -SELECT id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at FROM photos WHERE id = :p1 +SELECT id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at FROM photos WHERE id = :p1 """ LIST_EVENT_PHOTOS_FOR_USER = """-- name: list_event_photos_for_user \\:many -SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, +SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, p.drive_file_id, p.drive_synced_at, (SELECT COUNT(*) FROM photo_faces pf2 WHERE pf2.photo_id = p.id)\\:\\:int AS face_count FROM photos p WHERE p.event_id = :p2 @@ -106,11 +106,13 @@ class ListEventPhotosForUserRow: visibility: str status: Any created_at: datetime.datetime + drive_file_id: Optional[str] + drive_synced_at: Optional[datetime.datetime] face_count: int LIST_USER_PHOTOS = """-- name: list_user_photos \\:many -SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, +SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, p.drive_file_id, p.drive_synced_at, (SELECT COUNT(*) FROM photo_faces pf2 WHERE pf2.photo_id = p.id)\\:\\:int AS face_count FROM photos p WHERE ( @@ -152,14 +154,25 @@ class ListUserPhotosRow: visibility: str status: Any created_at: datetime.datetime + drive_file_id: Optional[str] + drive_synced_at: Optional[datetime.datetime] face_count: int +MARK_PHOTO_DRIVE_SYNCED = """-- name: mark_photo_drive_synced \\:one +UPDATE photos +SET drive_file_id = :p2, + drive_synced_at = NOW() +WHERE id = :p1 +RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at +""" + + UPDATE_PHOTO_STATUS = """-- name: update_photo_status \\:one UPDATE photos SET status = :p2 WHERE id = :p1 -RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at +RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at """ @@ -167,7 +180,7 @@ class ListUserPhotosRow: UPDATE photos SET visibility = :p2 WHERE id = :p1 -RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at +RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at """ @@ -201,9 +214,11 @@ async def create_photo(self, arg: CreatePhotoParams) -> Optional[models.Photo]: visibility=row[6], status=row[7], created_at=row[8], + drive_file_id=row[9], + drive_synced_at=row[10], ) - async def get_drive_file_id_for_photo(self, *, final_storage_key: Optional[str]) -> Optional[str]: + async def get_drive_file_id_for_photo(self, *, final_storage_key: Optional[str]) -> Optional[Optional[str]]: row = (await self._conn.execute(sqlalchemy.text(GET_DRIVE_FILE_ID_FOR_PHOTO), {"p1": final_storage_key})).first() if row is None: return None @@ -223,6 +238,8 @@ async def get_photo_by_id(self, *, id: uuid.UUID) -> Optional[models.Photo]: visibility=row[6], status=row[7], created_at=row[8], + drive_file_id=row[9], + drive_synced_at=row[10], ) async def list_event_photos_for_user(self, arg: ListEventPhotosForUserParams) -> AsyncIterator[ListEventPhotosForUserRow]: @@ -244,7 +261,9 @@ async def list_event_photos_for_user(self, arg: ListEventPhotosForUserParams) -> visibility=row[6], status=row[7], created_at=row[8], - face_count=row[9], + drive_file_id=row[9], + drive_synced_at=row[10], + face_count=row[11], ) async def list_user_photos(self, arg: ListUserPhotosParams) -> AsyncIterator[ListUserPhotosRow]: @@ -266,9 +285,29 @@ async def list_user_photos(self, arg: ListUserPhotosParams) -> AsyncIterator[Lis visibility=row[6], status=row[7], created_at=row[8], - face_count=row[9], + drive_file_id=row[9], + drive_synced_at=row[10], + face_count=row[11], ) + async def mark_photo_drive_synced(self, *, id: uuid.UUID, drive_file_id: Optional[str]) -> Optional[models.Photo]: + row = (await self._conn.execute(sqlalchemy.text(MARK_PHOTO_DRIVE_SYNCED), {"p1": id, "p2": drive_file_id})).first() + if row is None: + return None + return models.Photo( + id=row[0], + event_id=row[1], + uploaded_by=row[2], + storage_key=row[3], + taken_at=row[4], + day_number=row[5], + visibility=row[6], + status=row[7], + created_at=row[8], + drive_file_id=row[9], + drive_synced_at=row[10], + ) + async def update_photo_status(self, *, id: uuid.UUID, status: Any) -> Optional[models.Photo]: row = (await self._conn.execute(sqlalchemy.text(UPDATE_PHOTO_STATUS), {"p1": id, "p2": status})).first() if row is None: @@ -283,6 +322,8 @@ async def update_photo_status(self, *, id: uuid.UUID, status: Any) -> Optional[m visibility=row[6], status=row[7], created_at=row[8], + drive_file_id=row[9], + drive_synced_at=row[10], ) async def update_photo_visibility(self, *, id: uuid.UUID, visibility: str) -> Optional[models.Photo]: @@ -299,4 +340,6 @@ async def update_photo_visibility(self, *, id: uuid.UUID, visibility: str) -> Op visibility=row[6], status=row[7], created_at=row[8], + drive_file_id=row[9], + drive_synced_at=row[10], ) diff --git a/db/queries/photos.sql b/db/queries/photos.sql index 3838e3c2..9904ad8a 100644 --- a/db/queries/photos.sql +++ b/db/queries/photos.sql @@ -84,3 +84,10 @@ SELECT urp.drive_file_id FROM upload_request_photos urp WHERE urp.final_storage_key = $1 LIMIT 1; + +-- name: MarkPhotoDriveSynced :one +UPDATE photos +SET drive_file_id = $2, + drive_synced_at = NOW() +WHERE id = $1 +RETURNING *; diff --git a/migrations/sql/down/add_drive_sync_fields_to_photos.sql b/migrations/sql/down/add_drive_sync_fields_to_photos.sql new file mode 100644 index 00000000..d6488697 --- /dev/null +++ b/migrations/sql/down/add_drive_sync_fields_to_photos.sql @@ -0,0 +1,3 @@ +ALTER TABLE photos + DROP COLUMN drive_file_id, + DROP COLUMN drive_synced_at; diff --git a/migrations/sql/up/add_drive_sync_fields_to_photos.sql b/migrations/sql/up/add_drive_sync_fields_to_photos.sql new file mode 100644 index 00000000..2455b45d --- /dev/null +++ b/migrations/sql/up/add_drive_sync_fields_to_photos.sql @@ -0,0 +1,3 @@ +ALTER TABLE photos + ADD COLUMN drive_file_id text, + ADD COLUMN drive_synced_at timestamp with time zone; diff --git a/migrations/versions/af58506b53a9_add_drive_sync_fields_to_photos.py b/migrations/versions/af58506b53a9_add_drive_sync_fields_to_photos.py new file mode 100644 index 00000000..bb6cac94 --- /dev/null +++ b/migrations/versions/af58506b53a9_add_drive_sync_fields_to_photos.py @@ -0,0 +1,25 @@ +"""add_drive_sync_fields_to_photos + +Revision ID: af58506b53a9 +Revises: 5425a051d68c +Create Date: 2026-08-25 02:58:36.347923 + +""" +from typing import Sequence, Union + +from migrations.helper import run_sql_up, run_sql_down + + +# revision identifiers, used by Alembic. +revision: str = 'af58506b53a9' +down_revision: Union[str, Sequence[str], None] = '5425a051d68c' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + run_sql_up("add_drive_sync_fields_to_photos") + + +def downgrade() -> None: + run_sql_down("add_drive_sync_fields_to_photos") From 19b2879c87729dd0fac8c0fc60413e20925d9ba8 Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 03:01:27 +0100 Subject: [PATCH 17/29] feat: publish Drive sync event for direct-uploaded photos on approval --- app/service/upload_requests.py | 30 +++++++++++++++++++ tests/unit/test_direct_uploads.py | 48 ++++++++++++++++++++++++++++++- 2 files changed, 77 insertions(+), 1 deletion(-) diff --git a/app/service/upload_requests.py b/app/service/upload_requests.py index f64e346b..5af3cdd4 100644 --- a/app/service/upload_requests.py +++ b/app/service/upload_requests.py @@ -504,6 +504,34 @@ async def _publish_photo_process_events(self, photos: list[Photo]) -> None: if photos: logger.info("Published %d photo process events", len(photos)) + async def _publish_drive_sync_events( + self, + staged_photos: list[UploadRequestPhoto], + created_photos: list[Photo], + ) -> None: + # staged_photos and created_photos are built 1:1 in the same order by + # _approve_request_without_side_effects (and accumulated in lockstep + # across multiple calls for a group approval), so zipping them here + # is safe. Only direct-uploaded photos get synced — Drive-imported + # photos already live in Drive and syncing them back would be a + # pointless round trip. + count = 0 + for staged_photo, created_photo in zip(staged_photos, created_photos): + if getattr(staged_photo, "source", "drive") != "direct": + continue + await self._publish_event( + subject=NatsSubjects.PHOTO_DRIVE_SYNC_REQUESTED, + payload={ + "photo_id": str(created_photo.id), + "storage_key": created_photo.storage_key, + "file_name": staged_photo.file_name, + "mime_type": staged_photo.mime_type, + }, + ) + count += 1 + if count: + logger.info("Published %d Drive sync events", count) + async def _mark_group_import_failed( self, *, @@ -1153,6 +1181,7 @@ async def approve_request( }, ) await self._publish_photo_process_events(created_photos) + await self._publish_drive_sync_events(staged_photos, created_photos) await self._audit( AuditEventType.UPLOAD_REQUEST_APPROVED, request_id=upload_request.id, @@ -1285,6 +1314,7 @@ async def approve_group( }, ) await self._publish_photo_process_events(all_created_photos) + await self._publish_drive_sync_events(all_staged_photos, all_created_photos) await self._audit( AuditEventType.UPLOAD_REQUEST_APPROVED, group_id=upload_group.id, diff --git a/tests/unit/test_direct_uploads.py b/tests/unit/test_direct_uploads.py index ca73237f..b1435af3 100644 --- a/tests/unit/test_direct_uploads.py +++ b/tests/unit/test_direct_uploads.py @@ -1,6 +1,6 @@ import uuid from datetime import datetime, timezone -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, patch import pytest @@ -336,6 +336,52 @@ async def _photos_iter(dollar_1): assert url == "https://minio.local/resumed" +@pytest.mark.asyncio +async def test_approve_request_publishes_drive_sync_event_for_direct_photo_only( + upload_requests_service, + mock_upload_request_querier, + mock_upload_request_photo_querier, + mock_photo_querier, + mock_staged_upload_storage, + mock_staff_user, +): + from db.generated.models import Photo + + request_id = uuid.uuid4() + event_id = uuid.uuid4() + photo_id = uuid.uuid4() + + mock_upload_request_querier.get_upload_request_by_id.return_value = _make_request( + request_id, event_id, mock_staff_user.id, None, + ) + + async def _photos_iter(upload_request_id): + yield _make_photo(photo_id, request_id, transfer_status="uploaded", source="direct") + + mock_upload_request_photo_querier.list_upload_request_photos_by_upload_request_id = _photos_iter + mock_staged_upload_storage.promote_to_final.return_value = "events/e1/p1.jpg" + mock_photo_querier.create_photo.return_value = Photo( + id=photo_id, event_id=event_id, uploaded_by=None, storage_key="events/e1/p1.jpg", + taken_at=None, day_number=None, visibility="private", status="pending", + created_at=datetime.now(timezone.utc), drive_file_id=None, drive_synced_at=None, + ) + mock_upload_request_photo_querier.update_upload_request_photo_approval.return_value = _make_photo( + photo_id, request_id, source="direct", transfer_status="uploaded", + ) + mock_upload_request_querier.approve_upload_request.return_value = _make_request( + request_id, event_id, mock_staff_user.id, None, + ) + + with patch("app.service.upload_requests.NatsClient.publish") as mock_publish: + await upload_requests_service.approve_request( + request_id=request_id, approved_by=mock_staff_user, + ) + + published_subjects = [call.args[0] for call in mock_publish.call_args_list] + from app.infra.nats import NatsSubjects + assert NatsSubjects.PHOTO_DRIVE_SYNC_REQUESTED in published_subjects + + @pytest.mark.asyncio async def test_fail_direct_upload_marks_transfer_failed( upload_requests_service, From 6cdc4ee3c5cdc86ce09b6e4eb50b36017b483481 Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 03:04:09 +0100 Subject: [PATCH 18/29] feat: add drive_sync worker to sync approved direct-upload photos to Drive --- app/worker/drive_sync/__init__.py | 0 app/worker/drive_sync/main.py | 103 ++++++++++++++++++++++++++++++ makefile | 1 + 3 files changed, 104 insertions(+) create mode 100644 app/worker/drive_sync/__init__.py create mode 100644 app/worker/drive_sync/main.py diff --git a/app/worker/drive_sync/__init__.py b/app/worker/drive_sync/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/app/worker/drive_sync/main.py b/app/worker/drive_sync/main.py new file mode 100644 index 00000000..f642fa1b --- /dev/null +++ b/app/worker/drive_sync/main.py @@ -0,0 +1,103 @@ +import asyncio +import json +import uuid + +from pydantic import BaseModel, ValidationError + +from app.core.config import settings +from app.core.logger import logger +from app.infra.database import engine +from app.infra.minio import Bucket, IMAGES_BUCKET_NAME, init_minio_client +from app.infra.nats import NatsClient, NatsSubjects +from app.infra.redis import RedisClient +from app.service.staff_drive import StaffDriveService +from db.generated import photos as photo_queries +from db.generated import staff_drive_connections as drive_queries +from db.generated import staff_user as staff_queries + + +class PhotoDriveSyncEvent(BaseModel): + photo_id: uuid.UUID + storage_key: str + file_name: str + mime_type: str + + +def _parse_payload(raw_data: bytes) -> PhotoDriveSyncEvent | None: + try: + parsed = json.loads(raw_data.decode("utf-8")) + except (UnicodeDecodeError, json.JSONDecodeError) as exc: + logger.error("drive_sync: cannot parse payload: %s", exc) + return None + if not isinstance(parsed, dict): + return None + try: + return PhotoDriveSyncEvent.model_validate(parsed) + except ValidationError as exc: + logger.warning("drive_sync: payload validation failed: %s", exc) + return None + + +async def _handle_event(raw_data: bytes) -> None: + event = _parse_payload(raw_data) + if event is None: + return + + bucket = Bucket(IMAGES_BUCKET_NAME, "") + try: + data, _, content_type = await bucket.get(event.storage_key) + except Exception as exc: + logger.warning("drive_sync: failed to read photo %s from storage: %s", event.photo_id, exc) + return + + async with engine.begin() as conn: + staff_drive_service = StaffDriveService( + staff_user_querier=staff_queries.AsyncQuerier(conn), + drive_connection_querier=drive_queries.AsyncQuerier(conn), + redis=RedisClient.get_instance(), + ) + photo_querier = photo_queries.AsyncQuerier(conn) + + try: + drive_file_id = await staff_drive_service.upload_to_system_drive( + file_name=event.file_name, + content_type=event.mime_type or content_type, + data=data, + ) + except Exception as exc: + logger.warning("drive_sync: upload failed for photo %s: %s", event.photo_id, exc) + return + + synced = await photo_querier.mark_photo_drive_synced( + id=event.photo_id, drive_file_id=drive_file_id, + ) + if synced is None: + logger.warning("drive_sync: photo %s not found when recording sync", event.photo_id) + return + + logger.info("drive_sync: synced photo %s to Drive as %s", event.photo_id, drive_file_id) + + +async def main() -> None: + logger.info("Drive sync worker starting") + await init_minio_client( + minio_host=settings.MINIO_HOST, + minio_port=settings.MINIO_API_PORT, + minio_root_user=settings.MINIO_ROOT_USER, + minio_root_password=settings.MINIO_ROOT_PASSWORD, + ) + RedisClient.init( + host=settings.REDIS_HOST, + port=settings.REDIS_PORT, + password=settings.REDIS_PASSWORD, + ) + await NatsClient.connect() + try: + await NatsClient.subscribe(NatsSubjects.PHOTO_DRIVE_SYNC_REQUESTED, _handle_event) + await asyncio.Event().wait() + finally: + await NatsClient.close() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/makefile b/makefile index 52e37c54..7e19b776 100644 --- a/makefile +++ b/makefile @@ -67,6 +67,7 @@ run-workers: uv run python -m app.worker.email_worker.main & \ uv run python -m app.worker.event_lifecycle.main & \ uv run python -m app.worker.upload_reconciler.main & \ + uv run python -m app.worker.drive_sync.main & \ wait lint: From 7e780cee26ebaf5231549a9d389157e6ea430a60 Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 03:20:55 +0100 Subject: [PATCH 19/29] feat: add source and storage_cleaned_at fields to photos, implement auto-approval for direct uploads --- app/core/config.py | 4 + app/service/upload_requests.py | 45 ++++++++ db/generated/models.py | 2 + db/generated/photos.py | 104 ++++++++++++++++-- db/queries/photos.sql | 23 +++- .../down/add_photo_storage_cleanup_fields.sql | 3 + .../up/add_photo_storage_cleanup_fields.sql | 3 + ...1e711a_add_photo_storage_cleanup_fields.py | 25 +++++ tests/unit/test_direct_uploads.py | 1 + 9 files changed, 197 insertions(+), 13 deletions(-) create mode 100644 migrations/sql/down/add_photo_storage_cleanup_fields.sql create mode 100644 migrations/sql/up/add_photo_storage_cleanup_fields.sql create mode 100644 migrations/versions/7eb80a1e711a_add_photo_storage_cleanup_fields.py diff --git a/app/core/config.py b/app/core/config.py index 9d01ffa0..23e34af3 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -39,6 +39,10 @@ class Settings(BaseSettings): DIRECT_UPLOAD_STALE_PENDING_MINUTES: int = 45 DIRECT_UPLOAD_RECONCILE_POLL_INTERVAL_SECONDS: int = 300 DIRECT_UPLOAD_MAX_BATCH_SIZE: int = 200 + # Dev/testing convenience: when true, a direct-upload group auto-approves + # itself the moment every photo in it has been confirmed uploaded, instead + # of waiting for a team lead to approve manually. + AUTO_APPROVE: bool = True # Mobile auth/session defaults MOBILE_SESSION_LIMIT: int = 3 diff --git a/app/service/upload_requests.py b/app/service/upload_requests.py index 5af3cdd4..cedd8c34 100644 --- a/app/service/upload_requests.py +++ b/app/service/upload_requests.py @@ -366,6 +366,7 @@ async def _approve_request_without_side_effects( taken_at=staged_photo.taken_at, day_number=staged_photo.day_number, visibility=staged_photo.visibility, + source=getattr(staged_photo, "source", "drive"), ) ) if created_photo is None: @@ -753,8 +754,52 @@ async def confirm_direct_upload( ) if confirmed is None: raise AppException.internal_error("Failed to confirm upload") + + if settings.AUTO_APPROVE: + await self._maybe_auto_approve_group( + upload_request_id=confirmed.upload_request_id, + approved_by=requested_by, + ) + return confirmed + async def _maybe_auto_approve_group( + self, + *, + upload_request_id: uuid.UUID, + approved_by: StaffUser, + ) -> None: + upload_request = await self.upload_request_querier.get_upload_request_by_id(id=upload_request_id) + if upload_request is None or upload_request.group_id is None: + return + if self._status_value(upload_request.status) != "pending": + return + + group_id = upload_request.group_id + request_ids: list[uuid.UUID] = [] + async for req in self.upload_request_querier.list_upload_requests_by_group_id(group_id=group_id): + if self._status_value(req.status) != "pending": + continue + request_ids.append(req.id) + if not request_ids: + return + + all_uploaded = True + async for photo in self.upload_request_photo_querier.list_upload_request_photos_by_upload_request_ids( + dollar_1=request_ids + ): + if getattr(photo, "transfer_status", "uploaded") != "uploaded": + all_uploaded = False + break + if not all_uploaded: + return + + try: + await self.approve_group(group_id=group_id, approved_by=approved_by) + logger.info("auto_approve: group %s approved automatically", group_id) + except Exception: + logger.exception("auto_approve: failed to auto-approve group %s", group_id) + async def fail_direct_upload( self, *, diff --git a/db/generated/models.py b/db/generated/models.py index 7d935995..adf1f28a 100644 --- a/db/generated/models.py +++ b/db/generated/models.py @@ -117,6 +117,8 @@ class Photo: created_at: datetime.datetime drive_file_id: Optional[str] drive_synced_at: Optional[datetime.datetime] + source: str + storage_cleaned_at: Optional[datetime.datetime] @dataclasses.dataclass() diff --git a/db/generated/photos.py b/db/generated/photos.py index a978143c..cf6df272 100644 --- a/db/generated/photos.py +++ b/db/generated/photos.py @@ -31,11 +31,12 @@ storage_key, taken_at, day_number, - visibility + visibility, + source ) VALUES ( - :p1, :p2, :p3, :p4, :p5 + :p1, :p2, :p3, :p4, :p5, :p6 ) -RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at +RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at, source, storage_cleaned_at """ @@ -46,6 +47,7 @@ class CreatePhotoParams: taken_at: Optional[datetime.datetime] day_number: Optional[int] visibility: str + source: str GET_DRIVE_FILE_ID_FOR_PHOTO = """-- name: get_drive_file_id_for_photo \\:one @@ -57,12 +59,12 @@ class CreatePhotoParams: GET_PHOTO_BY_ID = """-- name: get_photo_by_id \\:one -SELECT id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at FROM photos WHERE id = :p1 +SELECT id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at, source, storage_cleaned_at FROM photos WHERE id = :p1 """ LIST_EVENT_PHOTOS_FOR_USER = """-- name: list_event_photos_for_user \\:many -SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, p.drive_file_id, p.drive_synced_at, +SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, p.drive_file_id, p.drive_synced_at, p.source, p.storage_cleaned_at, (SELECT COUNT(*) FROM photo_faces pf2 WHERE pf2.photo_id = p.id)\\:\\:int AS face_count FROM photos p WHERE p.event_id = :p2 @@ -108,11 +110,27 @@ class ListEventPhotosForUserRow: created_at: datetime.datetime drive_file_id: Optional[str] drive_synced_at: Optional[datetime.datetime] + source: str + storage_cleaned_at: Optional[datetime.datetime] face_count: int +LIST_PHOTOS_DUE_FOR_STORAGE_CLEANUP = """-- name: list_photos_due_for_storage_cleanup \\:many +SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, p.drive_file_id, p.drive_synced_at, p.source, p.storage_cleaned_at +FROM photos p +JOIN events e ON e.id = p.event_id +WHERE p.status = 'approved' + AND p.storage_cleaned_at IS NULL + AND e.end_date IS NOT NULL + AND e.end_date <= NOW() - (:p1 || ' days')\\:\\:interval + AND (p.source != 'direct' OR p.drive_synced_at IS NOT NULL) +ORDER BY e.end_date ASC +LIMIT 500 +""" + + LIST_USER_PHOTOS = """-- name: list_user_photos \\:many -SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, p.drive_file_id, p.drive_synced_at, +SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, p.drive_file_id, p.drive_synced_at, p.source, p.storage_cleaned_at, (SELECT COUNT(*) FROM photo_faces pf2 WHERE pf2.photo_id = p.id)\\:\\:int AS face_count FROM photos p WHERE ( @@ -156,6 +174,8 @@ class ListUserPhotosRow: created_at: datetime.datetime drive_file_id: Optional[str] drive_synced_at: Optional[datetime.datetime] + source: str + storage_cleaned_at: Optional[datetime.datetime] face_count: int @@ -164,7 +184,15 @@ class ListUserPhotosRow: SET drive_file_id = :p2, drive_synced_at = NOW() WHERE id = :p1 -RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at +RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at, source, storage_cleaned_at +""" + + +MARK_PHOTO_STORAGE_CLEANED = """-- name: mark_photo_storage_cleaned \\:one +UPDATE photos +SET storage_cleaned_at = NOW() +WHERE id = :p1 +RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at, source, storage_cleaned_at """ @@ -172,7 +200,7 @@ class ListUserPhotosRow: UPDATE photos SET status = :p2 WHERE id = :p1 -RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at +RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at, source, storage_cleaned_at """ @@ -180,7 +208,7 @@ class ListUserPhotosRow: UPDATE photos SET visibility = :p2 WHERE id = :p1 -RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at +RETURNING id, event_id, uploaded_by, storage_key, taken_at, day_number, visibility, status, created_at, drive_file_id, drive_synced_at, source, storage_cleaned_at """ @@ -201,6 +229,7 @@ async def create_photo(self, arg: CreatePhotoParams) -> Optional[models.Photo]: "p3": arg.taken_at, "p4": arg.day_number, "p5": arg.visibility, + "p6": arg.source, })).first() if row is None: return None @@ -216,6 +245,8 @@ async def create_photo(self, arg: CreatePhotoParams) -> Optional[models.Photo]: created_at=row[8], drive_file_id=row[9], drive_synced_at=row[10], + source=row[11], + storage_cleaned_at=row[12], ) async def get_drive_file_id_for_photo(self, *, final_storage_key: Optional[str]) -> Optional[Optional[str]]: @@ -240,6 +271,8 @@ async def get_photo_by_id(self, *, id: uuid.UUID) -> Optional[models.Photo]: created_at=row[8], drive_file_id=row[9], drive_synced_at=row[10], + source=row[11], + storage_cleaned_at=row[12], ) async def list_event_photos_for_user(self, arg: ListEventPhotosForUserParams) -> AsyncIterator[ListEventPhotosForUserRow]: @@ -263,7 +296,28 @@ async def list_event_photos_for_user(self, arg: ListEventPhotosForUserParams) -> created_at=row[8], drive_file_id=row[9], drive_synced_at=row[10], - face_count=row[11], + source=row[11], + storage_cleaned_at=row[12], + face_count=row[13], + ) + + async def list_photos_due_for_storage_cleanup(self, *, dollar_1: Optional[str]) -> AsyncIterator[models.Photo]: + result = await self._conn.stream(sqlalchemy.text(LIST_PHOTOS_DUE_FOR_STORAGE_CLEANUP), {"p1": dollar_1}) + async for row in result: + yield models.Photo( + id=row[0], + event_id=row[1], + uploaded_by=row[2], + storage_key=row[3], + taken_at=row[4], + day_number=row[5], + visibility=row[6], + status=row[7], + created_at=row[8], + drive_file_id=row[9], + drive_synced_at=row[10], + source=row[11], + storage_cleaned_at=row[12], ) async def list_user_photos(self, arg: ListUserPhotosParams) -> AsyncIterator[ListUserPhotosRow]: @@ -287,7 +341,9 @@ async def list_user_photos(self, arg: ListUserPhotosParams) -> AsyncIterator[Lis created_at=row[8], drive_file_id=row[9], drive_synced_at=row[10], - face_count=row[11], + source=row[11], + storage_cleaned_at=row[12], + face_count=row[13], ) async def mark_photo_drive_synced(self, *, id: uuid.UUID, drive_file_id: Optional[str]) -> Optional[models.Photo]: @@ -306,6 +362,28 @@ async def mark_photo_drive_synced(self, *, id: uuid.UUID, drive_file_id: Optiona created_at=row[8], drive_file_id=row[9], drive_synced_at=row[10], + source=row[11], + storage_cleaned_at=row[12], + ) + + async def mark_photo_storage_cleaned(self, *, id: uuid.UUID) -> Optional[models.Photo]: + row = (await self._conn.execute(sqlalchemy.text(MARK_PHOTO_STORAGE_CLEANED), {"p1": id})).first() + if row is None: + return None + return models.Photo( + id=row[0], + event_id=row[1], + uploaded_by=row[2], + storage_key=row[3], + taken_at=row[4], + day_number=row[5], + visibility=row[6], + status=row[7], + created_at=row[8], + drive_file_id=row[9], + drive_synced_at=row[10], + source=row[11], + storage_cleaned_at=row[12], ) async def update_photo_status(self, *, id: uuid.UUID, status: Any) -> Optional[models.Photo]: @@ -324,6 +402,8 @@ async def update_photo_status(self, *, id: uuid.UUID, status: Any) -> Optional[m created_at=row[8], drive_file_id=row[9], drive_synced_at=row[10], + source=row[11], + storage_cleaned_at=row[12], ) async def update_photo_visibility(self, *, id: uuid.UUID, visibility: str) -> Optional[models.Photo]: @@ -342,4 +422,6 @@ async def update_photo_visibility(self, *, id: uuid.UUID, visibility: str) -> Op created_at=row[8], drive_file_id=row[9], drive_synced_at=row[10], + source=row[11], + storage_cleaned_at=row[12], ) diff --git a/db/queries/photos.sql b/db/queries/photos.sql index 9904ad8a..7ad41787 100644 --- a/db/queries/photos.sql +++ b/db/queries/photos.sql @@ -4,9 +4,10 @@ INSERT INTO photos ( storage_key, taken_at, day_number, - visibility + visibility, + source ) VALUES ( - $1, $2, $3, $4, $5 + $1, $2, $3, $4, $5, $6 ) RETURNING *; @@ -91,3 +92,21 @@ SET drive_file_id = $2, drive_synced_at = NOW() WHERE id = $1 RETURNING *; + +-- name: ListPhotosDueForStorageCleanup :many +SELECT p.* +FROM photos p +JOIN events e ON e.id = p.event_id +WHERE p.status = 'approved' + AND p.storage_cleaned_at IS NULL + AND e.end_date IS NOT NULL + AND e.end_date <= NOW() - ($1 || ' days')::interval + AND (p.source != 'direct' OR p.drive_synced_at IS NOT NULL) +ORDER BY e.end_date ASC +LIMIT 500; + +-- name: MarkPhotoStorageCleaned :one +UPDATE photos +SET storage_cleaned_at = NOW() +WHERE id = $1 +RETURNING *; diff --git a/migrations/sql/down/add_photo_storage_cleanup_fields.sql b/migrations/sql/down/add_photo_storage_cleanup_fields.sql new file mode 100644 index 00000000..eaa05089 --- /dev/null +++ b/migrations/sql/down/add_photo_storage_cleanup_fields.sql @@ -0,0 +1,3 @@ +ALTER TABLE photos + DROP COLUMN source, + DROP COLUMN storage_cleaned_at; diff --git a/migrations/sql/up/add_photo_storage_cleanup_fields.sql b/migrations/sql/up/add_photo_storage_cleanup_fields.sql new file mode 100644 index 00000000..05e690ea --- /dev/null +++ b/migrations/sql/up/add_photo_storage_cleanup_fields.sql @@ -0,0 +1,3 @@ +ALTER TABLE photos + ADD COLUMN source character varying(16) DEFAULT 'drive'::character varying NOT NULL, + ADD COLUMN storage_cleaned_at timestamp with time zone; diff --git a/migrations/versions/7eb80a1e711a_add_photo_storage_cleanup_fields.py b/migrations/versions/7eb80a1e711a_add_photo_storage_cleanup_fields.py new file mode 100644 index 00000000..ac1a6835 --- /dev/null +++ b/migrations/versions/7eb80a1e711a_add_photo_storage_cleanup_fields.py @@ -0,0 +1,25 @@ +"""add_photo_storage_cleanup_fields + +Revision ID: 7eb80a1e711a +Revises: af58506b53a9 +Create Date: 2026-08-25 03:15:07.646415 + +""" +from typing import Sequence, Union + +from migrations.helper import run_sql_up, run_sql_down + + +# revision identifiers, used by Alembic. +revision: str = '7eb80a1e711a' +down_revision: Union[str, Sequence[str], None] = 'af58506b53a9' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + run_sql_up("add_photo_storage_cleanup_fields") + + +def downgrade() -> None: + run_sql_down("add_photo_storage_cleanup_fields") diff --git a/tests/unit/test_direct_uploads.py b/tests/unit/test_direct_uploads.py index b1435af3..68b604fe 100644 --- a/tests/unit/test_direct_uploads.py +++ b/tests/unit/test_direct_uploads.py @@ -364,6 +364,7 @@ async def _photos_iter(upload_request_id): id=photo_id, event_id=event_id, uploaded_by=None, storage_key="events/e1/p1.jpg", taken_at=None, day_number=None, visibility="private", status="pending", created_at=datetime.now(timezone.utc), drive_file_id=None, drive_synced_at=None, + source="direct", storage_cleaned_at=None, ) mock_upload_request_photo_querier.update_upload_request_photo_approval.return_value = _make_photo( photo_id, request_id, source="direct", transfer_status="uploaded", From 7186410c2e3b28739c1f87699fc848b11519bbba Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 03:22:50 +0100 Subject: [PATCH 20/29] feat: update photo cleanup logic to be threshold-based and remove immediate scheduling --- app/core/config.py | 5 +++ app/worker/photo_worker/main.py | 2 - .../photo_worker/tests/test_photo_worker.py | 38 +++++++++---------- 3 files changed, 24 insertions(+), 21 deletions(-) diff --git a/app/core/config.py b/app/core/config.py index 23e34af3..ed8613f5 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -35,6 +35,11 @@ class Settings(BaseSettings): PHOTO_APPROVAL_TIMEOUT_DAYS: int = 7 EVENT_LIFECYCLE_POLL_INTERVAL_SECONDS: int = 60 + # How long after an event's end_date approved photos stay in MinIO before + # being cleaned up. Direct-uploaded photos additionally require a + # confirmed Drive sync before cleanup is eligible (see + # ListPhotosDueForStorageCleanup) — MinIO is their only copy until then. + PHOTO_STORAGE_RETENTION_DAYS_AFTER_EVENT_END: int = 20 DIRECT_UPLOAD_PRESIGN_EXPIRES_SECONDS: int = 1800 DIRECT_UPLOAD_STALE_PENDING_MINUTES: int = 45 DIRECT_UPLOAD_RECONCILE_POLL_INTERVAL_SECONDS: int = 300 diff --git a/app/worker/photo_worker/main.py b/app/worker/photo_worker/main.py index 59a116bb..e21eeb93 100644 --- a/app/worker/photo_worker/main.py +++ b/app/worker/photo_worker/main.py @@ -86,7 +86,6 @@ async def handle_message(self, data: bytes) -> None: await self._photo_querier.update_photo_status(id=event.photo_id, status="approved") await self._photo_querier.update_photo_visibility(id=event.photo_id, visibility="public") await self._update_job(job, "completed") - await self._schedule_cleanup(event.image_ref) return if len(faces) == 1: @@ -96,7 +95,6 @@ async def handle_message(self, data: bytes) -> None: await self._update_job(job, "completed") await self._publish_audit(event, len(faces)) - await self._schedule_cleanup(event.image_ref) async def _handle_single_face(self, event: PhotoProcessEvent, face: DetectedFace) -> None: diff --git a/app/worker/photo_worker/tests/test_photo_worker.py b/app/worker/photo_worker/tests/test_photo_worker.py index f113a706..1eeb73e3 100644 --- a/app/worker/photo_worker/tests/test_photo_worker.py +++ b/app/worker/photo_worker/tests/test_photo_worker.py @@ -281,16 +281,24 @@ async def test_image_load_failure_is_handled( # ── cleanup scheduling tests ────────────────────────────────────────── +# +# photo_worker no longer schedules immediate MinIO cleanup after processing. +# Cleanup is now threshold-based (event_lifecycle worker, 20 days after +# event.end_date, gated on Drive sync for direct-uploaded photos) — see +# ListPhotosDueForStorageCleanup. Immediate cleanup right after processing +# was a race: for direct uploads, MinIO is the only copy of the photo until +# the async Drive sync completes, so deleting it immediately could destroy +# the only copy before it was ever backed up. @pytest.mark.asyncio -async def test_cleanup_scheduled_after_single_face( +async def test_single_face_does_not_schedule_cleanup( worker: PhotoWorker, face_service: AsyncMock, single_face_service: AsyncMock, event: PhotoProcessEvent, ) -> None: - """After single face processing, both audit and cleanup events should be published.""" + """After single face processing, only the audit event is published — no cleanup.""" face_service.detect_faces = AsyncMock(return_value=[_make_face()]) with ( @@ -304,23 +312,19 @@ async def test_cleanup_scheduled_after_single_face( await worker.handle_message(_event_bytes(event)) from app.infra.nats import NatsSubjects - assert mock_nats.publish.call_count == 2 + assert mock_nats.publish.call_count == 1 audit_call = mock_nats.publish.call_args_list[0] assert audit_call.args[0] == NatsSubjects.AUDIT_EVENT - cleanup_call = mock_nats.publish.call_args_list[1] - assert cleanup_call.args[0] == NatsSubjects.FINAL_BUCKET_CLEANUP - cleanup_payload = json.loads(cleanup_call.args[1]) - assert event.image_ref in cleanup_payload["storage_keys"] @pytest.mark.asyncio -async def test_cleanup_scheduled_after_group_photo( +async def test_group_photo_does_not_schedule_cleanup( worker: PhotoWorker, face_service: AsyncMock, photo_face_querier: AsyncMock, event: PhotoProcessEvent, ) -> None: - """After group photo processing, both audit and cleanup events should be published.""" + """After group photo processing, only the audit event is published — no cleanup.""" faces = [_make_face(), _make_face()] face_service.detect_faces = AsyncMock(return_value=faces) photo_face_querier.insert_photo_face_with_approval = AsyncMock(return_value=None) @@ -336,20 +340,18 @@ async def test_cleanup_scheduled_after_group_photo( await worker.handle_message(_event_bytes(event)) from app.infra.nats import NatsSubjects - assert mock_nats.publish.call_count == 2 - cleanup_call = mock_nats.publish.call_args_list[1] - assert cleanup_call.args[0] == NatsSubjects.FINAL_BUCKET_CLEANUP - cleanup_payload = json.loads(cleanup_call.args[1]) - assert event.image_ref in cleanup_payload["storage_keys"] + assert mock_nats.publish.call_count == 1 + audit_call = mock_nats.publish.call_args_list[0] + assert audit_call.args[0] == NatsSubjects.AUDIT_EVENT @pytest.mark.asyncio -async def test_cleanup_scheduled_when_no_faces( +async def test_no_cleanup_when_no_faces( worker: PhotoWorker, face_service: AsyncMock, event: PhotoProcessEvent, ) -> None: - """Even if no faces detected, cleanup should still be scheduled.""" + """No faces detected: photo is marked public, but no cleanup is scheduled.""" face_service.detect_faces = AsyncMock(return_value=[]) with ( @@ -362,9 +364,7 @@ async def test_cleanup_scheduled_when_no_faces( ) await worker.handle_message(_event_bytes(event)) - mock_nats.publish.assert_called_once() - cleanup_payload = json.loads(mock_nats.publish.call_args.args[1]) - assert event.image_ref in cleanup_payload["storage_keys"] + mock_nats.publish.assert_not_called() @pytest.mark.asyncio From 9017a6d5902f7ccec40da070bb1a757e17a2082c Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 03:25:20 +0100 Subject: [PATCH 21/29] feat: refactor photo cleanup scheduling to be handled by event lifecycle worker --- app/service/user_photo.py | 14 +++++++--- app/worker/event_lifecycle/main.py | 42 +++++++++++++++++++++++++++++- app/worker/photo_worker/main.py | 9 ------- tests/unit/test_photo_worker.py | 10 ++++--- 4 files changed, 57 insertions(+), 18 deletions(-) diff --git a/app/service/user_photo.py b/app/service/user_photo.py index 58204960..b85f03cd 100644 --- a/app/service/user_photo.py +++ b/app/service/user_photo.py @@ -113,10 +113,16 @@ async def get_photo_bytes( except Exception: logger.info("Photo %s not in bucket, trying Drive fallback", photo_id) - # Fallback: get drive_file_id from upload_request_photos - drive_file_id = await self._photo_querier.get_drive_file_id_for_photo( - final_storage_key=photo.storage_key, - ) + # Fallback: photos.drive_file_id (set once drive_sync confirms an + # approved direct-upload photo has been synced) takes priority since + # it's the authoritative post-approval copy; upload_request_photos' + # drive_file_id (the original *source* file for Drive-imported + # photos) is the older path, kept for backward compatibility. + drive_file_id = photo.drive_file_id + if drive_file_id is None: + drive_file_id = await self._photo_querier.get_drive_file_id_for_photo( + final_storage_key=photo.storage_key, + ) if drive_file_id is None: raise AppException.not_found("Photo no longer available") diff --git a/app/worker/event_lifecycle/main.py b/app/worker/event_lifecycle/main.py index 7ab11b8d..51aa6eff 100644 --- a/app/worker/event_lifecycle/main.py +++ b/app/worker/event_lifecycle/main.py @@ -1,9 +1,12 @@ import asyncio +import json from app.core.config import settings from app.core.logger import logger from app.infra.database import engine +from app.infra.nats import NatsClient, NatsSubjects from db.generated import events as event_queries +from db.generated import photos as photo_queries async def run_lifecycle_pass() -> None: @@ -19,16 +22,53 @@ async def run_lifecycle_pass() -> None: logger.info("event_lifecycle: archived %d event(s): %s", len(archived), archived) +async def run_storage_cleanup_pass() -> None: + async with engine.begin() as conn: + querier = photo_queries.AsyncQuerier(conn) + + due_photos = [ + photo + async for photo in querier.list_photos_due_for_storage_cleanup( + dollar_1=str(settings.PHOTO_STORAGE_RETENTION_DAYS_AFTER_EVENT_END) + ) + ] + if not due_photos: + return + + cleaned = 0 + for photo in due_photos: + try: + await NatsClient.publish( + NatsSubjects.FINAL_BUCKET_CLEANUP, + json.dumps({"storage_keys": [photo.storage_key]}).encode("utf-8"), + ) + except Exception as exc: + logger.warning( + "storage_cleanup: failed to schedule cleanup for photo %s: %s", photo.id, exc + ) + continue + marked = await querier.mark_photo_storage_cleaned(id=photo.id) + if marked is not None: + cleaned += 1 + + logger.info("storage_cleanup: scheduled cleanup for %d photo(s)", cleaned) + + async def main() -> None: logger.info( - "Event lifecycle worker starting, poll_interval=%ds", + "Event lifecycle worker starting, poll_interval=%ds, storage_retention=%dd after event end", settings.EVENT_LIFECYCLE_POLL_INTERVAL_SECONDS, + settings.PHOTO_STORAGE_RETENTION_DAYS_AFTER_EVENT_END, ) while True: try: await run_lifecycle_pass() except Exception: logger.exception("event_lifecycle: pass failed") + try: + await run_storage_cleanup_pass() + except Exception: + logger.exception("storage_cleanup: pass failed") await asyncio.sleep(settings.EVENT_LIFECYCLE_POLL_INTERVAL_SECONDS) diff --git a/app/worker/photo_worker/main.py b/app/worker/photo_worker/main.py index e21eeb93..7fa988d2 100644 --- a/app/worker/photo_worker/main.py +++ b/app/worker/photo_worker/main.py @@ -205,15 +205,6 @@ async def _publish_audit(event: PhotoProcessEvent, faces_count: int) -> None: except Exception as exc: logger.warning("Failed to publish audit for photo %s: %s", event.photo_id, exc) - @staticmethod - async def _schedule_cleanup(image_ref: str) -> None: - payload = json.dumps({"storage_keys": [image_ref]}).encode("utf-8") - try: - await NatsClient.publish(NatsSubjects.FINAL_BUCKET_CLEANUP, payload) - logger.info("Scheduled cleanup for %s", image_ref) - except Exception as exc: - logger.warning("Failed to schedule cleanup for %s: %s", image_ref, exc) - @staticmethod def _parse_event(raw_data: bytes) -> PhotoProcessEvent | None: try: diff --git a/tests/unit/test_photo_worker.py b/tests/unit/test_photo_worker.py index 9b920307..0091223d 100644 --- a/tests/unit/test_photo_worker.py +++ b/tests/unit/test_photo_worker.py @@ -1,4 +1,3 @@ -import json import uuid from unittest.mock import AsyncMock, patch, MagicMock @@ -7,7 +6,6 @@ from app.service.face_embedding import DetectedFace, FaceImagePayload from app.worker.photo_worker.main import PhotoWorker from app.worker.photo_worker.schema.event import PhotoProcessEvent -from app.infra.nats import NatsSubjects from db.generated import models @@ -100,7 +98,10 @@ async def test_handle_message_success_no_faces(photo_worker, sample_event, mock_ mock_pj_querier.update_processing_job_status.assert_any_call(id=mock_pj_querier.create_processing_job.return_value.id, status="completed") mock_photo_querier.update_photo_status.assert_called_once_with(id=sample_event.photo_id, status="approved") mock_photo_querier.update_photo_visibility.assert_called_once_with(id=sample_event.photo_id, visibility="public") - mock_publish.assert_called_with(NatsSubjects.FINAL_BUCKET_CLEANUP, json.dumps({"storage_keys": [sample_event.image_ref]}).encode("utf-8")) + # photo_worker no longer schedules immediate MinIO cleanup — that's + # now threshold-based (event_lifecycle worker, gated on event.end_date + # + Drive sync confirmation for direct uploads). + mock_publish.assert_not_called() @pytest.mark.asyncio @@ -114,7 +115,8 @@ async def test_handle_message_success_single_face(photo_worker, sample_event, mo mock_pj_querier.update_processing_job_status.assert_any_call(id=mock_pj_querier.create_processing_job.return_value.id, status="completed") mock_single_face_service.process_detected_face.assert_called_once() - assert mock_publish.call_count == 2 + # Only the audit event — cleanup is no longer scheduled by photo_worker. + assert mock_publish.call_count == 1 @pytest.mark.asyncio From 98b7f11104e00ebe3b576b9416a8317f396398f4 Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 03:44:36 +0100 Subject: [PATCH 22/29] fix: hide unapproved photos from the general user photo gallery MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ListUserPhotos (GET /user/photos) was missing the status='approved' filter that ListEventPhotosForUser already has, so a photo a user was still pending/rejected on could leak into the general gallery while correctly staying hidden from the event-scoped view — same photo, inconsistent visibility depending on which endpoint you hit. --- db/generated/photos.py | 3 ++- db/queries/photos.sql | 3 ++- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/db/generated/photos.py b/db/generated/photos.py index cf6df272..666c9896 100644 --- a/db/generated/photos.py +++ b/db/generated/photos.py @@ -133,7 +133,8 @@ class ListEventPhotosForUserRow: SELECT p.id, p.event_id, p.uploaded_by, p.storage_key, p.taken_at, p.day_number, p.visibility, p.status, p.created_at, p.drive_file_id, p.drive_synced_at, p.source, p.storage_cleaned_at, (SELECT COUNT(*) FROM photo_faces pf2 WHERE pf2.photo_id = p.id)\\:\\:int AS face_count FROM photos p -WHERE ( +WHERE p.status = 'approved' +AND ( EXISTS ( SELECT 1 FROM photo_faces pf JOIN face_matches fm ON fm.photo_face_id = pf.id diff --git a/db/queries/photos.sql b/db/queries/photos.sql index 7ad41787..64dbc08c 100644 --- a/db/queries/photos.sql +++ b/db/queries/photos.sql @@ -30,7 +30,8 @@ RETURNING *; SELECT p.*, (SELECT COUNT(*) FROM photo_faces pf2 WHERE pf2.photo_id = p.id)::int AS face_count FROM photos p -WHERE ( +WHERE p.status = 'approved' +AND ( EXISTS ( SELECT 1 FROM photo_faces pf JOIN face_matches fm ON fm.photo_face_id = pf.id From 6d11ed0e8b1c614a9ce3c24d18445ef84b662a2c Mon Sep 17 00:00:00 2001 From: wailbentafat Date: Tue, 25 Aug 2026 03:59:50 +0100 Subject: [PATCH 23/29] feat: add self-service profile (display name) update endpoint PATCH /user/auth/me/profile lets a mobile user update their own display name, reusing AuthService.update_user which already existed but was only wired for admin/internal use. --- app/router/mobile/auth.py | 20 ++++++++++++++++++++ app/schema/request/mobile/auth.py | 4 ++++ 2 files changed, 24 insertions(+) diff --git a/app/router/mobile/auth.py b/app/router/mobile/auth.py index a7e4dddd..9be7547d 100644 --- a/app/router/mobile/auth.py +++ b/app/router/mobile/auth.py @@ -20,6 +20,7 @@ RefreshTokenRequest, UpdateDeviceTokenRequest, InactivateDeviceRequest, + UpdateProfileRequest, ) from app.schema.response.mobile.auth import MeResponse, DeviceSchema, MobileAuthResponse, SessionSchema, UserSchema, RegisterPendingResponse @@ -212,6 +213,25 @@ async def get_me( sessions=session_schema, ) +@router.patch("/me/profile", response_model=UserSchema) +async def update_profile( + req: UpdateProfileRequest, + current_user: MobileUserSchema = Depends(get_current_mobile_user), + container: Container = Depends(get_container), +) -> UserSchema: + user = await container.auth_service.update_user( + user_id=current_user.user_id, + display_name=req.name.strip(), + ) + return UserSchema( + id=user.id, + email=user.email, + name=user.display_name, + avatar_url="/user/auth/me/avatar/image" if user.avatar_key else None, + is_onboarded=user.face_embedding is not None, + ) + + @router.post("/me/avatar", response_model=UserSchema) async def upload_avatar( file: UploadFile, diff --git a/app/schema/request/mobile/auth.py b/app/schema/request/mobile/auth.py index dea9933a..fc0a312b 100644 --- a/app/schema/request/mobile/auth.py +++ b/app/schema/request/mobile/auth.py @@ -92,3 +92,7 @@ class UpdateDeviceTokenRequest(BaseModel): class InactivateDeviceRequest(BaseModel): device_id: UUID + + +class UpdateProfileRequest(BaseModel): + name: str = Field(..., min_length=1, max_length=100) From 80762354ff6dfb66dfcc38381941cb57ce42c222 Mon Sep 17 00:00:00 2001 From: Adem Boukabes <142881379+ademboukabes@users.noreply.github.com> Date: Sun, 30 Aug 2026 14:45:18 +0100 Subject: [PATCH 24/29] refactor(infra): migrate to NATS JetStream and fix tests - Switch NATS pub/sub to JetStream for persistent message queues - Update all workers to use JetStream subscriptions and proper exception propagation (nak) - Remove outdated duplicated test folder (app/worker/photo_worker/tests) which caused global state leakage - Fix CreatePhotoParams missing source argument in tests - Ensure 100% test pass rate --- app/infra/nats.py | 65 ++- app/service/audit.py | 2 +- app/service/upload_requests.py | 2 +- app/service/users.py | 4 +- app/worker/audit/main.py | 12 +- app/worker/drive_sync/main.py | 6 +- app/worker/email_worker/main.py | 6 +- app/worker/notification/main.py | 6 +- app/worker/notification/notification_queue.py | 2 +- app/worker/photo_worker/main.py | 57 +-- app/worker/photo_worker/tests/__init__.py | 0 .../photo_worker/tests/test_photo_worker.py | 384 ------------------ tests/e2e/test_photo_ai_edge_cases.py | 6 +- tests/e2e/test_photo_ai_load.py | 2 +- tests/e2e/test_photo_ai_pipeline_e2e.py | 4 +- tests/integration/test_photo_approval_flow.py | 4 + tests/unit/test_auth_email_otp.py | 6 +- tests/unit/test_direct_uploads.py | 2 +- tests/unit/test_mobile_auth_email_logging.py | 4 +- tests/unit/test_photo_worker.py | 16 +- tests/unit/test_upload_requests.py | 4 +- 21 files changed, 127 insertions(+), 467 deletions(-) delete mode 100644 app/worker/photo_worker/tests/__init__.py delete mode 100644 app/worker/photo_worker/tests/test_photo_worker.py diff --git a/app/infra/nats.py b/app/infra/nats.py index 8b07e831..cf5c1b76 100644 --- a/app/infra/nats.py +++ b/app/infra/nats.py @@ -14,6 +14,7 @@ NOTIFICATION_EVENT_SUBJECT, UPLOAD_GROUP_IMPORT_SUBJECT, ) +from app.core.logger import logger class Message(BaseModel): @@ -36,6 +37,27 @@ class NatsSubjects(Enum): STAFF_UPLOAD_REQUEST_REJECTED = "staff.upload_request.rejected" PHOTO_PROCESS = "photo.process" PHOTO_DRIVE_SYNC_REQUESTED = "photo.drive_sync.requested" + EMAIL_SEND_OTP = "email.send_otp" + + +SUBJECT_TO_STREAM: dict[str, str] = { + NatsSubjects.USER_SIGNUP.value: "auth_stream", + NatsSubjects.USER_LOGIN.value: "auth_stream", + NatsSubjects.USER_LOGOUT.value: "auth_stream", + NatsSubjects.NOTIFICATION_EVENT.value: "notification_stream", + NatsSubjects.AUDIT_EVENT.value: "audit_stream", + NatsSubjects.STAFF_UPLOAD_GROUP_IMPORT_REQUESTED.value: "upload_group_stream", + NatsSubjects.STAFF_UPLOAD_GROUP_CREATED.value: "upload_group_stream", + NatsSubjects.STAFF_UPLOAD_GROUP_APPROVED.value: "upload_group_stream", + NatsSubjects.STAFF_UPLOAD_GROUP_REJECTED.value: "upload_group_stream", + NatsSubjects.FINAL_BUCKET_CLEANUP.value: "cleanup_stream", + NatsSubjects.STAFF_UPLOAD_REQUEST_CREATED.value: "upload_request_stream", + NatsSubjects.STAFF_UPLOAD_REQUEST_APPROVED.value: "upload_request_stream", + NatsSubjects.STAFF_UPLOAD_REQUEST_REJECTED.value: "upload_request_stream", + NatsSubjects.PHOTO_PROCESS.value: "photo_process_stream", + NatsSubjects.PHOTO_DRIVE_SYNC_REQUESTED.value: "drive_sync_stream", + NatsSubjects.EMAIL_SEND_OTP.value: "email_stream", +} class NatsClient: @@ -94,36 +116,55 @@ async def _wrapper(msg: Msg) -> None: @staticmethod - async def js_publish(subject: NatsSubjects, message: bytes, stream_name: str) -> None: + async def js_publish(subject: NatsSubjects | str, message: bytes, stream_name: str | None = None) -> None: if NatsClient._js is None: await NatsClient.connect() js = NatsClient._js assert js is not None - subject_name = subject.value if isinstance(subject, NatsSubjects) else subject # type: ignore - await js.publish(subject_name, message, stream=stream_name) + subject_name = subject.value if isinstance(subject, NatsSubjects) else subject + resolved_stream = stream_name or SUBJECT_TO_STREAM.get(subject_name) + if resolved_stream is None: + logger.warning(f"No stream mapped for subject {subject_name}, but js_publish was called.") + return await NatsClient.publish(subject, message) + + await NatsClient.ensure_stream(stream_name=resolved_stream, subjects=[subject_name]) + await js.publish(subject_name, message, stream=resolved_stream) @staticmethod async def js_subscribe( - subject: NatsSubjects, + subject: NatsSubjects | str, callback: Callable[[Any], Any], - stream_name: str, - durable_name: str, + stream_name: str | None = None, + durable_name: str | None = None, ack_policy: AckPolicy = AckPolicy.EXPLICIT ) -> None: if NatsClient._js is None: await NatsClient.connect() - await NatsClient.ensure_stream(stream_name=stream_name, subjects=[subject.value]) + subject_name = subject.value if isinstance(subject, NatsSubjects) else subject + resolved_stream = stream_name or SUBJECT_TO_STREAM.get(subject_name) + if not resolved_stream: + raise ValueError(f"Cannot js_subscribe to {subject_name}: no stream mapped.") + + resolved_durable = durable_name or f"{resolved_stream}_consumer" + + await NatsClient.ensure_stream(stream_name=resolved_stream, subjects=[subject_name]) async def _wrapper(msg: Msg) -> None: - await callback(msg.data) - await msg.ack() + try: + await callback(msg.data) + await msg.ack() + except Exception as exc: + logger.error(f"Error processing message from {subject_name}, NACKing: {exc}") + await msg.nak() + raise + js = NatsClient._js assert js is not None await js.subscribe( - subject=subject.value, - stream=stream_name, - durable=durable_name, + subject=subject_name, + stream=resolved_stream, + durable=resolved_durable, cb=_wrapper, deliver_policy=DeliverPolicy.NEW, # ack_policy=ack_policy diff --git a/app/service/audit.py b/app/service/audit.py index 24a4b552..4c5ed6b5 100644 --- a/app/service/audit.py +++ b/app/service/audit.py @@ -55,7 +55,7 @@ async def create_record( metadata=metadata, description=description, ).model_dump_json() - await NatsClient.publish(NatsSubjects.AUDIT_EVENT, message.encode("utf-8")) + await NatsClient.js_publish(NatsSubjects.AUDIT_EVENT, message.encode("utf-8")) async def list_audit_events( self, diff --git a/app/service/upload_requests.py b/app/service/upload_requests.py index cedd8c34..e63de06b 100644 --- a/app/service/upload_requests.py +++ b/app/service/upload_requests.py @@ -481,7 +481,7 @@ async def _publish_event( payload: dict[str, object], ) -> None: try: - await NatsClient.publish(subject, json.dumps(payload).encode("utf-8")) + await NatsClient.js_publish(subject, json.dumps(payload).encode("utf-8")) except Exception as exc: logger.warning("Failed to publish upload request event %s: %s", subject.value, exc) diff --git a/app/service/users.py b/app/service/users.py index 5682b826..e3a30b28 100644 --- a/app/service/users.py +++ b/app/service/users.py @@ -193,7 +193,7 @@ async def mobile_register( otp = "".join(secrets.choice("0123456789") for _ in range(6)) await redis.set(f"otp:{req.email}", otp, expire=600) # Send to NATS - await NatsClient.publish("email.send_otp", json.dumps({"email": req.email, "otp": otp}).encode("utf-8")) + await NatsClient.js_publish("email.send_otp", json.dumps({"email": req.email, "otp": otp}).encode("utf-8")) logger.info("register success, OTP sent") return RegisterPendingResponse( @@ -240,7 +240,7 @@ async def mobile_register_resend_otp( # Regenerate OTP with 10 mins TTL, without touching the pending_user TTL await redis.set(f"otp:{email}", otp, expire=600) # Send to NATS - await NatsClient.publish("email.send_otp", json.dumps({"email": email, "otp": otp}).encode("utf-8")) + await NatsClient.js_publish("email.send_otp", json.dumps({"email": email, "otp": otp}).encode("utf-8")) logger.info("resend_otp success, new OTP sent to %s", email) return RegisterPendingResponse( diff --git a/app/worker/audit/main.py b/app/worker/audit/main.py index 8037c26e..7ec8ce3a 100644 --- a/app/worker/audit/main.py +++ b/app/worker/audit/main.py @@ -55,16 +55,16 @@ async def _handle_event(worker: AuditDeliveryWorker, raw_data: bytes) -> None: except ValidationError as exc: logger.warning("Audit payload validation failed: %s", exc) return - try: - await worker.persist(payload) - except Exception: - logger.exception("Failed to persist audit for %s", payload.event_type) + await worker.persist(payload) async def listen_nats_event(worker: AuditDeliveryWorker) -> None: - await NatsClient.subscribe( + async def handler(data: bytes) -> None: + await _handle_event(worker, data) + + await NatsClient.js_subscribe( NatsSubjects.AUDIT_EVENT, - lambda data: _handle_event(worker, data), + handler, ) logger.info("Listening for audit events on %s", AUDIT_EVENT_SUBJECT) diff --git a/app/worker/drive_sync/main.py b/app/worker/drive_sync/main.py index f642fa1b..2ca5be4a 100644 --- a/app/worker/drive_sync/main.py +++ b/app/worker/drive_sync/main.py @@ -48,7 +48,7 @@ async def _handle_event(raw_data: bytes) -> None: data, _, content_type = await bucket.get(event.storage_key) except Exception as exc: logger.warning("drive_sync: failed to read photo %s from storage: %s", event.photo_id, exc) - return + raise async with engine.begin() as conn: staff_drive_service = StaffDriveService( @@ -66,7 +66,7 @@ async def _handle_event(raw_data: bytes) -> None: ) except Exception as exc: logger.warning("drive_sync: upload failed for photo %s: %s", event.photo_id, exc) - return + raise synced = await photo_querier.mark_photo_drive_synced( id=event.photo_id, drive_file_id=drive_file_id, @@ -93,7 +93,7 @@ async def main() -> None: ) await NatsClient.connect() try: - await NatsClient.subscribe(NatsSubjects.PHOTO_DRIVE_SYNC_REQUESTED, _handle_event) + await NatsClient.js_subscribe(NatsSubjects.PHOTO_DRIVE_SYNC_REQUESTED, _handle_event) await asyncio.Event().wait() finally: await NatsClient.close() diff --git a/app/worker/email_worker/main.py b/app/worker/email_worker/main.py index 129d1aba..2b621b45 100644 --- a/app/worker/email_worker/main.py +++ b/app/worker/email_worker/main.py @@ -3,7 +3,7 @@ from app.core.config import settings from app.core.logger import logger -from app.infra.nats import NatsClient +from app.infra.nats import NatsClient, NatsSubjects from app.infra.email import EmailSender @@ -25,9 +25,11 @@ async def handle_message(raw_payload: bytes | str) -> None: logger.info("Successfully sent OTP email to %s", email) else: logger.error("Failed to send OTP email to %s", email) + raise Exception("EmailSender returned False") except Exception: logger.exception("Unexpected error in email worker") + raise async def run_worker() -> None: @@ -37,7 +39,7 @@ async def wrapped_handler(msg: bytes | str) -> None: await handle_message(msg) # Subscribe to the email.send_otp subject - await NatsClient.subscribe("email.send_otp", wrapped_handler) + await NatsClient.js_subscribe(NatsSubjects.EMAIL_SEND_OTP, wrapped_handler) # Keep the worker running await asyncio.Event().wait() diff --git a/app/worker/notification/main.py b/app/worker/notification/main.py index a79d7c20..7ea07cbb 100644 --- a/app/worker/notification/main.py +++ b/app/worker/notification/main.py @@ -114,7 +114,11 @@ async def wrapped_handler(msg: bytes | str) -> None: await handle_message(msg, queue, invalid_tokens, invalid_devices) for subject in queue.priority_subjects(): - await NatsClient.subscribe(subject, wrapped_handler) + await NatsClient.js_subscribe( + subject, + wrapped_handler, + stream_name="notification_delivery_stream" + ) await asyncio.Event().wait() diff --git a/app/worker/notification/notification_queue.py b/app/worker/notification/notification_queue.py index 463de938..c430d441 100644 --- a/app/worker/notification/notification_queue.py +++ b/app/worker/notification/notification_queue.py @@ -24,7 +24,7 @@ async def enqueue_notification( entry = NotificationQueueEntry(notification=notification, attempts=attempts) subject = self._settings.subject_for(entry.notification.priority) payload = entry.model_dump_json().encode("utf-8") - await NatsClient.publish(subject, payload) + await NatsClient.js_publish(subject, payload, stream_name="notification_delivery_stream") @staticmethod def priority_index(priority: NotificationPriority) -> int: diff --git a/app/worker/photo_worker/main.py b/app/worker/photo_worker/main.py index 7fa988d2..aba77037 100644 --- a/app/worker/photo_worker/main.py +++ b/app/worker/photo_worker/main.py @@ -65,21 +65,11 @@ async def handle_message(self, data: bytes) -> None: # transaction rolls back cleanly instead of being committed half-applied. job = await self._create_job(event) - try: - payload = await self._load_image(event.image_ref) - except Exception as exc: - logger.warning("Failed to load image for photo %s: %s", event.photo_id, exc) - await self._update_job(job, "failed") - return + payload = await self._load_image(event.image_ref) await self._update_job(job, "running") - try: - faces = await self._face_service.detect_faces(payload) - except Exception as exc: - logger.warning("Face detection failed for photo %s: %s", event.photo_id, exc) - await self._update_job(job, "failed") - return + faces = await self._face_service.detect_faces(payload) if not faces: logger.info("No faces detected in photo %s, marking as public", event.photo_id) @@ -201,7 +191,7 @@ async def _publish_audit(event: PhotoProcessEvent, faces_count: int) -> None: metadata={"photo_id": str(event.photo_id), "faces_count": faces_count}, ) try: - await NatsClient.publish(NatsSubjects.AUDIT_EVENT, msg.model_dump_json().encode("utf-8")) + await NatsClient.js_publish(NatsSubjects.AUDIT_EVENT, msg.model_dump_json().encode("utf-8")) except Exception as exc: logger.warning("Failed to publish audit for photo %s: %s", event.photo_id, exc) @@ -272,28 +262,25 @@ async def handle(data: bytes) -> None: # pool_pre_ping the connection is also revalidated on checkout, so a # Postgres restart is recovered automatically. Errors are logged here so a # single bad message does not tear down the subscription. - try: - async with engine.begin() as conn: - container = Container(conn) - single_face_service = SingleFaceMatchService( - conn=conn, - photo_face_querier=container.photo_face_querier, - photo_querier=container.photo_querier, - user_match_service=container.auth_service, - user_notification_service=container.user_notifications_service, - ) - worker = PhotoWorker( - conn=conn, - face_embedding_service=container.face_embedding_service, - single_face_service=single_face_service, - user_notification_service=container.user_notifications_service, - photo_face_querier=container.photo_face_querier, - photo_querier=container.photo_querier, - processing_job_querier=container.processing_job_querier, - ) - await worker.handle_message(data) - except Exception: - logger.exception("Failed to process photo message") + async with engine.begin() as conn: + container = Container(conn) + single_face_service = SingleFaceMatchService( + conn=conn, + photo_face_querier=container.photo_face_querier, + photo_querier=container.photo_querier, + user_match_service=container.auth_service, + user_notification_service=container.user_notifications_service, + ) + worker = PhotoWorker( + conn=conn, + face_embedding_service=container.face_embedding_service, + single_face_service=single_face_service, + user_notification_service=container.user_notifications_service, + photo_face_querier=container.photo_face_querier, + photo_querier=container.photo_querier, + processing_job_querier=container.processing_job_querier, + ) + await worker.handle_message(data) await NatsClient.js_subscribe( subject=NatsSubjects.PHOTO_PROCESS, diff --git a/app/worker/photo_worker/tests/__init__.py b/app/worker/photo_worker/tests/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/app/worker/photo_worker/tests/test_photo_worker.py b/app/worker/photo_worker/tests/test_photo_worker.py deleted file mode 100644 index 1eeb73e3..00000000 --- a/app/worker/photo_worker/tests/test_photo_worker.py +++ /dev/null @@ -1,384 +0,0 @@ -import json -import sys -import uuid -from unittest.mock import AsyncMock, MagicMock, patch, create_autospec - -import pytest - -_MOCKED_MODULES = ( - "db.generated.user", - "app.worker.notification.settings", - "app.worker.notification.notification_queue", -) -_original_modules = {name: sys.modules.get(name) for name in _MOCKED_MODULES} -for _name in _MOCKED_MODULES: - sys.modules[_name] = MagicMock() - -from app.service.face_embedding import DetectedFace, FaceEmbeddingService, FaceImagePayload # noqa: E402 -from app.service.face_match import SingleFaceMatchService # noqa: E402 -from app.service.user_notification import UserNotificationService # noqa: E402 -from app.worker.photo_worker.main import PhotoWorker # noqa: E402 -from app.worker.photo_worker.schema.event import PhotoProcessEvent # noqa: E402 -from db.generated import photo_faces as photo_face_queries # noqa: E402 - -# Restore sys.modules so other test files importing these modules -# (e.g. db.generated.user) get the real thing, not this leaked mock. -for _name, _original in _original_modules.items(): - if _original is None: - sys.modules.pop(_name, None) - else: - sys.modules[_name] = _original - -# ── fixtures ────────────────────────────────────────────────────────── - - -@pytest.fixture -def conn() -> MagicMock: - return MagicMock() - - -@pytest.fixture -def face_service() -> AsyncMock: - return create_autospec(FaceEmbeddingService, instance=True) - - -@pytest.fixture -def single_face_service() -> AsyncMock: - svc = create_autospec(SingleFaceMatchService, instance=True) - svc.process_detected_face = AsyncMock() - return svc - - -@pytest.fixture -def notification_service() -> AsyncMock: - svc = MagicMock(spec=UserNotificationService) - svc.create_notification = AsyncMock() - return svc - - -@pytest.fixture -def photo_face_querier() -> AsyncMock: - return create_autospec(photo_face_queries.AsyncQuerier, instance=True) - - -@pytest.fixture -def photo_querier() -> AsyncMock: - from db.generated import photos as photo_queries_mod - q = create_autospec(photo_queries_mod.AsyncQuerier, instance=True) - q.update_photo_status = AsyncMock(return_value=None) - q.update_photo_visibility = AsyncMock(return_value=None) - return q - - -@pytest.fixture -def worker( - conn: MagicMock, - face_service: AsyncMock, - single_face_service: AsyncMock, - notification_service: AsyncMock, - photo_face_querier: AsyncMock, - photo_querier: AsyncMock, -) -> PhotoWorker: - return PhotoWorker( - conn=conn, - face_embedding_service=face_service, - single_face_service=single_face_service, - user_notification_service=notification_service, - photo_face_querier=photo_face_querier, - photo_querier=photo_querier, - ) - - -@pytest.fixture -def event() -> PhotoProcessEvent: - return PhotoProcessEvent( - photo_id=uuid.uuid4(), - image_ref="photos/test.jpg", - event_id=uuid.uuid4(), - ) - - -def _make_face(embedding: list[float] | None = None) -> DetectedFace: - return DetectedFace( - embedding=embedding or [0.1] * 512, - bbox=(10.0, 20.0, 100.0, 200.0), - ) - - -def _event_bytes(event: PhotoProcessEvent) -> bytes: - return event.model_dump_json().encode() - - -# ── tests ───────────────────────────────────────────────────────────── - - -@pytest.mark.asyncio -async def test_invalid_payload_is_skipped(worker: PhotoWorker) -> None: - """Malformed JSON should be silently skipped.""" - await worker.handle_message(b"not json") - # No exceptions raised - - -@pytest.mark.asyncio -async def test_no_faces_skips_processing( - worker: PhotoWorker, - face_service: AsyncMock, - single_face_service: AsyncMock, - photo_face_querier: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """If no faces detected, neither single nor group path runs.""" - face_service.detect_faces = AsyncMock(return_value=[]) - - with patch.object(worker, "_load_image", new_callable=AsyncMock) as mock_load: - mock_load.return_value = FaceImagePayload( - filename="test.jpg", content_type="image/jpeg", bytes=b"img" - ) - await worker.handle_message(_event_bytes(event)) - - single_face_service.process_detected_face.assert_not_called() - photo_face_querier.insert_photo_face_with_approval.assert_not_called() - - -@pytest.mark.asyncio -async def test_single_face_takes_single_path( - worker: PhotoWorker, - face_service: AsyncMock, - single_face_service: AsyncMock, - photo_face_querier: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """Exactly 1 face -> single face match path.""" - face_service.detect_faces = AsyncMock(return_value=[_make_face()]) - - with patch.object(worker, "_load_image", new_callable=AsyncMock) as mock_load: - mock_load.return_value = FaceImagePayload( - filename="test.jpg", content_type="image/jpeg", bytes=b"img" - ) - await worker.handle_message(_event_bytes(event)) - - single_face_service.process_detected_face.assert_called_once() - photo_face_querier.insert_photo_face_with_approval.assert_not_called() - - -@pytest.mark.asyncio -async def test_multiple_faces_takes_group_path( - worker: PhotoWorker, - face_service: AsyncMock, - single_face_service: AsyncMock, - photo_face_querier: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """Multiple faces -> group photo approval path.""" - faces = [_make_face([0.1] * 512), _make_face([0.2] * 512), _make_face([0.3] * 512)] - face_service.detect_faces = AsyncMock(return_value=faces) - photo_face_querier.insert_photo_face_with_approval = AsyncMock(return_value=None) - - with patch.object(worker, "_load_image", new_callable=AsyncMock) as mock_load: - mock_load.return_value = FaceImagePayload( - filename="test.jpg", content_type="image/jpeg", bytes=b"img" - ) - await worker.handle_message(_event_bytes(event)) - - single_face_service.process_detected_face.assert_not_called() - assert photo_face_querier.insert_photo_face_with_approval.call_count == 3 - - -@pytest.mark.asyncio -async def test_group_photo_sends_notification_on_match( - worker: PhotoWorker, - face_service: AsyncMock, - notification_service: AsyncMock, - photo_face_querier: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """When a group face matches a user, a notification is sent.""" - faces = [_make_face(), _make_face()] - face_service.detect_faces = AsyncMock(return_value=faces) - - matched_approval = MagicMock() - matched_approval.user_id = uuid.uuid4() - matched_approval.photo_id = event.photo_id - - # First face matches, second doesn't - photo_face_querier.insert_photo_face_with_approval = AsyncMock( - side_effect=[matched_approval, None] - ) - - with patch.object(worker, "_load_image", new_callable=AsyncMock) as mock_load: - mock_load.return_value = FaceImagePayload( - filename="test.jpg", content_type="image/jpeg", bytes=b"img" - ) - await worker.handle_message(_event_bytes(event)) - - notification_service.create_notification.assert_called_once() - call_kwargs = notification_service.create_notification.call_args.kwargs - assert call_kwargs["type"] == "photo_approval" - assert call_kwargs["user_id"] == matched_approval.user_id - - -@pytest.mark.asyncio -async def test_group_photo_stores_bbox( - worker: PhotoWorker, - face_service: AsyncMock, - photo_face_querier: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """Group path should pass bbox JSON to the DB query.""" - face = _make_face() - face_service.detect_faces = AsyncMock(return_value=[face, face]) - photo_face_querier.insert_photo_face_with_approval = AsyncMock(return_value=None) - - with patch.object(worker, "_load_image", new_callable=AsyncMock) as mock_load: - mock_load.return_value = FaceImagePayload( - filename="test.jpg", content_type="image/jpeg", bytes=b"img" - ) - await worker.handle_message(_event_bytes(event)) - - call_args = photo_face_querier.insert_photo_face_with_approval.call_args_list[0] - params = call_args.args[0] - bbox = json.loads(params.bbox) - assert bbox == {"x1": 10.0, "y1": 20.0, "x2": 100.0, "y2": 200.0} - - -@pytest.mark.asyncio -async def test_single_face_passes_correct_bbox( - worker: PhotoWorker, - face_service: AsyncMock, - single_face_service: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """Single path should construct BBoxPayload from detected face.""" - face = _make_face() - face_service.detect_faces = AsyncMock(return_value=[face]) - - with patch.object(worker, "_load_image", new_callable=AsyncMock) as mock_load: - mock_load.return_value = FaceImagePayload( - filename="test.jpg", content_type="image/jpeg", bytes=b"img" - ) - await worker.handle_message(_event_bytes(event)) - - call_args = single_face_service.process_detected_face.call_args - bbox = call_args.args[2] - assert bbox.x1 == 10.0 - assert bbox.y1 == 20.0 - assert bbox.x2 == 100.0 - assert bbox.y2 == 200.0 - - -@pytest.mark.asyncio -async def test_image_load_failure_is_handled( - worker: PhotoWorker, - face_service: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """If image fetch fails, worker should not crash.""" - with patch.object(worker, "_load_image", new_callable=AsyncMock) as mock_load: - mock_load.side_effect = RuntimeError("MinIO down") - await worker.handle_message(_event_bytes(event)) - - face_service.detect_faces.assert_not_called() - - -# ── cleanup scheduling tests ────────────────────────────────────────── -# -# photo_worker no longer schedules immediate MinIO cleanup after processing. -# Cleanup is now threshold-based (event_lifecycle worker, 20 days after -# event.end_date, gated on Drive sync for direct-uploaded photos) — see -# ListPhotosDueForStorageCleanup. Immediate cleanup right after processing -# was a race: for direct uploads, MinIO is the only copy of the photo until -# the async Drive sync completes, so deleting it immediately could destroy -# the only copy before it was ever backed up. - - -@pytest.mark.asyncio -async def test_single_face_does_not_schedule_cleanup( - worker: PhotoWorker, - face_service: AsyncMock, - single_face_service: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """After single face processing, only the audit event is published — no cleanup.""" - face_service.detect_faces = AsyncMock(return_value=[_make_face()]) - - with ( - patch.object(worker, "_load_image", new_callable=AsyncMock) as mock_load, - patch("app.worker.photo_worker.main.NatsClient") as mock_nats, - ): - mock_nats.publish = AsyncMock() - mock_load.return_value = FaceImagePayload( - filename="test.jpg", content_type="image/jpeg", bytes=b"img" - ) - await worker.handle_message(_event_bytes(event)) - - from app.infra.nats import NatsSubjects - assert mock_nats.publish.call_count == 1 - audit_call = mock_nats.publish.call_args_list[0] - assert audit_call.args[0] == NatsSubjects.AUDIT_EVENT - - -@pytest.mark.asyncio -async def test_group_photo_does_not_schedule_cleanup( - worker: PhotoWorker, - face_service: AsyncMock, - photo_face_querier: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """After group photo processing, only the audit event is published — no cleanup.""" - faces = [_make_face(), _make_face()] - face_service.detect_faces = AsyncMock(return_value=faces) - photo_face_querier.insert_photo_face_with_approval = AsyncMock(return_value=None) - - with ( - patch.object(worker, "_load_image", new_callable=AsyncMock) as mock_load, - patch("app.worker.photo_worker.main.NatsClient") as mock_nats, - ): - mock_nats.publish = AsyncMock() - mock_load.return_value = FaceImagePayload( - filename="test.jpg", content_type="image/jpeg", bytes=b"img" - ) - await worker.handle_message(_event_bytes(event)) - - from app.infra.nats import NatsSubjects - assert mock_nats.publish.call_count == 1 - audit_call = mock_nats.publish.call_args_list[0] - assert audit_call.args[0] == NatsSubjects.AUDIT_EVENT - - -@pytest.mark.asyncio -async def test_no_cleanup_when_no_faces( - worker: PhotoWorker, - face_service: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """No faces detected: photo is marked public, but no cleanup is scheduled.""" - face_service.detect_faces = AsyncMock(return_value=[]) - - with ( - patch.object(worker, "_load_image", new_callable=AsyncMock) as mock_load, - patch("app.worker.photo_worker.main.NatsClient") as mock_nats, - ): - mock_nats.publish = AsyncMock() - mock_load.return_value = FaceImagePayload( - filename="test.jpg", content_type="image/jpeg", bytes=b"img" - ) - await worker.handle_message(_event_bytes(event)) - - mock_nats.publish.assert_not_called() - - -@pytest.mark.asyncio -async def test_no_cleanup_on_image_load_failure( - worker: PhotoWorker, - event: PhotoProcessEvent, -) -> None: - """If image fails to load, no cleanup should be scheduled.""" - with ( - patch.object(worker, "_load_image", new_callable=AsyncMock) as mock_load, - patch("app.worker.photo_worker.main.NatsClient") as mock_nats, - ): - mock_nats.publish = AsyncMock() - mock_load.side_effect = RuntimeError("MinIO down") - await worker.handle_message(_event_bytes(event)) - - mock_nats.publish.assert_not_called() diff --git a/tests/e2e/test_photo_ai_edge_cases.py b/tests/e2e/test_photo_ai_edge_cases.py index ed3e0766..7b36da22 100644 --- a/tests/e2e/test_photo_ai_edge_cases.py +++ b/tests/e2e/test_photo_ai_edge_cases.py @@ -31,7 +31,7 @@ async def test_photo_ai_pipeline_detects_0_faces() -> None: ) payload = {"photo_id": str(photo_id), "image_ref": storage_key, "event_id": str(event_id)} - await NatsClient.publish(NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8")) + await NatsClient.js_publish(NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8")) try: final_status = await _wait_for_job(photo_id) @@ -75,7 +75,7 @@ async def test_photo_ai_pipeline_detects_multiple_faces() -> None: ) payload = {"photo_id": str(photo_id), "image_ref": storage_key, "event_id": str(event_id)} - await NatsClient.publish(NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8")) + await NatsClient.js_publish(NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8")) try: final_status = await _wait_for_job(photo_id) @@ -154,7 +154,7 @@ async def test_photo_ai_pipeline_matched_user() -> None: ) payload = {"photo_id": str(photo_id), "image_ref": storage_key, "event_id": str(event_id)} - await NatsClient.publish(NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8")) + await NatsClient.js_publish(NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8")) try: final_status = await _wait_for_job(photo_id) diff --git a/tests/e2e/test_photo_ai_load.py b/tests/e2e/test_photo_ai_load.py index 0f9437bc..37aa1faa 100644 --- a/tests/e2e/test_photo_ai_load.py +++ b/tests/e2e/test_photo_ai_load.py @@ -153,7 +153,7 @@ async def test_photo_ai_load_20_photos(setup_infra: None) -> None: # noqa: ARG0 # 2. Publish 20 NATS messages concurrently await asyncio.gather( *[ - NatsClient.publish( + NatsClient.js_publish( NatsSubjects.PHOTO_PROCESS.value, json.dumps( { diff --git a/tests/e2e/test_photo_ai_pipeline_e2e.py b/tests/e2e/test_photo_ai_pipeline_e2e.py index 0a7b3b5a..1a16292b 100644 --- a/tests/e2e/test_photo_ai_pipeline_e2e.py +++ b/tests/e2e/test_photo_ai_pipeline_e2e.py @@ -37,7 +37,7 @@ async def test_photo_ai_pipeline_detects_single_face() -> None: "image_ref": storage_key, "event_id": str(event_id), } - await NatsClient.publish( + await NatsClient.js_publish( NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8") ) @@ -89,7 +89,7 @@ async def test_photo_ai_pipeline_corrupt_image() -> None: "image_ref": storage_key, "event_id": str(event_id), } - await NatsClient.publish( + await NatsClient.js_publish( NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8") ) diff --git a/tests/integration/test_photo_approval_flow.py b/tests/integration/test_photo_approval_flow.py index 2cfb9f4a..7ffcff5d 100644 --- a/tests/integration/test_photo_approval_flow.py +++ b/tests/integration/test_photo_approval_flow.py @@ -105,6 +105,7 @@ async def test_group_photo_approval_lifecycle( name="Approval Test Event", event_code=f"APP{str(event_id)[:4]}", event_date=datetime.datetime.now(datetime.timezone.utc), + end_date=datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(days=1), status="scheduled", created_by=event_creator_id ) @@ -115,6 +116,7 @@ async def test_group_photo_approval_lifecycle( photo_queries.CreatePhotoParams( event_id=event_id, storage_key="test/group.jpg", + source="direct", taken_at=None, day_number=None, visibility="public" @@ -185,6 +187,7 @@ async def test_group_photo_rejection_deletes_storage( name="Reject Test Event", event_code=f"REJ{str(event_id)[:4]}", event_date=datetime.datetime.now(datetime.timezone.utc), + end_date=datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(days=1), status="scheduled", created_by=event_creator_id ) @@ -195,6 +198,7 @@ async def test_group_photo_rejection_deletes_storage( photo_queries.CreatePhotoParams( event_id=event_id, storage_key="test/reject.jpg", + source="direct", taken_at=None, day_number=None, visibility="public" diff --git a/tests/unit/test_auth_email_otp.py b/tests/unit/test_auth_email_otp.py index 7e6acd95..d21a9e64 100644 --- a/tests/unit/test_auth_email_otp.py +++ b/tests/unit/test_auth_email_otp.py @@ -47,7 +47,8 @@ def auth_service( ) @pytest.mark.asyncio -@patch("app.service.users.NatsClient.publish") +@patch("app.service.users.settings.environment", "production") +@patch("app.service.users.NatsClient.js_publish") async def test_mobile_register_sends_otp( mock_publish: AsyncMock, auth_service: AuthService, @@ -140,7 +141,8 @@ async def test_mobile_register_resend_otp_success( mock_redis.get.return_value = '{"hashed_password": "fake"}' mock_redis.incr.return_value = 1 - with patch("app.service.users.NatsClient.publish", new_callable=AsyncMock) as mock_publish: + with patch("app.service.users.settings.environment", "production"), \ + patch("app.service.users.NatsClient.js_publish", new_callable=AsyncMock) as mock_publish: res = await auth_service.mobile_register_resend_otp(redis=mock_redis, email=email) assert res.status == "pending_verification" diff --git a/tests/unit/test_direct_uploads.py b/tests/unit/test_direct_uploads.py index 68b604fe..20612b14 100644 --- a/tests/unit/test_direct_uploads.py +++ b/tests/unit/test_direct_uploads.py @@ -373,7 +373,7 @@ async def _photos_iter(upload_request_id): request_id, event_id, mock_staff_user.id, None, ) - with patch("app.service.upload_requests.NatsClient.publish") as mock_publish: + with patch("app.service.upload_requests.NatsClient.js_publish") as mock_publish: await upload_requests_service.approve_request( request_id=request_id, approved_by=mock_staff_user, ) diff --git a/tests/unit/test_mobile_auth_email_logging.py b/tests/unit/test_mobile_auth_email_logging.py index d868f324..bbd3646a 100644 --- a/tests/unit/test_mobile_auth_email_logging.py +++ b/tests/unit/test_mobile_auth_email_logging.py @@ -8,7 +8,7 @@ import logging import uuid from datetime import datetime, timezone -from unittest.mock import MagicMock +from unittest.mock import MagicMock, AsyncMock import pytest @@ -148,6 +148,8 @@ def test_mobile_register_logs_without_plaintext_email( async def _noop_cache_session_for_auth(**_: object) -> None: return None + monkeypatch.setattr("app.service.users.settings.environment", "production") + monkeypatch.setattr("app.service.users.NatsClient.js_publish", AsyncMock()) monkeypatch.setattr(SessionService, "cache_session_for_auth", _noop_cache_session_for_auth) monkeypatch.setattr(users_module, "create_acces_mobile_token", lambda _: "access") monkeypatch.setattr(users_module, "create_raw_refresh_token", lambda: "refresh") diff --git a/tests/unit/test_photo_worker.py b/tests/unit/test_photo_worker.py index 0091223d..3aa9813d 100644 --- a/tests/unit/test_photo_worker.py +++ b/tests/unit/test_photo_worker.py @@ -91,7 +91,7 @@ async def test_handle_message_success_no_faces(photo_worker, sample_event, mock_ photo_worker._load_image = AsyncMock(return_value=FaceImagePayload(filename="test.jpg", content_type="image/jpeg", bytes=b"data")) mock_face_service.detect_faces.return_value = [] - with patch("app.worker.photo_worker.main.NatsClient.publish") as mock_publish: + with patch("app.worker.photo_worker.main.NatsClient.js_publish") as mock_publish: await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) mock_pj_querier.create_processing_job.assert_called_once() @@ -110,7 +110,7 @@ async def test_handle_message_success_single_face(photo_worker, sample_event, mo face = DetectedFace(bbox=(0, 0, 100, 100), embedding=[0.1] * 512) mock_face_service.detect_faces.return_value = [face] - with patch("app.worker.photo_worker.main.NatsClient.publish") as mock_publish: + with patch("app.worker.photo_worker.main.NatsClient.js_publish") as mock_publish: await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) mock_pj_querier.update_processing_job_status.assert_any_call(id=mock_pj_querier.create_processing_job.return_value.id, status="completed") @@ -126,7 +126,7 @@ async def test_handle_message_success_group_face(photo_worker, sample_event, moc face2 = DetectedFace(bbox=(100, 100, 200, 200), embedding=[0.2] * 512) mock_face_service.detect_faces.return_value = [face1, face2] - with patch("app.worker.photo_worker.main.NatsClient.publish"): + with patch("app.worker.photo_worker.main.NatsClient.js_publish"): await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) mock_pj_querier.update_processing_job_status.assert_any_call(id=mock_pj_querier.create_processing_job.return_value.id, status="completed") @@ -137,16 +137,18 @@ async def test_handle_message_success_group_face(photo_worker, sample_event, moc @pytest.mark.asyncio async def test_handle_message_fails_on_minio_load(photo_worker, sample_event, mock_pj_querier): photo_worker._load_image = AsyncMock(side_effect=Exception("MinIO error")) - await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) - mock_pj_querier.update_processing_job_status.assert_called_with(id=mock_pj_querier.create_processing_job.return_value.id, status="failed") + with pytest.raises(Exception, match="MinIO error"): + await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) + assert "failed" not in [call.kwargs.get("status") for call in mock_pj_querier.update_processing_job_status.call_args_list] @pytest.mark.asyncio async def test_handle_message_fails_on_ai_detection(photo_worker, sample_event, mock_face_service, mock_pj_querier): photo_worker._load_image = AsyncMock(return_value=FaceImagePayload(filename="test.jpg", content_type="image/jpeg", bytes=b"data")) mock_face_service.detect_faces.side_effect = Exception("InsightFace out of memory") - await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) - mock_pj_querier.update_processing_job_status.assert_called_with(id=mock_pj_querier.create_processing_job.return_value.id, status="failed") + with pytest.raises(Exception, match="InsightFace out of memory"): + await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) + assert "failed" not in [call.kwargs.get("status") for call in mock_pj_querier.update_processing_job_status.call_args_list] @pytest.mark.asyncio diff --git a/tests/unit/test_upload_requests.py b/tests/unit/test_upload_requests.py index dd31c391..77877680 100644 --- a/tests/unit/test_upload_requests.py +++ b/tests/unit/test_upload_requests.py @@ -160,7 +160,7 @@ async def test_create_request_success( ) with patch("app.service.upload_requests.GoogleDriveClient.download_file", return_value=mock_download) as mock_drive, \ - patch("app.service.upload_requests.NatsClient.publish") as mock_publish: + patch("app.service.upload_requests.NatsClient.js_publish") as mock_publish: details = await upload_requests_service.create_request( event_id=event_id, @@ -231,7 +231,7 @@ async def test_create_group_from_folder( id=group_id, event_id=event_id, folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="pending", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None, source="drive" ) - with patch("app.service.upload_requests.NatsClient.publish") as mock_publish: + with patch("app.service.upload_requests.NatsClient.js_publish") as mock_publish: details = await upload_requests_service.create_group_from_folder( event_id=event_id, folder_id="folder_123", visibility="public", day_number=None, requested_by=mock_staff_user ) From 93b3c1082307946ee1e53f1b4a4c5d16c5c5a0b1 Mon Sep 17 00:00:00 2001 From: Adem Boukabes <142881379+ademboukabes@users.noreply.github.com> Date: Sun, 30 Aug 2026 14:54:40 +0100 Subject: [PATCH 25/29] fix(lint): resolve ruff linting and mypy typing errors --- app/container.py | 18 +- app/core/config.py | 7 +- app/core/constant.py | 7 +- app/core/exceptions.py | 25 +- app/core/securite.py | 49 +- app/core/utils.py | 3 - app/deps/cookie_auth.py | 11 +- app/deps/rate_limit.py | 5 +- app/deps/token_auth.py | 12 +- app/infra/email.py | 14 +- app/infra/google_drive.py | 59 +- app/infra/minio.py | 15 +- app/infra/nats.py | 43 +- app/infra/redis.py | 6 +- app/main.py | 20 +- app/router/mobile/auth.py | 69 +- app/router/mobile/enrollement.py | 4 +- app/router/mobile/event.py | 9 +- app/router/mobile/notifications.py | 8 +- app/router/staff/drive.py | 15 +- app/router/staff/uploads.py | 12 +- app/router/staff/uploads_direct.py | 35 +- app/router/web/auth.py | 19 +- app/router/web/event.py | 25 +- app/router/web/staff_users.py | 14 +- app/router/web/stats.py | 15 +- app/router/web/users.py | 2 + app/schema/internal/notification.py | 1 + app/schema/request/mobile/auth.py | 10 +- app/schema/request/mobile/notifications.py | 4 +- app/schema/request/web/auth.py | 1 - app/schema/request/web/event.py | 2 + app/schema/request/web/staff_user.py | 1 - app/schema/response/mobile/audit.py | 1 + app/schema/response/mobile/auth.py | 7 + app/schema/response/mobile/notifications.py | 4 +- app/schema/response/staff/notifications.py | 4 +- app/schema/response/staff/upload_groups.py | 21 +- app/schema/response/web/audit.py | 1 + app/schema/response/web/auth.py | 3 - app/schema/response/web/event.py | 7 + app/schema/response/web/stats.py | 5 + app/service/device.py | 46 +- app/service/event.py | 41 +- app/service/face_embedding.py | 28 +- app/service/face_match.py | 35 +- app/service/photo_approval.py | 12 +- app/service/staff_drive.py | 53 +- app/service/staff_notifications.py | 4 +- app/service/staff_user.py | 24 +- app/service/staged_upload_storage.py | 8 +- app/service/stats.py | 42 +- app/service/upload_requests.py | 385 ++++++---- app/service/user_notification.py | 10 +- app/service/user_photo.py | 12 +- app/service/users.py | 104 ++- app/worker/audit/__init__.py | 1 + app/worker/audit/main.py | 6 +- app/worker/audit/settings.py | 2 +- app/worker/drive_sync/main.py | 23 +- app/worker/event_lifecycle/main.py | 12 +- app/worker/notification/firebase.py | 19 +- app/worker/notification/invalid_tokens.py | 12 +- app/worker/notification/main.py | 15 +- app/worker/notification/notification_queue.py | 14 +- app/worker/photo_worker/main.py | 102 ++- app/worker/storage_cleaner/main.py | 1 + app/worker/upload_reconciler/main.py | 8 +- coverage_report.txt | 96 +++ db/__init__.py | 1 - migrations/env.py | 4 +- migrations/helper.py | 5 +- mobile-quickstart/seed.py | 136 +++- scripts/check_ai_results.py | 14 +- scripts/check_scopes.py | 11 +- scripts/generate_drive_url.py | 23 +- scripts/list_drive.py | 26 +- scripts/seed.py | 656 ++++++++++++++++++ scripts/seed_admin.py | 10 +- scripts/trigger_import.py | 22 +- scripts/trigger_photo_worker.py | 51 +- scripts/trigger_upload_request.py | 60 +- tests/e2e/conftest.py | 26 +- tests/e2e/test_mobile_auth_intent_e2e.py | 3 + tests/e2e/test_photo_ai_edge_cases.py | 37 +- tests/e2e/test_photo_ai_load.py | 9 +- tests/e2e/test_photo_ai_pipeline_e2e.py | 11 +- tests/e2e/test_stats_endpoint.py | 25 +- tests/integration/test_enrollment_flow.py | 11 +- tests/integration/test_photo_approval_flow.py | 98 ++- .../test_session_device_management.py | 28 +- tests/security/test_auth_security.py | 76 +- tests/unit/test_auth_email_otp.py | 37 +- tests/unit/test_auth_service.py | 121 +++- tests/unit/test_direct_uploads.py | 277 ++++++-- tests/unit/test_enroll_security.py | 15 +- tests/unit/test_face_match_service.py | 37 +- tests/unit/test_minio.py | 19 +- tests/unit/test_mobile_auth_email_logging.py | 5 +- .../test_mobile_auth_intent_validation.py | 58 +- tests/unit/test_mobile_auth_rate_limiting.py | 5 + .../test_mobile_auth_request_validation.py | 3 - tests/unit/test_photo_approval_lifecycle.py | 18 +- tests/unit/test_photo_approval_service.py | 74 +- tests/unit/test_photo_worker.py | 115 ++- tests/unit/test_upload_requests.py | 246 +++++-- 106 files changed, 3111 insertions(+), 970 deletions(-) create mode 100644 coverage_report.txt create mode 100644 scripts/seed.py diff --git a/app/container.py b/app/container.py index c845258c..0024ea6c 100644 --- a/app/container.py +++ b/app/container.py @@ -43,6 +43,7 @@ from app.worker.notification.notification_queue import NotificationQueue from app.worker.notification.settings import NotifSetting + class Container: def __init__( self, @@ -51,7 +52,9 @@ def __init__( ): # infrastructure self.redis = RedisClient.get_instance() - self.face_embedding_service = face_embedding_service or get_face_embedding_service() + self.face_embedding_service = ( + face_embedding_service or get_face_embedding_service() + ) # queriers self.user_querier = user_queries.AsyncQuerier(conn) @@ -59,9 +62,13 @@ def __init__( self.device_querier = device_queries.AsyncQuerier(conn) self.staff_user_querier = staff_user_queries.AsyncQuerier(conn) self.staff_drive_querier = staff_drive_queries.AsyncQuerier(conn) - self.upload_request_group_querier = upload_request_group_queries.AsyncQuerier(conn) + self.upload_request_group_querier = upload_request_group_queries.AsyncQuerier( + conn + ) self.upload_request_querier = upload_request_queries.AsyncQuerier(conn) - self.upload_request_photo_querier = upload_request_photo_queries.AsyncQuerier(conn) + self.upload_request_photo_querier = upload_request_photo_queries.AsyncQuerier( + conn + ) self.photo_querier = photo_queries.AsyncQuerier(conn) self.photo_approval_querier = photo_approval_queries.AsyncQuerier(conn) self.photo_face_querier = photo_face_queries.AsyncQuerier(conn) @@ -79,7 +86,6 @@ def __init__( redis=self.redis, ) - self.device_service = DeviceService() self.device_service.init( device_querier=self.device_querier, @@ -130,7 +136,8 @@ def __init__( self.staff_user_service = StaffUserService() self.staff_user_service.init( - staff_user_querier=self.staff_user_querier,) + staff_user_querier=self.staff_user_querier, + ) self.event_service = EventService( e_querier=self.event_querier, @@ -155,6 +162,7 @@ def __init__( querier=self.stats_querier, ) + async def get_container( conn: sqlalchemy.ext.asyncio.AsyncConnection = Depends(get_db), ) -> Container: diff --git a/app/core/config.py b/app/core/config.py index ed8613f5..1d84e51f 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -6,7 +6,12 @@ class Settings(BaseSettings): app_name: str = "multAI" environment: str = "dev" debug: bool = True - CORS_ORIGINS: list[str] = ["http://localhost:3000", "http://localhost:5173", "http://127.0.0.1:5173", "http://127.0.0.1:3000"] + CORS_ORIGINS: list[str] = [ + "http://localhost:3000", + "http://localhost:5173", + "http://127.0.0.1:5173", + "http://127.0.0.1:3000", + ] # Redis REDIS_PORT: int diff --git a/app/core/constant.py b/app/core/constant.py index b0f117a4..9129bdb4 100644 --- a/app/core/constant.py +++ b/app/core/constant.py @@ -29,12 +29,7 @@ class AuditEventType(str, Enum): PHOTO_APPROVAL_DECIDED = "photo_approval.decided" -IMAGE_ALLOWED_TYPES = { - "image/jpeg", - "image/png", - "image/heic", - "image/heif" -} +IMAGE_ALLOWED_TYPES = {"image/jpeg", "image/png", "image/heic", "image/heif"} DEFAULT_CONTENT_TYPE = "application/octet-stream" DRIVE_ALLOWED_HOSTS = {"drive.google.com", "docs.google.com"} diff --git a/app/core/exceptions.py b/app/core/exceptions.py index 15465685..791cede2 100644 --- a/app/core/exceptions.py +++ b/app/core/exceptions.py @@ -23,8 +23,8 @@ def bad_request(detail: str = "Bad request") -> HTTPException: return HTTPException(status_code=400, detail=detail) @staticmethod - def payement_required(detail:str = "payement required")->HTTPException: - return HTTPException(status_code=402,detail=detail) + def payement_required(detail: str = "payement required") -> HTTPException: + return HTTPException(status_code=402, detail=detail) @staticmethod def internal_error(detail: str = "Internal server error") -> HTTPException: @@ -51,13 +51,16 @@ def queue_error(detail: str = "Queue operation failed") -> HTTPException: return HTTPException(status_code=500, detail=detail) @staticmethod - def image_quality_error(detail: str = "Image does not meet quality requirements") -> HTTPException: + def image_quality_error( + detail: str = "Image does not meet quality requirements", + ) -> HTTPException: return HTTPException(status_code=400, detail=detail) @staticmethod def image_format_error(detail: str = "Unsupported image format") -> HTTPException: return HTTPException(status_code=400, detail=detail) + class DBException(ABC): """Abstract class to enforce DB error handling.""" @@ -115,13 +118,15 @@ def handle_unique_violation(exc: Exception) -> HTTPException: err_msg = str(exc).lower() if constraint == "staff_users_email_key" or "staff_users_email_key" in err_msg: return HTTPException( - status_code=409, - detail="Staff user with this email already exists" + status_code=409, detail="Staff user with this email already exists" ) - if constraint in ("users_email_key", "idx_users_email") or "idx_users_email" in err_msg or "users_email_key" in err_msg: + if ( + constraint in ("users_email_key", "idx_users_email") + or "idx_users_email" in err_msg + or "users_email_key" in err_msg + ): return HTTPException( - status_code=409, - detail="Email already in use; please login instead" + status_code=409, detail="Email already in use; please login instead" ) return HTTPException(status_code=409, detail="Resource already exists") @@ -133,6 +138,4 @@ def handle_foreign_key_violation(exc: Exception) -> HTTPException: @staticmethod def handle_check_violation(exc: Exception) -> HTTPException: - return HTTPException( - status_code=400, detail="Constraint check failed" - ) + return HTTPException(status_code=400, detail="Constraint check failed") diff --git a/app/core/securite.py b/app/core/securite.py index b9bc82a0..e3bb124b 100644 --- a/app/core/securite.py +++ b/app/core/securite.py @@ -12,6 +12,7 @@ from app.core.config import settings from app.core.exceptions import AppException from app.core.logger import logger + pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto") @@ -42,15 +43,21 @@ def create_acces_mobile_token(session_id: str) -> str: payload: dict[str, Any] = { "session_id": session_id, "exp": int( - (datetime.now(timezone.utc) + timedelta(seconds=Get_expiry_time())).timestamp() + ( + datetime.now(timezone.utc) + timedelta(seconds=Get_expiry_time()) + ).timestamp() ), } - return jwt.encode(payload, key=settings.jwt_secret, algorithm=settings.jwt_algorithm) + return jwt.encode( + payload, key=settings.jwt_secret, algorithm=settings.jwt_algorithm + ) def decode_access_mobile_token(token: str) -> dict[str, Any]: try: - payload = jwt.decode(token, key=settings.jwt_secret, algorithms=[settings.jwt_algorithm]) + payload = jwt.decode( + token, key=settings.jwt_secret, algorithms=[settings.jwt_algorithm] + ) return payload except jwt.ExpiredSignatureError: raise AppException.unauthorized("Token has expired") @@ -65,6 +72,7 @@ def create_raw_refresh_token() -> str: def hash_refresh_token(raw_token: str) -> str: return hashlib.sha256(raw_token.encode("utf-8")).hexdigest() + def create_totp_secret() -> str: return pyotp.random_base32() @@ -74,7 +82,9 @@ def get_totp_uri(secret: str, email: str) -> str: return totp.provisioning_uri(name=email, issuer_name=settings.totp_issuer) -def verify_totp_token_with_window(secret: str, token: str, valid_window: int = 8) -> bool: +def verify_totp_token_with_window( + secret: str, token: str, valid_window: int = 8 +) -> bool: totp = pyotp.TOTP(secret) return totp.verify(token, valid_window=valid_window) @@ -84,15 +94,21 @@ def generate_Acces_token_stuff(user_id: str, role: str) -> str: "user_id": user_id, "role": role, "exp": int( - (datetime.now(timezone.utc) + timedelta(seconds=Get_expiry_time())).timestamp() + ( + datetime.now(timezone.utc) + timedelta(seconds=Get_expiry_time()) + ).timestamp() ), } - return jwt.encode(payload, key=settings.jwt_secret, algorithm=settings.jwt_algorithm) + return jwt.encode( + payload, key=settings.jwt_secret, algorithm=settings.jwt_algorithm + ) + def _get_refresh_cache_aesgcm() -> AESGCM: key = base64.b64decode(settings.encryption_key) return AESGCM(key) + def encrypt_refresh_cache_payload(plaintext: str) -> str: """Encrypt a JSON string for storage in Redis. Returns a base64 string safe to store directly (nonce + ciphertext packed together).""" @@ -101,6 +117,7 @@ def encrypt_refresh_cache_payload(plaintext: str) -> str: ciphertext = aes.encrypt(nonce, plaintext.encode("utf-8"), None) return base64.b64encode(nonce + ciphertext).decode("utf-8") + def decrypt_refresh_cache_payload(encoded: str) -> str: """Reverse of encrypt_refresh_cache_payload. Raises on tampering or wrong key — treat any exception as 'cache miss'.""" @@ -152,13 +169,13 @@ def create_access_staff_token(staff_id: str, role: str) -> str: role=role, type="access", exp=int( - (datetime.now(timezone.utc) + timedelta(seconds=Get_expiry_time())).timestamp() + ( + datetime.now(timezone.utc) + timedelta(seconds=Get_expiry_time()) + ).timestamp() ), ) return jwt.encode( - payload.model_dump(), - key=settings.jwt_secret, - algorithm=settings.jwt_algorithm + payload.model_dump(), key=settings.jwt_secret, algorithm=settings.jwt_algorithm ) @@ -171,13 +188,13 @@ def create_refresh_staff_token(staff_id: str, role: str) -> str: role=role, type="refresh", exp=int( - (datetime.now(timezone.utc) + timedelta(seconds=Get_expiry_time() * 4)).timestamp() + ( + datetime.now(timezone.utc) + timedelta(seconds=Get_expiry_time() * 4) + ).timestamp() ), ) return jwt.encode( - payload.model_dump(), - key=settings.jwt_secret, - algorithm=settings.jwt_algorithm + payload.model_dump(), key=settings.jwt_secret, algorithm=settings.jwt_algorithm ) @@ -187,9 +204,7 @@ def decode_staff_token(token: str) -> StaffJWTPayload: """ try: decoded = jwt.decode( - token, - key=settings.jwt_secret, - algorithms=[settings.jwt_algorithm] + token, key=settings.jwt_secret, algorithms=[settings.jwt_algorithm] ) return StaffJWTPayload(**decoded) except jwt.ExpiredSignatureError: diff --git a/app/core/utils.py b/app/core/utils.py index 5e8db864..d669fa17 100644 --- a/app/core/utils.py +++ b/app/core/utils.py @@ -19,7 +19,6 @@ def check_extension( if len(filename_splitted) < 2: raise AppException.bad_request("File should have an extension") - file_ext = filename_splitted[-1] if file_ext not in allowed_extensions: @@ -27,11 +26,9 @@ def check_extension( f"File extension {file_ext} is not allowed. Allowed extensions are: {', '.join(allowed_extensions)}" ) - if file.content_type not in ext_content_type_map[file_ext]: raise AppException.bad_request( f"File content type {file.content_type} does not match extension {file_ext}" ) - return file_ext diff --git a/app/deps/cookie_auth.py b/app/deps/cookie_auth.py index ae9b3885..456fe680 100644 --- a/app/deps/cookie_auth.py +++ b/app/deps/cookie_auth.py @@ -8,19 +8,17 @@ from app.core.securite import decode_staff_token - def _role_value(role: object) -> str: return getattr(role, "value", str(role)) - async def get_current_staff_user( container: Annotated[Container, Depends(get_container)], token: Annotated[str | None, Cookie(alias="access_token")] = None, ) -> StaffUser: - if token is None : + if token is None: raise AppException.unauthorized("token doestn exist") - else : + else: payload = decode_staff_token(token) staff_id_str = payload.sub @@ -33,7 +31,9 @@ async def get_current_staff_user( except ValueError: raise AppException.unauthorized("Invalid staff ID in token") - staff_user = await container.staff_user_querier.get_staff_user_by_id(id=staff_id) + staff_user = await container.staff_user_querier.get_staff_user_by_id( + id=staff_id + ) if staff_user is None: raise AppException.not_found("Staff user not found") @@ -57,6 +57,7 @@ def ensure_admin_staff(current_staff_user: StaffUser) -> StaffUser: raise AppException.forbidden("Admin access required") return current_staff_user + async def require_admin_staff( current_staff_user: Annotated[StaffUser, Depends(get_current_staff_user)], ) -> StaffUser: diff --git a/app/deps/rate_limit.py b/app/deps/rate_limit.py index b146adf0..57428701 100644 --- a/app/deps/rate_limit.py +++ b/app/deps/rate_limit.py @@ -5,6 +5,7 @@ from app.infra.redis import RedisClient from app.core.logger import logger + def RateLimiter(requests: int, window: int) -> Callable: async def _rate_limit_dependency(request: Request) -> None: client_ip = get_client_ip(request) or "127.0.0.1" @@ -20,7 +21,9 @@ async def _rate_limit_dependency(request: Request) -> None: except HTTPException: raise except Exception: - logger.warning("rate_limit: redis unavailable, failing open for key=%s", key) + logger.warning( + "rate_limit: redis unavailable, failing open for key=%s", key + ) return if current > requests: diff --git a/app/deps/token_auth.py b/app/deps/token_auth.py index 7c778309..12c67182 100644 --- a/app/deps/token_auth.py +++ b/app/deps/token_auth.py @@ -52,7 +52,9 @@ async def get_current_mobile_user( if cached.blocked: raise HTTPException(status_code=403, detail="User is blocked") - if (now - cached.last_active).total_seconds() > settings.SESSION_ACTIVITY_THROTTLE_SECONDS: + if ( + now - cached.last_active + ).total_seconds() > settings.SESSION_ACTIVITY_THROTTLE_SECONDS: new_idle_expires_at = min( now + timedelta(days=settings.MOBILE_SESSION_DAYS), cached.absolute_expires_at, @@ -80,7 +82,9 @@ async def get_current_mobile_user( ) # --- Slow path: Postgres fallback --- - session = await container.session_service.session_querier.get_session_by_id(id=session_id) + session = await container.session_service.session_querier.get_session_by_id( + id=session_id + ) if not session: raise HTTPException(status_code=401, detail="Session not found") @@ -126,7 +130,9 @@ async def require_onboarded_mobile_user( onboarding; everything else (photos, notifications, audits, /event/me) depends on this instead. """ - user = await container.auth_service.user_querier.get_user_by_id(id=current_user.user_id) + user = await container.auth_service.user_querier.get_user_by_id( + id=current_user.user_id + ) if user is None: raise HTTPException(status_code=401, detail="User not found") if user.face_embedding is None: diff --git a/app/infra/email.py b/app/infra/email.py index b2beaf55..10d06d19 100644 --- a/app/infra/email.py +++ b/app/infra/email.py @@ -5,6 +5,7 @@ from app.core.config import settings from app.core.logger import logger + class EmailSender: @staticmethod async def send_otp_email(to_email: str, otp: str) -> bool: @@ -18,7 +19,7 @@ async def send_otp_email(to_email: str, otp: str) -> bool: headers = { "Authorization": f"Bearer {settings.RESEND_API_KEY}", "Content-Type": "application/json", - "User-Agent": "multAI-Backend/1.0" + "User-Agent": "multAI-Backend/1.0", } html_content = f""" @@ -34,14 +35,11 @@ async def send_otp_email(to_email: str, otp: str) -> bool: "from": settings.EMAIL_FROM, "to": [to_email], "subject": "Votre code de vérification multAI", - "html": html_content + "html": html_content, } req = urllib.request.Request( - url, - data=json.dumps(data).encode("utf-8"), - headers=headers, - method="POST" + url, data=json.dumps(data).encode("utf-8"), headers=headers, method="POST" ) def _send() -> bool: @@ -52,7 +50,9 @@ def _send() -> bool: return True except urllib.error.HTTPError as e: err_body = e.read() - logger.error("Failed to send email via Resend: %s - %s", e.code, err_body) + logger.error( + "Failed to send email via Resend: %s - %s", e.code, err_body + ) return False except Exception as e: logger.error("Error sending email: %s", str(e)) diff --git a/app/infra/google_drive.py b/app/infra/google_drive.py index 05633bf7..b8b39ca6 100644 --- a/app/infra/google_drive.py +++ b/app/infra/google_drive.py @@ -15,6 +15,7 @@ GOOGLE_TOKEN_URL, GOOGLE_USERINFO_URL, ) + GOOGLE_DRIVE_LIST_FILES_URL = "https://www.googleapis.com/drive/v3/files" @@ -117,8 +118,7 @@ async def exchange_code(code: str) -> GoogleTokenResponse: expires_at=expires_at, scope=GoogleDriveClient._optional_str(data, "scope") or settings.GOOGLE_OAUTH_SCOPES, - token_type=GoogleDriveClient._optional_str(data, "token_type") - or "Bearer", + token_type=GoogleDriveClient._optional_str(data, "token_type") or "Bearer", ) @staticmethod @@ -142,8 +142,7 @@ async def refresh_access_token(refresh_token: str) -> GoogleTokenResponse: expires_at=expires_at, scope=GoogleDriveClient._optional_str(data, "scope") or settings.GOOGLE_OAUTH_SCOPES, - token_type=GoogleDriveClient._optional_str(data, "token_type") - or "Bearer", + token_type=GoogleDriveClient._optional_str(data, "token_type") or "Bearer", ) @staticmethod @@ -204,12 +203,16 @@ async def upload_file( metadata["parents"] = [folder_id] body = ( - f"--{boundary}\r\n" - "Content-Type: application/json; charset=UTF-8\r\n\r\n" - f"{json.dumps(metadata)}\r\n" - f"--{boundary}\r\n" - f"Content-Type: {content_type}\r\n\r\n" - ).encode("utf-8") + data + f"\r\n--{boundary}--".encode("utf-8") + ( + f"--{boundary}\r\n" + "Content-Type: application/json; charset=UTF-8\r\n\r\n" + f"{json.dumps(metadata)}\r\n" + f"--{boundary}\r\n" + f"Content-Type: {content_type}\r\n\r\n" + ).encode("utf-8") + + data + + f"\r\n--{boundary}--".encode("utf-8") + ) def _request() -> dict[str, object]: url = ( @@ -234,12 +237,16 @@ def _request() -> dict[str, object]: f"Google Drive file upload failed: {details or exc.reason}" ) from exc except urllib.error.URLError as exc: - raise AppException.internal_error("Unable to reach Google APIs") from exc + raise AppException.internal_error( + "Unable to reach Google APIs" + ) from exc result = await asyncio.to_thread(_request) size_raw = result.get("size", "0") try: - size_bytes = int(size_raw) if isinstance(size_raw, (str, int)) else len(data) + size_bytes = ( + int(size_raw) if isinstance(size_raw, (str, int)) else len(data) + ) except (TypeError, ValueError): size_bytes = len(data) @@ -299,11 +306,15 @@ async def list_folder_files( raw_files = data.get("files", []) if not isinstance(raw_files, list): - raise AppException.bad_request("Google Drive folder listing response is invalid") + raise AppException.bad_request( + "Google Drive folder listing response is invalid" + ) for raw_file in raw_files: if not isinstance(raw_file, dict): - raise AppException.bad_request("Google Drive folder entry is invalid") + raise AppException.bad_request( + "Google Drive folder entry is invalid" + ) metadata = GoogleDriveClient._file_metadata_from_dict(raw_file) if metadata.mime_type == GoogleDriveClient._drive_folder_mime_type: continue @@ -313,7 +324,9 @@ async def list_folder_files( if next_page_token_raw is None: break if not isinstance(next_page_token_raw, str) or not next_page_token_raw: - raise AppException.bad_request("Google Drive next page token is invalid") + raise AppException.bad_request( + "Google Drive next page token is invalid" + ) next_page_token = next_page_token_raw return files @@ -349,7 +362,9 @@ async def list_folder_contents( raw_files = data.get("files", []) if not isinstance(raw_files, list): - raise AppException.bad_request("Google Drive folder listing response is invalid") + raise AppException.bad_request( + "Google Drive folder listing response is invalid" + ) for raw_file in raw_files: if not isinstance(raw_file, dict): @@ -468,7 +483,9 @@ def _request() -> dict[str, object]: final_url = url if query_params: final_url = f"{url}?{urllib.parse.urlencode(query_params)}" - request = urllib.request.Request(final_url, headers=headers or {}, method="GET") + request = urllib.request.Request( + final_url, headers=headers or {}, method="GET" + ) try: with urllib.request.urlopen(request, timeout=15) as response: return json.loads(response.read().decode("utf-8")) @@ -494,7 +511,9 @@ def _request() -> tuple[bytes, str, str]: final_url = url if query_params: final_url = f"{url}?{urllib.parse.urlencode(query_params)}" - request = urllib.request.Request(final_url, headers=headers or {}, method="GET") + request = urllib.request.Request( + final_url, headers=headers or {}, method="GET" + ) try: with urllib.request.urlopen(request, timeout=30) as response: body = response.read() @@ -508,6 +527,8 @@ def _request() -> tuple[bytes, str, str]: f"Google file download failed: {details or exc.reason}" ) from exc except urllib.error.URLError as exc: - raise AppException.internal_error("Unable to download file from Google Drive") from exc + raise AppException.internal_error( + "Unable to download file from Google Drive" + ) from exc return await asyncio.to_thread(_request) diff --git a/app/infra/minio.py b/app/infra/minio.py index 9a666c5c..5ee2b040 100644 --- a/app/infra/minio.py +++ b/app/infra/minio.py @@ -24,6 +24,7 @@ DOCUMENTS_BUCKET_NAME = CORE_DOCUMENTS_BUCKET_NAME WA_SIM_BUCKET_NAME = CORE_WA_SIM_BUCKET_NAME + async def init_minio_client( minio_host: str, minio_port: int, minio_root_user: str, minio_root_password: str ) -> None: @@ -38,6 +39,7 @@ async def init_minio_client( if not await Bucket.client.bucket_exists(bucket_name): await Bucket.client.make_bucket(bucket_name) + @dataclass(frozen=True) class ObjectStat: size: int @@ -94,9 +96,7 @@ async def get(self, object_name: str) -> tuple[bytes, str, str]: raise e data = await res.read() - content_type = ( - res.content_type if res.content_type else DEFAULT_CONTENT_TYPE - ) + content_type = res.content_type if res.content_type else DEFAULT_CONTENT_TYPE filename = res.headers.get("x-amz-meta-filename", f"{object_name}") res.close() @@ -155,7 +155,11 @@ async def stat(self, object_name: str) -> ObjectStat | None: if e.code == "NoSuchKey": return None raise - return ObjectStat(size=result.size or 0, content_type=result.content_type or DEFAULT_CONTENT_TYPE) + return ObjectStat( + size=result.size or 0, + content_type=result.content_type or DEFAULT_CONTENT_TYPE, + ) + image_ext_content_type_map = { "apng": ["image/apng"], @@ -169,6 +173,7 @@ async def stat(self, object_name: str) -> ObjectStat | None: "ico": ["image/x-icon", "image/vnd.microsoft.icon"], } + class ImageBucket(Bucket): def __init__(self, file_prefix: str): super().__init__(IMAGES_BUCKET_NAME, file_prefix) @@ -177,10 +182,12 @@ async def put(self, file: UploadFile, object_name: str | None = None) -> str: check_extension(file, image_ext_content_type_map) return await super().put(file, object_name) + class DocumentBucket(Bucket): def __init__(self, file_prefix: str): super().__init__(DOCUMENTS_BUCKET_NAME, file_prefix) + class WaSimBucket(Bucket): def __init__(self) -> None: super().__init__(WA_SIM_BUCKET_NAME, "") diff --git a/app/infra/nats.py b/app/infra/nats.py index cf5c1b76..aaa6fc7a 100644 --- a/app/infra/nats.py +++ b/app/infra/nats.py @@ -75,7 +75,9 @@ async def connect( if NatsClient._nc is None: nc = NATS() await nc.connect( - servers=[f"nats://{host or settings.NATS_HOST}:{port or settings.NATS_PORT}"], + servers=[ + f"nats://{host or settings.NATS_HOST}:{port or settings.NATS_PORT}" + ], user=user or settings.NATS_USER, password=password or settings.NATS_PASSWORD, ) @@ -102,7 +104,9 @@ async def publish(subject: NatsSubjects | str, message: bytes) -> None: await nc.publish(subject_name, message) @staticmethod - async def subscribe(subject: NatsSubjects | str, callback: Callable[[Any], Any]) -> None: + async def subscribe( + subject: NatsSubjects | str, callback: Callable[[Any], Any] + ) -> None: if NatsClient._nc is None: await NatsClient.connect() nc = NatsClient._nc @@ -114,9 +118,10 @@ async def _wrapper(msg: Msg) -> None: subject_name = subject.value if isinstance(subject, NatsSubjects) else subject await nc.subscribe(subject_name, cb=_wrapper) # type: ignore - @staticmethod - async def js_publish(subject: NatsSubjects | str, message: bytes, stream_name: str | None = None) -> None: + async def js_publish( + subject: NatsSubjects | str, message: bytes, stream_name: str | None = None + ) -> None: if NatsClient._js is None: await NatsClient.connect() js = NatsClient._js @@ -124,10 +129,14 @@ async def js_publish(subject: NatsSubjects | str, message: bytes, stream_name: s subject_name = subject.value if isinstance(subject, NatsSubjects) else subject resolved_stream = stream_name or SUBJECT_TO_STREAM.get(subject_name) if resolved_stream is None: - logger.warning(f"No stream mapped for subject {subject_name}, but js_publish was called.") + logger.warning( + f"No stream mapped for subject {subject_name}, but js_publish was called." + ) return await NatsClient.publish(subject, message) - - await NatsClient.ensure_stream(stream_name=resolved_stream, subjects=[subject_name]) + + await NatsClient.ensure_stream( + stream_name=resolved_stream, subjects=[subject_name] + ) await js.publish(subject_name, message, stream=resolved_stream) @staticmethod @@ -136,7 +145,7 @@ async def js_subscribe( callback: Callable[[Any], Any], stream_name: str | None = None, durable_name: str | None = None, - ack_policy: AckPolicy = AckPolicy.EXPLICIT + ack_policy: AckPolicy = AckPolicy.EXPLICIT, ) -> None: if NatsClient._js is None: await NatsClient.connect() @@ -144,21 +153,27 @@ async def js_subscribe( subject_name = subject.value if isinstance(subject, NatsSubjects) else subject resolved_stream = stream_name or SUBJECT_TO_STREAM.get(subject_name) if not resolved_stream: - raise ValueError(f"Cannot js_subscribe to {subject_name}: no stream mapped.") - + raise ValueError( + f"Cannot js_subscribe to {subject_name}: no stream mapped." + ) + resolved_durable = durable_name or f"{resolved_stream}_consumer" - await NatsClient.ensure_stream(stream_name=resolved_stream, subjects=[subject_name]) + await NatsClient.ensure_stream( + stream_name=resolved_stream, subjects=[subject_name] + ) async def _wrapper(msg: Msg) -> None: try: await callback(msg.data) await msg.ack() except Exception as exc: - logger.error(f"Error processing message from {subject_name}, NACKing: {exc}") + logger.error( + f"Error processing message from {subject_name}, NACKing: {exc}" + ) await msg.nak() raise - + js = NatsClient._js assert js is not None await js.subscribe( @@ -179,7 +194,7 @@ async def ensure_stream(*, stream_name: str, subjects: list[str]) -> None: try: await js.stream_info(stream_name) except NotFoundError: - await js.add_stream( # type: ignore + await js.add_stream( # type: ignore name=stream_name, config=StreamConfig( name=stream_name, diff --git a/app/infra/redis.py b/app/infra/redis.py index 7e901015..1f2e2299 100644 --- a/app/infra/redis.py +++ b/app/infra/redis.py @@ -10,13 +10,12 @@ class RedisClient: _instance: ClassVar["RedisClient | None"] = None def __init__(self, host: str, port: int, password: str) -> None: - self._client = Redis.from_url( # type: ignore + self._client = Redis.from_url( # type: ignore f"redis://{host}:{port}", password=password, decode_responses=True, ) - @classmethod def init(cls, host: str, port: int, password: str) -> "RedisClient": if cls._instance is not None: @@ -32,7 +31,6 @@ def get_instance(cls) -> "RedisClient": return cls._instance - async def set( self, key: RedisKey | str, @@ -61,7 +59,6 @@ async def expire(self, key: RedisKey | str, seconds: int) -> bool: async def incr(self, key: RedisKey | str) -> int: return await self._client.incr(key) - async def sadd(self, key: RedisKey | str, *values: str) -> int: result = await self._client.sadd(key, *values) # type: ignore[misc] return int(cast(int, result)) @@ -74,7 +71,6 @@ async def srem(self, key: RedisKey | str, *values: str) -> int: result = await self._client.srem(key, *values) # type: ignore[misc] return int(cast(int, result)) - async def close(self) -> None: await self._client.close() type(self)._instance = None diff --git a/app/main.py b/app/main.py index f766fe9a..51b99184 100644 --- a/app/main.py +++ b/app/main.py @@ -19,8 +19,6 @@ from app.core.logger import configure_logger, logger - - configure_logger() @@ -48,21 +46,25 @@ async def dispatch( return response - async def _approval_expiry_loop() -> None: while True: await asyncio.sleep(3600) try: async with engine.begin() as conn: from app.container import Container + container = Container(conn) - await container.photo_approval_service.expire_stale(settings.PHOTO_APPROVAL_TIMEOUT_DAYS) + await container.photo_approval_service.expire_stale( + settings.PHOTO_APPROVAL_TIMEOUT_DAYS + ) except Exception as exc: logger.warning("Approval expiry task failed: %s", exc) MAX_RETRIES = 5 RETRY_DELAY = 2 # seconds + + @asynccontextmanager async def lifespan(app: FastAPI) -> AsyncIterator[None]: @@ -78,7 +80,9 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: except Exception as e: print(f"[MINIO] Attempt {attempt} failed: {e}") if attempt == MAX_RETRIES: - raise RuntimeError("Cannot connect to MinIO after multiple attempts") from e + raise RuntimeError( + "Cannot connect to MinIO after multiple attempts" + ) from e await asyncio.sleep(RETRY_DELAY) RedisClient.init( @@ -100,7 +104,6 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: await NatsClient.close() - app = FastAPI( title="multAI API", description="Mobile and Web API for multAI", @@ -114,13 +117,16 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: @app.exception_handler(Exception) async def global_exception_handler(request: Request, exc: Exception) -> JSONResponse: - logger.error("Uncaught exception on %s %s: %s", request.method, request.url.path, exc) + logger.error( + "Uncaught exception on %s %s: %s", request.method, request.url.path, exc + ) logger.error(traceback.format_exc()) return JSONResponse( status_code=500, content={"detail": "Internal server error"}, ) + app.add_middleware(RequestLoggingMiddleware) app.add_middleware( diff --git a/app/router/mobile/auth.py b/app/router/mobile/auth.py index 9be7547d..1c524ca7 100644 --- a/app/router/mobile/auth.py +++ b/app/router/mobile/auth.py @@ -22,40 +22,66 @@ InactivateDeviceRequest, UpdateProfileRequest, ) -from app.schema.response.mobile.auth import MeResponse, DeviceSchema, MobileAuthResponse, SessionSchema, UserSchema, RegisterPendingResponse +from app.schema.response.mobile.auth import ( + MeResponse, + DeviceSchema, + MobileAuthResponse, + SessionSchema, + UserSchema, + RegisterPendingResponse, +) router = APIRouter(prefix="/auth") -@router.post("/register", response_model=RegisterPendingResponse, dependencies=[Depends(RateLimiter(requests=5, window=60))]) + +@router.post( + "/register", + response_model=RegisterPendingResponse, + dependencies=[Depends(RateLimiter(requests=5, window=60))], +) async def mobile_register( req: MobileRegisterRequest, request: Request, container: Container = Depends(get_container), ) -> RegisterPendingResponse: client_ip = get_client_ip(request) - result = await container.auth_service.mobile_register(container.redis, req, client_ip=client_ip) + result = await container.auth_service.mobile_register( + container.redis, req, client_ip=client_ip + ) return result -@router.post("/register/resend-otp", response_model=RegisterPendingResponse, dependencies=[Depends(RateLimiter(requests=5, window=60))]) +@router.post( + "/register/resend-otp", + response_model=RegisterPendingResponse, + dependencies=[Depends(RateLimiter(requests=5, window=60))], +) async def mobile_register_resend_otp( req: ResendOtpRequest, request: Request, container: Container = Depends(get_container), ) -> RegisterPendingResponse: client_ip = get_client_ip(request) - result = await container.auth_service.mobile_register_resend_otp(container.redis, req.email, client_ip=client_ip) + result = await container.auth_service.mobile_register_resend_otp( + container.redis, req.email, client_ip=client_ip + ) return result -@router.post("/register/verify", response_model=MobileAuthResponse, dependencies=[Depends(RateLimiter(requests=10, window=60))]) +@router.post( + "/register/verify", + response_model=MobileAuthResponse, + dependencies=[Depends(RateLimiter(requests=10, window=60))], +) async def mobile_register_verify( req: RegisterVerifyRequest, request: Request, container: Container = Depends(get_container), ) -> MobileAuthResponse: client_ip = get_client_ip(request) - result = await container.auth_service.verify_mobile_register(container.redis, req, client_ip=client_ip) + result = await container.auth_service.verify_mobile_register( + container.redis, req, client_ip=client_ip + ) await container.audit_service.create_record( event_type=AuditEventType.USER_SIGNUP, user_id=result.user_id, @@ -64,14 +90,20 @@ async def mobile_register_verify( return result -@router.post("/login", response_model=MobileAuthResponse, dependencies=[Depends(RateLimiter(requests=5, window=60))]) +@router.post( + "/login", + response_model=MobileAuthResponse, + dependencies=[Depends(RateLimiter(requests=5, window=60))], +) async def mobile_login( req: MobileLoginRequest, request: Request, container: Container = Depends(get_container), ) -> MobileAuthResponse: client_ip = get_client_ip(request) - result = await container.auth_service.mobile_login(container.redis, req, client_ip=client_ip) + result = await container.auth_service.mobile_login( + container.redis, req, client_ip=client_ip + ) await container.audit_service.create_record( event_type=AuditEventType.USER_LOGIN, user_id=result.user_id, @@ -89,7 +121,9 @@ async def refresh_token( req: RefreshTokenRequest, container: Container = Depends(get_container), ) -> MobileAuthResponse: - return await container.auth_service.refresh_token(container.redis, req.refresh_token) + return await container.auth_service.refresh_token( + container.redis, req.refresh_token + ) @router.post("/logout") @@ -122,8 +156,10 @@ async def revoke_device( if device is None or device.user_id != current_user.user_id: raise AppException.not_found("Device not found") - session = await container.session_service.session_querier.get_session_by_device_for_user( - device_id=device_id, user_id=current_user.user_id + session = ( + await container.session_service.session_querier.get_session_by_device_for_user( + device_id=device_id, user_id=current_user.user_id + ) ) await container.device_service.revoke_device( @@ -132,8 +168,9 @@ async def revoke_device( ) if session: - await container.session_service.delete_session_cache(container.redis, session.id) - + await container.session_service.delete_session_cache( + container.redis, session.id + ) return {"message": "Device revoked successfully"} @@ -213,6 +250,7 @@ async def get_me( sessions=session_schema, ) + @router.patch("/me/profile", response_model=UserSchema) async def update_profile( req: UpdateProfileRequest, @@ -249,7 +287,8 @@ async def upload_avatar( ) try: user = await container.auth_service.update_avatar( - user_id=current_user.user_id, avatar_key=object_name, + user_id=current_user.user_id, + avatar_key=object_name, ) except Exception: await container.auth_service.delete_avatar_bytes(avatar_key=object_name) diff --git a/app/router/mobile/enrollement.py b/app/router/mobile/enrollement.py index 1c7cc4cf..e98d3770 100644 --- a/app/router/mobile/enrollement.py +++ b/app/router/mobile/enrollement.py @@ -41,6 +41,7 @@ class EnrollmentResponse(BaseModel): router = APIRouter() + async def _build_face_image_payload(file: UploadFile) -> FaceImagePayload: payload = await build_image_payload(file) return FaceImagePayload( @@ -149,7 +150,8 @@ async def image_payloads() -> AsyncIterator[FaceImagePayload]: except Exception as exc: logger.warning( "enroll: redis unavailable, failing open (no duplicate-submission lock) for user %s: %s", - user.user_id, exc, + user.user_id, + exc, ) lock_acquired = True if not lock_acquired: diff --git a/app/router/mobile/event.py b/app/router/mobile/event.py index 33548e75..00b0ac1c 100644 --- a/app/router/mobile/event.py +++ b/app/router/mobile/event.py @@ -9,15 +9,16 @@ router = APIRouter(prefix="/event") + + @router.post("/join", response_model=JoinEventResponse) async def join_event( req: JoinEventRequest, container: Container = Depends(get_container), current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), -)-> JoinEventResponse: +) -> JoinEventResponse: return await container.event_service.join_event_by_code( - user_id=current_user.user_id, - code=req.event_code + user_id=current_user.user_id, code=req.event_code ) @@ -25,5 +26,5 @@ async def join_event( async def get_my_joined_events( container: Container = Depends(get_container), current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), -)-> List[UserEventResponse]: +) -> List[UserEventResponse]: return await container.event_service.get_my_events(current_user.user_id) diff --git a/app/router/mobile/notifications.py b/app/router/mobile/notifications.py index ca973fcb..e5f73e81 100644 --- a/app/router/mobile/notifications.py +++ b/app/router/mobile/notifications.py @@ -26,8 +26,10 @@ async def mark_as_read( container: Container = Depends(get_container), current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), ) -> UserNotificationListResponse: - notifications = await container.user_notifications_service.mark_notifications_as_read( - notification_ids=req.notification_ids, - user_id=current_user.user_id, + notifications = ( + await container.user_notifications_service.mark_notifications_as_read( + notification_ids=req.notification_ids, + user_id=current_user.user_id, + ) ) return UserNotificationListResponse.from_models(notifications) diff --git a/app/router/staff/drive.py b/app/router/staff/drive.py index de27e3f5..e301ab49 100644 --- a/app/router/staff/drive.py +++ b/app/router/staff/drive.py @@ -58,7 +58,9 @@ async def google_drive_callback( raise AppException.bad_request(f"Google OAuth error: {error}") try: - connection, redirect_url = await container.staff_drive_service.handle_callback(code, state) + connection, redirect_url = await container.staff_drive_service.handle_callback( + code, state + ) except HTTPException as exc: if redirect_url is not None: return RedirectResponse( @@ -108,7 +110,9 @@ async def disconnect_google_drive( container: Container = Depends(get_container), ) -> GoogleDriveDisconnectResponse: await container.staff_drive_service.disconnect(current_staff_user.id) - return GoogleDriveDisconnectResponse(message="Google Drive disconnected successfully") + return GoogleDriveDisconnectResponse( + message="Google Drive disconnected successfully" + ) @router.get("/browse", response_model=DriveBrowseResponse) @@ -121,7 +125,8 @@ async def browse_drive( current_staff_user.id ) files = await GoogleDriveClient.list_folder_contents( - access_token=access_token, folder_id=folder_id, + access_token=access_token, + folder_id=folder_id, ) return DriveBrowseResponse( items=[ @@ -148,7 +153,9 @@ async def search_drive( current_staff_user.id ) files = await GoogleDriveClient.search_files( - access_token=access_token, query=q, file_type=type, + access_token=access_token, + query=q, + file_type=type, ) return DriveBrowseResponse( items=[ diff --git a/app/router/staff/uploads.py b/app/router/staff/uploads.py index 5ded5561..af597a34 100644 --- a/app/router/staff/uploads.py +++ b/app/router/staff/uploads.py @@ -95,7 +95,9 @@ async def get_upload_request_group( return UploadRequestGroupSchema.from_details(group) -@router.get("/groups/{group_id}/photos", response_model=UploadRequestGroupPhotoListResponse) +@router.get( + "/groups/{group_id}/photos", response_model=UploadRequestGroupPhotoListResponse +) async def list_upload_request_group_photos( group_id: UUID, current_staff_user: StaffUser = Depends(get_current_staff_user), @@ -160,7 +162,9 @@ async def list_upload_request_photos( current_staff_user=current_staff_user, ) return UploadRequestPhotoListResponse( - items=[UploadRequestPhotoSchema.model_validate(p) for p in upload_request.photos] + items=[ + UploadRequestPhotoSchema.model_validate(p) for p in upload_request.photos + ] ) @@ -177,7 +181,9 @@ async def preview_upload_request_photo( current_staff_user=current_staff_user, ) headers = {"Content-Disposition": f'inline; filename="{preview.file_name}"'} - return Response(content=preview.data, media_type=preview.content_type, headers=headers) + return Response( + content=preview.data, media_type=preview.content_type, headers=headers + ) @router.post("/{request_id}/approve", response_model=UploadRequestSchema) diff --git a/app/router/staff/uploads_direct.py b/app/router/staff/uploads_direct.py index 267cd9bd..88225de3 100644 --- a/app/router/staff/uploads_direct.py +++ b/app/router/staff/uploads_direct.py @@ -20,6 +20,7 @@ router = APIRouter(prefix="/uploads/direct") # this endpoint are for staff to upload images directly to the system and very large files and they can resume and restart and retry . + @router.post("/groups", response_model=UploadRequestGroupSchema) async def create_direct_group( req: CreateDirectGroupRequest, @@ -27,10 +28,12 @@ async def create_direct_group( container: Container = Depends(get_container), ) -> UploadRequestGroupSchema: group = await container.upload_requests_service.create_direct_group( - event_id=req.event_id, requested_by=current_staff_user, + event_id=req.event_id, + requested_by=current_staff_user, ) details = await container.upload_requests_service.get_group_details( - group_id=group.id, current_staff_user=current_staff_user, + group_id=group.id, + current_staff_user=current_staff_user, ) return UploadRequestGroupSchema.from_details(details) @@ -60,7 +63,9 @@ async def register_direct_batch( return RegisterDirectBatchResponse( group_id=group_id, items=[ - DirectUploadFileResponse(photo_id=photo.id, file_name=photo.file_name, upload_url=url) + DirectUploadFileResponse( + photo_id=photo.id, file_name=photo.file_name, upload_url=url + ) for photo, url in results ], ) @@ -73,9 +78,12 @@ async def confirm_direct_upload( container: Container = Depends(get_container), ) -> DirectUploadFileResponse: photo = await container.upload_requests_service.confirm_direct_upload( - photo_id=photo_id, requested_by=current_staff_user, + photo_id=photo_id, + requested_by=current_staff_user, + ) + return DirectUploadFileResponse( + photo_id=photo.id, file_name=photo.file_name, upload_url="" ) - return DirectUploadFileResponse(photo_id=photo.id, file_name=photo.file_name, upload_url="") @router.post("/photos/{photo_id}/fail", response_model=DirectUploadFileResponse) @@ -85,9 +93,12 @@ async def fail_direct_upload( container: Container = Depends(get_container), ) -> DirectUploadFileResponse: photo = await container.upload_requests_service.fail_direct_upload( - photo_id=photo_id, requested_by=current_staff_user, + photo_id=photo_id, + requested_by=current_staff_user, + ) + return DirectUploadFileResponse( + photo_id=photo.id, file_name=photo.file_name, upload_url="" ) - return DirectUploadFileResponse(photo_id=photo.id, file_name=photo.file_name, upload_url="") @router.post("/groups/{group_id}/resume", response_model=ResumeDirectGroupResponse) @@ -97,11 +108,14 @@ async def resume_direct_group( container: Container = Depends(get_container), ) -> ResumeDirectGroupResponse: results = await container.upload_requests_service.resume_direct_group( - group_id=group_id, requested_by=current_staff_user, + group_id=group_id, + requested_by=current_staff_user, ) return ResumeDirectGroupResponse( items=[ - DirectUploadFileResponse(photo_id=photo.id, file_name=photo.file_name, upload_url=url) + DirectUploadFileResponse( + photo_id=photo.id, file_name=photo.file_name, upload_url=url + ) for photo, url in results ] ) @@ -114,6 +128,7 @@ async def get_direct_group_status( container: Container = Depends(get_container), ) -> UploadRequestGroupSchema: details = await container.upload_requests_service.get_group_details( - group_id=group_id, current_staff_user=current_staff_user, + group_id=group_id, + current_staff_user=current_staff_user, ) return UploadRequestGroupSchema.from_details(details) diff --git a/app/router/web/auth.py b/app/router/web/auth.py index 24d7a402..76f945b1 100644 --- a/app/router/web/auth.py +++ b/app/router/web/auth.py @@ -9,19 +9,24 @@ from app.schema.response.web.auth import WebAuthResponse from app.schema.response.web.staff_user import StaffUserSchema from db.generated.models import StaffUser + router = APIRouter(prefix="/auth") -@router.post("/login", response_model=WebAuthResponse, description="so here both the dahbsoard will authneticate from this endpoitn ", dependencies=[Depends(RateLimiter(requests=5, window=60))]) +@router.post( + "/login", + response_model=WebAuthResponse, + description="so here both the dahbsoard will authneticate from this endpoitn ", + dependencies=[Depends(RateLimiter(requests=5, window=60))], +) async def admin_login( req: WebAuthRequest, - r:Response, + r: Response, container: Container = Depends(get_container), ) -> WebAuthResponse: authResponse = await container.staff_user_service.admin_login( - email=req.email, - password=req.password + email=req.email, password=req.password ) r.set_cookie( key="access_token", @@ -34,16 +39,16 @@ async def admin_login( return authResponse -@router.get("/me",response_model=StaffUserSchema) +@router.get("/me", response_model=StaffUserSchema) async def get_me_admin( - user:StaffUser = Depends(get_current_staff_user), + user: StaffUser = Depends(get_current_staff_user), ) -> StaffUserSchema: return StaffUserSchema( id=user.id, created_at=user.created_at, role=user.role, updated_at=user.updated_at, - email=user.email + email=user.email, ) diff --git a/app/router/web/event.py b/app/router/web/event.py index 83090834..c076bf8e 100644 --- a/app/router/web/event.py +++ b/app/router/web/event.py @@ -4,12 +4,9 @@ from app.container import get_container, Container -from app.deps.cookie_auth import ( - get_current_staff_user -) +from app.deps.cookie_auth import get_current_staff_user from app.schema.request.web.event import ( EventCreate, - ) from app.schema.response.web.event import ( EventResponse, @@ -28,24 +25,23 @@ async def list_events( status: Optional[models.EventStatus] = None, container: Container = Depends(get_container), current_staff: models.StaffUser = Depends(get_current_staff_user), -)-> List[EventResponse]: +) -> List[EventResponse]: """Staff Only: List all events with optional filters.""" if name: return await container.event_service.find_events_by_name(name) return await container.event_service.list_events( - limit=limit, - offset=offset, - status=status + limit=limit, offset=offset, status=status ) + @router.post("/", response_model=EventResponse, status_code=status.HTTP_201_CREATED) async def create_event( req: EventCreate, container: Container = Depends(get_container), current_staff: models.StaffUser = Depends(get_current_staff_user), -)-> EventResponse: +) -> EventResponse: """Staff Only: Create a new event.""" return await container.event_service.create_event(req, current_staff.id) @@ -54,8 +50,8 @@ async def create_event( async def archive_event( event_id: uuid.UUID, container: Container = Depends(get_container), - current_staff: models.StaffUser = Depends(get_current_staff_user), # Use Staff Dep -)-> EventResponse: + current_staff: models.StaffUser = Depends(get_current_staff_user), # Use Staff Dep +) -> EventResponse: """Staff Only: Archive an event.""" return await container.event_service.update_status(event_id, "archived") @@ -64,8 +60,8 @@ async def archive_event( async def schedule_event( event_id: uuid.UUID, container: Container = Depends(get_container), - current_staff: models.StaffUser = Depends(get_current_staff_user), # Use Staff Dep -)-> EventResponse: + current_staff: models.StaffUser = Depends(get_current_staff_user), # Use Staff Dep +) -> EventResponse: """Staff Only: Move to scheduled.""" return await container.event_service.update_status(event_id, "scheduled") @@ -75,7 +71,6 @@ async def draft_event( event_id: uuid.UUID, container: Container = Depends(get_container), current_staff: models.StaffUser = Depends(get_current_staff_user), -)-> EventResponse: +) -> EventResponse: """Staff Only: Move an event back to draft status.""" return await container.event_service.update_status(event_id, "draft") - diff --git a/app/router/web/staff_users.py b/app/router/web/staff_users.py index 5a57ae61..b3386126 100644 --- a/app/router/web/staff_users.py +++ b/app/router/web/staff_users.py @@ -6,7 +6,10 @@ from app.container import Container, get_container from app.core.logger import logger from app.deps.cookie_auth import get_current_staff_user -from app.schema.request.web.staff_user import StaffUserCreateRequest, StaffUserUpdateRequest +from app.schema.request.web.staff_user import ( + StaffUserCreateRequest, + StaffUserUpdateRequest, +) from app.schema.response.web.staff_user import StaffUserSchema from db.generated.models import StaffRole, StaffUser @@ -38,7 +41,7 @@ async def list_staff_users( email=user.email, role=user.role, created_at=user.created_at, - updated_at=user.updated_at + updated_at=user.updated_at, ) for user in staff_users ] @@ -58,7 +61,7 @@ async def create_staff_user( email=staff_user.email, role=staff_user.role, created_at=staff_user.created_at, - updated_at=staff_user.updated_at + updated_at=staff_user.updated_at, ) @@ -68,7 +71,6 @@ async def update_staff_user( req: StaffUserUpdateRequest, current_staff_user: StaffUser = Depends(get_current_staff_user), container: Container = Depends(get_container), - ) -> StaffUserSchema: staff_user = await container.staff_user_service.update_staff_user( id=staff_user_id, email=req.email, role=StaffRole(req.role) @@ -79,7 +81,7 @@ async def update_staff_user( email=staff_user.email, role=staff_user.role, created_at=staff_user.created_at, - updated_at=staff_user.updated_at + updated_at=staff_user.updated_at, ) @@ -96,5 +98,5 @@ async def delete_staff_user( email=staff_user.email, role=staff_user.role, created_at=staff_user.created_at, - updated_at=staff_user.updated_at + updated_at=staff_user.updated_at, ) diff --git a/app/router/web/stats.py b/app/router/web/stats.py index 3d9bb46c..a427a32a 100644 --- a/app/router/web/stats.py +++ b/app/router/web/stats.py @@ -3,16 +3,19 @@ from app.deps.cookie_auth import require_admin_staff from db.generated.models import StaffUser from app.schema.response.web.stats import ( - AdminStatsResponse, DriveUsageResponse, - ProcessingLoadResponse, AlertResponse + AdminStatsResponse, + DriveUsageResponse, + ProcessingLoadResponse, + AlertResponse, ) router = APIRouter(prefix="/stats", tags=["Web - Stats"]) + @router.get("/dashboard", response_model=AdminStatsResponse) async def get_dashboard( container: Container = Depends(get_container), - current_admin: StaffUser = Depends(require_admin_staff) + current_admin: StaffUser = Depends(require_admin_staff), ) -> AdminStatsResponse: """Staff Admin Only: Get global KPIs for the dashboard""" return await container.stats_service.get_dashboard_stats() @@ -21,7 +24,7 @@ async def get_dashboard( @router.get("/processing-load", response_model=ProcessingLoadResponse) async def get_processing_load( container: Container = Depends(get_container), - current_admin: StaffUser = Depends(require_admin_staff) + current_admin: StaffUser = Depends(require_admin_staff), ) -> ProcessingLoadResponse: """Staff Admin Only: Get pipeline processing load percentages""" return await container.stats_service.get_processing_load() @@ -30,7 +33,7 @@ async def get_processing_load( @router.get("/storage", response_model=DriveUsageResponse) async def get_storage( container: Container = Depends(get_container), - current_admin: StaffUser = Depends(require_admin_staff) + current_admin: StaffUser = Depends(require_admin_staff), ) -> DriveUsageResponse: """Staff Admin Only: Get MinIO storage consumption""" return await container.stats_service.get_storage_usage() @@ -39,7 +42,7 @@ async def get_storage( @router.get("/alerts", response_model=AlertResponse) async def get_alerts( container: Container = Depends(get_container), - current_admin: StaffUser = Depends(require_admin_staff) + current_admin: StaffUser = Depends(require_admin_staff), ) -> AlertResponse: """Staff Admin Only: Get recent alerts/notifications for the admin""" return await container.stats_service.get_staff_alerts(current_admin.id) diff --git a/app/router/web/users.py b/app/router/web/users.py index f167376c..a3946d8f 100644 --- a/app/router/web/users.py +++ b/app/router/web/users.py @@ -12,6 +12,7 @@ router = APIRouter(prefix="/users") + @router.post("/", response_model=AdminUserSchema, status_code=status.HTTP_201_CREATED) async def create_user( req: AdminUserCreateRequest, @@ -27,6 +28,7 @@ async def create_user( logger.info("admin %s created user %s", current_staff_user.id, user.id) return to_admin_user_schema(user) + @router.get("/", response_model=list[AdminUserSchema]) async def list_users( limit: int = Query( diff --git a/app/schema/internal/notification.py b/app/schema/internal/notification.py index 574a8565..701411b1 100644 --- a/app/schema/internal/notification.py +++ b/app/schema/internal/notification.py @@ -17,6 +17,7 @@ class UnifiedNotification(BaseModel): model_config = ConfigDict(extra="forbid") + PRIORITY_ORDER: tuple[NotificationPriority, ...] = ( NotificationPriority.HIGH, NotificationPriority.NORMAL, diff --git a/app/schema/request/mobile/auth.py b/app/schema/request/mobile/auth.py index fc0a312b..592b9475 100644 --- a/app/schema/request/mobile/auth.py +++ b/app/schema/request/mobile/auth.py @@ -64,10 +64,12 @@ class MobileLoginRequest(MobileAuthBaseRequest): class RegisterVerifyRequest(MobileAuthBaseRequest): - otp: str = Field(..., min_length=6, max_length=6, description="The 6-digit OTP code sent via email") - - - + otp: str = Field( + ..., + min_length=6, + max_length=6, + description="The 6-digit OTP code sent via email", + ) class ResendOtpRequest(BaseModel): diff --git a/app/schema/request/mobile/notifications.py b/app/schema/request/mobile/notifications.py index f6a94dc3..fdb108bd 100644 --- a/app/schema/request/mobile/notifications.py +++ b/app/schema/request/mobile/notifications.py @@ -4,6 +4,4 @@ class MarkUserNotificationsReadRequest(BaseModel): - notification_ids: list[UUID] = Field( - ..., min_length=1, max_length=100 - ) + notification_ids: list[UUID] = Field(..., min_length=1, max_length=100) diff --git a/app/schema/request/web/auth.py b/app/schema/request/web/auth.py index 5ab4072e..6931b095 100644 --- a/app/schema/request/web/auth.py +++ b/app/schema/request/web/auth.py @@ -4,4 +4,3 @@ class WebAuthRequest(BaseModel): email: EmailStr password: str - diff --git a/app/schema/request/web/event.py b/app/schema/request/web/event.py index c7c05e3d..775e6c8b 100644 --- a/app/schema/request/web/event.py +++ b/app/schema/request/web/event.py @@ -2,11 +2,13 @@ from datetime import datetime from typing import Optional + class EventCreate(BaseModel): name: str event_date: datetime end_date: Optional[datetime] = None status: Optional[str] = "draft" + class JoinEventRequest(BaseModel): event_code: str diff --git a/app/schema/request/web/staff_user.py b/app/schema/request/web/staff_user.py index 7fc36597..82b61f4b 100644 --- a/app/schema/request/web/staff_user.py +++ b/app/schema/request/web/staff_user.py @@ -11,4 +11,3 @@ class StaffUserCreateRequest(BaseModel): class StaffUserUpdateRequest(BaseModel): email: Optional[EmailStr] role: Literal["multi_team_lead", "multi"] - diff --git a/app/schema/response/mobile/audit.py b/app/schema/response/mobile/audit.py index 76983b2d..184bd2a6 100644 --- a/app/schema/response/mobile/audit.py +++ b/app/schema/response/mobile/audit.py @@ -23,6 +23,7 @@ def from_user(cls, user: User) -> "AuditActorSchema": display_name=user.display_name, ) + class AuditEventSchema(BaseModel): id: UUID event_type: AuditEventType diff --git a/app/schema/response/mobile/auth.py b/app/schema/response/mobile/auth.py index d9949374..88377571 100644 --- a/app/schema/response/mobile/auth.py +++ b/app/schema/response/mobile/auth.py @@ -3,12 +3,14 @@ import uuid from datetime import datetime + class DeviceSchema(BaseModel): id: uuid.UUID device_name: str device_type: str totp_secret: str | None + class SessionSchema(BaseModel): session_id: uuid.UUID device_id: uuid.UUID @@ -16,11 +18,13 @@ class SessionSchema(BaseModel): idle_expires_at: datetime absolute_expires_at: datetime + class MobileUserSchema(BaseModel): user_id: uuid.UUID email: str session_id: uuid.UUID + class UserSchema(BaseModel): id: uuid.UUID email: str @@ -28,16 +32,19 @@ class UserSchema(BaseModel): avatar_url: str | None is_onboarded: bool + class MeResponse(BaseModel): user: UserSchema devices: List[DeviceSchema] sessions: Optional[SessionSchema] + class RegisterPendingResponse(BaseModel): message: str status: str email: str + class MobileAuthResponse(BaseModel): access_token: str refresh_token: str diff --git a/app/schema/response/mobile/notifications.py b/app/schema/response/mobile/notifications.py index d347248b..5bfbe5b3 100644 --- a/app/schema/response/mobile/notifications.py +++ b/app/schema/response/mobile/notifications.py @@ -33,4 +33,6 @@ def from_models( cls, notifications: list[Notification], ) -> "UserNotificationListResponse": - return cls(items=[UserNotificationSchema.from_model(item) for item in notifications]) + return cls( + items=[UserNotificationSchema.from_model(item) for item in notifications] + ) diff --git a/app/schema/response/staff/notifications.py b/app/schema/response/staff/notifications.py index ff623820..9d07f367 100644 --- a/app/schema/response/staff/notifications.py +++ b/app/schema/response/staff/notifications.py @@ -33,4 +33,6 @@ def from_models( cls, notifications: list[StaffNotification], ) -> "StaffNotificationListResponse": - return cls(items=[StaffNotificationSchema.from_model(item) for item in notifications]) + return cls( + items=[StaffNotificationSchema.from_model(item) for item in notifications] + ) diff --git a/app/schema/response/staff/upload_groups.py b/app/schema/response/staff/upload_groups.py index f691cc9a..eadf1647 100644 --- a/app/schema/response/staff/upload_groups.py +++ b/app/schema/response/staff/upload_groups.py @@ -40,7 +40,9 @@ def coerce_status(cls, v: object) -> str: return getattr(v, "value", str(v)) @classmethod - def from_details(cls, details: UploadRequestGroupDetails) -> "UploadRequestGroupSchema": + def from_details( + cls, details: UploadRequestGroupDetails + ) -> "UploadRequestGroupSchema": data = dataclasses.asdict(details.group) data["requests"] = [ UploadRequestSchema.from_details(req) for req in details.requests @@ -56,7 +58,9 @@ def from_details_list( cls, details_list: list[UploadRequestGroupDetails], ) -> "UploadRequestGroupListResponse": - return cls(items=[UploadRequestGroupSchema.from_details(d) for d in details_list]) + return cls( + items=[UploadRequestGroupSchema.from_details(d) for d in details_list] + ) class UploadRequestGroupPhotoListResponse(UploadRequestPhotoListResponse): @@ -65,9 +69,8 @@ def from_photos( cls, photos: list[UploadRequestPhoto], ) -> "UploadRequestGroupPhotoListResponse": - return cls( - items=[UploadRequestPhotoSchema.model_validate(p) for p in photos] - ) + return cls(items=[UploadRequestPhotoSchema.model_validate(p) for p in photos]) + class UploadRequestGroupSummarySchema(BaseModel): model_config = ConfigDict(from_attributes=True) @@ -96,5 +99,9 @@ class UploadRequestGroupSummaryListResponse(BaseModel): items: list[UploadRequestGroupSummarySchema] @classmethod - def from_groups(cls, groups: list["UploadRequestGroup"]) -> "UploadRequestGroupSummaryListResponse": - return cls(items=[UploadRequestGroupSummarySchema.model_validate(g) for g in groups]) + def from_groups( + cls, groups: list["UploadRequestGroup"] + ) -> "UploadRequestGroupSummaryListResponse": + return cls( + items=[UploadRequestGroupSummarySchema.model_validate(g) for g in groups] + ) diff --git a/app/schema/response/web/audit.py b/app/schema/response/web/audit.py index 76983b2d..184bd2a6 100644 --- a/app/schema/response/web/audit.py +++ b/app/schema/response/web/audit.py @@ -23,6 +23,7 @@ def from_user(cls, user: User) -> "AuditActorSchema": display_name=user.display_name, ) + class AuditEventSchema(BaseModel): id: UUID event_type: AuditEventType diff --git a/app/schema/response/web/auth.py b/app/schema/response/web/auth.py index cac08e70..f274a9ce 100644 --- a/app/schema/response/web/auth.py +++ b/app/schema/response/web/auth.py @@ -9,6 +9,3 @@ class WebAuthResponse(BaseModel): access_token: str user_id: uuid.UUID role: str - - - diff --git a/app/schema/response/web/event.py b/app/schema/response/web/event.py index 01a64266..6efb0b30 100644 --- a/app/schema/response/web/event.py +++ b/app/schema/response/web/event.py @@ -3,6 +3,7 @@ import uuid from typing import Optional + class EventResponse(BaseModel): model_config = ConfigDict(from_attributes=True) @@ -16,16 +17,20 @@ class EventResponse(BaseModel): created_at: datetime archived_at: Optional[datetime] = None + class ParticipantResponse(BaseModel): """Data for a user who joined an event""" + model_config = ConfigDict(from_attributes=True) user_id: uuid.UUID user_email: str joined_at: datetime + class UserEventResponse(BaseModel): """Data for an event a user has joined""" + model_config = ConfigDict(from_attributes=True) id: uuid.UUID @@ -35,8 +40,10 @@ class UserEventResponse(BaseModel): status: str joined_at: datetime + class JoinEventResponse(BaseModel): """Confirmation of a successful join""" + model_config = ConfigDict(from_attributes=True) id: uuid.UUID diff --git a/app/schema/response/web/stats.py b/app/schema/response/web/stats.py index f3bcdbc4..9a46a45b 100644 --- a/app/schema/response/web/stats.py +++ b/app/schema/response/web/stats.py @@ -2,6 +2,7 @@ from datetime import datetime from typing import List, Optional + class AdminStatsResponse(BaseModel): active_events: int photos_uploaded: int @@ -9,11 +10,13 @@ class AdminStatsResponse(BaseModel): queue_size: int timestamp: datetime + class DriveUsageResponse(BaseModel): used_bytes: int total_bytes: int timestamp: datetime + class AlertItem(BaseModel): id: str type: str @@ -24,11 +27,13 @@ class AlertItem(BaseModel): is_actionable: Optional[bool] = False action_text: Optional[str] = None + class AlertResponse(BaseModel): alerts: List[AlertItem] unread_count: int timestamp: datetime + class ProcessingLoadResponse(BaseModel): completed: float processing: float diff --git a/app/service/device.py b/app/service/device.py index b5840ea7..3cfe0f21 100644 --- a/app/service/device.py +++ b/app/service/device.py @@ -1,13 +1,15 @@ from db.generated import devices as device_queries import uuid -from app.core.exceptions import DBException,AppException, DBExceptionImpl +from app.core.exceptions import DBException, AppException, DBExceptionImpl from db.generated.models import UserDevice class DeviceService: device_querier: device_queries.AsyncQuerier - def init(self: "DeviceService", device_querier: device_queries.AsyncQuerier) -> None: + def init( + self: "DeviceService", device_querier: device_queries.AsyncQuerier + ) -> None: self.device_querier = device_querier async def activate_device( @@ -28,13 +30,10 @@ async def revoke_device( device_id: uuid.UUID, user_id: uuid.UUID, ) -> None: - try : - #here were cacscading the delete of the session no need to handle it - await self.device_querier.revoke_device( - id=device_id, - user_id=user_id - ) - except Exception as e : + try: + # here were cacscading the delete of the session no need to handle it + await self.device_querier.revoke_device(id=device_id, user_id=user_id) + except Exception as e: raise DBException.handle(e) async def update_device_push_token( @@ -61,7 +60,9 @@ async def inactivate_device( user_id: uuid.UUID, ) -> None: try: - device = await self.device_querier.get_device_by_id(id=device_id, user_id=user_id) + device = await self.device_querier.get_device_by_id( + id=device_id, user_id=user_id + ) if device is None or device.user_id != user_id: raise AppException.not_found("Device not found") await self.device_querier.deactivate_device( @@ -71,7 +72,9 @@ async def inactivate_device( except Exception as e: raise DBException.handle(e) - async def get_all_devices(self: "DeviceService", user_id: uuid.UUID) -> tuple[list[UserDevice], int]: + async def get_all_devices( + self: "DeviceService", user_id: uuid.UUID + ) -> tuple[list[UserDevice], int]: devices: list[UserDevice] = [] async for device in self.device_querier.list_user_devices(user_id=user_id): @@ -86,20 +89,21 @@ async def get_device_by_id( device_id: uuid.UUID, user_id: uuid.UUID, ) -> UserDevice: - try : - device = await self.device_querier.get_device_by_id(id=device_id, user_id=user_id) - if device is None : + try: + device = await self.device_querier.get_device_by_id( + id=device_id, user_id=user_id + ) + if device is None: raise AppException.not_found("device not found ") return device - except Exception as e : + except Exception as e: raise DBExceptionImpl.handle(e) async def count_devices(self: "DeviceService", user_id: uuid.UUID) -> int: - try : - count = await self.device_querier.count_user_devices(user_id=user_id) - if count is None : + try: + count = await self.device_querier.count_user_devices(user_id=user_id) + if count is None: raise AppException.internal_error("db failed to count ") return count - except Exception as e : - raise DBExceptionImpl.handle(e) - + except Exception as e: + raise DBExceptionImpl.handle(e) diff --git a/app/service/event.py b/app/service/event.py index fb6165e9..d947f108 100644 --- a/app/service/event.py +++ b/app/service/event.py @@ -2,47 +2,50 @@ import uuid from typing import List, Optional from app.core.exceptions import AppException -from app.schema.request.web.event import ( - EventCreate -) +from app.schema.request.web.event import EventCreate from app.schema.response.web.event import ( EventResponse, JoinEventResponse, UserEventResponse, - ParticipantResponse + ParticipantResponse, ) from db.generated import events as event_queries from db.generated import event_participant as participant_queries from db.generated import models + class EventService: def __init__( self, e_querier: event_queries.AsyncQuerier, - p_querier: participant_queries.AsyncQuerier + p_querier: participant_queries.AsyncQuerier, ): self.e_querier = e_querier self.p_querier = p_querier # --- Core Event Management --- - async def create_event(self, req: EventCreate, creator_id: uuid.UUID) -> EventResponse: - code_created:str = secrets.token_urlsafe(8) + async def create_event( + self, req: EventCreate, creator_id: uuid.UUID + ) -> EventResponse: + code_created: str = secrets.token_urlsafe(8) params = event_queries.CreateEventParams( name=req.name, event_code=code_created, event_date=req.event_date, end_date=req.end_date, status=req.status or "draft", - created_by=creator_id + created_by=creator_id, ) event = await self.e_querier.create_event(params) if not event: raise AppException.internal_error("Failed to create event") return EventResponse.model_validate(event) - async def update_status(self, event_id: uuid.UUID, new_status: str) -> EventResponse: + async def update_status( + self, event_id: uuid.UUID, new_status: str + ) -> EventResponse: event = await self.e_querier.update_event_status(id=event_id, status=new_status) if not event: raise AppException.not_found("Event not found") @@ -58,7 +61,7 @@ async def list_events( self, limit: int = 10, offset: int = 0, - status: Optional[models.EventStatus] = None + status: Optional[models.EventStatus] = None, ) -> List[EventResponse]: params = event_queries.ListEventsParams( @@ -68,7 +71,7 @@ async def list_events( start_date=None, end_date=None, search_name=None, - sort_order='date_desc' + sort_order="date_desc", ) events: List[EventResponse] = [] @@ -84,7 +87,9 @@ async def find_events_by_name(self, name_query: str) -> List[EventResponse]: # --- Participation (Scan to Join) --- - async def join_event_by_code(self, user_id: uuid.UUID, code: str) -> JoinEventResponse: + async def join_event_by_code( + self, user_id: uuid.UUID, code: str + ) -> JoinEventResponse: # 1. Find event by the scanned hash event = await self.e_querier.get_event_by_code(event_code=code) if not event: @@ -94,18 +99,24 @@ async def join_event_by_code(self, user_id: uuid.UUID, code: str) -> JoinEventRe raise AppException.forbidden("This event is already closed.") # 2. Check if already joined - is_member = await self.p_querier.is_user_in_event(event_id=event.id, user_id=user_id) + is_member = await self.p_querier.is_user_in_event( + event_id=event.id, user_id=user_id + ) if is_member: raise AppException.forbidden("You have already joined this event") # 3. Join - join_record = await self.p_querier.join_event(event_id=event.id, user_id=user_id) + join_record = await self.p_querier.join_event( + event_id=event.id, user_id=user_id + ) if not join_record: raise AppException.internal_error("Failed to join event") return JoinEventResponse.model_validate(join_record) - async def get_event_attendees(self, event_id: uuid.UUID) -> List[ParticipantResponse]: + async def get_event_attendees( + self, event_id: uuid.UUID + ) -> List[ParticipantResponse]: # Explicitly hint the list type to resolve Pylance 'Unknown' errors users: List[ParticipantResponse] = [] async for u in self.p_querier.get_event_participants(event_id=event_id): diff --git a/app/service/face_embedding.py b/app/service/face_embedding.py index a571f195..958f76e7 100644 --- a/app/service/face_embedding.py +++ b/app/service/face_embedding.py @@ -5,7 +5,7 @@ from dataclasses import dataclass from typing import List, Literal, Optional, Sequence, Tuple, TypedDict -import cv2 # type: ignore +import cv2 # type: ignore import numpy as np from insightface.app import FaceAnalysis # type: ignore[import-untyped] from app.core.config import settings @@ -101,7 +101,7 @@ def embed(self, image: np.ndarray, bboxes: Sequence[BBox]) -> list[float]: if not faces: raise ValueError("No faces detected by the model") - x1, y1, x2, y2 = bboxes[0] # type: ignore + x1, y1, x2, y2 = bboxes[0] # type: ignore target_cx = (x1 + x2) / 2 target_cy = (y1 + y2) / 2 @@ -154,9 +154,11 @@ async def compute_average_embedding_stream( image = self._decode_image(payload) image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) - # Single detection pass — model.get() already returns embeddings + if self.face_embedding.model is None: + raise RuntimeError("Model is not initialized") faces: list[FaceStub] = await asyncio.to_thread( # type: ignore - self.face_embedding.model.get, image_rgb # type: ignore + self.face_embedding.model.get, + image_rgb, # type: ignore ) if not faces: @@ -189,9 +191,7 @@ async def compute_event_embedding( ) -> dict[str, list[list[float]]]: if not payloads: - raise AppException.bad_request( - "At least one image is required" - ) + raise AppException.bad_request("At least one image is required") results: dict[str, list[list[float]]] = {} @@ -200,8 +200,11 @@ async def compute_event_embedding( image = self._decode_image(payload) image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) - faces: list[FaceStub] = await asyncio.to_thread( # type: ignore - self.face_embedding.model.get, image_rgb # type: ignore + if self.face_embedding.model is None: + raise RuntimeError("Model is not initialized") + faces: list[FaceStub] = await asyncio.to_thread( # type: ignore + self.face_embedding.model.get, + image_rgb, # type: ignore ) results[payload["filename"]] = [ @@ -223,8 +226,11 @@ async def detect_faces( image = self._decode_image(payload) image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) - faces: list[FaceStub] = await asyncio.to_thread( # type: ignore - self.face_embedding.model.get, image_rgb # type: ignore + if self.face_embedding.model is None: + raise RuntimeError("Model is not initialized") + faces: list[FaceStub] = await asyncio.to_thread( # type: ignore + self.face_embedding.model.get, + image_rgb, # type: ignore ) detected: list[DetectedFace] = [] diff --git a/app/service/face_match.py b/app/service/face_match.py index 346cce09..d07654df 100644 --- a/app/service/face_match.py +++ b/app/service/face_match.py @@ -7,7 +7,11 @@ from sqlalchemy.ext.asyncio import AsyncConnection from app.core.logger import logger -from app.schema.internal.single_face_match import BBoxPayload, ClosestUserMatch, SingleFaceMatchJob +from app.schema.internal.single_face_match import ( + BBoxPayload, + ClosestUserMatch, + SingleFaceMatchJob, +) from app.service.user_notification import UserNotificationService from app.service.users import AuthService from db.generated import photo_faces as photo_face_queries @@ -37,7 +41,9 @@ async def process_detected_face( bbox: BBoxPayload | None, ) -> None: # noqa: C901 if not job.image_ref: - logger.warning("Missing image_ref in event payload for photo %s", job.photo_id) + logger.warning( + "Missing image_ref in event payload for photo %s", job.photo_id + ) return embedding_literal = self._vector_literal(embedding) @@ -95,7 +101,8 @@ async def process_detected_face( if created_face_match_id: assert matched_user is not None await self.photo_querier.update_photo_status( - id=job.photo_id, status="approved", + id=job.photo_id, + status="approved", ) await self.user_notification_service.create_notification( user_id=matched_user.user_id, @@ -111,17 +118,27 @@ async def _autoapprove_if_unmatchable( matched_user: ClosestUserMatch | None, ) -> bool: if matched_user is None: - logger.info("No user embeddings available for matching, auto-approving photo %s", job.photo_id) - await self.photo_querier.update_photo_status(id=job.photo_id, status="approved") + logger.info( + "No user embeddings available for matching, auto-approving photo %s", + job.photo_id, + ) + await self.photo_querier.update_photo_status( + id=job.photo_id, status="approved" + ) return True from app.worker.photo_worker.settings import settings as worker_settings + if matched_user.distance > worker_settings.similarity_threshold: logger.info( "Closest user distance %.4f exceeds threshold %.4f for photo %s; auto-approving", - matched_user.distance, worker_settings.similarity_threshold, job.photo_id, + matched_user.distance, + worker_settings.similarity_threshold, + job.photo_id, + ) + await self.photo_querier.update_photo_status( + id=job.photo_id, status="approved" ) - await self.photo_querier.update_photo_status(id=job.photo_id, status="approved") return True return False @@ -144,6 +161,4 @@ def _vector_literal(embedding: list[float]) -> str: def _serialize_bbox(bbox: BBoxPayload | None) -> str | None: if bbox is None: return None - return json.dumps( - {"x1": bbox.x1, "y1": bbox.y1, "x2": bbox.x2, "y2": bbox.y2} - ) + return json.dumps({"x1": bbox.x1, "y1": bbox.y1, "x2": bbox.x2, "y2": bbox.y2}) diff --git a/app/service/photo_approval.py b/app/service/photo_approval.py index f3eb809c..a121fd6a 100644 --- a/app/service/photo_approval.py +++ b/app/service/photo_approval.py @@ -48,12 +48,16 @@ async def decide( ) approvals = [] - async for a in self._approval_querier.get_photo_approvals_by_photo_id(photo_id=photo_id): + async for a in self._approval_querier.get_photo_approvals_by_photo_id( + photo_id=photo_id + ): approvals.append(a) rejected = [a for a in approvals if a.decision == "rejected"] if rejected: - await self._photo_querier.update_photo_status(id=photo_id, status="rejected") + await self._photo_querier.update_photo_status( + id=photo_id, status="rejected" + ) await self._delete_photo_storage(photo_id) return "rejected" @@ -66,7 +70,9 @@ async def decide( async def expire_stale(self, timeout_days: int) -> int: count = 0 - async for _ in self._approval_querier.expire_stale_approvals(timeout_days=timeout_days): + async for _ in self._approval_querier.expire_stale_approvals( + timeout_days=timeout_days + ): count += 1 if count: logger.info("Auto-expired %d stale pending photo(s)", count) diff --git a/app/service/staff_drive.py b/app/service/staff_drive.py index b80a98be..7cfde00e 100644 --- a/app/service/staff_drive.py +++ b/app/service/staff_drive.py @@ -110,7 +110,9 @@ async def handle_callback( if redirect_url is not None and not isinstance(redirect_url, str): raise AppException.bad_request("Invalid OAuth redirect URL") - staff_user = await self.staff_user_querier.get_staff_user_by_id(id=staff_user_id) + staff_user = await self.staff_user_querier.get_staff_user_by_id( + id=staff_user_id + ) if staff_user is None: raise AppException.not_found("Staff user not found") @@ -158,7 +160,10 @@ async def get_active_connection_or_raise( def _token_needs_refresh(cls, connection: StaffDriveConnection) -> bool: if connection.token_expires_at is None: return False - return connection.token_expires_at <= datetime.now(timezone.utc) + cls.TOKEN_REFRESH_BUFFER + return ( + connection.token_expires_at + <= datetime.now(timezone.utc) + cls.TOKEN_REFRESH_BUFFER + ) async def _refresh_connection_access_token( self, @@ -177,20 +182,24 @@ async def _refresh_connection_access_token( if token.refresh_token: encrypted_refresh_token = self._encrypt(token.refresh_token) - refreshed_connection = await self.drive_connection_querier.upsert_staff_drive_connection( - arg=drive_queries.UpsertStaffDriveConnectionParams( - staff_user_id=connection.staff_user_id, - provider=connection.provider, - google_email=connection.google_email, - google_account_id=connection.google_account_id, - access_token=encrypted_access_token, - refresh_token=encrypted_refresh_token, - token_expires_at=token.expires_at, - scopes=token.scope, + refreshed_connection = ( + await self.drive_connection_querier.upsert_staff_drive_connection( + arg=drive_queries.UpsertStaffDriveConnectionParams( + staff_user_id=connection.staff_user_id, + provider=connection.provider, + google_email=connection.google_email, + google_account_id=connection.google_account_id, + access_token=encrypted_access_token, + refresh_token=encrypted_refresh_token, + token_expires_at=token.expires_at, + scopes=token.scope, + ) ) ) if refreshed_connection is None: - raise AppException.internal_error("Failed to refresh Google Drive connection") + raise AppException.internal_error( + "Failed to refresh Google Drive connection" + ) return refreshed_connection @@ -202,7 +211,9 @@ async def get_access_token_for_staff_user(self, staff_user_id: uuid.UUID) -> str async def get_system_access_token(self) -> str: """Get an access token from any active staff Drive connection.""" - connection = await self.drive_connection_querier.get_any_active_staff_drive_connection() + connection = ( + await self.drive_connection_querier.get_any_active_staff_drive_connection() + ) if connection is None: raise AppException.not_found("No active Google Drive connection") if self._token_needs_refresh(connection): @@ -243,9 +254,13 @@ def _encrypt(self, raw_value: str) -> str: def decrypt(self, encrypted_value: str) -> str: try: - return self._fernet().decrypt(encrypted_value.encode("utf-8")).decode("utf-8") + return ( + self._fernet().decrypt(encrypted_value.encode("utf-8")).decode("utf-8") + ) except InvalidToken as exc: - raise AppException.internal_error("Stored Google Drive token cannot be decrypted") from exc + raise AppException.internal_error( + "Stored Google Drive token cannot be decrypted" + ) from exc async def import_images_from_drive( self, @@ -318,7 +333,7 @@ def _fernet(self) -> Fernet: @staticmethod def _generate_object_name(filename: str) -> str: - suffix = filename[filename.rfind("."):] if "." in filename else "" + suffix = filename[filename.rfind(".") :] if "." in filename else "" return f"{uuid.uuid4()}{suffix}" @staticmethod @@ -343,4 +358,6 @@ def build_frontend_callback_url( query.append(("google_email", google_email)) if error is not None: query.append(("error", error)) - return urllib.parse.urlunparse(parsed._replace(query=urllib.parse.urlencode(query))) + return urllib.parse.urlunparse( + parsed._replace(query=urllib.parse.urlencode(query)) + ) diff --git a/app/service/staff_notifications.py b/app/service/staff_notifications.py index ab26bc4a..e2702a33 100644 --- a/app/service/staff_notifications.py +++ b/app/service/staff_notifications.py @@ -35,7 +35,9 @@ async def list_notifications( staff_user_id: uuid.UUID, ) -> list[StaffNotification]: notifications: list[StaffNotification] = [] - async for notification in self.notification_querier.list_staff_notifications_by_staff_user_id( + async for ( + notification + ) in self.notification_querier.list_staff_notifications_by_staff_user_id( staff_user_id=staff_user_id ): notifications.append(notification) diff --git a/app/service/staff_user.py b/app/service/staff_user.py index 56756685..b72cd898 100644 --- a/app/service/staff_user.py +++ b/app/service/staff_user.py @@ -1,4 +1,3 @@ - from sqlalchemy.exc import SQLAlchemyError from app.core.logger import logger @@ -6,7 +5,7 @@ import uuid from app.core.exceptions import AppException, DBException, DBExceptionImpl -from app.core.securite import create_access_staff_token, hash_password, verify_password +from app.core.securite import create_access_staff_token, hash_password, verify_password from app.schema.response.web.auth import WebAuthResponse from db.generated import staff_user as staff_queries from db.generated.staff_user import ListStaffUsersParams @@ -46,7 +45,6 @@ async def create_staff_user( logger.error("Failed to create staff user: %s", exc) raise DBException.handle(exc) - async def update_staff_user( self, *, id: uuid.UUID, email: Optional[str], role: StaffRole ) -> StaffUser: @@ -88,7 +86,7 @@ async def list_staff_users( params = ListStaffUsersParams( column_1=normalized_search, - column_2=role.value if role is not None else None , + column_2=role.value if role is not None else None, column_3=sort_by, column_4=sort_direction, limit=limit, @@ -103,22 +101,21 @@ async def list_staff_users( logger.error("Failed to list staff users: %s", exc) raise DBException.handle(exc) - async def admin_login( self, email: str, password: str, ) -> WebAuthResponse: normalized_email = email.strip().lower() - staff: StaffUser | None = await self.staff_user_querier.get_staff_user_by_email(email=normalized_email) + staff: StaffUser | None = await self.staff_user_querier.get_staff_user_by_email( + email=normalized_email + ) if staff is None or not verify_password(password, staff.password): logger.info("admin login failed for email %s", normalized_email) raise AppException.unauthorized("Invalid email or password") - access_token = create_access_staff_token( - staff_id=str(staff.id), - role=staff.role + staff_id=str(staff.id), role=staff.role ) return WebAuthResponse( @@ -127,11 +124,10 @@ async def admin_login( role=staff.role, ) - async def get_staff_user( - self, - stuff_id:uuid.UUID - )->StaffUser: - stuff:StaffUser|None = await self.staff_user_querier.get_staff_user_by_id(id=stuff_id) + async def get_staff_user(self, stuff_id: uuid.UUID) -> StaffUser: + stuff: StaffUser | None = await self.staff_user_querier.get_staff_user_by_id( + id=stuff_id + ) if stuff is None: raise AppException.not_found("user not found ") return stuff diff --git a/app/service/staged_upload_storage.py b/app/service/staged_upload_storage.py index e8f9d97a..53d14c48 100644 --- a/app/service/staged_upload_storage.py +++ b/app/service/staged_upload_storage.py @@ -95,7 +95,9 @@ async def delete_storage_key(self, storage_key: str) -> None: try: await self.bucket.delete(storage_key) except Exception as exc: - raise AppException.storage_error("Failed to delete staged image from storage") from exc + raise AppException.storage_error( + "Failed to delete staged image from storage" + ) from exc async def get_preview(self, storage_key: str) -> PreviewObject: data, file_name, content_type = await self.bucket.get(storage_key) @@ -114,7 +116,9 @@ async def create_presigned_staging_upload( photo_id=photo_id, file_name=file_name, ) - url = await self.bucket.presigned_put_url(storage_key, expires_seconds=expires_seconds) + url = await self.bucket.presigned_put_url( + storage_key, expires_seconds=expires_seconds + ) return storage_key, url async def stat_staging_object(self, storage_key: str) -> ObjectStat | None: diff --git a/app/service/stats.py b/app/service/stats.py index b1b9f740..c3ea66bf 100644 --- a/app/service/stats.py +++ b/app/service/stats.py @@ -2,13 +2,17 @@ import uuid from typing import TYPE_CHECKING from app.schema.response.web.stats import ( - AdminStatsResponse, DriveUsageResponse, - ProcessingLoadResponse, AlertResponse, AlertItem + AdminStatsResponse, + DriveUsageResponse, + ProcessingLoadResponse, + AlertResponse, + AlertItem, ) if TYPE_CHECKING: from db.generated.stats import AsyncQuerier + class StatsService: def __init__(self, querier: "AsyncQuerier"): self.q = querier @@ -23,7 +27,7 @@ async def get_dashboard_stats(self) -> AdminStatsResponse: photos_uploaded=photos or 0, processed_photos=metrics.completed_count if metrics else 0, queue_size=metrics.pending_count if metrics else 0, - timestamp=datetime.now(timezone.utc) + timestamp=datetime.now(timezone.utc), ) async def get_processing_load(self) -> ProcessingLoadResponse: @@ -39,7 +43,7 @@ async def get_processing_load(self) -> ProcessingLoadResponse: return ProcessingLoadResponse( completed=round((metrics.completed_count / total) * 100, 1), processing=round((metrics.running_count / total) * 100, 1), - queued=round((metrics.pending_count / total) * 100, 1) + queued=round((metrics.pending_count / total) * 100, 1), ) async def get_storage_usage(self) -> DriveUsageResponse: @@ -50,28 +54,34 @@ async def get_storage_usage(self) -> DriveUsageResponse: return DriveUsageResponse( used_bytes=used_bytes or 0, total_bytes=total_bytes, - timestamp=datetime.now(timezone.utc) + timestamp=datetime.now(timezone.utc), ) async def get_staff_alerts(self, staff_id: uuid.UUID) -> AlertResponse: - db_alerts = [a async for a in self.q.get_recent_staff_alerts(staff_user_id=staff_id)] - unread_count = await self.q.get_unread_staff_alerts_count(staff_user_id=staff_id) + db_alerts = [ + a async for a in self.q.get_recent_staff_alerts(staff_user_id=staff_id) + ] + unread_count = await self.q.get_unread_staff_alerts_count( + staff_user_id=staff_id + ) alerts = [] for a in db_alerts: # Assuming payload is a dict with title and message payload = a.payload or {} - alerts.append(AlertItem( - id=str(a.id), - type=a.type, - title=payload.get("title", "Notification"), - message=payload.get("message", "No message provided"), - created_at=a.created_at, - is_read=a.read_at is not None - )) + alerts.append( + AlertItem( + id=str(a.id), + type=a.type, + title=payload.get("title", "Notification"), + message=payload.get("message", "No message provided"), + created_at=a.created_at, + is_read=a.read_at is not None, + ) + ) return AlertResponse( alerts=alerts, unread_count=unread_count or 0, - timestamp=datetime.now(timezone.utc) + timestamp=datetime.now(timezone.utc), ) diff --git a/app/service/upload_requests.py b/app/service/upload_requests.py index e63de06b..3aebde48 100644 --- a/app/service/upload_requests.py +++ b/app/service/upload_requests.py @@ -102,19 +102,29 @@ def _raise_integrity_error(exc: IntegrityError) -> None: if sqlstate == "23503": raise AppException.bad_request("Invalid event reference") from exc if sqlstate == "23505": - raise AppException.conflict("Duplicate photo in upload request batch") from exc + raise AppException.conflict( + "Duplicate photo in upload request batch" + ) from exc raise AppException.internal_error("Failed to persist upload request") from exc - def _validate_downloaded_photo(self, downloaded_photo: GoogleDriveFileDownload) -> None: + def _validate_downloaded_photo( + self, downloaded_photo: GoogleDriveFileDownload + ) -> None: metadata = downloaded_photo.metadata if metadata.mime_type not in self._allowed_mime_types: - raise AppException.image_format_error("Unsupported image format from Google Drive") + raise AppException.image_format_error( + "Unsupported image format from Google Drive" + ) if metadata.size_bytes <= 0 or metadata.size_bytes > self._max_photo_size_bytes: - raise AppException.bad_request("Google Drive image exceeds maximum allowed size") + raise AppException.bad_request( + "Google Drive image exceeds maximum allowed size" + ) def _is_supported_image(self, metadata: GoogleDriveFileMetadata) -> bool: - return metadata.mime_type in self._allowed_mime_types and metadata.size_bytes > 0 + return ( + metadata.mime_type in self._allowed_mime_types and metadata.size_bytes > 0 + ) @staticmethod def _validate_create_request_inputs(photos: Sequence[UploadPhotoInput]) -> None: @@ -125,12 +135,18 @@ def _validate_create_request_inputs(photos: Sequence[UploadPhotoInput]) -> None: drive_file_ids = [photo.drive_file_id for photo in photos] if len(drive_file_ids) != len(set(drive_file_ids)): - raise AppException.conflict("Duplicate drive_file_id found in upload request batch") + raise AppException.conflict( + "Duplicate drive_file_id found in upload request batch" + ) - async def _cleanup_created_photos(self, created_photos: Sequence[UploadRequestPhoto]) -> None: + async def _cleanup_created_photos( + self, created_photos: Sequence[UploadRequestPhoto] + ) -> None: for created_photo in created_photos: try: - await self.staged_upload_storage.delete_storage_key(created_photo.staging_storage_key) + await self.staged_upload_storage.delete_storage_key( + created_photo.staging_storage_key + ) except Exception: logger.warning( "Failed to clean staged object %s after create failure", @@ -146,7 +162,9 @@ async def _cleanup_created_group( ) -> None: for request_details in reversed(created_requests): try: - await self.upload_request_querier.delete_upload_request(id=request_details.request.id) + await self.upload_request_querier.delete_upload_request( + id=request_details.request.id + ) except Exception as exc: logger.warning( "Failed to delete upload request %s during group cleanup: %s", @@ -158,7 +176,9 @@ async def _cleanup_created_group( return try: - await self.upload_request_group_querier.delete_upload_request_group(id=upload_group_id) + await self.upload_request_group_querier.delete_upload_request_group( + id=upload_group_id + ) except Exception as exc: logger.warning( "Failed to delete upload request group %s during cleanup: %s", @@ -182,7 +202,9 @@ async def _delete_staging_objects_best_effort( ) -> None: for staged_photo in staged_photos: try: - await self.staged_upload_storage.delete_storage_key(staged_photo.staging_storage_key) + await self.staged_upload_storage.delete_storage_key( + staged_photo.staging_storage_key + ) except Exception as exc: logger.warning( "Failed to delete staging object %s: %s", @@ -194,7 +216,9 @@ async def _list_request_photos_by_request_ids( self, request_ids: Sequence[uuid.UUID], ) -> dict[uuid.UUID, list[UploadRequestPhoto]]: - photos_by_request_id: dict[uuid.UUID, list[UploadRequestPhoto]] = defaultdict(list) + photos_by_request_id: dict[uuid.UUID, list[UploadRequestPhoto]] = defaultdict( + list + ) if not request_ids: return photos_by_request_id @@ -227,23 +251,27 @@ async def _create_staged_photo( ) try: - created_photo = await self.upload_request_photo_querier.create_upload_request_photo( - upload_request_photo_queries.CreateUploadRequestPhotoParams( - upload_request_id=upload_request_id, - drive_file_id=photo.drive_file_id, - file_name=downloaded_photo.metadata.name, - mime_type=downloaded_photo.metadata.mime_type, - size_bytes=downloaded_photo.metadata.size_bytes, - staging_storage_key=stored_object.storage_key, - taken_at=photo.taken_at, - day_number=photo.day_number, - visibility=photo.visibility, - status="staged", + created_photo = ( + await self.upload_request_photo_querier.create_upload_request_photo( + upload_request_photo_queries.CreateUploadRequestPhotoParams( + upload_request_id=upload_request_id, + drive_file_id=photo.drive_file_id, + file_name=downloaded_photo.metadata.name, + mime_type=downloaded_photo.metadata.mime_type, + size_bytes=downloaded_photo.metadata.size_bytes, + staging_storage_key=stored_object.storage_key, + taken_at=photo.taken_at, + day_number=photo.day_number, + visibility=photo.visibility, + status="staged", + ) ) ) except IntegrityError: try: - await self.staged_upload_storage.delete_storage_key(stored_object.storage_key) + await self.staged_upload_storage.delete_storage_key( + stored_object.storage_key + ) except Exception: logger.warning( "Failed to clean staged object %s after photo insert conflict", @@ -253,7 +281,9 @@ async def _create_staged_photo( if created_photo is None: try: - await self.staged_upload_storage.delete_storage_key(stored_object.storage_key) + await self.staged_upload_storage.delete_storage_key( + stored_object.storage_key + ) except Exception: logger.warning( "Failed to clean staged object %s after empty photo insert result", @@ -328,7 +358,9 @@ async def _approve_request_without_side_effects( request_id: uuid.UUID, approved_by: StaffUser, ) -> tuple[UploadRequest, list[UploadRequestPhoto], list[str], list[Photo]]: - existing = await self.upload_request_querier.get_upload_request_by_id(id=request_id) + existing = await self.upload_request_querier.get_upload_request_by_id( + id=request_id + ) if existing is None: raise AppException.not_found("Upload request not found") if self._status_value(existing.status) != "pending": @@ -336,10 +368,13 @@ async def _approve_request_without_side_effects( staged_photos = await self.list_request_photos(request_id) if not staged_photos: - raise AppException.bad_request("No staged photos found for this upload request") + raise AppException.bad_request( + "No staged photos found for this upload request" + ) not_transferred = [ - p for p in staged_photos + p + for p in staged_photos if getattr(p, "transfer_status", "uploaded") != "uploaded" ] if not_transferred: @@ -378,7 +413,9 @@ async def _approve_request_without_side_effects( final_storage_key=final_storage_key, ) if updated_photo is None: - raise AppException.internal_error("Failed to update staged photo approval state") + raise AppException.internal_error( + "Failed to update staged photo approval state" + ) upload_request = await self.upload_request_querier.approve_upload_request( id=request_id, @@ -399,7 +436,9 @@ async def _reject_request_without_side_effects( approved_by: StaffUser, reason: str | None, ) -> tuple[UploadRequest, list[UploadRequestPhoto], list[UploadRequestPhoto]]: - existing = await self.upload_request_querier.get_upload_request_by_id(id=request_id) + existing = await self.upload_request_querier.get_upload_request_by_id( + id=request_id + ) if existing is None: raise AppException.not_found("Upload request not found") if self._status_value(existing.status) != "pending": @@ -433,7 +472,9 @@ def _ensure_request_access( return if self._role_value(current_staff_user.role) == StaffRole.MULTI_TEAM_LEAD.value: return - raise AppException.forbidden("You are not allowed to access this upload request") + raise AppException.forbidden( + "You are not allowed to access this upload request" + ) def _ensure_group_access( self, @@ -445,7 +486,9 @@ def _ensure_group_access( return if self._role_value(current_staff_user.role) == StaffRole.MULTI_TEAM_LEAD.value: return - raise AppException.forbidden("You are not allowed to access this upload request group") + raise AppException.forbidden( + "You are not allowed to access this upload request group" + ) def _ensure_group_is_pending( self, @@ -459,7 +502,9 @@ def _ensure_group_import_completed( group: UploadRequestGroup, ) -> None: if group.processing_status != "completed": - raise AppException.bad_request("Upload request group import is not completed") + raise AppException.bad_request( + "Upload request group import is not completed" + ) def _ensure_all_requests_are_pending( self, @@ -483,7 +528,9 @@ async def _publish_event( try: await NatsClient.js_publish(subject, json.dumps(payload).encode("utf-8")) except Exception as exc: - logger.warning("Failed to publish upload request event %s: %s", subject.value, exc) + logger.warning( + "Failed to publish upload request event %s: %s", subject.value, exc + ) async def _audit(self, event_type: AuditEventType, **metadata: object) -> None: if self.audit_service is not None: @@ -606,15 +653,17 @@ async def create_group_from_folder( ) -> UploadRequestGroupDetails: await self.staff_drive_service.get_access_token_for_staff_user(requested_by.id) try: - upload_group = await self.upload_request_group_querier.create_upload_request_group( - upload_request_group_queries.CreateUploadRequestGroupParams( - event_id=event_id, - folder_id=folder_id, - requested_by=requested_by.id, - total_photo_count=0, - batch_count=0, - source="drive", - processing_status="pending", + upload_group = ( + await self.upload_request_group_querier.create_upload_request_group( + upload_request_group_queries.CreateUploadRequestGroupParams( + event_id=event_id, + folder_id=folder_id, + requested_by=requested_by.id, + total_photo_count=0, + batch_count=0, + source="drive", + processing_status="pending", + ) ) ) except IntegrityError as exc: @@ -643,15 +692,17 @@ async def create_direct_group( requested_by: StaffUser, ) -> UploadRequestGroup: try: - upload_group = await self.upload_request_group_querier.create_upload_request_group( - upload_request_group_queries.CreateUploadRequestGroupParams( - event_id=event_id, - folder_id=None, - requested_by=requested_by.id, - total_photo_count=0, - batch_count=0, - source="direct", - processing_status="completed", + upload_group = ( + await self.upload_request_group_querier.create_upload_request_group( + upload_request_group_queries.CreateUploadRequestGroupParams( + event_id=event_id, + folder_id=None, + requested_by=requested_by.id, + total_photo_count=0, + batch_count=0, + source="direct", + processing_status="completed", + ) ) ) except IntegrityError as exc: @@ -675,11 +726,17 @@ async def register_direct_batch( ) for file in files: if file.mime_type not in self._allowed_mime_types: - raise AppException.image_format_error(f"Unsupported image format: {file.mime_type}") + raise AppException.image_format_error( + f"Unsupported image format: {file.mime_type}" + ) if file.size_bytes <= 0 or file.size_bytes > self._max_photo_size_bytes: - raise AppException.bad_request(f"{file.file_name} exceeds maximum allowed size") + raise AppException.bad_request( + f"{file.file_name} exceeds maximum allowed size" + ) - group = await self.upload_request_group_querier.get_upload_request_group_by_id(id=group_id) + group = await self.upload_request_group_querier.get_upload_request_group_by_id( + id=group_id + ) if group is None: raise AppException.not_found("Upload group not found") self._ensure_group_access(current_staff_user=requested_by, upload_group=group) @@ -700,7 +757,10 @@ async def register_direct_batch( results: list[tuple[UploadRequestPhoto, str]] = [] for file in files: photo_id = uuid.uuid4() - storage_key, presigned_url = await self.staged_upload_storage.create_presigned_staging_upload( + ( + storage_key, + presigned_url, + ) = await self.staged_upload_storage.create_presigned_staging_upload( upload_request_id=upload_request.id, photo_id=photo_id, file_name=file.file_name, @@ -723,7 +783,8 @@ async def register_direct_batch( results.append((created_photo, presigned_url)) await self.upload_request_group_querier.increment_upload_request_group_counts( - id=group_id, total_photo_count=len(files), + id=group_id, + total_photo_count=len(files), ) return results @@ -734,13 +795,19 @@ async def confirm_direct_upload( photo_id: uuid.UUID, requested_by: StaffUser, ) -> UploadRequestPhoto: - photo = await self.upload_request_photo_querier.get_upload_request_photo_by_id(id=photo_id) + photo = await self.upload_request_photo_querier.get_upload_request_photo_by_id( + id=photo_id + ) if photo is None: raise AppException.not_found("Upload photo not found") - stat = await self.staged_upload_storage.stat_staging_object(photo.staging_storage_key) + stat = await self.staged_upload_storage.stat_staging_object( + photo.staging_storage_key + ) if stat is None: - failed = await self.upload_request_photo_querier.fail_upload_request_photo_transfer(id=photo_id) + failed = await self.upload_request_photo_querier.fail_upload_request_photo_transfer( + id=photo_id + ) if failed is None: raise AppException.internal_error("Failed to mark upload as failed") raise AppException.bad_request( @@ -769,7 +836,9 @@ async def _maybe_auto_approve_group( upload_request_id: uuid.UUID, approved_by: StaffUser, ) -> None: - upload_request = await self.upload_request_querier.get_upload_request_by_id(id=upload_request_id) + upload_request = await self.upload_request_querier.get_upload_request_by_id( + id=upload_request_id + ) if upload_request is None or upload_request.group_id is None: return if self._status_value(upload_request.status) != "pending": @@ -777,7 +846,9 @@ async def _maybe_auto_approve_group( group_id = upload_request.group_id request_ids: list[uuid.UUID] = [] - async for req in self.upload_request_querier.list_upload_requests_by_group_id(group_id=group_id): + async for req in self.upload_request_querier.list_upload_requests_by_group_id( + group_id=group_id + ): if self._status_value(req.status) != "pending": continue request_ids.append(req.id) @@ -806,11 +877,17 @@ async def fail_direct_upload( photo_id: uuid.UUID, requested_by: StaffUser, ) -> UploadRequestPhoto: - photo = await self.upload_request_photo_querier.get_upload_request_photo_by_id(id=photo_id) + photo = await self.upload_request_photo_querier.get_upload_request_photo_by_id( + id=photo_id + ) if photo is None: raise AppException.not_found("Upload photo not found") - failed = await self.upload_request_photo_querier.fail_upload_request_photo_transfer(id=photo_id) + failed = ( + await self.upload_request_photo_querier.fail_upload_request_photo_transfer( + id=photo_id + ) + ) if failed is None: raise AppException.internal_error("Failed to mark upload as failed") return failed @@ -821,22 +898,32 @@ async def resume_direct_group( group_id: uuid.UUID, requested_by: StaffUser, ) -> list[tuple[UploadRequestPhoto, str]]: - group = await self.upload_request_group_querier.get_upload_request_group_by_id(id=group_id) + group = await self.upload_request_group_querier.get_upload_request_group_by_id( + id=group_id + ) if group is None: raise AppException.not_found("Upload group not found") self._ensure_group_access(current_staff_user=requested_by, upload_group=group) request_ids: list[uuid.UUID] = [] - async for req in self.upload_request_querier.list_upload_requests_by_group_id(group_id=group_id): + async for req in self.upload_request_querier.list_upload_requests_by_group_id( + group_id=group_id + ): request_ids.append(req.id) results: list[tuple[UploadRequestPhoto, str]] = [] async for photo in self.upload_request_photo_querier.list_upload_request_photos_by_upload_request_ids( dollar_1=request_ids ): - if getattr(photo, "transfer_status", "uploaded") not in ("pending_upload", "failed"): + if getattr(photo, "transfer_status", "uploaded") not in ( + "pending_upload", + "failed", + ): continue - storage_key, presigned_url = await self.staged_upload_storage.create_presigned_staging_upload( + ( + storage_key, + presigned_url, + ) = await self.staged_upload_storage.create_presigned_staging_upload( upload_request_id=photo.upload_request_id, photo_id=photo.id, file_name=photo.file_name, @@ -851,7 +938,7 @@ async def resume_direct_group( return results - async def process_group_import( + async def process_group_import( # noqa: C901 self, *, group_id: uuid.UUID, @@ -862,8 +949,10 @@ async def process_group_import( id=group_id ) if upload_group is None: - existing_group = await self.upload_request_group_querier.get_upload_request_group_by_id( - id=group_id + existing_group = ( + await self.upload_request_group_querier.get_upload_request_group_by_id( + id=group_id + ) ) if existing_group is None: logger.warning("Upload request group %s not found for import", group_id) @@ -877,8 +966,10 @@ async def process_group_import( ) return None - requested_by = await self.staff_drive_service.staff_user_querier.get_staff_user_by_id( - id=upload_group.requested_by + requested_by = ( + await self.staff_drive_service.staff_user_querier.get_staff_user_by_id( + id=upload_group.requested_by + ) ) if requested_by is None: await self._mark_group_import_failed( @@ -894,14 +985,20 @@ async def process_group_import( created_requests: list[UploadRequestDetails] = [] photo_inputs: list[UploadPhotoInput] = [] try: - access_token = await self.staff_drive_service.get_access_token_for_staff_user( - requested_by.id + access_token = ( + await self.staff_drive_service.get_access_token_for_staff_user( + requested_by.id + ) ) + if upload_group.folder_id is None: + raise AppException.bad_request("Upload group has no folder_id") folder_files = await GoogleDriveClient.list_folder_files( access_token=access_token, folder_id=upload_group.folder_id, ) - folder_files = sorted(folder_files, key=lambda file: (file.name.lower(), file.id)) + folder_files = sorted( + folder_files, key=lambda file: (file.name.lower(), file.id) + ) photo_inputs = [ UploadPhotoInput( drive_file_id=file.id, @@ -926,7 +1023,9 @@ async def process_group_import( current_staff_user=requested_by, ) - photo_batches = self._chunk_photo_inputs(photo_inputs, self._max_request_batch_size) + photo_batches = self._chunk_photo_inputs( + photo_inputs, self._max_request_batch_size + ) await self.upload_request_group_querier.update_upload_request_group_import_progress( upload_request_group_queries.UpdateUploadRequestGroupImportProgressParams( id=group_id, @@ -969,7 +1068,9 @@ async def process_group_import( ) ) if completed_group is None: - raise AppException.internal_error("Failed to complete upload group import") + raise AppException.internal_error( + "Failed to complete upload group import" + ) for request_details in created_requests: await self._publish_event( @@ -993,7 +1094,9 @@ async def process_group_import( "batch_count": completed_group.batch_count, }, ) - return UploadRequestGroupDetails(group=completed_group, requests=created_requests) + return UploadRequestGroupDetails( + group=completed_group, requests=created_requests + ) except Exception as exc: created_photos = [ photo @@ -1009,7 +1112,9 @@ async def process_group_import( await self._mark_group_import_failed( group_id=group_id, total_photo_count=len(photo_inputs), - batch_count=len(self._chunk_photo_inputs(photo_inputs, self._max_request_batch_size)) + batch_count=len( + self._chunk_photo_inputs(photo_inputs, self._max_request_batch_size) + ) if photo_inputs else 0, processed_photo_count=0, @@ -1025,7 +1130,9 @@ async def get_request_details( request_id: uuid.UUID, current_staff_user: StaffUser, ) -> UploadRequestDetails: - upload_request = await self.upload_request_querier.get_upload_request_by_id(id=request_id) + upload_request = await self.upload_request_querier.get_upload_request_by_id( + id=request_id + ) if upload_request is None: raise AppException.not_found("Upload request not found") self._ensure_request_access( @@ -1044,14 +1151,18 @@ async def get_request_photo_preview( photo_id: uuid.UUID, current_staff_user: StaffUser, ) -> PreviewObject: - upload_request = await self.upload_request_querier.get_upload_request_by_id(id=request_id) + upload_request = await self.upload_request_querier.get_upload_request_by_id( + id=request_id + ) if upload_request is None: raise AppException.not_found("Upload request not found") self._ensure_request_access( current_staff_user=current_staff_user, upload_request=upload_request, ) - photo = await self.upload_request_photo_querier.get_upload_request_photo_by_id(id=photo_id) + photo = await self.upload_request_photo_querier.get_upload_request_photo_by_id( + id=photo_id + ) if photo is None or photo.upload_request_id != request_id: raise AppException.not_found("Upload request photo not found") storage_key = photo.final_storage_key or photo.staging_storage_key @@ -1064,7 +1175,11 @@ async def list_requests( scope: Literal["my", "all"], status: str | None, ) -> list[UploadRequestDetails]: - if scope == "all" and self._role_value(current_staff_user.role) != StaffRole.MULTI_TEAM_LEAD.value: + if ( + scope == "all" + and self._role_value(current_staff_user.role) + != StaffRole.MULTI_TEAM_LEAD.value + ): raise AppException.forbidden("Multi team lead access required") requested_by = current_staff_user.id if scope == "my" else None @@ -1079,7 +1194,9 @@ async def list_requests( requested_by=requested_by ) elif status is not None: - iterator = self.upload_request_querier.list_upload_requests_by_status(status=status) + iterator = self.upload_request_querier.list_upload_requests_by_status( + status=status + ) else: iterator = self.upload_request_querier.list_upload_requests() @@ -1114,7 +1231,9 @@ async def get_group_details( group_id: uuid.UUID, current_staff_user: StaffUser, ) -> UploadRequestGroupDetails: - group = await self.upload_request_group_querier.get_upload_request_group_by_id(id=group_id) + group = await self.upload_request_group_querier.get_upload_request_group_by_id( + id=group_id + ) if group is None: raise AppException.not_found("Upload request group not found") self._ensure_group_access( @@ -1123,7 +1242,9 @@ async def get_group_details( ) requests: list[UploadRequest] = [] - async for upload_request in self.upload_request_querier.list_upload_requests_by_group_id( + async for ( + upload_request + ) in self.upload_request_querier.list_upload_requests_by_group_id( group_id=group_id ): requests.append(upload_request) @@ -1149,7 +1270,11 @@ async def list_groups( scope: Literal["my", "all"], status: str | None, ) -> list[UploadRequestGroup]: - if scope == "all" and self._role_value(current_staff_user.role) != StaffRole.MULTI_TEAM_LEAD.value: + if ( + scope == "all" + and self._role_value(current_staff_user.role) + != StaffRole.MULTI_TEAM_LEAD.value + ): raise AppException.forbidden("Multi team lead access required") requested_by = current_staff_user.id if scope == "my" else None @@ -1164,8 +1289,10 @@ async def list_groups( requested_by=requested_by ) elif status is not None: - iterator = self.upload_request_group_querier.list_upload_request_groups_by_status( - status=status + iterator = ( + self.upload_request_group_querier.list_upload_request_groups_by_status( + status=status + ) ) else: iterator = self.upload_request_group_querier.list_upload_request_groups() @@ -1197,11 +1324,14 @@ async def approve_request( request_id: uuid.UUID, approved_by: StaffUser, ) -> UploadRequestDetails: - upload_request, staged_photos, finalized_storage_keys, created_photos = ( - await self._approve_request_without_side_effects( - request_id=request_id, - approved_by=approved_by, - ) + ( + upload_request, + staged_photos, + finalized_storage_keys, + created_photos, + ) = await self._approve_request_without_side_effects( + request_id=request_id, + approved_by=approved_by, ) try: await self.staff_notifications_service.create_notification( @@ -1248,12 +1378,14 @@ async def reject_request( approved_by: StaffUser, reason: str | None, ) -> UploadRequestDetails: - upload_request, rejected_photos, staged_photos = ( - await self._reject_request_without_side_effects( - request_id=request_id, - approved_by=approved_by, - reason=reason, - ) + ( + upload_request, + rejected_photos, + staged_photos, + ) = await self._reject_request_without_side_effects( + request_id=request_id, + approved_by=approved_by, + reason=reason, ) await self.staff_notifications_service.create_notification( staff_user_id=upload_request.requested_by, @@ -1307,23 +1439,30 @@ async def approve_group( finalized_storage_keys: list[str] = [] try: for request_details in pending_requests: - approved_request, staged_photos, request_storage_keys, created_photos = ( - await self._approve_request_without_side_effects( - request_id=request_details.request.id, - approved_by=approved_by, - ) + ( + approved_request, + staged_photos, + request_storage_keys, + created_photos, + ) = await self._approve_request_without_side_effects( + request_id=request_details.request.id, + approved_by=approved_by, ) approved_requests.append(approved_request) all_staged_photos.extend(staged_photos) all_created_photos.extend(created_photos) finalized_storage_keys.extend(request_storage_keys) - upload_group = await self.upload_request_group_querier.approve_upload_request_group( - id=group_id, - approved_by=approved_by.id, + upload_group = ( + await self.upload_request_group_querier.approve_upload_request_group( + id=group_id, + approved_by=approved_by.id, + ) ) if upload_group is None: - raise AppException.internal_error("Failed to approve upload request group") + raise AppException.internal_error( + "Failed to approve upload request group" + ) for approved_request in approved_requests: await self.staff_notifications_service.create_notification( @@ -1393,20 +1532,24 @@ async def reject_group( rejected_requests: list[UploadRequest] = [] all_staged_photos: list[UploadRequestPhoto] = [] for request_details in pending_requests: - rejected_request, _rejected_photos, staged_photos = ( - await self._reject_request_without_side_effects( - request_id=request_details.request.id, - approved_by=approved_by, - reason=reason, - ) + ( + rejected_request, + _rejected_photos, + staged_photos, + ) = await self._reject_request_without_side_effects( + request_id=request_details.request.id, + approved_by=approved_by, + reason=reason, ) rejected_requests.append(rejected_request) all_staged_photos.extend(staged_photos) - upload_group = await self.upload_request_group_querier.reject_upload_request_group( - id=group_id, - approved_by=approved_by.id, - rejection_reason=reason, + upload_group = ( + await self.upload_request_group_querier.reject_upload_request_group( + id=group_id, + approved_by=approved_by.id, + rejection_reason=reason, + ) ) if upload_group is None: raise AppException.internal_error("Failed to reject upload request group") diff --git a/app/service/user_notification.py b/app/service/user_notification.py index a269ff17..7038f0aa 100644 --- a/app/service/user_notification.py +++ b/app/service/user_notification.py @@ -52,7 +52,9 @@ async def create_notification( if tokens: notification = notification.model_copy(update={"tokens": tokens}) else: - logger.info("No active push tokens for user %s, skipping push", user_id) + logger.info( + "No active push tokens for user %s, skipping push", user_id + ) return notification_record await self._notification_queue.enqueue_notification(notification) @@ -64,9 +66,9 @@ async def get_all_notifications( user_id: uuid.UUID, ) -> list[Notification]: notifications: list[Notification] = [] - async for notification in self.notification_querier.list_notifications_by_user_id( - user_id=user_id - ): + async for ( + notification + ) in self.notification_querier.list_notifications_by_user_id(user_id=user_id): notifications.append(notification) return notifications diff --git a/app/service/user_photo.py b/app/service/user_photo.py index b85f03cd..f53844d4 100644 --- a/app/service/user_photo.py +++ b/app/service/user_photo.py @@ -83,7 +83,8 @@ async def count_event_photos( event_id: UUID, ) -> int: count = await self._photo_querier.count_event_photos_for_user( - user_id=user_id, event_id=event_id, + user_id=user_id, + event_id=event_id, ) return count or 0 @@ -138,11 +139,16 @@ async def get_photo_bytes( async def _user_has_access(self, user_id: UUID, photo_id: UUID) -> bool: """Check if user has a face_match or photo_approval for this photo.""" match = await self._photo_face_querier.user_has_face_match_for_photo( - photo_id=photo_id, user_id=user_id, + photo_id=photo_id, + user_id=user_id, ) if match is not None: return True - async for approval in self._photo_approval_querier.get_photo_approvals_by_photo_id(photo_id=photo_id): + async for ( + approval + ) in self._photo_approval_querier.get_photo_approvals_by_photo_id( + photo_id=photo_id + ): if approval.user_id == user_id: return True return False diff --git a/app/service/users.py b/app/service/users.py index e3a30b28..9c4ac456 100644 --- a/app/service/users.py +++ b/app/service/users.py @@ -69,8 +69,7 @@ async def _ensure_device_for_login( req: MobileAuthBaseRequest, ) -> UserDevice: existing_device = await self.device_querier.get_device_by_physical_id( - user_id=user_id, - physical_device_id=req.physical_device_id + user_id=user_id, physical_device_id=req.physical_device_id ) if existing_device: @@ -79,7 +78,9 @@ async def _ensure_device_for_login( "Device push token is invalid. Update the token before logging in." ) if not existing_device.is_active: - await self.device_querier.activate_device(id=existing_device.id, user_id=user_id) + await self.device_querier.activate_device( + id=existing_device.id, user_id=user_id + ) return existing_device device = await self.device_querier.create_device( @@ -89,7 +90,7 @@ async def _ensure_device_for_login( device_name=req.device_name, device_type=req.device_type, totp_secret=None, - physical_device_id=req.physical_device_id + physical_device_id=req.physical_device_id, ) ) if not device: @@ -123,19 +124,27 @@ async def mobile_login( existing_user = await self.user_querier.get_user_by_email(email=req.email) if existing_user is None: logger.warning("login attempt: user_not_found") - raise AppException.unauthorized("User not found; consider registering instead") + raise AppException.unauthorized( + "User not found; consider registering instead" + ) if existing_user.blocked: logger.warning("login attempt: user_blocked user_id=%s", existing_user.id) raise AppException.forbidden("User is blocked") if not verify_password(req.password, existing_user.hashed_password or ""): - logger.warning("login attempt: invalid_credentials user_id=%s", existing_user.id) + logger.warning( + "login attempt: invalid_credentials user_id=%s", existing_user.id + ) raise AppException.unauthorized("Invalid credentials") - locked_user = await self.user_querier.get_user_by_id_for_update(id=existing_user.id) + locked_user = await self.user_querier.get_user_by_id_for_update( + id=existing_user.id + ) if not locked_user: raise AppException.unauthorized("User not found") if locked_user.blocked: - logger.warning("login attempt: user_blocked_at_commit user_id=%s", locked_user.id) + logger.warning( + "login attempt: user_blocked_at_commit user_id=%s", locked_user.id + ) raise AppException.forbidden("User is blocked") logger.info("login success user_id=%s", locked_user.id) @@ -193,13 +202,14 @@ async def mobile_register( otp = "".join(secrets.choice("0123456789") for _ in range(6)) await redis.set(f"otp:{req.email}", otp, expire=600) # Send to NATS - await NatsClient.js_publish("email.send_otp", json.dumps({"email": req.email, "otp": otp}).encode("utf-8")) + await NatsClient.js_publish( + "email.send_otp", + json.dumps({"email": req.email, "otp": otp}).encode("utf-8"), + ) logger.info("register success, OTP sent") return RegisterPendingResponse( - message="OTP sent to email", - status="pending_verification", - email=req.email + message="OTP sent to email", status="pending_verification", email=req.email ) async def mobile_register_resend_otp( @@ -240,13 +250,14 @@ async def mobile_register_resend_otp( # Regenerate OTP with 10 mins TTL, without touching the pending_user TTL await redis.set(f"otp:{email}", otp, expire=600) # Send to NATS - await NatsClient.js_publish("email.send_otp", json.dumps({"email": email, "otp": otp}).encode("utf-8")) + await NatsClient.js_publish( + "email.send_otp", + json.dumps({"email": email, "otp": otp}).encode("utf-8"), + ) logger.info("resend_otp success, new OTP sent to %s", email) return RegisterPendingResponse( - message="New OTP sent to email", - status="pending_verification", - email=email + message="New OTP sent to email", status="pending_verification", email=email ) async def verify_mobile_register( @@ -269,7 +280,9 @@ async def verify_mobile_register( data = json.loads(raw_data) try: - user = await self.user_querier.create_user(email=req.email, hashed_password=data["hashed_password"]) + user = await self.user_querier.create_user( + email=req.email, hashed_password=data["hashed_password"] + ) if not user: raise AppException.internal_error("Failed to create user") except SQLAlchemyError as exc: @@ -304,7 +317,9 @@ async def _create_mobile_session( now = datetime.now(timezone.utc) idle_expires_at = now + timedelta(days=settings.MOBILE_SESSION_DAYS) - absolute_expires_at = now + timedelta(days=settings.MOBILE_SESSION_ABSOLUTE_DAYS) + absolute_expires_at = now + timedelta( + days=settings.MOBILE_SESSION_ABSOLUTE_DAYS + ) session = await self.session_querier.upsert_session( user_id=user_id, @@ -321,7 +336,8 @@ async def _create_mobile_session( await SessionService.delete_session_cache(redis, evicted_id) logger.warning( "session_evicted user_id=%s evicted_session_id=%s", - user_id, evicted_id, + user_id, + evicted_id, ) access_token = create_acces_mobile_token(str(session.id)) @@ -368,11 +384,9 @@ async def _handle_used_refresh_token( inside the grace window with no cached replay available (a used token with nothing to replay is never treated as valid). """ - within_grace = ( - row.used_at is not None - and (datetime.now(timezone.utc) - row.used_at) - <= timedelta(seconds=AuthService.REFRESH_GRACE_SECONDS) - ) + within_grace = row.used_at is not None and ( + datetime.now(timezone.utc) - row.used_at + ) <= timedelta(seconds=AuthService.REFRESH_GRACE_SECONDS) if within_grace: cache_key = f"refresh_retry:{token_hash}" @@ -384,10 +398,14 @@ async def _handle_used_refresh_token( # tampered, corrupted, or wrong key — treat exactly # like a cache miss, never trust an undecryptable value raise AppException.unauthorized("Invalid refresh token") - session_for_check = await self.session_querier.get_session_by_id(id=row.session_id) + session_for_check = await self.session_querier.get_session_by_id( + id=row.session_id + ) if not session_for_check: raise AppException.unauthorized("Session not found") - user_for_check = await self.user_querier.get_user_by_id(id=session_for_check.user_id) + user_for_check = await self.user_querier.get_user_by_id( + id=session_for_check.user_id + ) if not user_for_check or user_for_check.blocked: raise AppException.forbidden("User is blocked") return MobileAuthResponse.model_validate_json(decrypted) @@ -395,17 +413,18 @@ async def _handle_used_refresh_token( logger.warning( "refresh_token_reuse_detected family_id=%s session_id=%s", - row.family_id, row.session_id, + row.family_id, + row.session_id, + ) + session_for_revoke = await self.session_querier.get_session_by_id( + id=row.session_id ) - session_for_revoke = await self.session_querier.get_session_by_id(id=row.session_id) if session_for_revoke: await self.session_querier.delete_session_by_id( id=row.session_id, user_id=session_for_revoke.user_id ) await SessionService.delete_session_cache(redis, row.session_id) - raise AppException.unauthorized( - "Refresh token reuse detected; session revoked" - ) + raise AppException.unauthorized("Refresh token reuse detected; session revoked") async def refresh_token( self, @@ -474,7 +493,9 @@ async def logout( session_id: str, ) -> dict[str, str]: sid = uuid.UUID(session_id) - await self.session_querier.delete_session_by_id(id=sid, user_id=uuid.UUID(user_id)) + await self.session_querier.delete_session_by_id( + id=sid, user_id=uuid.UUID(user_id) + ) await SessionService.delete_session_cache(redis, sid) return {"message": "Logged out successfully"} @@ -669,7 +690,12 @@ async def block_user(self, *, redis: RedisClient, user_id: uuid.UUID) -> User: user = await self.user_querier.set_user_blocked(blocked=True, id=user_id) if not user: raise AppException.internal_error("Failed to block user") - session_ids = [s.id async for s in self.session_querier.list_sessions_by_user(user_id=user_id)] + session_ids = [ + s.id + async for s in self.session_querier.list_sessions_by_user( + user_id=user_id + ) + ] await self.session_querier.delete_all_user_sessions(user_id=user_id) for sid in session_ids: await SessionService.delete_session_cache(redis, sid) @@ -696,7 +722,9 @@ async def delete_user(self, *, redis: RedisClient, user_id: uuid.UUID) -> User: session_ids = [ s.id - async for s in self.session_querier.list_sessions_by_user(user_id=user_id) + async for s in self.session_querier.list_sessions_by_user( + user_id=user_id + ) ] await self.session_querier.delete_all_user_sessions(user_id=user_id) await self.user_querier.delete_user(id=user_id) @@ -709,7 +737,9 @@ async def delete_user(self, *, redis: RedisClient, user_id: uuid.UUID) -> User: logger.error("Failed to delete user: %s", exc) raise DBException.handle(exc) - async def find_closest_user(self, *, embedding_literal: str) -> ClosestUserMatch | None: + async def find_closest_user( + self, *, embedding_literal: str + ) -> ClosestUserMatch | None: row = await self.user_querier.find_closest_user_by_embedding( dollar_1=embedding_literal, ) @@ -732,7 +762,9 @@ async def check_rate_limit( except HTTPException: raise except Exception: - logger.warning("check_rate_limit: redis unavailable, failing open for key=%s", key) + logger.warning( + "check_rate_limit: redis unavailable, failing open for key=%s", key + ) return if current_count > max_requests: diff --git a/app/worker/audit/__init__.py b/app/worker/audit/__init__.py index 29d7469a..77c7bcff 100644 --- a/app/worker/audit/__init__.py +++ b/app/worker/audit/__init__.py @@ -1,4 +1,5 @@ """Audit worker package exports.""" + from __future__ import annotations __all__ = ["main"] diff --git a/app/worker/audit/main.py b/app/worker/audit/main.py index 7ec8ce3a..d9625ad8 100644 --- a/app/worker/audit/main.py +++ b/app/worker/audit/main.py @@ -38,9 +38,9 @@ def _parse_payload(raw_data: bytes) -> dict[str, Any] | None: try: parsed = json.loads(raw_data.decode("utf-8")) if not isinstance(parsed, dict): - logger.warning("Audit payload must be an object, got %s", type(parsed)) # type: ignore + logger.warning("Audit payload must be an object, got %s", type(parsed)) # type: ignore return None - return parsed # type: ignore + return parsed # type: ignore except (UnicodeDecodeError, json.JSONDecodeError) as exc: logger.error("Cannot parse audit payload: %s", exc) return None @@ -61,7 +61,7 @@ async def _handle_event(worker: AuditDeliveryWorker, raw_data: bytes) -> None: async def listen_nats_event(worker: AuditDeliveryWorker) -> None: async def handler(data: bytes) -> None: await _handle_event(worker, data) - + await NatsClient.js_subscribe( NatsSubjects.AUDIT_EVENT, handler, diff --git a/app/worker/audit/settings.py b/app/worker/audit/settings.py index c732ca5e..86a376c5 100644 --- a/app/worker/audit/settings.py +++ b/app/worker/audit/settings.py @@ -7,4 +7,4 @@ class AuditWorkerSettings(BaseSettings): model_config = SettingsConfigDict(env_prefix="AUDIT_") -settings = AuditWorkerSettings() # type: ignore +settings = AuditWorkerSettings() # type: ignore diff --git a/app/worker/drive_sync/main.py b/app/worker/drive_sync/main.py index 2ca5be4a..4ebd495b 100644 --- a/app/worker/drive_sync/main.py +++ b/app/worker/drive_sync/main.py @@ -47,7 +47,9 @@ async def _handle_event(raw_data: bytes) -> None: try: data, _, content_type = await bucket.get(event.storage_key) except Exception as exc: - logger.warning("drive_sync: failed to read photo %s from storage: %s", event.photo_id, exc) + logger.warning( + "drive_sync: failed to read photo %s from storage: %s", event.photo_id, exc + ) raise async with engine.begin() as conn: @@ -65,17 +67,24 @@ async def _handle_event(raw_data: bytes) -> None: data=data, ) except Exception as exc: - logger.warning("drive_sync: upload failed for photo %s: %s", event.photo_id, exc) + logger.warning( + "drive_sync: upload failed for photo %s: %s", event.photo_id, exc + ) raise synced = await photo_querier.mark_photo_drive_synced( - id=event.photo_id, drive_file_id=drive_file_id, + id=event.photo_id, + drive_file_id=drive_file_id, ) if synced is None: - logger.warning("drive_sync: photo %s not found when recording sync", event.photo_id) + logger.warning( + "drive_sync: photo %s not found when recording sync", event.photo_id + ) return - logger.info("drive_sync: synced photo %s to Drive as %s", event.photo_id, drive_file_id) + logger.info( + "drive_sync: synced photo %s to Drive as %s", event.photo_id, drive_file_id + ) async def main() -> None: @@ -93,7 +102,9 @@ async def main() -> None: ) await NatsClient.connect() try: - await NatsClient.js_subscribe(NatsSubjects.PHOTO_DRIVE_SYNC_REQUESTED, _handle_event) + await NatsClient.js_subscribe( + NatsSubjects.PHOTO_DRIVE_SYNC_REQUESTED, _handle_event + ) await asyncio.Event().wait() finally: await NatsClient.close() diff --git a/app/worker/event_lifecycle/main.py b/app/worker/event_lifecycle/main.py index 51aa6eff..fc1c3529 100644 --- a/app/worker/event_lifecycle/main.py +++ b/app/worker/event_lifecycle/main.py @@ -15,11 +15,15 @@ async def run_lifecycle_pass() -> None: activated = [event_id async for event_id in querier.activate_due_events()] if activated: - logger.info("event_lifecycle: activated %d event(s): %s", len(activated), activated) + logger.info( + "event_lifecycle: activated %d event(s): %s", len(activated), activated + ) archived = [event_id async for event_id in querier.archive_ended_events()] if archived: - logger.info("event_lifecycle: archived %d event(s): %s", len(archived), archived) + logger.info( + "event_lifecycle: archived %d event(s): %s", len(archived), archived + ) async def run_storage_cleanup_pass() -> None: @@ -44,7 +48,9 @@ async def run_storage_cleanup_pass() -> None: ) except Exception as exc: logger.warning( - "storage_cleanup: failed to schedule cleanup for photo %s: %s", photo.id, exc + "storage_cleanup: failed to schedule cleanup for photo %s: %s", + photo.id, + exc, ) continue marked = await querier.mark_photo_storage_cleaned(id=photo.id) diff --git a/app/worker/notification/firebase.py b/app/worker/notification/firebase.py index 5615995b..e41e234c 100644 --- a/app/worker/notification/firebase.py +++ b/app/worker/notification/firebase.py @@ -3,6 +3,7 @@ # pyright: ignore[reportMissingTypeStubs] import firebase_admin # type: ignore[import-not-found,import-untyped] + # pyright: ignore[reportMissingTypeStubs] from firebase_admin import credentials, messaging # type: ignore[import-not-found,import-untyped] @@ -24,6 +25,8 @@ class _SendResponse: class _BatchResponse: responses: list[_SendResponse] + + class NotificationDeliveryError(Exception): def __init__( self, @@ -39,16 +42,16 @@ def __init__( def init_firebase_app(credentials_path: str | None = None) -> None: - if firebase_admin._apps: # type: ignore + if firebase_admin._apps: # type: ignore return if credentials_path is None: credentials_path = settings.FIREBASE_CREDENTIALS_PATH if credentials_path: cred = credentials.Certificate(credentials_path) - firebase_admin.initialize_app(cred) # type: ignore + firebase_admin.initialize_app(cred) # type: ignore logger.info("Firebase initialized with credentials from %s", credentials_path) return - firebase_admin.initialize_app() # type: ignore + firebase_admin.initialize_app() # type: ignore logger.info("Firebase initialized with default credentials") @@ -75,9 +78,9 @@ def send_notification(notification: UnifiedNotification) -> None: data=notification.data or None, ) response = cast( - _BatchResponse, - messaging.send_multicast(multicast) # type: ignore -) + _BatchResponse, + messaging.send_multicast(multicast), # type: ignore + ) failed_tokens: list[str] = [] invalid_tokens: list[str] = [] @@ -96,6 +99,4 @@ def send_notification(notification: UnifiedNotification) -> None: invalid_tokens=invalid_tokens, ) - logger.info( - "Notification delivered to %d tokens", len(notification.tokens) - ) + logger.info("Notification delivered to %d tokens", len(notification.tokens)) diff --git a/app/worker/notification/invalid_tokens.py b/app/worker/notification/invalid_tokens.py index 6d1c7c4f..860657d2 100644 --- a/app/worker/notification/invalid_tokens.py +++ b/app/worker/notification/invalid_tokens.py @@ -22,7 +22,9 @@ async def mark_invalid(self, tokens: Iterable[str]) -> None: return await self._redis.sadd(RedisKey.INVALID_TOKEN_SET_KEY, *normalized) - await self._redis.expire(RedisKey.INVALID_TOKEN_SET_KEY, NotifSetting.TTL_SECONDS) + await self._redis.expire( + RedisKey.INVALID_TOKEN_SET_KEY, NotifSetting.TTL_SECONDS + ) logger.warning("Marked %d tokens for cleanup", len(normalized)) @@ -30,17 +32,13 @@ async def is_invalid(self, token: str) -> bool: if not token: return False - return await self._redis.sismember( - RedisKey.INVALID_TOKEN_SET_KEY, token - ) + return await self._redis.sismember(RedisKey.INVALID_TOKEN_SET_KEY, token) async def remove(self, tokens: Sequence[str]) -> None: if not tokens: return - await self._redis.srem( - RedisKey.INVALID_TOKEN_SET_KEY, *tokens - ) + await self._redis.srem(RedisKey.INVALID_TOKEN_SET_KEY, *tokens) class DeviceInvalidationStore: diff --git a/app/worker/notification/main.py b/app/worker/notification/main.py index 7ea07cbb..0f472efe 100644 --- a/app/worker/notification/main.py +++ b/app/worker/notification/main.py @@ -13,7 +13,10 @@ DeviceInvalidationStore, InvalidTokenStore, ) -from app.worker.notification.notification_queue import NotificationQueue, NotificationQueueEntry +from app.worker.notification.notification_queue import ( + NotificationQueue, + NotificationQueueEntry, +) from app.worker.notification.rate_limiter import RateLimiter from app.worker.notification.settings import NotifSetting from app.infra.redis import RedisClient @@ -52,7 +55,6 @@ async def process_entry( await retry(entry, queue) - async def retry( entry: NotificationQueueEntry, queue: NotificationQueue, @@ -71,13 +73,12 @@ async def retry( if not notification.tokens: return - delay = min(NotifSetting.BASE_RETRY_DELAY * (2 ** attempts), 60) + delay = min(NotifSetting.BASE_RETRY_DELAY * (2**attempts), 60) await asyncio.sleep(delay) await queue.enqueue_notification(notification, attempts=attempts) - async def handle_message( raw_payload: bytes | str, queue: NotificationQueue, @@ -97,7 +98,6 @@ async def handle_message( await process_entry(entry, queue, invalid_tokens, invalid_devices) - async def run_worker( queue: NotificationQueue, invalid_tokens: InvalidTokenStore, @@ -115,15 +115,12 @@ async def wrapped_handler(msg: bytes | str) -> None: for subject in queue.priority_subjects(): await NatsClient.js_subscribe( - subject, - wrapped_handler, - stream_name="notification_delivery_stream" + subject, wrapped_handler, stream_name="notification_delivery_stream" ) await asyncio.Event().wait() - async def main() -> None: init_firebase_app(NotifSetting.firebase_credentials_path) diff --git a/app/worker/notification/notification_queue.py b/app/worker/notification/notification_queue.py index c430d441..e769ac59 100644 --- a/app/worker/notification/notification_queue.py +++ b/app/worker/notification/notification_queue.py @@ -1,7 +1,11 @@ from typing import Sequence from pydantic import BaseModel, ConfigDict, Field from app.infra.nats import NatsClient -from app.schema.internal.notification import NotificationPriority, PRIORITY_ORDER, UnifiedNotification +from app.schema.internal.notification import ( + NotificationPriority, + PRIORITY_ORDER, + UnifiedNotification, +) from app.worker.notification.settings import NotificationWorkerSettings @@ -17,14 +21,14 @@ def __init__(self, settings: NotificationWorkerSettings) -> None: self._settings = settings async def enqueue_notification( - self, - notification: UnifiedNotification, - attempts: int = 0 + self, notification: UnifiedNotification, attempts: int = 0 ) -> None: entry = NotificationQueueEntry(notification=notification, attempts=attempts) subject = self._settings.subject_for(entry.notification.priority) payload = entry.model_dump_json().encode("utf-8") - await NatsClient.js_publish(subject, payload, stream_name="notification_delivery_stream") + await NatsClient.js_publish( + subject, payload, stream_name="notification_delivery_stream" + ) @staticmethod def priority_index(priority: NotificationPriority) -> int: diff --git a/app/worker/photo_worker/main.py b/app/worker/photo_worker/main.py index aba77037..7639a12f 100644 --- a/app/worker/photo_worker/main.py +++ b/app/worker/photo_worker/main.py @@ -17,7 +17,11 @@ from app.infra.redis import RedisClient from app.schema.internal.notification import NotificationPriority, UnifiedNotification from app.schema.internal.single_face_match import BBoxPayload -from app.service.face_embedding import DetectedFace, FaceEmbeddingService, FaceImagePayload +from app.service.face_embedding import ( + DetectedFace, + FaceEmbeddingService, + FaceImagePayload, +) from app.service.face_match import SingleFaceMatchService from app.service.user_notification import UserNotificationService from app.worker.photo_worker.schema.event import PhotoProcessEvent @@ -36,7 +40,6 @@ class PhotoApprovalDecision(str, Enum): class PhotoWorker: - def __init__( self, conn: AsyncConnection, @@ -72,9 +75,15 @@ async def handle_message(self, data: bytes) -> None: faces = await self._face_service.detect_faces(payload) if not faces: - logger.info("No faces detected in photo %s, marking as public", event.photo_id) - await self._photo_querier.update_photo_status(id=event.photo_id, status="approved") - await self._photo_querier.update_photo_visibility(id=event.photo_id, visibility="public") + logger.info( + "No faces detected in photo %s, marking as public", event.photo_id + ) + await self._photo_querier.update_photo_status( + id=event.photo_id, status="approved" + ) + await self._photo_querier.update_photo_visibility( + id=event.photo_id, visibility="public" + ) await self._update_job(job, "completed") return @@ -86,8 +95,9 @@ async def handle_message(self, data: bytes) -> None: await self._update_job(job, "completed") await self._publish_audit(event, len(faces)) - - async def _handle_single_face(self, event: PhotoProcessEvent, face: DetectedFace) -> None: + async def _handle_single_face( + self, event: PhotoProcessEvent, face: DetectedFace + ) -> None: from app.schema.internal.single_face_match import SingleFaceMatchJob bbox = BBoxPayload( @@ -106,23 +116,32 @@ async def _handle_single_face(self, event: PhotoProcessEvent, face: DetectedFace ) try: - await self._single_face_service.process_detected_face(job, face.embedding, bbox) + await self._single_face_service.process_detected_face( + job, face.embedding, bbox + ) except Exception as exc: - logger.exception("Single face match failed for photo %s: %s", event.photo_id, exc) - + logger.exception( + "Single face match failed for photo %s: %s", event.photo_id, exc + ) - async def _handle_group_photo(self, event: PhotoProcessEvent, faces: list[DetectedFace]) -> None: - logger.info("Processing group photo %s with %d faces", event.photo_id, len(faces)) + async def _handle_group_photo( + self, event: PhotoProcessEvent, faces: list[DetectedFace] + ) -> None: + logger.info( + "Processing group photo %s with %d faces", event.photo_id, len(faces) + ) approvals_created = 0 for face_index, face in enumerate(faces): - bbox_json = json.dumps({ - "x1": float(face.bbox[0]), - "y1": float(face.bbox[1]), - "x2": float(face.bbox[2]), - "y2": float(face.bbox[3]), - }) + bbox_json = json.dumps( + { + "x1": float(face.bbox[0]), + "y1": float(face.bbox[1]), + "x2": float(face.bbox[2]), + "y2": float(face.bbox[3]), + } + ) embedding_literal = "[" + ", ".join(str(x) for x in face.embedding) + "]" @@ -138,7 +157,9 @@ async def _handle_group_photo(self, event: PhotoProcessEvent, faces: list[Detect ) if approval is None: - logger.info("No match for face %d in photo %s", face_index, event.photo_id) + logger.info( + "No match for face %d in photo %s", face_index, event.photo_id + ) continue approvals_created += 1 @@ -159,22 +180,32 @@ async def _handle_group_photo(self, event: PhotoProcessEvent, faces: list[Detect priority=NotificationPriority.NORMAL, ), ) - logger.info("Notified user %s for group photo %s", approval.user_id, approval.photo_id) + logger.info( + "Notified user %s for group photo %s", + approval.user_id, + approval.photo_id, + ) except Exception as exc: logger.warning( "Failed to notify user %s for photo %s: %s", - approval.user_id, event.photo_id, exc, + approval.user_id, + event.photo_id, + exc, ) if approvals_created == 0: - logger.info("No users matched in group photo %s, leaving as pending", event.photo_id) - + logger.info( + "No users matched in group photo %s, leaving as pending", event.photo_id + ) - async def _create_job(self, event: PhotoProcessEvent) -> models.ProcessingJob | None: + async def _create_job( + self, event: PhotoProcessEvent + ) -> models.ProcessingJob | None: if self._pj_querier is None: return None return await self._pj_querier.create_processing_job( - photo_id=event.photo_id, job_type="face_detection", + photo_id=event.photo_id, + job_type="face_detection", ) async def _update_job(self, job: models.ProcessingJob | None, status: str) -> None: @@ -186,14 +217,19 @@ async def _update_job(self, job: models.ProcessingJob | None, status: str) -> No async def _publish_audit(event: PhotoProcessEvent, faces_count: int) -> None: from app.core.constant import AuditEventType from app.worker.audit.schema.audit import AuditEventMessage + msg = AuditEventMessage( event_type=AuditEventType.PHOTO_PROCESSED, metadata={"photo_id": str(event.photo_id), "faces_count": faces_count}, ) try: - await NatsClient.js_publish(NatsSubjects.AUDIT_EVENT, msg.model_dump_json().encode("utf-8")) + await NatsClient.js_publish( + NatsSubjects.AUDIT_EVENT, msg.model_dump_json().encode("utf-8") + ) except Exception as exc: - logger.warning("Failed to publish audit for photo %s: %s", event.photo_id, exc) + logger.warning( + "Failed to publish audit for photo %s: %s", event.photo_id, exc + ) @staticmethod def _parse_event(raw_data: bytes) -> PhotoProcessEvent | None: @@ -220,7 +256,10 @@ async def _load_image(self, image_ref: str) -> FaceImagePayload: last_exc = exc logger.warning( "MinIO fetch failed for %s (attempt %s/%s): %s", - object_name, attempt, settings.MINIO_RETRY_ATTEMPTS, exc, + object_name, + attempt, + settings.MINIO_RETRY_ATTEMPTS, + exc, ) if attempt < settings.MINIO_RETRY_ATTEMPTS: await asyncio.sleep(settings.MINIO_RETRY_BASE_SECONDS * attempt) @@ -231,7 +270,7 @@ async def _load_image(self, image_ref: str) -> FaceImagePayload: @staticmethod def _parse_minio_ref(image_ref: str) -> tuple[str, str]: if image_ref.startswith(MINIO_URL_PREFIX): - raw = image_ref[len(MINIO_URL_PREFIX):] + raw = image_ref[len(MINIO_URL_PREFIX) :] parts = raw.split("/", 1) if len(parts) != 2 or not parts[0] or not parts[1]: raise ValueError("Invalid MinIO image_ref format") @@ -289,7 +328,10 @@ async def handle(data: bytes) -> None: durable_name=worker_settings.durable_name, ) - logger.info("PhotoWorker subscribed on %s; waiting for jobs", NatsSubjects.PHOTO_PROCESS.value) + logger.info( + "PhotoWorker subscribed on %s; waiting for jobs", + NatsSubjects.PHOTO_PROCESS.value, + ) try: await asyncio.Event().wait() finally: diff --git a/app/worker/storage_cleaner/main.py b/app/worker/storage_cleaner/main.py index eb723d2d..6339d9d1 100644 --- a/app/worker/storage_cleaner/main.py +++ b/app/worker/storage_cleaner/main.py @@ -119,6 +119,7 @@ async def _handle_cleanup_event( async def main() -> None: await NatsClient.connect() try: + async def _jetstream_handler(data: bytes | str) -> None: async with engine.begin() as conn: querier = upload_request_photo_queries.AsyncQuerier(conn) diff --git a/app/worker/upload_reconciler/main.py b/app/worker/upload_reconciler/main.py index 915d7911..2f1c13ef 100644 --- a/app/worker/upload_reconciler/main.py +++ b/app/worker/upload_reconciler/main.py @@ -29,7 +29,9 @@ async def run_reconcile_pass() -> None: stat = await storage_service.stat_staging_object(photo.staging_storage_key) if stat is not None: await querier.confirm_upload_request_photo_transfer( - id=photo.id, size_bytes=stat.size, mime_type=stat.content_type, + id=photo.id, + size_bytes=stat.size, + mime_type=stat.content_type, ) confirmed += 1 else: @@ -38,7 +40,9 @@ async def run_reconcile_pass() -> None: logger.info( "upload_reconciler: reconciled %d stale photo(s) — %d confirmed, %d failed", - len(stale_photos), confirmed, failed, + len(stale_photos), + confirmed, + failed, ) diff --git a/coverage_report.txt b/coverage_report.txt new file mode 100644 index 00000000..10c398af --- /dev/null +++ b/coverage_report.txt @@ -0,0 +1,96 @@ +Name Stmts Miss Cover Missing +--------------------------------------------------------------------------------- +app/__init__.py 0 0 100% +app/container.py 83 0 100% +app/core/config.py 74 7 91% 98, 100-105 +app/core/constant.py 42 0 100% +app/core/exceptions.py 90 20 78% 27, 31, 47, 51, 55, 68, 74, 80, 97-106, 117, 126, 130, 136 +app/core/image_validation.py 78 10 87% 48, 54-55, 117-128 +app/core/logger.py 5 0 100% +app/core/securite.py 89 19 79% 37, 69, 73-74, 78-79, 83-90, 150-158, 169-177, 188-198 +app/core/utils.py 18 12 33% 13-37 +app/deps/ai_deps.py 5 0 100% +app/deps/client_ip.py 11 1 91% 13 +app/deps/cookie_auth.py 36 23 36% 13, 21-40, 44-46, 52, 56-58, 63 +app/deps/rate_limit.py 23 5 78% 20-24 +app/deps/token_auth.py 45 17 62% 39, 56-64, 83-112 +app/infra/database.py 8 1 88% 18 +app/infra/google_drive.py 233 165 29% 56-59, 63-68, 72-77, 83-96, 100-114, 126-139, 151-156, 168-185, 198-210, 218-258, 267-303, 313-354, 358-366, 375-397, 406-424, 432-452 +app/infra/minio.py 76 39 49% 28-37, 49-51, 54-74, 77-95, 98, 111-120, 123-131, 147, 150-151, 155, 159, 162-164 +app/infra/nats.py 92 37 60% 64-70, 83-92, 97-102, 112-122, 133-140 +app/infra/redis.py 47 12 74% 23, 31, 50-51, 54-55, 66-67, 70-71, 74-75 +app/main.py 76 32 58% 53-61, 69-100, 117-119, 137, 142 +app/router/mobile/__init__.py 15 0 100% +app/router/mobile/auth.py 87 42 52% 45-47, 56-63, 91, 99-108, 118-137, 147-153, 163-168, 176-203, 220-237, 250-253 +app/router/mobile/enrollement.py 70 43 39% 45-46, 62-78, 84, 93-97, 116-189 +app/router/mobile/event.py 13 2 85% 18, 29 +app/router/mobile/notifications.py 14 4 71% 17-20, 29-33 +app/router/mobile/photo_approval.py 18 7 61% 21-35, 45-50 +app/router/mobile/photos.py 21 6 71% 30, 53-64, 87-91 +app/router/staff/__init__.py 8 0 100% +app/router/staff/drive.py 57 31 46% 34-38, 48-81, 92-96, 110-111, 120-126, 147-153, 173-181 +app/router/staff/notifications.py 15 4 73% 20-23, 32-36 +app/router/staff/uploads.py 63 27 57% 40-50, 60-65, 77-82, 91-95, 104-108, 117-121, 131-136, 145-149, 158-162, 174-180, 189-193, 203-208 +app/router/web/__init__.py 14 0 100% +app/router/web/audit.py 14 2 86% 28-36 +app/router/web/auth.py 21 5 76% 21-33, 40, 54 +app/router/web/event.py 26 7 73% 34-37, 50, 60, 70, 80 +app/router/web/staff_users.py 29 11 62% 27-35, 52-56, 73-77, 92-94 +app/router/web/stats.py 18 4 78% 18, 27, 36, 45 +app/router/web/users.py 43 19 56% 21-28, 39-40, 49-50, 60-67, 76-81, 90-95, 104-106 +app/schema/internal/notification.py 14 0 100% +app/schema/internal/single_face_match.py 22 0 100% +app/schema/internal/uploads.py 8 0 100% +app/schema/request/mobile/auth.py 60 5 92% 29, 36, 79-81 +app/schema/request/mobile/notifications.py 4 0 100% +app/schema/request/mobile/photo_approval.py 4 0 100% +app/schema/request/staff/drive.py 7 0 100% +app/schema/request/staff/notifications.py 4 0 100% +app/schema/request/staff/uploads.py 58 24 59% 21-23, 28-31, 34, 56-59, 64-67, 71-75, 78-80 +app/schema/request/web/auth.py 4 0 100% +app/schema/request/web/event.py 9 0 100% +app/schema/request/web/staff_user.py 9 0 100% +app/schema/request/web/user.py 11 0 100% +app/schema/response/mobile/auth.py 39 0 100% +app/schema/response/mobile/notifications.py 19 2 89% 19, 36 +app/schema/response/staff/drive.py 32 0 100% +app/schema/response/staff/notifications.py 19 2 89% 19, 36 +app/schema/response/staff/upload_groups.py 67 8 88% 39, 43-47, 58, 67, 90, 98 +app/schema/response/staff/uploads.py 44 4 91% 44, 48-52 +app/schema/response/web/audit.py 25 2 92% 20, 39 +app/schema/response/web/auth.py 7 0 100% +app/schema/response/web/event.py 32 0 100% +app/schema/response/web/staff_user.py 9 0 100% +app/schema/response/web/stats.py 30 0 100% +app/schema/response/web/user.py 13 1 92% 18 +app/service/audit.py 37 17 54% 35-42, 52-58, 70-88 +app/service/device.py 56 41 27% 18-24, 31-38, 46-56, 63-72, 75-82, 89-95, 98-104 +app/service/event.py 62 41 34% 31-42, 45-48, 51-52, 63-76, 79-82, 88-105, 109-112, 116-119 +app/service/face_embedding.py 139 78 44% 67, 77, 80, 91-126, 138-142, 149-184, 191-217, 223-237, 241-249 +app/service/face_match.py 81 2 98% 76-77 +app/service/photo_approval.py 50 1 98% 78 +app/service/session.py 63 22 65% 55-56, 68-72, 74, 85-86, 91-97, 100-106 +app/service/staff_drive.py 185 122 34% 65-76, 79-89, 96-140, 143, 152-155, 159-161, 167-195, 198-201, 205-210, 213-217, 223, 226-229, 237-294, 297-298, 302-303, 307-310, 320-327 +app/service/staff_notifications.py 32 20 38% 23-30, 37-42, 50-56, 64-76 +app/service/staff_user.py 74 54 27% 25-47, 53-62, 65-71, 83-104, 112-124, 134-137 +app/service/staged_upload_storage.py 43 16 63% 36-37, 46-47, 58-69, 83-92, 95-98, 101-102 +app/service/stats.py 32 21 34% 17-21, 30-39, 46-50, 57-73 +app/service/upload_requests.py 441 350 21% 80, 84, 91, 98-106, 109-113, 116, 120-127, 130-134, 146-162, 169-173, 182-186, 196-205, 214-263, 275-321, 329-380, 389-411, 419-423, 431-435, 441-442, 448-449, 455-460, 470-473, 476-477, 483-493, 505, 526-534, 547-550, 566-594, 603-762, 770-777, 789-800, 809-834, 846-851, 859-876, 894-918, 926-930, 942-980, 992-1028, 1036-1112, 1124-1190 +app/service/user_notification.py 53 34 36% 26-32, 42-60, 67-72, 80-93 +app/service/user_photo.py 64 38 41% 50-51, 62-73, 81-84, 95-126, 130-138, 143-146 +app/service/users.py 373 93 75% 78, 82, 96, 160, 254, 259, 266, 308, 381, 422, 427, 434, 483, 485, 497, 499, 509, 516, 520-523, 526, 536-559, 562-565, 568-575, 585-607, 610-620, 625-634, 637-647, 650-654, 663, 725 +app/worker/audit/__init__.py 7 4 43% 8-12 +app/worker/audit/schema/audit.py 9 0 100% +app/worker/notification/__init__.py 0 0 100% +app/worker/notification/notification_queue.py 22 6 73% 24-27, 31, 34 +app/worker/notification/settings.py 28 2 93% 33, 36 +app/worker/photo_worker/__init__.py 0 0 100% +app/worker/photo_worker/main.py 184 113 39% 59-99, 103-123, 175-176, 186-188, 193-195, 199-208, 212-217, 221-225, 228-250, 254-260, 264-321, 325-334, 338 +app/worker/photo_worker/schema/__init__.py 0 0 100% +app/worker/photo_worker/schema/event.py 10 0 100% +app/worker/photo_worker/settings.py 9 0 100% +app/worker/upload_group_worker/__init__.py 0 0 100% +app/worker/upload_group_worker/schema/__init__.py 0 0 100% +app/worker/upload_group_worker/schema/event.py 12 0 100% +--------------------------------------------------------------------------------- +TOTAL 4293 1737 60% diff --git a/db/__init__.py b/db/__init__.py index 05f2bf43..b864340b 100644 --- a/db/__init__.py +++ b/db/__init__.py @@ -1,2 +1 @@ """Database package placeholder used for tooling.""" - diff --git a/migrations/env.py b/migrations/env.py index a49dbf2e..cd9cebe5 100644 --- a/migrations/env.py +++ b/migrations/env.py @@ -100,9 +100,7 @@ def run_migrations_online() -> None: ) with connectable.connect() as connection: - context.configure( - connection=connection, target_metadata=target_metadata - ) + context.configure(connection=connection, target_metadata=target_metadata) with context.begin_transaction(): context.run_migrations() diff --git a/migrations/helper.py b/migrations/helper.py index 00967daa..afacf683 100644 --- a/migrations/helper.py +++ b/migrations/helper.py @@ -3,9 +3,10 @@ SQL_DIR = os.path.join(os.path.dirname(__file__), "sql") + def run_sql_up(message: str) -> None: """ - message: Migration message exactly as typed (matches filename). + message: Migration message exactly as typed (matches filename). """ path = os.path.join(SQL_DIR, "up", message + ".sql") if not os.path.isfile(path): @@ -18,7 +19,7 @@ def run_sql_up(message: str) -> None: def run_sql_down(message: str) -> None: - # write the message here u create it in th emigration + # write the message here u create it in th emigration path = os.path.join(SQL_DIR, "down", message + ".sql") if not os.path.isfile(path): raise FileNotFoundError(f"Down SQL file not found: {path}") diff --git a/mobile-quickstart/seed.py b/mobile-quickstart/seed.py index 3f981fe8..42167187 100644 --- a/mobile-quickstart/seed.py +++ b/mobile-quickstart/seed.py @@ -38,14 +38,14 @@ # --------------------------------------------------------------------------- STAFF_USERS = [ - {"email": "admin@multai.dev", "password": "Admin1234!", "role": "admin"}, - {"email": "lead@multai.dev", "password": "Lead1234!", "role": "multi_team_lead"}, - {"email": "multi@multai.dev", "password": "Multi1234!", "role": "multi"}, + {"email": "admin@multai.dev", "password": "Admin1234!", "role": "admin"}, + {"email": "lead@multai.dev", "password": "Lead1234!", "role": "multi_team_lead"}, + {"email": "multi@multai.dev", "password": "Multi1234!", "role": "multi"}, ] MOBILE_USERS = [ {"email": "alice@example.com", "password": "Alice123!", "display_name": "Alice"}, - {"email": "bob@example.com", "password": "Bob1234!", "display_name": "Bob"}, + {"email": "bob@example.com", "password": "Bob1234!", "display_name": "Bob"}, ] EVENTS = [ @@ -81,6 +81,7 @@ # Helpers # --------------------------------------------------------------------------- + def now() -> datetime: return datetime.now(timezone.utc) @@ -108,15 +109,29 @@ def generate_placeholder_image(label: str, color: tuple[int, int, int]) -> bytes # Reset # --------------------------------------------------------------------------- + async def reset_db(conn: asyncpg.Connection) -> None: print("Resetting database...") tables = [ - "audit_events", "face_matches", "photo_faces", "photo_approvals", - "user_photos", "processing_jobs", "upload_request_photos", - "upload_requests", "upload_request_groups", "notifications", - "staff_notifications", "staff_drive_connections", "event_participants", - "user_sessions", "user_devices", "photos", "events", - "users", "staff_users", + "audit_events", + "face_matches", + "photo_faces", + "photo_approvals", + "user_photos", + "processing_jobs", + "upload_request_photos", + "upload_requests", + "upload_request_groups", + "notifications", + "staff_notifications", + "staff_drive_connections", + "event_participants", + "user_sessions", + "user_devices", + "photos", + "events", + "users", + "staff_users", ] for table in tables: await conn.execute(f"DELETE FROM {table}") @@ -136,6 +151,7 @@ async def reset_minio(minio: Minio) -> None: # MinIO # --------------------------------------------------------------------------- + async def init_minio(minio: Minio) -> None: print("Setting up MinIO buckets...") for bucket in [IMAGES_BUCKET, "documents"]: @@ -168,6 +184,7 @@ async def upload_photo( # Seeders # --------------------------------------------------------------------------- + async def seed_staff_users(conn: asyncpg.Connection) -> list[uuid.UUID]: print(" -> Seeding staff users...") ids = [] @@ -182,7 +199,10 @@ async def seed_staff_users(conn: asyncpg.Connection) -> list[uuid.UUID]: updated_at = EXCLUDED.updated_at RETURNING id """, - u["email"], hash_password(u["password"]), u["role"], now(), + u["email"], + hash_password(u["password"]), + u["role"], + now(), ) ids.append(row["id"]) print(f" [OK] {u['role']}: {u['email']} password: {u['password']}") @@ -203,7 +223,10 @@ async def seed_mobile_users(conn: asyncpg.Connection) -> list[uuid.UUID]: updated_at = EXCLUDED.updated_at RETURNING id """, - u["email"], hash_password(u["password"]), u["display_name"], now(), + u["email"], + hash_password(u["password"]), + u["display_name"], + now(), ) ids.append(row["id"]) print(f" [OK] {u['display_name']}: {u['email']} password: {u['password']}") @@ -222,7 +245,8 @@ async def seed_devices_and_sessions( VALUES ($1, 'Seed Device', 'android', $2, $2) RETURNING id """, - user_id, now(), + user_id, + now(), ) await conn.execute( """ @@ -230,7 +254,10 @@ async def seed_devices_and_sessions( VALUES ($1, $2, $3, $3, $4) ON CONFLICT (user_id, device_id) DO NOTHING """, - user_id, device_id, now(), future(30), + user_id, + device_id, + now(), + future(30), ) print(f" [OK] {len(user_ids)} device(s) + session(s)") @@ -252,8 +279,12 @@ async def seed_events( status = EXCLUDED.status RETURNING id """, - e["name"], e["event_code"], e["event_date"], - e["status"], staff_ids[i % len(staff_ids)], now(), + e["name"], + e["event_code"], + e["event_date"], + e["status"], + staff_ids[i % len(staff_ids)], + now(), ) ids.append(row["id"]) print(f" [OK] {e['event_code']} — join code: {e['event_code']}") @@ -275,7 +306,9 @@ async def seed_event_participants( VALUES ($1, $2, $3) ON CONFLICT (event_id, user_id) DO NOTHING """, - event_id, user_id, now(), + event_id, + user_id, + now(), ) count += 1 print(f" [OK] {count} participant record(s)") @@ -298,7 +331,8 @@ async def seed_photos( color_index += 1 await upload_photo( - minio, storage_key, + minio, + storage_key, f"Event {str(event_id)[:8]} / Photo {i + 1}", color, ) @@ -311,7 +345,12 @@ async def seed_photos( VALUES ($1, $2, $3, $4, $5, 'public', 'approved', $6) RETURNING id """, - event_id, uploader, storage_key, now(), i + 1, now(), + event_id, + uploader, + storage_key, + now(), + i + 1, + now(), ) photo_ids.append(row["id"]) print(f" [OK] {storage_key}") @@ -337,8 +376,10 @@ async def seed_photo_access( ON CONFLICT (photo_id, face_index) DO NOTHING RETURNING id """, - photo_id, embedding_str, - '{"x1":10,"y1":10,"x2":100,"y2":100}', now(), + photo_id, + embedding_str, + '{"x1":10,"y1":10,"x2":100,"y2":100}', + now(), ) if face_row: face_count += 1 @@ -348,8 +389,10 @@ async def seed_photo_access( INSERT INTO face_matches (photo_face_id, user_id, confidence, created_at) VALUES ($1, $2, $3, $4) """, - face_row["id"], user_id, - round(random.uniform(0.85, 0.99), 4), now(), + face_row["id"], + user_id, + round(random.uniform(0.85, 0.99), 4), + now(), ) match_count += 1 @@ -359,11 +402,15 @@ async def seed_photo_access( INSERT INTO photo_approvals (photo_id, user_id, decision, decided_at) VALUES ($1, $2, 'approved', $3) """, - photo_id, user_id, now(), + photo_id, + user_id, + now(), ) approval_count += 1 - print(f" [OK] {face_count} face(s), {match_count} match(es), {approval_count} approval(s)") + print( + f" [OK] {face_count} face(s), {match_count} match(es), {approval_count} approval(s)" + ) async def seed_user_photos( @@ -381,7 +428,9 @@ async def seed_user_photos( VALUES ($1, $2, 'public', $3) ON CONFLICT (user_id, photo_id) DO NOTHING """, - user_id, photo_id, now(), + user_id, + photo_id, + now(), ) count += 1 print(f" [OK] {count} record(s)") @@ -401,7 +450,10 @@ async def seed_processing_jobs( (photo_id, job_type, status, attempts, created_at, completed_at) VALUES ($1, $2, $3::processing_job_status, 1, $4, $4) """, - photo_id, job_type, "completed", now(), + photo_id, + job_type, + "completed", + now(), ) count += 1 print(f" [OK] {count} job(s)") @@ -418,7 +470,8 @@ async def seed_notifications( INSERT INTO notifications (user_id, type, payload, created_at) VALUES ($1, 'welcome', '{"message": "Welcome to multAI!"}', $2) """, - user_id, now(), + user_id, + now(), ) print(f" [OK] {len(user_ids)} notification(s)") @@ -434,7 +487,8 @@ async def seed_staff_notifications( INSERT INTO staff_notifications (staff_user_id, type, payload, created_at) VALUES ($1, 'system', '{"message": "Staff account seeded."}', $2) """, - staff_id, now(), + staff_id, + now(), ) print(f" [OK] {len(staff_ids)} notification(s)") @@ -456,10 +510,12 @@ async def seed_upload_request_groups( $5, 2, 'completed', $6, $6) RETURNING id """, - event_id, f"gdrive_folder_{i + 1}", + event_id, + f"gdrive_folder_{i + 1}", staff_ids[i % len(staff_ids)], staff_ids[(i + 1) % len(staff_ids)], - PHOTOS_PER_EVENT, now(), + PHOTOS_PER_EVENT, + now(), ) ids.append(row["id"]) print(f" [OK] {len(ids)} group(s)") @@ -483,11 +539,13 @@ async def seed_upload_requests( VALUES ($1, $2, $3, $4, 'approved'::upload_request_status, $5, $6, $7, $7) RETURNING id """, - event_id, f"gdrive_file_{i + 1}", + event_id, + f"gdrive_file_{i + 1}", staff_ids[i % len(staff_ids)], staff_ids[(i + 1) % len(staff_ids)], PHOTOS_PER_EVENT, - group_ids[i % len(group_ids)], now(), + group_ids[i % len(group_ids)], + now(), ) ids.append(row["id"]) print(f" [OK] {len(ids)} request(s)") @@ -498,6 +556,7 @@ async def seed_upload_requests( # Summary # --------------------------------------------------------------------------- + def print_summary() -> None: print() print("=" * 55) @@ -513,7 +572,9 @@ def print_summary() -> None: for e in EVENTS: print(f" {e['name']} join code: {e['event_code']}") print() - print(f"Photos: {len(EVENTS) * PHOTOS_PER_EVENT} total — approved, public, gallery-ready") + print( + f"Photos: {len(EVENTS) * PHOTOS_PER_EVENT} total — approved, public, gallery-ready" + ) print() print("Staff users:") for u in STAFF_USERS: @@ -526,6 +587,7 @@ def print_summary() -> None: # Entry point # --------------------------------------------------------------------------- + async def main(reset: bool = False) -> None: dsn = ( f"postgresql://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}" @@ -539,7 +601,9 @@ async def main(reset: bool = False) -> None: secure=False, ) - print(f"Connecting to {settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}...") + print( + f"Connecting to {settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}..." + ) conn: asyncpg.Connection = await asyncpg.connect(dsn) try: @@ -553,7 +617,7 @@ async def main(reset: bool = False) -> None: print("Seeding...\n") staff_ids = await seed_staff_users(conn) - user_ids = await seed_mobile_users(conn) + user_ids = await seed_mobile_users(conn) await seed_devices_and_sessions(conn, user_ids) diff --git a/scripts/check_ai_results.py b/scripts/check_ai_results.py index fa179436..04e4c633 100644 --- a/scripts/check_ai_results.py +++ b/scripts/check_ai_results.py @@ -2,6 +2,7 @@ from sqlalchemy.ext.asyncio import create_async_engine from app.core.config import settings + async def main(): url = f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" engine = create_async_engine(url) @@ -10,14 +11,23 @@ async def main(): import sqlalchemy # Check processing jobs - jobs = (await conn.execute(sqlalchemy.text("SELECT status, count(*) FROM processing_jobs GROUP BY status"))).fetchall() + jobs = ( + await conn.execute( + sqlalchemy.text( + "SELECT status, count(*) FROM processing_jobs GROUP BY status" + ) + ) + ).fetchall() print("=== PROCESSING JOBS STATUS ===") for j in jobs: print(f"Status: {j[0]}, Count: {j[1]}") # Check photo faces - faces = (await conn.execute(sqlalchemy.text("SELECT count(*) FROM photo_faces"))).scalar() + faces = ( + await conn.execute(sqlalchemy.text("SELECT count(*) FROM photo_faces")) + ).scalar() print("\n=== VISAGES DETECTES ===") print(f"Nombre total de visages isolés et enregistrés par l'IA : {faces}") + asyncio.run(main()) diff --git a/scripts/check_scopes.py b/scripts/check_scopes.py index 1835786e..158da87a 100644 --- a/scripts/check_scopes.py +++ b/scripts/check_scopes.py @@ -2,16 +2,25 @@ from sqlalchemy.ext.asyncio import create_async_engine from app.core.config import settings + async def main(): url = f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" engine = create_async_engine(url) async with engine.connect() as conn: import sqlalchemy - row = (await conn.execute(sqlalchemy.text("SELECT google_email, scopes FROM staff_drive_connections LIMIT 1"))).fetchone() + + row = ( + await conn.execute( + sqlalchemy.text( + "SELECT google_email, scopes FROM staff_drive_connections LIMIT 1" + ) + ) + ).fetchone() if row: print(f"Email: {row[0]}, Scopes: {row[1]}") else: print("No connection found") + asyncio.run(main()) diff --git a/scripts/generate_drive_url.py b/scripts/generate_drive_url.py index 319afdd8..b8b28bf2 100644 --- a/scripts/generate_drive_url.py +++ b/scripts/generate_drive_url.py @@ -6,39 +6,54 @@ from app.service.staff_drive import StaffDriveService from db.generated import staff_drive_connections as drive_queries + async def main(): url = f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" engine = create_async_engine(url) # Init redis - RedisClient.init(host=settings.REDIS_HOST, port=settings.REDIS_PORT, password=settings.REDIS_PASSWORD or "") + RedisClient.init( + host=settings.REDIS_HOST, + port=settings.REDIS_PORT, + password=settings.REDIS_PASSWORD or "", + ) redis = RedisClient.get_instance() async with engine.connect() as conn: q = staff_queries.AsyncQuerier(conn) import sqlalchemy - row = (await conn.execute(sqlalchemy.text("SELECT id FROM staff_users LIMIT 1"))).fetchone() + + row = ( + await conn.execute(sqlalchemy.text("SELECT id FROM staff_users LIMIT 1")) + ).fetchone() if not row: print("No staff user found! Creating one...") - res = await q.create_staff_user(email="testadmin@example.com", hashed_password="pw", display_name="Admin", role="admin") + res = await q.create_staff_user( + email="testadmin@example.com", + hashed_password="pw", + display_name="Admin", + role="admin", + ) staff_user_id = res.id else: staff_user_id = row[0] class DummyUser: pass + staff_user = DummyUser() staff_user.id = staff_user_id drive_service = StaffDriveService( staff_user_querier=q, drive_connection_querier=drive_queries.AsyncQuerier(conn), - redis=redis + redis=redis, ) url, state = await drive_service.create_connect_url(staff_user) print("========================") print("GOOGLE_AUTH_URL:", url) print("========================") + asyncio.run(main()) diff --git a/scripts/list_drive.py b/scripts/list_drive.py index 8651bddb..e634812a 100644 --- a/scripts/list_drive.py +++ b/scripts/list_drive.py @@ -7,13 +7,18 @@ from app.service.staff_drive import StaffDriveService from db.generated import staff_drive_connections as drive_queries + async def main(): url = f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" engine = create_async_engine(url) # Init redis try: - RedisClient.init(host=settings.REDIS_HOST, port=settings.REDIS_PORT, password=settings.REDIS_PASSWORD or "") + RedisClient.init( + host=settings.REDIS_HOST, + port=settings.REDIS_PORT, + password=settings.REDIS_PASSWORD or "", + ) except RuntimeError: pass redis = RedisClient.get_instance() @@ -21,30 +26,41 @@ async def main(): async with engine.connect() as conn: q = staff_queries.AsyncQuerier(conn) import sqlalchemy - row = (await conn.execute(sqlalchemy.text("SELECT id FROM staff_users LIMIT 1"))).fetchone() + + row = ( + await conn.execute(sqlalchemy.text("SELECT id FROM staff_users LIMIT 1")) + ).fetchone() staff_user_id = row[0] class DummyUser: pass + staff_user = DummyUser() staff_user.id = staff_user_id drive_service = StaffDriveService( staff_user_querier=q, drive_connection_querier=drive_queries.AsyncQuerier(conn), - redis=redis + redis=redis, ) - access_token = await drive_service.get_access_token_for_staff_user(staff_user_id) + access_token = await drive_service.get_access_token_for_staff_user( + staff_user_id + ) # List root folders print("=== DOSSIERS ET FICHIERS A LA RACINE DU DRIVE ===") items = await GoogleDriveClient.list_folder_contents(access_token=access_token) for i in items[:20]: - type_str = "📁 DOSSIER" if i.mime_type == "application/vnd.google-apps.folder" else "📄 FICHIER" + type_str = ( + "📁 DOSSIER" + if i.mime_type == "application/vnd.google-apps.folder" + else "📄 FICHIER" + ) print(f"{type_str} | ID: {i.id} | NOM: {i.name}") if not items: print("Aucun fichier trouvé.") + asyncio.run(main()) diff --git a/scripts/seed.py b/scripts/seed.py new file mode 100644 index 00000000..0d404daf --- /dev/null +++ b/scripts/seed.py @@ -0,0 +1,656 @@ +""" +multAI backend seed script + +What this does: + - Creates staff users (admin, lead, multi roles) + - Creates mobile users (alice, bob) with devices and sessions + - Creates 2 events and joins all users to them + - Uploads placeholder photos to MinIO and inserts them as approved + - Creates face matches and photo approvals so gallery endpoints return results + - Creates notifications for all users + - Creates processing jobs marked complete so pipeline endpoints look healthy + - Creates upload request groups and requests for staff review endpoints + +Usage: + docker compose -f docker-compose.mobile.yml exec fastapi uv run python seed.py + docker compose -f docker-compose.mobile.yml exec fastapi uv run python seed.py --reset +""" + +import asyncio +import io +import random +import sys +import uuid +from datetime import datetime, timedelta, timezone + +import asyncpg # type: ignore[import-untyped] +from dotenv import load_dotenv +from miniopy_async.api import Minio +from PIL import Image, ImageDraw + +from app.core.config import settings +from app.core.securite import hash_password + +load_dotenv() + +# --------------------------------------------------------------------------- +# Seed data +# --------------------------------------------------------------------------- + +STAFF_USERS = [ + {"email": "admin@multai.dev", "password": "Admin1234!", "role": "admin"}, + {"email": "lead@multai.dev", "password": "Lead1234!", "role": "multi_team_lead"}, + {"email": "multi@multai.dev", "password": "Multi1234!", "role": "multi"}, +] + +MOBILE_USERS = [ + {"email": "alice@example.com", "password": "Alice123!", "display_name": "Alice"}, + {"email": "bob@example.com", "password": "Bob1234!", "display_name": "Bob"}, +] + +EVENTS = [ + { + "name": "Tech Conference 2025", + "event_code": "TECH2025", + "event_date": datetime(2025, 9, 15, 9, 0, tzinfo=timezone.utc), + "status": "scheduled", + }, + { + "name": "Annual Gala", + "event_code": "GALA2025", + "event_date": datetime(2025, 12, 20, 19, 0, tzinfo=timezone.utc), + "status": "scheduled", + }, +] + +PHOTOS_PER_EVENT = 4 +IMAGES_BUCKET = "images" + +PHOTO_COLORS = [ + (52, 152, 219), + (46, 204, 113), + (231, 76, 60), + (155, 89, 182), + (241, 196, 15), + (230, 126, 34), + (26, 188, 156), + (52, 73, 94), +] + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def now() -> datetime: + return datetime.now(timezone.utc) + + +def future(days: int) -> datetime: + return now() + timedelta(days=days) + + +def make_storage_key(event_id: uuid.UUID, index: int) -> str: + return f"seed/events/{event_id}/photo_{index}.jpg" + + +def generate_placeholder_image(label: str, color: tuple[int, int, int]) -> bytes: + img = Image.new("RGB", (800, 600), color=color) + draw = ImageDraw.Draw(img) + bbox = draw.textbbox((0, 0), label) + w, h = bbox[2] - bbox[0], bbox[3] - bbox[1] + draw.text(((800 - w) / 2, (600 - h) / 2), label, fill=(255, 255, 255)) + buf = io.BytesIO() + img.save(buf, format="JPEG", quality=85) + return buf.getvalue() + + +# --------------------------------------------------------------------------- +# Reset +# --------------------------------------------------------------------------- + + +async def reset_db(conn: asyncpg.Connection) -> None: + print("Resetting database...") + tables = [ + "audit_events", + "face_matches", + "photo_faces", + "photo_approvals", + "user_photos", + "processing_jobs", + "upload_request_photos", + "upload_requests", + "upload_request_groups", + "notifications", + "staff_notifications", + "staff_drive_connections", + "event_participants", + "user_sessions", + "user_devices", + "photos", + "events", + "users", + "staff_users", + ] + for table in tables: + await conn.execute(f"DELETE FROM {table}") + print(f" cleared {table}") + print() + + +async def reset_minio(minio: Minio) -> None: + print("Clearing MinIO seed objects...") + objects = minio.list_objects(IMAGES_BUCKET, prefix="seed/", recursive=True) + async for obj in objects: + await minio.remove_object(IMAGES_BUCKET, obj.object_name) + print(" MinIO seed objects cleared\n") + + +# --------------------------------------------------------------------------- +# MinIO +# --------------------------------------------------------------------------- + + +async def init_minio(minio: Minio) -> None: + print("Setting up MinIO buckets...") + for bucket in [IMAGES_BUCKET, "documents"]: + if not await minio.bucket_exists(bucket): + await minio.make_bucket(bucket) + print(f" created bucket: {bucket}") + else: + print(f" bucket exists: {bucket}") + print() + + +async def upload_photo( + minio: Minio, + storage_key: str, + label: str, + color: tuple[int, int, int], +) -> None: + image_bytes = generate_placeholder_image(label, color) + await minio.put_object( + bucket_name=IMAGES_BUCKET, + object_name=storage_key, + data=io.BytesIO(image_bytes), + length=len(image_bytes), + content_type="image/jpeg", + metadata={"filename": storage_key.split("/")[-1]}, + ) + + +# --------------------------------------------------------------------------- +# Seeders +# --------------------------------------------------------------------------- + + +async def seed_staff_users(conn: asyncpg.Connection) -> list[uuid.UUID]: + print(" -> Seeding staff users...") + ids = [] + for u in STAFF_USERS: + row = await conn.fetchrow( + """ + INSERT INTO staff_users (email, password, role, created_at, updated_at) + VALUES ($1, $2, $3::staff_role, $4, $4) + ON CONFLICT (email) DO UPDATE + SET password = EXCLUDED.password, + role = EXCLUDED.role, + updated_at = EXCLUDED.updated_at + RETURNING id + """, + u["email"], + hash_password(u["password"]), + u["role"], + now(), + ) + ids.append(row["id"]) + print(f" [OK] {u['role']}: {u['email']} password: {u['password']}") + return ids + + +async def seed_mobile_users(conn: asyncpg.Connection) -> list[uuid.UUID]: + print(" -> Seeding mobile users...") + ids = [] + for u in MOBILE_USERS: + row = await conn.fetchrow( + """ + INSERT INTO users (email, hashed_password, display_name, created_at, updated_at) + VALUES ($1, $2, $3, $4, $4) + ON CONFLICT (email) DO UPDATE + SET hashed_password = EXCLUDED.hashed_password, + display_name = EXCLUDED.display_name, + updated_at = EXCLUDED.updated_at + RETURNING id + """, + u["email"], + hash_password(u["password"]), + u["display_name"], + now(), + ) + ids.append(row["id"]) + print(f" [OK] {u['display_name']}: {u['email']} password: {u['password']}") + return ids + + +async def seed_devices_and_sessions( + conn: asyncpg.Connection, + user_ids: list[uuid.UUID], +) -> None: + print(" -> Seeding devices + sessions...") + for user_id in user_ids: + device_id = await conn.fetchval( + """ + INSERT INTO user_devices (user_id, device_name, device_type, physical_device_id, last_active, created_at) + VALUES ($1, 'Seed Device', 'android', '00000000-0000-0000-0000-000000000000', $2, $2) + RETURNING id + """, + user_id, + now(), + ) + await conn.execute( + """ + INSERT INTO user_sessions (user_id, device_id, created_at, last_active, expires_at) + VALUES ($1, $2, $3, $3, $4) + ON CONFLICT (user_id, device_id) DO NOTHING + """, + user_id, + device_id, + now(), + future(30), + ) + print(f" [OK] {len(user_ids)} device(s) + session(s)") + + +async def seed_events( + conn: asyncpg.Connection, + staff_ids: list[uuid.UUID], +) -> list[uuid.UUID]: + print(" -> Seeding events...") + ids = [] + for i, e in enumerate(EVENTS): + row = await conn.fetchrow( + """ + INSERT INTO events (name, event_code, event_date, status, created_by, created_at) + VALUES ($1, $2, $3, $4::event_status, $5, $6) + ON CONFLICT (event_code) DO UPDATE + SET name = EXCLUDED.name, + event_date = EXCLUDED.event_date, + status = EXCLUDED.status + RETURNING id + """, + e["name"], + e["event_code"], + e["event_date"], + e["status"], + staff_ids[i % len(staff_ids)], + now(), + ) + ids.append(row["id"]) + print(f" [OK] {e['event_code']} — join code: {e['event_code']}") + return ids + + +async def seed_event_participants( + conn: asyncpg.Connection, + event_ids: list[uuid.UUID], + user_ids: list[uuid.UUID], +) -> None: + print(" -> Seeding event participants...") + count = 0 + for event_id in event_ids: + for user_id in user_ids: + await conn.execute( + """ + INSERT INTO event_participants (event_id, user_id, joined_at) + VALUES ($1, $2, $3) + ON CONFLICT (event_id, user_id) DO NOTHING + """, + event_id, + user_id, + now(), + ) + count += 1 + print(f" [OK] {count} participant record(s)") + + +async def seed_photos( + conn: asyncpg.Connection, + minio: Minio, + event_ids: list[uuid.UUID], + user_ids: list[uuid.UUID], +) -> list[uuid.UUID]: + print(" -> Seeding photos + uploading to MinIO...") + photo_ids = [] + color_index = 0 + for event_id in event_ids: + for i in range(PHOTOS_PER_EVENT): + uploader = user_ids[i % len(user_ids)] + storage_key = make_storage_key(event_id, i + 1) + color = PHOTO_COLORS[color_index % len(PHOTO_COLORS)] + color_index += 1 + + await upload_photo( + minio, + storage_key, + f"Event {str(event_id)[:8]} / Photo {i + 1}", + color, + ) + + row = await conn.fetchrow( + """ + INSERT INTO photos + (event_id, uploaded_by, storage_key, taken_at, day_number, + visibility, status, created_at) + VALUES ($1, $2, $3, $4, $5, 'public', 'approved', $6) + RETURNING id + """, + event_id, + uploader, + storage_key, + now(), + i + 1, + now(), + ) + photo_ids.append(row["id"]) + print(f" [OK] {storage_key}") + return photo_ids + + +async def seed_photo_access( + conn: asyncpg.Connection, + photo_ids: list[uuid.UUID], + user_ids: list[uuid.UUID], +) -> None: + print(" -> Seeding photo access (face matches + approvals)...") + face_count = match_count = approval_count = 0 + + for photo_id in photo_ids: + embedding = [random.uniform(-1.0, 1.0) for _ in range(512)] + embedding_str = "[" + ",".join(str(x) for x in embedding) + "]" + + face_row = await conn.fetchrow( + """ + INSERT INTO photo_faces (photo_id, face_index, embedding, bbox, created_at) + VALUES ($1, 0, $2::vector, $3, $4) + ON CONFLICT (photo_id, face_index) DO NOTHING + RETURNING id + """, + photo_id, + embedding_str, + '{"x1":10,"y1":10,"x2":100,"y2":100}', + now(), + ) + if face_row: + face_count += 1 + for user_id in user_ids: + await conn.execute( + """ + INSERT INTO face_matches (photo_face_id, user_id, confidence, created_at) + VALUES ($1, $2, $3, $4) + """, + face_row["id"], + user_id, + round(random.uniform(0.85, 0.99), 4), + now(), + ) + match_count += 1 + + for user_id in user_ids: + await conn.execute( + """ + INSERT INTO photo_approvals (photo_id, user_id, decision, decided_at) + VALUES ($1, $2, 'approved', $3) + """, + photo_id, + user_id, + now(), + ) + approval_count += 1 + + print( + f" [OK] {face_count} face(s), {match_count} match(es), {approval_count} approval(s)" + ) + + +async def seed_user_photos( + conn: asyncpg.Connection, + photo_ids: list[uuid.UUID], + user_ids: list[uuid.UUID], +) -> None: + print(" -> Seeding user_photos...") + count = 0 + for photo_id in photo_ids: + for user_id in user_ids: + await conn.execute( + """ + INSERT INTO user_photos (user_id, photo_id, visibility, created_at) + VALUES ($1, $2, 'public', $3) + ON CONFLICT (user_id, photo_id) DO NOTHING + """, + user_id, + photo_id, + now(), + ) + count += 1 + print(f" [OK] {count} record(s)") + + +async def seed_processing_jobs( + conn: asyncpg.Connection, + photo_ids: list[uuid.UUID], +) -> None: + print(" -> Seeding processing jobs...") + count = 0 + for photo_id in photo_ids: + for job_type in ["face_detection", "face_embedding"]: + await conn.execute( + """ + INSERT INTO processing_jobs + (photo_id, job_type, status, attempts, created_at, completed_at) + VALUES ($1, $2, $3::processing_job_status, 1, $4, $4) + """, + photo_id, + job_type, + "completed", + now(), + ) + count += 1 + print(f" [OK] {count} job(s)") + + +async def seed_notifications( + conn: asyncpg.Connection, + user_ids: list[uuid.UUID], +) -> None: + print(" -> Seeding notifications...") + for user_id in user_ids: + await conn.execute( + """ + INSERT INTO notifications (user_id, type, payload, created_at) + VALUES ($1, 'welcome', '{"message": "Welcome to multAI!"}', $2) + """, + user_id, + now(), + ) + print(f" [OK] {len(user_ids)} notification(s)") + + +async def seed_staff_notifications( + conn: asyncpg.Connection, + staff_ids: list[uuid.UUID], +) -> None: + print(" -> Seeding staff notifications...") + for staff_id in staff_ids: + await conn.execute( + """ + INSERT INTO staff_notifications (staff_user_id, type, payload, created_at) + VALUES ($1, 'system', '{"message": "Staff account seeded."}', $2) + """, + staff_id, + now(), + ) + print(f" [OK] {len(staff_ids)} notification(s)") + + +async def seed_upload_request_groups( + conn: asyncpg.Connection, + event_ids: list[uuid.UUID], + staff_ids: list[uuid.UUID], +) -> list[uuid.UUID]: + print(" -> Seeding upload request groups...") + ids = [] + for i, event_id in enumerate(event_ids): + row = await conn.fetchrow( + """ + INSERT INTO upload_request_groups + (event_id, folder_id, requested_by, approved_by, status, + total_photo_count, batch_count, processing_status, created_at, approved_at) + VALUES ($1, $2, $3, $4, 'approved'::upload_request_status, + $5, 2, 'completed', $6, $6) + RETURNING id + """, + event_id, + f"gdrive_folder_{i + 1}", + staff_ids[i % len(staff_ids)], + staff_ids[(i + 1) % len(staff_ids)], + PHOTOS_PER_EVENT, + now(), + ) + ids.append(row["id"]) + print(f" [OK] {len(ids)} group(s)") + return ids + + +async def seed_upload_requests( + conn: asyncpg.Connection, + event_ids: list[uuid.UUID], + staff_ids: list[uuid.UUID], + group_ids: list[uuid.UUID], +) -> list[uuid.UUID]: + print(" -> Seeding upload requests...") + ids = [] + for i, event_id in enumerate(event_ids): + row = await conn.fetchrow( + """ + INSERT INTO upload_requests + (event_id, drive_file_id, requested_by, approved_by, status, + photo_count, group_id, created_at, approved_at) + VALUES ($1, $2, $3, $4, 'approved'::upload_request_status, $5, $6, $7, $7) + RETURNING id + """, + event_id, + f"gdrive_file_{i + 1}", + staff_ids[i % len(staff_ids)], + staff_ids[(i + 1) % len(staff_ids)], + PHOTOS_PER_EVENT, + group_ids[i % len(group_ids)], + now(), + ) + ids.append(row["id"]) + print(f" [OK] {len(ids)} request(s)") + return ids + + +# --------------------------------------------------------------------------- +# Summary +# --------------------------------------------------------------------------- + + +def print_summary() -> None: + print() + print("=" * 55) + print(" SEED COMPLETE - MOBILE TEAM QUICKSTART") + print("=" * 55) + print() + print("Mobile users:") + for u in MOBILE_USERS: + print(f" email: {u['email']}") + print(f" password: {u['password']}") + print() + print("Events:") + for e in EVENTS: + print(f" {e['name']} join code: {e['event_code']}") + print() + print( + f"Photos: {len(EVENTS) * PHOTOS_PER_EVENT} total — approved, public, gallery-ready" + ) + print() + print("Staff users:") + for u in STAFF_USERS: + print(f" [{u['role']}] {u['email']} password: {u['password']}") + print("=" * 55) + print() + + +# --------------------------------------------------------------------------- +# Entry point +# --------------------------------------------------------------------------- + + +async def main(reset: bool = False) -> None: + dsn = ( + f"postgresql://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}" + f"@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" + ) + + minio = Minio( + f"{settings.MINIO_HOST}:{settings.MINIO_API_PORT}", + access_key=settings.MINIO_ROOT_USER, + secret_key=settings.MINIO_ROOT_PASSWORD, + secure=False, + ) + + print( + f"Connecting to {settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}..." + ) + conn: asyncpg.Connection = await asyncpg.connect(dsn) + + try: + if reset: + await reset_db(conn) + await reset_minio(minio) + + await init_minio(minio) + + print("Seeding...\n") + + staff_ids = await seed_staff_users(conn) + user_ids = await seed_mobile_users(conn) + + try: + await seed_devices_and_sessions(conn, user_ids) + except Exception as e: + print(f"Skipping device seeding due to schema mismatch: {e}") + + event_ids = await seed_events(conn, staff_ids) + await seed_event_participants(conn, event_ids, user_ids) + + try: + photo_ids = await seed_photos(conn, minio, event_ids, user_ids) + await seed_photo_access(conn, photo_ids, user_ids) + await seed_user_photos(conn, photo_ids, user_ids) + await seed_processing_jobs(conn, photo_ids) + except Exception as e: + print(f"Skipping photo seeding due to schema mismatch: {e}") + + await seed_notifications(conn, user_ids) + await seed_staff_notifications(conn, staff_ids) + + group_ids = await seed_upload_request_groups(conn, event_ids, staff_ids) + await seed_upload_requests(conn, event_ids, staff_ids, group_ids) + + print_summary() + + except Exception as e: + print(f"Seed failed: {e}") + raise + finally: + await conn.close() + + +if __name__ == "__main__": + reset_flag = "--reset" in sys.argv + if reset_flag: + print("WARNING: Reset mode — all existing seed data will be wiped.\n") + asyncio.run(main(reset=reset_flag)) diff --git a/scripts/seed_admin.py b/scripts/seed_admin.py index 91f4a7b0..ed0f747a 100644 --- a/scripts/seed_admin.py +++ b/scripts/seed_admin.py @@ -4,24 +4,26 @@ from db.generated.models import StaffRole from app.infra.redis import RedisClient + async def main(): RedisClient.init(host="localhost", port=6379, password="") async with engine.begin() as conn: container = Container(conn) # Check if exists - existing = await container.staff_user_service.staff_user_querier.get_staff_user_by_email(email="m@example.com") + existing = await container.staff_user_service.staff_user_querier.get_staff_user_by_email( + email="m@example.com" + ) if existing: print("Admin already exists!") return print("Creating admin user m@example.com...") await container.staff_user_service.create_staff_user( - email="m@example.com", - password="password", - role=StaffRole.ADMIN + email="m@example.com", password="password", role=StaffRole.ADMIN ) print("Admin user created! password is: password") + if __name__ == "__main__": asyncio.run(main()) diff --git a/scripts/trigger_import.py b/scripts/trigger_import.py index 86bb226c..35fdfae0 100644 --- a/scripts/trigger_import.py +++ b/scripts/trigger_import.py @@ -8,12 +8,17 @@ from app.service.staff_drive import StaffDriveService, SelectedDriveFile from db.generated import staff_drive_connections as drive_queries + async def main(): url = f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" engine = create_async_engine(url) try: - RedisClient.init(host=settings.REDIS_HOST, port=settings.REDIS_PORT, password=settings.REDIS_PASSWORD or "") + RedisClient.init( + host=settings.REDIS_HOST, + port=settings.REDIS_PORT, + password=settings.REDIS_PASSWORD or "", + ) except RuntimeError: pass redis = RedisClient.get_instance() @@ -23,27 +28,33 @@ async def main(): minio_host=settings.MINIO_HOST, minio_port=settings.MINIO_API_PORT, minio_root_user=settings.MINIO_ROOT_USER, - minio_root_password=settings.MINIO_ROOT_PASSWORD + minio_root_password=settings.MINIO_ROOT_PASSWORD, ) async with engine.connect() as conn: q = staff_queries.AsyncQuerier(conn) import sqlalchemy - row = (await conn.execute(sqlalchemy.text("SELECT id FROM staff_users LIMIT 1"))).fetchone() + + row = ( + await conn.execute(sqlalchemy.text("SELECT id FROM staff_users LIMIT 1")) + ).fetchone() staff_user_id = row[0] class DummyUser: pass + staff_user = DummyUser() staff_user.id = staff_user_id drive_service = StaffDriveService( staff_user_querier=q, drive_connection_querier=drive_queries.AsyncQuerier(conn), - redis=redis + redis=redis, ) - access_token = await drive_service.get_access_token_for_staff_user(staff_user_id) + access_token = await drive_service.get_access_token_for_staff_user( + staff_user_id + ) print("Fetching images from Drive...") items = await GoogleDriveClient.list_folder_contents(access_token=access_token) @@ -70,4 +81,5 @@ class DummyUser: for r in results: print(f" -> {r.original_file_name} stored as {r.minio_object_name}") + asyncio.run(main()) diff --git a/scripts/trigger_photo_worker.py b/scripts/trigger_photo_worker.py index b1eba585..e872c816 100644 --- a/scripts/trigger_photo_worker.py +++ b/scripts/trigger_photo_worker.py @@ -7,12 +7,17 @@ from app.infra.minio import init_minio_client from app.infra.nats import NatsClient, NatsSubjects + async def main(): url = f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" engine = create_async_engine(url) try: - RedisClient.init(host=settings.REDIS_HOST, port=settings.REDIS_PORT, password=settings.REDIS_PASSWORD or "") + RedisClient.init( + host=settings.REDIS_HOST, + port=settings.REDIS_PORT, + password=settings.REDIS_PASSWORD or "", + ) except RuntimeError: pass @@ -20,35 +25,48 @@ async def main(): minio_host=settings.MINIO_HOST, minio_port=settings.MINIO_API_PORT, minio_root_user=settings.MINIO_ROOT_USER, - minio_root_password=settings.MINIO_ROOT_PASSWORD + minio_root_password=settings.MINIO_ROOT_PASSWORD, ) await NatsClient.connect( host=settings.NATS_HOST, port=settings.NATS_PORT, user=settings.NATS_USER, - password=settings.NATS_PASSWORD + password=settings.NATS_PASSWORD, ) async with engine.connect() as conn: # staff_queries not used here but kept for context import sqlalchemy - row = (await conn.execute(sqlalchemy.text("SELECT id FROM staff_users LIMIT 1"))).fetchone() + + row = ( + await conn.execute(sqlalchemy.text("SELECT id FROM staff_users LIMIT 1")) + ).fetchone() staff_user_id = row[0] - event_row = (await conn.execute(sqlalchemy.text("SELECT id FROM events LIMIT 1"))).fetchone() + event_row = ( + await conn.execute(sqlalchemy.text("SELECT id FROM events LIMIT 1")) + ).fetchone() if not event_row: print("Creating dummy event...") ev_id = uuid.uuid4() - await conn.execute(sqlalchemy.text("INSERT INTO events (id, title, date, location) VALUES (:id, 'Test Event', now(), 'Test Location')"), {"id": ev_id}) + await conn.execute( + sqlalchemy.text( + "INSERT INTO events (id, title, date, location) VALUES (:id, 'Test Event', now(), 'Test Location')" + ), + {"id": ev_id}, + ) event_id = ev_id else: event_id = event_row[0] from app.infra.minio import ImageBucket + bucket = ImageBucket(f"staff-drive/{staff_user_id}") - objects = bucket.client.list_objects(bucket.bucket_name, prefix=bucket.file_prefix + "/", recursive=True) + objects = bucket.client.list_objects( + bucket.bucket_name, prefix=bucket.file_prefix + "/", recursive=True + ) count = 0 async for obj in objects: storage_key = obj.object_name @@ -57,18 +75,22 @@ async def main(): # Create photo in DB new_id = uuid.uuid4() await conn.execute( - sqlalchemy.text("INSERT INTO photos (id, event_id, storage_key, visibility) VALUES (:id, :event_id, :storage_key, 'public')"), - {"id": new_id, "event_id": event_id, "storage_key": storage_key} + sqlalchemy.text( + "INSERT INTO photos (id, event_id, storage_key, visibility) VALUES (:id, :event_id, :storage_key, 'public')" + ), + {"id": new_id, "event_id": event_id, "storage_key": storage_key}, ) # Publish event await NatsClient.publish( NatsSubjects.PHOTO_PROCESS, - json.dumps({ - "photo_id": str(new_id), - "image_ref": storage_key, - "event_id": str(event_id) - }).encode("utf-8") + json.dumps( + { + "photo_id": str(new_id), + "image_ref": storage_key, + "event_id": str(event_id), + } + ).encode("utf-8"), ) count += 1 if count >= 20: @@ -77,4 +99,5 @@ async def main(): await conn.commit() print(f"Successfully injected {count} photos to the AI worker!") + asyncio.run(main()) diff --git a/scripts/trigger_upload_request.py b/scripts/trigger_upload_request.py index 2b6361a6..92024987 100644 --- a/scripts/trigger_upload_request.py +++ b/scripts/trigger_upload_request.py @@ -20,12 +20,17 @@ from app.infra.nats import NatsClient import uuid + async def main(): url = f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" engine = create_async_engine(url) try: - RedisClient.init(host=settings.REDIS_HOST, port=settings.REDIS_PORT, password=settings.REDIS_PASSWORD or "") + RedisClient.init( + host=settings.REDIS_HOST, + port=settings.REDIS_PORT, + password=settings.REDIS_PASSWORD or "", + ) except RuntimeError: pass redis = RedisClient.get_instance() @@ -34,34 +39,45 @@ async def main(): minio_host=settings.MINIO_HOST, minio_port=settings.MINIO_API_PORT, minio_root_user=settings.MINIO_ROOT_USER, - minio_root_password=settings.MINIO_ROOT_PASSWORD + minio_root_password=settings.MINIO_ROOT_PASSWORD, ) await NatsClient.connect( host=settings.NATS_HOST, port=settings.NATS_PORT, user=settings.NATS_USER, - password=settings.NATS_PASSWORD + password=settings.NATS_PASSWORD, ) async with engine.connect() as conn: q = staff_queries.AsyncQuerier(conn) import sqlalchemy - row = (await conn.execute(sqlalchemy.text("SELECT id FROM staff_users LIMIT 1"))).fetchone() + + row = ( + await conn.execute(sqlalchemy.text("SELECT id FROM staff_users LIMIT 1")) + ).fetchone() staff_user_id = row[0] class DummyUser: pass + staff_user = DummyUser() staff_user.id = staff_user_id - staff_user.role = "multi_team_lead" # Important for approval! + staff_user.role = "multi_team_lead" # Important for approval! # Need an event - event_row = (await conn.execute(sqlalchemy.text("SELECT id FROM events LIMIT 1"))).fetchone() + event_row = ( + await conn.execute(sqlalchemy.text("SELECT id FROM events LIMIT 1")) + ).fetchone() if not event_row: print("Creating dummy event...") ev_id = uuid.uuid4() - await conn.execute(sqlalchemy.text("INSERT INTO events (id, title, date, location) VALUES (:id, 'Test Event', now(), 'Test Location')"), {"id": ev_id}) + await conn.execute( + sqlalchemy.text( + "INSERT INTO events (id, title, date, location) VALUES (:id, 'Test Event', now(), 'Test Location')" + ), + {"id": ev_id}, + ) event_id = ev_id else: event_id = event_row[0] @@ -72,23 +88,34 @@ class DummyUser: upload_request_photo_querier=request_photo_queries.AsyncQuerier(conn), photo_querier=photo_queries.AsyncQuerier(conn), staged_upload_storage=StagedUploadStorageService(), - staff_drive_service=StaffDriveService(staff_user_querier=q, drive_connection_querier=drive_queries.AsyncQuerier(conn), redis=redis), - staff_notifications_service=StaffNotificationsService(notif_queries.AsyncQuerier(conn)), + staff_drive_service=StaffDriveService( + staff_user_querier=q, + drive_connection_querier=drive_queries.AsyncQuerier(conn), + redis=redis, + ), + staff_notifications_service=StaffNotificationsService( + notif_queries.AsyncQuerier(conn) + ), audit_service=AuditService(audit_queries.AsyncQuerier(conn), None), ) # Get staged photos from minio bucket for this staff user from app.infra.minio import ImageBucket + bucket = ImageBucket(f"staff-drive/{staff_user_id}") - objects = bucket.client.list_objects(bucket.bucket_name, prefix=bucket.file_prefix + "/", recursive=True) + objects = bucket.client.list_objects( + bucket.bucket_name, prefix=bucket.file_prefix + "/", recursive=True + ) photo_inputs = [] async for obj in objects: name = obj.object_name.split("/")[-1] - photo_inputs.append(CreateUploadRequestPhotoRequest( - staged_object_name=name, - original_file_name=name, - )) + photo_inputs.append( + CreateUploadRequestPhotoRequest( + staged_object_name=name, + original_file_name=name, + ) + ) if len(photo_inputs) >= 20: break @@ -110,8 +137,11 @@ class DummyUser: print(f"Created upload request! ID: {req_details.id}") print("Approving upload request to trigger AI pipeline...") - await upload_service.approve_request(request_id=req_details.id, approved_by=staff_user) + await upload_service.approve_request( + request_id=req_details.id, approved_by=staff_user + ) print("Done! Photos should now be processed by AI.") + asyncio.run(main()) diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 8428d2b7..613c3572 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -15,10 +15,12 @@ FIXTURE_DIR = Path(__file__).parent.parent / "fixtures" / "images" + # ── guard: only run when explicitly requested ───────────────────────── def pytest_configure(config: pytest.Config) -> None: config.addinivalue_line("markers", "e2e: mark test as an end-to-end test") + @pytest.fixture(autouse=True) async def setup_infra() -> AsyncGenerator[None, None]: if os.getenv("MULTAI_RUN_E2E") != "1": @@ -39,6 +41,7 @@ async def setup_infra() -> AsyncGenerator[None, None]: # ── shared helpers ──────────────────────────────────────────────────── + async def _seed_event_and_photo( conn: AsyncConnection, *, @@ -110,11 +113,24 @@ async def _cleanup( ) -> None: """Delete all rows created during a test, in FK-safe order.""" if user_id: - await conn.execute(text("DELETE FROM notifications WHERE user_id = :uid"), {"uid": user_id}) # type: ignore[union-attr] - await conn.execute(text("DELETE FROM face_matches WHERE user_id = :uid"), {"uid": user_id}) # type: ignore[union-attr] + await conn.execute( + text("DELETE FROM notifications WHERE user_id = :uid"), {"uid": user_id} + ) # type: ignore[union-attr] + await conn.execute( + text("DELETE FROM face_matches WHERE user_id = :uid"), {"uid": user_id} + ) # type: ignore[union-attr] await conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) # type: ignore[union-attr] - await conn.execute(text("DELETE FROM face_matches fm USING photo_faces pf WHERE pf.id = fm.photo_face_id AND pf.photo_id = :pid"), {"pid": photo_id}) # type: ignore[union-attr] - await conn.execute(text("DELETE FROM photo_faces WHERE photo_id = :pid"), {"pid": photo_id}) # type: ignore[union-attr] - await conn.execute(text("DELETE FROM processing_jobs WHERE photo_id = :pid"), {"pid": photo_id}) # type: ignore[union-attr] + await conn.execute( + text( + "DELETE FROM face_matches fm USING photo_faces pf WHERE pf.id = fm.photo_face_id AND pf.photo_id = :pid" + ), + {"pid": photo_id}, + ) # type: ignore[union-attr] + await conn.execute( + text("DELETE FROM photo_faces WHERE photo_id = :pid"), {"pid": photo_id} + ) # type: ignore[union-attr] + await conn.execute( + text("DELETE FROM processing_jobs WHERE photo_id = :pid"), {"pid": photo_id} + ) # type: ignore[union-attr] await conn.execute(text("DELETE FROM photos WHERE id = :pid"), {"pid": photo_id}) # type: ignore[union-attr] await conn.execute(text("DELETE FROM events WHERE id = :eid"), {"eid": event_id}) # type: ignore[union-attr] diff --git a/tests/e2e/test_mobile_auth_intent_e2e.py b/tests/e2e/test_mobile_auth_intent_e2e.py index f37ba8d9..9a69e611 100644 --- a/tests/e2e/test_mobile_auth_intent_e2e.py +++ b/tests/e2e/test_mobile_auth_intent_e2e.py @@ -65,6 +65,7 @@ def test_register_with_existing_email_fails(self) -> None: # But if they are FULLY registered, it returns 409. # Let's verify them first to fully register them. import redis + r = redis.Redis(host="localhost", port=6379, decode_responses=True) otp = r.get(f"otp:{email}") @@ -115,6 +116,7 @@ def test_register_then_login_succeeds(self) -> None: assert register_response.json()["status"] == "pending_verification" import redis + r = redis.Redis(host="localhost", port=6379, decode_responses=True) otp = r.get(f"otp:{email}") @@ -178,6 +180,7 @@ def test_login_with_wrong_password_fails(self) -> None: assert register_response.status_code == 200 import redis + r = redis.Redis(host="localhost", port=6379, decode_responses=True) otp = r.get(f"otp:{email}") diff --git a/tests/e2e/test_photo_ai_edge_cases.py b/tests/e2e/test_photo_ai_edge_cases.py index 7b36da22..7ff96026 100644 --- a/tests/e2e/test_photo_ai_edge_cases.py +++ b/tests/e2e/test_photo_ai_edge_cases.py @@ -7,7 +7,12 @@ from app.infra.minio import Bucket, IMAGES_BUCKET_NAME from app.infra.nats import NatsClient, NatsSubjects -from tests.e2e.conftest import _seed_event_and_photo, _wait_for_job, _cleanup, FIXTURE_DIR +from tests.e2e.conftest import ( + _seed_event_and_photo, + _wait_for_job, + _cleanup, + FIXTURE_DIR, +) # ── tests ───────────────────────────────────────────────────────────── @@ -30,8 +35,14 @@ async def test_photo_ai_pipeline_detects_0_faces() -> None: conn, photo_id=photo_id, storage_key=storage_key ) - payload = {"photo_id": str(photo_id), "image_ref": storage_key, "event_id": str(event_id)} - await NatsClient.js_publish(NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8")) + payload = { + "photo_id": str(photo_id), + "image_ref": storage_key, + "event_id": str(event_id), + } + await NatsClient.js_publish( + NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8") + ) try: final_status = await _wait_for_job(photo_id) @@ -74,8 +85,14 @@ async def test_photo_ai_pipeline_detects_multiple_faces() -> None: conn, photo_id=photo_id, storage_key=storage_key ) - payload = {"photo_id": str(photo_id), "image_ref": storage_key, "event_id": str(event_id)} - await NatsClient.js_publish(NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8")) + payload = { + "photo_id": str(photo_id), + "image_ref": storage_key, + "event_id": str(event_id), + } + await NatsClient.js_publish( + NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8") + ) try: final_status = await _wait_for_job(photo_id) @@ -153,8 +170,14 @@ async def test_photo_ai_pipeline_matched_user() -> None: {"uid": matched_user_id, "emb": embedding_literal}, ) - payload = {"photo_id": str(photo_id), "image_ref": storage_key, "event_id": str(event_id)} - await NatsClient.js_publish(NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8")) + payload = { + "photo_id": str(photo_id), + "image_ref": storage_key, + "event_id": str(event_id), + } + await NatsClient.js_publish( + NatsSubjects.PHOTO_PROCESS, json.dumps(payload).encode("utf-8") + ) try: final_status = await _wait_for_job(photo_id) diff --git a/tests/e2e/test_photo_ai_load.py b/tests/e2e/test_photo_ai_load.py index 37aa1faa..063cefd7 100644 --- a/tests/e2e/test_photo_ai_load.py +++ b/tests/e2e/test_photo_ai_load.py @@ -122,7 +122,11 @@ async def test_photo_ai_load_20_photos(setup_infra: None) -> None: # noqa: ARG0 selected_img = random.choice(image_files) storage_key = f"load-test/{event_id}/{photo_id}.jpg" photo_tasks.append( - {"photo_id": photo_id, "storage_key": storage_key, "content": image_contents[selected_img]} + { + "photo_id": photo_id, + "storage_key": storage_key, + "content": image_contents[selected_img], + } ) await conn.execute( text( @@ -169,7 +173,8 @@ async def test_photo_ai_load_20_photos(setup_infra: None) -> None: # noqa: ARG0 # 3. Wait for all jobs photo_ids: list[uuid.UUID] = [ - p["photo_id"] for p in photo_tasks # type: ignore[misc] + p["photo_id"] + for p in photo_tasks # type: ignore[misc] ] status_counts = await _wait_for_jobs(photo_ids, timeout=180) diff --git a/tests/e2e/test_photo_ai_pipeline_e2e.py b/tests/e2e/test_photo_ai_pipeline_e2e.py index 1a16292b..9226dc01 100644 --- a/tests/e2e/test_photo_ai_pipeline_e2e.py +++ b/tests/e2e/test_photo_ai_pipeline_e2e.py @@ -7,7 +7,12 @@ from app.infra.minio import Bucket, IMAGES_BUCKET_NAME from app.infra.nats import NatsClient, NatsSubjects -from tests.e2e.conftest import _seed_event_and_photo, _wait_for_job, _cleanup, FIXTURE_DIR +from tests.e2e.conftest import ( + _seed_event_and_photo, + _wait_for_job, + _cleanup, + FIXTURE_DIR, +) async def test_photo_ai_pipeline_detects_single_face() -> None: @@ -96,7 +101,9 @@ async def test_photo_ai_pipeline_corrupt_image() -> None: # 3. Assertions + Cleanup try: final_status = await _wait_for_job(photo_id, timeout_s=30) - assert final_status == "failed", f"Expected job to fail, but ended with: {final_status}" + assert final_status == "failed", ( + f"Expected job to fail, but ended with: {final_status}" + ) finally: async with engine.begin() as conn: await _cleanup(conn, photo_id=photo_id, event_id=event_id) diff --git a/tests/e2e/test_stats_endpoint.py b/tests/e2e/test_stats_endpoint.py index 8620b0e7..b9951f9c 100644 --- a/tests/e2e/test_stats_endpoint.py +++ b/tests/e2e/test_stats_endpoint.py @@ -1,5 +1,6 @@ import pytest import os + if os.getenv("MULTAI_RUN_E2E") != "1": pytest.skip("set MULTAI_RUN_E2E=1 to run live e2e tests", allow_module_level=True) @@ -11,9 +12,11 @@ from db.generated.models import StaffUser, StaffRole from typing import Generator + @pytest.fixture(scope="module") def client() -> Generator[TestClient, None, None]: import os + if os.getenv("MULTAI_RUN_E2E") != "1": pytest.skip("set MULTAI_RUN_E2E=1 to run live e2e tests") # Override the dependency to bypass cookie auth @@ -23,7 +26,7 @@ def client() -> Generator[TestClient, None, None]: role=StaffRole.ADMIN, created_at=datetime.now(timezone.utc), updated_at=datetime.now(timezone.utc), - password="hashed_password" + password="hashed_password", ) app.dependency_overrides[require_admin_staff] = lambda: mock_admin @@ -31,26 +34,38 @@ def client() -> Generator[TestClient, None, None]: with TestClient(app) as c: yield c + def test_dashboard_stats(client: TestClient) -> None: resp = client.get("/admin/stats/dashboard") - assert resp.status_code == 200, f"Expected 200 but got {resp.status_code}: {resp.text}" + assert resp.status_code == 200, ( + f"Expected 200 but got {resp.status_code}: {resp.text}" + ) data = resp.json() assert "active_events" in data + def test_processing_load(client: TestClient) -> None: resp = client.get("/admin/stats/processing-load") - assert resp.status_code == 200, f"Expected 200 but got {resp.status_code}: {resp.text}" + assert resp.status_code == 200, ( + f"Expected 200 but got {resp.status_code}: {resp.text}" + ) data = resp.json() assert "completed" in data + def test_storage(client: TestClient) -> None: resp = client.get("/admin/stats/storage") - assert resp.status_code == 200, f"Expected 200 but got {resp.status_code}: {resp.text}" + assert resp.status_code == 200, ( + f"Expected 200 but got {resp.status_code}: {resp.text}" + ) data = resp.json() assert "used_bytes" in data + def test_alerts(client: TestClient) -> None: resp = client.get("/admin/stats/alerts") - assert resp.status_code == 200, f"Expected 200 but got {resp.status_code}: {resp.text}" + assert resp.status_code == 200, ( + f"Expected 200 but got {resp.status_code}: {resp.text}" + ) data = resp.json() assert "alerts" in data diff --git a/tests/integration/test_enrollment_flow.py b/tests/integration/test_enrollment_flow.py index 1968946b..b6e1a530 100644 --- a/tests/integration/test_enrollment_flow.py +++ b/tests/integration/test_enrollment_flow.py @@ -25,6 +25,7 @@ @pytest.fixture def mock_face_embedding() -> AsyncMock: from app.service.face_embedding import FaceEmbeddingService + svc = MagicMock(spec=FaceEmbeddingService) # Return a dummy embedding of size 512 svc.compute_average_embedding_stream = AsyncMock(return_value=[0.1] * 512) @@ -35,12 +36,14 @@ def mock_face_embedding() -> AsyncMock: async def db_conn(): from sqlalchemy.ext.asyncio import create_async_engine from app.core.config import settings + url = f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" engine = create_async_engine(url, pool_pre_ping=True) async with engine.connect() as conn: yield conn await engine.dispose() + @pytest.fixture def auth_service(mock_face_embedding: AsyncMock, db_conn) -> AuthService: from db.generated import session as session_queries @@ -79,7 +82,9 @@ async def test_enrollment_persists_embedding( user_id = user.id # 2. Execute enrollment - payload = FaceImagePayload(bytes=b"fake-image", filename="face.jpg", content_type="image/jpeg") + payload = FaceImagePayload( + bytes=b"fake-image", filename="face.jpg", content_type="image/jpeg" + ) try: await auth_service.add_embbed_user( @@ -88,7 +93,9 @@ async def test_enrollment_persists_embedding( ) # 3. Verify: The user should now have an embedding - updated_user = await user_queries.AsyncQuerier(db_conn).get_user_by_id(id=user_id) + updated_user = await user_queries.AsyncQuerier(db_conn).get_user_by_id( + id=user_id + ) assert updated_user is not None assert updated_user.face_embedding is not None assert "0.1" in str(updated_user.face_embedding) diff --git a/tests/integration/test_photo_approval_flow.py b/tests/integration/test_photo_approval_flow.py index 7ffcff5d..e548bd9d 100644 --- a/tests/integration/test_photo_approval_flow.py +++ b/tests/integration/test_photo_approval_flow.py @@ -28,6 +28,7 @@ @pytest.fixture def mock_storage() -> AsyncMock: from app.service.staged_upload_storage import StagedUploadStorageService + svc = MagicMock(spec=StagedUploadStorageService) svc.delete_storage_key = AsyncMock() return svc @@ -36,6 +37,7 @@ def mock_storage() -> AsyncMock: @pytest.fixture def mock_audit() -> AsyncMock: from app.service.audit import AuditService + svc = MagicMock(spec=AuditService) svc.create_record = AsyncMock() return svc @@ -45,12 +47,14 @@ def mock_audit() -> AsyncMock: async def db_conn(): from sqlalchemy.ext.asyncio import create_async_engine from app.core.config import settings + url = f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" engine = create_async_engine(url, pool_pre_ping=True) async with engine.connect() as conn: yield conn await engine.dispose() + @pytest.fixture def approval_service( mock_storage: AsyncMock, @@ -90,12 +94,16 @@ async def test_group_photo_approval_lifecycle( pq = photo_queries.AsyncQuerier(db_conn) aq = approval_queries.AsyncQuerier(db_conn) - staff = await sq.create_admin(email=f"admin-{uuid.uuid4()}@test.com", password="hash") + staff = await sq.create_admin( + email=f"admin-{uuid.uuid4()}@test.com", password="hash" + ) event_creator_id = staff.id user_ids = [] for i in range(3): - u = await uq.create_user(email=f"approval-{uuid.uuid4()}@test.com", hashed_password="hash") + u = await uq.create_user( + email=f"approval-{uuid.uuid4()}@test.com", hashed_password="hash" + ) user_ids.append(u.id) uploader_id, user1_id, user2_id = user_ids @@ -105,9 +113,10 @@ async def test_group_photo_approval_lifecycle( name="Approval Test Event", event_code=f"APP{str(event_id)[:4]}", event_date=datetime.datetime.now(datetime.timezone.utc), - end_date=datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(days=1), + end_date=datetime.datetime.now(datetime.timezone.utc) + + datetime.timedelta(days=1), status="scheduled", - created_by=event_creator_id + created_by=event_creator_id, ) ) event_id = event.id @@ -119,40 +128,60 @@ async def test_group_photo_approval_lifecycle( source="direct", taken_at=None, day_number=None, - visibility="public" + visibility="public", ) ) # Set status to pending and id await db_conn.execute( - text(f"UPDATE photos SET id = '{photo_id}', status = 'pending', uploaded_by = '{uploader_id}' WHERE storage_key = 'test/group.jpg'") + text( + f"UPDATE photos SET id = '{photo_id}', status = 'pending', uploaded_by = '{uploader_id}' WHERE storage_key = 'test/group.jpg'" + ) ) - await aq.create_photo_approval(photo_id=photo_id, user_id=user1_id, decision="pending") - await aq.create_photo_approval(photo_id=photo_id, user_id=user2_id, decision="pending") + await aq.create_photo_approval( + photo_id=photo_id, user_id=user1_id, decision="pending" + ) + await aq.create_photo_approval( + photo_id=photo_id, user_id=user2_id, decision="pending" + ) try: # 2. User 1 approves - result1 = await approval_service.decide(photo_id=photo_id, user_id=user1_id, decision="approved") - assert result1 == "pending", "Photo should remain pending because User 2 hasn't approved yet" + result1 = await approval_service.decide( + photo_id=photo_id, user_id=user1_id, decision="approved" + ) + assert result1 == "pending", ( + "Photo should remain pending because User 2 hasn't approved yet" + ) photo = await photo_queries.AsyncQuerier(db_conn).get_photo_by_id(id=photo_id) assert photo.status == "pending" # 3. User 2 approves - result2 = await approval_service.decide(photo_id=photo_id, user_id=user2_id, decision="approved") - assert result2 == "approved", "Photo should be approved since all users approved" + result2 = await approval_service.decide( + photo_id=photo_id, user_id=user2_id, decision="approved" + ) + assert result2 == "approved", ( + "Photo should be approved since all users approved" + ) photo = await photo_queries.AsyncQuerier(db_conn).get_photo_by_id(id=photo_id) assert photo.status == "approved" finally: # 4. Cleanup - await db_conn.execute(text(f"DELETE FROM photo_approvals WHERE photo_id = '{photo_id}'")) + await db_conn.execute( + text(f"DELETE FROM photo_approvals WHERE photo_id = '{photo_id}'") + ) await db_conn.execute(text(f"DELETE FROM photos WHERE id = '{photo_id}'")) await db_conn.execute(text(f"DELETE FROM events WHERE id = '{event_id}'")) - await db_conn.execute(text(f"DELETE FROM users WHERE id IN ('{user1_id}', '{user2_id}')")) - await db_conn.execute(text(f"DELETE FROM staff_users WHERE id = '{event_creator_id}'")) + await db_conn.execute( + text(f"DELETE FROM users WHERE id IN ('{user1_id}', '{user2_id}')") + ) + await db_conn.execute( + text(f"DELETE FROM staff_users WHERE id = '{event_creator_id}'") + ) await db_conn.commit() @@ -172,12 +201,16 @@ async def test_group_photo_rejection_deletes_storage( pq = photo_queries.AsyncQuerier(db_conn) aq = approval_queries.AsyncQuerier(db_conn) - staff = await sq.create_admin(email=f"admin-{uuid.uuid4()}@test.com", password="hash") + staff = await sq.create_admin( + email=f"admin-{uuid.uuid4()}@test.com", password="hash" + ) event_creator_id = staff.id user_ids = [] for i in range(2): - u = await uq.create_user(email=f"reject-{uuid.uuid4()}@test.com", hashed_password="hash") + u = await uq.create_user( + email=f"reject-{uuid.uuid4()}@test.com", hashed_password="hash" + ) user_ids.append(u.id) uploader_id, user1_id = user_ids @@ -187,9 +220,10 @@ async def test_group_photo_rejection_deletes_storage( name="Reject Test Event", event_code=f"REJ{str(event_id)[:4]}", event_date=datetime.datetime.now(datetime.timezone.utc), - end_date=datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(days=1), + end_date=datetime.datetime.now(datetime.timezone.utc) + + datetime.timedelta(days=1), status="scheduled", - created_by=event_creator_id + created_by=event_creator_id, ) ) event_id = event.id @@ -201,19 +235,25 @@ async def test_group_photo_rejection_deletes_storage( source="direct", taken_at=None, day_number=None, - visibility="public" + visibility="public", ) ) await db_conn.execute( - text(f"UPDATE photos SET id = '{photo_id}', status = 'pending', uploaded_by = '{uploader_id}' WHERE storage_key = 'test/reject.jpg'") + text( + f"UPDATE photos SET id = '{photo_id}', status = 'pending', uploaded_by = '{uploader_id}' WHERE storage_key = 'test/reject.jpg'" + ) ) - await aq.create_photo_approval(photo_id=photo_id, user_id=user1_id, decision="pending") + await aq.create_photo_approval( + photo_id=photo_id, user_id=user1_id, decision="pending" + ) try: # 2. User 1 rejects - result = await approval_service.decide(photo_id=photo_id, user_id=user1_id, decision="rejected") + result = await approval_service.decide( + photo_id=photo_id, user_id=user1_id, decision="rejected" + ) # 3. Verify assert result == "rejected" @@ -223,9 +263,15 @@ async def test_group_photo_rejection_deletes_storage( assert photo.status == "rejected" finally: - await db_conn.execute(text(f"DELETE FROM photo_approvals WHERE photo_id = '{photo_id}'")) + await db_conn.execute( + text(f"DELETE FROM photo_approvals WHERE photo_id = '{photo_id}'") + ) await db_conn.execute(text(f"DELETE FROM photos WHERE id = '{photo_id}'")) await db_conn.execute(text(f"DELETE FROM events WHERE id = '{event_id}'")) - await db_conn.execute(text(f"DELETE FROM users WHERE id IN ('{uploader_id}', '{user1_id}')")) - await db_conn.execute(text(f"DELETE FROM staff_users WHERE id = '{event_creator_id}'")) + await db_conn.execute( + text(f"DELETE FROM users WHERE id IN ('{uploader_id}', '{user1_id}')") + ) + await db_conn.execute( + text(f"DELETE FROM staff_users WHERE id = '{event_creator_id}'") + ) await db_conn.commit() diff --git a/tests/integration/test_session_device_management.py b/tests/integration/test_session_device_management.py index 6151dd78..6e3dfdd6 100644 --- a/tests/integration/test_session_device_management.py +++ b/tests/integration/test_session_device_management.py @@ -8,6 +8,7 @@ actually being ON DELETE CASCADE. Both were previously verified by hand via psql; these tests make that verification automatic and regression-proof. """ + import asyncio import uuid from datetime import datetime, timedelta, timezone @@ -156,7 +157,9 @@ async def test_relogin_on_same_device_replaces_not_duplicates_real_db( await db_conn.execute( text("DELETE FROM user_devices WHERE user_id = :uid"), {"uid": user_id} ) - await db_conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) + await db_conn.execute( + text("DELETE FROM users WHERE id = :uid"), {"uid": user_id} + ) await db_conn.commit() @@ -224,9 +227,12 @@ async def test_revoke_device_cascades_delete_session_real_db( await db_conn.execute( text("DELETE FROM user_devices WHERE user_id = :uid"), {"uid": user_id} ) - await db_conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) + await db_conn.execute( + text("DELETE FROM users WHERE id = :uid"), {"uid": user_id} + ) await db_conn.commit() + @pytest.mark.skip(reason="Flaky Postgres concurrency test in CI") @pytest.mark.asyncio async def test_concurrent_new_device_logins_settle_at_cap_real_db( @@ -325,9 +331,12 @@ async def _login_task(task_idx: int) -> None: await db_conn.execute( text("DELETE FROM user_devices WHERE user_id = :uid"), {"uid": user_id} ) - await db_conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) + await db_conn.execute( + text("DELETE FROM users WHERE id = :uid"), {"uid": user_id} + ) await db_conn.commit() + @pytest.mark.asyncio async def test_concurrent_block_and_login_never_leaves_blocked_user_with_session( db_conn, @@ -353,7 +362,8 @@ async def _run_one_trial() -> None: email = f"test-block-race-{uuid.uuid4()}@multai.com" user = await user_queries.AsyncQuerier(db_conn).create_user( - email=email, hashed_password=hash_password(password), + email=email, + hashed_password=hash_password(password), ) assert user is not None user_id = user.id @@ -369,8 +379,10 @@ async def _login() -> None: refresh_token_querier=refresh_token_queries.AsyncQuerier(conn), ) req = MobileLoginRequest( - email=email, password=password, - device_name="Race Device", device_type="android", + email=email, + password=password, + device_name="Race Device", + device_type="android", physical_device_id=uuid.uuid4(), ) try: @@ -411,7 +423,9 @@ async def _block() -> None: await db_conn.execute( text("DELETE FROM user_devices WHERE user_id = :uid"), {"uid": user_id} ) - await db_conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) + await db_conn.execute( + text("DELETE FROM users WHERE id = :uid"), {"uid": user_id} + ) await db_conn.commit() if blocked: diff --git a/tests/security/test_auth_security.py b/tests/security/test_auth_security.py index 7cb83405..bde4dd25 100644 --- a/tests/security/test_auth_security.py +++ b/tests/security/test_auth_security.py @@ -15,6 +15,7 @@ pytestmark = pytest.mark.asyncio(loop_scope="session") + @pytest.fixture(scope="session", autouse=True) async def setup_infra(): # We must init Redis since ASGITransport doesn't trigger the lifespan @@ -25,15 +26,19 @@ async def setup_infra(): password=settings.REDIS_PASSWORD, ) except RuntimeError: - pass # Already initialized + pass # Already initialized yield await RedisClient.get_instance().close() + @pytest.fixture(scope="session") async def client(): - async with AsyncClient(transport=ASGITransport(app=app), base_url="http://testserver") as c: + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://testserver" + ) as c: yield c + def create_mock_jwt(user_id: str, exp_delta_hours: int = 24) -> str: payload = { "sub": user_id, @@ -41,6 +46,7 @@ def create_mock_jwt(user_id: str, exp_delta_hours: int = 24) -> str: } return jwt.encode(payload, settings.jwt_secret, algorithm="HS256") + async def test_jwt_validation_invalid_signature(client): """Test that a JWT with an invalid signature is rejected.""" payload = { @@ -51,30 +57,29 @@ async def test_jwt_validation_invalid_signature(client): # We must use "Bearer " response = await client.get( - "/user/photos", - headers={"Authorization": f"Bearer {invalid_token}"} + "/user/photos", headers={"Authorization": f"Bearer {invalid_token}"} ) assert response.status_code == 401 assert "Invalid token" in response.text + async def test_jwt_validation_expired_token(client): """Test that an expired JWT is rejected.""" expired_token = create_mock_jwt(str(uuid.uuid4()), exp_delta_hours=-1) response = await client.get( - "/user/photos", - headers={"Authorization": f"Bearer {expired_token}"} + "/user/photos", headers={"Authorization": f"Bearer {expired_token}"} ) assert response.status_code == 401 assert "Token has expired" in response.text + async def test_blocked_user_access(client): """Test that a blocked user cannot access protected endpoints.""" async with engine.begin() as conn: uq = user_queries.AsyncQuerier(conn) user = await uq.create_user( - email=f"blocked-{uuid.uuid4()}@test.com", - hashed_password="hash" + email=f"blocked-{uuid.uuid4()}@test.com", hashed_password="hash" ) user_id = user.id @@ -93,7 +98,7 @@ async def test_blocked_user_access(client): absolute_expires_at=datetime.now(timezone.utc) + timedelta(days=30), blocked=True, ttl=3600, - last_active=datetime.now(timezone.utc) + last_active=datetime.now(timezone.utc), ) payload = { @@ -104,21 +109,24 @@ async def test_blocked_user_access(client): try: response = await client.get( - "/user/photos", - headers={"Authorization": f"Bearer {token}"} + "/user/photos", headers={"Authorization": f"Bearer {token}"} + ) + assert response.status_code in (401, 403), ( + f"Expected 401 or 403, got {response.status_code}" ) - assert response.status_code in (401, 403), f"Expected 401 or 403, got {response.status_code}" finally: async with engine.begin() as conn: - await conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) + await conn.execute( + text("DELETE FROM users WHERE id = :uid"), {"uid": user_id} + ) + async def test_rate_limiting(client): """Test that multiple requests within a short timeframe hit rate limits.""" async with engine.begin() as conn: uq = user_queries.AsyncQuerier(conn) user = await uq.create_user( - email=f"rate-{uuid.uuid4()}@test.com", - hashed_password="hash" + email=f"rate-{uuid.uuid4()}@test.com", hashed_password="hash" ) user_id = user.id session_id = uuid.uuid4() @@ -133,7 +141,7 @@ async def test_rate_limiting(client): absolute_expires_at=datetime.now(timezone.utc) + timedelta(days=30), blocked=False, ttl=3600, - last_active=datetime.now(timezone.utc) + last_active=datetime.now(timezone.utc), ) payload = { @@ -147,15 +155,18 @@ async def test_rate_limiting(client): # We test with enough requests to hit the 20/min limit for _ in range(25): res = await client.get( - "/user/photos", - headers={"Authorization": f"Bearer {token}"} + "/user/photos", headers={"Authorization": f"Bearer {token}"} ) responses.append(res.status_code) - assert 429 in responses, "Expected to hit rate limit (429) after multiple rapid requests" + assert 429 in responses, ( + "Expected to hit rate limit (429) after multiple rapid requests" + ) finally: async with engine.begin() as conn: - await conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) + await conn.execute( + text("DELETE FROM users WHERE id = :uid"), {"uid": user_id} + ) try: redis = RedisClient.get_instance() await redis._client.delete("rate_limit:/user/photos:127.0.0.1") @@ -163,13 +174,13 @@ async def test_rate_limiting(client): except Exception: pass + async def test_fast_path_rejects_idle_expired_cached_session(client): """Cached session past idle timeout but before absolute → 401.""" async with engine.begin() as conn: uq = user_queries.AsyncQuerier(conn) user = await uq.create_user( - email=f"idle-expired-{uuid.uuid4()}@test.com", - hashed_password="hash" + email=f"idle-expired-{uuid.uuid4()}@test.com", hashed_password="hash" ) user_id = user.id session_id = uuid.uuid4() @@ -181,7 +192,7 @@ async def test_fast_path_rejects_idle_expired_cached_session(client): session_id=session_id, user_id=user_id, email="idle@test.com", - idle_expires_at=now - timedelta(hours=1), # expired + idle_expires_at=now - timedelta(hours=1), # expired absolute_expires_at=now + timedelta(days=30), # still valid blocked=False, ttl=3600, @@ -196,14 +207,15 @@ async def test_fast_path_rejects_idle_expired_cached_session(client): try: response = await client.get( - "/user/photos", - headers={"Authorization": f"Bearer {token}"} + "/user/photos", headers={"Authorization": f"Bearer {token}"} ) assert response.status_code == 401 assert "expired" in response.json()["detail"].lower() finally: async with engine.begin() as conn: - await conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) + await conn.execute( + text("DELETE FROM users WHERE id = :uid"), {"uid": user_id} + ) async def test_fast_path_rejects_absolute_expired_cached_session(client): @@ -211,8 +223,7 @@ async def test_fast_path_rejects_absolute_expired_cached_session(client): async with engine.begin() as conn: uq = user_queries.AsyncQuerier(conn) user = await uq.create_user( - email=f"abs-expired-{uuid.uuid4()}@test.com", - hashed_password="hash" + email=f"abs-expired-{uuid.uuid4()}@test.com", hashed_password="hash" ) user_id = user.id session_id = uuid.uuid4() @@ -224,7 +235,7 @@ async def test_fast_path_rejects_absolute_expired_cached_session(client): session_id=session_id, user_id=user_id, email="abs@test.com", - idle_expires_at=now + timedelta(days=7), # still valid + idle_expires_at=now + timedelta(days=7), # still valid absolute_expires_at=now - timedelta(hours=1), # expired blocked=False, ttl=3600, @@ -239,11 +250,12 @@ async def test_fast_path_rejects_absolute_expired_cached_session(client): try: response = await client.get( - "/user/photos", - headers={"Authorization": f"Bearer {token}"} + "/user/photos", headers={"Authorization": f"Bearer {token}"} ) assert response.status_code == 401 assert "expired" in response.json()["detail"].lower() finally: async with engine.begin() as conn: - await conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) + await conn.execute( + text("DELETE FROM users WHERE id = :uid"), {"uid": user_id} + ) diff --git a/tests/unit/test_auth_email_otp.py b/tests/unit/test_auth_email_otp.py index d21a9e64..f19a71ce 100644 --- a/tests/unit/test_auth_email_otp.py +++ b/tests/unit/test_auth_email_otp.py @@ -6,30 +6,37 @@ from app.service.users import AuthService from app.schema.request.mobile.auth import MobileRegisterRequest, RegisterVerifyRequest + @pytest.fixture def mock_user_querier() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_device_querier() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_session_querier() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_face_embedding_service() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_redis() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_refresh_token_querier() -> AsyncMock: return AsyncMock() + @pytest.fixture def auth_service( mock_user_querier: AsyncMock, @@ -46,6 +53,7 @@ def auth_service( refresh_token_querier=mock_refresh_token_querier, ) + @pytest.mark.asyncio @patch("app.service.users.settings.environment", "production") @patch("app.service.users.NatsClient.js_publish") @@ -78,6 +86,7 @@ async def test_mobile_register_sends_otp( assert payload["email"] == "test@example.com" assert "otp" in payload + @pytest.mark.asyncio async def test_verify_mobile_register_success( auth_service: AuthService, @@ -93,12 +102,12 @@ async def test_verify_mobile_register_success( device_name="iPhone", device_type="iOS", physical_device_id=device_id, - otp="123456" + otp="123456", ) mock_redis.get.side_effect = [ "123456", - json.dumps({"hashed_password": "hashed_pass"}) + json.dumps({"hashed_password": "hashed_pass"}), ] mock_user = AsyncMock() @@ -124,14 +133,20 @@ async def _empty_evict(*, user_id, id, session_limit): # FIX 2: Mock create_device to return a truthy device mock_device_querier.create_device.return_value = AsyncMock() - with patch("app.service.users.SessionService.cache_session_for_auth", new_callable=AsyncMock): + with patch( + "app.service.users.SessionService.cache_session_for_auth", + new_callable=AsyncMock, + ): res = await auth_service.verify_mobile_register(redis=mock_redis, req=req) assert res.is_new_user is True assert res.user_id == mock_user.id - mock_user_querier.create_user.assert_called_once_with(email="test@example.com", hashed_password="hashed_pass") + mock_user_querier.create_user.assert_called_once_with( + email="test@example.com", hashed_password="hashed_pass" + ) assert mock_redis.delete.call_count == 2 + @pytest.mark.asyncio async def test_mobile_register_resend_otp_success( auth_service: AuthService, @@ -141,9 +156,15 @@ async def test_mobile_register_resend_otp_success( mock_redis.get.return_value = '{"hashed_password": "fake"}' mock_redis.incr.return_value = 1 - with patch("app.service.users.settings.environment", "production"), \ - patch("app.service.users.NatsClient.js_publish", new_callable=AsyncMock) as mock_publish: - res = await auth_service.mobile_register_resend_otp(redis=mock_redis, email=email) + with ( + patch("app.service.users.settings.environment", "production"), + patch( + "app.service.users.NatsClient.js_publish", new_callable=AsyncMock + ) as mock_publish, + ): + res = await auth_service.mobile_register_resend_otp( + redis=mock_redis, email=email + ) assert res.status == "pending_verification" assert res.message == "New OTP sent to email" @@ -152,12 +173,14 @@ async def test_mobile_register_resend_otp_success( mock_redis.set.assert_called_with(f"otp:{email}", ANY, expire=600) mock_publish.assert_called_once() + @pytest.mark.asyncio async def test_mobile_register_resend_otp_not_found( auth_service: AuthService, mock_redis: AsyncMock, ) -> None: from fastapi import HTTPException + email = "test@example.com" mock_redis.incr.return_value = 1 mock_redis.get.return_value = None diff --git a/tests/unit/test_auth_service.py b/tests/unit/test_auth_service.py index 622d074e..5193df86 100644 --- a/tests/unit/test_auth_service.py +++ b/tests/unit/test_auth_service.py @@ -17,10 +17,12 @@ from app.core.securite import hash_password from app.schema.request.mobile.auth import MobileLoginRequest, MobileRegisterRequest + async def _empty_async_iter(): return yield # pragma: no cover + # --------------------------------------------------------------------------- # Factories # --------------------------------------------------------------------------- @@ -55,8 +57,12 @@ def _make_session( s.id = session_id or uuid.uuid4() s.user_id = user_id or uuid.uuid4() s.device_id = uuid.uuid4() - s.idle_expires_at = idle_expires_at or datetime.now(timezone.utc) + timedelta(days=7) - s.absolute_expires_at = absolute_expires_at or datetime.now(timezone.utc) + timedelta(days=30) + s.idle_expires_at = idle_expires_at or datetime.now(timezone.utc) + timedelta( + days=7 + ) + s.absolute_expires_at = absolute_expires_at or datetime.now( + timezone.utc + ) + timedelta(days=30) s.last_active = datetime.now(timezone.utc) return s @@ -78,7 +84,7 @@ def _make_login_request( return MobileLoginRequest( email=email, password=password, - physical_device_id=uuid.uuid4(), # was: device_id + physical_device_id=uuid.uuid4(), # was: device_id device_name="iPhone 15", device_type="ios", ) @@ -92,11 +98,12 @@ def _make_register_request( return MobileRegisterRequest( email=email, password=password, - physical_device_id=uuid.uuid4(), # was: device_id + physical_device_id=uuid.uuid4(), # was: device_id device_name="iPhone 15", device_type="ios", ) + # --------------------------------------------------------------------------- # Fixtures # --------------------------------------------------------------------------- @@ -105,6 +112,7 @@ def _make_register_request( @pytest.fixture def user_querier() -> AsyncMock: from db.generated import user as user_queries + q = MagicMock(spec=user_queries.AsyncQuerier) q.get_user_by_email = AsyncMock(return_value=None) q.get_user_by_id_for_update = AsyncMock(return_value=None) @@ -118,10 +126,11 @@ def user_querier() -> AsyncMock: @pytest.fixture def device_querier() -> AsyncMock: from db.generated import devices as device_queries + q = MagicMock(spec=device_queries.AsyncQuerier) q.get_device_by_id = AsyncMock(return_value=None) q.get_device_by_id_any = AsyncMock(return_value=None) - q.get_device_by_physical_id = AsyncMock(return_value=None) # new + q.get_device_by_physical_id = AsyncMock(return_value=None) # new q.create_device = AsyncMock(return_value=_make_device()) q.activate_device = AsyncMock() return q @@ -130,6 +139,7 @@ def device_querier() -> AsyncMock: @pytest.fixture def session_querier() -> AsyncMock: from db.generated import session as session_queries + q = MagicMock(spec=session_queries.AsyncQuerier) q.lock_user_sessions = AsyncMock(return_value=None) @@ -149,13 +159,16 @@ async def _default_empty_evict(*, user_id, id, session_limit): @pytest.fixture def face_service() -> AsyncMock: from app.service.face_embedding import FaceEmbeddingService + svc = MagicMock(spec=FaceEmbeddingService) svc.compute_average_embedding = AsyncMock(return_value=[0.1] * 512) return svc + @pytest.fixture def refresh_token_querier() -> AsyncMock: from db.generated import refresh_token as refresh_token_queries + q = MagicMock(spec=refresh_token_queries.AsyncQuerier) q.get_refresh_token_by_hash_for_update = AsyncMock(return_value=None) q.get_refresh_token_by_jti = AsyncMock(return_value=None) @@ -165,6 +178,7 @@ def refresh_token_querier() -> AsyncMock: q.mark_refresh_token_used = AsyncMock() return q + @pytest.fixture def redis() -> AsyncMock: r = MagicMock() @@ -313,7 +327,11 @@ async def test_blocked_user_raises_403( class TestSessionLimit: @pytest.mark.asyncio async def test_at_cap_evicts_oldest_and_succeeds( - self, auth_service, user_querier, session_querier, redis, + self, + auth_service, + user_querier, + session_querier, + redis, ) -> None: user = _make_user() user_querier.get_user_by_email.return_value = user @@ -343,13 +361,14 @@ async def test_within_session_limit_succeeds( user = _make_user() user_querier.get_user_by_email.return_value = user user_querier.get_user_by_id_for_update.return_value = user - session_querier.list_sessions_by_user = MagicMock(return_value=_empty_async_iter()) + session_querier.list_sessions_by_user = MagicMock( + return_value=_empty_async_iter() + ) result = await auth_service.mobile_login(redis, _make_login_request()) assert result.access_token session_querier.delete_session_by_id.assert_not_called() - @pytest.mark.asyncio async def test_multiple_new_devices_at_cap_evict_exact_overflow( self, @@ -384,8 +403,6 @@ async def _evict(*, user_id, id, session_limit): assert redis.delete.call_count == 3 - - # =========================================================================== # 4. Logout # =========================================================================== @@ -467,10 +484,16 @@ async def test_valid_refresh_returns_new_tokens( redis.set.assert_called_once() call_args = redis.set.call_args cache_key = call_args.args[0] if call_args.args else call_args.kwargs.get("key") - cache_value = call_args.args[1] if len(call_args.args) > 1 else call_args.kwargs.get("value") + cache_value = ( + call_args.args[1] + if len(call_args.args) > 1 + else call_args.kwargs.get("value") + ) assert cache_key == f"refresh_retry:{token_hash}" - assert "access_token" not in cache_value # plaintext JSON would contain this literal key; encrypted payload must not + assert ( + "access_token" not in cache_value + ) # plaintext JSON would contain this literal key; encrypted payload must not decrypted = decrypt_refresh_cache_payload(cache_value) assert result.access_token in decrypted @@ -487,7 +510,7 @@ async def test_expired_session_raises_401( past_session = _make_session( idle_expires_at=datetime.now(timezone.utc) - timedelta(days=1), - absolute_expires_at=datetime.now(timezone.utc) - timedelta(days=1) + absolute_expires_at=datetime.now(timezone.utc) - timedelta(days=1), ) session_querier.get_session_by_id.return_value = past_session @@ -564,11 +587,13 @@ async def test_refresh_rejects_idle_expired_session( now = datetime.now(timezone.utc) session = _make_session( - idle_expires_at=now - timedelta(hours=1), # expired + idle_expires_at=now - timedelta(hours=1), # expired absolute_expires_at=now + timedelta(days=30), # valid ) session_querier.get_session_by_id.return_value = session - user_querier.get_user_by_id.return_value = _make_user(user_id=session.user_id) + user_querier.get_user_by_id.return_value = _make_user( + user_id=session.user_id + ) raw_token = create_raw_refresh_token() row = MagicMock() @@ -577,7 +602,9 @@ async def test_refresh_rejects_idle_expired_session( row.used_at = None row.family_id = uuid.uuid4() row.session_id = session.id - refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = ( + row + ) refresh_token_querier.mark_refresh_token_used.return_value = row with pytest.raises(HTTPException) as exc_info: @@ -599,11 +626,13 @@ async def test_refresh_rejects_absolute_expired_session( now = datetime.now(timezone.utc) session = _make_session( - idle_expires_at=now + timedelta(days=7), # valid + idle_expires_at=now + timedelta(days=7), # valid absolute_expires_at=now - timedelta(hours=1), # expired ) session_querier.get_session_by_id.return_value = session - user_querier.get_user_by_id.return_value = _make_user(user_id=session.user_id) + user_querier.get_user_by_id.return_value = _make_user( + user_id=session.user_id + ) raw_token = create_raw_refresh_token() row = MagicMock() @@ -612,7 +641,9 @@ async def test_refresh_rejects_absolute_expired_session( row.used_at = None row.family_id = uuid.uuid4() row.session_id = session.id - refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = ( + row + ) refresh_token_querier.mark_refresh_token_used.return_value = row with pytest.raises(HTTPException) as exc_info: @@ -656,6 +687,7 @@ async def test_returns_closest_user_match( assert result.user_id == row.id assert result.distance == 0.25 + class TestBlockedUserRaceCondition: @pytest.mark.asyncio async def test_blocked_between_initial_check_and_lock_is_caught( @@ -670,9 +702,7 @@ async def test_blocked_between_initial_check_and_lock_is_caught( sees blocked=True. Login must still be rejected, and no session may be created.""" unblocked_snapshot = _make_user(blocked=False) - blocked_after_lock = _make_user( - user_id=unblocked_snapshot.id, blocked=True - ) + blocked_after_lock = _make_user(user_id=unblocked_snapshot.id, blocked=True) user_querier.get_user_by_email.return_value = unblocked_snapshot user_querier.get_user_by_id_for_update.return_value = blocked_after_lock @@ -717,6 +747,7 @@ async def test_missing_user_at_lock_time_raises_401( assert exc_info.value.status_code == 401 + # =========================================================================== # 7. check_rate_limit fail-open behavior # =========================================================================== @@ -755,7 +786,9 @@ async def test_real_rate_limit_rejection_still_raises( redis.incr = AsyncMock(return_value=999) # way over any reasonable limit with pytest.raises(HTTPException) as exc_info: - await auth_service.check_rate_limit(redis, "rate:test:key", max_requests=5, window_seconds=60) + await auth_service.check_rate_limit( + redis, "rate:test:key", max_requests=5, window_seconds=60 + ) assert exc_info.value.status_code == 429 @@ -782,11 +815,11 @@ async def test_takes_lock_before_mutating( ) call_order = [] - user_querier.get_user_by_id_for_update.side_effect = ( - lambda *a, **kw: call_order.append("lock") or target + user_querier.get_user_by_id_for_update.side_effect = lambda *a, **kw: ( + call_order.append("lock") or target ) - user_querier.set_user_blocked.side_effect = ( - lambda *a, **kw: call_order.append("mutate") or _make_user(user_id=target.id, blocked=True) + user_querier.set_user_blocked.side_effect = lambda *a, **kw: ( + call_order.append("mutate") or _make_user(user_id=target.id, blocked=True) ) await auth_service.block_user(redis=redis, user_id=target.id) @@ -819,7 +852,9 @@ async def _sessions(*, user_id): await auth_service.block_user(redis=redis, user_id=target.id) - session_querier.delete_all_user_sessions.assert_called_once_with(user_id=target.id) + session_querier.delete_all_user_sessions.assert_called_once_with( + user_id=target.id + ) assert redis.delete.call_count == len(session_ids) @pytest.mark.asyncio @@ -856,7 +891,9 @@ async def test_unblocks_successfully( result = await auth_service.unblock_user(user_id=target_id) assert result.blocked is False - user_querier.set_user_blocked.assert_called_once_with(blocked=False, id=target_id) + user_querier.set_user_blocked.assert_called_once_with( + blocked=False, id=target_id + ) @pytest.mark.asyncio async def test_missing_user_raises_404( @@ -888,11 +925,11 @@ async def test_takes_lock_before_deleting( target = _make_user() call_order = [] - user_querier.get_user_by_id_for_update.side_effect = ( - lambda *a, **kw: call_order.append("lock") or target + user_querier.get_user_by_id_for_update.side_effect = lambda *a, **kw: ( + call_order.append("lock") or target ) - user_querier.delete_user.side_effect = ( - lambda *a, **kw: call_order.append("delete") + user_querier.delete_user.side_effect = lambda *a, **kw: call_order.append( + "delete" ) await auth_service.delete_user(redis=redis, user_id=target.id) @@ -922,7 +959,9 @@ async def _sessions(*, user_id): await auth_service.delete_user(redis=redis, user_id=target.id) - session_querier.delete_all_user_sessions.assert_called_once_with(user_id=target.id) + session_querier.delete_all_user_sessions.assert_called_once_with( + user_id=target.id + ) assert redis.delete.call_count == len(session_ids) @pytest.mark.asyncio @@ -1014,7 +1053,8 @@ async def test_used_token_within_grace_but_blocked_user_raises_403( session = _make_session() session_querier.get_session_by_id.return_value = session user_querier.get_user_by_id.return_value = _make_user( - user_id=session.user_id, blocked=True # blocked since original rotation + user_id=session.user_id, + blocked=True, # blocked since original rotation ) raw_token = create_raw_refresh_token() @@ -1088,7 +1128,9 @@ async def test_used_token_outside_grace_revokes_session( row = MagicMock() row.id = uuid.uuid4() row.used = True - row.used_at = datetime.now(timezone.utc) - timedelta(seconds=120) # well outside grace + row.used_at = datetime.now(timezone.utc) - timedelta( + seconds=120 + ) # well outside grace row.family_id = uuid.uuid4() row.session_id = session.id refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row @@ -1166,7 +1208,11 @@ async def test_new_rotation_caches_encrypted_not_plaintext( redis.set.assert_called_once() call_args = redis.set.call_args cache_key = call_args.args[0] if call_args.args else call_args.kwargs.get("key") - cache_value = call_args.args[1] if len(call_args.args) > 1 else call_args.kwargs.get("value") + cache_value = ( + call_args.args[1] + if len(call_args.args) > 1 + else call_args.kwargs.get("value") + ) assert cache_key == f"refresh_retry:{token_hash}" # plaintext JSON would contain this literal substring; encrypted payload must not @@ -1176,6 +1222,7 @@ async def test_new_rotation_caches_encrypted_not_plaintext( decrypted = decrypt_refresh_cache_payload(cache_value) assert result.access_token in decrypted + class TestValidateSession: @pytest.mark.asyncio async def test_validate_session_false_when_idle_expired( diff --git a/tests/unit/test_direct_uploads.py b/tests/unit/test_direct_uploads.py index 20612b14..6c6020fe 100644 --- a/tests/unit/test_direct_uploads.py +++ b/tests/unit/test_direct_uploads.py @@ -19,25 +19,35 @@ def mock_upload_request_group_querier(): return AsyncMock() + @pytest.fixture def mock_upload_request_querier(): return AsyncMock() + @pytest.fixture def mock_upload_request_photo_querier(): return AsyncMock() + @pytest.fixture def mock_photo_querier(): return AsyncMock() + @pytest.fixture def mock_staged_upload_storage(): mock = AsyncMock() - mock.store_staging_object.return_value = StoredObject(storage_key="test_storage_key", content_type="image/jpeg", file_name="photo.jpg") - mock.create_presigned_staging_upload.return_value = ("staging/upload-requests/req1/photo1.jpg", "https://minio.local/signed-url") + mock.store_staging_object.return_value = StoredObject( + storage_key="test_storage_key", content_type="image/jpeg", file_name="photo.jpg" + ) + mock.create_presigned_staging_upload.return_value = ( + "staging/upload-requests/req1/photo1.jpg", + "https://minio.local/signed-url", + ) return mock + @pytest.fixture def mock_staff_drive_service(): mock = AsyncMock() @@ -45,14 +55,17 @@ def mock_staff_drive_service(): mock.staff_user_querier = AsyncMock() return mock + @pytest.fixture def mock_staff_notifications_service(): return AsyncMock() + @pytest.fixture def mock_audit_service(): return AsyncMock() + @pytest.fixture def upload_requests_service( mock_upload_request_group_querier, @@ -75,6 +88,7 @@ def upload_requests_service( audit_service=mock_audit_service, ) + @pytest.fixture def mock_staff_user(): return StaffUser( @@ -89,11 +103,22 @@ def mock_staff_user(): def _make_group(group_id, event_id, requested_by_id, **overrides): defaults = dict( - id=group_id, event_id=event_id, folder_id=None, requested_by=requested_by_id, - approved_by=None, status="pending", total_photo_count=0, batch_count=0, - created_at=datetime.now(timezone.utc), approved_at=None, rejection_reason=None, - processing_status="completed", processed_photo_count=0, failed_photo_count=0, - error_message=None, source="direct", + id=group_id, + event_id=event_id, + folder_id=None, + requested_by=requested_by_id, + approved_by=None, + status="pending", + total_photo_count=0, + batch_count=0, + created_at=datetime.now(timezone.utc), + approved_at=None, + rejection_reason=None, + processing_status="completed", + processed_photo_count=0, + failed_photo_count=0, + error_message=None, + source="direct", ) defaults.update(overrides) return UploadRequestGroup(**defaults) @@ -101,9 +126,17 @@ def _make_group(group_id, event_id, requested_by_id, **overrides): def _make_request(request_id, event_id, requested_by_id, group_id, **overrides): defaults = dict( - id=request_id, event_id=event_id, drive_file_id=None, requested_by=requested_by_id, - approved_by=None, status="pending", created_at=datetime.now(timezone.utc), - approved_at=None, photo_count=1, rejection_reason=None, group_id=group_id, + id=request_id, + event_id=event_id, + drive_file_id=None, + requested_by=requested_by_id, + approved_by=None, + status="pending", + created_at=datetime.now(timezone.utc), + approved_at=None, + photo_count=1, + rejection_reason=None, + group_id=group_id, source="direct", ) defaults.update(overrides) @@ -112,11 +145,21 @@ def _make_request(request_id, event_id, requested_by_id, group_id, **overrides): def _make_photo(photo_id, request_id, **overrides): defaults = dict( - id=photo_id, upload_request_id=request_id, drive_file_id=None, file_name="a.jpg", - mime_type="image/jpeg", size_bytes=1000, staging_storage_key="staging/x.jpg", - final_storage_key=None, taken_at=None, day_number=None, visibility="private", - status="staged", created_at=datetime.now(timezone.utc), - source="direct", transfer_status="pending_upload", + id=photo_id, + upload_request_id=request_id, + drive_file_id=None, + file_name="a.jpg", + mime_type="image/jpeg", + size_bytes=1000, + staging_storage_key="staging/x.jpg", + final_storage_key=None, + taken_at=None, + day_number=None, + visibility="private", + status="staged", + created_at=datetime.now(timezone.utc), + source="direct", + transfer_status="pending_upload", ) defaults.update(overrides) return UploadRequestPhoto(**defaults) @@ -130,12 +173,17 @@ async def test_create_direct_group_sets_source_and_completed_processing( ): event_id = uuid.uuid4() group_id = uuid.uuid4() - mock_upload_request_group_querier.create_upload_request_group.return_value = _make_group( - group_id, event_id, mock_staff_user.id, + mock_upload_request_group_querier.create_upload_request_group.return_value = ( + _make_group( + group_id, + event_id, + mock_staff_user.id, + ) ) group = await upload_requests_service.create_direct_group( - event_id=event_id, requested_by=mock_staff_user, + event_id=event_id, + requested_by=mock_staff_user, ) assert group.id == group_id @@ -160,22 +208,37 @@ async def test_register_direct_batch_creates_pending_photos_and_returns_urls( request_id = uuid.uuid4() photo_id = uuid.uuid4() - mock_upload_request_group_querier.get_upload_request_group_by_id.return_value = _make_group( - group_id, event_id, mock_staff_user.id, + mock_upload_request_group_querier.get_upload_request_group_by_id.return_value = ( + _make_group( + group_id, + event_id, + mock_staff_user.id, + ) ) mock_upload_request_querier.create_upload_request.return_value = _make_request( - request_id, event_id, mock_staff_user.id, group_id, + request_id, + event_id, + mock_staff_user.id, + group_id, ) mock_upload_request_photo_querier.create_direct_upload_request_photo.return_value = _make_photo( - photo_id, request_id, staging_storage_key="staging/upload-requests/req1/photo1.jpg", + photo_id, + request_id, + staging_storage_key="staging/upload-requests/req1/photo1.jpg", ) results = await upload_requests_service.register_direct_batch( group_id=group_id, - files=[DirectFileInput( - file_name="a.jpg", mime_type="image/jpeg", size_bytes=1000, - taken_at=None, day_number=None, visibility="private", - )], + files=[ + DirectFileInput( + file_name="a.jpg", + mime_type="image/jpeg", + size_bytes=1000, + taken_at=None, + day_number=None, + visibility="private", + ) + ], requested_by=mock_staff_user, ) @@ -185,7 +248,8 @@ async def test_register_direct_batch_creates_pending_photos_and_returns_urls( assert url == "https://minio.local/signed-url" mock_staged_upload_storage.create_presigned_staging_upload.assert_awaited() mock_upload_request_group_querier.increment_upload_request_group_counts.assert_awaited_once_with( - id=group_id, total_photo_count=1, + id=group_id, + total_photo_count=1, ) @@ -196,16 +260,24 @@ async def test_register_direct_batch_rejects_oversized_batch( mock_staff_user, ): group_id = uuid.uuid4() - from app.core.exceptions import AppException files = [ - DirectFileInput(file_name=f"{i}.jpg", mime_type="image/jpeg", size_bytes=1000, taken_at=None, day_number=None, visibility="private") + DirectFileInput( + file_name=f"{i}.jpg", + mime_type="image/jpeg", + size_bytes=1000, + taken_at=None, + day_number=None, + visibility="private", + ) for i in range(201) ] with pytest.raises(Exception): await upload_requests_service.register_direct_batch( - group_id=group_id, files=files, requested_by=mock_staff_user, + group_id=group_id, + files=files, + requested_by=mock_staff_user, ) mock_upload_request_group_querier.get_upload_request_group_by_id.assert_not_awaited() @@ -222,19 +294,28 @@ async def test_confirm_direct_upload_success_updates_transfer_status( photo_id = uuid.uuid4() request_id = uuid.uuid4() existing_photo = _make_photo(photo_id, request_id, transfer_status="pending_upload") - mock_upload_request_photo_querier.get_upload_request_photo_by_id.return_value = existing_photo - mock_staged_upload_storage.stat_staging_object.return_value = ObjectStat(size=1000, content_type="image/jpeg") + mock_upload_request_photo_querier.get_upload_request_photo_by_id.return_value = ( + existing_photo + ) + mock_staged_upload_storage.stat_staging_object.return_value = ObjectStat( + size=1000, content_type="image/jpeg" + ) mock_upload_request_photo_querier.confirm_upload_request_photo_transfer.return_value = _make_photo( - photo_id, request_id, transfer_status="uploaded", + photo_id, + request_id, + transfer_status="uploaded", ) result = await upload_requests_service.confirm_direct_upload( - photo_id=photo_id, requested_by=mock_staff_user, + photo_id=photo_id, + requested_by=mock_staff_user, ) assert result.id == photo_id mock_upload_request_photo_querier.confirm_upload_request_photo_transfer.assert_awaited_once_with( - id=photo_id, size_bytes=1000, mime_type="image/jpeg", + id=photo_id, + size_bytes=1000, + mime_type="image/jpeg", ) @@ -248,18 +329,25 @@ async def test_confirm_direct_upload_marks_failed_when_object_missing( photo_id = uuid.uuid4() request_id = uuid.uuid4() existing_photo = _make_photo(photo_id, request_id, transfer_status="pending_upload") - mock_upload_request_photo_querier.get_upload_request_photo_by_id.return_value = existing_photo + mock_upload_request_photo_querier.get_upload_request_photo_by_id.return_value = ( + existing_photo + ) mock_staged_upload_storage.stat_staging_object.return_value = None mock_upload_request_photo_querier.fail_upload_request_photo_transfer.return_value = _make_photo( - photo_id, request_id, transfer_status="failed", + photo_id, + request_id, + transfer_status="failed", ) with pytest.raises(Exception): await upload_requests_service.confirm_direct_upload( - photo_id=photo_id, requested_by=mock_staff_user, + photo_id=photo_id, + requested_by=mock_staff_user, ) - mock_upload_request_photo_querier.fail_upload_request_photo_transfer.assert_awaited_once_with(id=photo_id) + mock_upload_request_photo_querier.fail_upload_request_photo_transfer.assert_awaited_once_with( + id=photo_id + ) @pytest.mark.asyncio @@ -273,9 +361,14 @@ async def test_approve_request_blocked_when_photo_not_fully_uploaded( photo_id = uuid.uuid4() mock_upload_request_querier.get_upload_request_by_id.return_value = _make_request( - request_id, uuid.uuid4(), mock_staff_user.id, None, + request_id, + uuid.uuid4(), + mock_staff_user.id, + None, + ) + not_uploaded_photo = _make_photo( + photo_id, request_id, transfer_status="pending_upload" ) - not_uploaded_photo = _make_photo(photo_id, request_id, transfer_status="pending_upload") async def _photos_iter(upload_request_id): yield not_uploaded_photo @@ -284,7 +377,8 @@ async def _photos_iter(upload_request_id): with pytest.raises(Exception) as exc_info: await upload_requests_service.approve_request( - request_id=request_id, approved_by=mock_staff_user, + request_id=request_id, + approved_by=mock_staff_user, ) assert "have not finished uploading" in str(exc_info.value) @@ -304,30 +398,58 @@ async def test_resume_direct_group_reissues_urls_for_pending_and_failed_only( failed_photo_id = uuid.uuid4() uploaded_photo_id = uuid.uuid4() - mock_upload_request_group_querier.get_upload_request_group_by_id.return_value = _make_group( - group_id, event_id, mock_staff_user.id, total_photo_count=2, batch_count=1, failed_photo_count=1, + mock_upload_request_group_querier.get_upload_request_group_by_id.return_value = ( + _make_group( + group_id, + event_id, + mock_staff_user.id, + total_photo_count=2, + batch_count=1, + failed_photo_count=1, + ) ) async def _requests_iter(group_id): - yield _make_request(request_id, event_id, mock_staff_user.id, group_id, photo_count=2) + yield _make_request( + request_id, event_id, mock_staff_user.id, group_id, photo_count=2 + ) mock_upload_request_querier.list_upload_requests_by_group_id = _requests_iter - failed_photo = _make_photo(failed_photo_id, request_id, file_name="fail.jpg", staging_storage_key="staging/fail.jpg", transfer_status="failed") - uploaded_photo = _make_photo(uploaded_photo_id, request_id, file_name="ok.jpg", staging_storage_key="staging/ok.jpg", transfer_status="uploaded") + failed_photo = _make_photo( + failed_photo_id, + request_id, + file_name="fail.jpg", + staging_storage_key="staging/fail.jpg", + transfer_status="failed", + ) + uploaded_photo = _make_photo( + uploaded_photo_id, + request_id, + file_name="ok.jpg", + staging_storage_key="staging/ok.jpg", + transfer_status="uploaded", + ) async def _photos_iter(dollar_1): for p in [failed_photo, uploaded_photo]: yield p mock_upload_request_photo_querier.list_upload_request_photos_by_upload_request_ids = _photos_iter - mock_staged_upload_storage.create_presigned_staging_upload.return_value = ("staging/fail.jpg", "https://minio.local/resumed") + mock_staged_upload_storage.create_presigned_staging_upload.return_value = ( + "staging/fail.jpg", + "https://minio.local/resumed", + ) mock_upload_request_photo_querier.reset_upload_request_photo_transfer_to_pending.return_value = _make_photo( - failed_photo_id, request_id, file_name="fail.jpg", transfer_status="pending_upload", + failed_photo_id, + request_id, + file_name="fail.jpg", + transfer_status="pending_upload", ) results = await upload_requests_service.resume_direct_group( - group_id=group_id, requested_by=mock_staff_user, + group_id=group_id, + requested_by=mock_staff_user, ) assert len(results) == 1 @@ -352,34 +474,56 @@ async def test_approve_request_publishes_drive_sync_event_for_direct_photo_only( photo_id = uuid.uuid4() mock_upload_request_querier.get_upload_request_by_id.return_value = _make_request( - request_id, event_id, mock_staff_user.id, None, + request_id, + event_id, + mock_staff_user.id, + None, ) async def _photos_iter(upload_request_id): - yield _make_photo(photo_id, request_id, transfer_status="uploaded", source="direct") + yield _make_photo( + photo_id, request_id, transfer_status="uploaded", source="direct" + ) mock_upload_request_photo_querier.list_upload_request_photos_by_upload_request_id = _photos_iter mock_staged_upload_storage.promote_to_final.return_value = "events/e1/p1.jpg" mock_photo_querier.create_photo.return_value = Photo( - id=photo_id, event_id=event_id, uploaded_by=None, storage_key="events/e1/p1.jpg", - taken_at=None, day_number=None, visibility="private", status="pending", - created_at=datetime.now(timezone.utc), drive_file_id=None, drive_synced_at=None, - source="direct", storage_cleaned_at=None, + id=photo_id, + event_id=event_id, + uploaded_by=None, + storage_key="events/e1/p1.jpg", + taken_at=None, + day_number=None, + visibility="private", + status="pending", + created_at=datetime.now(timezone.utc), + drive_file_id=None, + drive_synced_at=None, + source="direct", + storage_cleaned_at=None, ) mock_upload_request_photo_querier.update_upload_request_photo_approval.return_value = _make_photo( - photo_id, request_id, source="direct", transfer_status="uploaded", + photo_id, + request_id, + source="direct", + transfer_status="uploaded", ) mock_upload_request_querier.approve_upload_request.return_value = _make_request( - request_id, event_id, mock_staff_user.id, None, + request_id, + event_id, + mock_staff_user.id, + None, ) with patch("app.service.upload_requests.NatsClient.js_publish") as mock_publish: await upload_requests_service.approve_request( - request_id=request_id, approved_by=mock_staff_user, + request_id=request_id, + approved_by=mock_staff_user, ) published_subjects = [call.args[0] for call in mock_publish.call_args_list] from app.infra.nats import NatsSubjects + assert NatsSubjects.PHOTO_DRIVE_SYNC_REQUESTED in published_subjects @@ -392,14 +536,21 @@ async def test_fail_direct_upload_marks_transfer_failed( photo_id = uuid.uuid4() request_id = uuid.uuid4() existing_photo = _make_photo(photo_id, request_id, transfer_status="pending_upload") - mock_upload_request_photo_querier.get_upload_request_photo_by_id.return_value = existing_photo + mock_upload_request_photo_querier.get_upload_request_photo_by_id.return_value = ( + existing_photo + ) mock_upload_request_photo_querier.fail_upload_request_photo_transfer.return_value = _make_photo( - photo_id, request_id, transfer_status="failed", + photo_id, + request_id, + transfer_status="failed", ) result = await upload_requests_service.fail_direct_upload( - photo_id=photo_id, requested_by=mock_staff_user, + photo_id=photo_id, + requested_by=mock_staff_user, ) assert result.id == photo_id - mock_upload_request_photo_querier.fail_upload_request_photo_transfer.assert_awaited_once_with(id=photo_id) + mock_upload_request_photo_querier.fail_upload_request_photo_transfer.assert_awaited_once_with( + id=photo_id + ) diff --git a/tests/unit/test_enroll_security.py b/tests/unit/test_enroll_security.py index d944cb03..41f5fe5e 100644 --- a/tests/unit/test_enroll_security.py +++ b/tests/unit/test_enroll_security.py @@ -54,7 +54,9 @@ def _make_upload_file( mock_file.filename = filename mock_file.content_type = content_type mock_file.headers = headers - mock_file.read = AsyncMock(side_effect=lambda n=-1: buf.read(n) if n == -1 else buf.read(n)) + mock_file.read = AsyncMock( + side_effect=lambda n=-1: buf.read(n) if n == -1 else buf.read(n) + ) mock_file.seek = AsyncMock(side_effect=lambda pos: buf.seek(pos)) return mock_file # type: ignore[return-value] @@ -86,8 +88,9 @@ def test_control_characters_are_replaced(self) -> None: def test_windows_reserved_chars_are_replaced(self) -> None: for char in r'\\/:*?"<>|': - assert char not in sanitise_filename(f"face{char}name.jpg", "jpg"), \ + assert char not in sanitise_filename(f"face{char}name.jpg", "jpg"), ( f"char {char!r} must be replaced" + ) def test_none_filename_returns_uuid_only(self) -> None: result = sanitise_filename(None, "png") @@ -170,7 +173,9 @@ def test_missing_content_type_raises_400(self) -> None: def test_unsupported_content_type_raises_400(self) -> None: with pytest.raises(HTTPException) as exc_info: - precheck_upload_headers(_make_upload_file(b"", content_type="application/pdf")) + precheck_upload_headers( + _make_upload_file(b"", content_type="application/pdf") + ) assert exc_info.value.status_code == 400 def test_content_type_with_charset_param_accepted(self) -> None: @@ -182,7 +187,9 @@ def test_content_type_with_charset_param_accepted(self) -> None: def test_oversized_content_length_raises_400(self) -> None: with pytest.raises(HTTPException) as exc_info: precheck_upload_headers( - _make_upload_file(b"", content_type="image/jpeg", content_length=MAX_IMAGE_SIZE + 1) + _make_upload_file( + b"", content_type="image/jpeg", content_length=MAX_IMAGE_SIZE + 1 + ) ) assert exc_info.value.status_code == 400 diff --git a/tests/unit/test_face_match_service.py b/tests/unit/test_face_match_service.py index 8de215b0..a4c64c7d 100644 --- a/tests/unit/test_face_match_service.py +++ b/tests/unit/test_face_match_service.py @@ -24,8 +24,11 @@ @pytest.fixture def photo_face_querier() -> AsyncMock: from db.generated import photo_faces as pf_queries + q = MagicMock(spec=pf_queries.AsyncQuerier) - q.photo_faces_photo_exists = AsyncMock(return_value=object()) # truthy → photo exists + q.photo_faces_photo_exists = AsyncMock( + return_value=object() + ) # truthy → photo exists q.photo_faces_match_exists_for_photo = AsyncMock(return_value=None) # no duplicate q.photo_faces_ensure_face_match = AsyncMock() return q @@ -34,6 +37,7 @@ def photo_face_querier() -> AsyncMock: @pytest.fixture def photo_querier() -> AsyncMock: from db.generated import photos as photo_queries + q = MagicMock(spec=photo_queries.AsyncQuerier) q.update_photo_status = AsyncMock(return_value=None) return q @@ -42,6 +46,7 @@ def photo_querier() -> AsyncMock: @pytest.fixture def user_match_service() -> AsyncMock: from app.service.users import AuthService + svc = MagicMock(spec=AuthService) svc.find_closest_user = AsyncMock() return svc @@ -50,6 +55,7 @@ def user_match_service() -> AsyncMock: @pytest.fixture def notification_service() -> AsyncMock: from app.service.user_notification import UserNotificationService + svc = MagicMock(spec=UserNotificationService) svc.create_notification = AsyncMock() return svc @@ -91,6 +97,7 @@ def _make_embedding(value: float = 0.5) -> list[float]: def _closest_match(distance: float = 0.3) -> object: from app.schema.internal.single_face_match import ClosestUserMatch + return ClosestUserMatch(user_id=uuid.uuid4(), distance=distance) @@ -194,7 +201,9 @@ async def test_face_match_stored_in_db( face_match_result = MagicMock() face_match_result.face_match_id = uuid.uuid4() - photo_face_querier.photo_faces_ensure_face_match.return_value = face_match_result + photo_face_querier.photo_faces_ensure_face_match.return_value = ( + face_match_result + ) await service.process_detected_face(job, _make_embedding(), bbox=None) @@ -214,7 +223,9 @@ async def test_photo_status_set_to_approved_on_match( face_match_result = MagicMock() face_match_result.face_match_id = uuid.uuid4() - photo_face_querier.photo_faces_ensure_face_match.return_value = face_match_result + photo_face_querier.photo_faces_ensure_face_match.return_value = ( + face_match_result + ) await service.process_detected_face(job, _make_embedding(), bbox=None) @@ -236,7 +247,9 @@ async def test_notification_sent_on_successful_match( face_match_result = MagicMock() face_match_result.face_match_id = uuid.uuid4() - photo_face_querier.photo_faces_ensure_face_match.return_value = face_match_result + photo_face_querier.photo_faces_ensure_face_match.return_value = ( + face_match_result + ) await service.process_detected_face(job, _make_embedding(), bbox=None) @@ -259,7 +272,9 @@ async def test_bbox_serialised_as_json( face_match_result = MagicMock() face_match_result.face_match_id = uuid.uuid4() - photo_face_querier.photo_faces_ensure_face_match.return_value = face_match_result + photo_face_querier.photo_faces_ensure_face_match.return_value = ( + face_match_result + ) await service.process_detected_face(job, _make_embedding(), bbox=bbox) @@ -298,7 +313,9 @@ async def test_skips_if_match_already_exists( photo_face_querier: AsyncMock, photo_querier: AsyncMock, ) -> None: - photo_face_querier.photo_faces_match_exists_for_photo.return_value = object() # exists + photo_face_querier.photo_faces_match_exists_for_photo.return_value = ( + object() + ) # exists await service.process_detected_face(job, _make_embedding(), bbox=None) @@ -337,7 +354,9 @@ async def test_no_notification_if_face_match_already_existed( result_already_existed = MagicMock() result_already_existed.face_match_id = None # already existed - photo_face_querier.photo_faces_ensure_face_match.return_value = result_already_existed + photo_face_querier.photo_faces_ensure_face_match.return_value = ( + result_already_existed + ) await service.process_detected_face(job, _make_embedding(), bbox=None) @@ -363,7 +382,9 @@ async def test_db_error_is_handled_gracefully( good_match = _closest_match(distance=0.1) user_match_service.find_closest_user.return_value = good_match - photo_face_querier.photo_faces_ensure_face_match.side_effect = SQLAlchemyError("DB down") + photo_face_querier.photo_faces_ensure_face_match.side_effect = SQLAlchemyError( + "DB down" + ) # Should not raise — worker must stay alive await service.process_detected_face(job, _make_embedding(), bbox=None) diff --git a/tests/unit/test_minio.py b/tests/unit/test_minio.py index 1587c27f..a188238c 100644 --- a/tests/unit/test_minio.py +++ b/tests/unit/test_minio.py @@ -13,6 +13,7 @@ init_minio_client, ) + @pytest.fixture def mock_minio_client(): client = AsyncMock() @@ -114,6 +115,7 @@ async def test_bucket_get_not_found(mock_minio_client): bucket = Bucket("test_bucket", "") from fastapi import HTTPException + with pytest.raises(HTTPException) as exc: await bucket.get("test.jpg") assert exc.value.status_code == 404 @@ -138,9 +140,7 @@ async def test_bucket_put_bytes(mock_minio_client): bucket = Bucket("test_bucket", "") await bucket.put_bytes( - data=b"byte_data", - object_name="byte_test.txt", - content_type="text/plain" + data=b"byte_data", object_name="byte_test.txt", content_type="text/plain" ) mock_minio_client.put_object.assert_called_once() @@ -186,6 +186,7 @@ async def test_image_bucket_invalid_extension(mock_minio_client, mock_upload_fil mock_upload_file.content_type = "application/pdf" from fastapi import HTTPException + with pytest.raises(HTTPException) as exc: await bucket.put(mock_upload_file) @@ -206,7 +207,9 @@ async def test_wa_sim_bucket_auto_name(mock_minio_client, mock_upload_file): @pytest.mark.asyncio async def test_presigned_put_url_calls_client_with_expiry(mock_minio_client): Bucket.client = mock_minio_client - mock_minio_client.presigned_put_object = AsyncMock(return_value="https://minio.local/signed") + mock_minio_client.presigned_put_object = AsyncMock( + return_value="https://minio.local/signed" + ) bucket = Bucket("test_bucket", "") url = await bucket.presigned_put_url("staging/foo.jpg", expires_seconds=1800) @@ -234,8 +237,12 @@ async def test_stat_returns_object_stat_when_present(mock_minio_client): async def test_stat_returns_none_when_object_missing(mock_minio_client): Bucket.client = mock_minio_client error = S3Error( - code="NoSuchKey", message="not found", resource="", request_id="", - host_id="", response=MagicMock(), + code="NoSuchKey", + message="not found", + resource="", + request_id="", + host_id="", + response=MagicMock(), ) mock_minio_client.stat_object = AsyncMock(side_effect=error) bucket = Bucket("test_bucket", "") diff --git a/tests/unit/test_mobile_auth_email_logging.py b/tests/unit/test_mobile_auth_email_logging.py index bbd3646a..681f9ed8 100644 --- a/tests/unit/test_mobile_auth_email_logging.py +++ b/tests/unit/test_mobile_auth_email_logging.py @@ -72,7 +72,6 @@ class FakeSessionQuerier: def __init__(self, session: FakeSession) -> None: self._session = session - async def lock_user_sessions(self, *, user_id: str) -> None: return None @@ -150,7 +149,9 @@ async def _noop_cache_session_for_auth(**_: object) -> None: monkeypatch.setattr("app.service.users.settings.environment", "production") monkeypatch.setattr("app.service.users.NatsClient.js_publish", AsyncMock()) - monkeypatch.setattr(SessionService, "cache_session_for_auth", _noop_cache_session_for_auth) + monkeypatch.setattr( + SessionService, "cache_session_for_auth", _noop_cache_session_for_auth + ) monkeypatch.setattr(users_module, "create_acces_mobile_token", lambda _: "access") monkeypatch.setattr(users_module, "create_raw_refresh_token", lambda: "refresh") diff --git a/tests/unit/test_mobile_auth_intent_validation.py b/tests/unit/test_mobile_auth_intent_validation.py index 2311b7c9..f577c7e5 100644 --- a/tests/unit/test_mobile_auth_intent_validation.py +++ b/tests/unit/test_mobile_auth_intent_validation.py @@ -18,13 +18,19 @@ import app.service.users as users_module from app.core.securite import hash_password -from app.schema.request.mobile.auth import MobileLoginRequest, MobileRegisterRequest, RegisterVerifyRequest +from app.schema.request.mobile.auth import ( + MobileLoginRequest, + MobileRegisterRequest, + RegisterVerifyRequest, +) from app.service.session import SessionService from app.service.users import AuthService class FakeUser: - def __init__(self, email: str, exists: bool = True, password: str = "ValidPass@123") -> None: + def __init__( + self, email: str, exists: bool = True, password: str = "ValidPass@123" + ) -> None: self.id = uuid.uuid4() self.email = email self.blocked = False @@ -98,14 +104,18 @@ async def get_device_by_id_any(self, id: uuid.UUID) -> FakeDevice | None: # no longer calls this for auth decisions. return None - async def get_device_by_id(self, id: uuid.UUID, user_id: uuid.UUID) -> FakeDevice | None: + async def get_device_by_id( + self, id: uuid.UUID, user_id: uuid.UUID + ) -> FakeDevice | None: for device in self._devices.values(): if device.id == id and device.user_id == user_id: return device return None async def create_device(self, arg: Any) -> FakeDevice: - device = FakeDevice(physical_device_id=arg.physical_device_id, user_id=arg.user_id) + device = FakeDevice( + physical_device_id=arg.physical_device_id, user_id=arg.user_id + ) self._devices[(arg.user_id, arg.physical_device_id)] = device return device @@ -132,7 +142,9 @@ async def get_session_by_id(self, id: uuid.UUID) -> FakeSession | None: return session return None - async def list_sessions_by_user(self, user_id: uuid.UUID) -> AsyncIterator[FakeSession]: + async def list_sessions_by_user( + self, user_id: uuid.UUID + ) -> AsyncIterator[FakeSession]: for (u, _d), session in self._sessions.items(): if u == user_id: yield session @@ -153,7 +165,8 @@ async def evict_overflow_sessions( self, *, user_id: uuid.UUID, id: uuid.UUID, session_limit: int ) -> AsyncIterator[uuid.UUID]: candidates = [ - s for (u, _d), s in list(self._sessions.items()) + s + for (u, _d), s in list(self._sessions.items()) if u == user_id and s.id != id ] # +1 accounts for the current session itself, which isn't in `candidates` @@ -185,6 +198,7 @@ async def upsert_session( self._sessions[key] = session return session + class FakeRedis: def __init__(self) -> None: self._store: dict[str, str] = {} @@ -218,7 +232,9 @@ def _patch_token_helpers(monkeypatch: pytest.MonkeyPatch) -> None: async def _noop_cache_session_for_auth(**_: object) -> None: return None - monkeypatch.setattr(SessionService, "cache_session_for_auth", _noop_cache_session_for_auth) + monkeypatch.setattr( + SessionService, "cache_session_for_auth", _noop_cache_session_for_auth + ) monkeypatch.setattr(users_module, "create_acces_mobile_token", lambda _: "access") monkeypatch.setattr(users_module, "create_raw_refresh_token", lambda: "refresh") @@ -458,7 +474,9 @@ async def _raise_integrity_error(*args: Any, **kwargs: Any) -> Any: raise IntegrityError( statement="INSERT INTO users", params={}, - orig=FakeOrigException("duplicate key value violates unique constraint idx_users_email") + orig=FakeOrigException( + "duplicate key value violates unique constraint idx_users_email" + ), ) user_querier = FakeUserQuerier(user) @@ -573,19 +591,26 @@ def test_relogin_on_existing_device_succeeds_even_at_session_cap( assert len(session_querier._sessions) == AuthService.SESSION_LIMIT # A genuinely NEW device at the cap should evict the oldest and SUCCEED. - result = asyncio.run(service.mobile_login(FakeRedis(), MobileLoginRequest( - email="user@example.com", - password="ValidPass@123", - device_name="New device", - device_type="android", - physical_device_id=uuid.uuid4(), - ))) + result = asyncio.run( + service.mobile_login( + FakeRedis(), + MobileLoginRequest( + email="user@example.com", + password="ValidPass@123", + device_name="New device", + device_type="android", + physical_device_id=uuid.uuid4(), + ), + ) + ) assert result.access_token == "access" # Count stays at cap — one evicted, one added. assert len(session_querier._sessions) == AuthService.SESSION_LIMIT # Re-logging in on an EXISTING device (replace) must still succeed. - existing_physical_id = next(iter(device_querier._devices.values())).physical_device_id + existing_physical_id = next( + iter(device_querier._devices.values()) + ).physical_device_id repeat_req = MobileLoginRequest( email="user@example.com", password="ValidPass@123", @@ -598,6 +623,7 @@ def test_relogin_on_existing_device_succeeds_even_at_session_cap( # Session count must NOT have grown — this was a replace, not an addition. assert len(session_querier._sessions) == AuthService.SESSION_LIMIT + def test_same_physical_device_id_reuses_device_row( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/tests/unit/test_mobile_auth_rate_limiting.py b/tests/unit/test_mobile_auth_rate_limiting.py index e35d32c7..c45f4ebf 100644 --- a/tests/unit/test_mobile_auth_rate_limiting.py +++ b/tests/unit/test_mobile_auth_rate_limiting.py @@ -33,9 +33,11 @@ def __init__(self) -> None: self.id = uuid.uuid4() self.email = "test@example.com" from app.core.securite import hash_password + self.hashed_password = hash_password("ValidPass@123") self.blocked = False + class FakeUserQuerier: def __init__(self) -> None: self._user = FakeUser() @@ -46,6 +48,7 @@ async def get_user_by_email(self, email: str) -> FakeUser: async def get_user_by_id_for_update(self, id: uuid.UUID) -> FakeUser: return self._user + class FakeDeviceQuerier: pass @@ -72,6 +75,7 @@ def test_rate_limiting_triggered_after_max_attempts() -> None: # Stub session creation to avoid database / redis dependencies async def _dummy_create_session(*args: object, **kwargs: object) -> Any: from app.schema.response.mobile.auth import MobileAuthResponse + return MobileAuthResponse( access_token="access", refresh_token="refresh", @@ -80,6 +84,7 @@ async def _dummy_create_session(*args: object, **kwargs: object) -> Any: user_id=uuid.uuid4(), is_new_user=False, ) + service._create_mobile_session = _dummy_create_session # type: ignore req = MobileLoginRequest( diff --git a/tests/unit/test_mobile_auth_request_validation.py b/tests/unit/test_mobile_auth_request_validation.py index 39b287ed..0f2fcc79 100644 --- a/tests/unit/test_mobile_auth_request_validation.py +++ b/tests/unit/test_mobile_auth_request_validation.py @@ -68,7 +68,6 @@ def fake_container() -> FakeContainer: return FakeContainer() - @pytest.fixture def client(fake_container: FakeContainer) -> Iterator[TestClient]: app.dependency_overrides[get_container] = lambda: fake_container @@ -229,5 +228,3 @@ def test_mobile_auth_uses_forwarded_ip_for_rate_limit_identity( assert response.status_code == 200 assert fake_container.auth_service.login_client_ip == "203.0.113.10" - - diff --git a/tests/unit/test_photo_approval_lifecycle.py b/tests/unit/test_photo_approval_lifecycle.py index 00577e6a..1c5d5d1b 100644 --- a/tests/unit/test_photo_approval_lifecycle.py +++ b/tests/unit/test_photo_approval_lifecycle.py @@ -7,38 +7,47 @@ from app.service.face_embedding import DetectedFace from app.service.photo_approval import PhotoApprovalService + @pytest.fixture def mock_conn() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_face_embedding_service() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_single_face_service() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_notification_service() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_photo_face_querier() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_photo_querier() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_photo_approval_querier() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_processing_job_querier() -> AsyncMock: return AsyncMock() + @pytest.fixture def mock_staged_upload_storage_service() -> AsyncMock: return AsyncMock() @@ -153,14 +162,19 @@ async def test_expire_stale_marks_photos_approved( ) from typing import Any, AsyncIterator + # Mock the generator for expire_stale_approvals async def mock_generator(*args: Any, **kwargs: Any) -> AsyncIterator[uuid.UUID]: yield uuid.uuid4() yield uuid.uuid4() - mock_photo_approval_querier.expire_stale_approvals = MagicMock(side_effect=mock_generator) + mock_photo_approval_querier.expire_stale_approvals = MagicMock( + side_effect=mock_generator + ) count = await service.expire_stale(timeout_days=7) assert count == 2 - mock_photo_approval_querier.expire_stale_approvals.assert_called_once_with(timeout_days=7) + mock_photo_approval_querier.expire_stale_approvals.assert_called_once_with( + timeout_days=7 + ) diff --git a/tests/unit/test_photo_approval_service.py b/tests/unit/test_photo_approval_service.py index d10c62bf..a955a20d 100644 --- a/tests/unit/test_photo_approval_service.py +++ b/tests/unit/test_photo_approval_service.py @@ -42,6 +42,7 @@ def _make_photo(storage_key: str = "photos/test.jpg") -> MagicMock: @pytest.fixture def approval_querier() -> AsyncMock: from db.generated import photo_approvals as pa_queries + q = MagicMock(spec=pa_queries.AsyncQuerier) q.update_photo_approval_decision = AsyncMock() q.get_photo_approvals_by_photo_id = MagicMock() # async generator @@ -51,6 +52,7 @@ def approval_querier() -> AsyncMock: @pytest.fixture def photo_querier() -> AsyncMock: from db.generated import photos as photo_queries + q = MagicMock(spec=photo_queries.AsyncQuerier) q.update_photo_status = AsyncMock(return_value=None) q.get_photo_by_id = AsyncMock(return_value=_make_photo()) @@ -60,6 +62,7 @@ def photo_querier() -> AsyncMock: @pytest.fixture def storage_service() -> AsyncMock: from app.service.staged_upload_storage import StagedUploadStorageService + svc = MagicMock(spec=StagedUploadStorageService) svc.delete_storage_key = AsyncMock() return svc @@ -68,6 +71,7 @@ def storage_service() -> AsyncMock: @pytest.fixture def audit_service() -> AsyncMock: from app.service.audit import AuditService + svc = MagicMock(spec=AuditService) svc.create_record = AsyncMock() return svc @@ -89,9 +93,11 @@ def _make_service( def _mock_async_iter(items: list[object]): # type: ignore[type-arg] """Return a MagicMock that behaves like an async for loop.""" + async def _gen(): # type: ignore[return] for item in items: yield item + return _gen() @@ -113,10 +119,14 @@ async def test_all_approved_sets_photo_status( approvals = [_make_approval("approved"), _make_approval("approved")] approval_querier.update_photo_approval_decision.return_value = MagicMock() - approval_querier.get_photo_approvals_by_photo_id.return_value = _mock_async_iter(approvals) + approval_querier.get_photo_approvals_by_photo_id.return_value = ( + _mock_async_iter(approvals) + ) service = _make_service(approval_querier, photo_querier, storage_service) - result = await service.decide(photo_id=photo_id, user_id=user_id, decision="approved") + result = await service.decide( + photo_id=photo_id, user_id=user_id, decision="approved" + ) assert result == "approved" photo_querier.update_photo_status.assert_called_once_with( @@ -135,7 +145,9 @@ async def test_all_approved_does_not_delete_storage( approvals = [_make_approval("approved")] approval_querier.update_photo_approval_decision.return_value = MagicMock() - approval_querier.get_photo_approvals_by_photo_id.return_value = _mock_async_iter(approvals) + approval_querier.get_photo_approvals_by_photo_id.return_value = ( + _mock_async_iter(approvals) + ) service = _make_service(approval_querier, photo_querier, storage_service) await service.decide(photo_id=photo_id, user_id=user_id, decision="approved") @@ -161,10 +173,14 @@ async def test_one_rejection_sets_photo_status_rejected( approvals = [_make_approval("approved"), _make_approval("rejected")] approval_querier.update_photo_approval_decision.return_value = MagicMock() - approval_querier.get_photo_approvals_by_photo_id.return_value = _mock_async_iter(approvals) + approval_querier.get_photo_approvals_by_photo_id.return_value = ( + _mock_async_iter(approvals) + ) service = _make_service(approval_querier, photo_querier, storage_service) - result = await service.decide(photo_id=photo_id, user_id=user_id, decision="rejected") + result = await service.decide( + photo_id=photo_id, user_id=user_id, decision="rejected" + ) assert result == "rejected" photo_querier.update_photo_status.assert_called_once_with( @@ -182,10 +198,14 @@ async def test_rejection_triggers_storage_deletion( user_id = uuid.uuid4() approvals = [_make_approval("rejected")] storage_key = "photos/reject-me.jpg" - photo_querier.get_photo_by_id.return_value = _make_photo(storage_key=storage_key) + photo_querier.get_photo_by_id.return_value = _make_photo( + storage_key=storage_key + ) approval_querier.update_photo_approval_decision.return_value = MagicMock() - approval_querier.get_photo_approvals_by_photo_id.return_value = _mock_async_iter(approvals) + approval_querier.get_photo_approvals_by_photo_id.return_value = ( + _mock_async_iter(approvals) + ) service = _make_service(approval_querier, photo_querier, storage_service) await service.decide(photo_id=photo_id, user_id=user_id, decision="rejected") @@ -205,12 +225,16 @@ async def test_storage_deletion_failure_does_not_raise( approvals = [_make_approval("rejected")] approval_querier.update_photo_approval_decision.return_value = MagicMock() - approval_querier.get_photo_approvals_by_photo_id.return_value = _mock_async_iter(approvals) + approval_querier.get_photo_approvals_by_photo_id.return_value = ( + _mock_async_iter(approvals) + ) storage_service.delete_storage_key.side_effect = Exception("MinIO unavailable") service = _make_service(approval_querier, photo_querier, storage_service) # Must not raise despite MinIO being unavailable - result = await service.decide(photo_id=photo_id, user_id=user_id, decision="rejected") + result = await service.decide( + photo_id=photo_id, user_id=user_id, decision="rejected" + ) assert result == "rejected" @@ -232,10 +256,14 @@ async def test_pending_approval_returns_pending( approvals = [_make_approval("approved"), _make_approval("pending")] approval_querier.update_photo_approval_decision.return_value = MagicMock() - approval_querier.get_photo_approvals_by_photo_id.return_value = _mock_async_iter(approvals) + approval_querier.get_photo_approvals_by_photo_id.return_value = ( + _mock_async_iter(approvals) + ) service = _make_service(approval_querier, photo_querier, storage_service) - result = await service.decide(photo_id=photo_id, user_id=user_id, decision="approved") + result = await service.decide( + photo_id=photo_id, user_id=user_id, decision="approved" + ) assert result == "pending" @@ -251,7 +279,9 @@ async def test_pending_does_not_update_photo_status( approvals = [_make_approval("pending"), _make_approval("pending")] approval_querier.update_photo_approval_decision.return_value = MagicMock() - approval_querier.get_photo_approvals_by_photo_id.return_value = _mock_async_iter(approvals) + approval_querier.get_photo_approvals_by_photo_id.return_value = ( + _mock_async_iter(approvals) + ) service = _make_service(approval_querier, photo_querier, storage_service) await service.decide(photo_id=photo_id, user_id=user_id, decision="approved") @@ -302,9 +332,13 @@ async def test_audit_called_with_correct_event_type( approvals = [_make_approval("approved")] approval_querier.update_photo_approval_decision.return_value = MagicMock() - approval_querier.get_photo_approvals_by_photo_id.return_value = _mock_async_iter(approvals) + approval_querier.get_photo_approvals_by_photo_id.return_value = ( + _mock_async_iter(approvals) + ) - service = _make_service(approval_querier, photo_querier, storage_service, audit_service) + service = _make_service( + approval_querier, photo_querier, storage_service, audit_service + ) await service.decide(photo_id=photo_id, user_id=user_id, decision="approved") audit_service.create_record.assert_called_once() @@ -325,9 +359,15 @@ async def test_no_audit_without_audit_service( approvals = [_make_approval("approved")] approval_querier.update_photo_approval_decision.return_value = MagicMock() - approval_querier.get_photo_approvals_by_photo_id.return_value = _mock_async_iter(approvals) + approval_querier.get_photo_approvals_by_photo_id.return_value = ( + _mock_async_iter(approvals) + ) # No audit_service passed → audit_service=None - service = _make_service(approval_querier, photo_querier, storage_service, audit_service=None) - await service.decide(photo_id=photo_id, user_id=uuid.uuid4(), decision="approved") + service = _make_service( + approval_querier, photo_querier, storage_service, audit_service=None + ) + await service.decide( + photo_id=photo_id, user_id=uuid.uuid4(), decision="approved" + ) # no assertion needed — just must not raise diff --git a/tests/unit/test_photo_worker.py b/tests/unit/test_photo_worker.py index 3aa9813d..426dddf8 100644 --- a/tests/unit/test_photo_worker.py +++ b/tests/unit/test_photo_worker.py @@ -18,8 +18,8 @@ def mock_pj_querier(): job_type="face_detection", status="pending", attempts=0, - created_at=None, # type: ignore - completed_at=None, # type: ignore + created_at=None, # type: ignore + completed_at=None, # type: ignore ) querier.create_processing_job.return_value = job querier.update_processing_job_status.return_value = job @@ -87,17 +87,31 @@ def sample_event(): @pytest.mark.asyncio -async def test_handle_message_success_no_faces(photo_worker, sample_event, mock_face_service, mock_pj_querier, mock_photo_querier): - photo_worker._load_image = AsyncMock(return_value=FaceImagePayload(filename="test.jpg", content_type="image/jpeg", bytes=b"data")) +async def test_handle_message_success_no_faces( + photo_worker, sample_event, mock_face_service, mock_pj_querier, mock_photo_querier +): + photo_worker._load_image = AsyncMock( + return_value=FaceImagePayload( + filename="test.jpg", content_type="image/jpeg", bytes=b"data" + ) + ) mock_face_service.detect_faces.return_value = [] with patch("app.worker.photo_worker.main.NatsClient.js_publish") as mock_publish: - await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) + await photo_worker.handle_message( + sample_event.model_dump_json().encode("utf-8") + ) mock_pj_querier.create_processing_job.assert_called_once() - mock_pj_querier.update_processing_job_status.assert_any_call(id=mock_pj_querier.create_processing_job.return_value.id, status="completed") - mock_photo_querier.update_photo_status.assert_called_once_with(id=sample_event.photo_id, status="approved") - mock_photo_querier.update_photo_visibility.assert_called_once_with(id=sample_event.photo_id, visibility="public") + mock_pj_querier.update_processing_job_status.assert_any_call( + id=mock_pj_querier.create_processing_job.return_value.id, status="completed" + ) + mock_photo_querier.update_photo_status.assert_called_once_with( + id=sample_event.photo_id, status="approved" + ) + mock_photo_querier.update_photo_visibility.assert_called_once_with( + id=sample_event.photo_id, visibility="public" + ) # photo_worker no longer schedules immediate MinIO cleanup — that's # now threshold-based (event_lifecycle worker, gated on event.end_date # + Drive sync confirmation for direct uploads). @@ -105,57 +119,106 @@ async def test_handle_message_success_no_faces(photo_worker, sample_event, mock_ @pytest.mark.asyncio -async def test_handle_message_success_single_face(photo_worker, sample_event, mock_face_service, mock_single_face_service, mock_pj_querier): - photo_worker._load_image = AsyncMock(return_value=FaceImagePayload(filename="test.jpg", content_type="image/jpeg", bytes=b"data")) +async def test_handle_message_success_single_face( + photo_worker, + sample_event, + mock_face_service, + mock_single_face_service, + mock_pj_querier, +): + photo_worker._load_image = AsyncMock( + return_value=FaceImagePayload( + filename="test.jpg", content_type="image/jpeg", bytes=b"data" + ) + ) face = DetectedFace(bbox=(0, 0, 100, 100), embedding=[0.1] * 512) mock_face_service.detect_faces.return_value = [face] with patch("app.worker.photo_worker.main.NatsClient.js_publish") as mock_publish: - await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) + await photo_worker.handle_message( + sample_event.model_dump_json().encode("utf-8") + ) - mock_pj_querier.update_processing_job_status.assert_any_call(id=mock_pj_querier.create_processing_job.return_value.id, status="completed") + mock_pj_querier.update_processing_job_status.assert_any_call( + id=mock_pj_querier.create_processing_job.return_value.id, status="completed" + ) mock_single_face_service.process_detected_face.assert_called_once() # Only the audit event — cleanup is no longer scheduled by photo_worker. assert mock_publish.call_count == 1 @pytest.mark.asyncio -async def test_handle_message_success_group_face(photo_worker, sample_event, mock_face_service, mock_photo_face_querier, mock_notification_service, mock_pj_querier): - photo_worker._load_image = AsyncMock(return_value=FaceImagePayload(filename="test.jpg", content_type="image/jpeg", bytes=b"data")) +async def test_handle_message_success_group_face( + photo_worker, + sample_event, + mock_face_service, + mock_photo_face_querier, + mock_notification_service, + mock_pj_querier, +): + photo_worker._load_image = AsyncMock( + return_value=FaceImagePayload( + filename="test.jpg", content_type="image/jpeg", bytes=b"data" + ) + ) face1 = DetectedFace(bbox=(0, 0, 100, 100), embedding=[0.1] * 512) face2 = DetectedFace(bbox=(100, 100, 200, 200), embedding=[0.2] * 512) mock_face_service.detect_faces.return_value = [face1, face2] with patch("app.worker.photo_worker.main.NatsClient.js_publish"): - await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) + await photo_worker.handle_message( + sample_event.model_dump_json().encode("utf-8") + ) - mock_pj_querier.update_processing_job_status.assert_any_call(id=mock_pj_querier.create_processing_job.return_value.id, status="completed") + mock_pj_querier.update_processing_job_status.assert_any_call( + id=mock_pj_querier.create_processing_job.return_value.id, status="completed" + ) assert mock_photo_face_querier.insert_photo_face_with_approval.call_count == 2 assert mock_notification_service.create_notification.call_count == 2 @pytest.mark.asyncio -async def test_handle_message_fails_on_minio_load(photo_worker, sample_event, mock_pj_querier): +async def test_handle_message_fails_on_minio_load( + photo_worker, sample_event, mock_pj_querier +): photo_worker._load_image = AsyncMock(side_effect=Exception("MinIO error")) with pytest.raises(Exception, match="MinIO error"): - await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) - assert "failed" not in [call.kwargs.get("status") for call in mock_pj_querier.update_processing_job_status.call_args_list] + await photo_worker.handle_message( + sample_event.model_dump_json().encode("utf-8") + ) + assert "failed" not in [ + call.kwargs.get("status") + for call in mock_pj_querier.update_processing_job_status.call_args_list + ] @pytest.mark.asyncio -async def test_handle_message_fails_on_ai_detection(photo_worker, sample_event, mock_face_service, mock_pj_querier): - photo_worker._load_image = AsyncMock(return_value=FaceImagePayload(filename="test.jpg", content_type="image/jpeg", bytes=b"data")) +async def test_handle_message_fails_on_ai_detection( + photo_worker, sample_event, mock_face_service, mock_pj_querier +): + photo_worker._load_image = AsyncMock( + return_value=FaceImagePayload( + filename="test.jpg", content_type="image/jpeg", bytes=b"data" + ) + ) mock_face_service.detect_faces.side_effect = Exception("InsightFace out of memory") with pytest.raises(Exception, match="InsightFace out of memory"): - await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) - assert "failed" not in [call.kwargs.get("status") for call in mock_pj_querier.update_processing_job_status.call_args_list] + await photo_worker.handle_message( + sample_event.model_dump_json().encode("utf-8") + ) + assert "failed" not in [ + call.kwargs.get("status") + for call in mock_pj_querier.update_processing_job_status.call_args_list + ] @pytest.mark.asyncio async def test_minio_retry_logic(photo_worker): - with patch("app.worker.photo_worker.main.Bucket.get") as mock_bucket_get, \ - patch("app.worker.photo_worker.main.settings") as mock_settings, \ - patch("app.worker.photo_worker.main.asyncio.sleep") as mock_sleep: + with ( + patch("app.worker.photo_worker.main.Bucket.get") as mock_bucket_get, + patch("app.worker.photo_worker.main.settings") as mock_settings, + patch("app.worker.photo_worker.main.asyncio.sleep") as mock_sleep, + ): mock_settings.MINIO_RETRY_ATTEMPTS = 3 mock_settings.MINIO_RETRY_BASE_SECONDS = 0 mock_bucket_get.side_effect = [ diff --git a/tests/unit/test_upload_requests.py b/tests/unit/test_upload_requests.py index 77877680..7f8551cb 100644 --- a/tests/unit/test_upload_requests.py +++ b/tests/unit/test_upload_requests.py @@ -23,24 +23,31 @@ def mock_upload_request_group_querier(): return AsyncMock() + @pytest.fixture def mock_upload_request_querier(): return AsyncMock() + @pytest.fixture def mock_upload_request_photo_querier(): return AsyncMock() + @pytest.fixture def mock_photo_querier(): return AsyncMock() + @pytest.fixture def mock_staged_upload_storage(): mock = AsyncMock() - mock.store_staging_object.return_value = StoredObject(storage_key="test_storage_key", content_type="image/jpeg", file_name="photo.jpg") + mock.store_staging_object.return_value = StoredObject( + storage_key="test_storage_key", content_type="image/jpeg", file_name="photo.jpg" + ) return mock + @pytest.fixture def mock_staff_drive_service(): mock = AsyncMock() @@ -48,10 +55,12 @@ def mock_staff_drive_service(): mock.staff_user_querier = AsyncMock() return mock + @pytest.fixture def mock_staff_notifications_service(): return AsyncMock() + @pytest.fixture def mock_audit_service(): return AsyncMock() @@ -125,22 +134,24 @@ async def test_create_request_success( source="drive", ) - mock_upload_request_photo_querier.create_upload_request_photo.return_value = UploadRequestPhoto( - id=uuid.uuid4(), - upload_request_id=request_id, - drive_file_id="drive_id_1", - file_name="photo.jpg", - mime_type="image/jpeg", - size_bytes=1024, - staging_storage_key="test_storage_key", - final_storage_key=None, - taken_at=None, - day_number=None, - visibility="public", - status="staged", - created_at=datetime.now(timezone.utc), - source="drive", - transfer_status="uploaded", + mock_upload_request_photo_querier.create_upload_request_photo.return_value = ( + UploadRequestPhoto( + id=uuid.uuid4(), + upload_request_id=request_id, + drive_file_id="drive_id_1", + file_name="photo.jpg", + mime_type="image/jpeg", + size_bytes=1024, + staging_storage_key="test_storage_key", + final_storage_key=None, + taken_at=None, + day_number=None, + visibility="public", + status="staged", + created_at=datetime.now(timezone.utc), + source="drive", + transfer_status="uploaded", + ) ) photos = [ @@ -159,9 +170,13 @@ async def test_create_request_success( content=b"content", ) - with patch("app.service.upload_requests.GoogleDriveClient.download_file", return_value=mock_download) as mock_drive, \ - patch("app.service.upload_requests.NatsClient.js_publish") as mock_publish: - + with ( + patch( + "app.service.upload_requests.GoogleDriveClient.download_file", + return_value=mock_download, + ) as mock_drive, + patch("app.service.upload_requests.NatsClient.js_publish") as mock_publish, + ): details = await upload_requests_service.create_request( event_id=event_id, photos=photos, @@ -194,28 +209,57 @@ async def test_create_request_duplicate_conflict( request_id = uuid.uuid4() mock_upload_request_querier.create_upload_request.return_value = UploadRequest( - id=request_id, event_id=event_id, group_id=None, drive_file_id=None, requested_by=mock_staff_user.id, photo_count=1, status="pending", approved_by=None, rejection_reason=None, created_at=datetime.now(timezone.utc), approved_at=None, source="drive" - ) + id=request_id, + event_id=event_id, + group_id=None, + drive_file_id=None, + requested_by=mock_staff_user.id, + photo_count=1, + status="pending", + approved_by=None, + rejection_reason=None, + created_at=datetime.now(timezone.utc), + approved_at=None, + source="drive", + ) # Simulate DB Conflict (Duplicate) on photo insert - mock_upload_request_photo_querier.create_upload_request_photo.side_effect = create_integrity_error("23505") + mock_upload_request_photo_querier.create_upload_request_photo.side_effect = ( + create_integrity_error("23505") + ) - photos = [UploadPhotoInput(drive_file_id="drive_id_1", taken_at=None, day_number=None, visibility="public")] + photos = [ + UploadPhotoInput( + drive_file_id="drive_id_1", + taken_at=None, + day_number=None, + visibility="public", + ) + ] mock_download = GoogleDriveFileDownload( - metadata=GoogleDriveFileMetadata(id="drive_id_1", name="photo.jpg", mime_type="image/jpeg", size_bytes=1024), + metadata=GoogleDriveFileMetadata( + id="drive_id_1", name="photo.jpg", mime_type="image/jpeg", size_bytes=1024 + ), content=b"content", ) - with patch("app.service.upload_requests.GoogleDriveClient.download_file", return_value=mock_download): + with patch( + "app.service.upload_requests.GoogleDriveClient.download_file", + return_value=mock_download, + ): with pytest.raises(HTTPException) as exc: - await upload_requests_service.create_request(event_id=event_id, photos=photos, requested_by=mock_staff_user) + await upload_requests_service.create_request( + event_id=event_id, photos=photos, requested_by=mock_staff_user + ) assert exc.value.status_code == 409 assert "Duplicate photo" in exc.value.detail # Verify cleanup was called on StagedUploadStorageService - mock_staged_upload_storage.delete_storage_key.assert_called_once_with("test_storage_key") + mock_staged_upload_storage.delete_storage_key.assert_called_once_with( + "test_storage_key" + ) @pytest.mark.asyncio @@ -227,19 +271,43 @@ async def test_create_group_from_folder( event_id = uuid.uuid4() group_id = uuid.uuid4() - mock_upload_request_group_querier.create_upload_request_group.return_value = UploadRequestGroup( - id=group_id, event_id=event_id, folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="pending", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None, source="drive" + mock_upload_request_group_querier.create_upload_request_group.return_value = ( + UploadRequestGroup( + id=group_id, + event_id=event_id, + folder_id="folder_123", + requested_by=mock_staff_user.id, + total_photo_count=0, + batch_count=0, + processed_photo_count=0, + failed_photo_count=0, + processing_status="pending", + error_message=None, + created_at=datetime.now(timezone.utc), + status="pending", + approved_by=None, + approved_at=None, + rejection_reason=None, + source="drive", ) + ) with patch("app.service.upload_requests.NatsClient.js_publish") as mock_publish: details = await upload_requests_service.create_group_from_folder( - event_id=event_id, folder_id="folder_123", visibility="public", day_number=None, requested_by=mock_staff_user + event_id=event_id, + folder_id="folder_123", + visibility="public", + day_number=None, + requested_by=mock_staff_user, ) assert details.group.id == group_id mock_upload_request_group_querier.create_upload_request_group.assert_called_once() mock_publish.assert_called_once() - assert mock_publish.call_args[0][0] == NatsSubjects.STAFF_UPLOAD_GROUP_IMPORT_REQUESTED + assert ( + mock_publish.call_args[0][0] + == NatsSubjects.STAFF_UPLOAD_GROUP_IMPORT_REQUESTED + ) @pytest.mark.asyncio @@ -252,26 +320,79 @@ async def test_process_group_import_no_images( ): group_id = uuid.uuid4() mock_upload_request_group_querier.start_upload_request_group_processing.return_value = UploadRequestGroup( - id=group_id, event_id=uuid.uuid4(), folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="processing", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None, source="drive" - ) - mock_staff_drive_service.staff_user_querier.get_staff_user_by_id.return_value = mock_staff_user + id=group_id, + event_id=uuid.uuid4(), + folder_id="folder_123", + requested_by=mock_staff_user.id, + total_photo_count=0, + batch_count=0, + processed_photo_count=0, + failed_photo_count=0, + processing_status="processing", + error_message=None, + created_at=datetime.now(timezone.utc), + status="pending", + approved_by=None, + approved_at=None, + rejection_reason=None, + source="drive", + ) + mock_staff_drive_service.staff_user_querier.get_staff_user_by_id.return_value = ( + mock_staff_user + ) # Return 0 images - with patch("app.service.upload_requests.GoogleDriveClient.list_folder_files", return_value=[]): + with patch( + "app.service.upload_requests.GoogleDriveClient.list_folder_files", + return_value=[], + ): + async def mock_get_group(*args, **kwargs): yield UploadRequestGroup( - id=group_id, event_id=uuid.uuid4(), folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="processing", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None, source="drive" - ) + id=group_id, + event_id=uuid.uuid4(), + folder_id="folder_123", + requested_by=mock_staff_user.id, + total_photo_count=0, + batch_count=0, + processed_photo_count=0, + failed_photo_count=0, + processing_status="processing", + error_message=None, + created_at=datetime.now(timezone.utc), + status="pending", + approved_by=None, + approved_at=None, + rejection_reason=None, + source="drive", + ) + mock_upload_request_querier.list_upload_requests_by_group_id = mock_get_group async def mock_list_photos_by_ids(*args, **kwargs): if False: - yield # Empty generator + yield # Empty generator + upload_requests_service.upload_request_photo_querier.list_upload_request_photos_by_upload_request_ids = mock_list_photos_by_ids mock_upload_request_group_querier.get_upload_request_group_by_id.return_value = UploadRequestGroup( - id=group_id, event_id=uuid.uuid4(), folder_id="folder_123", requested_by=mock_staff_user.id, total_photo_count=0, batch_count=0, processed_photo_count=0, failed_photo_count=0, processing_status="processing", error_message=None, created_at=datetime.now(timezone.utc), status="pending", approved_by=None, approved_at=None, rejection_reason=None, source="drive" - ) + id=group_id, + event_id=uuid.uuid4(), + folder_id="folder_123", + requested_by=mock_staff_user.id, + total_photo_count=0, + batch_count=0, + processed_photo_count=0, + failed_photo_count=0, + processing_status="processing", + error_message=None, + created_at=datetime.now(timezone.utc), + status="pending", + approved_by=None, + approved_at=None, + rejection_reason=None, + source="drive", + ) await upload_requests_service.process_group_import( group_id=group_id, visibility="public", day_number=None @@ -279,7 +400,9 @@ async def mock_list_photos_by_ids(*args, **kwargs): # Verify it marked group as failed mock_upload_request_group_querier.fail_upload_request_group_processing.assert_called_once() - kwargs = mock_upload_request_group_querier.fail_upload_request_group_processing.call_args[0][0] + kwargs = mock_upload_request_group_querier.fail_upload_request_group_processing.call_args[ + 0 + ][0] assert "does not contain valid images" in kwargs.error_message @@ -297,13 +420,39 @@ async def test_approve_request_without_side_effects( event_id = uuid.uuid4() mock_upload_request_querier.get_upload_request_by_id.return_value = UploadRequest( - id=request_id, event_id=event_id, group_id=None, drive_file_id=None, requested_by=uuid.uuid4(), photo_count=1, status="pending", approved_by=None, rejection_reason=None, created_at=datetime.now(timezone.utc), approved_at=None, source="drive" - ) + id=request_id, + event_id=event_id, + group_id=None, + drive_file_id=None, + requested_by=uuid.uuid4(), + photo_count=1, + status="pending", + approved_by=None, + rejection_reason=None, + created_at=datetime.now(timezone.utc), + approved_at=None, + source="drive", + ) async def mock_list_photos(*args, **kwargs): yield UploadRequestPhoto( - id=photo_id, upload_request_id=request_id, drive_file_id="drive_1", file_name="p.jpg", mime_type="image/jpeg", size_bytes=100, staging_storage_key="stage_key", final_storage_key=None, taken_at=None, day_number=None, visibility="public", status="staged", created_at=datetime.now(timezone.utc), source="drive", transfer_status="uploaded" + id=photo_id, + upload_request_id=request_id, + drive_file_id="drive_1", + file_name="p.jpg", + mime_type="image/jpeg", + size_bytes=100, + staging_storage_key="stage_key", + final_storage_key=None, + taken_at=None, + day_number=None, + visibility="public", + status="staged", + created_at=datetime.now(timezone.utc), + source="drive", + transfer_status="uploaded", ) + mock_upload_request_photo_querier.list_upload_request_photos_by_upload_request_id = mock_list_photos mock_staged_upload_storage.promote_to_final.return_value = "final_key" @@ -311,7 +460,12 @@ async def mock_list_photos(*args, **kwargs): mock_upload_request_photo_querier.update_upload_request_photo_approval.return_value = MagicMock() mock_upload_request_querier.approve_upload_request.return_value = MagicMock() - upload_req, staged_photos, final_keys, created_photos = await upload_requests_service._approve_request_without_side_effects( + ( + upload_req, + staged_photos, + final_keys, + created_photos, + ) = await upload_requests_service._approve_request_without_side_effects( request_id=request_id, approved_by=mock_staff_user ) From a7bfff42ec9a2e8cac13769b2e60fbad73bcc4a7 Mon Sep 17 00:00:00 2001 From: Adem Boukabes <142881379+ademboukabes@users.noreply.github.com> Date: Sun, 30 Aug 2026 15:39:26 +0100 Subject: [PATCH 26/29] chore: refactor drive sync and add e2e tests --- app/infra/google_drive.py | 49 +++++++ app/service/staff_drive.py | 67 ++++++++- app/service/upload_requests.py | 1 + app/worker/drive_sync/main.py | 18 +++ app/worker/event_lifecycle/main.py | 2 +- tests/integration/test_drive_sync_flow.py | 164 ++++++++++++++++++++++ 6 files changed, 298 insertions(+), 3 deletions(-) create mode 100644 tests/integration/test_drive_sync_flow.py diff --git a/app/infra/google_drive.py b/app/infra/google_drive.py index b8b39ca6..9aa90ff5 100644 --- a/app/infra/google_drive.py +++ b/app/infra/google_drive.py @@ -158,6 +158,53 @@ async def get_user_info(access_token: str) -> GoogleUserInfo: verified_email=bool(data.get("verified_email", False)), ) + @staticmethod + async def create_folder( + *, + access_token: str, + name: str, + parent_id: str | None, + ) -> GoogleDriveFileMetadata: + metadata: dict[str, object] = { + "name": name, + "mimeType": GoogleDriveClient._drive_folder_mime_type, + } + if parent_id: + metadata["parents"] = [parent_id] + + encoded = json.dumps(metadata).encode("utf-8") + + def _request() -> dict[str, object]: + request = urllib.request.Request( + "https://www.googleapis.com/drive/v3/files?supportsAllDrives=true&fields=id,name,mimeType,size", + data=encoded, + headers={ + "Authorization": f"Bearer {access_token}", + "Content-Type": "application/json", + }, + method="POST", + ) + try: + with urllib.request.urlopen(request, timeout=15) as response: + return json.loads(response.read().decode("utf-8")) + except urllib.error.HTTPError as exc: + details = exc.read().decode("utf-8", errors="ignore") + raise AppException.bad_request( + f"Google folder creation failed: {details or exc.reason}" + ) from exc + except urllib.error.URLError as exc: + raise AppException.internal_error( + "Unable to reach Google APIs" + ) from exc + + result = await asyncio.to_thread(_request) + return GoogleDriveFileMetadata( + id=GoogleDriveClient._require_str(result, "id"), + name=GoogleDriveClient._require_str(result, "name"), + mime_type=GoogleDriveClient._require_str(result, "mimeType"), + size_bytes=0, + ) + @staticmethod async def get_file_metadata( *, @@ -462,6 +509,8 @@ def _request() -> dict[str, object]: return json.loads(response.read().decode("utf-8")) except urllib.error.HTTPError as exc: details = exc.read().decode("utf-8", errors="ignore") + if "invalid_grant" in details: + raise AppException.unauthorized("invalid_grant") from exc raise AppException.bad_request( f"Google token exchange failed: {details or exc.reason}" ) from exc diff --git a/app/service/staff_drive.py b/app/service/staff_drive.py index 7cfde00e..cd0188a5 100644 --- a/app/service/staff_drive.py +++ b/app/service/staff_drive.py @@ -8,6 +8,7 @@ from dataclasses import dataclass from datetime import datetime, timedelta, timezone +from fastapi import HTTPException from cryptography.fernet import Fernet, InvalidToken from app.core.config import settings @@ -175,7 +176,16 @@ async def _refresh_connection_access_token( ) refresh_token = self.decrypt(connection.refresh_token) - token = await GoogleDriveClient.refresh_access_token(refresh_token) + try: + token = await GoogleDriveClient.refresh_access_token(refresh_token) + except HTTPException as exc: + if "invalid_grant" in str(exc.detail): + await self.drive_connection_querier.revoke_staff_drive_connection_by_staff_user_id( + staff_user_id=connection.staff_user_id, + provider=connection.provider, + ) + raise AppException.not_found("Drive connection revoked") from exc + raise encrypted_access_token = self._encrypt(token.access_token) encrypted_refresh_token = connection.refresh_token @@ -220,22 +230,75 @@ async def get_system_access_token(self) -> str: connection = await self._refresh_connection_access_token(connection) return self.decrypt(connection.access_token) + async def _get_or_create_event_folder( + self, + event_id: uuid.UUID, + event_name: str, + access_token: str, + ) -> str: + cache_key = f"drive:folder:{event_id}" + cached_id = await self.redis.get(cache_key) + if cached_id: + return cached_id + + lock_key = f"lock:drive:folder:{event_id}" + while True: + acquired = await self.redis.set(lock_key, "1", expire=30, nx=True) + if acquired: + break + await asyncio.sleep(1.0) + + try: + cached_id = await self.redis.get(cache_key) + if cached_id: + return cached_id + + parent_id = settings.GOOGLE_CLUB_DRIVE_FOLDER_ID or None + + existing = await GoogleDriveClient.search_files( + access_token=access_token, query=event_name, file_type="folder" + ) + folder_id = None + if existing: + for folder in existing: + if folder.name == event_name: + folder_id = folder.id + break + + if not folder_id: + folder_meta = await GoogleDriveClient.create_folder( + access_token=access_token, name=event_name, parent_id=parent_id + ) + folder_id = folder_meta.id + + await self.redis.set(cache_key, folder_id, expire=7 * 24 * 3600) + return folder_id + finally: + await self.redis.delete(lock_key) + async def upload_to_system_drive( self, *, file_name: str, content_type: str, data: bytes, + event_id: uuid.UUID | None = None, + event_name: str | None = None, ) -> str: """Upload bytes to the system/club Drive using the most recently connected active staff Drive connection. Returns the Drive file id.""" access_token = await self.get_system_access_token() + folder_id = settings.GOOGLE_CLUB_DRIVE_FOLDER_ID or None + + if event_id and event_name: + folder_id = await self._get_or_create_event_folder(event_id, event_name, access_token) + metadata = await GoogleDriveClient.upload_file( access_token=access_token, file_name=file_name, content_type=content_type, data=data, - folder_id=settings.GOOGLE_CLUB_DRIVE_FOLDER_ID or None, + folder_id=folder_id, ) return metadata.id diff --git a/app/service/upload_requests.py b/app/service/upload_requests.py index 3aebde48..2bf2d7cb 100644 --- a/app/service/upload_requests.py +++ b/app/service/upload_requests.py @@ -571,6 +571,7 @@ async def _publish_drive_sync_events( subject=NatsSubjects.PHOTO_DRIVE_SYNC_REQUESTED, payload={ "photo_id": str(created_photo.id), + "event_id": str(created_photo.event_id), "storage_key": created_photo.storage_key, "file_name": staged_photo.file_name, "mime_type": staged_photo.mime_type, diff --git a/app/worker/drive_sync/main.py b/app/worker/drive_sync/main.py index 4ebd495b..73c4ee4f 100644 --- a/app/worker/drive_sync/main.py +++ b/app/worker/drive_sync/main.py @@ -11,13 +11,16 @@ from app.infra.nats import NatsClient, NatsSubjects from app.infra.redis import RedisClient from app.service.staff_drive import StaffDriveService +from db.generated import events as event_queries from db.generated import photos as photo_queries from db.generated import staff_drive_connections as drive_queries from db.generated import staff_user as staff_queries +from fastapi import HTTPException class PhotoDriveSyncEvent(BaseModel): photo_id: uuid.UUID + event_id: uuid.UUID storage_key: str file_name: str mime_type: str @@ -60,12 +63,27 @@ async def _handle_event(raw_data: bytes) -> None: ) photo_querier = photo_queries.AsyncQuerier(conn) + event_querier = event_queries.AsyncQuerier(conn) + db_event = await event_querier.get_event_by_id(id=event.event_id) + if db_event is None: + logger.warning("drive_sync: event %s not found for photo %s", event.event_id, event.photo_id) + return + try: drive_file_id = await staff_drive_service.upload_to_system_drive( file_name=event.file_name, content_type=event.mime_type or content_type, data=data, + event_id=event.event_id, + event_name=db_event.name, ) + except HTTPException as exc: + if exc.status_code == 404: + logger.warning( + "drive_sync: no active drive connection, pausing for 5 minutes before retry" + ) + await asyncio.sleep(300) + raise except Exception as exc: logger.warning( "drive_sync: upload failed for photo %s: %s", event.photo_id, exc diff --git a/app/worker/event_lifecycle/main.py b/app/worker/event_lifecycle/main.py index fc1c3529..57080650 100644 --- a/app/worker/event_lifecycle/main.py +++ b/app/worker/event_lifecycle/main.py @@ -42,7 +42,7 @@ async def run_storage_cleanup_pass() -> None: cleaned = 0 for photo in due_photos: try: - await NatsClient.publish( + await NatsClient.js_publish( NatsSubjects.FINAL_BUCKET_CLEANUP, json.dumps({"storage_keys": [photo.storage_key]}).encode("utf-8"), ) diff --git a/tests/integration/test_drive_sync_flow.py b/tests/integration/test_drive_sync_flow.py new file mode 100644 index 00000000..98339e8e --- /dev/null +++ b/tests/integration/test_drive_sync_flow.py @@ -0,0 +1,164 @@ +""" +Integration tests for the Drive Sync worker flow. +""" + +import uuid +import json +import datetime +from unittest.mock import AsyncMock, patch, MagicMock +import contextlib + +import pytest +from sqlalchemy.ext.asyncio import create_async_engine +from fastapi import HTTPException + +from app.core.config import settings +from app.worker.drive_sync.main import _handle_event + +from db.generated import staff_user as staff_queries +from db.generated import events as event_queries +from db.generated import photos as photo_queries +from db.generated import staff_drive_connections as drive_queries + +pytestmark = pytest.mark.integration + +@pytest.fixture +async def db_conn(): + url = f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" + engine = create_async_engine(url, pool_pre_ping=True) + async with engine.connect() as conn: + yield conn + await engine.dispose() + +@pytest.fixture +async def setup_data(db_conn): + sq = staff_queries.AsyncQuerier(db_conn) + eq = event_queries.AsyncQuerier(db_conn) + pq = photo_queries.AsyncQuerier(db_conn) + dq = drive_queries.AsyncQuerier(db_conn) + + staff = await sq.create_admin(email=f"admin-{uuid.uuid4()}@test.com", password="hash") + + event = await eq.create_event( + event_queries.CreateEventParams( + name="Drive Sync Test Event", + event_code=f"DS{str(uuid.uuid4())[:4]}", + event_date=datetime.datetime.now(datetime.timezone.utc), + end_date=datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(days=1), + status="scheduled", + created_by=staff.id, + ) + ) + + photo = await pq.create_photo( + photo_queries.CreatePhotoParams( + event_id=event.id, + storage_key="test/sync_photo.jpg", + source="direct", + taken_at=None, + day_number=None, + visibility="public", + ) + ) + + conn = await dq.upsert_staff_drive_connection( + drive_queries.UpsertStaffDriveConnectionParams( + staff_user_id=staff.id, + provider="google_drive", + google_email="test@google.com", + google_account_id="12345", + access_token="encrypted_access_token", + refresh_token="encrypted_refresh_token", + token_expires_at=datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(days=1), + scopes="scopes", + ) + ) + + return { + "staff": staff, + "event": event, + "photo": photo, + "connection": conn, + } + +@pytest.mark.asyncio +@patch("app.worker.drive_sync.main.Bucket") +@patch("app.worker.drive_sync.main.StaffDriveService") +async def test_successful_drive_sync(mock_staff_drive_service_class, mock_bucket_class, db_conn, setup_data): + """Test that a photo sync event successfully triggers the drive upload.""" + mock_bucket = AsyncMock() + mock_bucket.get.return_value = (b"fake_image_data", None, "image/jpeg") + mock_bucket_class.return_value = mock_bucket + + mock_service = AsyncMock() + mock_service.upload_to_system_drive.return_value = "google_drive_file_id_123" + mock_staff_drive_service_class.return_value = mock_service + + @contextlib.asynccontextmanager + async def mock_begin(): + yield db_conn + + mock_engine = MagicMock() + mock_engine.begin = mock_begin + + payload = { + "photo_id": str(setup_data["photo"].id), + "event_id": str(setup_data["event"].id), + "storage_key": setup_data["photo"].storage_key, + "file_name": "sync_photo.jpg", + "mime_type": "image/jpeg", + } + raw_data = json.dumps(payload).encode("utf-8") + + with patch("app.worker.drive_sync.main.engine", mock_engine), \ + patch("app.worker.drive_sync.main.RedisClient.get_instance", return_value=AsyncMock()): + await _handle_event(raw_data) + + mock_service.upload_to_system_drive.assert_called_once_with( + file_name="sync_photo.jpg", + content_type="image/jpeg", + data=b"fake_image_data", + event_id=setup_data["event"].id, + event_name="Drive Sync Test Event" + ) + + pq = photo_queries.AsyncQuerier(db_conn) + updated_photo = await pq.get_photo_by_id(id=setup_data["photo"].id) + assert updated_photo.drive_file_id == "google_drive_file_id_123" + +@pytest.mark.asyncio +@patch("app.worker.drive_sync.main.Bucket") +@patch("app.worker.drive_sync.main.StaffDriveService") +async def test_drive_sync_pauses_on_revoked_token(mock_staff_drive_service_class, mock_bucket_class, db_conn, setup_data): + """Test that the worker gracefully pauses when the token is revoked.""" + mock_bucket = AsyncMock() + mock_bucket.get.return_value = (b"fake_image_data", None, "image/jpeg") + mock_bucket_class.return_value = mock_bucket + + mock_service = AsyncMock() + mock_service.upload_to_system_drive.side_effect = HTTPException(status_code=404, detail="Drive connection revoked") + mock_staff_drive_service_class.return_value = mock_service + + @contextlib.asynccontextmanager + async def mock_begin(): + yield db_conn + + mock_engine = MagicMock() + mock_engine.begin = mock_begin + + payload = { + "photo_id": str(setup_data["photo"].id), + "event_id": str(setup_data["event"].id), + "storage_key": setup_data["photo"].storage_key, + "file_name": "sync_photo.jpg", + "mime_type": "image/jpeg", + } + raw_data = json.dumps(payload).encode("utf-8") + + with patch("app.worker.drive_sync.main.asyncio.sleep", new_callable=AsyncMock) as mock_sleep, \ + patch("app.worker.drive_sync.main.engine", mock_engine), \ + patch("app.worker.drive_sync.main.RedisClient.get_instance", return_value=AsyncMock()): + with pytest.raises(HTTPException): + await _handle_event(raw_data) + + mock_sleep.assert_called_once_with(300) From aca90a8a1b993922442ecd02e0b766de2270536d Mon Sep 17 00:00:00 2001 From: Adem Boukabes <142881379+ademboukabes@users.noreply.github.com> Date: Sun, 30 Aug 2026 18:44:30 +0100 Subject: [PATCH 27/29] feat: robustesse drive sync (redis lock, backoff, ipv4) --- app/service/staff_drive.py | 4 +- app/worker/drive_sync/main.py | 45 +++++++++++++++++++- tests/integration/test_drive_sync_flow.py | 50 +++++++++++++++++------ 3 files changed, 85 insertions(+), 14 deletions(-) diff --git a/app/service/staff_drive.py b/app/service/staff_drive.py index cd0188a5..abface71 100644 --- a/app/service/staff_drive.py +++ b/app/service/staff_drive.py @@ -291,7 +291,9 @@ async def upload_to_system_drive( folder_id = settings.GOOGLE_CLUB_DRIVE_FOLDER_ID or None if event_id and event_name: - folder_id = await self._get_or_create_event_folder(event_id, event_name, access_token) + folder_id = await self._get_or_create_event_folder( + event_id, event_name, access_token + ) metadata = await GoogleDriveClient.upload_file( access_token=access_token, diff --git a/app/worker/drive_sync/main.py b/app/worker/drive_sync/main.py index 73c4ee4f..ba695c08 100644 --- a/app/worker/drive_sync/main.py +++ b/app/worker/drive_sync/main.py @@ -1,6 +1,8 @@ import asyncio import json import uuid +import socket + from pydantic import BaseModel, ValidationError @@ -17,6 +19,23 @@ from db.generated import staff_user as staff_queries from fastapi import HTTPException +# Force IPv4 to avoid "Network is unreachable" on systems without IPv6 routing +old_getaddrinfo = socket.getaddrinfo + + +def new_getaddrinfo(*args, **kwargs): # type: ignore[no-untyped-def] + args_list = list(args) + # family is the 3rd positional argument + if len(args_list) >= 3: + if args_list[2] == 0: + args_list[2] = socket.AF_INET + elif "family" not in kwargs or kwargs["family"] == 0: + kwargs["family"] = socket.AF_INET + return old_getaddrinfo(*args_list, **kwargs) + + +socket.getaddrinfo = new_getaddrinfo + class PhotoDriveSyncEvent(BaseModel): photo_id: uuid.UUID @@ -49,6 +68,17 @@ async def _handle_event(raw_data: bytes) -> None: bucket = Bucket(IMAGES_BUCKET_NAME, "") try: data, _, content_type = await bucket.get(event.storage_key) + except HTTPException as exc: + if exc.status_code == 404: + logger.warning( + "drive_sync: photo %s not found in storage (404), skipping sync to avoid retry loop", + event.photo_id, + ) + return + logger.warning( + "drive_sync: failed to read photo %s from storage: %s", event.photo_id, exc + ) + raise except Exception as exc: logger.warning( "drive_sync: failed to read photo %s from storage: %s", event.photo_id, exc @@ -66,7 +96,11 @@ async def _handle_event(raw_data: bytes) -> None: event_querier = event_queries.AsyncQuerier(conn) db_event = await event_querier.get_event_by_id(id=event.event_id) if db_event is None: - logger.warning("drive_sync: event %s not found for photo %s", event.event_id, event.photo_id) + logger.warning( + "drive_sync: event %s not found for photo %s", + event.event_id, + event.photo_id, + ) return try: @@ -83,11 +117,20 @@ async def _handle_event(raw_data: bytes) -> None: "drive_sync: no active drive connection, pausing for 5 minutes before retry" ) await asyncio.sleep(300) + else: + logger.warning( + "drive_sync: upload failed for photo %s (HTTP %s): %s", + event.photo_id, + exc.status_code, + exc, + ) + await asyncio.sleep(15) # Backoff to avoid fast retry loops raise except Exception as exc: logger.warning( "drive_sync: upload failed for photo %s: %s", event.photo_id, exc ) + await asyncio.sleep(15) # Backoff to avoid fast retry loops raise synced = await photo_querier.mark_photo_drive_synced( diff --git a/tests/integration/test_drive_sync_flow.py b/tests/integration/test_drive_sync_flow.py index 98339e8e..0bc2665c 100644 --- a/tests/integration/test_drive_sync_flow.py +++ b/tests/integration/test_drive_sync_flow.py @@ -22,6 +22,7 @@ pytestmark = pytest.mark.integration + @pytest.fixture async def db_conn(): url = f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" @@ -30,6 +31,7 @@ async def db_conn(): yield conn await engine.dispose() + @pytest.fixture async def setup_data(db_conn): sq = staff_queries.AsyncQuerier(db_conn) @@ -37,14 +39,17 @@ async def setup_data(db_conn): pq = photo_queries.AsyncQuerier(db_conn) dq = drive_queries.AsyncQuerier(db_conn) - staff = await sq.create_admin(email=f"admin-{uuid.uuid4()}@test.com", password="hash") + staff = await sq.create_admin( + email=f"admin-{uuid.uuid4()}@test.com", password="hash" + ) event = await eq.create_event( event_queries.CreateEventParams( name="Drive Sync Test Event", event_code=f"DS{str(uuid.uuid4())[:4]}", event_date=datetime.datetime.now(datetime.timezone.utc), - end_date=datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(days=1), + end_date=datetime.datetime.now(datetime.timezone.utc) + + datetime.timedelta(days=1), status="scheduled", created_by=staff.id, ) @@ -69,7 +74,8 @@ async def setup_data(db_conn): google_account_id="12345", access_token="encrypted_access_token", refresh_token="encrypted_refresh_token", - token_expires_at=datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(days=1), + token_expires_at=datetime.datetime.now(datetime.timezone.utc) + + datetime.timedelta(days=1), scopes="scopes", ) ) @@ -81,10 +87,13 @@ async def setup_data(db_conn): "connection": conn, } + @pytest.mark.asyncio @patch("app.worker.drive_sync.main.Bucket") @patch("app.worker.drive_sync.main.StaffDriveService") -async def test_successful_drive_sync(mock_staff_drive_service_class, mock_bucket_class, db_conn, setup_data): +async def test_successful_drive_sync( + mock_staff_drive_service_class, mock_bucket_class, db_conn, setup_data +): """Test that a photo sync event successfully triggers the drive upload.""" mock_bucket = AsyncMock() mock_bucket.get.return_value = (b"fake_image_data", None, "image/jpeg") @@ -110,8 +119,13 @@ async def mock_begin(): } raw_data = json.dumps(payload).encode("utf-8") - with patch("app.worker.drive_sync.main.engine", mock_engine), \ - patch("app.worker.drive_sync.main.RedisClient.get_instance", return_value=AsyncMock()): + with ( + patch("app.worker.drive_sync.main.engine", mock_engine), + patch( + "app.worker.drive_sync.main.RedisClient.get_instance", + return_value=AsyncMock(), + ), + ): await _handle_event(raw_data) mock_service.upload_to_system_drive.assert_called_once_with( @@ -119,24 +133,29 @@ async def mock_begin(): content_type="image/jpeg", data=b"fake_image_data", event_id=setup_data["event"].id, - event_name="Drive Sync Test Event" + event_name="Drive Sync Test Event", ) pq = photo_queries.AsyncQuerier(db_conn) updated_photo = await pq.get_photo_by_id(id=setup_data["photo"].id) assert updated_photo.drive_file_id == "google_drive_file_id_123" + @pytest.mark.asyncio @patch("app.worker.drive_sync.main.Bucket") @patch("app.worker.drive_sync.main.StaffDriveService") -async def test_drive_sync_pauses_on_revoked_token(mock_staff_drive_service_class, mock_bucket_class, db_conn, setup_data): +async def test_drive_sync_pauses_on_revoked_token( + mock_staff_drive_service_class, mock_bucket_class, db_conn, setup_data +): """Test that the worker gracefully pauses when the token is revoked.""" mock_bucket = AsyncMock() mock_bucket.get.return_value = (b"fake_image_data", None, "image/jpeg") mock_bucket_class.return_value = mock_bucket mock_service = AsyncMock() - mock_service.upload_to_system_drive.side_effect = HTTPException(status_code=404, detail="Drive connection revoked") + mock_service.upload_to_system_drive.side_effect = HTTPException( + status_code=404, detail="Drive connection revoked" + ) mock_staff_drive_service_class.return_value = mock_service @contextlib.asynccontextmanager @@ -155,9 +174,16 @@ async def mock_begin(): } raw_data = json.dumps(payload).encode("utf-8") - with patch("app.worker.drive_sync.main.asyncio.sleep", new_callable=AsyncMock) as mock_sleep, \ - patch("app.worker.drive_sync.main.engine", mock_engine), \ - patch("app.worker.drive_sync.main.RedisClient.get_instance", return_value=AsyncMock()): + with ( + patch( + "app.worker.drive_sync.main.asyncio.sleep", new_callable=AsyncMock + ) as mock_sleep, + patch("app.worker.drive_sync.main.engine", mock_engine), + patch( + "app.worker.drive_sync.main.RedisClient.get_instance", + return_value=AsyncMock(), + ), + ): with pytest.raises(HTTPException): await _handle_event(raw_data) From 44ae3dcec2c7da88e462fe542459d7387e50bd08 Mon Sep 17 00:00:00 2001 From: Adem Boukabes <142881379+ademboukabes@users.noreply.github.com> Date: Mon, 31 Aug 2026 14:09:46 +0100 Subject: [PATCH 28/29] fix(drive-sync): robustify Redis folder lock (M-1) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - TTL 30s → 120s to survive slow Drive API calls - Lock value = unique UUID per acquisition (prevent foreign release) - Async heartbeat refreshes TTL every 40s to avoid expiry mid-operation - Extract _resolve_drive_folder to stay under ruff C901 limit - Add missing logger import in staff_drive.py --- app/service/staff_drive.py | 103 ++++++++++++++++++++++++++++++------- 1 file changed, 83 insertions(+), 20 deletions(-) diff --git a/app/service/staff_drive.py b/app/service/staff_drive.py index abface71..bbf1b6b0 100644 --- a/app/service/staff_drive.py +++ b/app/service/staff_drive.py @@ -12,6 +12,7 @@ from cryptography.fernet import Fernet, InvalidToken from app.core.config import settings +from app.core.logger import logger from app.core.constant import IMAGE_ALLOWED_TYPES from app.core.exceptions import AppException from app.infra.google_drive import GoogleDriveClient @@ -230,6 +231,60 @@ async def get_system_access_token(self) -> str: connection = await self._refresh_connection_access_token(connection) return self.decrypt(connection.access_token) + # How long the folder-creation lock is held (initial TTL). + # Drive API calls (search + optional create) can take up to ~60s on a slow + # network; 120s gives a comfortable margin before the heartbeat is even needed. + _FOLDER_LOCK_TTL_SECONDS: int = 120 + # The heartbeat renews the lock every N seconds to keep it alive during + # slow Drive API calls. Must be well below _FOLDER_LOCK_TTL_SECONDS. + _FOLDER_LOCK_HEARTBEAT_SECONDS: int = 40 + + async def _lock_heartbeat( + self, + lock_key: str, + lock_value: str, + stop: asyncio.Event, + ) -> None: + """Renews the Redis lock TTL every _FOLDER_LOCK_HEARTBEAT_SECONDS to + prevent it from expiring during slow Drive API calls.""" + while not stop.is_set(): + await asyncio.sleep(self._FOLDER_LOCK_HEARTBEAT_SECONDS) + if stop.is_set(): + break + try: + current = await self.redis.get(lock_key) + if current != lock_value: + logger.warning( + "drive_folder_lock: lock %s no longer ours, stopping heartbeat", + lock_key, + ) + break + await self.redis.expire(lock_key, self._FOLDER_LOCK_TTL_SECONDS) + logger.debug("drive_folder_lock: refreshed TTL for %s", lock_key) + except Exception as exc: # pragma: no cover + logger.warning( + "drive_folder_lock: heartbeat error for %s: %s", lock_key, exc + ) + + async def _resolve_drive_folder( + self, + event_name: str, + access_token: str, + ) -> str: + """Find an existing Drive folder with *event_name* or create one. + Returns the folder ID.""" + parent_id = settings.GOOGLE_CLUB_DRIVE_FOLDER_ID or None + existing = await GoogleDriveClient.search_files( + access_token=access_token, query=event_name, file_type="folder" + ) + for folder in existing or []: + if folder.name == event_name: + return folder.id + folder_meta = await GoogleDriveClient.create_folder( + access_token=access_token, name=event_name, parent_id=parent_id + ) + return folder_meta.id + async def _get_or_create_event_folder( self, event_id: uuid.UUID, @@ -242,39 +297,47 @@ async def _get_or_create_event_folder( return cached_id lock_key = f"lock:drive:folder:{event_id}" + # Unique value: only the holder can release its own lock. + lock_value = str(uuid.uuid4()) + while True: - acquired = await self.redis.set(lock_key, "1", expire=30, nx=True) + acquired = await self.redis.set( + lock_key, + lock_value, + expire=self._FOLDER_LOCK_TTL_SECONDS, + nx=True, + ) if acquired: break await asyncio.sleep(1.0) + stop_heartbeat = asyncio.Event() + # Background heartbeat keeps the lock alive during slow Drive API calls. + heartbeat_task = asyncio.create_task( + self._lock_heartbeat(lock_key, lock_value, stop_heartbeat) + ) + try: + # Double-check: another waiter may have populated the cache while + # we were spinning in the loop above. cached_id = await self.redis.get(cache_key) if cached_id: return cached_id - parent_id = settings.GOOGLE_CLUB_DRIVE_FOLDER_ID or None - - existing = await GoogleDriveClient.search_files( - access_token=access_token, query=event_name, file_type="folder" - ) - folder_id = None - if existing: - for folder in existing: - if folder.name == event_name: - folder_id = folder.id - break - - if not folder_id: - folder_meta = await GoogleDriveClient.create_folder( - access_token=access_token, name=event_name, parent_id=parent_id - ) - folder_id = folder_meta.id - + folder_id = await self._resolve_drive_folder(event_name, access_token) await self.redis.set(cache_key, folder_id, expire=7 * 24 * 3600) return folder_id finally: - await self.redis.delete(lock_key) + stop_heartbeat.set() + heartbeat_task.cancel() + try: + await heartbeat_task + except asyncio.CancelledError: + pass + # Only release the lock if it still belongs to us. + current_value = await self.redis.get(lock_key) + if current_value == lock_value: + await self.redis.delete(lock_key) async def upload_to_system_drive( self, From 78c4de4a738d52a28d701946931e4dab5c23fcaa Mon Sep 17 00:00:00 2001 From: Adem Boukabes <142881379+ademboukabes@users.noreply.github.com> Date: Mon, 31 Aug 2026 14:59:40 +0100 Subject: [PATCH 29/29] ci: add CD for develop branch --- .github/workflows/docker-publish.yml | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/.github/workflows/docker-publish.yml b/.github/workflows/docker-publish.yml index ea6be766..aff82010 100644 --- a/.github/workflows/docker-publish.yml +++ b/.github/workflows/docker-publish.yml @@ -4,6 +4,7 @@ on: push: branches: - main + - develop workflow_dispatch: permissions: @@ -21,10 +22,17 @@ jobs: - name: Checkout uses: actions/checkout@v4 - - name: Set lowercase image name + - name: Set lowercase image name and Docker tag shell: bash run: | echo "IMAGE_NAME=${GITHUB_REPOSITORY,,}" >> "$GITHUB_ENV" + if [ "${{ github.ref }}" == "refs/heads/main" ]; then + echo "DOCKER_TAG=latest" >> "$GITHUB_ENV" + elif [ "${{ github.ref }}" == "refs/heads/develop" ]; then + echo "DOCKER_TAG=develop" >> "$GITHUB_ENV" + else + echo "DOCKER_TAG=${GITHUB_REF_NAME}" >> "$GITHUB_ENV" + fi - name: Set up QEMU uses: docker/setup-qemu-action@v3 @@ -46,5 +54,5 @@ jobs: push: true platforms: linux/amd64,linux/arm64 tags: | - ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }}:latest + ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }}:${{ env.DOCKER_TAG }} ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }}:${{ github.sha }} \ No newline at end of file