Skip to content

Commit d967239

Browse files
manav252JohnnyWilson16Rohith27216
authored
Add progress callback for clean pipeline (#114)
* Add progress callback for clean pipeline * Align progress callback event schema * Update src/freshdata/cleaner.py Co-authored-by: Rohith <mamidalarohith7@gmail.com> --------- Co-authored-by: Johnny Wilson Dougherty <johnnydougherty09@gmail.com> Co-authored-by: Rohith <mamidalarohith7@gmail.com>
1 parent 331892f commit d967239

4 files changed

Lines changed: 90 additions & 5 deletions

File tree

‎src/freshdata/api.py‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -226,8 +226,9 @@ def clean(
226226
``strategy`` (``"balanced"`` default / ``"aggressive"`` / ``"conservative"``),
227227
``missing_threshold_low``/``_medium``/``_high``, ``duplicate_threshold``,
228228
``outlier_method``, ``outlier_action``, ``preserve_original``, ``verbose``,
229-
``preserve_columns``, ``target_column``, ``duplicate_keep``, ``impute``,
230-
``outliers``. Unknown names raise :class:`TypeError`.
229+
``progress_callback``, ``preserve_columns``, ``target_column``,
230+
``duplicate_keep``, ``impute``, ``outliers``. Unknown names raise
231+
:class:`TypeError`.
231232
232233
Examples
233234
--------

‎src/freshdata/cleaner.py‎

Lines changed: 42 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44

55
import dataclasses
66
import time
7-
from collections.abc import Mapping
7+
from collections.abc import Callable, Mapping
88

99
import pandas as pd
1010

@@ -23,6 +23,8 @@
2323
from .steps.prune import drop_constant_columns, drop_empty_columns, drop_empty_rows
2424
from .steps.strings import clean_strings
2525

26+
ProgressCallback = Callable[[dict[str, object]], None]
27+
2628

2729
def _validate_input(df: object, config: CleanConfig) -> pd.DataFrame:
2830
if isinstance(df, pd.Series):
@@ -42,9 +44,27 @@ def _validate_input(df: object, config: CleanConfig) -> pd.DataFrame:
4244
return frame
4345

4446

47+
def _emit_progress(
48+
callback: ProgressCallback | None,
49+
step: str,
50+
status: str,
51+
frame: pd.DataFrame,
52+
) -> None:
53+
if callback is None:
54+
return
55+
callback(
56+
{
57+
"step": step,
58+
"status": status,
59+
"rows": len(frame),
60+
"columns": frame.shape[1],
61+
}
62+
)
63+
64+
4565
def run_pipeline(
4666
df: pd.DataFrame,
47-
config: CleanConfig,
67+
def run_pipeline( # noqa: PLR0915 - fixed-order pipeline orchestration
4868
*,
4969
memory: object | None = None,
5070
profile: object | None = None,
@@ -69,6 +89,8 @@ def run_pipeline(
6989
forwarded so the profile backend can replay learned value maps.
7090
"""
7191
df = _validate_input(df, config)
92+
progress_callback = config.progress_callback
93+
_emit_progress(progress_callback, "input", "after", df)
7294
report = CleanReport(
7395
rows_before=len(df),
7496
cols_before=df.shape[1],
@@ -84,10 +106,12 @@ def run_pipeline(
84106
from .context import apply_policy_to_config # noqa: PLC0415
85107

86108
config = apply_policy_to_config(config, df=df, report=report)
109+
_emit_progress(progress_callback, "context", "after", df)
87110

88111
out = df.copy(deep=False) if config.preserve_original else df
89112
if config.column_names:
90113
out = normalize_column_names(out, report)
114+
_emit_progress(progress_callback, "column_names", "after", out)
91115

92116
# Hard protected-column guard (context policy / mutable=False): fold the
93117
# protected set into preserve_columns so drop/impute logic honors it, and
@@ -110,44 +134,60 @@ def run_pipeline(
110134
else:
111135
guard_snapshot = {}
112136
out = clean_strings(out, config, report)
137+
_emit_progress(progress_callback, "strings", "after", out)
113138
if config.drop_empty_columns:
114139
out = drop_empty_columns(out, report, config)
140+
_emit_progress(progress_callback, "empty_columns", "after", out)
115141
if config.drop_empty_rows:
116142
out = drop_empty_rows(out, report)
143+
_emit_progress(progress_callback, "empty_rows", "after", out)
117144
if config.fix_dtypes:
118145
out = fix_dtypes(out, config, report)
146+
_emit_progress(progress_callback, "dtypes", "after", out)
119147
if config.drop_constant_columns:
120148
out = drop_constant_columns(out, config, report)
149+
_emit_progress(progress_callback, "constant_columns", "after", out)
121150
if config.drop_duplicates:
122151
out = drop_duplicate_rows(out, config, report)
152+
_emit_progress(progress_callback, "duplicates", "after", out)
123153
if config.semantic_enabled:
124154
# Semantic cleaning runs after representation repair and before the
125155
# statistical engine, so missing/outlier logic sees repaired values.
126156
# Lazily imported to keep ``import freshdata`` light.
127157
from .semantic.apply import run_semantic # noqa: PLC0415
128158

129159
out = run_semantic(out, config, report, memory=memory, profile=profile)
160+
_emit_progress(progress_callback, "semantic", "after", out)
130161
if config.engine_mode is not None:
131162
cache = build_engine_cache(out, config)
163+
_emit_progress(progress_callback, "engine_cache", "after", out)
132164
out = auto_missing(
133165
out, config, report, contexts=cache.contexts, numeric_corr=cache.numeric_corr
134166
)
167+
_emit_progress(progress_callback, "engine_missing", "after", out)
135168
out = auto_outliers(out, config, report, contexts=cache.contexts)
169+
_emit_progress(progress_callback, "engine_outliers", "after", out)
136170
out = impute_missing(out, config, report)
171+
_emit_progress(progress_callback, "missing", "after", out)
137172
out = handle_outliers(out, config, report)
173+
_emit_progress(progress_callback, "outliers", "after", out)
138174
out = optimize_memory(out, config, report)
175+
_emit_progress(progress_callback, "memory", "after", out)
139176
if guard_snapshot:
140177
# Physical byte-identity check, before reset_index so row survivors
141178
# can still be aligned by their original index labels.
142179
verify_protected(out, guard_snapshot, report)
180+
_emit_progress(progress_callback, "protected_columns", "after", out)
143181
if config.reset_index:
144182
out = out.reset_index(drop=True)
183+
_emit_progress(progress_callback, "index", "after", out)
145184

146185
report.rows_after = len(out)
147186
report.cols_after = out.shape[1]
148187
report.memory_after = memory_bytes(out)
149188
report.missing_after = int(out.isna().sum().sum())
150189
report.duration_seconds = time.perf_counter() - started
190+
_emit_progress(progress_callback, "complete", "after", out)
151191
return out, report
152192

153193

‎src/freshdata/config.py‎

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
import dataclasses
1212
import difflib
1313
import warnings
14-
from collections.abc import Mapping
14+
from collections.abc import Callable, Mapping
1515
from dataclasses import dataclass
1616

1717
_STRATEGY_CHOICES = ("conservative", "balanced", "aggressive", "auto")
@@ -130,6 +130,9 @@ class CleanConfig:
130130
preserve_original: bool = True
131131
#: Print a one-line cleaning summary (plus warnings) after each clean.
132132
verbose: bool = True
133+
#: Optional hook called with a small progress event after each enabled
134+
#: pipeline stage.
135+
progress_callback: Callable[[dict[str, object]], None] | None = None
133136
#: Columns that must never be dropped by the engine (post-rename names).
134137
preserve_columns: tuple[str, ...] = ()
135138
#: The label/target column; never modified by the engine. Columns named
@@ -303,6 +306,8 @@ def __post_init__(self) -> None:
303306
raise ValueError(f"outlier_factor must be > 0, got {self.outlier_factor!r}")
304307
if self.sample_size < 1:
305308
raise ValueError(f"sample_size must be >= 1, got {self.sample_size!r}")
309+
if self.progress_callback is not None and not callable(self.progress_callback):
310+
raise TypeError("progress_callback must be callable")
306311
self._validate_semantic()
307312
self._validate_context()
308313
extra = _coerce_str_tuple(self.extra_sentinels)

‎tests/test_api.py‎

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,45 @@ def test_clean_report_tuple(messy):
3131
assert len(report) > 0
3232

3333

34+
def test_clean_progress_callback_reports_pipeline_stages(messy):
35+
events = []
36+
37+
out = fd.clean(messy, verbose=False, progress_callback=events.append)
38+
39+
assert isinstance(out, pd.DataFrame)
40+
assert events
41+
assert all({"step", "status", "rows", "columns"} <= set(event) for event in events)
42+
43+
steps = [(event["step"], event["status"]) for event in events]
44+
assert ("input", "after") in steps
45+
assert ("column_names", "after") in steps
46+
assert ("strings", "after") in steps
47+
assert ("dtypes", "after") in steps
48+
assert ("duplicates", "after") in steps
49+
assert ("engine_missing", "after") in steps
50+
assert ("engine_outliers", "after") in steps
51+
assert steps[-1] == ("complete", "after")
52+
assert events[-1]["rows"] == len(out)
53+
assert events[-1]["columns"] == out.shape[1]
54+
55+
56+
def test_cleaner_progress_callback_reports_pipeline_stages(messy):
57+
events = []
58+
cleaner = fd.Cleaner(verbose=False, progress_callback=events.append)
59+
60+
out = cleaner.clean(messy)
61+
62+
assert isinstance(out, pd.DataFrame)
63+
assert cleaner.report_ is not None
64+
assert [event["step"] for event in events][-1] == "complete"
65+
66+
67+
def test_clean_progress_callback_must_be_callable():
68+
df = pd.DataFrame({"a": [1]})
69+
with pytest.raises(TypeError, match="progress_callback"):
70+
fd.clean(df, progress_callback="not-callable")
71+
72+
3473
def test_clean_rejects_non_dataframe():
3574
with pytest.raises(TypeError, match="DataFrame"):
3675
fd.clean([1, 2, 3])

0 commit comments

Comments
 (0)