diff --git a/src/cdm_reader_mapper/common/__init__.py b/src/cdm_reader_mapper/common/__init__.py index ff175e2f..074d0414 100755 --- a/src/cdm_reader_mapper/common/__init__.py +++ b/src/cdm_reader_mapper/common/__init__.py @@ -2,12 +2,12 @@ from __future__ import annotations +from .dataframe_helpers import restore_columns, standardize_object_columns from .getting_files import load_file from .inspect import count_by_cat, get_length from .io_files import get_filename from .iterators import ParquetStreamReader, ProcessFunction, is_valid_iterator, parquet_stream_from_iterable, process_disk_backed, process_function from .json_dict import collect_json_files, combine_dicts, open_json_file -from .object_types import standardize_object_columns from .replace import replace_columns from .select import ( split_by_boolean, diff --git a/src/cdm_reader_mapper/common/dataframe_helpers.py b/src/cdm_reader_mapper/common/dataframe_helpers.py new file mode 100755 index 00000000..8266c8a7 --- /dev/null +++ b/src/cdm_reader_mapper/common/dataframe_helpers.py @@ -0,0 +1,84 @@ +"""Utility function for reading and writing files.""" + +from __future__ import annotations +from typing import Any + +import pandas as pd + + +def standardize_object_columns(df: pd.DataFrame) -> pd.DataFrame: + """ + Convert string columns to object dtype and replace NaNs with None. + + Parameters + ---------- + df : pd.DataFrame + The input DataFrame to be standardized. + + Returns + ------- + pd.DataFrame + The same DataFrame instance after the dtype conversion and NaN handling. + """ + df = df.copy() + string_cols = df.select_dtypes(include="string").columns + df[string_cols] = df[string_cols].astype(object) + object_cols = df.select_dtypes(include="object").columns + df[object_cols] = df[object_cols].fillna(None) + return df + + +def restore_columns(item: Any) -> Any: + """ + Restore columns from string literals if `item` is a pandas DataFrame or Series. + + Parameters + ---------- + item : Any + Object to restore. + + Returns + ------- + Any + Restored object. + """ + + def _literal_eval(column: Any) -> Any: + """ + Evaluate a string literal if possible. + + Parameters + ---------- + column : Any + Column that is possibly a string literal. + + Returns + ------- + Any + Evaluated column. + """ + if not isinstance(column, str): + return column + try: + from ast import literal_eval + + return literal_eval(column) + except (ValueError, SyntaxError): + return column + + if isinstance(item, pd.DataFrame): + columns = item.columns + new_columns = [] + for column in columns: + column = _literal_eval(column) + new_columns.append(column) + + if new_columns and all(isinstance(c, tuple) for c in new_columns) and len({len(c) for c in new_columns}) == 1: + item.columns = pd.MultiIndex.from_tuples(new_columns) + else: + item.columns = new_columns + + if isinstance(item, pd.Series): + item.name = _literal_eval(item.name) + + return item diff --git a/src/cdm_reader_mapper/common/iterators.py b/src/cdm_reader_mapper/common/iterators.py index 94e97bb2..1688b3e5 100755 --- a/src/cdm_reader_mapper/common/iterators.py +++ b/src/cdm_reader_mapper/common/iterators.py @@ -18,7 +18,7 @@ import pyarrow.parquet as pq import xarray as xr -from .object_types import standardize_object_columns +from .dataframe_helpers import restore_columns, standardize_object_columns class ProcessFunction: @@ -291,7 +291,8 @@ def read(self) -> pd.DataFrame: if not chunks: return pd.DataFrame() - return pd.concat(chunks) + df = pd.concat(chunks) + return restore_columns(df) def copy(self) -> ParquetStreamReader: """ @@ -835,6 +836,8 @@ def _process_chunks( for items in zip(*readers, strict=True): _validate_chunk(items, requested_types) + items = tuple([restore_columns(item) for item in items]) + result = func(*items, *static_args, **static_kwargs) data, meta = _process_result(result, requested_types, non_data_output, chunk_counter) diff --git a/src/cdm_reader_mapper/common/object_types.py b/src/cdm_reader_mapper/common/object_types.py deleted file mode 100755 index afe4e18e..00000000 --- a/src/cdm_reader_mapper/common/object_types.py +++ /dev/null @@ -1,27 +0,0 @@ -"""Utility function for reading and writing files.""" - -from __future__ import annotations - -import pandas as pd - - -def standardize_object_columns(df: pd.DataFrame) -> pd.DataFrame: - """ - Convert string columns to object dtype and replace NaNs with None. - - Parameters - ---------- - df : pd.DataFrame - The input DataFrame to be standardized. - - Returns - ------- - pd.DataFrame - The same DataFrame instance after the dtype conversion and NaN handling. - """ - df = df.copy() - string_cols = df.select_dtypes(include="string").columns - df[string_cols] = df[string_cols].astype(object) - object_cols = df.select_dtypes(include="object").columns - df[object_cols] = df[object_cols].fillna(None) - return df diff --git a/tests/test_common_utils.py b/tests/test_common_utils.py index 8d059279..ac3e7d4d 100755 --- a/tests/test_common_utils.py +++ b/tests/test_common_utils.py @@ -13,6 +13,7 @@ import pytest import requests +from cdm_reader_mapper.common.dataframe_helpers import restore_columns, standardize_object_columns from cdm_reader_mapper.common.getting_files import ( _check_md5s, _file_md5_checksum, @@ -30,7 +31,6 @@ open_json_file, ) from cdm_reader_mapper.common.logging_hdlr import init_logger -from cdm_reader_mapper.common.object_types import standardize_object_columns def compute_md5(content: bytes) -> str: @@ -508,3 +508,55 @@ def test_standardize_object_columns(): ) pd.testing.assert_frame_equal(result, expected) + + +def test_restore_columns_index(): + df = pd.DataFrame(columns=["a", "b", "c"]) + + result = restore_columns(df) + expected = pd.Index(["a", "b", "c"]) + + pd.testing.assert_index_equal(result.columns, expected) + + +def test_restore_columns_evaluate(): + df = pd.DataFrame(columns=["1", "b", "True"]) + + result = restore_columns(df) + expected = pd.Index([1, "b", True]) + + pd.testing.assert_index_equal(result.columns, expected) + + +def test_restore_columns_multiindex(): + df = pd.DataFrame(columns=["('a','d')", "('b','e')", "('c','f')"]) + + result = restore_columns(df) + expected = pd.MultiIndex.from_tuples([("a", "d"), ("b", "e"), ("c", "f")]) + + pd.testing.assert_index_equal(result.columns, expected) + + +def test_restore_columns_mixed(): + df = pd.DataFrame(columns=["a", "('b','e')", "('c','f')"]) + + result = restore_columns(df) + expected = pd.Index(["a", ("b", "e"), ("c", "f")]) + + pd.testing.assert_index_equal(result.columns, expected) + + +def test_restore_columns_series(): + s = pd.Series([1, 2, 3], name="('a', 'b')") + + result = restore_columns(s) + + assert result.name == ("a", "b") + + +def test_restore_columns_no_pandas(): + obj = {"a": 1} + + result = restore_columns(obj) + + assert result is obj