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 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 dea2f69a..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 @@ -34,6 +39,20 @@ class Settings(BaseSettings): POSTGRES_PORT: int = 5432 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 + 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 @@ -54,6 +73,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 @@ -75,9 +97,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/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 a14de6c4..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") @@ -114,3 +118,26 @@ 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/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 f25c4ef6..9aa90ff5 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 @@ -159,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( *, @@ -189,6 +235,75 @@ 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( *, @@ -238,11 +353,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 @@ -252,7 +371,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 @@ -288,7 +409,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): @@ -386,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 @@ -407,7 +532,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")) @@ -433,7 +560,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() @@ -447,6 +576,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 e6249dae..5ee2b040 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 @@ -22,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: @@ -36,6 +39,13 @@ 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 @@ -86,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() @@ -130,6 +138,29 @@ 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"], @@ -142,6 +173,7 @@ async def copy(self, *, source_object_name: str, target_object_name: str) -> str "ico": ["image/x-icon", "image/vnd.microsoft.icon"], } + class ImageBucket(Bucket): def __init__(self, file_prefix: str): super().__init__(IMAGES_BUCKET_NAME, file_prefix) @@ -150,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 e17bd500..aaa6fc7a 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): @@ -35,6 +36,28 @@ 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" + 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: @@ -52,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, ) @@ -79,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 @@ -91,38 +118,68 @@ 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, 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, - ack_policy: AckPolicy = AckPolicy.EXPLICIT + 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 @@ -137,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/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..1c524ca7 100644 --- a/app/router/mobile/auth.py +++ b/app/router/mobile/auth.py @@ -20,41 +20,68 @@ RefreshTokenRequest, UpdateDeviceTokenRequest, 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, @@ -63,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, @@ -88,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") @@ -121,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( @@ -131,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"} @@ -206,11 +244,32 @@ 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, ) + +@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, @@ -228,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) @@ -239,6 +299,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/enrollement.py b/app/router/mobile/enrollement.py index f8e2ac10..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( @@ -139,12 +140,20 @@ 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." diff --git a/app/router/mobile/event.py b/app/router/mobile/event.py index adada86a..00b0ac1c 100644 --- a/app/router/mobile/event.py +++ b/app/router/mobile/event.py @@ -3,27 +3,28 @@ 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 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(get_current_mobile_user), -)-> JoinEventResponse: + current_user: MobileUserSchema = Depends(require_onboarded_mobile_user), +) -> 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 ) @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), -)-> List[UserEventResponse]: + 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..e5f73e81 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,10 +24,12 @@ 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, - 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/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..5d0f357e 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( @@ -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 ] @@ -47,7 +48,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( @@ -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 ], @@ -81,7 +83,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/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/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 new file mode 100644 index 00000000..88225de3 --- /dev/null +++ b/app/router/staff/uploads_direct.py @@ -0,0 +1,134 @@ +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 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, + 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) diff --git a/app/router/web/auth.py b/app/router/web/auth.py index 7629cf93..76f945b1 100644 --- a/app/router/web/auth.py +++ b/app/router/web/auth.py @@ -2,47 +2,53 @@ 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 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", value=authResponse.access_token, httponly=True, - secure=True, + secure=settings.environment != "dev", samesite="strict", max_age=60 * 60 * 24 * 7, ) 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/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/schema/request/mobile/auth.py b/app/schema/request/mobile/auth.py index dea9933a..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): @@ -92,3 +94,7 @@ class UpdateDeviceTokenRequest(BaseModel): class InactivateDeviceRequest(BaseModel): device_id: UUID + + +class UpdateProfileRequest(BaseModel): + name: str = Field(..., min_length=1, max_length=100) 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/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/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 d03b46ff..775e6c8b 100644 --- a/app/schema/request/web/event.py +++ b/app/schema/request/web/event.py @@ -2,10 +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 67bf1398..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,27 +18,33 @@ 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 name: str | None 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 7a264c2d..eadf1647 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 @@ -39,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 @@ -55,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): @@ -64,16 +69,16 @@ 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) 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 @@ -94,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/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] 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 4334fc44..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) @@ -10,31 +11,39 @@ 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 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 name: str event_date: datetime + end_date: Optional[datetime] = None 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/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/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 e697ddd8..d947f108 100644 --- a/app/service/event.py +++ b/app/service/event.py @@ -2,46 +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") @@ -57,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( @@ -67,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] = [] @@ -83,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: @@ -93,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 f3b2ba00..bbf1b6b0 100644 --- a/app/service/staff_drive.py +++ b/app/service/staff_drive.py @@ -8,9 +8,11 @@ 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 +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 @@ -110,7 +112,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 +162,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, @@ -170,27 +177,40 @@ 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 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,13 +222,151 @@ 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): 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, + 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}" + # Unique value: only the holder can release its own lock. + lock_value = str(uuid.uuid4()) + + while 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 + + 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: + 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, + *, + 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=folder_id, + ) + 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: @@ -224,9 +382,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, @@ -299,7 +461,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 @@ -324,4 +486,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 813fa2bf..53d14c48 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) @@ -95,8 +95,31 @@ 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) 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/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 0ab404d6..2bf2d7cb 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 @@ -101,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: @@ -124,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", @@ -145,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", @@ -157,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", @@ -181,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", @@ -193,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 @@ -226,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", @@ -252,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", @@ -282,6 +313,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: @@ -326,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": @@ -334,7 +368,20 @@ 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 + 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] = [] @@ -354,6 +401,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: @@ -365,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, @@ -386,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": @@ -420,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, @@ -432,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, @@ -446,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, @@ -468,9 +526,11 @@ 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) + 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: @@ -492,6 +552,35 @@ 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), + "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, + }, + ) + count += 1 + if count: + logger.info("Published %d Drive sync events", count) + async def _mark_group_import_failed( self, *, @@ -565,13 +654,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, + 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: @@ -593,7 +686,260 @@ async def create_group_from_folder( ) return UploadRequestGroupDetails(group=upload_group, requests=[]) - async def process_group_import( + 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") + + 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, + *, + 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 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( # noqa: C901 self, *, group_id: uuid.UUID, @@ -604,8 +950,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) @@ -619,8 +967,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( @@ -636,14 +986,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, @@ -668,7 +1024,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, @@ -711,7 +1069,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( @@ -735,7 +1095,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 @@ -751,7 +1113,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, @@ -767,7 +1131,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( @@ -786,14 +1152,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 @@ -806,7 +1176,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 @@ -821,7 +1195,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() @@ -856,7 +1232,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( @@ -865,7 +1243,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) @@ -891,7 +1271,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 @@ -906,8 +1290,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() @@ -939,11 +1325,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( @@ -968,6 +1357,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, @@ -989,12 +1379,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, @@ -1048,23 +1440,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( @@ -1100,6 +1499,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, @@ -1133,20 +1533,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 f8181bc0..7038f0aa 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") @@ -53,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) @@ -65,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 f10a6303..f53844d4 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, @@ -79,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 @@ -109,10 +114,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") @@ -128,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 3864757b..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) @@ -176,7 +185,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,16 +193,23 @@ 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.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( @@ -226,19 +241,23 @@ 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.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( @@ -261,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: @@ -296,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, @@ -313,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)) @@ -360,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}" @@ -376,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) @@ -387,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, @@ -466,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"} @@ -661,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) @@ -688,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) @@ -701,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, ) @@ -724,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 8037c26e..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 @@ -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/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/photo_worker/tests/__init__.py b/app/worker/drive_sync/__init__.py similarity index 100% rename from app/worker/photo_worker/tests/__init__.py rename to app/worker/drive_sync/__init__.py diff --git a/app/worker/drive_sync/main.py b/app/worker/drive_sync/main.py new file mode 100644 index 00000000..ba695c08 --- /dev/null +++ b/app/worker/drive_sync/main.py @@ -0,0 +1,175 @@ +import asyncio +import json +import uuid +import socket + + +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 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 + +# 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 + event_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 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 + ) + raise + + 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) + + 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) + 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( + 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.js_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/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/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..57080650 --- /dev/null +++ b/app/worker/event_lifecycle/main.py @@ -0,0 +1,82 @@ +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: + 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 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.js_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, 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) + + +if __name__ == "__main__": + asyncio.run(main()) 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 a79d7c20..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, @@ -114,12 +114,13 @@ 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() - 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 463de938..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.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 59a116bb..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, @@ -65,28 +68,23 @@ 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) - 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") - await self._schedule_cleanup(event.image_ref) return if len(faces) == 1: @@ -96,10 +94,10 @@ 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: + async def _handle_single_face( + self, event: PhotoProcessEvent, face: DetectedFace + ) -> None: from app.schema.internal.single_face_match import SingleFaceMatchJob bbox = BBoxPayload( @@ -118,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) + "]" @@ -150,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 @@ -171,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: @@ -198,23 +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.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) - - @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) + await NatsClient.js_publish( + NatsSubjects.AUDIT_EVENT, msg.model_dump_json().encode("utf-8") + ) except Exception as exc: - logger.warning("Failed to schedule cleanup for %s: %s", image_ref, exc) + logger.warning( + "Failed to publish audit for photo %s: %s", event.photo_id, exc + ) @staticmethod def _parse_event(raw_data: bytes) -> PhotoProcessEvent | None: @@ -241,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) @@ -252,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") @@ -283,28 +301,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, @@ -313,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/photo_worker/tests/test_photo_worker.py b/app/worker/photo_worker/tests/test_photo_worker.py deleted file mode 100644 index f113a706..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 ────────────────────────────────────────── - - -@pytest.mark.asyncio -async def test_cleanup_scheduled_after_single_face( - worker: PhotoWorker, - face_service: AsyncMock, - single_face_service: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """After single face processing, both audit and cleanup events should be published.""" - 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 == 2 - 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( - worker: PhotoWorker, - face_service: AsyncMock, - photo_face_querier: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """After group photo processing, both audit and cleanup events should be published.""" - 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 == 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"] - - -@pytest.mark.asyncio -async def test_cleanup_scheduled_when_no_faces( - worker: PhotoWorker, - face_service: AsyncMock, - event: PhotoProcessEvent, -) -> None: - """Even if no faces detected, cleanup should still be 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_called_once() - cleanup_payload = json.loads(mock_nats.publish.call_args.args[1]) - assert event.image_ref in cleanup_payload["storage_keys"] - - -@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/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/__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..2f1c13ef --- /dev/null +++ b/app/worker/upload_reconciler/main.py @@ -0,0 +1,64 @@ +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/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/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..adf1f28a 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() @@ -114,6 +115,10 @@ class Photo: visibility: str status: Any 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() @@ -207,13 +212,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 @@ -226,13 +232,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 @@ -243,6 +250,8 @@ class UploadRequestPhoto: visibility: str status: str created_at: datetime.datetime + source: str + transfer_status: str @dataclasses.dataclass() diff --git a/db/generated/photos.py b/db/generated/photos.py index f757a9fb..666c9896 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 +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,13 @@ 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, 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 +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 AND p.status = 'approved' @@ -94,10 +97,44 @@ 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 + 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 +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 @@ -125,11 +162,46 @@ 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 + drive_file_id: Optional[str] + drive_synced_at: Optional[datetime.datetime] + source: str + storage_cleaned_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, 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 +""" + + 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, source, storage_cleaned_at """ @@ -137,7 +209,7 @@ class ListUserPhotosParams: 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, source, storage_cleaned_at """ @@ -158,6 +230,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 @@ -171,9 +244,13 @@ 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], + source=row[11], + storage_cleaned_at=row[12], ) - 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 @@ -193,9 +270,13 @@ 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], + source=row[11], + storage_cleaned_at=row[12], ) - 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, @@ -203,6 +284,26 @@ async def list_event_photos_for_user(self, arg: ListEventPhotosForUserParams) -> "p4": arg.limit, "p5": arg.offset, }) + async for row in result: + yield ListEventPhotosForUserRow( + 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], + 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], @@ -214,9 +315,13 @@ async def list_event_photos_for_user(self, arg: ListEventPhotosForUserParams) -> 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[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 +330,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,8 +340,53 @@ async def list_user_photos(self, arg: ListUserPhotosParams) -> AsyncIterator[mod 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], + face_count=row[13], ) + 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], + 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]: row = (await self._conn.execute(sqlalchemy.text(UPDATE_PHOTO_STATUS), {"p1": id, "p2": status})).first() if row is None: @@ -251,6 +401,10 @@ 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], + source=row[11], + storage_cleaned_at=row[12], ) async def update_photo_visibility(self, *, id: uuid.UUID, visibility: str) -> Optional[models.Photo]: @@ -267,4 +421,8 @@ 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], + source=row[11], + storage_cleaned_at=row[12], ) 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/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/db/queries/photos.sql b/db/queries/photos.sql index 993c397a..64dbc08c 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 *; @@ -26,9 +27,11 @@ 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 ( +WHERE p.status = 'approved' +AND ( EXISTS ( SELECT 1 FROM photo_faces pf JOIN face_matches fm ON fm.photo_face_id = pf.id @@ -46,7 +49,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' @@ -82,3 +86,28 @@ 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 *; + +-- 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/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 *; diff --git a/makefile b/makefile index 241764ae..7e19b776 100644 --- a/makefile +++ b/makefile @@ -65,6 +65,9 @@ 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 & \ + uv run python -m app.worker.upload_reconciler.main & \ + uv run python -m app.worker.drive_sync.main & \ wait lint: 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/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/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/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/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_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/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/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/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/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") 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/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") 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") 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 ed3e0766..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.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.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.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 0f9437bc..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( @@ -153,7 +157,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( { @@ -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 0a7b3b5a..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: @@ -37,7 +42,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,14 +94,16 @@ 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") ) # 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_drive_sync_flow.py b/tests/integration/test_drive_sync_flow.py new file mode 100644 index 00000000..0bc2665c --- /dev/null +++ b/tests/integration/test_drive_sync_flow.py @@ -0,0 +1,190 @@ +""" +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) 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 2cfb9f4a..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,8 +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), status="scheduled", - created_by=event_creator_id + created_by=event_creator_id, ) ) event_id = event.id @@ -115,42 +125,63 @@ 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" + 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() @@ -170,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 @@ -185,8 +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), status="scheduled", - created_by=event_creator_id + created_by=event_creator_id, ) ) event_id = event.id @@ -195,21 +232,28 @@ 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" + 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" @@ -219,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 7e6acd95..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,8 +53,10 @@ def auth_service( refresh_token_querier=mock_refresh_token_querier, ) + @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, @@ -77,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, @@ -92,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() @@ -123,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, @@ -140,8 +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.NatsClient.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" @@ -150,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 new file mode 100644 index 00000000..6c6020fe --- /dev/null +++ b/tests/unit/test_direct_uploads.py @@ -0,0 +1,556 @@ +import uuid +from datetime import datetime, timezone +from unittest.mock import AsyncMock, patch + +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() + + 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_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_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, + 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", + ) + 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.js_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, + 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 + ) 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 66115611..a188238c 100644 --- a/tests/unit/test_minio.py +++ b/tests/unit/test_minio.py @@ -8,10 +8,12 @@ from app.infra.minio import ( Bucket, ImageBucket, + ObjectStat, WaSimBucket, init_minio_client, ) + @pytest.fixture def mock_minio_client(): client = AsyncMock() @@ -113,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 @@ -137,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() @@ -185,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) @@ -200,3 +202,51 @@ 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_mobile_auth_email_logging.py b/tests/unit/test_mobile_auth_email_logging.py index d868f324..681f9ed8 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 @@ -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 @@ -148,7 +147,11 @@ def test_mobile_register_logs_without_plaintext_email( 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("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_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 9b920307..426dddf8 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 @@ -20,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 @@ -89,69 +87,138 @@ 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.publish") as mock_publish: - await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) + 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() - 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")) + 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). + mock_publish.assert_not_called() @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.publish") as mock_publish: - await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) + 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") + 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 -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.publish"): - await photo_worker.handle_message(sample_event.model_dump_json().encode("utf-8")) + 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") + 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")) - 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")) +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 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 691c30c8..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() @@ -122,22 +131,27 @@ 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( - 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), + 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 = [ @@ -156,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.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, @@ -191,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 - ) + 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 @@ -224,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 + 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.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 + 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 @@ -249,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 - ) - 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 - ) + 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 - ) + 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 @@ -276,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 @@ -294,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 - ) + 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 mock_staged_upload_storage.promote_to_final.return_value = "final_key" @@ -308,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 )