from __future__ import annotations from collections.abc import Sequence from datetime import datetime from typing import Any, ClassVar, Protocol, Self, cast import sqlalchemy as sa from sqlalchemy import ColumnExpressionArgument from sqlalchemy.exc import NoResultFound from sqlalchemy.orm import InstrumentedAttribute, Session, joinedload from sqlalchemy.sql.selectable import ForUpdateParameter from fansifter_common.adapters.db.base import Database from fansifter_common.adapters.db.exceptions import ObjectNotFoundError class HasDBProtocol(Protocol): __db__: ClassVar[Database] class Query[T]: def __init__(self, model: type[T], session: Session) -> None: self.model = model self.session = session self._query = sa.select(model) self._options: list[Any] = [] # --- Descriptor --- @classmethod def as_descriptor(cls) -> Self: return cast(Self, QueryDescriptor(query_cls=cls)) # -- Builder --- def where(self, *whereclause: ColumnExpressionArgument[bool]) -> Self: self._query = self._query.where(*whereclause) return self def join( self, target: Any, onclause: Any | None = None, *, isouter: bool = False, full: bool = False, ) -> Self: """Add an INNER JOIN (or OUTER JOIN if `isouter=True`) to the query.""" self._query = self._query.join( target, onclause=onclause, isouter=isouter, full=full ) return self def limit(self, limit: int) -> Self: """Limit the number of rows returned.""" self._query = self._query.limit(limit) return self def offset(self, offset: int) -> Self: """Skip a number of rows before returning results.""" self._query = self._query.offset(offset) return self def order_by(self, *order_by: Any) -> Self: """Order the query results by one or more columns.""" self._query = self._query.order_by(*order_by) return self def joinedload(self, key: Any, *subkeys: Any) -> Self: """Apply SQLAlchemy joinedload() option for eager loading a relationship. Additional keys are chained as nested joinedloads. E.g. joinedload(Campaign.domain, Domain.esp) produces joinedload(Campaign.domain).joinedload(Domain.esp). """ option = joinedload(key) for subkey in subkeys: option = option.joinedload(subkey) self._options.append(option) self._query = self._query.options(option) return self # --- Execution methods --- def get( self, ident: Any, *, with_for_update: ForUpdateParameter | None = None ) -> T | None: return self.session.get( self.model, ident, options=self._options or None, with_for_update=with_for_update, ) def get_one( self, ident: Any, *, with_for_update: ForUpdateParameter | None = None ) -> T: try: return self.session.get_one( self.model, ident, options=self._options or None, with_for_update=with_for_update, ) except NoResultFound: raise ObjectNotFoundError from None def all(self) -> Sequence[T]: """Execute the query and return all matching results.""" result = self.session.execute(self._query) return result.scalars().all() def one(self) -> T: """Execute the query and return exactly one result.""" result = self.session.execute(self._query) try: return result.scalars().one() except NoResultFound: raise ObjectNotFoundError from None def one_or_none(self) -> T | None: """Execute the query and return one or zero results.""" result = self.session.execute(self._query) return result.scalars().one_or_none() def first(self) -> T | None: """Execute the query and return the first result or None.""" query = self._query.limit(1) result = self.session.execute(query) return result.scalars().first() def latest(self, key: InstrumentedAttribute[datetime]) -> T | None: """Execute the query and return the latest result or None.""" statement = self._query.order_by(key.desc()).limit(1) result = self.session.execute(statement) return result.scalars().first() def count(self) -> int: """Count the number of rows matching the query.""" statement = sa.select(sa.func.count()).select_from(self._query.subquery()) result = self.session.execute(statement) return result.scalar_one() def exists(self) -> bool: """Check whether any rows match the current query.""" statement = sa.select(sa.exists(self._query.limit(1).subquery())) result = self.session.execute(statement) return bool(result.scalar()) class QueryDescriptor[T: Query[Any]]: def __init__(self, query_cls: type[Query[Any]]) -> None: self.query_cls = query_cls def __get__(self, instance: Any, owner: type[HasDBProtocol]) -> Any: return self.query_cls( model=cast(type[Any], owner), session=owner.__db__.session )