"""Test for neo4j DB connector.""" from unittest import mock import pytest from flask import g from permissions import config from permissions.connectors import neo4j from permissions.constants import constants def _get_mock_driver(): """Return a mock object for drivers.""" mock_session = mock.MagicMock() # have to mock the whole driver to avoid TypeError: 'Mock' object is not iterable. mock_session.run.return_value = [ '1', ] mock_driver = mock.MagicMock() mock_driver.session.return_value = mock_session return mock_driver def test_db_read_session_success(context): """Test db_read_session when query exec is success.""" with context: neo4j.neo4j_driver = _get_mock_driver() neo4j.aura_driver = _get_mock_driver() with neo4j.db_session() as session1: session1.run('MATCH (n) RETURN n') assert not session1.close.called with neo4j.db_session() as session2: session2.run('MATCH (n) RETURN n') # auto commit happens automatically at the driver level. assert not session2.close.called assert not session2.commit.called assert not session2.rollback.called assert session2.run.called # db_session() returns same session object. assert session1 == session2 # at the end of app context, session would close. assert session2.close.call_count == 1 def test_db_read_session_creates_new_when_write_exists(context): """Test that read session does not re-use existing write session.""" with context: mock_driver = _get_mock_driver() # Configure to return different sessions for each call mock_driver.session.side_effect = [ mock.Mock(name='read_session'), mock.Mock(name='write_session'), ] neo4j.neo4j_driver = mock_driver neo4j.aura_driver = mock_driver # First create a write session with neo4j.db_session(access_mode=constants.NEO4j_WRITE_ACCESS) as write_session: pass assert not write_session.close.called # Then request a read session with neo4j.db_session(access_mode=constants.NEO4j_READ_ACCESS) as read_session: pass # Verify we got separate sessions assert write_session != read_session # at the end of app context, session should close once assert read_session.close.call_count == 1 def test_db_write_session_creates_new_when_read_exists(context): """Test that write session creates new session when only read session exists.""" with context: mock_driver = _get_mock_driver() # Configure to return different sessions for each call mock_driver.session.side_effect = [ mock.Mock(name='read_session'), mock.Mock(name='write_session'), ] neo4j.neo4j_driver = mock_driver neo4j.aura_driver = mock_driver # First create a read session with neo4j.db_session(access_mode=constants.NEO4j_READ_ACCESS) as read_session: pass assert not read_session.close.called # Then request a write session - should get new one with neo4j.db_session(access_mode=constants.NEO4j_WRITE_ACCESS) as write_session: pass assert not write_session.close.called assert not write_session.commit.called assert not write_session.rollback.called # Verify we got different sessions assert write_session != read_session # at the end of app context, both sessions should close assert read_session.close.call_count == 1 assert write_session.close.call_count == 1 def test_close_db(context): """Test close_db successfully closes all sessions.""" with context: g.neo4j_sessions = { constants.NEO4j_READ_ACCESS: mock.Mock(close=mock.Mock()), constants.NEO4j_WRITE_ACCESS: mock.Mock(close=mock.Mock()), } neo4j.close_db(None) # Both sessions should be closed assert g.neo4j_sessions[constants.NEO4j_READ_ACCESS].close.call_count == 1 assert g.neo4j_sessions[constants.NEO4j_WRITE_ACCESS].close.call_count == 1 def test_close_db_no_session(context): """Test close_db do nothing when there is no open session.""" with context: neo4j.close_db(None) assert not hasattr(g, 'neo4j_db') @pytest.mark.parametrize( 'access_mode', [ None, constants.NEO4j_READ_ACCESS, constants.NEO4j_WRITE_ACCESS, ], ) def test_db_session_access_mode(context, access_mode): """Test db_session.""" with context: mock_driver = _get_mock_driver() neo4j.neo4j_driver = mock_driver neo4j.aura_driver = mock_driver with neo4j.db_session(access_mode=access_mode): # run query 1 pass mock_driver.session.assert_called_with( default_access_mode=access_mode, database=config.NEO4J_DATABASE_NAME )