"""Tests for snowflake_proxy fetch functions. These tests verify the streaming, cursor, and chunked fetch methods added for INT-2638 to prevent OOM issues with large aggregate queries. """ import pytest from unittest.mock import MagicMock, patch class MockResult: """Mock SQLAlchemy result object.""" def __init__(self, rows): self._rows = list(rows) self._index = 0 def fetchmany(self, size): chunk = self._rows[self._index:self._index + size] self._index += size return chunk def fetchall(self): return self._rows def mappings(self): return self._rows def __iter__(self): return iter(self._rows) class MockSession: """Mock SQLAlchemy session.""" def __init__(self, result): self._result = result self.closed = False def execute(self, sql, params=None): return self._result def close(self): self.closed = True @pytest.fixture def mock_sessionmaker(): """Fixture that returns a session factory mock.""" def _factory(rows): result = MockResult(rows) session = MockSession(result) maker = MagicMock(return_value=session) return maker, session return _factory @pytest.fixture def sample_rows(): """Sample rows for testing.""" return [ {"id": 1, "name": "row1"}, {"id": 2, "name": "row2"}, {"id": 3, "name": "row3"}, {"id": 4, "name": "row4"}, {"id": 5, "name": "row5"}, ] @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') def test_fetchproxy_stream_yields_result(mock_get_sessionmaker, mock_validator): """Test fetchproxy_stream yields result and closes session.""" from snowflake_proxy.snowflake_proxy import fetchproxy_stream rows = [{"id": 1}, {"id": 2}] result = MockResult(rows) session = MockSession(result) mock_get_sessionmaker.return_value = MagicMock(return_value=session) mock_validator.format_identifiers.return_value = ("SELECT *", {}) sql = "SELECT * FROM table" params = {"param": "value"} with fetchproxy_stream(sql, params) as res: fetched = list(res) assert fetched == rows assert session.closed @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') def test_fetchproxy_stream_without_params(mock_get_sessionmaker, mock_validator): """Test fetchproxy_stream works without params.""" from snowflake_proxy.snowflake_proxy import fetchproxy_stream rows = [{"id": 1}] result = MockResult(rows) session = MockSession(result) mock_get_sessionmaker.return_value = MagicMock(return_value=session) sql = "SELECT * FROM table" with fetchproxy_stream(sql) as res: fetched = list(res) assert fetched == rows # validator.format_identifiers should not be called without params mock_validator.format_identifiers.assert_not_called() @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') def test_fetchproxy_cursor_yields_chunks(mock_get_sessionmaker, mock_validator, sample_rows): """Test fetchproxy_cursor yields results in chunks.""" from snowflake_proxy.snowflake_proxy import fetchproxy_cursor result = MockResult(sample_rows) session = MockSession(result) mock_get_sessionmaker.return_value = MagicMock(return_value=session) mock_validator.format_identifiers.return_value = ("SELECT *", {}) sql = "SELECT * FROM table" params = {"param": "value"} chunks = list(fetchproxy_cursor(sql, params, chunk_size=2)) # Should yield 3 chunks: [2 rows], [2 rows], [1 row] assert len(chunks) == 3 assert len(chunks[0]) == 2 assert len(chunks[1]) == 2 assert len(chunks[2]) == 1 assert session.closed @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') def test_fetchproxy_cursor_empty_result(mock_get_sessionmaker, mock_validator): """Test fetchproxy_cursor handles empty results.""" from snowflake_proxy.snowflake_proxy import fetchproxy_cursor result = MockResult([]) session = MockSession(result) mock_get_sessionmaker.return_value = MagicMock(return_value=session) sql = "SELECT * FROM table" chunks = list(fetchproxy_cursor(sql, chunk_size=10)) assert len(chunks) == 0 assert session.closed @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') def test_fetchproxy_chunked_yields_paginated(mock_get_sessionmaker, mock_validator): """Test fetchproxy_chunked yields paginated results.""" from snowflake_proxy.snowflake_proxy import fetchproxy_chunked # Simulate pagination - each call returns a different page page1 = [{"id": 1}, {"id": 2}] page2 = [{"id": 3}] call_count = [0] sessions = [] def create_session(): idx = call_count[0] call_count[0] += 1 if idx == 0: result = MockResult(page1) elif idx == 1: result = MockResult(page2) else: result = MockResult([]) session = MockSession(result) sessions.append(session) return session mock_get_sessionmaker.return_value = create_session sql = "SELECT * FROM table ORDER BY id" chunks = list(fetchproxy_chunked(sql, chunk_size=2)) # Should get 2 chunks before empty result stops iteration assert len(chunks) == 2 assert chunks[0] == page1 assert chunks[1] == page2 # All sessions should be closed for s in sessions: assert s.closed @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') def test_fetchproxy_chunked_empty_first_page(mock_get_sessionmaker, mock_validator): """Test fetchproxy_chunked handles empty first page.""" from snowflake_proxy.snowflake_proxy import fetchproxy_chunked result = MockResult([]) session = MockSession(result) mock_get_sessionmaker.return_value = MagicMock(return_value=session) sql = "SELECT * FROM table ORDER BY id" chunks = list(fetchproxy_chunked(sql, chunk_size=10)) assert len(chunks) == 0 assert session.closed @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') def test_fetchproxy_returns_mappings(mock_get_sessionmaker, mock_validator): """Test fetchproxy returns mappings for dict-like access.""" from snowflake_proxy.snowflake_proxy import fetchproxy rows = [{"id": 1, "name": "test"}] result = MockResult(rows) session = MockSession(result) mock_get_sessionmaker.return_value = MagicMock(return_value=session) sql = "SELECT * FROM table" res = fetchproxy(sql) assert res == rows assert session.closed @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') def test_fetchproxy_with_params_calls_validator(mock_get_sessionmaker, mock_validator): """Test fetchproxy calls validator.format_identifiers with params.""" from snowflake_proxy.snowflake_proxy import fetchproxy rows = [{"id": 1}] result = MockResult(rows) session = MockSession(result) mock_get_sessionmaker.return_value = MagicMock(return_value=session) mock_validator.format_identifiers.return_value = ("SELECT * FROM tbl", {"p": "v"}) sql = "SELECT * FROM :table_name:" params = {"table_name": "tbl", "p": "v"} fetchproxy(sql, params) mock_validator.format_identifiers.assert_called_once_with(sql, params) @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') @patch('snowflake_proxy.snowflake_proxy.logger') def test_fetchproxy_retries_on_database_error( mock_logger, mock_get_sessionmaker, mock_validator ): """Test fetchproxy retries once on DatabaseError.""" from snowflake_proxy.snowflake_proxy import fetchproxy, DBAPIError, DatabaseError rows = [{"id": 1}] result = MockResult(rows) call_count = [0] sessions = [] def create_session(): idx = call_count[0] call_count[0] += 1 session = MagicMock() sessions.append(session) if idx == 0: # First call raises DatabaseError db_err = DatabaseError("Connection lost") session.execute.side_effect = DBAPIError( "test", {}, db_err, connection_invalidated=True ) session.execute.side_effect.orig = db_err else: # Second call succeeds session.execute.return_value = result return session mock_get_sessionmaker.return_value = create_session sql = "SELECT * FROM table" res = fetchproxy(sql) assert res == rows assert call_count[0] == 2 mock_logger.warning.assert_called_once() @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') def test_fetchproxy_raises_non_database_error(mock_get_sessionmaker, mock_validator): """Test fetchproxy raises non-DatabaseError exceptions.""" from snowflake_proxy.snowflake_proxy import fetchproxy, DBAPIError session = MagicMock() err = ValueError("Some other error") session.execute.side_effect = DBAPIError("test", {}, err) session.execute.side_effect.orig = err mock_get_sessionmaker.return_value = MagicMock(return_value=session) sql = "SELECT * FROM table" with pytest.raises(DBAPIError): fetchproxy(sql) @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') @patch('snowflake_proxy.snowflake_proxy.logger') def test_fetchproxy_stream_retries_on_database_error( mock_logger, mock_get_sessionmaker, mock_validator ): """Test fetchproxy_stream retries once on DatabaseError.""" from snowflake_proxy.snowflake_proxy import fetchproxy_stream, DBAPIError, DatabaseError rows = [{"id": 1}] result = MockResult(rows) call_count = [0] sessions = [] def create_session(): idx = call_count[0] call_count[0] += 1 session = MagicMock() session.closed = False session.close = MagicMock(side_effect=lambda: setattr(session, 'closed', True)) sessions.append(session) if idx == 0: db_err = DatabaseError("Connection lost") session.execute.side_effect = DBAPIError( "test", {}, db_err, connection_invalidated=True ) session.execute.side_effect.orig = db_err else: session.execute.return_value = result return session mock_get_sessionmaker.return_value = create_session sql = "SELECT * FROM table" with fetchproxy_stream(sql) as res: fetched = list(res) assert fetched == rows assert call_count[0] == 2 mock_logger.warning.assert_called_once() @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') def test_fetchproxy_cursor_large_chunk_single_fetch(mock_get_sessionmaker, mock_validator): """Test fetchproxy_cursor fetches all in one chunk when chunk_size > rows.""" from snowflake_proxy.snowflake_proxy import fetchproxy_cursor rows = [{"id": 1}, {"id": 2}] result = MockResult(rows) session = MockSession(result) mock_get_sessionmaker.return_value = MagicMock(return_value=session) sql = "SELECT * FROM table" chunks = list(fetchproxy_cursor(sql, chunk_size=100)) assert len(chunks) == 1 assert len(chunks[0]) == 2 @patch('snowflake_proxy.snowflake_proxy.validator') @patch('snowflake_proxy.snowflake_proxy._get_sessionmaker') def test_fetchproxy_chunked_exact_chunk_boundary(mock_get_sessionmaker, mock_validator): """Test fetchproxy_chunked handles exact chunk boundaries.""" from snowflake_proxy.snowflake_proxy import fetchproxy_chunked page1 = [{"id": 1}, {"id": 2}] page2 = [] # Empty page signals end call_count = [0] def create_session(): idx = call_count[0] call_count[0] += 1 if idx == 0: result = MockResult(page1) else: result = MockResult(page2) session = MockSession(result) return session mock_get_sessionmaker.return_value = create_session sql = "SELECT * FROM table ORDER BY id" chunks = list(fetchproxy_chunked(sql, chunk_size=2)) # When first chunk is exactly chunk_size, it makes another call # which returns empty, so we get 1 chunk assert len(chunks) == 1 assert chunks[0] == page1