"""Localization & media-ref repositories — Phase 11.5.""" from __future__ import annotations from typing import Sequence from uuid import UUID from sqlalchemy import func, select, update from app.models.localization import ( ExperienceLocalization, ExperienceMediaRef, SiteLocaleBinding, ) from app.models.types import LocalizationTarget from app.repositories.base import TenantBaseRepository from app.specifications.localization import ( LocalizationListSpec, MediaRefListSpec, SiteLocaleBindingListSpec, ) class SiteLocaleBindingRepository(TenantBaseRepository[SiteLocaleBinding]): model = SiteLocaleBinding async def get_by_site_profile( self, tenant_id: UUID, site_id: UUID, locale_profile_id: UUID ) -> SiteLocaleBinding | None: stmt = select(self.model).where( self.model.tenant_id == tenant_id, self.model.site_id == site_id, self.model.locale_profile_id == locale_profile_id, self.model.is_deleted.is_(False), ) result = await self.session.execute(stmt) return result.scalar_one_or_none() async def clear_default_for_site( self, tenant_id: UUID, site_id: UUID, *, except_id: UUID | None = None ) -> None: stmt = ( update(self.model) .where( self.model.tenant_id == tenant_id, self.model.site_id == site_id, self.model.is_deleted.is_(False), self.model.is_default.is_(True), ) .values(is_default=False) ) if except_id is not None: stmt = stmt.where(self.model.id != except_id) await self.session.execute(stmt) async def list_filtered( self, tenant_id: UUID, *, offset: int, limit: int, spec: SiteLocaleBindingListSpec | None = None, ) -> tuple[Sequence[SiteLocaleBinding], int]: base = select(self.model).where( self.model.tenant_id == tenant_id, self.model.is_deleted.is_(False), ) count_base = select(func.count()).select_from(self.model).where( self.model.tenant_id == tenant_id, self.model.is_deleted.is_(False), ) if spec is not None: base = spec.apply(base) for clause in spec.filter_clauses(): count_base = count_base.where(clause) else: base = base.order_by(self.model.sort_order.asc()) total = int((await self.session.execute(count_base)).scalar_one()) result = await self.session.execute(base.offset(offset).limit(limit)) return result.scalars().all(), total class ExperienceLocalizationRepository( TenantBaseRepository[ExperienceLocalization] ): model = ExperienceLocalization async def get_by_target_locale( self, tenant_id: UUID, target_type: LocalizationTarget, target_id: UUID, locale: str, ) -> ExperienceLocalization | None: stmt = select(self.model).where( self.model.tenant_id == tenant_id, self.model.target_type == target_type, self.model.target_id == target_id, self.model.locale == locale, self.model.is_deleted.is_(False), ) result = await self.session.execute(stmt) return result.scalar_one_or_none() async def list_filtered( self, tenant_id: UUID, *, offset: int, limit: int, spec: LocalizationListSpec | None = None, ) -> tuple[Sequence[ExperienceLocalization], int]: base = select(self.model).where( self.model.tenant_id == tenant_id, self.model.is_deleted.is_(False), ) count_base = select(func.count()).select_from(self.model).where( self.model.tenant_id == tenant_id, self.model.is_deleted.is_(False), ) if spec is not None: base = spec.apply(base) for clause in spec.filter_clauses(): count_base = count_base.where(clause) else: base = base.order_by(self.model.created_at.desc()) total = int((await self.session.execute(count_base)).scalar_one()) result = await self.session.execute(base.offset(offset).limit(limit)) return result.scalars().all(), total class ExperienceMediaRefRepository(TenantBaseRepository[ExperienceMediaRef]): model = ExperienceMediaRef async def list_filtered( self, tenant_id: UUID, *, offset: int, limit: int, spec: MediaRefListSpec | None = None, ) -> tuple[Sequence[ExperienceMediaRef], int]: base = select(self.model).where( self.model.tenant_id == tenant_id, self.model.is_deleted.is_(False), ) count_base = select(func.count()).select_from(self.model).where( self.model.tenant_id == tenant_id, self.model.is_deleted.is_(False), ) if spec is not None: base = spec.apply(base) for clause in spec.filter_clauses(): count_base = count_base.where(clause) else: base = base.order_by(self.model.sort_order.asc()) total = int((await self.session.execute(count_base)).scalar_one()) result = await self.session.execute(base.offset(offset).limit(limit)) return result.scalars().all(), total