diff --git a/burr/integrations/serde/pandas.py b/burr/integrations/serde/pandas.py index 22ee32b01..434eb7351 100644 --- a/burr/integrations/serde/pandas.py +++ b/burr/integrations/serde/pandas.py @@ -16,52 +16,100 @@ # under the License. # try to import to serialize Pandas Objects +# +# DataFrames are stored as parquet files under a configured directory +# (``pandas_kwargs["path"]``). The persisted state records the file name +# relative to that directory; on load the name is resolved against the +# configured directory and must land inside it. import hashlib import os -from typing import Optional +from typing import Optional, Union import pandas as pd from burr.core import serde +def _pandas_kwargs_with_path(pandas_kwargs: Optional[dict], operation: str) -> dict: + """Returns a copy of ``pandas_kwargs``, requiring it to carry a ``path`` entry.""" + if not isinstance(pandas_kwargs, dict) or "path" not in pandas_kwargs: + raise ValueError( + f"{operation} a pandas DataFrame requires a `path` entry in `pandas_kwargs` -- " + "this is the directory where the dataframe is stored as a parquet file. " + "Pass it through whatever triggers serialization, e.g. " + '`LocalTrackingClient(..., serde_kwargs={"pandas_kwargs": {"path": "/some/dir"}})` ' + 'or `state.serialize(pandas_kwargs={"path": "/some/dir"})`. ' + f"Got pandas_kwargs={pandas_kwargs!r}." + ) + return pandas_kwargs.copy() + + +def _resolve_within_base(base_path: Union[str, os.PathLike], stored_path: object) -> str: + """Resolves a persisted parquet path against the configured directory. + + Serialization records the parquet file name relative to the configured directory. + Earlier versions recorded the joined path instead, so an absolute value is still + accepted as long as it resolves to a location inside the configured directory. + URL-style values and anything resolving outside the directory are rejected. + + :param base_path: the configured directory (``pandas_kwargs["path"]``). + :param stored_path: the ``path`` value read from persisted state. + :return: the resolved absolute path of the parquet file. + """ + if not isinstance(stored_path, str) or not stored_path: + raise ValueError(f"Persisted parquet path must be a non-empty string, got {stored_path!r}.") + if "://" in stored_path: + raise ValueError( + f"Persisted parquet path {stored_path!r} must be a file name relative to the " + "configured directory, not a URL." + ) + real_base = os.path.realpath(os.fspath(base_path)) + candidate = stored_path if os.path.isabs(stored_path) else os.path.join(real_base, stored_path) + real_candidate = os.path.realpath(candidate) + try: + inside = ( + real_candidate != real_base + and os.path.commonpath([real_base, real_candidate]) == real_base + ) + except ValueError: + # commonpath cannot compare the two (e.g. different drives on Windows) + inside = False + if not inside: + raise ValueError( + f"Persisted parquet path {stored_path!r} resolves outside the configured " + f"directory {os.fspath(base_path)!r}." + ) + return real_candidate + + @serde.serialize.register(pd.DataFrame) def serialize_pandas_df( value: pd.DataFrame, pandas_kwargs: Optional[dict] = None, **kwargs ) -> dict: """Custom serde for pandas dataframes. - Saves the dataframe to a parquet file and returns the path to the file. + Saves the dataframe to a parquet file under ``pandas_kwargs["path"]`` and records the + file name, relative to that directory, in the returned dictionary. Pass the same + ``path`` when deserializing; the recorded name is resolved against it. Requires a `path` key in the `pandas_kwargs` dictionary. :param value: the pandas dataframe to serialize. - :param pandas_kwargs: `path` key is required -- this is the base path to save the parquet file. As \ - well as any other kwargs to pass to the pandas to_parquet function. + :param pandas_kwargs: `path` key is required -- this is the directory to save the parquet \ + file in. As well as any other kwargs to pass to the pandas to_parquet function. :param kwargs: :return: """ - if not isinstance(pandas_kwargs, dict) or "path" not in pandas_kwargs: - raise ValueError( - "Serializing a pandas DataFrame requires a `path` entry in `pandas_kwargs` -- " - "this is the base path where the dataframe is saved as a parquet file. " - "Pass it through whatever triggers serialization, e.g. " - '`LocalTrackingClient(..., serde_kwargs={"pandas_kwargs": {"path": "/some/dir"}})` ' - 'or `state.serialize(pandas_kwargs={"path": "/some/dir"})`. ' - f"Got pandas_kwargs={pandas_kwargs!r}." - ) + kwargs = _pandas_kwargs_with_path(pandas_kwargs, "Serializing") hash_object = hashlib.sha256() hash_value = str(value.columns) + str(value.shape) + str(value.dtypes) hash_object.update(hash_value.encode()) # Return the hexadecimal representation of the hash file_name = f"df_{hash_object.hexdigest()}.parquet" - kwargs = pandas_kwargs.copy() - base_path: str = kwargs.pop("path") - if not os.path.exists(base_path): - os.makedirs(base_path) - saved_to = os.path.join(base_path, file_name) - value.to_parquet(path=saved_to, **kwargs) - return {serde.KEY: "pandas.DataFrame", "path": saved_to} + base_path = kwargs.pop("path") + os.makedirs(base_path, exist_ok=True) + value.to_parquet(path=os.path.join(base_path, file_name), **kwargs) + return {serde.KEY: "pandas.DataFrame", "path": file_name} @serde.deserializer.register("pandas.DataFrame") @@ -70,13 +118,19 @@ def deserialize_pandas_df( ) -> pd.DataFrame: """Custom deserializer for pandas dataframes. + The persisted ``path`` is resolved against ``pandas_kwargs["path"]`` -- the same + directory used during serialization -- and must resolve to a file inside it. + Values recorded by earlier versions as absolute paths are accepted when they resolve + inside that directory. URL-style values and paths outside the directory raise + ``ValueError``. + :param value: the dictionary to pull the path from to load the parquet file. - :param pandas_kwargs: other args to pass to the pandas read_parquet function. Optional. + :param pandas_kwargs: `path` key is required -- the directory the parquet file was saved \ + in. As well as any other kwargs to pass to the pandas read_parquet function. :param kwargs: :return: pandas dataframe """ - kwargs = pandas_kwargs.copy() if pandas_kwargs is not None else {} - if "path" in kwargs: - # remove this to not clash; we already have the full path. - kwargs.pop("path") - return pd.read_parquet(value["path"], **kwargs) + kwargs = _pandas_kwargs_with_path(pandas_kwargs, "Deserializing") + base_path = kwargs.pop("path") + resolved = _resolve_within_base(base_path, value.get("path")) + return pd.read_parquet(resolved, **kwargs) diff --git a/tests/integrations/serde/test_pandas.py b/tests/integrations/serde/test_pandas.py index 90ce9950b..0cc8868f0 100644 --- a/tests/integrations/serde/test_pandas.py +++ b/tests/integrations/serde/test_pandas.py @@ -16,6 +16,7 @@ # under the License. import os +import sys import pandas as pd import pytest @@ -23,19 +24,36 @@ from burr.core import serde, state +def _write_parquet(path, df) -> str: + """Writes ``df`` to ``path`` (creating parent dirs) and returns the path as a string.""" + os.makedirs(os.path.dirname(path), exist_ok=True) + df.to_parquet(path) + assert os.path.exists(path) + return str(path) + + +def _persisted(path) -> dict: + """Builds a persisted-state dict pointing the dataframe field at ``path``.""" + return {"df": {serde.KEY: "pandas.DataFrame", "path": path}} + + def test_serde_of_pandas_dataframe(tmp_path): df = pd.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]}) og = state.State({"df": df}) serialized = og.serialize(pandas_kwargs={"path": tmp_path}) assert serialized["df"][serde.KEY] == "pandas.DataFrame" - assert serialized["df"]["path"].startswith(str(tmp_path)) + + # The recorded path is the file name relative to the configured directory. + stored = serialized["df"]["path"] + assert not os.path.isabs(stored) + assert os.path.basename(stored) == stored + assert os.path.exists(tmp_path / stored) # Verify filename pattern instead of exact hash (hash may change with pandas versions) - filename = os.path.basename(serialized["df"]["path"]) - assert filename.startswith("df_") - assert filename.endswith(".parquet") + assert stored.startswith("df_") + assert stored.endswith(".parquet") # Verify it's a valid SHA256 hash (64 hex characters) - hash_part = filename[3:-8] # Remove 'df_' prefix and '.parquet' suffix + hash_part = stored[3:-8] # Remove 'df_' prefix and '.parquet' suffix assert len(hash_part) == 64 assert all(c in "0123456789abcdef" for c in hash_part) @@ -44,6 +62,15 @@ def test_serde_of_pandas_dataframe(tmp_path): pd.testing.assert_frame_equal(ng["df"], df) +def test_serde_of_pandas_dataframe_with_relative_base_path(tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + df = pd.DataFrame({"a": [1, 2, 3]}) + serialized = state.State({"df": df}).serialize(pandas_kwargs={"path": "data"}) + assert os.path.exists(tmp_path / "data" / serialized["df"]["path"]) + ng = state.State.deserialize(serialized, pandas_kwargs={"path": "data"}) + pd.testing.assert_frame_equal(ng["df"], df) + + def test_serialize_pandas_df_without_pandas_kwargs_raises_informative_error(): df = pd.DataFrame({"a": [1, 2, 3]}) og = state.State({"df": df}) @@ -62,9 +89,99 @@ def test_serialize_pandas_df_without_path_raises_informative_error(): assert "path" in str(exc_info.value) -def test_deserialize_pandas_df_without_pandas_kwargs(tmp_path): +def test_deserialize_pandas_df_without_path_raises_informative_error(tmp_path): df = pd.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]}) - og = state.State({"df": df}) - serialized = og.serialize(pandas_kwargs={"path": tmp_path}) - ng = state.State.deserialize(serialized) + serialized = state.State({"df": df}).serialize(pandas_kwargs={"path": tmp_path}) + for pandas_kwargs in ({}, None, {"columns": ["a"]}): + with pytest.raises(ValueError) as exc_info: + state.State.deserialize(serialized, pandas_kwargs=pandas_kwargs) + assert "Failed to deserialize state field 'df'" in str(exc_info.value) + assert "pandas_kwargs" in str(exc_info.value) + assert "path" in str(exc_info.value) + + +def test_deserialize_pandas_df_passes_through_read_kwargs(tmp_path): + df = pd.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]}) + serialized = state.State({"df": df}).serialize(pandas_kwargs={"path": tmp_path}) + ng = state.State.deserialize(serialized, pandas_kwargs={"path": tmp_path, "columns": ["a"]}) + pd.testing.assert_frame_equal(ng["df"], df[["a"]]) + + +def test_deserialize_pandas_df_accepts_legacy_absolute_path_inside_base(tmp_path): + # Earlier versions recorded the joined path; it still loads when it is inside the base. + df = pd.DataFrame({"a": [1, 2, 3]}) + legacy = _write_parquet(tmp_path / "df_legacy.parquet", df) + assert os.path.isabs(legacy) + ng = state.State.deserialize(_persisted(legacy), pandas_kwargs={"path": tmp_path}) + pd.testing.assert_frame_equal(ng["df"], df) + + +def test_deserialize_pandas_df_accepts_nested_relative_path_inside_base(tmp_path): + df = pd.DataFrame({"a": [1, 2, 3]}) + _write_parquet(tmp_path / "nested" / "df.parquet", df) + ng = state.State.deserialize( + _persisted(os.path.join("nested", "df.parquet")), pandas_kwargs={"path": tmp_path} + ) pd.testing.assert_frame_equal(ng["df"], df) + + +def test_deserialize_pandas_df_rejects_parent_directory_path(tmp_path): + df = pd.DataFrame({"a": [1, 2, 3]}) + base = tmp_path / "base" + base.mkdir() + # The target exists, so the failure is the resolution rule rather than a missing file. + _write_parquet(tmp_path / "outside.parquet", df) + with pytest.raises(ValueError) as exc_info: + state.State.deserialize(_persisted("../outside.parquet"), pandas_kwargs={"path": base}) + assert "resolves outside the configured directory" in str(exc_info.value) + + +def test_deserialize_pandas_df_rejects_absolute_path_outside_base(tmp_path): + df = pd.DataFrame({"a": [1, 2, 3]}) + base = tmp_path / "base" + base.mkdir() + outside = _write_parquet(tmp_path / "outside.parquet", df) + with pytest.raises(ValueError) as exc_info: + state.State.deserialize(_persisted(outside), pandas_kwargs={"path": base}) + assert "resolves outside the configured directory" in str(exc_info.value) + + +def test_deserialize_pandas_df_rejects_base_directory_itself(tmp_path): + for stored in (".", "", str(tmp_path)): + with pytest.raises(ValueError): + state.State.deserialize(_persisted(stored), pandas_kwargs={"path": tmp_path}) + + +@pytest.mark.parametrize( + "stored", + [ + "http://example.com/df.parquet", + "https://example.com/df.parquet", + "s3://bucket/df.parquet", + "gs://bucket/df.parquet", + "file:///tmp/df.parquet", + ], +) +def test_deserialize_pandas_df_rejects_url_paths(tmp_path, stored): + with pytest.raises(ValueError) as exc_info: + state.State.deserialize(_persisted(stored), pandas_kwargs={"path": tmp_path}) + assert "not a URL" in str(exc_info.value) + + +@pytest.mark.parametrize("stored", [None, 123, ["df.parquet"]]) +def test_deserialize_pandas_df_rejects_non_string_path(tmp_path, stored): + with pytest.raises(ValueError) as exc_info: + state.State.deserialize(_persisted(stored), pandas_kwargs={"path": tmp_path}) + assert "non-empty string" in str(exc_info.value) + + +@pytest.mark.skipif(sys.platform == "win32", reason="symlink creation needs privileges on Windows") +def test_deserialize_pandas_df_rejects_symlink_escaping_base(tmp_path): + df = pd.DataFrame({"a": [1, 2, 3]}) + base = tmp_path / "base" + base.mkdir() + outside = _write_parquet(tmp_path / "outside.parquet", df) + os.symlink(outside, base / "link.parquet") + with pytest.raises(ValueError) as exc_info: + state.State.deserialize(_persisted("link.parquet"), pandas_kwargs={"path": base}) + assert "resolves outside the configured directory" in str(exc_info.value)