Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
50 changes: 50 additions & 0 deletions tests/test_connections.py
Original file line number Diff line number Diff line change
Expand Up @@ -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())
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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)
Expand Down
32 changes: 20 additions & 12 deletions ydb_dbapi/connections.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
Loading