"""Driver repositories — Phase 10.1.""" from __future__ import annotations from typing import Sequence from uuid import UUID from sqlalchemy import func, select from app.models.drivers import ( Driver, DriverCredential, DriverDocument, DriverLifecycleEvent, ) from app.repositories.base import TenantBaseRepository from app.specifications.drivers import DriverListSpec class DriverRepository(TenantBaseRepository[Driver]): model = Driver async def get_by_code(self, tenant_id: UUID, code: str) -> Driver | 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: DriverListSpec | None = None, ) -> Sequence[Driver]: 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: DriverListSpec | 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 DriverCredentialRepository(TenantBaseRepository[DriverCredential]): model = DriverCredential async def list_for_driver( self, tenant_id: UUID, driver_id: UUID ) -> Sequence[DriverCredential]: stmt = select(self.model).where( self.model.tenant_id == tenant_id, self.model.driver_id == driver_id, self.model.is_deleted.is_(False), ) result = await self.session.execute(stmt) return result.scalars().all() class DriverDocumentRepository(TenantBaseRepository[DriverDocument]): model = DriverDocument async def list_for_driver( self, tenant_id: UUID, driver_id: UUID ) -> Sequence[DriverDocument]: stmt = select(self.model).where( self.model.tenant_id == tenant_id, self.model.driver_id == driver_id, self.model.is_deleted.is_(False), ) result = await self.session.execute(stmt) return result.scalars().all() class DriverLifecycleEventRepository(TenantBaseRepository[DriverLifecycleEvent]): model = DriverLifecycleEvent async def list_for_driver( self, tenant_id: UUID, driver_id: UUID, *, limit: int = 100 ) -> Sequence[DriverLifecycleEvent]: stmt = ( select(self.model) .where( self.model.tenant_id == tenant_id, self.model.driver_id == driver_id, ) .order_by(self.model.created_at.desc()) .limit(limit) ) result = await self.session.execute(stmt) return result.scalars().all()