from __future__ import annotations from typing import Any, Generic, Sequence, Type, TypeVar, cast from sqlalchemy import Result, and_, delete, exists, insert, select, update from sqlalchemy.orm import MapperProperty, class_mapper from sqlalchemy.sql.base import Executable, elements from src.core.models.base import Base as BaseModel from src.db.connector import db_session __all__ = ["BaseService"] T = TypeVar("T", bound=BaseModel) class BaseService(Generic[T]): model: Type[T] def _get_pk_props(self) -> list[MapperProperty]: mapper = class_mapper(self.model) props = [] for key in mapper.primary_key: props.append(mapper.get_property_by_column(key)) return props def _build_pk_criteria(self, pk_value: Any | dict[str, Any]) -> elements.ColumnElement: pk_pops = self._get_pk_props() criteria = [] if len(pk_pops) == 1: prop = pk_pops[0] if isinstance(pk_value, dict): if prop.key not in pk_value: raise KeyError(f"No value for '{prop}'") value = pk_value[prop.key] else: value = pk_value criteria.append(prop.class_attribute == value) else: if not isinstance(pk_value, dict): raise ValueError("'pk_value' must be dict in case of composite primary key.") for prop in pk_pops: if prop.key not in pk_value: raise KeyError(f"No value for '{prop}'") criteria.append(prop.class_attribute == pk_value[prop.key]) return and_(*criteria) def _execute(self, stmt: Executable) -> Result: with db_session() as session: return session.execute(stmt) def _select_many(self, stmt: Executable) -> Sequence[T]: with db_session() as session: return session.execute(stmt).scalars().all() def _select_one(self, stmt: Executable) -> T | None: with db_session() as session: return session.execute(stmt).scalars().one_or_none() def select_all(self) -> Sequence[T]: with db_session() as session: return session.execute(select(self.model)).scalars().all() def select_by_pk(self, pk_value: Any) -> T | None: stmt = select(self.model).where(self._build_pk_criteria(pk_value)) return self._select_one(stmt) def exists_by_pk(self, pk_value: Any) -> bool: stmt = select(exists(self.model)).where(self._build_pk_criteria(pk_value)) return cast(bool, self._execute(stmt).scalar()) def delete_by_pk(self, pk_value: Any): stmt = delete(self.model).where(self._build_pk_criteria(pk_value)) self._execute(stmt) def delete_all(self): self._execute(delete(self.model)) def update_by_pk(self, pk_value: Any, **values): stmt = update(self.model).where(self._build_pk_criteria(pk_value)).values(**values) self._execute(stmt) def _insert_instances(self, instances: Sequence[T]): with db_session() as session: session.bulk_save_objects(instances) def insert_instances(self, instances: Sequence[T], batch_size: int = 100): if not batch_size: self._insert_instances(instances) return for i in range(0, len(instances), batch_size): self._insert_instances(instances[i : i + batch_size]) def insert_instance(self, instance: T): with db_session() as session: session.add(instance) def delete_instance(self, instance: T): with db_session() as session: session.delete(instance) def insert(self, **values): stmt = insert(self.model).values(**values) self._execute(stmt)