Initial release: async repository layer over SQLModel
This commit is contained in:
26
metaorm/__init__.py
Normal file
26
metaorm/__init__.py
Normal 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
58
metaorm/container.py
Normal 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
17
metaorm/exceptions.py
Normal 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
214
metaorm/repositories.py
Normal 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
8
metaorm/settings.py
Normal 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
18
metaorm/tables.py
Normal 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
|
||||
}
|
||||
Reference in New Issue
Block a user