"""Settlement repositories — Phase 10.8.""" from __future__ import annotations from typing import Sequence from uuid import UUID from sqlalchemy import func, select from app.models.settlement import SettlementIntent, SettlementLine from app.repositories.base import TenantBaseRepository from app.specifications.settlement import SettlementIntentListSpec class SettlementIntentRepository(TenantBaseRepository[SettlementIntent]): model = SettlementIntent async def get_by_code( self, tenant_id: UUID, code: str ) -> SettlementIntent | None: stmt = select(self.model).where( self.model.tenant_id == tenant_id, self.model.code == code, 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 = 0, limit: int = 20, spec: SettlementIntentListSpec | None = None, ) -> Sequence[SettlementIntent]: stmt = select(self.model).where( self.model.tenant_id == tenant_id, self.model.is_deleted.is_(False), ) if spec is not None: stmt = spec.apply(stmt) else: stmt = stmt.order_by(self.model.created_at.desc()) stmt = stmt.offset(offset).limit(limit) result = await self.session.execute(stmt) return result.scalars().all() async def count_filtered( self, tenant_id: UUID, *, spec: SettlementIntentListSpec | None = None ) -> int: stmt = 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: clauses = spec.filter_clauses() if clauses: stmt = stmt.where(*clauses) result = await self.session.execute(stmt) return int(result.scalar_one()) class SettlementLineRepository(TenantBaseRepository[SettlementLine]): model = SettlementLine async def list_for_intent( self, tenant_id: UUID, settlement_intent_id: UUID ) -> Sequence[SettlementLine]: stmt = select(self.model).where( self.model.tenant_id == tenant_id, self.model.settlement_intent_id == settlement_intent_id, self.model.is_deleted.is_(False), ) result = await self.session.execute(stmt) return result.scalars().all()