diff --git a/packages/data-designer-config/src/data_designer/config/preview_results.py b/packages/data-designer-config/src/data_designer/config/preview_results.py index d7804e0da..5805c50d3 100644 --- a/packages/data-designer-config/src/data_designer/config/preview_results.py +++ b/packages/data-designer-config/src/data_designer/config/preview_results.py @@ -3,7 +3,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any from data_designer.config.analysis.dataset_profiler import DatasetProfilerResults from data_designer.config.config_builder import DataDesignerConfigBuilder @@ -23,6 +23,7 @@ def __init__( dataset: pd.DataFrame | None = None, analysis: DatasetProfilerResults | None = None, processor_artifacts: dict[str, list[dict]] | None = None, + task_traces: list[Any] | None = None, ): """Creates a new instance with results from a Data Designer preview run. @@ -32,9 +33,11 @@ def __init__( dataset: Dataset of the preview run. analysis: Analysis of the preview run. processor_artifacts: Artifacts generated by the processors. + task_traces: Async scheduler task traces (when DATA_DESIGNER_ASYNC_TRACE=1). """ self.dataset: pd.DataFrame | None = dataset self.analysis: DatasetProfilerResults | None = analysis self.processor_artifacts: dict[str, list[dict]] | None = processor_artifacts self.dataset_metadata: DatasetMetadata | None = dataset_metadata + self.task_traces: list[Any] | None = task_traces self._config_builder = config_builder diff --git a/packages/data-designer-engine/src/data_designer/engine/dataset_builders/async_scheduler.py b/packages/data-designer-engine/src/data_designer/engine/dataset_builders/async_scheduler.py index 067795363..c62ff60d9 100644 --- a/packages/data-designer-engine/src/data_designer/engine/dataset_builders/async_scheduler.py +++ b/packages/data-designer-engine/src/data_designer/engine/dataset_builders/async_scheduler.py @@ -15,6 +15,7 @@ import data_designer.lazy_heavy_imports as lazy from data_designer.config.column_configs import GenerationStrategy from data_designer.engine.context import current_row_group +from data_designer.engine.dataset_builders.errors import DatasetGenerationError from data_designer.engine.dataset_builders.multi_column_configs import MultiColumnConfig from data_designer.engine.dataset_builders.utils.async_progress_reporter import ( DEFAULT_REPORT_INTERVAL, @@ -216,6 +217,10 @@ def _setup_async_progress_reporter( progress_bar=self._progress_bar, ) + @property + def active_worker_count(self) -> int: + return sum(1 for t in self._worker_tasks if not t.done()) + def _spawn_worker(self, coro: Coroutine[Any, Any, None]) -> asyncio.Task: """Create a tracked worker task that auto-removes itself on completion.""" task = asyncio.create_task(coro) @@ -265,33 +270,32 @@ async def run(self) -> None: # Launch admission as a background task so it interleaves with dispatch. admission_task = asyncio.create_task(self._admit_row_groups()) + dispatch_error: BaseException | None = None try: # Main dispatch loop await self._main_dispatch_loop(seed_cols, has_pre_batch, all_columns) - - # Cancel admission if still running + except BaseException as exc: + dispatch_error = exc + raise + finally: + # Always cancel admission + drain in-flight workers, regardless + # of how the dispatch loop exited (normal, early shutdown, + # CancelledError, or processor failure). if not admission_task.done(): admission_task.cancel() with contextlib.suppress(asyncio.CancelledError): await admission_task + await asyncio.shield(self._cancel_workers()) - if self._reporter: - self._reporter.log_final() - - if self._rg_states: - incomplete = list(self._rg_states) - logger.error( - f"Scheduler exited with {len(self._rg_states)} unfinished row group(s): {incomplete}. " - "These row groups were not checkpointed." - ) + if self._reporter: + self._reporter.log_final() - except asyncio.CancelledError: - if not admission_task.done(): - admission_task.cancel() - with contextlib.suppress(asyncio.CancelledError): - await admission_task - await asyncio.shield(self._cancel_workers()) - raise + if self._rg_states and dispatch_error is None: + incomplete = list(self._rg_states) + logger.error( + f"Scheduler exited with {len(self._rg_states)} unfinished row group(s): {incomplete}. " + "These row groups were not checkpointed." + ) async def _main_dispatch_loop( self, @@ -500,29 +504,26 @@ def _checkpoint_completed_row_groups(self, all_columns: list[str]) -> None: if self._tracker.is_row_group_complete(rg_id, state.size, all_columns) ] for rg_id, rg_size in completed: - dropped = False try: - del self._rg_states[rg_id] if self._on_before_checkpoint: try: self._on_before_checkpoint(rg_id, rg_size) - except Exception: - # Post-batch is mandatory; drop rather than checkpoint unprocessed data. - logger.error( - f"on_before_checkpoint failed for row group {rg_id}, dropping row group.", - exc_info=True, - ) - self._drop_row_group(rg_id, rg_size) - if self._buffer_manager: - self._buffer_manager.free_row_group(rg_id) - dropped = True + except DatasetGenerationError: + raise + except Exception as exc: + raise DatasetGenerationError( + f"Post-batch processor failed for row group {rg_id}: {exc}" + ) from exc + # Remove from tracking only after the callback succeeds. + del self._rg_states[rg_id] # If all rows were dropped (e.g. seed failure), free instead of finalizing - if not dropped and all(self._tracker.is_dropped(rg_id, ri) for ri in range(rg_size)): + if all(self._tracker.is_dropped(rg_id, ri) for ri in range(rg_size)): if self._buffer_manager: self._buffer_manager.free_row_group(rg_id) - dropped = True - if not dropped and self._on_finalize_row_group is not None: + elif self._on_finalize_row_group is not None: self._on_finalize_row_group(rg_id) + except DatasetGenerationError: + raise except Exception: logger.error(f"Failed to checkpoint row group {rg_id}.", exc_info=True) finally: @@ -543,19 +544,19 @@ def _run_seeds_complete_check(self, seed_cols: frozenset[str]) -> None: if self._on_seeds_complete: try: self._on_seeds_complete(rg_id, state.size) - # The callback may drop rows (e.g. pre-batch filtering). - # Record skipped tasks for any newly-dropped rows so - # progress reporting stays accurate. - if self._reporter: - for ri in range(state.size): - if self._tracker.is_dropped(rg_id, ri): - self._record_skipped_tasks_for_row(rg_id, ri) - except Exception: - logger.warning( - f"Pre-batch processor failed for row group {rg_id}, skipping.", - exc_info=True, - ) - self._drop_row_group(rg_id, state.size) + except DatasetGenerationError: + raise + except Exception as exc: + raise DatasetGenerationError( + f"Pre-batch processor failed for row group {rg_id}: {exc}" + ) from exc + # The callback may drop rows (e.g. pre-batch filtering). + # Record skipped tasks for any newly-dropped rows so + # progress reporting stays accurate. + if self._reporter: + for ri in range(state.size): + if self._tracker.is_dropped(rg_id, ri): + self._record_skipped_tasks_for_row(rg_id, ri) def _drop_row(self, row_group: int, row_index: int, *, exclude_columns: set[str] | None = None) -> None: if self._tracker.is_dropped(row_group, row_index): diff --git a/packages/data-designer-engine/src/data_designer/engine/dataset_builders/dataset_builder.py b/packages/data-designer-engine/src/data_designer/engine/dataset_builders/dataset_builder.py index 182f0438e..c76be0980 100644 --- a/packages/data-designer-engine/src/data_designer/engine/dataset_builders/dataset_builder.py +++ b/packages/data-designer-engine/src/data_designer/engine/dataset_builders/dataset_builder.py @@ -9,6 +9,7 @@ import os import time import uuid +import warnings from pathlib import Path from typing import TYPE_CHECKING, Any, Callable @@ -58,6 +59,7 @@ if TYPE_CHECKING: import pandas as pd + from data_designer.config.run_config import RunConfig from data_designer.engine.column_generators.generators.base import ColumnGeneratorWithModelRegistry from data_designer.engine.dataset_builders.utils.task_model import TaskTrace from data_designer.engine.models.usage import ModelUsageStats @@ -91,6 +93,10 @@ _CLIENT_VERSION: str = get_library_version() +def _is_async_trace_enabled(settings: RunConfig) -> bool: + return settings.async_trace or os.environ.get("DATA_DESIGNER_ASYNC_TRACE", "0") == "1" + + class DatasetBuilder: def __init__( self, @@ -106,6 +112,7 @@ def __init__( self._task_traces: list[TaskTrace] = [] self._registry = registry or DataDesignerRegistry() self._graph: ExecutionGraph | None = None + self._use_async: bool = DATA_DESIGNER_ASYNC_ENGINE self._data_designer_config = compile_data_designer_config(data_designer_config, resource_provider) self._column_configs = compile_dataset_builder_column_configs(self._data_designer_config) @@ -185,8 +192,8 @@ def build( start_time = time.perf_counter() buffer_size = self._resource_provider.run_config.buffer_size - if DATA_DESIGNER_ASYNC_ENGINE: - self._validate_async_compatibility() + self._use_async = DATA_DESIGNER_ASYNC_ENGINE and self._resolve_async_compatibility() + if self._use_async: self._build_async(generators, num_records, buffer_size, on_batch_complete) else: group_id = uuid.uuid4().hex @@ -218,8 +225,8 @@ def build_preview(self, *, num_records: int) -> pd.DataFrame: generators, self._graph = self._initialize_generators_and_graph() start_time = time.perf_counter() - if DATA_DESIGNER_ASYNC_ENGINE: - self._validate_async_compatibility() + self._use_async = DATA_DESIGNER_ASYNC_ENGINE and self._resolve_async_compatibility() + if self._use_async: dataset = self._build_async_preview(generators, num_records) else: group_id = uuid.uuid4().hex @@ -236,11 +243,15 @@ def _build_async_preview(self, generators: list[ColumnGenerator], num_records: i """Async preview path - single row group, no disk writes, returns in-memory DataFrame.""" logger.info("⚡ DATA_DESIGNER_ASYNC_ENGINE is enabled - using async task-queue preview") + settings = self._resource_provider.run_config + trace_enabled = _is_async_trace_enabled(settings) + scheduler, buffer_manager = self._prepare_async_run( generators, num_records, buffer_size=num_records, run_post_batch_in_scheduler=False, + trace=trace_enabled, ) loop = ensure_async_engine_loop() @@ -256,15 +267,23 @@ def _build_async_preview(self, generators: list[ColumnGenerator], num_records: i buffer_manager.free_row_group(0) return dataset - def _validate_async_compatibility(self) -> None: - """Raise if any column uses allow_resize=True with the async scheduler.""" + def _resolve_async_compatibility(self) -> bool: + """Check if the async engine can be used; auto-fallback to sync if not. + + Returns True if async is usable, False if allow_resize forces sync fallback. + """ offending = [config.name for config in self.single_column_configs if getattr(config, "allow_resize", False)] if offending: - raise DatasetGenerationError( - f"allow_resize=True is not supported with DATA_DESIGNER_ASYNC_ENGINE=1. " - f"Offending column(s): {offending}. Either remove allow_resize=True or " - f"disable the async scheduler." + msg = ( + f"allow_resize=True detected on column(s) {offending}. " + "Falling back to sync engine for this run. " + "allow_resize is deprecated and will be removed in a future release; " + "use workflow chaining instead (see issue #552)." ) + logger.warning(f"⚠️ {msg}") + warnings.warn(msg, DeprecationWarning, stacklevel=4) + return False + return True def _build_async( self, @@ -277,7 +296,7 @@ def _build_async( logger.info("⚡ DATA_DESIGNER_ASYNC_ENGINE is enabled - using async task-queue builder") settings = self._resource_provider.run_config - trace_enabled = settings.async_trace or os.environ.get("DATA_DESIGNER_ASYNC_TRACE", "0") == "1" + trace_enabled = _is_async_trace_enabled(settings) def finalize_row_group(rg_id: int) -> None: def on_complete(final_path: Path | str | None) -> None: @@ -318,6 +337,15 @@ def on_complete(final_path: Path | str | None) -> None: # Write metadata buffer_manager.write_metadata(target_num_records=num_records, buffer_size=buffer_size) + # Surface partial completion + actual = buffer_manager.actual_num_records + if actual < num_records: + pct = actual / num_records * 100 if num_records > 0 else 0 + logger.warning( + f"⚠️ Generated {actual} of {num_records} requested records ({pct:.0f}%). " + "The dataset may be incomplete due to errors or early shutdown." + ) + def _prepare_async_run( self, generators: list[ColumnGenerator], @@ -366,10 +394,10 @@ def _prepare_async_run( buffer_manager = RowGroupBufferManager(self.artifact_storage) # Pre-batch processor callback: runs after seed tasks complete for a row group. - # If it raises, the scheduler drops all rows in the row group (skips it). + # If it raises, the scheduler propagates the error as DatasetGenerationError (fail-fast). def on_seeds_complete(rg_id: int, rg_size: int) -> None: df = buffer_manager.get_dataframe(rg_id) - df = self._processor_runner.run_pre_batch_on_df(df) + df = self._processor_runner.run_pre_batch_on_df(df, strict_row_count=True) buffer_manager.replace_dataframe(rg_id, df) for ri in range(rg_size): if buffer_manager.is_dropped(rg_id, ri) and not tracker.is_dropped(rg_id, ri): @@ -378,7 +406,7 @@ def on_seeds_complete(rg_id: int, rg_size: int) -> None: # Post-batch processor callback: runs after all columns, before finalization. def on_before_checkpoint(rg_id: int, rg_size: int) -> None: df = buffer_manager.get_dataframe(rg_id) - df = self._processor_runner.run_post_batch(df, current_batch_number=rg_id) + df = self._processor_runner.run_post_batch(df, current_batch_number=rg_id, strict_row_count=True) buffer_manager.replace_dataframe(rg_id, df) # Coarse upper bound: sums all registered aliases, not just those used @@ -505,7 +533,7 @@ def _run_cell_by_cell_generator(self, generator: ColumnGenerator) -> None: max_workers = self._resource_provider.run_config.non_inference_max_parallel_workers if isinstance(generator, ColumnGeneratorWithModel): max_workers = generator.inference_parameters.max_parallel_requests - if DATA_DESIGNER_ASYNC_ENGINE: + if self._use_async: logger.info("⚡ Using async engine for concurrent execution") self._fan_out_with_async(generator, max_workers=max_workers) else: diff --git a/packages/data-designer-engine/src/data_designer/engine/dataset_builders/utils/processor_runner.py b/packages/data-designer-engine/src/data_designer/engine/dataset_builders/utils/processor_runner.py index 61bba8919..e8ea468a1 100644 --- a/packages/data-designer-engine/src/data_designer/engine/dataset_builders/utils/processor_runner.py +++ b/packages/data-designer-engine/src/data_designer/engine/dataset_builders/utils/processor_runner.py @@ -72,13 +72,44 @@ def run_pre_batch(self, batch_manager: DatasetBatchManager) -> None: df = self._run_stage(df, ProcessorStage.PRE_BATCH) batch_manager.replace_buffer(df.to_dict(orient="records"), allow_resize=True) - def run_pre_batch_on_df(self, df: pd.DataFrame) -> pd.DataFrame: - """Run PRE_BATCH processors on a DataFrame and return the result.""" - return self._run_stage(df, ProcessorStage.PRE_BATCH) + def run_pre_batch_on_df(self, df: pd.DataFrame, *, strict_row_count: bool = False) -> pd.DataFrame: + """Run PRE_BATCH processors on a DataFrame and return the result. + + Args: + df: Input DataFrame. + strict_row_count: If True, raise ``DatasetProcessingError`` when a + processor changes the row count. Used by the async engine where + row-count changes are not supported. + """ + original_len = len(df) + df = self._run_stage(df, ProcessorStage.PRE_BATCH) + if strict_row_count and len(df) != original_len: + raise DatasetProcessingError( + f"Pre-batch processor changed row count from {original_len} to {len(df)}. " + "Row-count changes in pre-batch processors are not supported with the async engine." + ) + return df - def run_post_batch(self, df: pd.DataFrame, current_batch_number: int | None) -> pd.DataFrame: - """Run process_after_batch() on processors that implement it.""" - return self._run_stage(df, ProcessorStage.POST_BATCH, current_batch_number=current_batch_number) + def run_post_batch( + self, df: pd.DataFrame, current_batch_number: int | None, *, strict_row_count: bool = False + ) -> pd.DataFrame: + """Run process_after_batch() on processors that implement it. + + Args: + df: Input DataFrame. + current_batch_number: Batch index passed to processors. + strict_row_count: If True, raise ``DatasetProcessingError`` when a + processor changes the row count. Used by the async engine where + row-count changes are not supported. + """ + original_len = len(df) + df = self._run_stage(df, ProcessorStage.POST_BATCH, current_batch_number=current_batch_number) + if strict_row_count and len(df) != original_len: + raise DatasetProcessingError( + f"Post-batch processor changed row count from {original_len} to {len(df)}. " + "Row-count changes in post-batch processors are not supported with the async engine." + ) + return df def run_after_generation_on_df(self, df: pd.DataFrame) -> pd.DataFrame: """Run process_after_generation() on a DataFrame (for preview mode).""" diff --git a/packages/data-designer-engine/tests/engine/dataset_builders/test_async_builder_integration.py b/packages/data-designer-engine/tests/engine/dataset_builders/test_async_builder_integration.py index 3ec3f2b42..846089095 100644 --- a/packages/data-designer-engine/tests/engine/dataset_builders/test_async_builder_integration.py +++ b/packages/data-designer-engine/tests/engine/dataset_builders/test_async_builder_integration.py @@ -4,6 +4,7 @@ from __future__ import annotations import math +import warnings from unittest.mock import MagicMock, Mock import pytest @@ -23,7 +24,6 @@ ) from data_designer.engine.dataset_builders.async_scheduler import AsyncTaskScheduler from data_designer.engine.dataset_builders.dataset_builder import DatasetBuilder -from data_designer.engine.dataset_builders.errors import DatasetGenerationError from data_designer.engine.dataset_builders.utils.completion_tracker import CompletionTracker from data_designer.engine.dataset_builders.utils.execution_graph import ExecutionGraph from data_designer.engine.dataset_builders.utils.row_group_buffer import RowGroupBufferManager @@ -75,29 +75,34 @@ def generate(self, data: lazy.pd.DataFrame) -> lazy.pd.DataFrame: @pytest.mark.parametrize( - "configs,should_raise", + "configs,expected", [ pytest.param( [Mock(name="col_a", allow_resize=True), Mock(name="col_b", allow_resize=False)], - True, - id="raises_on_allow_resize", + False, + id="fallback_on_allow_resize", ), pytest.param( [Mock(name="col_a", allow_resize=False), Mock(name="col_b", allow_resize=False)], - False, - id="passes_without_allow_resize", + True, + id="async_without_allow_resize", ), ], ) -def test_validate_async_compatibility(configs: list[Mock], should_raise: bool) -> None: - """Validation rejects allow_resize=True with the async engine.""" +def test_resolve_async_compatibility(configs: list[Mock], expected: bool) -> None: + """allow_resize=True triggers auto-fallback to sync with a deprecation warning.""" builder = Mock(spec=DatasetBuilder) builder.single_column_configs = configs - if should_raise: - with pytest.raises(DatasetGenerationError, match="allow_resize=True"): - DatasetBuilder._validate_async_compatibility(builder) + with warnings.catch_warnings(record=True) as w: + warnings.simplefilter("always") + result = DatasetBuilder._resolve_async_compatibility(builder) + assert result is expected + if not expected: + assert len(w) == 1 + assert issubclass(w[0].category, DeprecationWarning) + assert "allow_resize" in str(w[0].message) else: - DatasetBuilder._validate_async_compatibility(builder) + assert len(w) == 0 # -- _build_async integration test with mock generators ----------------------- @@ -262,3 +267,55 @@ async def test_checkpoint_produces_correct_parquet_calls() -> None: assert storage.write_batch_to_parquet_file.call_count == 2 assert storage.move_partial_result_to_final_file_path.call_count == 2 assert buffer_manager.actual_num_records == 5 + + +# -- Partial completion warning ------------------------------------------------ + + +@pytest.mark.asyncio(loop_scope="session") +async def test_dropped_rows_reduce_actual_record_count() -> None: + """When all rows in a row group are dropped, actual_num_records reflects the shortfall + and write_metadata records the correct actual vs target counts.""" + provider = _mock_provider() + seed_gen = MockSeed(config=_expr_config("seed"), resource_provider=provider) + + configs = [SamplerColumnConfig(name="seed", sampler_type=SamplerType.CATEGORY, params={"values": ["X"]})] + strategies = {"seed": GenerationStrategy.FULL_COLUMN} + gen_map = {"seed": seed_gen} + + graph = ExecutionGraph.create(configs, strategies) + num_records = 6 + row_groups = [(0, 3), (1, 3)] + tracker = CompletionTracker.with_graph(graph, row_groups) + + storage = MagicMock() + storage.dataset_name = "test" + storage.get_file_paths.return_value = {} + storage.write_batch_to_parquet_file.return_value = "/fake.parquet" + storage.move_partial_result_to_final_file_path.return_value = "/fake_final.parquet" + + buffer_manager = RowGroupBufferManager(storage) + + def drop_all_in_rg1(rg_id: int, rg_size: int) -> None: + if rg_id == 1: + for ri in range(rg_size): + tracker.drop_row(rg_id, ri) + buffer_manager.drop_row(rg_id, ri) + + scheduler = AsyncTaskScheduler( + generators=gen_map, + graph=graph, + tracker=tracker, + row_groups=row_groups, + buffer_manager=buffer_manager, + on_finalize_row_group=lambda rg_id: buffer_manager.checkpoint_row_group(rg_id), + on_seeds_complete=drop_all_in_rg1, + ) + await scheduler.run() + + assert buffer_manager.actual_num_records < num_records + + buffer_manager.write_metadata(target_num_records=num_records, buffer_size=3) + written = storage.write_metadata.call_args[0][0] + assert written["actual_num_records"] == buffer_manager.actual_num_records + assert written["target_num_records"] == num_records diff --git a/packages/data-designer-engine/tests/engine/dataset_builders/test_async_scheduler.py b/packages/data-designer-engine/tests/engine/dataset_builders/test_async_scheduler.py index cea3e4ccf..0c6ec4e4d 100644 --- a/packages/data-designer-engine/tests/engine/dataset_builders/test_async_scheduler.py +++ b/packages/data-designer-engine/tests/engine/dataset_builders/test_async_scheduler.py @@ -27,6 +27,7 @@ ) from data_designer.engine.column_generators.generators.custom import CustomColumnGenerator from data_designer.engine.dataset_builders.async_scheduler import AsyncTaskScheduler, build_llm_bound_lookup +from data_designer.engine.dataset_builders.errors import DatasetGenerationError from data_designer.engine.dataset_builders.utils.completion_tracker import CompletionTracker from data_designer.engine.dataset_builders.utils.execution_graph import ExecutionGraph from data_designer.engine.dataset_builders.utils.row_group_buffer import RowGroupBufferManager @@ -611,8 +612,8 @@ async def test_scheduler_non_retryable_seed_failure_no_keyerror_on_downstream() @pytest.mark.asyncio(loop_scope="session") -async def test_scheduler_pre_batch_failure_marks_downstream_tasks_skipped() -> None: - """Pre-batch row-group drops count downstream cell tasks as skipped.""" +async def test_scheduler_pre_batch_failure_raises() -> None: + """Pre-batch processor failure propagates as DatasetGenerationError.""" provider = _mock_provider() configs = [ SamplerColumnConfig(name="seed", sampler_type=SamplerType.CATEGORY, params={"values": ["A"]}), @@ -643,14 +644,8 @@ def fail_pre_batch(row_group: int, row_group_size: int) -> None: num_records=3, buffer_size=3, ) - await scheduler.run() - - for row_index in range(3): - assert tracker.is_dropped(0, row_index) - - assert scheduler._reporter is not None - assert scheduler._reporter._trackers["cell_out"].skipped == 3 - assert scheduler._reporter._trackers["cell_out"].completed == 3 + with pytest.raises(DatasetGenerationError, match="Pre-batch processor failed"): + await scheduler.run() @pytest.mark.asyncio(loop_scope="session") @@ -941,8 +936,8 @@ async def test_scheduler_on_finalize_skips_empty_row_group() -> None: @pytest.mark.asyncio(loop_scope="session") -async def test_scheduler_pre_batch_failure_skips_row_group() -> None: - """Pre-batch processor failure drops all rows in the row group; other row groups continue.""" +async def test_scheduler_pre_batch_failure_propagates_across_row_groups() -> None: + """Pre-batch processor failure propagates even when other row groups exist.""" provider = _mock_provider() seed_gen = MockSeedGenerator(config=_expr_config("seed"), resource_provider=provider) cell_gen = MockCellGenerator(config=_expr_config("cell_out"), resource_provider=provider) @@ -981,12 +976,8 @@ def failing_pre_batch(rg_id: int, rg_size: int) -> None: buffer_manager=buffer_mgr, on_seeds_complete=failing_pre_batch, ) - await scheduler.run() - - # Row group 0: all rows dropped due to pre-batch failure - assert all(tracker.is_dropped(0, ri) for ri in range(3)) - # Row group 1: completed normally - assert tracker.is_row_group_complete(1, 2, ["seed", "cell_out"]) + with pytest.raises(DatasetGenerationError, match="Pre-batch processor failed"): + await scheduler.run() class _SlowSeedGenerator(FromScratchColumnGenerator[ExpressionColumnConfig]): @@ -1838,3 +1829,87 @@ def generate_from_scratch(self, num_records: int) -> lazy.pd.DataFrame: assert row.get("review") is None, f"row {ri}: review should be skipped (seed={seed_val})" else: assert row["review"] == "batch_val", f"row {ri}: review should be generated (seed={seed_val})" + + +# -- Post-batch (on_before_checkpoint) failure propagation -------------------- + + +@pytest.mark.asyncio(loop_scope="session") +async def test_scheduler_post_batch_failure_raises() -> None: + """Post-batch processor failure propagates as DatasetGenerationError.""" + provider = _mock_provider() + configs = [ + SamplerColumnConfig(name="seed", sampler_type=SamplerType.CATEGORY, params={"values": ["A"]}), + LLMTextColumnConfig(name="cell_out", prompt="{{ seed }}", model_alias=MODEL_ALIAS), + ] + strategies = { + "seed": GenerationStrategy.FULL_COLUMN, + "cell_out": GenerationStrategy.CELL_BY_CELL, + } + generators = { + "seed": MockSeedGenerator(config=_expr_config("seed"), resource_provider=provider), + "cell_out": MockCellGenerator(config=_expr_config("cell_out"), resource_provider=provider), + } + + graph = ExecutionGraph.create(configs, strategies) + row_groups = [(0, 3)] + tracker = CompletionTracker.with_graph(graph, row_groups) + + storage = MagicMock() + storage.dataset_name = "test" + storage.get_file_paths.return_value = {} + buffer_mgr = RowGroupBufferManager(storage) + + def fail_post_batch(rg_id: int, rg_size: int) -> None: + raise RuntimeError("post-batch processor exploded") + + scheduler = AsyncTaskScheduler( + generators=generators, + graph=graph, + tracker=tracker, + row_groups=row_groups, + buffer_manager=buffer_mgr, + on_before_checkpoint=fail_post_batch, + ) + with pytest.raises(DatasetGenerationError, match="Post-batch processor failed"): + await scheduler.run() + + +# -- Early shutdown drains workers ------------------------------------------- + + +@pytest.mark.asyncio(loop_scope="session") +async def test_early_shutdown_drains_workers() -> None: + """Workers are cancelled after early shutdown, not left dangling.""" + provider = _mock_provider() + configs = [ + SamplerColumnConfig(name="seed", sampler_type=SamplerType.CATEGORY, params={"values": ["A"]}), + LLMTextColumnConfig(name="fail_col", prompt="{{ seed }}", model_alias=MODEL_ALIAS), + ] + strategies = { + "seed": GenerationStrategy.FULL_COLUMN, + "fail_col": GenerationStrategy.CELL_BY_CELL, + } + generators = { + "seed": MockSeedGenerator(config=_expr_config("seed"), resource_provider=provider), + "fail_col": MockFailingGenerator(config=_expr_config("fail_col"), resource_provider=provider), + } + + graph = ExecutionGraph.create(configs, strategies) + row_groups = [(0, 5)] + tracker = CompletionTracker.with_graph(graph, row_groups) + + scheduler = AsyncTaskScheduler( + generators=generators, + graph=graph, + tracker=tracker, + row_groups=row_groups, + shutdown_error_rate=0.5, + shutdown_error_window=5, + num_records=5, + buffer_size=5, + ) + await scheduler.run() + + # After run() returns, no worker tasks should remain. + assert scheduler.active_worker_count == 0 diff --git a/packages/data-designer-engine/tests/engine/models/test_async_engine_switch.py b/packages/data-designer-engine/tests/engine/models/test_async_engine_switch.py index 5cba94623..b1a749155 100644 --- a/packages/data-designer-engine/tests/engine/models/test_async_engine_switch.py +++ b/packages/data-designer-engine/tests/engine/models/test_async_engine_switch.py @@ -3,9 +3,7 @@ from __future__ import annotations -from unittest.mock import MagicMock, patch - -import pytest +from unittest.mock import MagicMock from data_designer.config.column_configs import GenerationStrategy from data_designer.engine.dataset_builders.dataset_builder import DatasetBuilder @@ -26,9 +24,8 @@ def test_model_facade_has_sync_methods() -> None: assert hasattr(ModelFacade, "generate_text_embeddings") -def test_async_engine_env_controls_builder_execution_path(monkeypatch: pytest.MonkeyPatch) -> None: +def test_async_engine_env_controls_builder_execution_path() -> None: """When DATA_DESIGNER_ASYNC_ENGINE is set, _run_cell_by_cell_generator dispatches to async fan-out.""" - import data_designer.engine.dataset_builders.dataset_builder as cwb_module mock_generator = MagicMock() mock_generator.get_generation_strategy.return_value = GenerationStrategy.CELL_BY_CELL @@ -38,15 +35,15 @@ def test_async_engine_env_controls_builder_execution_path(monkeypatch: pytest.Mo builder._resource_provider.run_config.non_inference_max_parallel_workers = 4 # Test with async enabled — uses max_parallel_requests from generator (same as sync) - with patch.object(cwb_module, "DATA_DESIGNER_ASYNC_ENGINE", True): - DatasetBuilder._run_cell_by_cell_generator(builder, mock_generator) - builder._fan_out_with_async.assert_called_once_with(mock_generator, max_workers=4) - builder._fan_out_with_threads.assert_not_called() + builder._use_async = True + DatasetBuilder._run_cell_by_cell_generator(builder, mock_generator) + builder._fan_out_with_async.assert_called_once_with(mock_generator, max_workers=4) + builder._fan_out_with_threads.assert_not_called() builder.reset_mock() # Test with async disabled — uses max_parallel_requests from generator - with patch.object(cwb_module, "DATA_DESIGNER_ASYNC_ENGINE", False): - DatasetBuilder._run_cell_by_cell_generator(builder, mock_generator) - builder._fan_out_with_threads.assert_called_once_with(mock_generator, max_workers=4) - builder._fan_out_with_async.assert_not_called() + builder._use_async = False + DatasetBuilder._run_cell_by_cell_generator(builder, mock_generator) + builder._fan_out_with_threads.assert_called_once_with(mock_generator, max_workers=4) + builder._fan_out_with_async.assert_not_called() diff --git a/packages/data-designer/src/data_designer/interface/data_designer.py b/packages/data-designer/src/data_designer/interface/data_designer.py index 7913bbbee..30f8fe108 100644 --- a/packages/data-designer/src/data_designer/interface/data_designer.py +++ b/packages/data-designer/src/data_designer/interface/data_designer.py @@ -224,6 +224,8 @@ def create( try: builder = self._create_dataset_builder(config_builder.build(), resource_provider) builder.build(num_records=num_records) + except DeprecationWarning: + raise except Exception as e: raise DataDesignerGenerationError(f"🛑 Error generating dataset: {e}") from e @@ -296,6 +298,8 @@ def preview( builder = self._create_dataset_builder(config_builder.build(), resource_provider) raw_dataset = builder.build_preview(num_records=num_records) processed_dataset = builder.process_preview(raw_dataset) + except DeprecationWarning: + raise except Exception as e: raise DataDesignerGenerationError(f"🛑 Error generating preview dataset: {e}") from e @@ -333,6 +337,7 @@ def preview( processor_artifacts=processor_artifacts, config_builder=config_builder, dataset_metadata=dataset_metadata, + task_traces=builder.task_traces or None, ) def _log_jinja_rendering_engine_mode(self) -> None: