"""Availability repositories — Phase 10.3.""" from __future__ import annotations from typing import Sequence from uuid import UUID from sqlalchemy import func, select from app.models.availability import ( DriverAvailability, Shift, ShiftAssignment, WorkingZone, ) from app.repositories.base import TenantBaseRepository from app.specifications.availability import ShiftListSpec, WorkingZoneListSpec class DriverAvailabilityRepository(TenantBaseRepository[DriverAvailability]): model = DriverAvailability async def get_for_driver( self, tenant_id: UUID, driver_id: UUID ) -> DriverAvailability | None: stmt = ( select(self.model) .where( self.model.tenant_id == tenant_id, self.model.driver_id == driver_id, self.model.is_deleted.is_(False), ) .order_by(self.model.updated_at.desc()) .limit(1) ) result = await self.session.execute(stmt) return result.scalar_one_or_none() class ShiftRepository(TenantBaseRepository[Shift]): model = Shift async def get_by_code(self, tenant_id: UUID, code: str) -> Shift | 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: ShiftListSpec | None = None, ) -> Sequence[Shift]: 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.starts_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: ShiftListSpec | 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 ShiftAssignmentRepository(TenantBaseRepository[ShiftAssignment]): model = ShiftAssignment async def list_for_shift( self, tenant_id: UUID, shift_id: UUID ) -> Sequence[ShiftAssignment]: stmt = select(self.model).where( self.model.tenant_id == tenant_id, self.model.shift_id == shift_id, self.model.is_deleted.is_(False), ) result = await self.session.execute(stmt) return result.scalars().all() class WorkingZoneRepository(TenantBaseRepository[WorkingZone]): model = WorkingZone async def get_by_code(self, tenant_id: UUID, code: str) -> WorkingZone | 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: WorkingZoneListSpec | None = None, ) -> Sequence[WorkingZone]: 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: WorkingZoneListSpec | 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())