# mypy: ignore-errors """Tests for SQL connector.""" import os from unittest.mock import AsyncMock import pytest import pytest_asyncio from sqlalchemy.sql import text from delivery_metadata.config import AR_DB_URL from delivery_metadata.connectors.sql import SqlConnector @pytest_asyncio.fixture def sql_connector(): return SqlConnector(AR_DB_URL) @pytest.mark.asyncio async def test_sql_connector_initialization(sql_connector): assert str(sql_connector._db_engine.url) == ( "mysql+aiomysql://{user}:{password}@{host}/{db_name}?charset={charset}".format( user=os.environ.get("ART_RELATIONS_DB_USER"), password="***", host=os.environ.get("ART_RELATIONS_DB_HOST"), db_name=os.environ.get("ART_RELATIONS_DB_DATABASE"), charset="utf8mb4", ) ) @pytest.mark.asyncio(loop_scope="session") async def test_sql_insert_transaction(sql_connector): async with sql_connector.db_session(transaction=True) as session: await session.execute(text("DROP TABLE IF EXISTS widgets")) await session.execute(text("CREATE TABLE widgets (id int(10))")) with pytest.raises(Exception, match="Test error"): async with sql_connector.db_session(transaction=True) as session: await session.execute(text("INSERT INTO widgets (id) VALUES (1)")) raise Exception("Test error") async with sql_connector.db_session() as session: result = await session.execute(text("SELECT * FROM widgets")) assert len(result.mappings().all()) == 0 @pytest.mark.asyncio(loop_scope="session") async def test_sql_default_connector_db_session(sql_connector, mocker): mocker.patch.object(sql_connector, "_db_session_maker", return_value=AsyncMock()) async with sql_connector.db_session() as session: assert session == sql_connector._db_session_maker.return_value assert not session.commit.called assert not session.rollback.called assert session.close.called @pytest.mark.asyncio(loop_scope="session") async def test_sql_default_connector_db_session_exception(sql_connector, mocker): mocker.patch.object(sql_connector, "_db_session_maker", return_value=AsyncMock()) with pytest.raises(Exception, match="Test error"): async with sql_connector.db_session() as session: raise Exception("Test error") assert not session.commit.called assert not session.rollback.called assert session.close.called @pytest.mark.asyncio(loop_scope="session") async def test_sql_transaction_connector_db_session(sql_connector, mocker): mocker.patch.object(sql_connector, "_db_session_maker", return_value=AsyncMock()) async with sql_connector.db_session(transaction=True) as session: assert session == sql_connector._db_session_maker.return_value assert session.commit.called assert not session.rollback.called assert session.close.called @pytest.mark.asyncio(loop_scope="session") async def test_sql_transaction_connector_db_session_exception(sql_connector, mocker): mocker.patch.object(sql_connector, "_db_session_maker", return_value=AsyncMock()) with pytest.raises(Exception, match="Test rollback"): async with sql_connector.db_session(transaction=True) as session: raise Exception("Test rollback") assert not session.commit.called assert session.rollback.called assert session.close.called @pytest.mark.asyncio(loop_scope="session") async def test_sql_query_and_close_exception(sql_connector, mocker): mock_session = AsyncMock() mock_session.close.side_effect = Exception("Close error") mocker.patch.object(sql_connector, "_db_session_maker", return_value=mock_session) with pytest.raises(Exception, match="Test error"): async with sql_connector.db_session(transaction=True) as session: raise Exception("Test error") assert not session.commit.called assert session.rollback.called assert session.close.called @pytest.mark.asyncio(loop_scope="session") async def test_sql_query_and_no_close_exception(sql_connector, mocker): mock_session = AsyncMock() mock_session.close.side_effect = Exception("Close error") mocker.patch.object(sql_connector, "_db_session_maker", return_value=mock_session) with pytest.raises(Exception, match="Close error"): async with sql_connector.db_session(transaction=True) as session: pass assert session.commit.called assert not session.rollback.called assert session.close.called @pytest.mark.asyncio(loop_scope="session") async def test_sql_connector_close(sql_connector, mocker): sql_connector._db_engine = mocker.AsyncMock() await sql_connector.close() sql_connector._db_engine.dispose.assert_called_once()