diff --git a/packages/common-library/src/common_library/unit_of_work.py b/packages/common-library/src/common_library/unit_of_work.py new file mode 100644 index 00000000000..9e729e61adb --- /dev/null +++ b/packages/common-library/src/common_library/unit_of_work.py @@ -0,0 +1,26 @@ +from abc import ABC, abstractmethod +from contextlib import AbstractAsyncContextManager + + +class ReadUnitOfWork: + """Opaque persistence scope for sequential reads.""" + + +class TransactionalUnitOfWork(ReadUnitOfWork): + """Opaque persistence scope for sequential reads and writes.""" + + +class UnitOfWorkFactory(ABC): + @abstractmethod + def read( + self, + *, + existing: ReadUnitOfWork | None = None, + ) -> AbstractAsyncContextManager[ReadUnitOfWork]: ... + + @abstractmethod + def transaction( + self, + *, + existing: TransactionalUnitOfWork | None = None, + ) -> AbstractAsyncContextManager[TransactionalUnitOfWork]: ... diff --git a/packages/common-library/tests/test_unit_of_work.py b/packages/common-library/tests/test_unit_of_work.py new file mode 100644 index 00000000000..f844bc6fc86 --- /dev/null +++ b/packages/common-library/tests/test_unit_of_work.py @@ -0,0 +1,75 @@ +import inspect +from contextlib import AbstractAsyncContextManager +from types import TracebackType + +from common_library.unit_of_work import ( + ReadUnitOfWork, + TransactionalUnitOfWork, + UnitOfWorkFactory, +) + + +class _ReadUnitOfWork(ReadUnitOfWork): ... + + +class _TransactionalUnitOfWork(TransactionalUnitOfWork): ... + + +class _UnitOfWorkContext[UnitOfWorkT: ReadUnitOfWork]: + def __init__(self, unit_of_work: UnitOfWorkT) -> None: + self._unit_of_work = unit_of_work + + async def __aenter__(self) -> UnitOfWorkT: + return self._unit_of_work + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool: + return False + + +class _IncompleteUnitOfWorkFactory(UnitOfWorkFactory): ... + + +class _UnitOfWorkFactory(UnitOfWorkFactory): + def read( + self, + *, + existing: ReadUnitOfWork | None = None, + ) -> AbstractAsyncContextManager[ReadUnitOfWork]: + self._read_calls += 1 + return _UnitOfWorkContext(existing or _ReadUnitOfWork()) + + def transaction( + self, + *, + existing: TransactionalUnitOfWork | None = None, + ) -> AbstractAsyncContextManager[TransactionalUnitOfWork]: + self._transaction_calls += 1 + return _UnitOfWorkContext(existing or _TransactionalUnitOfWork()) + + def __init__(self) -> None: + self._read_calls = 0 + self._transaction_calls = 0 + + +def test_incomplete_unit_of_work_factory_cannot_be_instantiated(): + assert inspect.isabstract(_IncompleteUnitOfWorkFactory) + + +async def test_unit_of_work_factory_contract_supports_new_and_existing_scopes(): + factory = _UnitOfWorkFactory() + + async with factory.read() as read_uow: + assert isinstance(read_uow, ReadUnitOfWork) + async with factory.read(existing=read_uow) as reused_read_uow: + assert reused_read_uow is read_uow + + async with factory.transaction() as transactional_uow: + assert isinstance(transactional_uow, TransactionalUnitOfWork) + assert isinstance(transactional_uow, ReadUnitOfWork) + async with factory.transaction(existing=transactional_uow) as reused_transactional_uow: + assert reused_transactional_uow is transactional_uow diff --git a/packages/postgres-database/src/simcore_postgres_database/unit_of_work.py b/packages/postgres-database/src/simcore_postgres_database/unit_of_work.py new file mode 100644 index 00000000000..8bc6c5553be --- /dev/null +++ b/packages/postgres-database/src/simcore_postgres_database/unit_of_work.py @@ -0,0 +1,86 @@ +from collections.abc import AsyncIterator +from contextlib import AbstractAsyncContextManager, asynccontextmanager +from dataclasses import dataclass + +from common_library.unit_of_work import ( + ReadUnitOfWork, + TransactionalUnitOfWork, + UnitOfWorkFactory, +) +from sqlalchemy.ext.asyncio import AsyncConnection, AsyncEngine + + +@dataclass(frozen=True, kw_only=True, slots=True) +class _SqlAlchemyReadUnitOfWork(ReadUnitOfWork): + connection: AsyncConnection + + +@dataclass(frozen=True, kw_only=True, slots=True) +class _SqlAlchemyTransactionalUnitOfWork(TransactionalUnitOfWork): + connection: AsyncConnection + + +def get_sqlalchemy_connection(unit_of_work: ReadUnitOfWork) -> AsyncConnection: + if isinstance( + unit_of_work, + (_SqlAlchemyReadUnitOfWork, _SqlAlchemyTransactionalUnitOfWork), + ): + return unit_of_work.connection + msg = f"Expected a SQLAlchemy unit of work, got {type(unit_of_work).__name__}" + raise TypeError(msg) + + +def get_sqlalchemy_transaction_connection( + unit_of_work: TransactionalUnitOfWork, +) -> AsyncConnection: + if isinstance(unit_of_work, _SqlAlchemyTransactionalUnitOfWork): + return unit_of_work.connection + msg = f"Expected a SQLAlchemy transactional unit of work, got {type(unit_of_work).__name__}" + raise TypeError(msg) + + +@asynccontextmanager +async def _read_scope( + engine: AsyncEngine, + existing: ReadUnitOfWork | None, +) -> AsyncIterator[ReadUnitOfWork]: + if existing is not None: + get_sqlalchemy_connection(existing) + yield existing + return + + async with engine.connect() as connection: + yield _SqlAlchemyReadUnitOfWork(connection=connection) + + +@asynccontextmanager +async def _transaction_scope( + engine: AsyncEngine, + existing: TransactionalUnitOfWork | None, +) -> AsyncIterator[TransactionalUnitOfWork]: + if existing is not None: + get_sqlalchemy_transaction_connection(existing) + yield existing + return + + async with engine.begin() as connection: + yield _SqlAlchemyTransactionalUnitOfWork(connection=connection) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class SqlAlchemyUnitOfWorkFactory(UnitOfWorkFactory): + engine: AsyncEngine + + def read( + self, + *, + existing: ReadUnitOfWork | None = None, + ) -> AbstractAsyncContextManager[ReadUnitOfWork]: + return _read_scope(self.engine, existing) + + def transaction( + self, + *, + existing: TransactionalUnitOfWork | None = None, + ) -> AbstractAsyncContextManager[TransactionalUnitOfWork]: + return _transaction_scope(self.engine, existing) diff --git a/packages/postgres-database/tests/unit_of_work/conftest.py b/packages/postgres-database/tests/unit_of_work/conftest.py new file mode 100644 index 00000000000..d2eb18f9799 --- /dev/null +++ b/packages/postgres-database/tests/unit_of_work/conftest.py @@ -0,0 +1,68 @@ +from dataclasses import dataclass +from types import TracebackType +from typing import cast + +import pytest +from simcore_postgres_database.unit_of_work import SqlAlchemyUnitOfWorkFactory +from sqlalchemy.ext.asyncio import AsyncConnection, AsyncEngine + + +@dataclass +class _ScopeState: + closed: bool = False + committed: bool = False + rolled_back: bool = False + + +class _ConnectionScope: + def __init__(self, connection: AsyncConnection, state: _ScopeState) -> None: + self._connection = connection + self._state = state + + async def __aenter__(self) -> AsyncConnection: + return self._connection + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool: + self._state.closed = True + return False + + +class _TransactionScope(_ConnectionScope): + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool: + self._state.closed = True + self._state.committed = exc_type is None + self._state.rolled_back = exc_type is not None + return False + + +class _Engine: + def __init__(self) -> None: + self.connection = cast(AsyncConnection, object()) + self.read_scopes: list[_ScopeState] = [] + self.transaction_scopes: list[_ScopeState] = [] + + def connect(self) -> _ConnectionScope: + state = _ScopeState() + self.read_scopes.append(state) + return _ConnectionScope(self.connection, state) + + def begin(self) -> _TransactionScope: + state = _ScopeState() + self.transaction_scopes.append(state) + return _TransactionScope(self.connection, state) + + +@pytest.fixture +def sqlalchemy_uow_factory() -> tuple[SqlAlchemyUnitOfWorkFactory, _Engine]: + engine = _Engine() + return SqlAlchemyUnitOfWorkFactory(engine=cast(AsyncEngine, engine)), engine diff --git a/packages/postgres-database/tests/unit_of_work/test_sqlalchemy_unit_of_work.py b/packages/postgres-database/tests/unit_of_work/test_sqlalchemy_unit_of_work.py new file mode 100644 index 00000000000..b63d430eb01 --- /dev/null +++ b/packages/postgres-database/tests/unit_of_work/test_sqlalchemy_unit_of_work.py @@ -0,0 +1,146 @@ +from typing import Protocol, cast + +import pytest +from common_library.unit_of_work import ReadUnitOfWork, TransactionalUnitOfWork +from simcore_postgres_database.unit_of_work import ( + SqlAlchemyUnitOfWorkFactory, + get_sqlalchemy_connection, + get_sqlalchemy_transaction_connection, +) +from sqlalchemy.ext.asyncio import AsyncConnection + + +class _ScopeState(Protocol): + closed: bool + committed: bool + rolled_back: bool + + +class _Engine(Protocol): + connection: AsyncConnection + read_scopes: list[_ScopeState] + transaction_scopes: list[_ScopeState] + + +type SqlAlchemyUowFixture = tuple[SqlAlchemyUnitOfWorkFactory, _Engine] + + +class _ForeignReadUnitOfWork(ReadUnitOfWork): ... + + +class _ForeignTransactionalUnitOfWork(TransactionalUnitOfWork): ... + + +async def test_read_scope_acquires_lazily_and_closes_owned_connection( + sqlalchemy_uow_factory: SqlAlchemyUowFixture, +): + factory, engine = sqlalchemy_uow_factory + + scope = factory.read() + assert engine.read_scopes == [] + + async with scope as unit_of_work: + assert len(engine.read_scopes) == 1 + assert engine.read_scopes[0].closed is False + assert get_sqlalchemy_connection(unit_of_work) is engine.connection + + assert engine.read_scopes[0].closed is True + + +async def test_read_scope_closes_owned_connection_on_error( + sqlalchemy_uow_factory: SqlAlchemyUowFixture, +): + factory, engine = sqlalchemy_uow_factory + expected_error = RuntimeError("read failed") + + with pytest.raises(RuntimeError, match="read failed") as exc_info: + async with factory.read(): + raise expected_error + + assert exc_info.value is expected_error + assert engine.read_scopes[0].closed is True + + +async def test_read_scope_reuses_existing_read_or_transaction_scope( + sqlalchemy_uow_factory: SqlAlchemyUowFixture, +): + factory, engine = sqlalchemy_uow_factory + + async with ( + factory.read() as read_unit_of_work, + factory.read(existing=read_unit_of_work) as reused_unit_of_work, + ): + assert reused_unit_of_work is read_unit_of_work + assert len(engine.read_scopes) == 1 + + async with ( + factory.transaction() as transaction_unit_of_work, + factory.read(existing=transaction_unit_of_work) as reused_unit_of_work, + ): + assert reused_unit_of_work is transaction_unit_of_work + assert len(engine.read_scopes) == 1 + + +async def test_transaction_scope_commits_or_rolls_back_and_closes_owned_connection( + sqlalchemy_uow_factory: SqlAlchemyUowFixture, +): + factory, engine = sqlalchemy_uow_factory + + async with factory.transaction() as unit_of_work: + assert get_sqlalchemy_transaction_connection(unit_of_work) is engine.connection + + committed_state = engine.transaction_scopes[0] + assert committed_state.closed is True + assert committed_state.committed is True + assert committed_state.rolled_back is False + + expected_error = RuntimeError("rollback") + with pytest.raises(RuntimeError, match="rollback") as exc_info: + async with factory.transaction(): + raise expected_error + + assert exc_info.value is expected_error + rolled_back_state = engine.transaction_scopes[1] + assert rolled_back_state.closed is True + assert rolled_back_state.committed is False + assert rolled_back_state.rolled_back is True + + +async def test_transaction_scope_reuses_existing_scope_without_owning_it( + sqlalchemy_uow_factory: SqlAlchemyUowFixture, +): + factory, engine = sqlalchemy_uow_factory + + async with factory.transaction() as unit_of_work: + outer_state = engine.transaction_scopes[0] + + async with factory.transaction(existing=unit_of_work) as reused_unit_of_work: + assert reused_unit_of_work is unit_of_work + assert len(engine.transaction_scopes) == 1 + + assert outer_state.closed is False + assert outer_state.committed is False + assert outer_state.rolled_back is False + + assert outer_state.closed is True + assert outer_state.committed is True + assert outer_state.rolled_back is False + + +async def test_transaction_scope_rejects_read_only_and_foreign_units_of_work( + sqlalchemy_uow_factory: SqlAlchemyUowFixture, +): + factory, _ = sqlalchemy_uow_factory + + async with factory.read() as read_unit_of_work: + with pytest.raises(TypeError, match="transactional unit of work"): + async with factory.transaction(existing=cast(TransactionalUnitOfWork, read_unit_of_work)): + pytest.fail("read-only unit of work was accepted for a transaction") + + with pytest.raises(TypeError, match="SQLAlchemy unit of work"): + async with factory.read(existing=_ForeignReadUnitOfWork()): + pytest.fail("foreign read unit of work was accepted") + + with pytest.raises(TypeError, match="transactional unit of work"): + async with factory.transaction(existing=_ForeignTransactionalUnitOfWork()): + pytest.fail("foreign transactional unit of work was accepted")