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
106 changes: 80 additions & 26 deletions burr/integrations/serde/pandas.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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)
135 changes: 126 additions & 9 deletions tests/integrations/serde/test_pandas.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,26 +16,44 @@
# under the License.

import os
import sys

import pandas as pd
import pytest

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)

Expand All @@ -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})
Expand All @@ -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)
Loading