"""Tracking repositories — Phase 10.7.""" from __future__ import annotations from typing import Sequence from uuid import UUID from sqlalchemy import func, select from app.models.tracking import ( CustomerTrackingToken, ProofOfDelivery, TrackingPoint, TrackingSession, ) from app.repositories.base import TenantBaseRepository from app.specifications.tracking import ( CustomerTrackingTokenListSpec, ProofOfDeliveryListSpec, TrackingSessionListSpec, ) class TrackingSessionRepository(TenantBaseRepository[TrackingSession]): model = TrackingSession async def list_filtered( self, tenant_id: UUID, *, offset: int = 0, limit: int = 20, spec: TrackingSessionListSpec | None = None, ) -> Sequence[TrackingSession]: 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: TrackingSessionListSpec | 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 TrackingPointRepository(TenantBaseRepository[TrackingPoint]): model = TrackingPoint async def list_for_session( self, tenant_id: UUID, tracking_session_id: UUID, *, limit: int = 500 ) -> Sequence[TrackingPoint]: stmt = ( select(self.model) .where( self.model.tenant_id == tenant_id, self.model.tracking_session_id == tracking_session_id, ) .order_by(self.model.recorded_at.asc()) .limit(limit) ) result = await self.session.execute(stmt) return result.scalars().all() class ProofOfDeliveryRepository(TenantBaseRepository[ProofOfDelivery]): model = ProofOfDelivery async def list_filtered( self, tenant_id: UUID, *, offset: int = 0, limit: int = 20, spec: ProofOfDeliveryListSpec | None = None, ) -> Sequence[ProofOfDelivery]: 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: ProofOfDeliveryListSpec | 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 CustomerTrackingTokenRepository(TenantBaseRepository[CustomerTrackingToken]): model = CustomerTrackingToken async def get_by_token( self, tenant_id: UUID, token: str ) -> CustomerTrackingToken | None: stmt = select(self.model).where( self.model.tenant_id == tenant_id, self.model.token == token, 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: CustomerTrackingTokenListSpec | None = None, ) -> Sequence[CustomerTrackingToken]: 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: CustomerTrackingTokenListSpec | 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())