"""Test database utilities.""" import sqlalchemy from collaborator.utils import db def test_create_engine_queue_pool(mocker): """Test creating a DB engine using the QueuePool and associated params.""" mocker.patch.multiple( db.config, POOL_CLASS=sqlalchemy.pool.QueuePool, POOL_MAX_OVERFLOW=1, POOL_PRE_PING=2, POOL_RECYCLE_MS=3, POOL_SIZE=4, ) mock_create_engine = mocker.patch.object(db.sqlalchemy, "create_engine") db.create_engine("testurl") mock_create_engine.assert_called_with( "testurl", max_overflow=1, pool_pre_ping=2, pool_recycle=3, pool_size=4 ) def test_create_engine_static_pool(mocker): """Test creating a DB engine using the StaticPool.""" mocker.patch.object(db.config, "POOL_CLASS", sqlalchemy.pool.StaticPool) mock_create_engine = mocker.patch.object(db.sqlalchemy, "create_engine") db.create_engine("testurl") mock_create_engine.assert_called_with( "testurl", poolclass=sqlalchemy.pool.StaticPool ) def test_create_session(mocker): """Test creating a session.""" session_mock = mocker.patch("sqlalchemy.orm.session.Session") session_maker_stub = mocker.stub() session_maker_stub.return_value = session_mock with db.create_session(session_maker_stub): pass session_maker_stub.assert_called_once() session_mock.commit.assert_called_once() session_mock.close.assert_called_once() session_mock.rollback.assert_not_called() def test_create_session_with_error(mocker): """Test creating a session which encounters an error.""" session_mock = mocker.patch("sqlalchemy.orm.session.Session") session_maker_stub = mocker.stub() session_maker_stub.return_value = session_mock try: with db.create_session(session_maker_stub): raise Exception("Test error") except Exception: pass session_maker_stub.assert_called_once() session_mock.rollback.assert_called_once() session_mock.close.assert_called_once() session_mock.commit.assert_not_called()