diff --git a/tests/test_connections.py b/tests/test_connections.py index d87f9ed..029a064 100644 --- a/tests/test_connections.py +++ b/tests/test_connections.py @@ -66,6 +66,43 @@ def _test_commit_rollback_after_begin( maybe_await(connection.rollback()) + def _test_close_releases_session_after_rollback_error( + self, + connection: dbapi.Connection, + ) -> None: + connection.set_isolation_level(dbapi.IsolationLevel.SERIALIZABLE) + maybe_await(connection.begin()) + + acquired_session = connection._session + assert acquired_session is not None + + released = [] + session_pool = connection._session_pool + tx_context = connection._tx_context + original_release = session_pool.release + original_rollback = tx_context.rollback + + def release(session: ydb.QuerySession) -> any: + released.append(session) + return original_release(session) + + def failing_rollback(*args: any, **kwargs: any) -> None: + raise ydb.issues.BadSession("session is invalidated") + + session_pool.release = release + tx_context.rollback = failing_rollback + + try: + with pytest.raises(dbapi.Error): + maybe_await(connection.close()) + + assert released == [acquired_session] + assert connection._session is None + assert connection._tx_context is None + finally: + session_pool.release = original_release + tx_context.rollback = original_rollback + def _test_connection(self, connection: dbapi.Connection) -> None: maybe_await(connection.commit()) maybe_await(connection.rollback()) @@ -466,6 +503,11 @@ def test_commit_rollback_after_begin( connection, isolation_level ) + def test_close_releases_session_after_rollback_error( + self, connection: dbapi.Connection + ) -> None: + self._test_close_releases_session_after_rollback_error(connection) + def test_connection(self, connection: dbapi.Connection) -> None: self._test_connection(connection) @@ -580,6 +622,14 @@ async def test_commit_rollback_after_begin( isolation_level ) + @pytest.mark.asyncio + async def test_close_releases_session_after_rollback_error( + self, connection: dbapi.AsyncConnection + ) -> None: + await greenlet_spawn( + self._test_close_releases_session_after_rollback_error, connection + ) + @pytest.mark.asyncio async def test_connection(self, connection: dbapi.AsyncConnection) -> None: await greenlet_spawn(self._test_connection, connection) diff --git a/ydb_dbapi/connections.py b/ydb_dbapi/connections.py index cc7a117..285fd1b 100644 --- a/ydb_dbapi/connections.py +++ b/ydb_dbapi/connections.py @@ -291,14 +291,18 @@ def rollback(self) -> None: @handle_ydb_errors def close(self) -> None: - self.rollback() + try: + self.rollback() + finally: + self._tx_context = None - if self._session: - self._session_pool.release(self._session) + if self._session: + self._session_pool.release(self._session) + self._session = None - if not self._shared_session_pool: - self._session_pool.stop() - self._driver.stop() + if not self._shared_session_pool: + self._session_pool.stop() + self._driver.stop() @handle_ydb_errors def describe(self, table_path: str) -> ydb.TableSchemeEntry: @@ -489,14 +493,18 @@ async def rollback(self) -> None: @handle_ydb_errors async def close(self) -> None: - await self.rollback() + try: + await self.rollback() + finally: + self._tx_context = None - if self._session: - await self._session_pool.release(self._session) + if self._session: + await self._session_pool.release(self._session) + self._session = None - if not self._shared_session_pool: - await self._session_pool.stop() - await self._driver.stop() + if not self._shared_session_pool: + await self._session_pool.stop() + await self._driver.stop() @handle_ydb_errors async def describe(self, table_path: str) -> ydb.TableSchemeEntry: