"""Tenant-aware base repository with soft-delete helpers.""" from __future__ import annotations from datetime import datetime, timezone from typing import Generic, Sequence, TypeVar from uuid import UUID from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession from app.core.database import Base ModelT = TypeVar("ModelT", bound=Base) class TenantBaseRepository(Generic[ModelT]): model: type[ModelT] def __init__(self, session: AsyncSession) -> None: self.session = session async def get(self, tenant_id: UUID, entity_id: UUID) -> ModelT | None: clauses = [ self.model.tenant_id == tenant_id, # type: ignore[attr-defined] self.model.id == entity_id, # type: ignore[attr-defined] ] if hasattr(self.model, "is_deleted"): clauses.append(self.model.is_deleted.is_(False)) # type: ignore[attr-defined] stmt = select(self.model).where(*clauses) result = await self.session.execute(stmt) return result.scalar_one_or_none() async def get_including_deleted( self, tenant_id: UUID, entity_id: UUID ) -> ModelT | None: stmt = select(self.model).where( self.model.tenant_id == tenant_id, # type: ignore[attr-defined] self.model.id == entity_id, # type: ignore[attr-defined] ) result = await self.session.execute(stmt) return result.scalar_one_or_none() async def add(self, entity: ModelT) -> ModelT: self.session.add(entity) await self.session.flush() return entity async def delete(self, entity: ModelT) -> None: await self.session.delete(entity) await self.session.flush() async def soft_delete(self, entity: ModelT, *, deleted_by: str | None = None) -> None: entity.is_deleted = True # type: ignore[attr-defined] entity.deleted_at = datetime.now(timezone.utc) # type: ignore[attr-defined] if deleted_by is not None and hasattr(entity, "deleted_by"): entity.deleted_by = deleted_by # type: ignore[attr-defined] await self.session.flush() async def restore(self, entity: ModelT) -> None: entity.is_deleted = False # type: ignore[attr-defined] entity.deleted_at = None # type: ignore[attr-defined] if hasattr(entity, "deleted_by"): entity.deleted_by = None # type: ignore[attr-defined] await self.session.flush() async def list_by_tenant( self, tenant_id: UUID, *, offset: int = 0, limit: int = 20 ) -> Sequence[ModelT]: clauses = [self.model.tenant_id == tenant_id] # type: ignore[attr-defined] if hasattr(self.model, "is_deleted"): clauses.append(self.model.is_deleted.is_(False)) # type: ignore[attr-defined] stmt = select(self.model).where(*clauses).offset(offset).limit(limit) result = await self.session.execute(stmt) return result.scalars().all() async def count_by_tenant(self, tenant_id: UUID) -> int: clauses = [self.model.tenant_id == tenant_id] # type: ignore[attr-defined] if hasattr(self.model, "is_deleted"): clauses.append(self.model.is_deleted.is_(False)) # type: ignore[attr-defined] stmt = select(func.count()).select_from(self.model).where(*clauses) result = await self.session.execute(stmt) return int(result.scalar_one())