Skip to content

Commit f325241

Browse files
feat: support polars dataframe in CachedFoundryClient.save_dataset (#103)
* feat: support polars dataframe in CachedFoundryClient.save_dataset * Update docstring Co-authored-by: Nicolas Renkamp <nicornk@users.noreply.github.com> --------- Co-authored-by: Nicolas Renkamp <nicornk@users.noreply.github.com>
1 parent 03f84a6 commit f325241

2 files changed

Lines changed: 29 additions & 8 deletions

File tree

libs/foundry-dev-tools/src/foundry_dev_tools/cached_foundry_client.py

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,7 @@
2727

2828
if TYPE_CHECKING:
2929
import pandas as pd
30+
import polars as pl
3031
import pyspark.sql
3132

3233
from foundry_dev_tools.utils import api_types
@@ -185,7 +186,7 @@ def _download_dataset_to_cache_dir(
185186

186187
def save_dataset(
187188
self,
188-
df: pd.DataFrame | pyspark.sql.DataFrame,
189+
df: pd.DataFrame | pyspark.sql.DataFrame | pl.DataFrame,
189190
dataset_path_or_rid: str,
190191
branch: str = "master",
191192
exists_ok: bool = False,
@@ -198,8 +199,8 @@ def save_dataset(
198199
Creates SNAPSHOT transactions by default.
199200
200201
Args:
201-
df (:external+pandas:py:class:`pandas.DataFrame` | :external+spark:py:class:`pyspark.sql.DataFrame`): A
202-
pyspark or pandas DataFrame to upload
202+
df (:external+pandas:py:class:`pandas.DataFrame` | :external+polars:py:class:`polars.DataFrame` | :external+spark:py:class:`pyspark.sql.DataFrame`):
203+
A pyspark, pandas or polars DataFrame to upload
203204
dataset_path_or_rid (str): Path or Rid of the dataset in which the object should be stored.
204205
branch (str): Branch of the dataset in which the object should be stored
205206
exists_ok (bool): By default, this method creates a new dataset.
@@ -217,7 +218,7 @@ def save_dataset(
217218
ValueError: when dataframe is None
218219
ValueError: when branch is None
219220
220-
"""
221+
""" # noqa: E501
221222
if df is None:
222223
msg = "Please provide a spark or pandas dataframe object with parameter 'df'"
223224
raise ValueError(msg)
@@ -227,6 +228,7 @@ def save_dataset(
227228

228229
with tempfile.TemporaryDirectory() as path:
229230
from foundry_dev_tools._optional.pandas import pd
231+
from foundry_dev_tools._optional.polars import pl
230232

231233
if not pd.__fake__ and isinstance(df, pd.DataFrame):
232234
df.to_parquet(
@@ -235,6 +237,12 @@ def save_dataset(
235237
compression="snappy",
236238
flavor="spark",
237239
)
240+
elif not pl.__fake__ and isinstance(df, pl.DataFrame):
241+
df.write_parquet(
242+
os.sep.join([path + "/dataset.parquet"]), # noqa: PTH118
243+
use_pyarrow=True,
244+
compression="snappy",
245+
)
238246
else:
239247
df.write.format("parquet").option("compression", "snappy").save(path=path, mode="overwrite")
240248

tests/unit/test_cached_foundry_client.py

Lines changed: 17 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
from unittest.mock import MagicMock, patch
66

77
import pandas as pd
8+
import polars as pl
89
import pytest
910
from pandas.testing import assert_frame_equal
1011
from pyspark.sql import SparkSession
@@ -34,22 +35,34 @@
3435
)
3536
@patch(API + ".infer_dataset_schema")
3637
@patch(API + ".upload_dataset_schema")
37-
def test_save_pandas(upload_dataset_schema, infer_dataset_schema, save_objects, temporary_directory, test_context_mock):
38+
@pytest.mark.parametrize(
39+
"df",
40+
[
41+
pd.DataFrame(data={"col1": [1, 2], "col2": [3, 4]}),
42+
pl.DataFrame(data={"col1": [1, 2], "col2": [3, 4]}),
43+
],
44+
)
45+
def test_save_pandas(
46+
upload_dataset_schema, infer_dataset_schema, save_objects, temporary_directory, test_context_mock, df
47+
):
3848
mtp = "mock_tmp_path"
3949
dsn = "dataset.parquet"
4050
dsp = mtp + "/" + dsn
4151
temporary_directory.return_value.__enter__.return_value = mtp
4252
fdt = CachedFoundryClient(ctx=test_context_mock)
43-
df = pd.DataFrame(data={"col1": [1, 2], "col2": [3, 4]})
44-
with patch.object(df, "to_parquet") as pd_to_parquet:
53+
54+
method_to_patch = "to_parquet" if isinstance(df, pd.DataFrame) else "write_parquet"
55+
56+
with patch.object(df, method_to_patch) as df_to_parquet:
4557
dataset_rid, transaction_id = fdt.save_dataset(
4658
df,
4759
dataset_path_or_rid=DATASET_PATH,
4860
branch="master",
4961
exists_ok=True,
5062
mode="SNAPSHOT",
5163
)
52-
assert pd_to_parquet.call_args[0][0] == dsp
64+
65+
assert df_to_parquet.call_args[0][0] == dsp
5366

5467
assert save_objects.call_args[0][0] == {"spark/" + dsn: Path(dsp)}
5568
assert save_objects.call_args[0][1] == DATASET_PATH

0 commit comments

Comments
 (0)