diff --git a/README.md b/README.md index 2ad75ff..5d88712 100644 --- a/README.md +++ b/README.md @@ -157,6 +157,15 @@ Here's the nicely formatted error message: ![DataFramesNotEqualError](https://raw.githubusercontent.com/MrPowers/chispa/main/images/dfs_not_equal_error.png) +### Disable the full diff log + +If your DataFrames are large and the diff table is too noisy for your test output, set `full_log=False` to raise a compact error without printing the row-level diff. + +```python +assert_df_equality(actual_df, expected_df, full_log=False) +assert_approx_df_equality(actual_df, expected_df, 0.01, full_log=False) +``` + ### Ignore row order You can easily compare DataFrames, ignoring the order of the rows. The content of the DataFrames is usually what matters, not the order of the rows. diff --git a/chispa/__init__.py b/chispa/__init__.py index 57b35c6..a528691 100644 --- a/chispa/__init__.py +++ b/chispa/__init__.py @@ -42,6 +42,7 @@ def assert_df_equality( underline_cells: bool = False, ignore_metadata: bool = False, ignore_columns: list[str] | None = None, + full_log: bool = True, ) -> None: return assert_df_equality( df1, @@ -55,6 +56,7 @@ def assert_df_equality( ignore_metadata, ignore_columns, self.formats, + full_log, ) diff --git a/chispa/dataframe_comparer.py b/chispa/dataframe_comparer.py index da8b1f9..f686d0f 100644 --- a/chispa/dataframe_comparer.py +++ b/chispa/dataframe_comparer.py @@ -74,6 +74,7 @@ def assert_df_equality( ignore_metadata: bool = False, ignore_columns: list[str] | None = None, formats: FormattingConfig | None = None, + full_log: bool = True, ) -> None: if not formats: formats = FormattingConfig() @@ -102,6 +103,7 @@ def assert_df_equality( {"allow_nan_equality": True}, underline_cells=underline_cells, formats=formats, + full_log=full_log, ) else: assert_basic_rows_equality( @@ -109,6 +111,7 @@ def assert_df_equality( df2.collect(), underline_cells=underline_cells, formats=formats, + full_log=full_log, ) @@ -131,6 +134,7 @@ def assert_approx_df_equality( ignore_row_order: bool = False, ignore_columns: list[str] | None = None, formats: FormattingConfig | None = None, + full_log: bool = True, ) -> None: if not formats: formats = FormattingConfig() @@ -158,10 +162,16 @@ def assert_approx_df_equality( are_rows_approx_equal, {"precision": precision, "allow_nan_equality": allow_nan_equality}, formats=formats, + full_log=full_log, ) elif allow_nan_equality: assert_generic_rows_equality( - df1.collect(), df2.collect(), are_rows_equal_enhanced, {"allow_nan_equality": True}, formats=formats + df1.collect(), + df2.collect(), + are_rows_equal_enhanced, + {"allow_nan_equality": True}, + formats=formats, + full_log=full_log, ) else: - assert_basic_rows_equality(df1.collect(), df2.collect(), formats=formats) + assert_basic_rows_equality(df1.collect(), df2.collect(), formats=formats, full_log=full_log) diff --git a/chispa/rows_comparer.py b/chispa/rows_comparer.py index 2019527..7fe4216 100644 --- a/chispa/rows_comparer.py +++ b/chispa/rows_comparer.py @@ -12,7 +12,11 @@ def assert_basic_rows_equality( - rows1: list[Row], rows2: list[Row], underline_cells: bool = False, formats: FormattingConfig | None = None + rows1: list[Row], + rows2: list[Row], + underline_cells: bool = False, + formats: FormattingConfig | None = None, + full_log: bool = True, ) -> None: if not formats: formats = FormattingConfig() @@ -48,7 +52,9 @@ def assert_basic_rows_equality( t.add_row([r1_res, r2_res]) if all_rows_equal is False: - raise chispa.DataFramesNotEqualError("\n" + t.get_string()) + if full_log: + raise chispa.DataFramesNotEqualError("\n" + t.get_string()) + raise chispa.DataFramesNotEqualError("DataFrames are not equal") def assert_generic_rows_equality( @@ -58,6 +64,7 @@ def assert_generic_rows_equality( row_equality_fun_args: dict[str, Any], underline_cells: bool = False, formats: FormattingConfig | None = None, + full_log: bool = True, ) -> None: if not formats: formats = FormattingConfig() @@ -103,4 +110,6 @@ def assert_generic_rows_equality( t.add_row([r1_res, r2_res]) if all_rows_equal is False: - raise chispa.DataFramesNotEqualError("\n" + t.get_string()) + if full_log: + raise chispa.DataFramesNotEqualError("\n" + t.get_string()) + raise chispa.DataFramesNotEqualError("DataFrames are not equal") diff --git a/tests/test_dataframe_comparer.py b/tests/test_dataframe_comparer.py index d21f4da..59c413f 100644 --- a/tests/test_dataframe_comparer.py +++ b/tests/test_dataframe_comparer.py @@ -152,6 +152,20 @@ def it_throws_with_content_mismatches(spark: SparkSession): with pytest.raises(DataFramesNotEqualError): assert_df_equality(df1, df2) + def it_can_raise_without_the_full_diff_log(spark: SparkSession): + data1 = [("jose", "jose"), ("li", "li")] + df1 = spark.createDataFrame(data1, ["name", "expected_name"]) + data2 = [("bob", "jose"), ("li", "li")] + df2 = spark.createDataFrame(data2, ["name", "expected_name"]) + + with pytest.raises(DataFramesNotEqualError) as exc_info: + assert_df_equality(df1, df2, full_log=False) + + message = str(exc_info.value) + assert message == "DataFrames are not equal" + assert "bob" not in message + assert "PrettyTable" not in message + def it_throws_with_length_mismatches(spark: SparkSession): data1 = [("jose", "jose"), ("li", "li"), ("laura", "laura")] df1 = spark.createDataFrame(data1, ["name", "expected_name"]) @@ -282,6 +296,19 @@ def it_throws_with_content_mismatch(spark: SparkSession): with pytest.raises(DataFramesNotEqualError): assert_approx_df_equality(df1, df2, 0.1) + def it_can_raise_approx_without_the_full_diff_log(spark: SparkSession): + data1 = [(1.0, "jose"), (1.1, "li")] + df1 = spark.createDataFrame(data1, ["num", "expected_name"]) + data2 = [(1.0, "jose"), (9.9, "li")] + df2 = spark.createDataFrame(data2, ["num", "expected_name"]) + + with pytest.raises(DataFramesNotEqualError) as exc_info: + assert_approx_df_equality(df1, df2, 0.1, full_log=False) + + message = str(exc_info.value) + assert message == "DataFrames are not equal" + assert "9.9" not in message + def it_throws_with_with_length_mismatch(spark: SparkSession): data1 = [(1.0, "jose"), (1.1, "li"), (1.2, "laura"), (None, None)] df1 = spark.createDataFrame(data1, ["num", "expected_name"])