From fe4f8f6e78e35f24953fdadc1244cde3e57f0a81 Mon Sep 17 00:00:00 2001 From: Andre Manoel Date: Wed, 12 Aug 2026 13:11:50 -0300 Subject: [PATCH] feat: expose terminal failure locations Add opt-in terminal failure capture for create and preview results, including zero-row generation errors. Surface early shutdown so callers can distinguish cancelled rows from attributable terminal failures. Closes #860 Signed-off-by: Andre Manoel --- .../src/data_designer/config/__init__.py | 2 + .../src/data_designer/config/interface.py | 2 + .../data_designer/config/preview_results.py | 9 + .../data_designer/config/terminal_failure.py | 19 ++ .../dataset_builders/async_scheduler.py | 26 +++ .../dataset_builders/dataset_builder.py | 42 ++++- .../dataset_builders/test_async_scheduler.py | 167 +++++++++++++++++- .../dataset_builders/test_dataset_builder.py | 4 + .../data_designer/interface/data_designer.py | 79 +++++++-- .../src/data_designer/interface/errors.py | 11 ++ .../src/data_designer/interface/results.py | 10 ++ .../tests/interface/test_acreate.py | 5 +- .../tests/interface/test_data_designer.py | 102 ++++++++++- .../tests/interface/test_results.py | 29 +++ 14 files changed, 479 insertions(+), 28 deletions(-) create mode 100644 packages/data-designer-config/src/data_designer/config/terminal_failure.py diff --git a/packages/data-designer-config/src/data_designer/config/__init__.py b/packages/data-designer-config/src/data_designer/config/__init__.py index c3fc1195f..6680a9f9a 100644 --- a/packages/data-designer-config/src/data_designer/config/__init__.py +++ b/packages/data-designer-config/src/data_designer/config/__init__.py @@ -107,6 +107,7 @@ LocalFileSeedSource, ) from data_designer.config.seed_source_dataframe import DataFrameSeedSource # noqa: F401 + from data_designer.config.terminal_failure import TerminalTaskFailure # noqa: F401 from data_designer.config.utils.code_lang import CodeLang # noqa: F401 from data_designer.config.utils.info import InfoType # noqa: F401 from data_designer.config.utils.media_helpers import AudioFormat, ImageFormat, VideoFormat # noqa: F401 @@ -194,6 +195,7 @@ "ResumeMode": (f"{_MOD_BASE}.run_config", "ResumeMode"), "RunConfig": (f"{_MOD_BASE}.run_config", "RunConfig"), "ThrottleConfig": (f"{_MOD_BASE}.run_config", "ThrottleConfig"), + "TerminalTaskFailure": (f"{_MOD_BASE}.terminal_failure", "TerminalTaskFailure"), # script_params "DataDesignerScriptParams": (f"{_MOD_BASE}.script_params", "DataDesignerScriptParams"), # scheduling metadata diff --git a/packages/data-designer-config/src/data_designer/config/interface.py b/packages/data-designer-config/src/data_designer/config/interface.py index 499fcbba5..662b365ee 100644 --- a/packages/data-designer-config/src/data_designer/config/interface.py +++ b/packages/data-designer-config/src/data_designer/config/interface.py @@ -33,6 +33,7 @@ def create( config_builder: DataDesignerConfigBuilder, *, num_records: int = DEFAULT_NUM_RECORDS, + capture_terminal_failures: bool = False, ) -> ResultsT: ... @abstractmethod @@ -41,6 +42,7 @@ def preview( config_builder: DataDesignerConfigBuilder, *, num_records: int = DEFAULT_NUM_RECORDS, + capture_terminal_failures: bool = False, ) -> PreviewResults: ... @abstractmethod 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 65880dccb..a44a614c7 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 @@ -9,6 +9,7 @@ from data_designer.config.config_builder import DataDesignerConfigBuilder from data_designer.config.dataset_metadata import DatasetMetadata from data_designer.config.seed_source_dataframe import DataFrameSeedSource +from data_designer.config.terminal_failure import TerminalTaskFailure from data_designer.config.utils.visualization import WithRecordSamplerMixin if TYPE_CHECKING: @@ -25,6 +26,8 @@ def __init__( analysis: DatasetProfilerResults | None = None, processor_artifacts: dict[str, list[dict]] | None = None, task_traces: list[Any] | None = None, + terminal_failures: list[TerminalTaskFailure] | None = None, + early_shutdown: bool = False, ): """Creates a new instance with results from a Data Designer preview run. @@ -35,12 +38,18 @@ def __init__( 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). + terminal_failures: Terminal column failures captured for omitted seed rows. + Check ``early_shutdown`` before treating this list as complete. + early_shutdown: Whether generation stopped at the global error-rate threshold. + Cancelled rows are not included in ``terminal_failures``. """ 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.terminal_failures: list[TerminalTaskFailure] = list(terminal_failures or []) + self.early_shutdown = early_shutdown self._config_builder = config_builder def to_config_builder(self, columns: list[str] | None = None) -> DataDesignerConfigBuilder: diff --git a/packages/data-designer-config/src/data_designer/config/terminal_failure.py b/packages/data-designer-config/src/data_designer/config/terminal_failure.py new file mode 100644 index 000000000..d39ce861c --- /dev/null +++ b/packages/data-designer-config/src/data_designer/config/terminal_failure.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True, order=True, slots=True) +class TerminalTaskFailure: + """Terminal column failure for an omitted seed row. + + ``seed_row_index`` is the zero-based position in the requested generation + sequence. It is not necessarily the raw source index for shuffled, selected, + or cycled seed datasets. + """ + + seed_row_index: int + column: str 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 483bb7338..7418c7fc2 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 @@ -17,6 +17,7 @@ import data_designer.lazy_heavy_imports as lazy from data_designer.config.column_configs import ExpressionColumnConfig, GenerationStrategy +from data_designer.config.terminal_failure import TerminalTaskFailure from data_designer.engine.capacity import ( AsyncCapacityConfigured, AsyncCapacityObservedMaxima, @@ -214,6 +215,7 @@ def __init__( adaptive_row_group_initial_target: int = 1, request_pressure_provider: RequestPressureSnapshotProvider | None = None, request_pressure_advisory: bool = False, + capture_terminal_failures: bool = False, ) -> None: self._generators = generators self._graph = graph @@ -339,6 +341,7 @@ def __init__( # context naturally because the from_scratch task raised; the async # engine drops rows and continues, losing the cause unless we capture it. self._first_non_retryable_error: Exception | None = None + self._terminal_failures: list[TerminalTaskFailure] | None = [] if capture_terminal_failures else None self._fatal_worker_error: BaseException | None = None self._cancel_requested = Event() self._run_loop: asyncio.AbstractEventLoop | None = None @@ -446,6 +449,11 @@ def first_non_retryable_error(self) -> Exception | None: """ return self._first_non_retryable_error + @property + def terminal_failures(self) -> list[TerminalTaskFailure]: + """Terminal column failures captured for omitted seed rows.""" + return sorted(self._terminal_failures or []) + @property def retryable_outcome_metrics(self) -> dict[str, object]: """Return sanitized rolling and cumulative model-task outcome counts.""" @@ -1542,6 +1550,7 @@ async def _salvage_stalled_row_groups( already_dropped = task.row_index is not None and self._tracker.is_dropped(task.row_group, task.row_index) if not already_dropped and self._reporter: self._reporter.record_failure(task.column) + self._record_terminal_failure(task) if task.row_index is not None: self._drop_row(task.row_group, task.row_index, exclude_columns={task.column}) else: @@ -1767,6 +1776,22 @@ def _drop_row(self, row_group: int, row_index: int, *, exclude_columns: set[str] if self._buffer_manager: self._buffer_manager.drop_row(row_group, row_index) + def _record_terminal_failure(self, task: Task) -> None: + if self._terminal_failures is None: + return + + start_offset = self._get_rg_start_offset(task.row_group) + if start_offset is None: + return + row_indices = (task.row_index,) if task.row_index is not None else range(self._get_rg_size(task.row_group)) + for row_index in row_indices: + # Preserve the failure that actually caused the row to be omitted. + if self._tracker.is_dropped(task.row_group, row_index): + continue + self._terminal_failures.append( + TerminalTaskFailure(seed_row_index=start_offset + row_index, column=task.column) + ) + def _drop_row_group(self, row_group: int, row_group_size: int, *, exclude_columns: set[str] | None = None) -> None: for row_index in range(row_group_size): self._drop_row(row_group, row_index, exclude_columns=exclude_columns) @@ -2091,6 +2116,7 @@ async def _execute_task_inner_impl(self, task: Task, lease: TaskAdmissionLease, logger.error("Unexpected %s", log_message, exc_info=True) # Non-retryable data/user/provider failures drop the affected row(s); # internal bug-shaped failures above abort the run instead. + self._record_terminal_failure(task) if task.row_index is not None: self._drop_row(task.row_group, task.row_index, exclude_columns={task.column}) else: 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 5615571f9..b13f38e5b 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 @@ -27,6 +27,7 @@ ProcessorConfig, ProcessorType, ) +from data_designer.config.terminal_failure import TerminalTaskFailure from data_designer.config.utils.type_helpers import StrEnum from data_designer.config.version import get_library_version from data_designer.engine.column_generators.generators.base import ( @@ -193,6 +194,7 @@ def __init__( # async run, if any. Used by the interface to surface the original cause # when a run produces 0 records due to deterministic failures. self._first_non_retryable_error: Exception | None = None + self._terminal_failures: list[TerminalTaskFailure] = [] 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) @@ -239,6 +241,11 @@ def first_non_retryable_error(self) -> Exception | None: """First non-retryable error captured by the scheduler in the most recent run.""" return self._first_non_retryable_error + @property + def terminal_failures(self) -> list[TerminalTaskFailure]: + """Terminal column failures captured during the most recent run.""" + return list(self._terminal_failures) + @functools.cached_property def single_column_configs(self) -> list[ColumnConfigT]: configs = [] @@ -256,6 +263,7 @@ def build( on_batch_complete: Callable[[Path], None] | None = None, save_multimedia_to_disk: bool = True, resume: ResumeMode = ResumeMode.NEVER, + capture_terminal_failures: bool = False, ) -> Path: """Build the dataset. @@ -279,6 +287,7 @@ def build( In all resume modes, in-flight partial results from the interrupted run are discarded before generation continues. + capture_terminal_failures: Capture the terminal column for omitted seed rows. Returns: Path to the generated dataset directory. @@ -351,7 +360,14 @@ def build( resume = ResumeMode.NEVER self.artifact_storage.resume = ResumeMode.NEVER - self._build_async(generators, num_records, buffer_size, on_batch_complete, resume=resume) + self._build_async( + generators, + num_records, + buffer_size, + on_batch_complete, + resume=resume, + capture_terminal_failures=capture_terminal_failures, + ) # After-generation processors run unconditionally on the on-disk dataset # (not gated on ``generated``). When resume sees every row group already @@ -537,7 +553,7 @@ def _load_resume_state(self, num_records: int, buffer_size: int) -> _ResumeState completed_row_groups=completed_row_groups, ) - def build_preview(self, *, num_records: int) -> pd.DataFrame: + def build_preview(self, *, num_records: int, capture_terminal_failures: bool = False) -> pd.DataFrame: self._reset_run_state() run_readiness_check( self.single_column_configs, @@ -551,7 +567,11 @@ def build_preview(self, *, num_records: int) -> pd.DataFrame: generators, self._graph = self._initialize_generators_and_graph() start_time = time.perf_counter() - dataset = self._build_async_preview(generators, num_records) + dataset = self._build_async_preview( + generators, + num_records, + capture_terminal_failures=capture_terminal_failures, + ) self._resource_provider.model_registry.log_model_usage(time.perf_counter() - start_time) @@ -564,8 +584,15 @@ def _reset_run_state(self) -> None: self._actual_num_records = -1 self._first_non_retryable_error = None self._task_traces = [] + self._terminal_failures = [] - def _build_async_preview(self, generators: list[ColumnGenerator], num_records: int) -> pd.DataFrame: + def _build_async_preview( + self, + generators: list[ColumnGenerator], + num_records: int, + *, + capture_terminal_failures: bool = False, + ) -> pd.DataFrame: """Async preview path - single row group, no disk writes, returns in-memory DataFrame.""" logger.info("⚡ Using async task-queue preview") @@ -578,6 +605,7 @@ def _build_async_preview(self, generators: list[ColumnGenerator], num_records: i buffer_size=num_records, run_post_batch_in_scheduler=False, trace=trace_enabled, + capture_terminal_failures=capture_terminal_failures, ) loop = ensure_async_engine_loop() @@ -590,6 +618,7 @@ def _build_async_preview(self, generators: list[ColumnGenerator], num_records: i self._partial_row_groups = scheduler.partial_row_groups self._actual_num_records = buffer_manager.actual_num_records self._first_non_retryable_error = scheduler.first_non_retryable_error + self._terminal_failures = scheduler.terminal_failures if not buffer_manager.has_row_group(0): return lazy.pd.DataFrame() @@ -746,6 +775,7 @@ def _build_async( on_batch_complete: Callable[[Path], None] | None = None, *, resume: ResumeMode = ResumeMode.NEVER, + capture_terminal_failures: bool = False, ) -> bool: """Async task-queue builder path - dispatches tasks based on dependency readiness. @@ -851,6 +881,7 @@ def on_complete(final_path: Path | str | None) -> None: initial_actual_num_records=initial_actual_num_records, initial_total_num_batches=initial_total_num_batches, scheduler_event_sink=scheduler_event_sink, + capture_terminal_failures=capture_terminal_failures, ) # Run on background event loop. Capture scheduler state in `finally` @@ -867,6 +898,7 @@ def on_complete(final_path: Path | str | None) -> None: self._partial_row_groups = scheduler.partial_row_groups self._actual_num_records = buffer_manager.actual_num_records self._first_non_retryable_error = scheduler.first_non_retryable_error + self._terminal_failures = scheduler.terminal_failures # Emit telemetry try: @@ -917,6 +949,7 @@ def _prepare_async_run( initial_actual_num_records: int = 0, initial_total_num_batches: int = 0, scheduler_event_sink: SchedulerAdmissionEventSink | None = None, + capture_terminal_failures: bool = False, ) -> tuple[AsyncTaskScheduler, RowGroupBufferManager]: """Build a fully-wired scheduler and buffer manager for async generation. @@ -1008,6 +1041,7 @@ def on_before_checkpoint(rg_id: int, rg_size: int) -> None: ), request_pressure_provider=self._resource_provider.model_registry.request_admission, request_pressure_advisory=True, + capture_terminal_failures=capture_terminal_failures, ) return scheduler, buffer_manager 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 ed2a7fa0a..f6b8ab391 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 @@ -28,6 +28,7 @@ from data_designer.config.models import ChatCompletionInferenceParams, ModelConfig from data_designer.config.sampler_params import SamplerType from data_designer.config.scheduling import SchedulingMetadata +from data_designer.config.terminal_failure import TerminalTaskFailure from data_designer.engine.column_generators.generators.base import ( ColumnGenerator, ColumnGeneratorFullColumn, @@ -149,6 +150,11 @@ def generate(self, data: lazy.pd.DataFrame) -> lazy.pd.DataFrame: return data +class MockFailingFullColumnGenerator(ColumnGeneratorFullColumn[ExpressionColumnConfig]): + def generate(self, data: lazy.pd.DataFrame) -> lazy.pd.DataFrame: + raise ValueError("permanent batch failure") + + class MockStatefulSeed(FromScratchColumnGenerator[ExpressionColumnConfig]): """Stateful mock seed generator.""" @@ -418,6 +424,7 @@ def _build_simple_pipeline( max_concurrent_row_groups: int = 3, adaptive_row_group_admission: bool = False, adaptive_row_group_initial_target: int = 1, + capture_terminal_failures: bool = False, ) -> tuple[AsyncTaskScheduler, CompletionTracker]: """Build a simple seed → cell pipeline for testing.""" if configs is None: @@ -460,6 +467,7 @@ def _build_simple_pipeline( scheduler_event_sink=scheduler_event_sink, adaptive_row_group_admission=adaptive_row_group_admission, adaptive_row_group_initial_target=adaptive_row_group_initial_target, + capture_terminal_failures=capture_terminal_failures, ) return scheduler, tracker @@ -770,6 +778,57 @@ async def test_scheduler_auto_computes_row_group_start_offsets_for_fresh_runs() assert buffer_manager.get_row(2, 0)["seed"] == 4 +@pytest.mark.asyncio(loop_scope="session") +async def test_scheduler_expands_batch_failure_to_global_seed_row_indices() -> None: + provider = _mock_provider() + configs = [SamplerColumnConfig(name="seed", sampler_type=SamplerType.CATEGORY, params={"values": ["A"]})] + graph = ExecutionGraph.create(configs, {"seed": GenerationStrategy.FULL_COLUMN}) + row_groups = [(0, 2), (1, 2), (2, 1)] + tracker = CompletionTracker.with_graph(graph, row_groups) + scheduler = AsyncTaskScheduler( + generators={ + "seed": MockFailingSeedGenerator(config=_expr_config("seed"), resource_provider=provider), + }, + graph=graph, + tracker=tracker, + row_groups=row_groups, + capture_terminal_failures=True, + ) + + await scheduler.run() + + assert scheduler.terminal_failures == [ + TerminalTaskFailure(seed_row_index=index, column="seed") for index in range(5) + ] + + +@pytest.mark.asyncio(loop_scope="session") +async def test_scheduler_reports_only_current_resume_invocation() -> None: + provider = _mock_provider() + configs = [SamplerColumnConfig(name="seed", sampler_type=SamplerType.CATEGORY, params={"values": ["A"]})] + graph = ExecutionGraph.create(configs, {"seed": GenerationStrategy.FULL_COLUMN}) + row_groups = CompactRowGroupPlan.resume( + original_target=5, + num_records=5, + buffer_size=2, + completed_ids={0, 1}, + ) + tracker = CompletionTracker.with_graph(graph, row_groups) + scheduler = AsyncTaskScheduler( + generators={ + "seed": MockFailingSeedGenerator(config=_expr_config("seed"), resource_provider=provider), + }, + graph=graph, + tracker=tracker, + row_groups=row_groups, + capture_terminal_failures=True, + ) + + await scheduler.run() + + assert scheduler.terminal_failures == [TerminalTaskFailure(seed_row_index=4, column="seed")] + + @pytest.mark.asyncio(loop_scope="session") async def test_scheduler_non_retryable_failure_drops_row() -> None: """Non-retryable failure drops the row.""" @@ -796,6 +855,7 @@ async def test_scheduler_non_retryable_failure_drops_row() -> None: graph=graph, tracker=tracker, row_groups=row_groups, + capture_terminal_failures=True, ) await scheduler.run() @@ -805,6 +865,26 @@ async def test_scheduler_non_retryable_failure_drops_row() -> None: # Row group is "complete" because all non-dropped rows have all columns # (there are no non-dropped rows) assert tracker.is_row_group_complete(0, 2, ["seed", "fail_col"]) + assert scheduler.terminal_failures == [ + TerminalTaskFailure(seed_row_index=0, column="fail_col"), + TerminalTaskFailure(seed_row_index=1, column="fail_col"), + ] + + +@pytest.mark.asyncio(loop_scope="session") +async def test_scheduler_terminal_failure_capture_is_opt_in() -> None: + provider = _mock_provider() + scheduler, _ = _build_simple_pipeline( + num_records=1, + generators={ + "seed": MockSeedGenerator(config=_expr_config("seed"), resource_provider=provider), + "cell_out": MockFailingGenerator(config=_expr_config("cell_out"), resource_provider=provider), + }, + ) + + await scheduler.run() + + assert scheduler.terminal_failures == [] @pytest.mark.asyncio(loop_scope="session") @@ -1198,7 +1278,11 @@ async def test_scheduler_retryable_failure_recovers_in_salvage() -> None: "fail_col": fail_gen, } scheduler, tracker = _build_simple_pipeline( - num_records=2, generators=generators, configs=configs, strategies=strategies + num_records=2, + generators=generators, + configs=configs, + strategies=strategies, + capture_terminal_failures=True, ) await scheduler.run() @@ -1206,6 +1290,33 @@ async def test_scheduler_retryable_failure_recovers_in_salvage() -> None: assert not tracker.is_dropped(0, 0) assert not tracker.is_dropped(0, 1) assert tracker.is_row_group_complete(0, 2, ["seed", "fail_col"]) + assert scheduler.terminal_failures == [] + + +@pytest.mark.asyncio(loop_scope="session") +async def test_scheduler_captures_retryable_failures_only_after_salvage_exhaustion() -> None: + provider = _mock_provider() + scheduler, tracker = _build_simple_pipeline( + num_records=2, + generators={ + "seed": MockSeedGenerator(config=_expr_config("seed"), resource_provider=provider), + "cell_out": MockFailingGenerator( + config=_expr_config("cell_out"), + resource_provider=provider, + transient_failures=100, + ), + }, + capture_terminal_failures=True, + ) + + await scheduler.run() + + assert tracker.is_dropped(0, 0) + assert tracker.is_dropped(0, 1) + assert scheduler.terminal_failures == [ + TerminalTaskFailure(seed_row_index=0, column="cell_out"), + TerminalTaskFailure(seed_row_index=1, column="cell_out"), + ] @pytest.mark.asyncio(loop_scope="session") @@ -1244,6 +1355,7 @@ async def test_scheduler_eager_row_drop_skips_downstream_of_failed_column() -> N trace=True, num_records=2, buffer_size=2, + capture_terminal_failures=True, ) await scheduler.run() @@ -1253,6 +1365,10 @@ async def test_scheduler_eager_row_drop_skips_downstream_of_failed_column() -> N # downstream was never dispatched for the dropped rows downstream_traces = [t for t in scheduler.traces if t.column == "downstream"] assert len(downstream_traces) == 0 + assert scheduler.terminal_failures == [ + TerminalTaskFailure(seed_row_index=0, column="fail_col"), + TerminalTaskFailure(seed_row_index=1, column="fail_col"), + ] # Row group is still "complete" (no non-dropped rows remain) assert tracker.is_row_group_complete(0, 2, ["seed", "fail_col", "downstream"]) assert scheduler._reporter is not None @@ -1460,6 +1576,7 @@ def on_finalize(rg_id: int) -> None: on_finalize_row_group=on_finalize, shutdown_error_rate=0.5, shutdown_error_window=4, + capture_terminal_failures=True, ) await scheduler.run() @@ -1471,6 +1588,8 @@ def on_finalize(rg_id: int) -> None: assert 0 in finalized assert scheduler.partial_row_groups == (0,) assert 1 <= buffer_mgr.actual_num_records <= 7 + assert scheduler.terminal_failures + assert {failure.seed_row_index for failure in scheduler.terminal_failures} <= {5, 6, 7} @pytest.mark.asyncio(loop_scope="session") @@ -2647,6 +2766,50 @@ def drop_middle_row(row_group: int, row_group_size: int) -> FrontierDelta: assert tracker.is_row_group_complete(0, 3, ["seed", "cell_out"]) +@pytest.mark.asyncio(loop_scope="session") +async def test_batch_failure_excludes_row_dropped_by_pre_batch_processor() -> None: + provider = _mock_provider() + configs = [ + SamplerColumnConfig(name="seed", sampler_type=SamplerType.CATEGORY, params={"values": ["A"]}), + ExpressionColumnConfig(name="full_out", expr="{{ seed }}"), + ] + strategies = { + "seed": GenerationStrategy.FULL_COLUMN, + "full_out": GenerationStrategy.FULL_COLUMN, + } + graph = ExecutionGraph.create(configs, strategies) + tracker = CompletionTracker.with_graph(graph, [(0, 3)]) + buffer_manager = RowGroupBufferManager(_make_storage()) + + def drop_middle_row(row_group: int, row_group_size: int) -> FrontierDelta: + del row_group_size + buffer_manager.drop_row(row_group, 1) + return tracker.drop_row(row_group, 1) + + scheduler = AsyncTaskScheduler( + generators={ + "seed": MockSeedGenerator(config=_expr_config("seed"), resource_provider=provider), + "full_out": MockFailingFullColumnGenerator( + config=_expr_config("full_out"), + resource_provider=provider, + ), + }, + graph=graph, + tracker=tracker, + row_groups=[(0, 3)], + buffer_manager=buffer_manager, + on_seeds_complete=drop_middle_row, + capture_terminal_failures=True, + ) + + await scheduler.run() + + assert scheduler.terminal_failures == [ + TerminalTaskFailure(seed_row_index=0, column="full_out"), + TerminalTaskFailure(seed_row_index=2, column="full_out"), + ] + + def test_apply_frontier_delta_enqueues_ready_tasks_in_one_queue_operation(monkeypatch: pytest.MonkeyPatch) -> None: provider = _mock_provider() configs = [ @@ -5116,11 +5279,13 @@ def generate_from_scratch(self, num_records: int) -> lazy.pd.DataFrame: trace=True, num_records=num_records, buffer_size=num_records, + capture_terminal_failures=True, ) await asyncio.wait_for(scheduler.run(), timeout=10.0) assert tracker.is_row_group_complete(0, num_records, ["seed", "review", "complaint"]) assert scheduler.retryable_outcome_metrics["cumulative_counts"] == {"success": 2} + assert scheduler.terminal_failures == [] for ri in range(num_records): row = buffer_mgr.get_row(0, ri) diff --git a/packages/data-designer-engine/tests/engine/dataset_builders/test_dataset_builder.py b/packages/data-designer-engine/tests/engine/dataset_builders/test_dataset_builder.py index 0ecc321ef..36143a48c 100644 --- a/packages/data-designer-engine/tests/engine/dataset_builders/test_dataset_builder.py +++ b/packages/data-designer-engine/tests/engine/dataset_builders/test_dataset_builder.py @@ -30,6 +30,7 @@ from data_designer.config.seed import IndexRange, PartitionBlock, SamplingStrategy from data_designer.config.seed_source import LocalFileSeedSource from data_designer.config.seed_source_dataframe import DataFrameSeedSource +from data_designer.config.terminal_failure import TerminalTaskFailure from data_designer.engine.column_generators.generators.base import GenerationStrategy from data_designer.engine.dataset_builders.dataset_builder import DatasetBuilder, build_row_group_resume_plan from data_designer.engine.dataset_builders.errors import DatasetGenerationError, DatasetProcessingError @@ -440,6 +441,7 @@ class StubScheduler: early_shutdown: bool = False partial_row_groups: tuple[int, ...] = () first_non_retryable_error: Exception | None = None + terminal_failures: list[TerminalTaskFailure] = [] async def run(self) -> None: return None @@ -506,6 +508,7 @@ def test_reset_run_state_clears_per_run_signals(stub_resource_provider, stub_tes builder._partial_row_groups = (0, 1) builder._actual_num_records = 42 builder._task_traces = ["trace"] # type: ignore[list-item] + builder._terminal_failures = [TerminalTaskFailure(seed_row_index=3, column="failed_col")] builder._reset_run_state() @@ -513,6 +516,7 @@ def test_reset_run_state_clears_per_run_signals(stub_resource_provider, stub_tes assert builder.partial_row_groups == () assert builder.actual_num_records == -1 assert builder.task_traces == [] + assert builder.terminal_failures == [] # Processor tests 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 e1b6961f7..5d2d4e72b 100644 --- a/packages/data-designer/src/data_designer/interface/data_designer.py +++ b/packages/data-designer/src/data_designer/interface/data_designer.py @@ -223,6 +223,7 @@ def create( resume: ResumeMode = ResumeMode.NEVER, artifact_path: Path | str | None = None, on_batch_complete: Callable[[Path], None] | None = None, + capture_terminal_failures: bool = False, ) -> DatasetCreationResults: """Create dataset and save results to the local artifact storage. @@ -263,6 +264,8 @@ def create( path, so it is recommended to keep it lightweight or delegate slow work to a queue, e.g. ``on_batch_complete=lambda path: queue_upload(path)``. Callback exceptions abort the run and are wrapped as ``DataDesignerGenerationError``. + capture_terminal_failures: Capture the terminal column for each omitted seed row. + When results report ``early_shutdown=True``, cancelled rows are not included. Returns: DatasetCreationResults object with methods for loading the generated dataset, @@ -296,11 +299,19 @@ def create( raise DataDesignerGenerationError(f"🛑 Error generating dataset: {e}") from e try: - builder.build(num_records=num_records, on_batch_complete=on_batch_complete, resume=resume) + builder.build( + num_records=num_records, + on_batch_complete=on_batch_complete, + resume=resume, + capture_terminal_failures=capture_terminal_failures, + ) except DeprecationWarning: raise except Exception as e: - raise DataDesignerGenerationError(f"🛑 Error generating dataset: {e}") from e + raise DataDesignerGenerationError( + f"🛑 Error generating dataset: {e}", + terminal_failures=builder.terminal_failures, + ) from e task_traces = builder.task_traces @@ -319,17 +330,22 @@ def create( "🛑 Generation produced zero records — early shutdown was triggered. " "The non-retryable error rate exceeded the configured threshold; check the " "warnings above (and any 'Provider showing degraded performance' logs) for " - "the contributing failures." + "the contributing failures.", + terminal_failures=builder.terminal_failures, ) from e # Surface the original task error when the run produced 0 records due to a # deterministic non-retryable failure (e.g. bad seed source). Without this, # the user sees a generic FileNotFoundError-on-parquet that obscures the cause. root_cause = builder.first_non_retryable_error if root_cause is not None and builder.actual_num_records == 0: - raise DataDesignerGenerationError(f"🛑 {type(root_cause).__name__}: {root_cause}") from root_cause + raise DataDesignerGenerationError( + f"🛑 {type(root_cause).__name__}: {root_cause}", + terminal_failures=builder.terminal_failures, + ) from root_cause raise DataDesignerGenerationError( f"🛑 Failed to load generated dataset — all records may have been dropped " - f"due to generation failures. Check the warnings above for details. Original error: {e}" + f"due to generation failures. Check the warnings above for details. Original error: {e}", + terminal_failures=builder.terminal_failures, ) from e # Defensive: the batch manager skips writing when the buffer is empty, so in @@ -343,14 +359,19 @@ def create( if builder.early_shutdown and builder.actual_num_records == 0: raise DataDesignerEarlyShutdownError( "🛑 Dataset is empty — early shutdown was triggered before any records " - "could complete. Check the warnings above for the contributing failures." + "could complete. Check the warnings above for the contributing failures.", + terminal_failures=builder.terminal_failures, ) root_cause = builder.first_non_retryable_error if root_cause is not None and builder.actual_num_records == 0: - raise DataDesignerGenerationError(f"🛑 {type(root_cause).__name__}: {root_cause}") from root_cause + raise DataDesignerGenerationError( + f"🛑 {type(root_cause).__name__}: {root_cause}", + terminal_failures=builder.terminal_failures, + ) from root_cause raise DataDesignerGenerationError( "🛑 Dataset is empty — all records were dropped due to generation failures. " - "Check the warnings above for details on which columns failed." + "Check the warnings above for details on which columns failed.", + terminal_failures=builder.terminal_failures, ) try: @@ -378,6 +399,8 @@ def create( config_builder=config_builder, dataset_metadata=dataset_metadata, task_traces=task_traces, + terminal_failures=builder.terminal_failures, + early_shutdown=builder.early_shutdown, ) async def acreate( @@ -388,15 +411,25 @@ async def acreate( dataset_name: str = "dataset", resume: ResumeMode = ResumeMode.NEVER, artifact_path: Path | str | None = None, + capture_terminal_failures: bool = False, ) -> DatasetCreationResults: """Async wrapper for creating a dataset without blocking the caller's event loop.""" - kwargs = {"num_records": num_records, "dataset_name": dataset_name, "resume": resume} + kwargs = { + "num_records": num_records, + "dataset_name": dataset_name, + "resume": resume, + "capture_terminal_failures": capture_terminal_failures, + } if artifact_path is not None: kwargs["artifact_path"] = artifact_path return await asyncio.to_thread(self.create, config_builder, **kwargs) def preview( - self, config_builder: DataDesignerConfigBuilder, *, num_records: int = DEFAULT_NUM_RECORDS + self, + config_builder: DataDesignerConfigBuilder, + *, + num_records: int = DEFAULT_NUM_RECORDS, + capture_terminal_failures: bool = False, ) -> PreviewResults: """Generate preview dataset for fast iteration on your Data Designer configuration. @@ -407,6 +440,8 @@ def preview( config_builder: The DataDesignerConfigBuilder containing the dataset configuration (columns, constraints, seed data, etc.). num_records: Number of records to generate. + capture_terminal_failures: Capture the terminal column for each omitted seed row. + When results report ``early_shutdown=True``, cancelled rows are not included. Returns: PreviewResults object with methods for inspecting the results. @@ -421,9 +456,13 @@ def preview( self._log_jinja_rendering_engine_mode() resource_provider = self._create_resource_provider("preview-dataset", config_builder) + builder: DatasetBuilder | None = None try: builder = self._create_dataset_builder(config_builder.build(), resource_provider) - raw_dataset = builder.build_preview(num_records=num_records) + raw_dataset = builder.build_preview( + num_records=num_records, + capture_terminal_failures=capture_terminal_failures, + ) processed_dataset = builder.process_preview(raw_dataset) except DeprecationWarning: # See comment in create() — strict warning filters convert engine-level @@ -431,7 +470,10 @@ def preview( # propagate untouched. raise except Exception as e: - raise DataDesignerGenerationError(f"🛑 Error generating preview dataset: {e}") from e + raise DataDesignerGenerationError( + f"🛑 Error generating preview dataset: {e}", + terminal_failures=builder.terminal_failures if builder is not None else None, + ) from e if len(processed_dataset) == 0: # Mirror the create() path: distinguish "early shutdown produced zero @@ -440,14 +482,19 @@ def preview( if builder.early_shutdown and builder.actual_num_records == 0: raise DataDesignerEarlyShutdownError( "🛑 Preview is empty — early shutdown was triggered before any records " - "could complete. Check the warnings above for the contributing failures." + "could complete. Check the warnings above for the contributing failures.", + terminal_failures=builder.terminal_failures, ) root_cause = builder.first_non_retryable_error if root_cause is not None and builder.actual_num_records == 0: - raise DataDesignerGenerationError(f"🛑 {type(root_cause).__name__}: {root_cause}") from root_cause + raise DataDesignerGenerationError( + f"🛑 {type(root_cause).__name__}: {root_cause}", + terminal_failures=builder.terminal_failures, + ) from root_cause raise DataDesignerGenerationError( "🛑 Dataset is empty — all records were dropped due to generation or processing failures. " - "Check the warnings above for details on which columns failed." + "Check the warnings above for details on which columns failed.", + terminal_failures=builder.terminal_failures, ) dropped_columns = raw_dataset.columns.difference(processed_dataset.columns) @@ -479,6 +526,8 @@ def preview( config_builder=config_builder, dataset_metadata=dataset_metadata, task_traces=builder.task_traces or None, + terminal_failures=builder.terminal_failures, + early_shutdown=builder.early_shutdown, ) def compose_workflow(self, *, name: str) -> CompositeWorkflow: diff --git a/packages/data-designer/src/data_designer/interface/errors.py b/packages/data-designer/src/data_designer/interface/errors.py index 8a7fb248f..274c57aa4 100644 --- a/packages/data-designer/src/data_designer/interface/errors.py +++ b/packages/data-designer/src/data_designer/interface/errors.py @@ -3,6 +3,9 @@ from __future__ import annotations +from collections.abc import Sequence + +from data_designer.config.terminal_failure import TerminalTaskFailure from data_designer.errors import DataDesignerError @@ -13,6 +16,14 @@ class DataDesignerProfilingError(DataDesignerError): class DataDesignerGenerationError(DataDesignerError): """Raised for errors related to a Data Designer dataset generation.""" + def __init__( + self, + *args: object, + terminal_failures: Sequence[TerminalTaskFailure] | None = None, + ) -> None: + super().__init__(*args) + self.terminal_failures = list(terminal_failures or []) + class DataDesignerWorkflowError(DataDesignerError): """Raised for errors related to composite workflow orchestration.""" diff --git a/packages/data-designer/src/data_designer/interface/results.py b/packages/data-designer/src/data_designer/interface/results.py index 6c5b076a3..18e72d770 100644 --- a/packages/data-designer/src/data_designer/interface/results.py +++ b/packages/data-designer/src/data_designer/interface/results.py @@ -12,6 +12,7 @@ from data_designer.config.dataset_metadata import DatasetMetadata from data_designer.config.errors import InvalidFileFormatError from data_designer.config.seed_source_dataframe import DataFrameSeedSource +from data_designer.config.terminal_failure import TerminalTaskFailure from data_designer.config.utils.visualization import WithRecordSamplerMixin from data_designer.engine.dataset_builders.errors import ArtifactStorageError from data_designer.engine.storage.artifact_storage import ArtifactStorage @@ -50,6 +51,8 @@ def __init__( config_builder: DataDesignerConfigBuilder, dataset_metadata: DatasetMetadata, task_traces: list[TaskTrace] | None = None, + terminal_failures: list[TerminalTaskFailure] | None = None, + early_shutdown: bool = False, ): """Creates a new instance with results based on a dataset creation run. @@ -62,12 +65,19 @@ def __init__( Resume note: only contains traces for the current invocation; traces from earlier ``create()`` calls that this run resumed are not retained. + terminal_failures: Terminal column failures for seed rows omitted by + the current invocation. Check ``early_shutdown`` before treating + this list as complete. + early_shutdown: Whether generation stopped at the global error-rate threshold. + Cancelled rows are not included in ``terminal_failures``. """ self.artifact_storage = artifact_storage self._analysis = analysis self._config_builder = config_builder self.dataset_metadata = dataset_metadata self.task_traces: list[TaskTrace] = task_traces or [] + self.terminal_failures: list[TerminalTaskFailure] = list(terminal_failures or []) + self.early_shutdown = early_shutdown def load_analysis(self) -> DatasetProfilerResults: """Load the profiling analysis results for the generated dataset. diff --git a/packages/data-designer/tests/interface/test_acreate.py b/packages/data-designer/tests/interface/test_acreate.py index 30c452fb2..b192401fb 100644 --- a/packages/data-designer/tests/interface/test_acreate.py +++ b/packages/data-designer/tests/interface/test_acreate.py @@ -60,6 +60,7 @@ async def test_acreate_delegates_to_create( num_records=1, dataset_name="async-dataset", resume=ResumeMode.IF_POSSIBLE, + capture_terminal_failures=True, ) assert result is expected @@ -68,6 +69,7 @@ async def test_acreate_delegates_to_create( num_records=1, dataset_name="async-dataset", resume=ResumeMode.IF_POSSIBLE, + capture_terminal_failures=True, ) @@ -90,9 +92,10 @@ def fake_create( num_records: int, dataset_name: str, resume: ResumeMode = ResumeMode.NEVER, + capture_terminal_failures: bool = False, ) -> DatasetCreationResults: nonlocal started_count - del num_records, dataset_name, resume + del num_records, dataset_name, resume, capture_terminal_failures with started_lock: started_count += 1 if started_count == 2: diff --git a/packages/data-designer/tests/interface/test_data_designer.py b/packages/data-designer/tests/interface/test_data_designer.py index 0a0ce5a9c..6ef21e125 100644 --- a/packages/data-designer/tests/interface/test_data_designer.py +++ b/packages/data-designer/tests/interface/test_data_designer.py @@ -36,6 +36,7 @@ FileContentsSeedSource, HuggingFaceSeedSource, ) +from data_designer.config.terminal_failure import TerminalTaskFailure from data_designer.engine.models.clients.adapters.http_model_client import ClientConcurrencyMode from data_designer.engine.models.errors import ( RETRYABLE_MODEL_ERRORS, @@ -625,7 +626,10 @@ def on_batch_complete(path: Path) -> None: mock_builder = MagicMock() mock_builder.build.return_value = None + terminal_failures = [TerminalTaskFailure(seed_row_index=2, column="failed_col")] mock_builder.task_traces = [] + mock_builder.terminal_failures = terminal_failures + mock_builder.early_shutdown = True mock_builder.artifact_storage.load_dataset_with_dropped_columns.return_value = lazy.pd.DataFrame({"col": [1]}) mock_builder_method.return_value = mock_builder @@ -633,11 +637,66 @@ def on_batch_complete(path: Path) -> None: mock_profiler.profile_dataset.return_value = None mock_profiler_method.return_value = mock_profiler - data_designer.create(stub_sampler_only_config_builder, num_records=1, on_batch_complete=on_batch_complete) + results = data_designer.create( + stub_sampler_only_config_builder, + num_records=1, + on_batch_complete=on_batch_complete, + capture_terminal_failures=True, + ) _, build_kwargs = mock_builder.build.call_args assert build_kwargs["num_records"] == 1 assert build_kwargs["on_batch_complete"] is on_batch_complete + assert build_kwargs["capture_terminal_failures"] is True + assert results.terminal_failures == terminal_failures + assert results.early_shutdown is True + + +def test_preview_forwards_and_returns_terminal_failures( + stub_artifact_path: Path, + stub_model_providers: list[ModelProvider], + stub_sampler_only_config_builder: DataDesignerConfigBuilder, + stub_managed_assets_path: Path, +) -> None: + data_designer = DataDesigner( + artifact_path=stub_artifact_path, + model_providers=stub_model_providers, + secret_resolver=PlaintextResolver(), + managed_assets_path=stub_managed_assets_path, + ) + terminal_failures = [TerminalTaskFailure(seed_row_index=2, column="failed_col")] + + with ( + patch.object(data_designer, "_create_resource_provider") as mock_resource_provider_method, + patch.object(data_designer, "_create_dataset_builder") as mock_builder_method, + patch.object(data_designer, "_create_dataset_profiler") as mock_profiler_method, + ): + mock_resource_provider = MagicMock() + mock_resource_provider.get_dataset_metadata.return_value = {} + mock_resource_provider_method.return_value = mock_resource_provider + + mock_builder = MagicMock() + mock_builder.build_preview.return_value = lazy.pd.DataFrame({"col": [1]}) + mock_builder.process_preview.return_value = lazy.pd.DataFrame({"col": [1]}) + mock_builder.task_traces = [] + mock_builder.terminal_failures = terminal_failures + mock_builder.early_shutdown = True + mock_builder.artifact_storage.list_processor_names.return_value = [] + mock_builder_method.return_value = mock_builder + + mock_profiler = MagicMock() + mock_profiler.profile_dataset.return_value = None + mock_profiler_method.return_value = mock_profiler + + results = data_designer.preview( + stub_sampler_only_config_builder, + num_records=1, + capture_terminal_failures=True, + ) + + mock_builder.build_preview.assert_called_once_with(num_records=1, capture_terminal_failures=True) + assert results.terminal_failures == terminal_failures + assert results.early_shutdown is True def test_run_config_rejects_invalid_buffer_size() -> None: @@ -876,6 +935,7 @@ def _patch_builder_state( early_shutdown: bool, actual_num_records: int = 0, first_non_retryable_error: Exception | None = None, + terminal_failures: list[TerminalTaskFailure] | None = None, ) -> contextlib.ExitStack: """Patch DatasetBuilder.early_shutdown / actual_num_records / first_non_retryable_error.""" stack = contextlib.ExitStack() @@ -900,6 +960,13 @@ def _patch_builder_state( return_value=first_non_retryable_error, ) ) + stack.enter_context( + patch( + "data_designer.engine.dataset_builders.dataset_builder.DatasetBuilder.terminal_failures", + new_callable=PropertyMock, + return_value=terminal_failures or [], + ) + ) return stack @@ -1052,6 +1119,7 @@ def test_create_surfaces_first_non_retryable_error_when_zero_records( """ data_designer = _make_data_designer(stub_artifact_path, stub_model_providers, stub_managed_assets_path) root_cause = ValueError("invalid seed source: no rows after hydration") + terminal_failures = [TerminalTaskFailure(seed_row_index=0, column="invalid_seed")] if load_side_effect == "raises": load_patch = patch( @@ -1070,15 +1138,21 @@ def test_create_surfaces_first_non_retryable_error_when_zero_records( early_shutdown=False, actual_num_records=0, first_non_retryable_error=root_cause, + terminal_failures=terminal_failures, ), ): with pytest.raises(DataDesignerGenerationError, match="invalid seed source") as exc_info: - data_designer.create(stub_sampler_only_config_builder, num_records=10) + data_designer.create( + stub_sampler_only_config_builder, + num_records=10, + capture_terminal_failures=True, + ) # Original cause is preserved via __cause__, not lost behind the parquet error. assert exc_info.value.__cause__ is root_cause # The typed DataDesignerEarlyShutdownError must NOT fire here — the gate didn't trip. assert not isinstance(exc_info.value, DataDesignerEarlyShutdownError) + assert exc_info.value.terminal_failures == terminal_failures def test_preview_raises_generation_error_when_dataset_is_empty( @@ -1093,13 +1167,27 @@ def test_preview_raises_generation_error_when_dataset_is_empty( secret_resolver=PlaintextResolver(), managed_assets_path=stub_managed_assets_path, ) + terminal_failures = [TerminalTaskFailure(seed_row_index=0, column="failed_col")] - with patch( - "data_designer.engine.dataset_builders.dataset_builder.DatasetBuilder.process_preview", - return_value=lazy.pd.DataFrame(), + with ( + patch( + "data_designer.engine.dataset_builders.dataset_builder.DatasetBuilder.process_preview", + return_value=lazy.pd.DataFrame(), + ), + _patch_builder_state( + early_shutdown=False, + actual_num_records=0, + terminal_failures=terminal_failures, + ), ): - with pytest.raises(DataDesignerGenerationError, match="Dataset is empty"): - data_designer.preview(stub_sampler_only_config_builder, num_records=1) + with pytest.raises(DataDesignerGenerationError, match="Dataset is empty") as exc_info: + data_designer.preview( + stub_sampler_only_config_builder, + num_records=1, + capture_terminal_failures=True, + ) + + assert exc_info.value.terminal_failures == terminal_failures def test_preview_raises_early_shutdown_error_on_empty_after_shutdown( diff --git a/packages/data-designer/tests/interface/test_results.py b/packages/data-designer/tests/interface/test_results.py index 200d43810..06bf24b27 100644 --- a/packages/data-designer/tests/interface/test_results.py +++ b/packages/data-designer/tests/interface/test_results.py @@ -10,6 +10,7 @@ import pytest import data_designer.lazy_heavy_imports as lazy +from data_designer.config import TerminalTaskFailure from data_designer.config.analysis.dataset_profiler import DatasetProfilerResults from data_designer.config.config_builder import DataDesignerConfigBuilder from data_designer.config.dataset_metadata import DatasetMetadata @@ -61,6 +62,29 @@ def test_init(stub_artifact_storage, stub_dataset_profiler_results, stub_complet assert results._analysis == stub_dataset_profiler_results assert results._config_builder == stub_complete_builder assert results.dataset_metadata == stub_dataset_metadata + assert results.terminal_failures == [] + assert results.early_shutdown is False + + +def test_results_expose_terminal_failures( + stub_artifact_storage, + stub_dataset_profiler_results, + stub_complete_builder, + stub_dataset_metadata, +) -> None: + terminal_failures = [TerminalTaskFailure(seed_row_index=3, column="review")] + + results = DatasetCreationResults( + artifact_storage=stub_artifact_storage, + analysis=stub_dataset_profiler_results, + config_builder=stub_complete_builder, + dataset_metadata=stub_dataset_metadata, + terminal_failures=terminal_failures, + early_shutdown=True, + ) + + assert results.terminal_failures == terminal_failures + assert results.early_shutdown is True def test_load_dataset(stub_dataset_creation_results, stub_artifact_storage, stub_dataframe): @@ -430,16 +454,21 @@ def test_preview_results_dataset_metadata() -> None: config_builder.get_columns_of_type.return_value = [] dataset_metadata = DatasetMetadata(seed_column_names=["name", "age"]) + terminal_failures = [TerminalTaskFailure(seed_row_index=1, column="greeting")] results = PreviewResults( config_builder=config_builder, dataset=lazy.pd.DataFrame({"name": ["Alice"], "age": [25], "greeting": ["Hello"]}), dataset_metadata=dataset_metadata, + terminal_failures=terminal_failures, + early_shutdown=True, ) # Verify metadata is stored as public attribute assert results.dataset_metadata == dataset_metadata assert results.dataset_metadata.seed_column_names == ["name", "age"] + assert results.terminal_failures == terminal_failures + assert results.early_shutdown is True # Patch display_sample_record to capture the seed_column_names argument with patch("data_designer.config.utils.visualization.display_sample_record", wraps=display_fn) as mock_display: