fix: make lock_tables generic over session type (#270)

This commit is contained in:
d3vyce
2026-05-04 20:25:00 +02:00
committed by GitHub
parent af4c57c293
commit e4e3f0ec60
+16 -13
View File
@@ -151,14 +151,13 @@ class LockMode(str, Enum):
ACCESS_EXCLUSIVE = "ACCESS EXCLUSIVE" ACCESS_EXCLUSIVE = "ACCESS EXCLUSIVE"
@asynccontextmanager def lock_tables(
async def lock_tables( session_maker: async_sessionmaker[_SessionT],
session_maker: async_sessionmaker[AsyncSession],
tables: list[type[DeclarativeBase]], tables: list[type[DeclarativeBase]],
*, *,
mode: LockMode = LockMode.SHARE_UPDATE_EXCLUSIVE, mode: LockMode = LockMode.SHARE_UPDATE_EXCLUSIVE,
timeout: str = "5s", timeout: str = "5s",
) -> AsyncGenerator[AsyncSession, None]: ) -> AbstractAsyncContextManager[_SessionT]:
"""Lock PostgreSQL tables for the duration of a transaction. """Lock PostgreSQL tables for the duration of a transaction.
Args: Args:
@@ -190,15 +189,19 @@ async def lock_tables(
""" """
table_names = ",".join(table.__tablename__ for table in tables) table_names = ",".join(table.__tablename__ for table in tables)
async with session_maker() as session: @asynccontextmanager
try: async def _lock() -> AsyncGenerator[_SessionT, None]:
await session.execute(text(f"SET LOCAL lock_timeout='{timeout}'")) async with session_maker() as session:
await session.execute(text(f"LOCK {table_names} IN {mode.value} MODE")) try:
yield session await session.execute(text(f"SET LOCAL lock_timeout='{timeout}'"))
await session.commit() await session.execute(text(f"LOCK {table_names} IN {mode.value} MODE"))
except BaseException: yield session
await session.rollback() await session.commit()
raise except BaseException:
await session.rollback()
raise
return _lock()
async def create_database( async def create_database(