Skip to content
Merged
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 coded-flows/coded_flows/utils/media.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import pyarrow.parquet as pq
import numpy as np
import pandas as pd
import polars as pl
from typing import Union, List, Dict, Any
from io import BytesIO
from PIL import Image
Expand Down Expand Up @@ -58,19 +59,36 @@ def _save_arrow_table_to_json(table: pa.Table, filename: str = None) -> str:
records = table.to_pylist()
with open(file_path, "w") as f:
json.dump(records, f, separators=(",", ":"))
return file_path


def _save_polars_to_json(df: pl.DataFrame, filename: str = None) -> str:
random_filename = f"cfdata_{filename if filename else uuid.uuid4().hex}.json"
temp_dir = os.path.join(tempfile.gettempdir(), "coded-flows-media")
os.makedirs(temp_dir, exist_ok=True)
file_path = os.path.join(temp_dir, random_filename)
df.write_json(file_path)
return file_path


# List, DataSeries, NDArray, DataRecords, DataFrame, Arrow
# List, DataSeries, NDArray, DataRecords, DataFrame, Arrow, Polars
def save_data_to_json(
*data_args: Union[
pd.DataFrame, pd.Series, pa.Table, np.ndarray, List[Dict[str, Any]], List[Any]
pd.DataFrame,
pd.Series,
pa.Table,
np.ndarray,
List[Dict[str, Any]],
List[Any],
pl.DataFrame,
pl.Series,
pl.LazyFrame,
],
labels: List[str] = [],
is_table: bool = False,
filename: str = None,
) -> str:

labels = ["values"] if is_table else labels

if not is_table and len(data_args) != len(labels):
Expand All @@ -86,6 +104,8 @@ def save_data_to_json(
if is_table and (
isinstance(data, pd.DataFrame)
or isinstance(data, pa.Table)
or isinstance(data, pl.DataFrame)
or isinstance(data, pl.LazyFrame)
or (
isinstance(data, list) and all(isinstance(item, dict) for item in data[:50])
)
Expand All @@ -95,6 +115,11 @@ def save_data_to_json(
table_df = data
elif isinstance(data, pa.Table):
return _save_arrow_table_to_json(data, filename)
elif isinstance(data, pl.DataFrame):
return _save_polars_to_json(data, filename)
elif isinstance(data, pl.LazyFrame):
# Collect LazyFrame to DataFrame first
return _save_polars_to_json(data.collect(), filename)
else:
table_df = pd.DataFrame.from_records(data)

Expand All @@ -112,6 +137,17 @@ def save_data_to_json(
if label not in data.column_names:
raise ValueError(f"Label '{label}' not found in Arrow table columns.")
col_data = data.column(label).to_pylist()
elif isinstance(data, pl.DataFrame):
if label not in data.columns:
raise ValueError(f"Label '{label}' not found in DataFrame columns.")
col_data = data[label].to_list()
elif isinstance(data, pl.LazyFrame):
collected_data = data.collect()
if label not in collected_data.columns:
raise ValueError(f"Label '{label}' not found in DataFrame columns.")
col_data = collected_data[label].to_list()
elif isinstance(data, pl.Series):
col_data = data.to_list()
elif isinstance(data, pd.Series):
col_data = data.values
elif isinstance(data, np.ndarray):
Expand Down Expand Up @@ -142,7 +178,15 @@ def save_data_to_json(

def save_data_to_parquet(
data: Union[
pd.DataFrame, pd.Series, pa.Table, np.ndarray, List[Dict[str, Any]], List[Any]
pd.DataFrame,
pd.Series,
pl.DataFrame,
pl.Series,
pl.LazyFrame,
pa.Table,
np.ndarray,
List[Dict[str, Any]],
List[Any],
],
filename=None,
) -> str:
Expand All @@ -155,6 +199,9 @@ def save_data_to_parquet(
if (
isinstance(data, pd.DataFrame)
or isinstance(data, pd.Series)
or isinstance(data, pl.DataFrame)
or isinstance(data, pl.LazyFrame)
or isinstance(data, pl.Series)
or isinstance(data, pa.Table)
or (
isinstance(data, list) and all(isinstance(item, dict) for item in data[:50])
Expand All @@ -169,9 +216,18 @@ def save_data_to_parquet(
data.to_frame().to_parquet(
file_path, row_group_size=50000, index=False, engine="pyarrow"
)
elif isinstance(data, pl.DataFrame):
data.write_parquet(file_path, row_group_size=50000, use_pyarrow=True)
elif isinstance(data, pl.LazyFrame):
data.collect().write_parquet(
file_path, row_group_size=50000, use_pyarrow=True
)
elif isinstance(data, pl.Series):
data.to_frame().write_parquet(
file_path, row_group_size=50000, use_pyarrow=True
)
elif isinstance(data, pa.Table):
pq.write_table(data, file_path, row_group_size=50000)

else:
table = pa.Table.from_pylist(data)
pq.write_table(table, file_path, row_group_size=50000)
Expand Down
2 changes: 1 addition & 1 deletion coded-flows/pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[tool.poetry]
name = "coded-flows"
version = "0.6.4"
version = "0.6.5"
description = "Various utilities for Coded Flows"
authors = ["COLOR CODED CODES <contact@colorcoded.codes>"]
readme = "README.md"
Expand Down
115 changes: 109 additions & 6 deletions coded-flows/tests/test_data_save.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import os
import pytest
import pandas as pd
import polars as pl
import pandas.testing as pdt
import pyarrow as pa
import pyarrow.parquet as pq
Expand All @@ -14,10 +15,11 @@ def test_valid_inputs():
arr = np.array([20, 21, 22])
data_records = [{"key1": 30}, {"key1": 31}, {"key1": 32}]
arrow_table = pa.table({"col1": [1, 2, 3], "col2": ["a", "b", "c"]})
dfp = pl.DataFrame({"uu": [1, 2, 3]})

labels = ["x", "y", "z", "w", "col1"]
labels = ["x", "y", "z", "w", "col1", "uu"]
file_path = save_data_to_json(
df, series, arr, data_records, arrow_table, labels=labels
df, series, arr, data_records, arrow_table, dfp, labels=labels
)

# Check file existence
Expand All @@ -27,7 +29,7 @@ def test_valid_inputs():
with open(file_path, "r") as f:
json_content = f.read()

expected_content = '[{"x":1,"y":100,"z":20,"w":null,"col1":1},{"x":2,"y":200,"z":21,"w":null,"col1":2},{"x":3,"y":300,"z":22,"w":null,"col1":3}]'
expected_content = '[{"x":1,"y":100,"z":20,"w":null,"col1":1,"uu":1},{"x":2,"y":200,"z":21,"w":null,"col1":2,"uu":2},{"x":3,"y":300,"z":22,"w":null,"col1":3,"uu":3}]'
assert json_content.strip() == expected_content, "JSON content mismatch."
os.remove(file_path)

Expand All @@ -51,6 +53,23 @@ def test_missing_column_in_dataframe():
save_data_to_json(df, labels=labels)


def test_missing_column_in_polars_dataframe():
df = pl.DataFrame({"x": [1, 2, 3]})
labels = ["z"]

with pytest.raises(ValueError, match="Label 'z' not found in DataFrame columns."):
save_data_to_json(df, labels=labels)


def test_missing_column_in_polars_lazyframe():
df = pl.DataFrame({"x": [1, 2, 3]})
df_lazy = df.lazy()
labels = ["z"]

with pytest.raises(ValueError, match="Label 'z' not found in DataFrame columns."):
save_data_to_json(df_lazy, labels=labels)


def test_invalid_numpy_array():
arr = np.array([[1, 2], [3, 4]]) # 2D array, invalid
labels = ["x"]
Expand Down Expand Up @@ -90,15 +109,18 @@ def test_variable_lengths():
def test_series_input():
series1 = pd.Series([1, 2, 3], name="x")
series2 = pd.Series([10, 20, 30], name="y")
labels = ["x", "y"]
series3 = pl.Series("z", [100, 200, 300])
labels = ["x", "y", "z"]

file_path = save_data_to_json(series1, series2, labels=labels)
file_path = save_data_to_json(series1, series2, series3, labels=labels)

# Check JSON content
with open(file_path, "r") as f:
json_content = f.read()

expected_content = '[{"x":1,"y":10},{"x":2,"y":20},{"x":3,"y":30}]'
expected_content = (
'[{"x":1,"y":10,"z":100},{"x":2,"y":20,"z":200},{"x":3,"y":30,"z":300}]'
)
assert (
json_content.strip() == expected_content
), "JSON content mismatch for Series inputs."
Expand All @@ -124,6 +146,25 @@ def test_save_data_to_json_is_table_true_with_dataframe():
os.remove(json_path)


def test_save_data_to_json_is_table_true_with_polars_dataframe():
# Create a sample DataFrame
df = pl.DataFrame({"col1": [1, 2, 3], "col2": ["a", "b", "c"]})

# Call the function with `is_table=True`
json_path = save_data_to_json(df, is_table=True)

# Assert the file exists
assert os.path.exists(json_path), "JSON file was not created."

# Assert the file content matches the DataFrame content
with open(json_path, "r") as f:
saved_data = pl.read_json(f)
assert df.equals(saved_data)

# Clean up
os.remove(json_path)


def test_save_data_to_json_is_table_true_with_arrow_table():
# Create a sample Arrow Table
table = pa.table({"col1": [1, 2, 3], "col2": ["a", "b", "c"]})
Expand Down Expand Up @@ -226,6 +267,20 @@ def test_save_data_to_json_is_table_true_series():
os.remove(json_path)


def test_save_data_to_json_is_table_true_polars_series():
# Series input
series = pl.Series("col1", [1, 2, 3])
json_path = save_data_to_json(series, is_table=True)
assert os.path.exists(json_path)
with open(json_path, "r") as f:
saved_data = pl.read_json(f)
expected_df = pl.DataFrame({"values": series})
assert expected_df.equals(saved_data)

# Clean up
os.remove(json_path)


def test_dataframe_to_parquet():
df = pd.DataFrame({"x": [1, 2, 3], "z": [7, 8, 9]})

Expand All @@ -241,6 +296,37 @@ def test_dataframe_to_parquet():
os.remove(file_path)


def test_polars_dataframe_to_parquet():
df = pl.DataFrame({"x": [1, 2, 3], "z": [7, 8, 9]})

file_path = save_data_to_parquet(df)

# Check file existence
assert os.path.exists(file_path), "Output file does not exist."

# Read back from Parquet into DataFrame
df_loaded = pl.read_parquet(file_path, use_pyarrow=True)

assert df.equals(df_loaded)
os.remove(file_path)


def test_polars_lazyframe_to_parquet():
df = pl.DataFrame({"x": [1, 2, 3], "z": [7, 8, 9]})
dfl = df.lazy()

file_path = save_data_to_parquet(dfl)

# Check file existence
assert os.path.exists(file_path), "Output file does not exist."

# Read back from Parquet into DataFrame
df_loaded = pl.read_parquet(file_path, use_pyarrow=True)

assert df.equals(df_loaded)
os.remove(file_path)


def test_arrow_table_to_parquet():
table = pa.table({"x": [1, 2, 3], "z": [7, 8, 9]})

Expand Down Expand Up @@ -310,3 +396,20 @@ def test_pandas_series_to_parquet():

pdt.assert_frame_equal(expected_df.sort_index(axis=1), df_loaded.sort_index(axis=1))
os.remove(file_path)


def test_polars_series_to_parquet():
polars_series = pl.Series("values", [100, 200, 300, 400, 500])
file_path = save_data_to_parquet(polars_series)

# Check file existence
assert os.path.exists(file_path), "Output file does not exist."

# Read back from Parquet into DataFrame
df_loaded = pl.read_parquet(file_path, use_pyarrow=True)

# Convert original Series to DataFrame for comparison
expected_df = polars_series.to_frame()

assert expected_df.equals(df_loaded)
os.remove(file_path)