"""Tracking query specifications — Phase 10.7.""" from __future__ import annotations from dataclasses import dataclass from uuid import UUID from sqlalchemy import Select, asc, desc from sqlalchemy.sql import ColumnElement from app.models.tracking import CustomerTrackingToken, ProofOfDelivery, TrackingSession from app.models.types import ( CustomerTrackingTokenStatus, ProofOfDeliveryStatus, TrackingSessionStatus, ) SESSION_SORT = frozenset({"created_at", "updated_at", "status"}) POD_SORT = frozenset({"created_at", "updated_at", "status"}) TOKEN_SORT = frozenset({"created_at", "updated_at", "status"}) @dataclass(frozen=True, slots=True) class TrackingSessionListSpec: status: TrackingSessionStatus | None = None dispatch_job_id: UUID | None = None driver_id: UUID | None = None sort_by: str = "created_at" sort_dir: str = "desc" def filter_clauses(self) -> list[ColumnElement]: clauses: list[ColumnElement] = [] if self.status is not None: clauses.append(TrackingSession.status == self.status) if self.dispatch_job_id is not None: clauses.append(TrackingSession.dispatch_job_id == self.dispatch_job_id) if self.driver_id is not None: clauses.append(TrackingSession.driver_id == self.driver_id) return clauses def apply(self, stmt: Select) -> Select: clauses = self.filter_clauses() if clauses: stmt = stmt.where(*clauses) sort_key = self.sort_by if self.sort_by in SESSION_SORT else "created_at" column = getattr(TrackingSession, sort_key) order = desc(column) if self.sort_dir.lower() != "asc" else asc(column) return stmt.order_by(order) @dataclass(frozen=True, slots=True) class ProofOfDeliveryListSpec: status: ProofOfDeliveryStatus | None = None dispatch_job_id: UUID | None = None sort_by: str = "created_at" sort_dir: str = "desc" def filter_clauses(self) -> list[ColumnElement]: clauses: list[ColumnElement] = [] if self.status is not None: clauses.append(ProofOfDelivery.status == self.status) if self.dispatch_job_id is not None: clauses.append(ProofOfDelivery.dispatch_job_id == self.dispatch_job_id) return clauses def apply(self, stmt: Select) -> Select: clauses = self.filter_clauses() if clauses: stmt = stmt.where(*clauses) sort_key = self.sort_by if self.sort_by in POD_SORT else "created_at" column = getattr(ProofOfDelivery, sort_key) order = desc(column) if self.sort_dir.lower() != "asc" else asc(column) return stmt.order_by(order) @dataclass(frozen=True, slots=True) class CustomerTrackingTokenListSpec: status: CustomerTrackingTokenStatus | None = None dispatch_job_id: UUID | None = None sort_by: str = "created_at" sort_dir: str = "desc" def filter_clauses(self) -> list[ColumnElement]: clauses: list[ColumnElement] = [] if self.status is not None: clauses.append(CustomerTrackingToken.status == self.status) if self.dispatch_job_id is not None: clauses.append( CustomerTrackingToken.dispatch_job_id == self.dispatch_job_id ) return clauses def apply(self, stmt: Select) -> Select: clauses = self.filter_clauses() if clauses: stmt = stmt.where(*clauses) sort_key = self.sort_by if self.sort_by in TOKEN_SORT else "created_at" column = getattr(CustomerTrackingToken, sort_key) order = desc(column) if self.sort_dir.lower() != "asc" else asc(column) return stmt.order_by(order)