from collections.abc import Iterable, Sequence from typing import Any, ClassVar, get_args, get_origin import sqlalchemy as sa from sqlalchemy.orm import DeclarativeBase from sqlalchemy.orm.interfaces import ORMOption from sqlalchemy.sql.selectable import ForUpdateParameter from fansifter_common.adapters.db.base import Database from .utils import delete_instance, refresh_instance, save_instance class BaseRepository[BaseDB_T: Database, Base_T: DeclarativeBase]: model: type[Base_T] default_options: ClassVar[Sequence[ORMOption]] = [] __provided__: ClassVar[dict[str, Any]] def __init__(self, db: BaseDB_T) -> None: self.db = db def __init_subclass__(cls, **kwargs: Any) -> None: super().__init_subclass__(**kwargs) for base in getattr(cls, "__orig_bases__", ()): origin = get_origin(base) if origin is not None: args = get_args(base) if ( args and isinstance(origin, type) and issubclass(origin, BaseRepository) ): cls.model = args[0] cls.__provided__ = {"scope": "singleton", "alias": base} return # Instance methods def add(self, instance: Any) -> None: self.db.session.add(instance) def add_all(self, instances: Any) -> None: self.db.session.add_all(instances) def save(self, instance: Base_T, *, flush: bool = False) -> Base_T: """Save instance.""" return save_instance(self.db, instance=instance, flush=flush) def delete(self, instance: Base_T, *, flush: bool = False) -> None: """Delete an object from the database.""" delete_instance(self.db, instance=instance, flush=flush) def refresh( self, instance: Base_T, *, attribute_names: Iterable[str] | None = None, with_for_update: ForUpdateParameter = None, ) -> None: """Refresh the instance from the database.""" refresh_instance( self.db, instance=instance, attribute_names=attribute_names, with_for_update=with_for_update, ) # Query methods def all(self) -> Sequence[Base_T]: query = sa.select(self.model).options(*self.default_options) result = self.db.session.execute(query) return result.scalars().all() def get(self, ident: Any) -> Base_T | None: return self.db.session.get( self.model, ident=ident, options=self.default_options ) def filter_by(self, **kwargs: Any) -> Sequence[Base_T]: query = sa.select(self.model).filter_by(**kwargs).options(*self.default_options) result = self.db.session.execute(query) return result.scalars().all() def first(self) -> Base_T | None: query = sa.select(self.model).options(*self.default_options) result = self.db.session.execute(query) return result.scalars().first() def count(self) -> int: query = sa.select(sa.func.count()).select_from(sa.select(self.model).subquery()) result = self.db.session.execute(query) return result.scalar_one()