Initial release: async repository layer over SQLModel

This commit is contained in:
2026-08-13 22:25:59 +03:00
commit 3d85c2b5dd
26 changed files with 2252 additions and 0 deletions

26
metaorm/__init__.py Normal file
View File

@@ -0,0 +1,26 @@
from .container import RepositoriesContainer
from .exceptions import (
AlreadyExistsError,
DatabaseException,
HaveNoSessionError,
NotFoundError,
)
from .repositories import BaseRepository
from .settings import DatabaseSettings
from .tables import BaseTable
__all__ = (
# container
"RepositoriesContainer",
# exceptions
"AlreadyExistsError",
"DatabaseException",
"HaveNoSessionError",
"NotFoundError",
# repositories
"BaseRepository",
# settings
"DatabaseSettings",
# tables
"BaseTable",
)

58
metaorm/container.py Normal file
View File

@@ -0,0 +1,58 @@
import contextvars
from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager
from typing import TypeVar
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlmodel.ext.asyncio.session import AsyncSession
from .settings import DatabaseSettings
RepositoryType = TypeVar("RepositoryType", bound="BaseRepository") # noqa: F821
class RepositoriesContainer:
def __init__(self, settings: DatabaseSettings):
engine_parameters = {
"url": settings.dsn,
"pool_recycle": settings.pool_recycle,
}
if not settings.dsn.startswith("sqlite"):
engine_parameters["pool_timeout"] = settings.pool_timeout
engine_parameters["pool_size"] = settings.pool_size
self._engine = create_async_engine(**engine_parameters)
self._session_context = contextvars.ContextVar(
"session_context",
default=None,
)
@property
def engine(self) -> AsyncEngine:
return self._engine
@property
def session(self) -> AsyncSession | None:
return self._session_context.get(None)
@asynccontextmanager
async def transaction(self) -> AsyncGenerator[AsyncSession, None]:
existing_session = self._session_context.get()
if existing_session is not None:
yield existing_session
return
session_parameters = {
"bind": self._engine,
"expire_on_commit": False,
}
async with AsyncSession(**session_parameters) as session:
token = self._session_context.set(session)
try:
async with session.begin():
yield session
finally:
self._session_context.reset(token)
def get_repository(self, repository_class: type[RepositoryType]) -> RepositoryType:
return repository_class(container=self)

17
metaorm/exceptions.py Normal file
View File

@@ -0,0 +1,17 @@
class DatabaseException(Exception):
pass
class NotFoundError(DatabaseException):
def __init__(self, detail: str = "Not found"):
super().__init__(detail)
class HaveNoSessionError(DatabaseException):
def __init__(self):
super().__init__("Have no actual session")
class AlreadyExistsError(DatabaseException):
def __init__(self, detail: str = "Record already exists"):
super().__init__(detail)

214
metaorm/repositories.py Normal file
View File

@@ -0,0 +1,214 @@
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 .tables import BaseTable
class BaseRepository:
def __init__(
self,
settings: DatabaseSettings | None = None,
container: RepositoriesContainer | None = None,
):
if container is not None:
self._container = container
elif settings is not None:
self._container = RepositoriesContainer(settings=settings)
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()
statement = select(func.count()).select_from(table)
statement = append_to_statement(
statement=statement,
model=table,
filter_=filter_,
)
async with self.transaction():
result = await self.session.exec(statement)
items_count = result.first()
return items_count
async def get_items(
self,
filter_: BaseFilter | None = None,
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()
statement = select(table)
statement = append_to_statement(
statement=statement,
model=table,
filter_=filter_,
pagination=pagination,
sort=sort,
)
if options:
statement = statement.options(*options)
async with self.transaction():
result = await self.session.execute(statement)
result = result.yield_per(100)
for db_item in result.scalars():
yield self._convert_from_table(db_item)
async def create_item(self, item: BaseModel) -> BaseModel:
table = self.get_db_table()
values = self._convert_to_table(item).to_values()
statement = insert(table).values(values).returning(table)
async with self.transaction():
try:
result = await self.session.exec(statement)
except IntegrityError as error:
error_message = str(error.orig).lower()
if "unique" in error_message or "duplicate" in error_message:
raise AlreadyExistsError(
f"Record for model '{table.__name__}' already exists",
) from error
raise DatabaseException(
f"Integrity error for model '{table.__name__}'",
) from error
db_item = result.scalar_one()
return self._convert_from_table(db_item)
async def update_items(
self,
filter_: BaseFilter | None = None,
options: Sequence[Any] | None = None,
**values,
) -> AsyncGenerator[BaseModel]:
if filter_ is not None:
self._ensure_filter_fields(type(filter_))
table = self.get_db_table()
statement = update(table)
statement = append_to_statement(
statement=statement,
model=table,
filter_=filter_,
)
statement = statement.values(**values).returning(table)
if options:
statement = statement.options(*options)
async with self.transaction():
result = await self.session.exec(statement)
result = result.yield_per(100)
for db_item in result.scalars():
yield self._convert_from_table(db_item)
async def delete_items(
self,
filter_: BaseFilter | None = None,
) -> None:
if filter_ is not None:
self._ensure_filter_fields(type(filter_))
table = self.get_db_table()
statement = delete(table)
statement = append_to_statement(
statement=statement,
model=table,
filter_=filter_,
)
async with self.transaction():
await self.session.exec(statement)
def _convert_to_table(self, item: BaseModel) -> BaseModel:
table = self.get_db_table()
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:
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
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]
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

8
metaorm/settings.py Normal file
View File

@@ -0,0 +1,8 @@
from pydantic import BaseModel, Field
class DatabaseSettings(BaseModel):
dsn: str = Field(default="sqlite+aiosqlite:///db.sqlite3", pattern=r"^.+://")
pool_size: int = Field(default=5, ge=1)
pool_recycle: int = Field(default=60, ge=1) # in seconds: 1 minute
pool_timeout: int = Field(default=60, ge=1) # in seconds: 1 minute

18
metaorm/tables.py Normal file
View File

@@ -0,0 +1,18 @@
from typing import Any, Self
from pydantic import BaseModel
from sqlmodel import SQLModel
class BaseTable[ItemType: BaseModel](SQLModel):
@classmethod
def from_item(cls, item: ItemType) -> Self:
raise NotImplementedError
def to_item(self) -> ItemType:
raise NotImplementedError
def to_values(self) -> dict[str, Any]:
return {
column.name: getattr(self, column.name) for column in self.__table__.columns
}