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
64 changes: 60 additions & 4 deletions src/pytest_regressions/dataframe_regression.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import csv
import os
from pathlib import Path
from typing import Any
Expand Down Expand Up @@ -93,7 +94,15 @@ def _check_data_shapes(self, obtained_column: Any, expected_column: Any) -> None
)
raise AssertionError(error_msg)

def _check_fn(self, obtained_filename: Path, expected_filename: Path) -> None:
def _check_fn(
self,
obtained_filename: Path,
expected_filename: Path,
*,
column_levels: int = 1,
has_multiindex_columns: bool = False,
has_index_names: bool = False,
) -> None:
"""
Check if dict contents dumped to a file match the contents in expected file.
"""
Expand All @@ -108,8 +117,39 @@ def _check_fn(self, obtained_filename: Path, expected_filename: Path) -> None:

__tracebackhide__ = True

obtained_data = pd.read_csv(str(obtained_filename))
expected_data = pd.read_csv(str(expected_filename))
header = list(range(column_levels))
read_csv_options: dict[str, Any] = {"header": header}
if has_index_names:

def read_index_names(filename: Path) -> list[str]:
with filename.open(encoding="utf-8", newline="") as stream:
rows = csv.reader(stream)
for _ in header:
next(rows, None)
names = next(rows, None)
assert names is not None, "Could not find the row index names."
return names

assert read_index_names(obtained_filename) == read_index_names(
expected_filename
), "Row index names are not the same."
# MultiIndex CSVs store named row indexes in an extra metadata row.
read_csv_options["skiprows"] = [column_levels]
# The Python parser skips CSV records, including quoted newlines.
read_csv_options["engine"] = "python"

obtained_data = pd.read_csv(str(obtained_filename), **read_csv_options)
try:
expected_data = pd.read_csv(str(expected_filename), **read_csv_options)
except pd.errors.ParserError as error:
raise AssertionError(
f"Could not parse the expected results.\n{error}\n"
"To update values, use --force-regen option.\n"
) from error
if has_multiindex_columns and column_levels == 1:
# A one-row CSV header is otherwise parsed as a flat Index.
for frame in (obtained_data, expected_data):
frame.columns = pd.MultiIndex.from_arrays([frame.columns])

comparison_tables_dict = {}
for k in obtained_data.keys():
Expand Down Expand Up @@ -282,13 +322,29 @@ def check(
self._default_tolerance = default_tolerance

dump_fn = functools.partial(self._dump_fn, data_frame)
has_multiindex_columns = isinstance(data_frame.columns, pd.MultiIndex)
has_index_names = False
if has_multiindex_columns:
# Match the index labels emitted by pandas' MultiIndex CSV writer.
if isinstance(data_frame.index, pd.MultiIndex):
index_labels = [name or "" for name in data_frame.index.names]
else:
index_labels = [
"" if name is None else name for name in data_frame.index.names
]
has_index_names = set(index_labels) != {""}

with pd.option_context(*self._pandas_display_options):
perform_regression_check(
datadir=self.datadir,
original_datadir=self.original_datadir,
request=self.request,
check_fn=self._check_fn,
check_fn=functools.partial(
self._check_fn,
column_levels=data_frame.columns.nlevels,
has_multiindex_columns=has_multiindex_columns,
has_index_names=has_index_names,
),
dump_fn=dump_fn,
extension=".csv",
basename=basename,
Expand Down
Loading
Loading