"""Unit tests for src/assert_snowflake_sync/snowflake_sync.py.""" import datetime as dt import decimal from unittest.mock import MagicMock import pytest import config from src.assert_snowflake_sync.snowflake_sync import ( _build_field_specs, _coerce_to_datetime, _is_empty, _row_value, build_field_specs, build_neo4j_query, build_snowflake_query, build_window, check_window_config, compare_row_maps, fetch_neo4j_rows, fetch_snowflake_rows, format_value, normalize_value, run_row_level_check, values_equal, ) class TestBuildWindow: """Tests for build_window.""" def test_returns_tuple_of_two_utc_datetimes(self): """Both returned values are UTC-aware datetime objects.""" window_start, window_end = build_window() assert window_start.tzinfo == dt.timezone.utc assert window_end.tzinfo == dt.timezone.utc def test_window_gap_matches_config(self): """Gap between window_start and window_end matches WINDOW_START_HOURS - WINDOW_END_HOURS.""" window_start, window_end = build_window() expected = dt.timedelta(hours=config.WINDOW_START_HOURS - config.WINDOW_END_HOURS) assert window_end - window_start == expected class TestNormalizeValue: """Tests for normalize_value.""" def test_none_returns_none(self): """None input returns None.""" assert normalize_value(None) is None def test_empty_dict_returns_none(self): """Empty dict is treated as empty and returns None.""" assert normalize_value({}) is None def test_empty_list_returns_none(self): """Empty list is treated as empty and returns None.""" assert normalize_value([]) is None def test_empty_tuple_returns_none(self): """Empty tuple is treated as empty and returns None.""" assert normalize_value(()) is None def test_non_empty_container_returns_as_is(self): """Non-empty list is returned unchanged.""" assert normalize_value([1, 2]) == [1, 2] def test_naive_datetime_gets_utc_timezone(self): """Naive datetime has UTC timezone attached.""" naive = dt.datetime(2024, 1, 1, 12, 0, 0) result = normalize_value(naive) assert result.tzinfo == dt.timezone.utc def test_aware_datetime_unchanged(self): """Timezone-aware datetime is returned as-is.""" aware = dt.datetime(2024, 1, 1, 12, 0, 0, tzinfo=dt.timezone.utc) assert normalize_value(aware) == aware def test_date_unchanged(self): """date object is returned unchanged.""" d = dt.date(2024, 3, 15) assert normalize_value(d) == d def test_whole_decimal_returns_int(self): """Whole-number Decimal is returned as int.""" result = normalize_value(decimal.Decimal('5')) assert result == 5 assert isinstance(result, int) def test_fractional_decimal_returns_float(self): """Fractional Decimal is returned as float.""" result = normalize_value(decimal.Decimal('5.5')) assert result == 5.5 assert isinstance(result, float) def test_to_native_called_on_neo4j_type(self): """to_native() is called when available and result is used.""" neo4j_obj = MagicMock() neo4j_obj.to_native.return_value = 'native_value' result = normalize_value(neo4j_obj) neo4j_obj.to_native.assert_called_once() assert result == 'native_value' def test_to_native_exception_falls_through(self): """If to_native() raises, the original object is returned unchanged.""" neo4j_obj = MagicMock() neo4j_obj.to_native.side_effect = RuntimeError('fail') result = normalize_value(neo4j_obj) assert result is neo4j_obj def test_plain_string_unchanged(self): """Plain string is returned unchanged.""" assert normalize_value('hello') == 'hello' class TestPrivateHelpers: """Tests for private helper functions _coerce_to_datetime and _is_empty.""" # --- _coerce_to_datetime --- def test_coerce_none_returns_none(self): """_coerce_to_datetime(None) returns None.""" assert _coerce_to_datetime(None) is None def test_coerce_naive_datetime_gets_utc(self): """Naive datetime gets UTC timezone.""" naive = dt.datetime(2024, 6, 1, 8, 0, 0) result = _coerce_to_datetime(naive) assert result.tzinfo == dt.timezone.utc assert result == dt.datetime(2024, 6, 1, 8, 0, 0, tzinfo=dt.timezone.utc) def test_coerce_aware_datetime_unchanged(self): """Aware datetime is returned unchanged.""" aware = dt.datetime(2024, 6, 1, 8, 0, 0, tzinfo=dt.timezone.utc) assert _coerce_to_datetime(aware) == aware def test_coerce_date_converts_to_midnight_utc(self): """date is converted to midnight UTC datetime.""" d = dt.date(2024, 3, 20) result = _coerce_to_datetime(d) assert result == dt.datetime(2024, 3, 20, 0, 0, 0, tzinfo=dt.timezone.utc) def test_coerce_non_date_type_returns_none(self): """Non-date type returns None.""" assert _coerce_to_datetime('2024-01-01') is None assert _coerce_to_datetime(42) is None # --- _is_empty --- def test_none_is_empty(self): """None is considered empty.""" assert _is_empty(None) is True def test_empty_string_is_empty(self): """Empty string is considered empty.""" assert _is_empty('') is True def test_zero_is_not_empty(self): """Zero is not empty.""" assert _is_empty(0) is False def test_false_is_not_empty(self): """False is not empty.""" assert _is_empty(False) is False def test_non_empty_string_is_not_empty(self): """Non-empty string is not empty.""" assert _is_empty('hello') is False class TestValuesEqual: """Tests for values_equal.""" def test_both_none_equal(self): """Two None values are equal.""" assert values_equal(None, None) is True def test_both_empty_string_equal(self): """Two empty strings are equal.""" assert values_equal('', '') is True def test_none_and_empty_string_equal(self): """None and empty string are equal (both treated as empty).""" assert values_equal(None, '') is True assert values_equal('', None) is True def test_none_and_value_not_equal(self): """None and a non-empty value are not equal.""" assert values_equal(None, 'foo') is False def test_identical_strings_equal(self): """Identical strings are equal.""" assert values_equal('abc', 'abc') is True def test_different_strings_not_equal(self): """Different strings are not equal.""" assert values_equal('abc', 'xyz') is False def test_date_and_matching_datetime_equal(self): """A date and its corresponding midnight UTC datetime are equal.""" d = dt.date(2024, 5, 10) dtime = dt.datetime(2024, 5, 10, 0, 0, 0, tzinfo=dt.timezone.utc) assert values_equal(d, dtime) is True def test_date_and_different_datetime_not_equal(self): """A date and a non-matching datetime are not equal.""" d = dt.date(2024, 5, 10) dtime = dt.datetime(2024, 5, 11, 0, 0, 0, tzinfo=dt.timezone.utc) assert values_equal(d, dtime) is False class TestFormatValue: """Tests for format_value.""" def test_none_returns_none(self): """None returns None.""" assert format_value(None) is None def test_datetime_returns_iso_string(self): """datetime returns ISO 8601 string.""" dtime = dt.datetime(2024, 6, 15, 10, 30, 0, tzinfo=dt.timezone.utc) assert format_value(dtime) == dtime.isoformat() def test_date_returns_iso_string(self): """date returns ISO 8601 string.""" d = dt.date(2024, 6, 15) assert format_value(d) == '2024-06-15' def test_string_unchanged(self): """Plain string is returned unchanged.""" assert format_value('hello') == 'hello' def test_int_unchanged(self): """Integer is returned unchanged.""" assert format_value(42) == 42 class TestBuildFieldSpecs: """Tests for _build_field_specs and build_field_specs.""" def test_key_prefix_applied(self): """Key field aliases start with 'k'.""" fields = [{'neo4j': 'n.id', 'snowflake': 'ID'}] specs = _build_field_specs(fields, 'k') assert specs[0]['alias'] == 'k0' def test_compare_prefix_applied(self): """Compare field aliases start with 'f'.""" fields = [{'neo4j': 'n.name', 'snowflake': 'NAME'}] specs = _build_field_specs(fields, 'f') assert specs[0]['alias'] == 'f0' def test_label_defaults_to_snowflake_column(self): """Label defaults to the Snowflake column name when not specified.""" fields = [{'neo4j': 'n.id', 'snowflake': 'MY_COL'}] specs = _build_field_specs(fields, 'k') assert specs[0]['label'] == 'MY_COL' def test_explicit_label_used(self): """Explicit label is used when provided.""" fields = [{'neo4j': 'n.id', 'snowflake': 'MY_COL', 'label': 'Custom Label'}] specs = _build_field_specs(fields, 'k') assert specs[0]['label'] == 'Custom Label' def test_missing_key_fields_raises(self): """build_field_specs raises ValueError when key_fields is missing.""" with pytest.raises(ValueError): build_field_specs({'name': 'no_keys'}) def test_empty_key_fields_raises(self): """build_field_specs raises ValueError when key_fields is empty.""" with pytest.raises(ValueError): build_field_specs({'name': 'empty_keys', 'key_fields': []}) def test_no_compare_fields_returns_empty_list(self): """build_field_specs returns empty compare_specs when compare_fields is absent.""" check = {'name': 'x', 'key_fields': [{'neo4j': 'n.id', 'snowflake': 'ID'}]} _, compare_specs = build_field_specs(check) assert compare_specs == [] class TestCheckWindowConfig: """Tests for check_window_config.""" def test_both_window_fields_returns_true(self): """Returns True when both window fields are defined.""" check = { 'window_property': 'n.updatedAt', 'snowflake_window_column': 'UPDATED_AT', } assert check_window_config(check) is True def test_neither_window_field_returns_false(self): """Returns False when neither window field is defined.""" assert check_window_config({}) is False def test_only_neo4j_window_raises(self): """Raises ValueError when only window_property is defined.""" with pytest.raises(ValueError): check_window_config({'window_property': 'n.updatedAt'}) def test_only_snowflake_window_raises(self): """Raises ValueError when only snowflake_window_column is defined.""" with pytest.raises(ValueError): check_window_config({'snowflake_window_column': 'UPDATED_AT'}) class TestBuildNeo4jQuery: """Tests for build_neo4j_query.""" def test_no_window_omits_where_clause(self): """Query has no WHERE clause when use_window is False.""" check = {'cypher_pattern': '(n:Node)'} specs = [{'alias': 'k0', 'neo4j': 'n.id'}] query = build_neo4j_query(check, specs, use_window=False) assert 'WHERE' not in query assert query == 'MATCH (n:Node) RETURN n.id AS k0' def test_window_adds_where_clause_with_params(self): """Query includes WHERE clause with $window_start and $window_end params.""" check = {'cypher_pattern': '(n:Node)', 'window_property': 'n.updatedAt'} specs = [{'alias': 'k0', 'neo4j': 'n.id'}] query = build_neo4j_query(check, specs, use_window=True) assert 'WHERE' in query assert '$window_start' in query assert '$window_end' in query def test_field_aliases_in_return_clause(self): """All field aliases appear in the RETURN clause.""" check = {'cypher_pattern': '(n:Node)'} specs = [ {'alias': 'k0', 'neo4j': 'n.id'}, {'alias': 'f0', 'neo4j': 'n.name'}, ] query = build_neo4j_query(check, specs, use_window=False) assert 'n.id AS k0' in query assert 'n.name AS f0' in query class TestBuildSnowflakeQuery: """Tests for build_snowflake_query.""" def test_fully_qualified_table_used_as_is(self): """Table with two dots is used as-is without adding database/schema prefix.""" check = {'snowflake_table': 'MY_DB.MY_SCHEMA.MY_TABLE'} specs = [{'alias': 'k0', 'snowflake': 'ID'}] query = build_snowflake_query(check, specs, use_window=False) assert 'MY_DB.MY_SCHEMA.MY_TABLE' in query assert query.count('MY_DB') == 1 def test_unqualified_table_gets_db_schema_prefix(self, mocker): """Unqualified table name gets SNOWFLAKE_DATABASE.SNOWFLAKE_SCHEMA prefix.""" mock_config = mocker.patch('src.assert_snowflake_sync.snowflake_sync.config') mock_config.SNOWFLAKE_DATABASE = 'DB' mock_config.SNOWFLAKE_SCHEMA = 'SCH' check = {'snowflake_table': 'RAW_TABLE'} specs = [{'alias': 'k0', 'snowflake': 'ID'}] query = build_snowflake_query(check, specs, use_window=False) assert 'DB.SCH.RAW_TABLE' in query def test_no_window_omits_where(self): """Query has no WHERE clause when use_window is False.""" check = {'snowflake_table': 'DB.SCH.TBL'} specs = [{'alias': 'k0', 'snowflake': 'ID'}] query = build_snowflake_query(check, specs, use_window=False) assert 'WHERE' not in query def test_window_adds_where_with_pyformat_params(self): """Query includes WHERE clause with %(window_start)s and %(window_end)s.""" check = { 'snowflake_table': 'DB.SCH.TBL', 'snowflake_window_column': 'UPDATED_AT', } specs = [{'alias': 'k0', 'snowflake': 'ID'}] query = build_snowflake_query(check, specs, use_window=True) assert 'WHERE' in query assert '%(window_start)s' in query assert '%(window_end)s' in query class TestFetchNeo4jRows: """Tests for fetch_neo4j_rows.""" def _make_record(self, data): """Create a mock record whose .data() returns the given dict.""" record = MagicMock() record.data.return_value = data return record def test_empty_result_returns_empty_dict(self, neo4j_driver): """Returns {} when the query yields no records.""" result = fetch_neo4j_rows(neo4j_driver, 'MATCH (n) RETURN n', {}, ['k0'], ['f0']) assert result == {} def test_rows_indexed_by_key_tuple(self, neo4j_driver): """Rows are indexed by a tuple of key values.""" session = neo4j_driver.session.return_value.__enter__.return_value session.run.return_value = iter([ self._make_record({'k0': 'id-1', 'f0': 'Alice'}), ]) result = fetch_neo4j_rows(neo4j_driver, 'Q', {}, ['k0'], ['f0']) assert ('id-1',) in result assert result[('id-1',)]['f0'] == 'Alice' def test_null_key_skipped(self, neo4j_driver): """Rows with a None key component are skipped.""" session = neo4j_driver.session.return_value.__enter__.return_value session.run.return_value = iter([ self._make_record({'k0': None, 'f0': 'Alice'}), ]) result = fetch_neo4j_rows(neo4j_driver, 'Q', {}, ['k0'], ['f0']) assert result == {} def test_duplicate_key_first_wins(self, neo4j_driver): """When two rows share a key, the first is kept.""" session = neo4j_driver.session.return_value.__enter__.return_value session.run.return_value = iter([ self._make_record({'k0': 'id-1', 'f0': 'first'}), self._make_record({'k0': 'id-1', 'f0': 'second'}), ]) result = fetch_neo4j_rows(neo4j_driver, 'Q', {}, ['k0'], ['f0']) assert result[('id-1',)]['f0'] == 'first' def test_values_normalized(self, neo4j_driver): """Values in compare fields are normalized (e.g. Decimal to int).""" session = neo4j_driver.session.return_value.__enter__.return_value session.run.return_value = iter([ self._make_record({'k0': 'id-1', 'f0': decimal.Decimal('7')}), ]) result = fetch_neo4j_rows(neo4j_driver, 'Q', {}, ['k0'], ['f0']) assert result[('id-1',)]['f0'] == 7 assert isinstance(result[('id-1',)]['f0'], int) class TestRowValue: """Tests for _row_value private helper.""" def test_exact_match(self): """Exact alias match returns the value.""" assert _row_value({'k0': 'val'}, 'k0') == 'val' def test_upper_case_match(self): """UPPER-case dict key is found when alias is lower-case.""" assert _row_value({'K0': 'val'}, 'k0') == 'val' def test_lower_case_match(self): """lower-case dict key is found when alias is UPPER-case.""" assert _row_value({'k0': 'val'}, 'K0') == 'val' def test_missing_returns_none(self): """Missing alias returns None.""" assert _row_value({}, 'k0') is None class TestFetchSnowflakeRows: """Tests for fetch_snowflake_rows.""" def test_empty_result_returns_empty_dict(self, snowflake_executor): """Returns {} when fetchall returns no rows.""" result = fetch_snowflake_rows(snowflake_executor, 'SELECT 1', {}, ['k0'], ['f0']) assert result == {} def test_rows_indexed_by_key_tuple(self, snowflake_executor): """Rows are indexed by a tuple of key values.""" snowflake_executor.fetchall.return_value = [{'k0': 'id-1', 'f0': 'Bob'}] result = fetch_snowflake_rows(snowflake_executor, 'Q', {}, ['k0'], ['f0']) assert ('id-1',) in result assert result[('id-1',)]['f0'] == 'Bob' def test_null_key_skipped(self, snowflake_executor): """Rows with a None key component are skipped.""" snowflake_executor.fetchall.return_value = [{'k0': None, 'f0': 'Bob'}] result = fetch_snowflake_rows(snowflake_executor, 'Q', {}, ['k0'], ['f0']) assert result == {} def test_duplicate_key_first_wins(self, snowflake_executor): """When two rows share a key, the first is kept.""" snowflake_executor.fetchall.return_value = [ {'k0': 'id-1', 'f0': 'first'}, {'k0': 'id-1', 'f0': 'second'}, ] result = fetch_snowflake_rows(snowflake_executor, 'Q', {}, ['k0'], ['f0']) assert result[('id-1',)]['f0'] == 'first' def test_dict_cursor_true_passed_to_fetchall(self, snowflake_executor): """fetchall is called with dict_cursor=True.""" fetch_snowflake_rows(snowflake_executor, 'Q', {}, ['k0'], ['f0']) snowflake_executor.fetchall.assert_called_once_with('Q', {}, dict_cursor=True) def test_case_insensitive_column_lookup(self, snowflake_executor): """Snowflake rows use case-insensitive column lookup via _row_value.""" snowflake_executor.fetchall.return_value = [{'K0': 'id-1', 'F0': 'Bob'}] result = fetch_snowflake_rows(snowflake_executor, 'Q', {}, ['k0'], ['f0']) assert ('id-1',) in result assert result[('id-1',)]['f0'] == 'Bob' class TestCompareRowMaps: """Tests for compare_row_maps.""" def _make_compare_specs(self): """Return a single compare spec for field 'f0'.""" return [{'alias': 'f0', 'label': 'NAME'}] def test_identical_maps_no_discrepancies(self): """Identical row maps produce no discrepancies.""" neo4j_rows = {('id-1',): {'f0': 'Alice'}} snowflake_rows = {('id-1',): {'f0': 'Alice'}} result = compare_row_maps(neo4j_rows, snowflake_rows, self._make_compare_specs()) assert result['missing_in_snowflake'] == [] assert result['missing_in_neo4j'] == [] assert result['field_mismatches'] == [] def test_missing_in_snowflake_detected(self): """Key in Neo4j but not Snowflake appears in missing_in_snowflake.""" neo4j_rows = {('id-1',): {'f0': 'Alice'}} result = compare_row_maps(neo4j_rows, {}, self._make_compare_specs()) assert ('id-1',) in result['missing_in_snowflake'] def test_missing_in_neo4j_detected(self): """Key in Snowflake but not Neo4j appears in missing_in_neo4j.""" snowflake_rows = {('id-1',): {'f0': 'Alice'}} result = compare_row_maps({}, snowflake_rows, self._make_compare_specs()) assert ('id-1',) in result['missing_in_neo4j'] def test_field_mismatch_detected(self): """Differing field values appear in field_mismatches.""" neo4j_rows = {('id-1',): {'f0': 'Alice'}} snowflake_rows = {('id-1',): {'f0': 'Bob'}} result = compare_row_maps(neo4j_rows, snowflake_rows, self._make_compare_specs()) assert len(result['field_mismatches']) == 1 def test_mismatch_entry_has_key_field_neo4j_snowflake(self): """Mismatch entry contains key, field, neo4j, and snowflake keys.""" neo4j_rows = {('id-1',): {'f0': 'Alice'}} snowflake_rows = {('id-1',): {'f0': 'Bob'}} result = compare_row_maps(neo4j_rows, snowflake_rows, self._make_compare_specs()) entry = result['field_mismatches'][0] assert 'key' in entry assert 'field' in entry assert 'neo4j' in entry assert 'snowflake' in entry def test_sample_number_truncates_results(self): """sample_number limits the number of results in each category.""" neo4j_rows = {(f'id-{i}',): {'f0': 'val'} for i in range(10)} result = compare_row_maps(neo4j_rows, {}, self._make_compare_specs(), sample_number=3) assert len(result['missing_in_snowflake']) == 3 def test_none_sample_number_returns_all(self): """None sample_number returns all discrepancies.""" neo4j_rows = {(f'id-{i}',): {'f0': 'val'} for i in range(10)} result = compare_row_maps( neo4j_rows, {}, self._make_compare_specs(), sample_number=None, ) assert len(result['missing_in_snowflake']) == 10 class TestRunRowLevelCheck: """Tests for run_row_level_check.""" _BASE_CHECK = { 'name': 'test', 'cypher_pattern': '(n:Node)', 'snowflake_table': 'DB.SCH.TBL', 'key_fields': [{'neo4j': 'n.id', 'snowflake': 'ID'}], 'compare_fields': [{'neo4j': 'n.name', 'snowflake': 'NAME'}], } def _window(self): """Return a fixed (window_start, window_end) tuple.""" return ( dt.datetime(2025, 1, 1, tzinfo=dt.timezone.utc), dt.datetime(2025, 1, 8, tzinfo=dt.timezone.utc), ) def test_passing_check_summary_should_be_true(self, mocker, neo4j_driver, snowflake_executor): """summary['shouldBeTrue'] is True when there are no mismatches.""" mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_neo4j_rows', return_value={}, ) mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_snowflake_rows', return_value={}, ) ws, we = self._window() summary, _ = run_row_level_check( self._BASE_CHECK, neo4j_driver, snowflake_executor, ws, we, ) assert summary['shouldBeTrue'] is True def test_failing_check_summary_should_be_false(self, mocker, neo4j_driver, snowflake_executor): """summary['shouldBeTrue'] is False when mismatches exist.""" mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_neo4j_rows', return_value={('id-1',): {}}, ) mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_snowflake_rows', return_value={}, ) ws, we = self._window() summary, _ = run_row_level_check( self._BASE_CHECK, neo4j_driver, snowflake_executor, ws, we, ) assert summary['shouldBeTrue'] is False def test_summary_counts_match_comparison_lengths( self, mocker, neo4j_driver, snowflake_executor): """Summary mismatch counts reflect the comparison result lengths.""" mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_neo4j_rows', return_value={('id-1',): {}, ('id-2',): {}}, ) mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_snowflake_rows', return_value={}, ) ws, we = self._window() summary, comparison = run_row_level_check( self._BASE_CHECK, neo4j_driver, snowflake_executor, ws, we, ) assert summary['missing_in_snowflake_count'] == len(comparison['missing_in_snowflake']) def test_window_params_passed_when_use_window(self, mocker, neo4j_driver, snowflake_executor): """fetch_neo4j_rows receives window params when the check uses a time window.""" check = dict(self._BASE_CHECK) check['window_property'] = 'n.updatedAt' check['snowflake_window_column'] = 'UPDATED_AT' mock_fetch = mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_neo4j_rows', return_value={}, ) mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_snowflake_rows', return_value={}, ) ws, we = self._window() run_row_level_check(check, neo4j_driver, snowflake_executor, ws, we) params = mock_fetch.call_args[0][2] assert params == {'window_start': ws, 'window_end': we} def test_no_params_when_no_window(self, mocker, neo4j_driver, snowflake_executor): """fetch_neo4j_rows receives empty params when no window is configured.""" mock_fetch = mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_neo4j_rows', return_value={}, ) mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_snowflake_rows', return_value={}, ) ws, we = self._window() run_row_level_check(self._BASE_CHECK, neo4j_driver, snowflake_executor, ws, we) params = mock_fetch.call_args[0][2] assert params == {} def test_default_sample_number_used(self, mocker, neo4j_driver, snowflake_executor): """DEFAULT_SAMPLE_NUMBER is used when sample_number is absent from check.""" mock_compare = mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.compare_row_maps', return_value={ 'missing_in_snowflake': [], 'missing_in_neo4j': [], 'field_mismatches': [], }, ) mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_neo4j_rows', return_value={}, ) mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_snowflake_rows', return_value={}, ) ws, we = self._window() run_row_level_check(self._BASE_CHECK, neo4j_driver, snowflake_executor, ws, we) sample_number = mock_compare.call_args[0][3] assert sample_number == config.DEFAULT_SAMPLE_NUMBER def test_custom_sample_number_respected(self, mocker, neo4j_driver, snowflake_executor): """Custom sample_number in check config is passed to compare_row_maps.""" mock_compare = mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.compare_row_maps', return_value={ 'missing_in_snowflake': [], 'missing_in_neo4j': [], 'field_mismatches': [], }, ) mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_neo4j_rows', return_value={}, ) mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_snowflake_rows', return_value={}, ) check = dict(self._BASE_CHECK) check['sample_number'] = 42 ws, we = self._window() run_row_level_check(check, neo4j_driver, snowflake_executor, ws, we) sample_number = mock_compare.call_args[0][3] assert sample_number == 42 def test_row_counts_in_summary(self, mocker, neo4j_driver, snowflake_executor): """Summary contains accurate neo4j_row_count and snowflake_row_count.""" mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_neo4j_rows', return_value={('a',): {}, ('b',): {}, ('c',): {}}, ) mocker.patch( 'src.assert_snowflake_sync.snowflake_sync.fetch_snowflake_rows', return_value={('a',): {}, ('b',): {}}, ) ws, we = self._window() summary, _ = run_row_level_check( self._BASE_CHECK, neo4j_driver, snowflake_executor, ws, we, ) assert summary['neo4j_row_count'] == 3 assert summary['snowflake_row_count'] == 2