refactor: migrate to keyword-based repository config, add nested transactions and tests

This commit is contained in:
2026-08-14 19:58:32 +03:00
parent 3d85c2b5dd
commit 7abb513b30
22 changed files with 658 additions and 324 deletions

View File

@@ -1,26 +1,46 @@
from collections.abc import AsyncGenerator, Sequence
from contextlib import asynccontextmanager
from typing import Any
from pydantic import BaseModel
from pydantic_filters import BaseFilter, BasePagination, BaseSort
from pydantic_filters.drivers.sqlalchemy import append_to_statement
from pydantic_filters.filter._fields import FilterFieldInfo
from sqlalchemy import func
from sqlalchemy.exc import IntegrityError
from sqlmodel import delete, insert, select, update
from sqlmodel.ext.asyncio.session import AsyncSession
from .container import RepositoriesContainer
from .exceptions import AlreadyExistsError, DatabaseException, HaveNoSessionError
from .settings import DatabaseSettings
from .exceptions import AlreadyExistsError, DatabaseException
from .settings import RepositorySettings
from .tables import BaseTable
class BaseRepository:
def __init_subclass__(
cls,
table: type[BaseTable] | None = None,
filter_: type[BaseFilter] | None = None,
dto: type[BaseModel] | None = None,
**kwargs,
):
super().__init_subclass__(**kwargs)
if (table := table or getattr(cls, "_table_type", None)) is None:
raise TypeError(
f"{cls.__name__} must specify 'table' keyword argument",
)
if (filter_ := filter_ or getattr(cls, "_filter_type", None)) is None:
raise TypeError(
f"{cls.__name__} must specify 'filter_' keyword argument",
)
cls._table_type = table
cls._filter_type = filter_
cls._dto_type = dto or getattr(cls, "_dto_type", None)
def __init__(
self,
settings: DatabaseSettings | None = None,
settings: RepositorySettings | None = None,
container: RepositoriesContainer | None = None,
):
if container is not None:
@@ -30,47 +50,11 @@ class BaseRepository:
else:
raise TypeError("Either 'container' or 'settings' must be provided")
def get_db_table(self) -> type[BaseTable]:
raise NotImplementedError
def get_dto_type(self) -> type[BaseModel] | None:
return None
def get_filter_type(self) -> type[BaseFilter] | None:
return None
@property
def session(self) -> AsyncSession:
session = self._container.session
if session is None:
raise HaveNoSessionError()
return session
@asynccontextmanager
async def transaction(self) -> AsyncGenerator[None, None]:
existing_session = self._container.session
if existing_session is not None:
yield
return
async with self._container.transaction():
yield
async def create_tables(self) -> None:
table = self.get_db_table()
async with self._container.engine.begin() as connection:
await connection.run_sync(
table.metadata.create_all,
tables=[table.__table__],
)
async def get_items_count(
self,
filter_: BaseFilter | None = None,
) -> int:
if filter_ is not None:
self._ensure_filter_fields(type(filter_))
table = self.get_db_table()
table = self.get_table_type()
statement = select(func.count()).select_from(table)
statement = append_to_statement(
statement=statement,
@@ -90,10 +74,8 @@ class BaseRepository:
pagination: BasePagination | None = None,
sort: BaseSort | None = None,
options: Sequence[Any] | None = None,
) -> AsyncGenerator[BaseModel]:
if filter_ is not None:
self._ensure_filter_fields(type(filter_))
table = self.get_db_table()
) -> AsyncGenerator[Any]:
table = self.get_table_type()
statement = select(table)
statement = append_to_statement(
statement=statement,
@@ -106,13 +88,13 @@ class BaseRepository:
statement = statement.options(*options)
async with self.transaction():
result = await self.session.execute(statement)
result = await self.session.exec(statement)
result = result.yield_per(100)
for db_item in result.scalars():
for db_item in result:
yield self._convert_from_table(db_item)
async def create_item(self, item: BaseModel) -> BaseModel:
table = self.get_db_table()
table = self.get_table_type()
values = self._convert_to_table(item).to_values()
statement = insert(table).values(values).returning(table)
@@ -138,9 +120,7 @@ class BaseRepository:
options: Sequence[Any] | None = None,
**values,
) -> AsyncGenerator[BaseModel]:
if filter_ is not None:
self._ensure_filter_fields(type(filter_))
table = self.get_db_table()
table = self.get_table_type()
statement = update(table)
statement = append_to_statement(
statement=statement,
@@ -161,9 +141,7 @@ class BaseRepository:
self,
filter_: BaseFilter | None = None,
) -> None:
if filter_ is not None:
self._ensure_filter_fields(type(filter_))
table = self.get_db_table()
table = self.get_table_type()
statement = delete(table)
statement = append_to_statement(
statement=statement,
@@ -174,41 +152,45 @@ class BaseRepository:
async with self.transaction():
await self.session.exec(statement)
@property
def session(self) -> AsyncSession | None:
return self._container.session
@property
def transaction(self):
return self._container.transaction
@property
def nested_transaction(self):
return self._container.nested_transaction
async def create_tables(self) -> None:
table = self.get_table_type()
async with self._container.engine.begin() as connection:
await connection.run_sync(
table.metadata.create_all,
tables=[table.__table__],
)
def _convert_to_table(self, item: BaseModel) -> BaseModel:
table = self.get_db_table()
table = self.get_table_type()
if isinstance(item, table):
return item
return table.from_item(item=item)
def _convert_from_table(self, table: BaseTable) -> BaseModel:
dto_type = self.get_dto_type()
if dto_type is None:
if self.get_dto_type() is None:
return table
return table.to_item()
def _ensure_filter_fields(self, filter_class: type[BaseFilter]) -> None:
"""Workaround for pydantic-filters not registering filter_fields with pydantic v2."""
if getattr(filter_class, "filter_fields", None):
return
@classmethod
def get_table_type(cls) -> type[BaseTable] | None:
return cls._table_type
filter_fields: dict[str, FilterFieldInfo] = {}
for field_name, field_info in filter_class.model_fields.items():
if field_info.annotation is None:
continue
annotation = field_info.annotation
# unwrap Optional[X] -> X
origin = getattr(annotation, "__origin__", None)
if origin is type | None:
args = getattr(annotation, "__args__", ())
if args and args[0] is not type(None):
annotation = args[0]
@classmethod
def get_filter_type(cls) -> type[BaseFilter] | None:
return cls._filter_type
is_sequence = hasattr(
annotation, "__origin__"
) and annotation.__origin__ in (list, set)
filter_fields[field_name] = FilterFieldInfo(
target=field_name,
type_="eq",
is_sequence=is_sequence,
)
filter_class.filter_fields = filter_fields
@classmethod
def get_dto_type(cls) -> type[BaseModel] | None:
return cls._dto_type