Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
2 changes: 2 additions & 0 deletions chispa/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -55,6 +56,7 @@ def assert_df_equality(
ignore_metadata,
ignore_columns,
self.formats,
full_log,
)


Expand Down
14 changes: 12 additions & 2 deletions chispa/dataframe_comparer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -102,13 +103,15 @@ def assert_df_equality(
{"allow_nan_equality": True},
underline_cells=underline_cells,
formats=formats,
full_log=full_log,
)
else:
assert_basic_rows_equality(
df1.collect(),
df2.collect(),
underline_cells=underline_cells,
formats=formats,
full_log=full_log,
)


Expand All @@ -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()
Expand Down Expand Up @@ -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)
15 changes: 12 additions & 3 deletions chispa/rows_comparer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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(
Expand All @@ -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()
Expand Down Expand Up @@ -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")
27 changes: 27 additions & 0 deletions tests/test_dataframe_comparer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"])
Expand Down Expand Up @@ -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"])
Expand Down
Loading