Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 22 additions & 0 deletions packages/data-designer/src/data_designer/interface/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,19 +22,41 @@
DataDesignerWorkflowError,
)
from data_designer.interface.results import DatasetCreationResults # noqa: F401
from data_designer.interface.workflow_metadata import ( # noqa: F401
CompletedWorkflowStageMetadata,
FailedWorkflowStageMetadata,
RunningWorkflowStageMetadata,
SkippedWorkflowStageMetadata,
WorkflowMetadata,
WorkflowStageMetadata,
WorkflowStageMetadataVariant,
)

_LAZY_IMPORTS: dict[str, tuple[str, str]] = {
"CompositeWorkflow": ("data_designer.interface.composite_workflow", "CompositeWorkflow"),
"CompositeWorkflowResults": ("data_designer.interface.composite_workflow", "CompositeWorkflowResults"),
"CompletedWorkflowStageMetadata": (
"data_designer.interface.workflow_metadata",
"CompletedWorkflowStageMetadata",
),
"DataDesigner": ("data_designer.interface.data_designer", "DataDesigner"),
"DataDesignerEarlyShutdownError": ("data_designer.interface.errors", "DataDesignerEarlyShutdownError"),
"DataDesignerGenerationError": ("data_designer.interface.errors", "DataDesignerGenerationError"),
"DataDesignerProfilingError": ("data_designer.interface.errors", "DataDesignerProfilingError"),
"DataDesignerWorkflowError": ("data_designer.interface.errors", "DataDesignerWorkflowError"),
"DatasetCreationResults": ("data_designer.interface.results", "DatasetCreationResults"),
"FailedWorkflowStageMetadata": ("data_designer.interface.workflow_metadata", "FailedWorkflowStageMetadata"),
"ResumeMode": ("data_designer.config.run_config", "ResumeMode"),
"RunningWorkflowStageMetadata": ("data_designer.interface.workflow_metadata", "RunningWorkflowStageMetadata"),
"SkippedStageResult": ("data_designer.interface.composite_workflow", "SkippedStageResult"),
"SkippedStageStatus": ("data_designer.interface.composite_workflow", "SkippedStageStatus"),
"SkippedWorkflowStageMetadata": ("data_designer.interface.workflow_metadata", "SkippedWorkflowStageMetadata"),
"WorkflowMetadata": ("data_designer.interface.workflow_metadata", "WorkflowMetadata"),
"WorkflowStageMetadata": ("data_designer.interface.workflow_metadata", "WorkflowStageMetadata"),
"WorkflowStageMetadataVariant": (
"data_designer.interface.workflow_metadata",
"WorkflowStageMetadataVariant",
),
}

__all__ = list(_LAZY_IMPORTS.keys())
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
_export_jsonl,
_export_parquet,
)
from data_designer.interface.workflow_metadata import WorkflowMetadata, WorkflowStageMetadata

if TYPE_CHECKING:
import pandas as pd
Expand Down Expand Up @@ -219,8 +220,19 @@ def add_stage(
_validate_dir_name(name, "stage name")
if any(stage.name == name for stage in self._stages):
raise DataDesignerWorkflowError(f"Stage name {name!r} is already used in workflow {self.name!r}.")
if num_records is not None and num_records < 1:
raise DataDesignerWorkflowError("Stage num_records must be at least 1.")
if num_records is not None:
if not isinstance(num_records, int) or isinstance(num_records, bool):
raise DataDesignerWorkflowError("Stage num_records must be an integer.")
if num_records < 1:
raise DataDesignerWorkflowError("Stage num_records must be at least 1.")
if on_success_version is not None and not isinstance(on_success_version, str):
raise DataDesignerWorkflowError("Stage on_success_version must be a string.")
if not isinstance(allow_empty, bool):
raise DataDesignerWorkflowError("Stage allow_empty must be a boolean.")
if not isinstance(sampling_strategy, SamplingStrategy):
raise DataDesignerWorkflowError("Stage sampling_strategy must be a SamplingStrategy.")
if selection_strategy is not None and not isinstance(selection_strategy, (IndexRange, PartitionBlock)):
raise DataDesignerWorkflowError("Stage selection_strategy must be an IndexRange or PartitionBlock.")
_validate_stage_output(output)
output_processors = output_processors or []
_validate_distinct_output_processors(config_builder, output_processors)
Expand Down Expand Up @@ -276,6 +288,7 @@ def run(
workflow_path.mkdir(parents=True, exist_ok=True)
prior_metadata = _read_prior_workflow_metadata(workflow_path, self.name, resume)
metadata: dict[str, Any] = {
**(prior_metadata or {}),
"name": self.name,
"library_version": get_library_version(),
"stages": [],
Expand Down Expand Up @@ -388,6 +401,10 @@ def run(
f"Cannot resume workflow {self.name!r}: stage {stage.name!r} is not reusable."
)

if stage_resume == ResumeMode.ALWAYS and prior_stage_metadata is not None:
prior_stage = WorkflowStageMetadata.model_validate(prior_stage_metadata).root
stage_metadata.update(prior_stage.model_extra or {})

if stage_resume == ResumeMode.NEVER and stage_path.exists():
shutil.rmtree(stage_path)

Expand Down Expand Up @@ -598,14 +615,30 @@ def _read_prior_workflow_metadata(
raise DataDesignerWorkflowError(
f"Cannot resume workflow {workflow_name!r}: workflow metadata has invalid shape."
)
try:
validated_metadata = WorkflowMetadata.model_validate(metadata)
except ValidationError as exc:
if resume != ResumeMode.ALWAYS:
error = exc.errors(include_url=False, include_input=False)[0]
location = ".".join(str(part) for part in error["loc"])
logger.warning(
"Workflow metadata for %r has invalid field %s (%s); starting fresh.",
workflow_name,
location,
error["msg"],
)
return None
raise DataDesignerWorkflowError(
f"Cannot resume workflow {workflow_name!r}: workflow metadata has invalid shape."
) from exc
if metadata.get("name") != workflow_name:
if resume != ResumeMode.ALWAYS:
logger.warning("Workflow metadata for %r has a different name; starting fresh.", workflow_name)
return None
raise DataDesignerWorkflowError(
f"Cannot resume workflow {workflow_name!r}: workflow metadata name does not match."
)
return metadata
return validated_metadata.model_dump(mode="json", exclude_unset=True)


def _get_prior_stage_metadata(
Expand Down Expand Up @@ -870,9 +903,13 @@ def _parquet_files(path: Path) -> list[Path]:
def _write_workflow_metadata(workflow_path: Path, metadata: dict[str, Any]) -> None:
path = workflow_path / WORKFLOW_METADATA_FILENAME
tmp_path = path.with_name(f"{path.name}.tmp.{os.getpid()}.{uuid.uuid4().hex}")
try:
validated_metadata = WorkflowMetadata.model_validate(metadata)
except ValidationError as exc:
raise DataDesignerWorkflowError("Workflow metadata has invalid shape.") from exc
try:
with tmp_path.open("w", encoding="utf-8") as f:
json.dump(metadata, f, indent=2, sort_keys=True)
json.dump(validated_metadata.model_dump(mode="json", exclude_unset=True), f, indent=2, sort_keys=True)
f.flush()
os.fsync(f.fileno())
os.replace(tmp_path, path)
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,90 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

from __future__ import annotations

from typing import Annotated, Any, Literal, TypeAlias

from pydantic import BaseModel, ConfigDict, Field, RootModel


class _WorkflowStageMetadataBase(BaseModel):
"""Common metadata for a composite workflow stage."""

model_config = ConfigDict(extra="allow")

status: Literal["running", "failed", "completed", "completed_empty", "skipped_empty_upstream"]
index: int
name: str
stage_dir: str
depends_on: list[str]
allow_empty: bool
on_success_version: str | None
output_processors: list[dict[str, Any]]
output: str
sampling_strategy: str
selection_strategy: dict[str, Any] | None


class _StartedWorkflowStageMetadata(_WorkflowStageMetadataBase):
fingerprint: str
num_records_requested: int
seeded_from_stage: str | None
seed_path: str | None
config: dict[str, Any]


class RunningWorkflowStageMetadata(_StartedWorkflowStageMetadata):
"""Metadata for a running workflow stage."""

status: Literal["running"]


class FailedWorkflowStageMetadata(_StartedWorkflowStageMetadata):
"""Metadata for a failed workflow stage."""

status: Literal["failed"]
duration_sec: float | None = None


class CompletedWorkflowStageMetadata(_StartedWorkflowStageMetadata):
"""Metadata for a completed workflow stage."""

status: Literal["completed", "completed_empty"]
num_records_actual: int
output_records: int
output_seed_path: str
callback_output_path: str | None
stage_output_override_path: str | None = None
output_processor_output_path: str | None
duration_sec: float


class SkippedWorkflowStageMetadata(_WorkflowStageMetadataBase):
"""Metadata for a stage skipped after an empty upstream stage."""

status: Literal["skipped_empty_upstream"]
upstream_stage: str


WorkflowStageMetadataVariant: TypeAlias = Annotated[
RunningWorkflowStageMetadata
| FailedWorkflowStageMetadata
| CompletedWorkflowStageMetadata
| SkippedWorkflowStageMetadata,
Field(discriminator="status"),
]


class WorkflowStageMetadata(RootModel[WorkflowStageMetadataVariant]):
"""Status-specific stage metadata, with the concrete model available through ``root``."""


class WorkflowMetadata(BaseModel):
"""Metadata persisted for a composite workflow run."""

model_config = ConfigDict(extra="allow")

name: str
library_version: str
stages: list[WorkflowStageMetadataVariant]
Loading
Loading