refactor: migrate to keyword-based repository config, add nested transactions and tests
This commit is contained in:
@@ -3,7 +3,7 @@ from collections.abc import AsyncGenerator
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
from metaorm import DatabaseSettings, RepositoriesContainer
|
||||
from metaorm import RepositoriesContainer, RepositorySettings
|
||||
|
||||
from .models import (
|
||||
AuthorRepository,
|
||||
@@ -14,8 +14,8 @@ from .models import (
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def database_settings() -> DatabaseSettings:
|
||||
return DatabaseSettings(
|
||||
def database_settings() -> RepositorySettings:
|
||||
return RepositorySettings(
|
||||
dsn="sqlite+aiosqlite:///:memory:",
|
||||
pool_size=1,
|
||||
pool_recycle=60,
|
||||
@@ -25,7 +25,7 @@ def database_settings() -> DatabaseSettings:
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def repositories_container(
|
||||
database_settings: DatabaseSettings,
|
||||
database_settings: RepositorySettings,
|
||||
) -> AsyncGenerator[RepositoriesContainer, None]:
|
||||
container = RepositoriesContainer(settings=database_settings)
|
||||
yield container
|
||||
@@ -43,7 +43,7 @@ async def user_repository(
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def product_repository_settings(
|
||||
database_settings: DatabaseSettings,
|
||||
database_settings: RepositorySettings,
|
||||
) -> AsyncGenerator[ProductRepository, None]:
|
||||
repository = ProductRepository(settings=database_settings)
|
||||
await repository.create_tables()
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
from pydantic import BaseModel
|
||||
from pydantic_filters import BaseFilter
|
||||
from sqlmodel import Field, Relationship
|
||||
|
||||
from metaorm import BaseRepository, BaseTable
|
||||
from metaorm import BaseFilter, BaseRepository, BaseTable, Field, Relationship
|
||||
|
||||
|
||||
class User(BaseModel):
|
||||
@@ -25,12 +23,13 @@ class UserTable(BaseTable[User], table=True):
|
||||
return User(id=self.id, name=self.name, email=self.email)
|
||||
|
||||
|
||||
class UserRepository(BaseRepository):
|
||||
def get_db_table(self) -> type[UserTable]:
|
||||
return UserTable
|
||||
class UserFilter(BaseFilter):
|
||||
name: str | None = None
|
||||
email: str | None = None
|
||||
|
||||
def get_dto_type(self) -> type[User]:
|
||||
return User
|
||||
|
||||
class UserRepository(BaseRepository, table=UserTable, filter_=UserFilter, dto=User):
|
||||
pass
|
||||
|
||||
|
||||
class ProductTable(BaseTable, table=True):
|
||||
@@ -45,12 +44,8 @@ class ProductFilter(BaseFilter):
|
||||
price: int | None = None
|
||||
|
||||
|
||||
class ProductRepository(BaseRepository):
|
||||
def get_db_table(self) -> type[ProductTable]:
|
||||
return ProductTable
|
||||
|
||||
def get_filter_type(self) -> type[ProductFilter]:
|
||||
return ProductFilter
|
||||
class ProductRepository(BaseRepository, table=ProductTable, filter_=ProductFilter):
|
||||
pass
|
||||
|
||||
|
||||
class AuthorTable(BaseTable, table=True):
|
||||
@@ -68,11 +63,18 @@ class BookTable(BaseTable, table=True):
|
||||
author: AuthorTable = Relationship(back_populates="books")
|
||||
|
||||
|
||||
class AuthorRepository(BaseRepository):
|
||||
def get_db_table(self) -> type[AuthorTable]:
|
||||
return AuthorTable
|
||||
class AuthorFilter(BaseFilter):
|
||||
name: str | None = None
|
||||
|
||||
|
||||
class BookRepository(BaseRepository):
|
||||
def get_db_table(self) -> type[BookTable]:
|
||||
return BookTable
|
||||
class BookFilter(BaseFilter):
|
||||
title: str | None = None
|
||||
author_id: int | None = None
|
||||
|
||||
|
||||
class AuthorRepository(BaseRepository, table=AuthorTable, filter_=AuthorFilter):
|
||||
pass
|
||||
|
||||
|
||||
class BookRepository(BaseRepository, table=BookTable, filter_=BookFilter):
|
||||
pass
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
|
||||
from metaorm import RepositoriesContainer
|
||||
from tests.models import UserRepository
|
||||
from tests.models import User, UserRepository
|
||||
|
||||
|
||||
class TestRepositoriesContainer:
|
||||
@@ -39,6 +40,69 @@ class TestRepositoriesContainer:
|
||||
):
|
||||
assert inner_session is outer_session
|
||||
|
||||
async def test_nested_transaction_creates_session(
|
||||
self,
|
||||
repositories_container: RepositoriesContainer,
|
||||
) -> None:
|
||||
assert repositories_container.session is None
|
||||
|
||||
async with repositories_container.nested_transaction() as session:
|
||||
assert session is not None
|
||||
assert repositories_container.session is session
|
||||
|
||||
assert repositories_container.session is None
|
||||
|
||||
async def test_nested_transaction_reuses_outer_session(
|
||||
self,
|
||||
repositories_container: RepositoriesContainer,
|
||||
) -> None:
|
||||
async with (
|
||||
repositories_container.transaction() as outer_session,
|
||||
repositories_container.nested_transaction() as inner_session,
|
||||
):
|
||||
assert inner_session is outer_session
|
||||
|
||||
async def test_nested_transaction_rollbacks_on_exception(
|
||||
self,
|
||||
repositories_container: RepositoriesContainer,
|
||||
) -> None:
|
||||
repository = repositories_container.get_repository(UserRepository)
|
||||
await repository.create_tables()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
async with repositories_container.nested_transaction():
|
||||
await repository.create_item(
|
||||
User(name="Alice", email="alice@example.com"),
|
||||
)
|
||||
raise ValueError("boom")
|
||||
|
||||
count = await repository.get_items_count()
|
||||
assert count == 0
|
||||
|
||||
async def test_nested_transaction_in_outer_transaction_rollbacks_only_inner(
|
||||
self,
|
||||
repositories_container: RepositoriesContainer,
|
||||
) -> None:
|
||||
repository = repositories_container.get_repository(UserRepository)
|
||||
await repository.create_tables()
|
||||
|
||||
async with repositories_container.transaction():
|
||||
await repository.create_item(
|
||||
User(name="Bob", email="bob@example.com"),
|
||||
)
|
||||
with pytest.raises(ValueError):
|
||||
async with repositories_container.nested_transaction():
|
||||
await repository.create_item(
|
||||
User(name="Alice", email="alice@example.com"),
|
||||
)
|
||||
raise ValueError("boom")
|
||||
|
||||
count = await repository.get_items_count()
|
||||
assert count == 1
|
||||
|
||||
items = [item async for item in repository.get_items()]
|
||||
assert items[0].name == "Bob"
|
||||
|
||||
async def test_get_repository_returns_repository_instance(
|
||||
self,
|
||||
repositories_container: RepositoriesContainer,
|
||||
|
||||
@@ -1,8 +1,13 @@
|
||||
import pytest
|
||||
from pydantic_filters import BaseSort, OffsetPagination
|
||||
from sqlalchemy.orm import joinedload
|
||||
|
||||
from metaorm import AlreadyExistsError, DatabaseSettings, HaveNoSessionError
|
||||
from metaorm import (
|
||||
AlreadyExistsError,
|
||||
BaseRepository,
|
||||
BaseSort,
|
||||
OffsetPagination,
|
||||
RepositorySettings,
|
||||
)
|
||||
from tests.models import (
|
||||
AuthorRepository,
|
||||
AuthorTable,
|
||||
@@ -12,18 +17,13 @@ from tests.models import (
|
||||
ProductRepository,
|
||||
ProductTable,
|
||||
User,
|
||||
UserFilter,
|
||||
UserRepository,
|
||||
UserTable,
|
||||
)
|
||||
|
||||
|
||||
class TestBaseRepository:
|
||||
async def test_session_raises_error_without_transaction(
|
||||
self,
|
||||
user_repository: UserRepository,
|
||||
) -> None:
|
||||
with pytest.raises(HaveNoSessionError):
|
||||
_ = user_repository.session
|
||||
|
||||
async def test_create_item(self, user_repository: UserRepository) -> None:
|
||||
user = User(name="Alice", email="alice@example.com")
|
||||
|
||||
@@ -93,12 +93,29 @@ class TestBaseRepository:
|
||||
with pytest.raises(AlreadyExistsError):
|
||||
await user_repository.create_item(user)
|
||||
|
||||
async def test_transaction_reuses_existing_session(
|
||||
async def test_transaction_scope_allows_crud(
|
||||
self,
|
||||
user_repository: UserRepository,
|
||||
) -> None:
|
||||
async with user_repository.transaction():
|
||||
_ = user_repository.session
|
||||
count = await user_repository.get_items_count()
|
||||
assert count == 0
|
||||
|
||||
async def test_session_is_none_without_transaction(
|
||||
self,
|
||||
user_repository: UserRepository,
|
||||
) -> None:
|
||||
assert user_repository.session is None
|
||||
|
||||
async def test_session_returns_session_inside_transaction(
|
||||
self,
|
||||
user_repository: UserRepository,
|
||||
) -> None:
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
async with user_repository.transaction() as session:
|
||||
assert isinstance(user_repository.session, AsyncSession)
|
||||
assert user_repository.session is session
|
||||
|
||||
async def test_get_items_with_pagination(
|
||||
self,
|
||||
@@ -147,11 +164,83 @@ class TestBaseRepository:
|
||||
assert created.name == "Widget"
|
||||
|
||||
async def test_get_filter_type(self) -> None:
|
||||
repository = ProductRepository(settings=DatabaseSettings())
|
||||
product_repository = ProductRepository(settings=RepositorySettings())
|
||||
user_repository = UserRepository(settings=RepositorySettings())
|
||||
|
||||
filter_type = repository.get_filter_type()
|
||||
assert product_repository.get_filter_type() is ProductFilter
|
||||
assert user_repository.get_filter_type() is UserFilter
|
||||
|
||||
assert filter_type is ProductFilter
|
||||
async def test_get_dto_type(self) -> None:
|
||||
product_repository = ProductRepository(settings=RepositorySettings())
|
||||
user_repository = UserRepository(settings=RepositorySettings())
|
||||
|
||||
assert product_repository.get_dto_type() is None
|
||||
assert user_repository.get_dto_type() is User
|
||||
|
||||
async def test_get_items_count_with_filter(
|
||||
self,
|
||||
product_repository_settings: ProductRepository,
|
||||
) -> None:
|
||||
await product_repository_settings.create_item(
|
||||
ProductTable(name="Alpha", price=10.0),
|
||||
)
|
||||
await product_repository_settings.create_item(
|
||||
ProductTable(name="Beta", price=20.0),
|
||||
)
|
||||
|
||||
count = await product_repository_settings.get_items_count(
|
||||
filter_=ProductFilter(name="Alpha"),
|
||||
)
|
||||
|
||||
assert count == 1
|
||||
|
||||
async def test_delete_items_with_filter(
|
||||
self,
|
||||
product_repository_settings: ProductRepository,
|
||||
) -> None:
|
||||
await product_repository_settings.create_item(
|
||||
ProductTable(name="Alpha", price=10.0),
|
||||
)
|
||||
await product_repository_settings.create_item(
|
||||
ProductTable(name="Beta", price=20.0),
|
||||
)
|
||||
|
||||
await product_repository_settings.delete_items(
|
||||
filter_=ProductFilter(name="Alpha"),
|
||||
)
|
||||
|
||||
count = await product_repository_settings.get_items_count()
|
||||
assert count == 1
|
||||
|
||||
remaining = [item async for item in product_repository_settings.get_items()]
|
||||
assert remaining[0].name == "Beta"
|
||||
|
||||
async def test_update_items_with_filter(
|
||||
self,
|
||||
product_repository_settings: ProductRepository,
|
||||
) -> None:
|
||||
await product_repository_settings.create_item(
|
||||
ProductTable(name="Alpha", price=10.0),
|
||||
)
|
||||
await product_repository_settings.create_item(
|
||||
ProductTable(name="Beta", price=20.0),
|
||||
)
|
||||
|
||||
updated = [
|
||||
item
|
||||
async for item in product_repository_settings.update_items(
|
||||
filter_=ProductFilter(name="Alpha"),
|
||||
name="Gamma",
|
||||
)
|
||||
]
|
||||
|
||||
assert len(updated) == 1
|
||||
assert updated[0].name == "Gamma"
|
||||
|
||||
all_items = [item async for item in product_repository_settings.get_items()]
|
||||
assert len(all_items) == 2
|
||||
names = {item.name for item in all_items}
|
||||
assert names == {"Gamma", "Beta"}
|
||||
|
||||
async def test_get_items_with_filter(
|
||||
self,
|
||||
@@ -217,3 +306,72 @@ class TestBaseRepository:
|
||||
|
||||
assert len(updated) == 1
|
||||
assert updated[0].title == "Updated"
|
||||
|
||||
async def test_init_raises_without_container_or_settings(self) -> None:
|
||||
with pytest.raises(TypeError):
|
||||
BaseRepository()
|
||||
|
||||
async def test_repository_without_table_raises(self) -> None:
|
||||
with pytest.raises(TypeError):
|
||||
class BadRepository(BaseRepository, filter_=UserFilter):
|
||||
pass
|
||||
|
||||
async def test_nested_transaction_property_allows_crud(
|
||||
self,
|
||||
user_repository: UserRepository,
|
||||
) -> None:
|
||||
async with user_repository.nested_transaction():
|
||||
count = await user_repository.get_items_count()
|
||||
assert count == 0
|
||||
|
||||
async def test_nested_transaction_rollbacks_inner_scope(
|
||||
self,
|
||||
user_repository: UserRepository,
|
||||
) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
async with user_repository.nested_transaction():
|
||||
await user_repository.create_item(
|
||||
User(name="Alice", email="alice@example.com"),
|
||||
)
|
||||
raise ValueError("boom")
|
||||
|
||||
count = await user_repository.get_items_count()
|
||||
assert count == 0
|
||||
|
||||
async def test_nested_transaction_in_outer_transaction_rollbacks_only_inner(
|
||||
self,
|
||||
user_repository: UserRepository,
|
||||
) -> None:
|
||||
async with user_repository.transaction():
|
||||
await user_repository.create_item(
|
||||
User(name="Bob", email="bob@example.com"),
|
||||
)
|
||||
with pytest.raises(ValueError):
|
||||
async with user_repository.nested_transaction():
|
||||
await user_repository.create_item(
|
||||
User(name="Alice", email="alice@example.com"),
|
||||
)
|
||||
raise ValueError("boom")
|
||||
|
||||
count = await user_repository.get_items_count()
|
||||
assert count == 1
|
||||
|
||||
items = [item async for item in user_repository.get_items()]
|
||||
assert items[0].name == "Bob"
|
||||
|
||||
async def test_params_via_intermediate_base_class(self) -> None:
|
||||
class IntermediateRepository(
|
||||
BaseRepository,
|
||||
table=UserTable,
|
||||
filter_=UserFilter,
|
||||
dto=User,
|
||||
):
|
||||
pass
|
||||
|
||||
class ConcreteRepository(IntermediateRepository):
|
||||
pass
|
||||
|
||||
repository = ConcreteRepository(settings=RepositorySettings())
|
||||
|
||||
assert repository.get_filter_type() is UserFilter
|
||||
assert repository.get_dto_type() is User
|
||||
|
||||
@@ -1,12 +1,12 @@
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from metaorm import DatabaseSettings
|
||||
from metaorm import RepositorySettings
|
||||
|
||||
|
||||
class TestDatabaseSettings:
|
||||
class TestRepositorySettings:
|
||||
def test_default_values(self) -> None:
|
||||
settings = DatabaseSettings()
|
||||
settings = RepositorySettings()
|
||||
|
||||
assert settings.dsn == "sqlite+aiosqlite:///db.sqlite3"
|
||||
assert settings.pool_size == 5
|
||||
@@ -14,7 +14,7 @@ class TestDatabaseSettings:
|
||||
assert settings.pool_timeout == 60
|
||||
|
||||
def test_custom_values(self) -> None:
|
||||
settings = DatabaseSettings(
|
||||
settings = RepositorySettings(
|
||||
dsn="postgresql+asyncpg://user:pass@localhost/db",
|
||||
pool_size=10,
|
||||
pool_recycle=120,
|
||||
@@ -28,7 +28,7 @@ class TestDatabaseSettings:
|
||||
|
||||
def test_dsn_must_match_pattern(self) -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
DatabaseSettings(dsn="invalid_dsn")
|
||||
RepositorySettings(dsn="invalid_dsn")
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"field_name,invalid_value",
|
||||
@@ -44,4 +44,4 @@ class TestDatabaseSettings:
|
||||
invalid_value: int,
|
||||
) -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
DatabaseSettings(**{field_name: invalid_value})
|
||||
RepositorySettings(**{field_name: invalid_value})
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
import pytest
|
||||
from sqlmodel import Field
|
||||
|
||||
from metaorm import BaseTable
|
||||
from metaorm import BaseTable, Field
|
||||
from tests.models import User
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user