"""Carveouts related tables.""" from typing import Any from sqlalchemy import BigInteger, Column, Integer, SmallInteger, select, update from sqlalchemy.orm import Session from video.connectors import mysql class ReleaseDmsMasterRestriction(mysql.ArModel): """Release DMS Master restriction.""" __tablename__ = "release_dms_master_restriction" primary_key = Column( Integer, name="restriction_id", nullable=False, primary_key=True, autoincrement=True, ) release_id = Column(Integer) upc = Column(BigInteger, nullable=False) customer_master_master_id = Column(Integer) distribution_type_id = Column(Integer) updated_by = Column(Integer) def to_dict(self) -> dict[str, Any]: """Return a dictionary of this release's properties.""" return { "release_id": self.release_id, "upc": self.upc, "customer_master_master_id": self.customer_master_master_id, "distribution_type_id": self.distribution_type_id, } class ReleaseDmsRestriction(mysql.ArModel): """Release DM Restriction.""" __tablename__ = "release_dms_restriction" primary_key = Column( Integer, name="restriction_id", nullable=False, primary_key=True, autoincrement=True, ) release_id = Column(Integer) upc = Column(BigInteger, nullable=False) dms_customer_id = Column(Integer) def to_dict(self) -> dict[str, Any]: """Return a dictionary of this release's properties.""" return { "release_id": self.release_id, "upc": self.upc, "dms_customer_id": self.dms_customer_id, } class ReleaseTerritoryRestriction(mysql.ArModel): """Release Territory Restriction.""" __tablename__ = "release_territory_restriction" primary_key = Column( Integer, name="restriction_id", nullable=False, primary_key=True, autoincrement=True, ) release_id = Column(Integer) upc = Column(BigInteger, nullable=False) country_id = Column(SmallInteger) def to_dict(self) -> dict[str, Any]: """Return a dictionary of this release's properties.""" return { "release_id": self.release_id, "upc": self.upc, "country_id": self.country_id, } def update_upc(release_id: int, upc: int, session: Session) -> None: """Update upc for carveouts table rows.""" session.execute( update(ReleaseDmsMasterRestriction) .where(ReleaseDmsMasterRestriction.release_id == release_id) .values(upc=upc) ) session.execute( update(ReleaseDmsRestriction) .where(ReleaseDmsRestriction.release_id == release_id) .values(upc=upc) ) session.execute( update(ReleaseTerritoryRestriction) .where(ReleaseTerritoryRestriction.release_id == release_id) .values(upc=upc) ) def add_release_dms_master_restrictions( distribution_type_id: int, upc: int, product_id: int, updated_by: int, customer_master_master_ids: list[int], session: Session, ) -> None: """Add carveouts for specified stores for auto approval workflow.""" release_restrictions = [] stores_to_auto_carveout = _deduplicate_dms_master_restrictions( customer_master_master_ids, product_id, session=session ) for customer_master_master_id in stores_to_auto_carveout: release_restriction_model = ReleaseDmsMasterRestriction( customer_master_master_id=customer_master_master_id, distribution_type_id=distribution_type_id, upc=upc, release_id=product_id, updated_by=updated_by, ) release_restrictions.append(release_restriction_model) session.add_all(release_restrictions) def _deduplicate_dms_master_restrictions( all_stores_to_be_carved_out: list[int], product_id: int, session: Session, ) -> list[int]: current_carveouts = _get_all_for_model( product_id, ReleaseDmsMasterRestriction, session=session ) current_carved_out_stores = [] for carveout in current_carveouts: if carveout["distribution_type_id"] == 3: current_carved_out_stores.append(carveout["customer_master_master_id"]) return list( set(all_stores_to_be_carved_out).difference(set(current_carved_out_stores)) ) def _get_all_for_model( release_id: int, model: type[Any], session: Session ) -> list[dict[str, Any]]: results = ( session.execute(select(model).where(model.release_id == release_id)) .scalars() .all() ) return [result.to_dict() for result in results]