refactor: migrate to keyword-based repository config, add nested transactions and tests
This commit is contained in:
@@ -1,8 +1,6 @@
|
||||
import asyncio
|
||||
|
||||
from sqlmodel import Field
|
||||
|
||||
from metaorm import BaseRepository, BaseTable, DatabaseSettings
|
||||
from metaorm import BaseFilter, BaseRepository, BaseTable, Field, RepositorySettings
|
||||
|
||||
|
||||
class UserTable(BaseTable, table=True):
|
||||
@@ -13,13 +11,17 @@ class UserTable(BaseTable, table=True):
|
||||
email: str = Field(unique=True)
|
||||
|
||||
|
||||
class UserRepository(BaseRepository):
|
||||
def get_db_table(self) -> type[UserTable]:
|
||||
return UserTable
|
||||
class UserFilter(BaseFilter):
|
||||
name: str | None = None
|
||||
email: str | None = None
|
||||
|
||||
|
||||
class UserRepository(BaseRepository, table=UserTable, filter_=UserFilter):
|
||||
pass
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
settings = DatabaseSettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
settings = RepositorySettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
repository = UserRepository(settings=settings)
|
||||
|
||||
await repository.create_tables()
|
||||
|
||||
@@ -1,8 +1,13 @@
|
||||
import asyncio
|
||||
|
||||
from sqlmodel import Field
|
||||
|
||||
from metaorm import BaseRepository, BaseTable, DatabaseSettings, RepositoriesContainer
|
||||
from metaorm import (
|
||||
BaseFilter,
|
||||
BaseRepository,
|
||||
BaseTable,
|
||||
Field,
|
||||
RepositoriesContainer,
|
||||
RepositorySettings,
|
||||
)
|
||||
|
||||
|
||||
class UserTable(BaseTable, table=True):
|
||||
@@ -20,18 +25,24 @@ class OrderTable(BaseTable, table=True):
|
||||
total: float
|
||||
|
||||
|
||||
class UserRepository(BaseRepository):
|
||||
def get_db_table(self) -> type[UserTable]:
|
||||
return UserTable
|
||||
class UserFilter(BaseFilter):
|
||||
name: str | None = None
|
||||
|
||||
|
||||
class OrderRepository(BaseRepository):
|
||||
def get_db_table(self) -> type[OrderTable]:
|
||||
return OrderTable
|
||||
class OrderFilter(BaseFilter):
|
||||
user_id: int | None = None
|
||||
|
||||
|
||||
class UserRepository(BaseRepository, table=UserTable, filter_=UserFilter):
|
||||
pass
|
||||
|
||||
|
||||
class OrderRepository(BaseRepository, table=OrderTable, filter_=OrderFilter):
|
||||
pass
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
settings = DatabaseSettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
settings = RepositorySettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
container = RepositoriesContainer(settings=settings)
|
||||
|
||||
user_repo = container.get_repository(UserRepository)
|
||||
@@ -46,6 +57,19 @@ async def main() -> None:
|
||||
await order_repo.create_item(OrderTable(user_id=user.id, total=100.00))
|
||||
await order_repo.create_item(OrderTable(user_id=user.id, total=250.50))
|
||||
|
||||
# Nested transaction inside outer transaction (savepoint)
|
||||
async with container.transaction():
|
||||
user = await user_repo.create_item(UserTable(name="Bob"))
|
||||
try:
|
||||
async with container.nested_transaction():
|
||||
await order_repo.create_item(
|
||||
OrderTable(user_id=user.id, total=999.99),
|
||||
)
|
||||
raise ValueError("Rollback nested order")
|
||||
except ValueError:
|
||||
pass
|
||||
# Bob stays, the order is rolled back
|
||||
|
||||
# Verify results
|
||||
users = [item async for item in user_repo.get_items()]
|
||||
orders = [item async for item in order_repo.get_items()]
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
import asyncio
|
||||
|
||||
from pydantic import BaseModel
|
||||
from sqlmodel import Field
|
||||
|
||||
from metaorm import BaseRepository, BaseTable, DatabaseSettings
|
||||
from metaorm import BaseFilter, BaseRepository, BaseTable, Field, RepositorySettings
|
||||
|
||||
|
||||
class User(BaseModel):
|
||||
@@ -27,16 +26,17 @@ 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
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
settings = DatabaseSettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
settings = RepositorySettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
repository = UserRepository(settings=settings)
|
||||
|
||||
await repository.create_tables()
|
||||
|
||||
@@ -1,9 +1,14 @@
|
||||
import asyncio
|
||||
|
||||
from pydantic_filters import BaseFilter, BaseSort, OffsetPagination
|
||||
from sqlmodel import Field
|
||||
|
||||
from metaorm import BaseRepository, BaseTable, DatabaseSettings
|
||||
from metaorm import (
|
||||
BaseFilter,
|
||||
BaseRepository,
|
||||
BaseSort,
|
||||
BaseTable,
|
||||
Field,
|
||||
OffsetPagination,
|
||||
RepositorySettings,
|
||||
)
|
||||
|
||||
|
||||
class BookTable(BaseTable, table=True):
|
||||
@@ -19,16 +24,12 @@ class BookFilter(BaseFilter):
|
||||
year: int | None = None
|
||||
|
||||
|
||||
class BookRepository(BaseRepository):
|
||||
def get_db_table(self) -> type[BookTable]:
|
||||
return BookTable
|
||||
|
||||
def get_filter_type(self) -> type[BookFilter]:
|
||||
return BookFilter
|
||||
class BookRepository(BaseRepository, table=BookTable, filter_=BookFilter):
|
||||
pass
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
settings = DatabaseSettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
settings = RepositorySettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
repository = BookRepository(settings=settings)
|
||||
|
||||
await repository.create_tables()
|
||||
|
||||
64
examples/nested_transactions.py
Normal file
64
examples/nested_transactions.py
Normal file
@@ -0,0 +1,64 @@
|
||||
import asyncio
|
||||
|
||||
from metaorm import BaseFilter, BaseRepository, BaseTable, Field, RepositorySettings
|
||||
|
||||
|
||||
class ProductTable(BaseTable, table=True):
|
||||
__tablename__ = "products"
|
||||
|
||||
id: int | None = Field(default=None, primary_key=True)
|
||||
name: str
|
||||
price: float
|
||||
|
||||
|
||||
class ProductFilter(BaseFilter):
|
||||
name: str | None = None
|
||||
|
||||
|
||||
class ProductRepository(BaseRepository, table=ProductTable, filter_=ProductFilter):
|
||||
pass
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
settings = RepositorySettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
repository = ProductRepository(settings=settings)
|
||||
|
||||
await repository.create_tables()
|
||||
|
||||
# Standalone nested transaction: rollback only the inner scope
|
||||
try:
|
||||
async with repository.nested_transaction():
|
||||
await repository.create_item(ProductTable(name="Laptop", price=999.99))
|
||||
await repository.create_item(ProductTable(name="Mouse", price=29.99))
|
||||
raise ValueError("Simulated error inside nested transaction")
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
count = await repository.get_items_count()
|
||||
print(f"Items after standalone nested rollback: {count}") # 0
|
||||
|
||||
# Nested transaction inside an outer transaction
|
||||
async with repository.transaction():
|
||||
await repository.create_item(ProductTable(name="Keyboard", price=79.99))
|
||||
|
||||
try:
|
||||
async with repository.nested_transaction():
|
||||
await repository.create_item(
|
||||
ProductTable(name="Monitor", price=299.99),
|
||||
)
|
||||
raise ValueError("Nested rollback")
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
# Monitor is rolled back, Keyboard stays in the outer transaction
|
||||
items = [item async for item in repository.get_items()]
|
||||
print(f"Items after partial rollback: {len(items)}") # 1
|
||||
print(items[0].name) # Keyboard
|
||||
|
||||
# Verify committed results
|
||||
all_items = [item async for item in repository.get_items()]
|
||||
print(f"Final items: {[item.name for item in all_items]}") # ["Keyboard"]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
@@ -1,9 +1,16 @@
|
||||
import asyncio
|
||||
|
||||
from sqlalchemy.orm import joinedload
|
||||
from sqlmodel import Field, Relationship
|
||||
|
||||
from metaorm import BaseRepository, BaseTable, DatabaseSettings, RepositoriesContainer
|
||||
from metaorm import (
|
||||
BaseFilter,
|
||||
BaseRepository,
|
||||
BaseTable,
|
||||
Field,
|
||||
Relationship,
|
||||
RepositoriesContainer,
|
||||
RepositorySettings,
|
||||
)
|
||||
|
||||
|
||||
class AuthorTable(BaseTable, table=True):
|
||||
@@ -23,18 +30,25 @@ class BookTable(BaseTable, table=True):
|
||||
author: AuthorTable = Relationship(back_populates="books")
|
||||
|
||||
|
||||
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):
|
||||
def get_db_table(self) -> type[AuthorTable]:
|
||||
return AuthorTable
|
||||
class AuthorFilter(BaseFilter):
|
||||
name: str | None = None
|
||||
|
||||
|
||||
class BookRepository(BaseRepository, table=BookTable, filter_=BookFilter):
|
||||
pass
|
||||
|
||||
|
||||
class AuthorRepository(BaseRepository, table=AuthorTable, filter_=AuthorFilter):
|
||||
pass
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
settings = DatabaseSettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
settings = RepositorySettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
container = RepositoriesContainer(settings=settings)
|
||||
|
||||
author_repo = AuthorRepository(container=container)
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
import asyncio
|
||||
|
||||
from sqlmodel import Field
|
||||
|
||||
from metaorm import BaseRepository, BaseTable, DatabaseSettings
|
||||
from metaorm import BaseFilter, BaseRepository, BaseTable, Field, RepositorySettings
|
||||
|
||||
|
||||
class ProductTable(BaseTable, table=True):
|
||||
@@ -13,13 +11,17 @@ class ProductTable(BaseTable, table=True):
|
||||
price: float
|
||||
|
||||
|
||||
class ProductRepository(BaseRepository):
|
||||
def get_db_table(self) -> type[ProductTable]:
|
||||
return ProductTable
|
||||
class ProductFilter(BaseFilter):
|
||||
name: str | None = None
|
||||
price: int | None = None
|
||||
|
||||
|
||||
class ProductRepository(BaseRepository, table=ProductTable, filter_=ProductFilter):
|
||||
pass
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
settings = DatabaseSettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
settings = RepositorySettings(dsn="sqlite+aiosqlite:///:memory:")
|
||||
repository = ProductRepository(settings=settings)
|
||||
|
||||
await repository.create_tables()
|
||||
@@ -34,10 +36,23 @@ async def main() -> None:
|
||||
)
|
||||
print(f"Created in transaction: {product1.name}, {product2.name}")
|
||||
|
||||
# Nested transaction reuses existing session
|
||||
# Reusing an existing session (no new savepoint)
|
||||
async with repository.transaction(), repository.transaction():
|
||||
items = [item async for item in repository.get_items()]
|
||||
print(f"Items in nested transaction: {len(items)}")
|
||||
print(f"Items in reused session: {len(items)}")
|
||||
|
||||
# True nested transaction (savepoint) via repository
|
||||
try:
|
||||
async with repository.nested_transaction():
|
||||
await repository.create_item(
|
||||
ProductTable(name="Keyboard", price=79.99),
|
||||
)
|
||||
raise ValueError("Rollback nested")
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
count = await repository.get_items_count()
|
||||
print(f"Items after nested rollback: {count}") # 2
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user