"""Bulk Session CRUD operation.""" import uuid from collections.abc import Sequence from datetime import datetime, timedelta from typing import Any, Dict from pydantic import UUID1, UUID4 from sqlalchemy import ( CHAR, TIMESTAMP, Boolean, Enum, and_, case, func, literal, literal_column, or_, select, text, update, INT, ) from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, mapped_column, relationship from product_staging.api.schemas.bulk_session import ( AssetStatus, MetadataStatus, UpdateBulkSessionRequest, ) from product_staging.connectors import db class BulkSession(db.base_model): """Table definition for bulk_session.""" __tablename__ = "bulk_sessions" __allow_unmapped__ = True bulk_session_id = mapped_column(CHAR(36), primary_key=True) vendor_uuid = mapped_column(CHAR(36), nullable=False) subaccount_id = mapped_column(INT) metadata_status: Mapped[MetadataStatus] = mapped_column( Enum(MetadataStatus), nullable=False, default="uploading", ) asset_status: Mapped[AssetStatus] = mapped_column( Enum(AssetStatus), nullable=False, default="default", ) metadata_error_report_json = mapped_column(CHAR(61)) created_on = mapped_column( TIMESTAMP(timezone=True), nullable=False, server_default="CURRENT_TIMESTAMP" ) created_by = mapped_column(CHAR(36), nullable=False) updated_on = mapped_column( TIMESTAMP(timezone=True), nullable=False, server_default=text("CURRENT_TIMESTAMP"), onupdate=text("CURRENT_TIMESTAMP"), ) updated_by = mapped_column(CHAR(36), nullable=False) asset_error_report = mapped_column(CHAR(50)) asset_report_json = mapped_column(CHAR(55)) is_cancelled = mapped_column(Boolean, nullable=False, default=False) bulk_session_metadata_files = relationship( "BulkSessionMetadataFile", primaryjoin="BulkSession.bulk_session_id==BulkSessionMetadataFile.bulk_session_id", foreign_keys="[BulkSessionMetadataFile.bulk_session_id]", uselist=True, order_by="desc(BulkSessionMetadataFile.created_on)", lazy="select", ) bulk_session_assets = relationship( "BulkSessionAsset", primaryjoin="BulkSession.bulk_session_id==foreign(BulkSessionAsset.bulk_session_id)", uselist=True, lazy="select", ) download_link: str | None = None vendor_id: int | None = None _cached_success_file_id = None _cached_failure_file_id = None _cached_validated_metadata_json_file = None _cached_latest_upload_file = None _cached_latest_upload_file_id = None _cached_latest_upload_original_filename = None _cached_total_products = None _cached_assets_complete = None _cached_failed_product_uploads = None @db.db_session_wrap async def latest_upload_file(self, session=None): from product_staging.models.bulk_session_metadata_file import ( BulkSessionMetadataFile, ) assert session if self._cached_latest_upload_file is not None: return self._cached_latest_upload_file stmt = ( select(BulkSessionMetadataFile) .where( BulkSessionMetadataFile.bulk_session_id == self.bulk_session_id, BulkSessionMetadataFile.upload_token.isnot(None), ) .order_by(BulkSessionMetadataFile.created_on.desc()) # optional .limit(1) ) result = await session.execute(stmt) self._cached_latest_upload_file = result.scalar_one_or_none() return ( self._cached_latest_upload_file if self._cached_latest_upload_file else None ) async def latest_upload_file_id(self): if self._cached_latest_upload_file_id is not None: return self._cached_latest_upload_file_id latest_upload = await self.latest_upload_file() self._cached_latest_upload_file_id = ( latest_upload.bulk_session_metadata_file_id if latest_upload else None ) return self._cached_latest_upload_file_id async def latest_upload_original_filename(self): if self._cached_latest_upload_original_filename is not None: return self._cached_latest_upload_original_filename latest_upload = await self.latest_upload_file() self._cached_latest_upload_original_filename = ( latest_upload.original_filename if latest_upload else None ) return self._cached_latest_upload_original_filename @db.db_session_wrap async def success_file_id(self, session=None): assert session from product_staging.models.bulk_session_metadata_file import ( BulkSessionMetadataFile, ) if self._cached_success_file_id is not None: return self._cached_success_file_id stmt = ( select(BulkSessionMetadataFile.bulk_session_metadata_file_id) .where( BulkSessionMetadataFile.bulk_session_id == self.bulk_session_id, BulkSessionMetadataFile.file_status == "success", ) .order_by(BulkSessionMetadataFile.created_on.desc()) .limit(1) ) result = await session.execute(stmt) self._cached_success_file_id = result.scalar_one_or_none() return self._cached_success_file_id @db.db_session_wrap async def failure_file_id(self, session=None): assert session from product_staging.models.bulk_session_metadata_file import ( BulkSessionMetadataFile, ) if self._cached_failure_file_id is not None: return self._cached_failure_file_id stmt = ( select(BulkSessionMetadataFile.bulk_session_metadata_file_id) .where( BulkSessionMetadataFile.bulk_session_id == self.bulk_session_id, BulkSessionMetadataFile.file_status == "failure", ) .order_by(BulkSessionMetadataFile.created_on.desc()) .limit(1) ) result = await session.execute(stmt) self._cached_failure_file_id = result.scalar_one_or_none() return self._cached_failure_file_id @db.db_session_wrap async def validated_metadata_json_file(self, session=None): assert session from product_staging.models.bulk_session_metadata_file import ( BulkSessionMetadataFile, ) if self._cached_validated_metadata_json_file is not None: return self._cached_validated_metadata_json_file stmt = ( select(BulkSessionMetadataFile.json_file_name) .where( BulkSessionMetadataFile.bulk_session_id == self.bulk_session_id, BulkSessionMetadataFile.file_status == "success", BulkSessionMetadataFile.json_file_name.isnot(None), ) .order_by(BulkSessionMetadataFile.created_on.desc()) .limit(1) ) result = await session.execute(stmt) self._cached_validated_metadata_json_file = result.scalar_one_or_none() return self._cached_validated_metadata_json_file @db.db_session_wrap async def total_products(self, session=None): assert session from product_staging.models.bulk_session_metadata_file import ( BulkSessionMetadataFile, ) if self._cached_total_products is not None: return self._cached_total_products file_id = await self.latest_upload_file_id() if not file_id: return None stmt = ( select(BulkSessionMetadataFile.total_products) .where( BulkSessionMetadataFile.bulk_session_id == self.bulk_session_id, BulkSessionMetadataFile.bulk_session_metadata_file_id == file_id, ) .order_by(BulkSessionMetadataFile.created_on.desc()) .limit(1) ) result = await session.execute(stmt) self._cached_total_products = result.scalar_one_or_none() return self._cached_total_products @db.db_session_wrap async def assets_complete(self, session=None) -> bool: """Check if all required assets are complete (have associated asset files).""" from product_staging.models.bulk_session_asset import ( BulkSessionAsset, ) assert session if self._cached_assets_complete is not None: return self._cached_assets_complete # Check if there are any required assets without asset files stmt = ( select(BulkSessionAsset) .where( BulkSessionAsset.bulk_session_id == self.bulk_session_id, BulkSessionAsset.bulk_session_asset_file_id.is_(None), BulkSessionAsset.required.is_(True), ) .limit(1) ) incomplete_assets = await session.execute(stmt) self._cached_assets_complete = incomplete_assets.scalar_one_or_none() is None return self._cached_assets_complete @db.db_session_wrap async def failed_product_uploads(self, session=None) -> int | None: """Count distinct products with at least one failed or missing required asset. Returns None if asset_status is not 'complete' or 'incomplete'. Otherwise returns the count of products with failed assets. """ from product_staging.models.bulk_session_asset import ( BulkSessionAsset, ) assert session if self.asset_status not in ("complete", "incomplete"): return None if self._cached_failed_product_uploads is not None: return self._cached_failed_product_uploads products_with_failures = set() stmt = select(BulkSessionAsset.product_code).where( BulkSessionAsset.bulk_session_id == self.bulk_session_id, BulkSessionAsset.bulk_session_asset_file_id.is_(None), BulkSessionAsset.required.is_(True), ) result = await session.execute(stmt) products_with_failures.update(result.scalars().all()) self._cached_failed_product_uploads = len(products_with_failures) return self._cached_failed_product_uploads def to_dict(self): """Return a dictionary of bulk_session.""" return { "bulk_session_id": self.bulk_session_id, "vendor_uuid": self.vendor_uuid, "vendor_id": self.vendor_id, "subaccount_id": self.subaccount_id, "metadata_status": self.metadata_status, "asset_status": self.asset_status, "metadata_error_report_json": self.metadata_error_report_json, "created_on": self.created_on, "created_by": self.created_by, "updated_on": self.updated_on, "updated_by": self.updated_by, "asset_error_report": self.asset_error_report, "asset_report_json": self.asset_report_json, "is_cancelled": self.is_cancelled, "success_file_id": self._cached_success_file_id, "failure_file_id": self._cached_failure_file_id, "validated_metadata_json_file": self._cached_validated_metadata_json_file, "latest_upload_file_id": self._cached_latest_upload_file_id, "latest_upload_original_filename": self._cached_latest_upload_original_filename, "total_products": self._cached_total_products, "assets_complete": self._cached_assets_complete, "failed_product_uploads": self._cached_failed_product_uploads, "download_link": self.download_link, "latest_cloud_transfer_job_id": getattr( self, "latest_cloud_transfer_job_id", None ), } @db.db_session_wrap async def _get(bulk_session_id: UUID4, is_cancelled: bool | None = False, session=None): """Retrieve BulkSession by bulk_session_id is_cancelled parameter is used to filter out "cancelled" sessions by default: to reset the filter, pass is_cancelled=None; to get only cancelled sessions pass is_cancelled=True. """ assert session stmt = select(BulkSession).where( BulkSession.bulk_session_id == str(bulk_session_id) ) if is_cancelled is not None: stmt = stmt.where(BulkSession.is_cancelled == is_cancelled) result = await session.execute(stmt) return result.scalar_one_or_none() @db.db_session_wrap async def get_bulk_session( bulk_session_id: UUID4, is_cancelled: bool | None = False, session=None ): """Retrieve BulkSession by bulk_session_id is_cancelled parameter is used to filter out "cancelled" sessions by default: to reset the filter, pass is_cancelled=None; to get only cancelled sessions pass is_cancelled=True. """ return await _get(bulk_session_id, is_cancelled, session=session) @db.db_session_wrap async def get_bulk_session_by_vendor_and_created_on( vendor_uuid: str | UUID1 | UUID4, created_on: datetime, session=None ): """Retrieve a BulkSession by vendor + created-on second (for slug resolution).""" assert session stmt = ( select(BulkSession) .where( BulkSession.vendor_uuid == str(vendor_uuid), BulkSession.created_on >= created_on, BulkSession.created_on < created_on + timedelta(seconds=1), ) .order_by(BulkSession.created_on.desc()) .limit(1) ) result = await session.execute(stmt) return result.scalar_one_or_none() async def _get_many( bulk_session_ids: list[UUID4], is_cancelled: bool | None = False, session=None ) -> Sequence[BulkSession]: """Retrieve BulkSessions by bulk_session_ids is_cancelled parameter is used to filter out "cancelled" sessions by default: to reset the filter, pass is_cancelled=None; to get only cancelled sessions pass is_cancelled=True. """ assert session stmt = select(BulkSession).where( BulkSession.bulk_session_id.in_( [str(bulk_session_id) for bulk_session_id in bulk_session_ids] ) ) if is_cancelled is not None: stmt = stmt.where(BulkSession.is_cancelled == is_cancelled) result = await session.execute(stmt) return result.scalars().all() @db.db_session_wrap async def get_bulk_sessions( bulk_session_ids: list[UUID4], is_cancelled: bool | None = False, session=None ) -> Sequence[BulkSession]: """Retrieve BulkSessions by bulk_session_ids is_cancelled parameter is used to filter out "cancelled" sessions by default: to reset the filter, pass is_cancelled=None; to get only cancelled sessions pass is_cancelled=True. """ return await _get_many(bulk_session_ids, is_cancelled, session=session) @db.db_session_wrap @db.wrap_db_errors async def create_bulk_session( vendor_uuid: UUID1 | UUID4, identity_uuid: UUID4, subaccount_id: int | None = None, session: AsyncSession | None = None, ): """Create a bulk_session object.""" assert session bulk_session_id = uuid.uuid4() bulk_session = BulkSession( bulk_session_id=str(bulk_session_id), created_by=str(identity_uuid), updated_by=str(identity_uuid), vendor_uuid=str(vendor_uuid), subaccount_id=subaccount_id, metadata_status=MetadataStatus.uploading, asset_status=AssetStatus.default, ) session.add(bulk_session) await session.flush() return await _get(bulk_session_id, session=session) @db.db_session_wrap async def update_bulk_session( bulk_session_id: UUID4, identity_uuid: UUID4, update_req: UpdateBulkSessionRequest, session: AsyncSession | None = None, ): """Update a bulk_session object.""" from product_staging.models.bulk_session_metadata_file import ( create_bulk_session_failure_file, update_bulk_session_metadata_file, ) assert session update_dict: Dict[str, Any] = {} if update_req.metadata_status is not None: update_dict["metadata_status"] = update_req.metadata_status.value if update_req.asset_status is not None: update_dict["asset_status"] = update_req.asset_status.value if update_req.failure_file is not None: await create_bulk_session_failure_file( bulk_session_id=bulk_session_id, s3_filename=update_req.failure_file, identity_uuid=identity_uuid, json_file_name=update_req.metadata_error_report_json, ) if update_req.metadata_error_report_json is not None: update_dict["metadata_error_report_json"] = ( update_req.metadata_error_report_json ) if update_req.success_file is not None: await update_bulk_session_metadata_file( bulk_session_id=bulk_session_id, s3_filename=update_req.success_file, status="success", ) if update_req.is_cancelled is not None: update_dict["is_cancelled"] = update_req.is_cancelled if update_req.asset_error_report is not None: update_dict["asset_error_report"] = update_req.asset_error_report if update_req.asset_report_json is not None: update_dict["asset_report_json"] = update_req.asset_report_json update_dict["updated_by"] = str(identity_uuid) stmt = ( update(BulkSession) .where(BulkSession.bulk_session_id == str(bulk_session_id)) .values(**update_dict) ) await session.execute(stmt) return await _get(bulk_session_id, session=session, is_cancelled=None) @db.db_session_wrap async def get_created_bulk_sessions( identity_uuid: UUID4, session=None, is_cancelled: bool | None = False, limit: int = 10, offset: int = 0, ): """Retrieve BulkSessions created by given identity id with computed asset_status and failed_product_uploads. The asset_status is computed in the SQL query by checking: - If there are NO required assets without a valid asset file (file_status='success'), and the session's current asset_status is 'incomplete', then computed_asset_status = 'complete' - Otherwise, computed_asset_status follows the current asset_status The failed_product_uploads is also computed in the same query by counting distinct products with at least one failed asset (file_status != 'success' or NULL). is_cancelled parameter is used to filter out "cancelled" sessions by default: to reset the filter, pass is_cancelled=None; to get only cancelled sessions pass is_cancelled=True. """ from product_staging.models.bulk_session_asset import ( BulkSessionAsset, ) from product_staging.models.bulk_session_asset_file import BulkSessionAssetFile from product_staging.models.bulk_session_metadata_file import ( BulkSessionMetadataFile, ) assert session latest_upload_file_id_subq = ( select(BulkSessionMetadataFile.bulk_session_metadata_file_id) .where( BulkSessionMetadataFile.bulk_session_id == BulkSession.bulk_session_id, BulkSessionMetadataFile.upload_token.isnot(None), ) .order_by(BulkSessionMetadataFile.created_on.desc()) .limit(1) .scalar_subquery() ) total_products_subq = ( select(BulkSessionMetadataFile.total_products) .where( BulkSessionMetadataFile.bulk_session_id == BulkSession.bulk_session_id, BulkSessionMetadataFile.upload_token.isnot(None), ) .order_by(BulkSessionMetadataFile.created_on.desc()) .limit(1) .scalar_subquery() ) success_file_id_subq = ( select(BulkSessionMetadataFile.bulk_session_metadata_file_id) .where( BulkSessionMetadataFile.bulk_session_id == BulkSession.bulk_session_id, BulkSessionMetadataFile.file_status == "success", ) .order_by(BulkSessionMetadataFile.created_on.desc()) .limit(1) .scalar_subquery() ) failure_file_id_subq = ( select(BulkSessionMetadataFile.bulk_session_metadata_file_id) .where( BulkSessionMetadataFile.bulk_session_id == BulkSession.bulk_session_id, BulkSessionMetadataFile.file_status == "failure", ) .order_by(BulkSessionMetadataFile.created_on.desc()) .limit(1) .scalar_subquery() ) incomplete_assets = ( select(literal(1)) .select_from(BulkSessionAsset) .outerjoin( BulkSessionAssetFile, BulkSessionAsset.bulk_session_asset_file_id == BulkSessionAssetFile.bulk_session_asset_file_id, ) .where( BulkSessionAsset.bulk_session_id == BulkSession.bulk_session_id, BulkSessionAsset.required.is_(True), or_( BulkSessionAsset.bulk_session_asset_file_id.is_(None), BulkSessionAssetFile.file_status != "success", ), ) .exists() ) computed_asset_status = case( ( and_( BulkSession.asset_status == "incomplete", ~incomplete_assets, ), "complete", ), else_=BulkSession.asset_status, ).label("computed_asset_status") computed_assets_complete = case( ( ~incomplete_assets, True, ), else_=False, ).label("assets_complete") user_sessions_select = select(BulkSession.bulk_session_id).where( BulkSession.created_by == str(identity_uuid) ) if is_cancelled is not None: user_sessions_select = user_sessions_select.where( BulkSession.is_cancelled == is_cancelled ) user_sessions_subq = user_sessions_select.subquery() assets_subq = ( select( BulkSessionAsset.bulk_session_id.label("bulk_session_id"), BulkSessionAsset.product_code.label("product_code"), BulkSessionAsset.bulk_session_asset_file_id.label("file_id"), ) .where( BulkSessionAsset.bulk_session_id.in_( select(user_sessions_subq.c.bulk_session_id) ), BulkSessionAsset.required.is_(True), ) .alias("assets") ) failed_products_subq = ( select( literal_column("assets.bulk_session_id").label("session_id"), literal_column("assets.product_code").label("product_code"), func.sum( case( (BulkSessionAssetFile.file_status != "success", 1), (BulkSessionAssetFile.file_status.is_(None), 1), else_=0, ) ).label("failed_assets"), ) .select_from(assets_subq) .outerjoin( BulkSessionAssetFile, literal_column("assets.file_id") == BulkSessionAssetFile.bulk_session_asset_file_id, ) .group_by( literal_column("assets.bulk_session_id"), literal_column("assets.product_code"), ) .subquery() ) failed_products_count_subq = ( select( failed_products_subq.c.session_id, func.count(func.distinct(failed_products_subq.c.product_code)).label( "failed_count" ), ) .filter(failed_products_subq.c.failed_assets > 0) .group_by(failed_products_subq.c.session_id) .subquery() ) query = ( select( BulkSession, computed_asset_status, func.coalesce(failed_products_count_subq.c.failed_count, 0).label( "failed_product_uploads" ), computed_assets_complete, latest_upload_file_id_subq.label("latest_upload_file_id"), total_products_subq.label("total_products"), success_file_id_subq.label("success_file_id"), failure_file_id_subq.label("failure_file_id"), ) .outerjoin( failed_products_count_subq, BulkSession.bulk_session_id == failed_products_count_subq.c.session_id, ) .where(BulkSession.created_by == str(identity_uuid)) ) if is_cancelled is not None: query = query.where(BulkSession.is_cancelled == is_cancelled) query = query.order_by(BulkSession.created_on.desc()) if offset: query = query.offset(offset) if limit: query = query.limit(limit) result = (await session.execute(query)).all() if not result: return [] sessions = [] for ( row, computed_status, failed_count, assets_complete, latest_upload_file_id, total_products, success_file_id, failure_file_id, ) in result: session_dict = row.to_dict() session_dict["computed_asset_status"] = computed_status session_dict["failed_product_uploads"] = ( failed_count if row.asset_status in ("complete", "incomplete") else None ) session_dict["assets_complete"] = assets_complete session_dict["latest_upload_file_id"] = latest_upload_file_id session_dict["total_products"] = total_products session_dict["success_file_id"] = success_file_id session_dict["failure_file_id"] = failure_file_id sessions.append(session_dict) return sessions @db.db_session_wrap async def has_created_bulk_sessions(identity_uuid: UUID4, session=None): """Check if the identity has created at least one bulk session.""" assert session query = ( select(BulkSession).where(BulkSession.created_by == str(identity_uuid)).limit(1) ) return bool((await session.execute(query)).scalar_one_or_none())