"""Test for Neo4jSession decorator with use_v2 flag.""" import importlib from ssl import SSLError from unittest import mock import pytest import connector_neo4j from flask import Flask from neo4j.exceptions import ServiceUnavailable from neo4j.exceptions import SessionExpired app = Flask(__name__) @mock.patch("neo4j.GraphDatabase.driver") def test_read_connection(driver_create): """Test standard read only session. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) with app.test_request_context() as context: verify_connectivity = driver_create.return_value.verify_connectivity session_create = driver_create.return_value.session session = session_create.return_value session._transaction = False session_create.assert_not_called() with connector_neo4j.Neo4jSession(use_v2=True): verify_connectivity.assert_called_once() session_create.assert_called_once_with( default_access_mode="READ", database=None ) assert not connector_neo4j.neo4j_session assert context.g.neo4j_session == session assert connector_neo4j.get_session() == session session.run.assert_not_called() session.close.assert_not_called() session.begin_transaction.assert_not_called() session.close.assert_called_once() session.commit.assert_not_called() session.rollback.assert_not_called() assert context.g.neo4j_session is None @mock.patch("neo4j.GraphDatabase.driver") def test_force_write_connection(driver_create): """Test override non-transaction session to write server. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) with app.test_request_context() as context: verify_connectivity = driver_create.return_value.verify_connectivity session_create = driver_create.return_value.session session = session_create.return_value session._transaction = False session_create.assert_not_called() with connector_neo4j.Neo4jSession(force_write_server=True, use_v2=True): verify_connectivity.assert_called_once() session_create.assert_called_once_with( default_access_mode="WRITE", database=None ) assert not connector_neo4j.neo4j_session assert context.g.neo4j_session == session assert connector_neo4j.get_session() == session session.run.assert_not_called() session.close.assert_not_called() session.begin_transaction.assert_not_called() session.close.assert_called_once() session.commit.assert_not_called() session.rollback.assert_not_called() assert context.g.neo4j_session is None @mock.patch("neo4j.GraphDatabase.driver") def test_commit(driver_create): """Test open and commit transaction. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) with app.test_request_context() as context: verify_connectivity = driver_create.return_value.verify_connectivity session_create = driver_create.return_value.session session = session_create.return_value transaction = session._transaction assert "neo4j_session" not in context.g with connector_neo4j.Neo4jSession(transaction=True, use_v2=True): verify_connectivity.assert_called_once() session_create.assert_called_once_with( default_access_mode="WRITE", database=None ) assert context.g.neo4j_session == session assert connector_neo4j.get_session() == transaction session.run.assert_not_called() session.begin_transaction.assert_called_once() session.close.assert_not_called() transaction.commit.assert_not_called() session.close.assert_called_once() transaction.commit.assert_called() transaction.rollback.assert_not_called() # Verify get_session handles context.g.neo4j_session is None assert context.g.neo4j_session is None with connector_neo4j.Neo4jSession(transaction=True, use_v2=True): pass @mock.patch("neo4j.GraphDatabase.driver") def test_rollback(driver_create): """Test open and rollback transaction. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) with app.test_request_context() as context: verify_connectivity = driver_create.return_value.verify_connectivity session_create = driver_create.return_value.session session = session_create.return_value transaction = session._transaction @connector_neo4j.Neo4jSession(transaction=True, use_v2=True) def func(): verify_connectivity.assert_called_once() session_create.assert_called_once_with( default_access_mode="WRITE", database=None ) assert context.g.neo4j_session == transaction assert connector_neo4j.get_session() == transaction session.run.assert_not_called() session.begin_transaction.assert_called_once() session.close.assert_not_called() transaction.rollback.assert_not_called() raise Exception("something went wrong") try: func() except: # noqa: E722 pass session.close.assert_called_once() transaction.commit.assert_not_called() # Verify get_session handles context.g.neo4j_session is None assert context.g.neo4j_session is None with connector_neo4j.Neo4jSession(transaction=True, use_v2=True): pass @mock.patch("neo4j.GraphDatabase.driver") def test_commit_failure(driver_create): """Test error during transaction commit. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) commit_error_msg = "commit error" with app.test_request_context() as context: session_create = driver_create.return_value.session session = session_create.return_value transaction = session._transaction transaction.commit.side_effect = Exception(commit_error_msg) @connector_neo4j.Neo4jSession(transaction=True, use_v2=True) def func(): pass with pytest.raises(Exception) as exc: func() session.close.assert_called_once() assert str(exc.value) == commit_error_msg assert context.g.neo4j_session is None @mock.patch("neo4j.GraphDatabase.driver") def test_commit_skipped_when_transaction_torn_down(driver_create): """Test exit does not commit when the driver already closed the transaction. A query that fails inside a transaction makes the driver set session._transaction to None. If that error was swallowed upstream, exc_type is None on exit; exit must not call commit() on the missing transaction. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) with app.test_request_context() as context: session_create = driver_create.return_value.session session = session_create.return_value transaction = session._transaction @connector_neo4j.Neo4jSession(transaction=True, use_v2=True) def func(): # simulate the driver tearing down the transaction on a failed query session._transaction = None func() session.close.assert_called_once() transaction.commit.assert_not_called() assert context.g.neo4j_session is None @mock.patch("neo4j.GraphDatabase.driver") def test_create_session_twice(driver_create): """Test creating session inside of session. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) with app.test_request_context(): @connector_neo4j.Neo4jSession(use_v2=True) def func_one(): pass @connector_neo4j.Neo4jSession(use_v2=True) def func_two(): func_one() with pytest.raises(connector_neo4j.exceptions.SessionAlreadyOpen): func_two() @mock.patch("neo4j.GraphDatabase.driver") def test_default_retry(driver_create): """Test exhausting pool according to size.""" importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) verify_connectivity = driver_create.return_value.verify_connectivity session_create = driver_create.return_value.session session = session_create.return_value pool_start_size = 10 driver_create.return_value._pool.address = "localhost" driver_create.return_value._pool.connections = { "localhost": [mock.MagicMock() for x in range(0, pool_start_size)] } # simulate connection failed and removal from pool def mock_session_failed(): connections = driver_create.return_value._pool.connections["localhost"] if connections: connections.pop() raise SessionExpired("test") verify_connectivity.side_effect = mock_session_failed with app.test_request_context(): @connector_neo4j.Neo4jSession(use_v2=True) def func(): pass with pytest.raises(SessionExpired): func() assert verify_connectivity.call_count == pool_start_size + 1 session.run.assert_not_called() @mock.patch("neo4j.GraphDatabase.driver") def test_auto_refresh_session(driver_create): """Test session retry on expiration. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) verify_connectivity = driver_create.return_value.verify_connectivity session_create = driver_create.return_value.session session = session_create.return_value session._transaction = False verify_connectivity.side_effect = [ SessionExpired("test"), ServiceUnavailable, mock.MagicMock(), ] with app.test_request_context() as context: @connector_neo4j.Neo4jSession(use_v2=True) def func(): assert verify_connectivity.call_count == 3 assert context.g.neo4j_session == session assert connector_neo4j.get_session() == session func() @pytest.mark.parametrize( "root_error, final_error_type", [ (OSError(9, "Bad file descriptor"), ServiceUnavailable), (SSLError(), ServiceUnavailable), ( AttributeError("'NoneType' object has no attribute 'complete'"), ServiceUnavailable, ), # noqa:E501 (OSError(9, "Blerg"), OSError), (OSError(9), OSError), ], ) @mock.patch("neo4j.GraphDatabase.driver") def test_verify_connectivity_mapped_exceptions( driver_create, root_error, final_error_type ): """Test when neo4j_driver.verify_connectivity has special exceptions. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) verify_connectivity = driver_create.return_value.verify_connectivity driver_create.return_value._pool.connections = {} verify_connectivity.side_effect = root_error with app.test_request_context(): @connector_neo4j.Neo4jSession(use_v2=True) def func(): pass with pytest.raises(final_error_type): func() @pytest.mark.parametrize( "root_error, final_error_type", [ (OSError(9, "Bad file descriptor"), OSError), (SSLError(), SSLError), ( AttributeError("'NoneType' object has no attribute 'complete'"), AttributeError, ), # noqa:E501 (OSError(9, "Blerg"), OSError), (OSError(9), OSError), ], ) @mock.patch("neo4j.GraphDatabase.driver") def test_session_create_mapped_exceptions(driver_create, root_error, final_error_type): """Test when neo4j_driver.session has special exceptions. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) driver_create.return_value._pool.connections = {} driver_create.return_value.session = mock.MagicMock(side_effect=root_error) with app.test_request_context(): @connector_neo4j.Neo4jSession(use_v2=True) def func(): pass with pytest.raises(final_error_type): func() @pytest.mark.parametrize( "root_error, final_error_type", [ (OSError(9, "Bad file descriptor"), ServiceUnavailable), (SSLError(), ServiceUnavailable), (OSError(9, "Blerg"), OSError), (OSError(9), OSError), ], ) @mock.patch("neo4j.GraphDatabase.driver") def test_session_run_mapped_exceptions(driver_create, root_error, final_error_type): """Test exception handling during context open. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) driver_create.return_value._pool.connections = {} with app.test_request_context(): @connector_neo4j.Neo4jSession() def func(): raise root_error with pytest.raises(final_error_type): func() @pytest.mark.parametrize( "root_error, exc_line, exc_file, final_error_type", [ ( TypeError, 980, "/var/venv/lib64/python3.8/site-packages/neo4j/io/__init__.py", TypeError, ), ( TypeError, 980, "/var/venv/lib32/python3.8/site-packages/neo4j/io/__init__.py", TypeError, ), ( AttributeError, 980, "/var/venv/lib64/python3.8/site-packages/neo4j/io/__init__.py", AttributeError, ), (TypeError, 980, "__init__.py", TypeError), ( TypeError, 777, "/var/venv/lib64/python3.8/site-packages/neo4j/io/__init__.py", TypeError, ), ( TypeError, 980, "/var/venv/lib64/python3.8/site-packages/sql/io/__init__.py", TypeError, ), ], ) @mock.patch("sys.exc_info") @mock.patch("neo4j.GraphDatabase.driver") def test_session_run_traceback_exception( driver_create, mock_exc_info, root_error, exc_line, exc_file, final_error_type ): """Test extra traceback exception handling does not run for session v2. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) driver_create.return_value._pool.connections = {} mock_file_info = mock.MagicMock() mock_file_info.tb_lineno = exc_line mock_file_info.tb_frame.f_code.co_filename = exc_file mock_exc_info.return_value = (None, None, mock_file_info) with app.test_request_context(): @connector_neo4j.Neo4jSession(use_v2=True) def func(): raise root_error with pytest.raises(final_error_type): func() @mock.patch("neo4j.GraphDatabase.driver") def test_session_create_fail(driver_create): """Test session failure on create. Args: driver_create (MagicMock): mock of neo4j driver constructor """ importlib.reload(connector_neo4j) num_retries = 5 connector_neo4j.configure( "url", "username", "password", connector_neo4j.SessionStorageMode.FLASK ) verify_connectivity = driver_create.return_value.verify_connectivity session_create = driver_create.return_value.session session = session_create.return_value verify_connectivity.side_effect = [ SessionExpired("test") for _ in range(0, num_retries + 1) ] session._pool.connections = {} with app.test_request_context(): @connector_neo4j.Neo4jSession(use_v2=True) def func(): pass with pytest.raises(SessionExpired): func() assert verify_connectivity.call_count == num_retries + 1