Files
AutoShop-Djimbo-Simple/tgbot/database/repository.py
T

174 lines
6.4 KiB
Python

# - *- coding: utf- 8 - *-
from dataclasses import fields
from typing import Any, Dict, Generic, List, Optional, Type, TypeVar
from sqlalchemy import delete as sqlalchemy_delete
from sqlalchemy import select
from sqlalchemy import update as sqlalchemy_update
from tgbot.database.core import Base, database_path, engine, session_scope
from tgbot.utils.misc.bot_logging import bot_logger
ModelTranslator = TypeVar("ModelTranslator", bound=Base)
EntityTranslator = TypeVar("EntityTranslator")
# Базовый репозиторий для SQLAlchemy-моделей
class BaseRepository(Generic[ModelTranslator, EntityTranslator]):
# Настройка модели репозитория
def __init__(self):
self.storage_name = "storage"
self.table_model: Optional[Type[ModelTranslator]] = None
self.entity_model: Optional[Type[EntityTranslator]] = None
# Получение SQLAlchemy-модели репозитория
def _model(self) -> Type[ModelTranslator]:
if self.table_model is None:
raise RuntimeError("Модель базы данных не настроена")
return self.table_model
# Преобразование ORM-строки в безопасный DTO
def _to_entity(self, instance: ModelTranslator) -> EntityTranslator:
if self.entity_model is None:
raise RuntimeError("DTO-модель базы данных не настроена")
values = {
field.name: getattr(instance, field.name)
for field in fields(self.entity_model)
}
return self.entity_model(**values)
# Добавление записи через текущую модель
async def _insert(self, **kwargs) -> EntityTranslator:
model = self._model()
instance = model(**kwargs)
async with session_scope() as session:
session.add(instance)
await session.flush()
await session.refresh(instance)
return self._to_entity(instance)
# Добавление записи с возвратом DTO
async def add(self, **kwargs) -> EntityTranslator:
return await self._insert(**kwargs)
# Удаление записи по явному фильтру
async def delete(self, **kwargs) -> None:
if not kwargs:
raise ValueError("Для удаления нужен хотя бы один фильтр")
model = self._model()
async with session_scope() as session:
await session.execute(sqlalchemy_delete(model).filter_by(**kwargs))
# Очистка текущей таблицы
async def clear(self) -> None:
await self.delete_all_rows()
# Удаление всех строк текущей таблицы
async def delete_all_rows(self) -> None:
model = self._model()
async with session_scope() as session:
await session.execute(sqlalchemy_delete(model))
# Получение первой записи по фильтру
async def get(self, **kwargs) -> Optional[EntityTranslator]:
model = self._model()
statement = select(model)
if kwargs:
statement = statement.filter_by(**kwargs)
if hasattr(model, "increment"):
statement = statement.order_by(getattr(model, "increment"))
async with session_scope() as session:
response = await session.execute(statement)
instance = response.scalars().first()
if instance is None:
return None
return self._to_entity(instance)
# Получение записи или явно падает
async def get_required(self, **kwargs) -> EntityTranslator:
entity = await self.get(**kwargs)
if entity is None:
raise LookupError(f"Запись не найдена в {self.storage_name}: {kwargs}")
return entity
# Получение списка записей по фильтру
async def gets(self, **kwargs) -> List[EntityTranslator]:
model = self._model()
statement = select(model)
if kwargs:
statement = statement.filter_by(**kwargs)
if hasattr(model, "increment"):
statement = statement.order_by(getattr(model, "increment"))
async with session_scope() as session:
response = await session.execute(statement)
return [self._to_entity(instance) for instance in response.scalars().all()]
# Получение всех записей таблицы
async def get_all(self) -> List[EntityTranslator]:
return await self.gets()
# Обновление записи через текущую модель
async def _update(self, where: Optional[Dict[str, Any]] = None, **kwargs) -> int:
if not kwargs:
return 0
model = self._model()
statement = sqlalchemy_update(model).values(**kwargs)
if where:
statement = statement.filter_by(**where)
async with session_scope() as session:
result = await session.execute(statement)
rowcount = getattr(result, "rowcount", 0)
return int(rowcount or 0)
# Обновление записи по фильтру
async def update(self, where: Optional[Dict[str, Any]] = None, **kwargs) -> int:
return await self._update(where=where, **kwargs)
# Применение миграций и создание дефолтных строк
async def prepare_database() -> None:
database_path.parent.mkdir(parents=True, exist_ok=True)
# Импорт всех моделей до create_all, чтобы они зарегистрировались в Base.metadata
import tgbot.database.db_category # noqa: F401
import tgbot.database.db_item # noqa: F401
import tgbot.database.db_payments # noqa: F401
import tgbot.database.db_position # noqa: F401
import tgbot.database.db_purchases # noqa: F401
import tgbot.database.db_refill # noqa: F401
import tgbot.database.db_settings # noqa: F401
import tgbot.database.db_users # noqa: F401
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
from tgbot.database.db_payments import Paymentsx
from tgbot.database.db_settings import Settingsx
await Settingsx().ensure_default()
await Paymentsx().ensure_default()
bot_logger.info("База данных готова")