diff --git a/packages/data-designer/src/data_designer/interface/__init__.py b/packages/data-designer/src/data_designer/interface/__init__.py index d4a112f74..183d932f1 100644 --- a/packages/data-designer/src/data_designer/interface/__init__.py +++ b/packages/data-designer/src/data_designer/interface/__init__.py @@ -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()) diff --git a/packages/data-designer/src/data_designer/interface/composite_workflow.py b/packages/data-designer/src/data_designer/interface/composite_workflow.py index 408083be5..23c9ca622 100644 --- a/packages/data-designer/src/data_designer/interface/composite_workflow.py +++ b/packages/data-designer/src/data_designer/interface/composite_workflow.py @@ -40,6 +40,7 @@ _export_jsonl, _export_parquet, ) +from data_designer.interface.workflow_metadata import WorkflowMetadata, WorkflowStageMetadata if TYPE_CHECKING: import pandas as pd @@ -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) @@ -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": [], @@ -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) @@ -598,6 +615,22 @@ 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) @@ -605,7 +638,7 @@ def _read_prior_workflow_metadata( 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( @@ -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) diff --git a/packages/data-designer/src/data_designer/interface/workflow_metadata.py b/packages/data-designer/src/data_designer/interface/workflow_metadata.py new file mode 100644 index 000000000..84970ff6c --- /dev/null +++ b/packages/data-designer/src/data_designer/interface/workflow_metadata.py @@ -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] diff --git a/packages/data-designer/tests/interface/test_composite_workflow.py b/packages/data-designer/tests/interface/test_composite_workflow.py index eb9a7abad..48e0c749b 100644 --- a/packages/data-designer/tests/interface/test_composite_workflow.py +++ b/packages/data-designer/tests/interface/test_composite_workflow.py @@ -6,6 +6,7 @@ import json import shutil from pathlib import Path +from typing import Any from unittest.mock import MagicMock import pytest @@ -22,6 +23,7 @@ from data_designer.config.seed_source_dataframe import DataFrameSeedSource from data_designer.engine.secret_resolver import PlaintextResolver from data_designer.engine.storage.artifact_storage import ArtifactStorage, BatchStage, ResumeMode +from data_designer.interface import WorkflowMetadata from data_designer.interface.composite_workflow import SkippedStageResult, SkippedStageStatus from data_designer.interface.data_designer import DataDesigner from data_designer.interface.errors import DataDesignerWorkflowError @@ -194,6 +196,9 @@ def test_composite_workflow_runs_linear_stages_with_disk_handoff( assert data_designer.artifact_path == stub_artifact_path metadata = _load_workflow_metadata(stub_artifact_path, "linear-chain") + metadata_path = stub_artifact_path / "linear-chain" / "workflow-metadata.json" + validated_metadata = WorkflowMetadata.model_validate_json(metadata_path.read_text()) + assert validated_metadata.model_dump(mode="json", exclude_unset=True) == metadata assert [stage["status"] for stage in metadata["stages"]] == ["completed", "completed"] assert metadata["stages"][1]["seeded_from_stage"] == "base" assert metadata["stages"][1]["depends_on"] == ["base"] @@ -364,6 +369,33 @@ def test_composite_workflow_rejects_invalid_stage_outputs( workflow.add_stage("base", _category_builder(stub_model_configs), output=output) +@pytest.mark.parametrize( + ("kwargs", "match"), + [ + ({"num_records": "3"}, "num_records must be an integer"), + ({"num_records": True}, "num_records must be an integer"), + ({"num_records": 0}, "num_records must be at least 1"), + ({"on_success_version": 1}, "on_success_version must be a string"), + ({"allow_empty": "yes"}, "allow_empty must be a boolean"), + ({"sampling_strategy": "ordered"}, "sampling_strategy must be a SamplingStrategy"), + ({"selection_strategy": {}}, "selection_strategy must be an IndexRange or PartitionBlock"), + ], +) +def test_composite_workflow_rejects_invalid_stage_metadata_before_artifacts( + stub_artifact_path: Path, + stub_model_providers: list[ModelProvider], + stub_model_configs: list[ModelConfig], + kwargs: dict[str, Any], + match: str, +) -> None: + workflow = _data_designer(stub_artifact_path, stub_model_providers).compose_workflow(name="invalid-version") + + with pytest.raises(DataDesignerWorkflowError, match=match): + workflow.add_stage("base", _category_builder(stub_model_configs), **kwargs) + + assert not stub_artifact_path.exists() + + def test_composite_workflow_rejects_unknown_processor_stage_output( stub_artifact_path: Path, stub_model_providers: list[ModelProvider], @@ -568,7 +600,7 @@ def test_composite_workflow_stage_output_override_path_must_exist( assert create_mock.call_count == 0 -def test_composite_workflow_resume_if_possible_skips_completed_stages( +def test_composite_workflow_resume_if_possible_skips_legacy_completed_stages( stub_artifact_path: Path, stub_model_providers: list[ModelProvider], stub_model_configs: list[ModelConfig], @@ -580,6 +612,13 @@ def test_composite_workflow_resume_if_possible_skips_completed_stages( workflow.add_stage("base", _category_builder(stub_model_configs), num_records=3) workflow.add_stage("copy", _copy_builder(stub_model_configs)) workflow.run() + metadata_path = stub_artifact_path / "resume-skip" / "workflow-metadata.json" + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + metadata["workflow_extension"] = True + for stage in metadata["stages"]: + stage.pop("stage_output_override_path") + stage["stage_extension"] = {"value": stage["index"]} + metadata_path.write_text(json.dumps(metadata), encoding="utf-8") create_mock.reset_mock() resumed = data_designer.compose_workflow(name="resume-skip") @@ -590,6 +629,10 @@ def test_composite_workflow_resume_if_possible_skips_completed_stages( assert create_mock.call_count == 0 assert results.count_records() == 3 assert results.load_dataset()["category"].tolist() == ["alpha", "alpha", "alpha"] + resumed_metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + assert resumed_metadata["workflow_extension"] is True + assert [stage["stage_extension"] for stage in resumed_metadata["stages"]] == [{"value": 0}, {"value": 1}] + assert all("stage_output_override_path" not in stage for stage in resumed_metadata["stages"]) def test_composite_workflow_resume_if_possible_skips_stage_with_output_processors( @@ -825,6 +868,30 @@ def test_composite_workflow_resume_if_possible_invalid_metadata_shape_starts_fre assert [call.kwargs["dataset_name"] for call in create_mock.call_args_list] == ["stage-0-base", "stage-1-copy"] +def test_composite_workflow_resume_if_possible_invalid_stage_metadata_starts_fresh( + stub_artifact_path: Path, + stub_model_providers: list[ModelProvider], + stub_model_configs: list[ModelConfig], + stub_dataset_profiler_results, +) -> None: + data_designer = _data_designer(stub_artifact_path, stub_model_providers) + create_mock = _patch_create(data_designer, stub_dataset_profiler_results) + workflow = data_designer.compose_workflow(name="resume-invalid-stage") + workflow.add_stage("base", _category_builder(stub_model_configs), num_records=2) + workflow.run() + metadata_path = stub_artifact_path / "resume-invalid-stage" / "workflow-metadata.json" + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + metadata["stages"][0]["status"] = "unknown" + metadata_path.write_text(json.dumps(metadata), encoding="utf-8") + create_mock.reset_mock() + + resumed = data_designer.compose_workflow(name="resume-invalid-stage") + resumed.add_stage("base", _category_builder(stub_model_configs), num_records=2) + resumed.run(resume=ResumeMode.IF_POSSIBLE) + + assert [call.kwargs["dataset_name"] for call in create_mock.call_args_list] == ["stage-0-base"] + + @pytest.mark.parametrize("status", ["running", "failed"]) def test_composite_workflow_resume_if_possible_delegates_matching_resumable_stage( stub_artifact_path: Path, @@ -842,8 +909,20 @@ def test_composite_workflow_resume_if_possible_delegates_matching_resumable_stag metadata_path = stub_artifact_path / "resume-partial" / "workflow-metadata.json" metadata = json.loads(metadata_path.read_text(encoding="utf-8")) _mark_stage_resumable(metadata, 0, status) + # duration_sec is schema-owned by failed metadata, but is an extension on running metadata. + if status == "failed": + metadata["stages"][0]["duration_sec"] = 123.0 + metadata["stages"][0]["stage_extension"] = {"value": status} metadata_path.write_text(json.dumps(metadata), encoding="utf-8") create_mock.reset_mock() + create_side_effect = create_mock.side_effect + metadata_at_create: list[dict[str, Any]] = [] + + def capture_create(*args: Any, **kwargs: Any) -> DatasetCreationResults: + metadata_at_create.append(json.loads(metadata_path.read_text(encoding="utf-8"))) + return create_side_effect(*args, **kwargs) + + create_mock.side_effect = capture_create resumed = data_designer.compose_workflow(name="resume-partial") resumed.add_stage("base", _category_builder(stub_model_configs), num_records=2) @@ -852,6 +931,10 @@ def test_composite_workflow_resume_if_possible_delegates_matching_resumable_stag assert [call.kwargs["dataset_name"] for call in create_mock.call_args_list] == ["stage-0-base", "stage-1-copy"] assert [call.kwargs["resume"] for call in create_mock.call_args_list] == [ResumeMode.ALWAYS, ResumeMode.NEVER] + assert metadata_at_create[0]["stages"][0]["stage_extension"] == {"value": status} + assert "duration_sec" not in metadata_at_create[0]["stages"][0] + resumed_metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + assert resumed_metadata["stages"][0]["stage_extension"] == {"value": status} def test_composite_workflow_resume_always_reruns_descendants_after_partial_stage( @@ -936,6 +1019,29 @@ def test_composite_workflow_resume_always_rejects_invalid_metadata_shape( resumed.run(resume=ResumeMode.ALWAYS) +def test_composite_workflow_resume_always_rejects_invalid_stage_metadata( + stub_artifact_path: Path, + stub_model_providers: list[ModelProvider], + stub_model_configs: list[ModelConfig], + stub_dataset_profiler_results, +) -> None: + data_designer = _data_designer(stub_artifact_path, stub_model_providers) + _patch_create(data_designer, stub_dataset_profiler_results) + workflow = data_designer.compose_workflow(name="resume-invalid-stage-always") + workflow.add_stage("base", _category_builder(stub_model_configs), num_records=2) + workflow.run() + metadata_path = stub_artifact_path / "resume-invalid-stage-always" / "workflow-metadata.json" + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + metadata["stages"][0]["status"] = "unknown" + metadata_path.write_text(json.dumps(metadata), encoding="utf-8") + + resumed = data_designer.compose_workflow(name="resume-invalid-stage-always") + resumed.add_stage("base", _category_builder(stub_model_configs), num_records=2) + + with pytest.raises(DataDesignerWorkflowError, match="workflow metadata has invalid shape"): + resumed.run(resume=ResumeMode.ALWAYS) + + def test_composite_workflow_resume_always_rejects_changed_stage( stub_artifact_path: Path, stub_model_providers: list[ModelProvider], diff --git a/packages/data-designer/tests/interface/test_workflow_metadata.py b/packages/data-designer/tests/interface/test_workflow_metadata.py new file mode 100644 index 000000000..7e6e82a21 --- /dev/null +++ b/packages/data-designer/tests/interface/test_workflow_metadata.py @@ -0,0 +1,189 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Any + +import pytest +from pydantic import ValidationError + +import data_designer.interface as interface +from data_designer.interface import ( + CompletedWorkflowStageMetadata, + FailedWorkflowStageMetadata, + RunningWorkflowStageMetadata, + SkippedWorkflowStageMetadata, + WorkflowMetadata, + WorkflowStageMetadata, + WorkflowStageMetadataVariant, +) + + +@pytest.fixture +def base_stage_metadata() -> dict[str, Any]: + return { + "index": 0, + "name": "base", + "stage_dir": "stage-0-base", + "depends_on": [], + "allow_empty": False, + "on_success_version": None, + "output_processors": [], + "output": "final", + "sampling_strategy": "ordered", + "selection_strategy": None, + } + + +@pytest.fixture +def started_stage_metadata() -> dict[str, Any]: + return { + "fingerprint": "abc123", + "num_records_requested": 2, + "seeded_from_stage": None, + "seed_path": None, + "config": {"columns": []}, + } + + +def test_workflow_metadata_models_are_public() -> None: + assert interface.WorkflowMetadata is WorkflowMetadata + assert interface.WorkflowStageMetadata is WorkflowStageMetadata + assert interface.WorkflowStageMetadataVariant is WorkflowStageMetadataVariant + assert interface.RunningWorkflowStageMetadata is RunningWorkflowStageMetadata + assert interface.FailedWorkflowStageMetadata is FailedWorkflowStageMetadata + assert interface.CompletedWorkflowStageMetadata is CompletedWorkflowStageMetadata + assert interface.SkippedWorkflowStageMetadata is SkippedWorkflowStageMetadata + + +@pytest.mark.parametrize( + ("status_fields", "expected_type"), + [ + ({"status": "running"}, RunningWorkflowStageMetadata), + ({"status": "failed", "duration_sec": 0.25}, FailedWorkflowStageMetadata), + ( + { + "status": "completed", + "num_records_actual": 2, + "output_records": 2, + "output_seed_path": "stage-0-base/parquet-files", + "callback_output_path": None, + "stage_output_override_path": None, + "output_processor_output_path": None, + "duration_sec": 1.5, + }, + CompletedWorkflowStageMetadata, + ), + ( + { + "status": "completed_empty", + "num_records_actual": 0, + "output_records": 0, + "output_seed_path": "stage-0-base/parquet-files", + "callback_output_path": None, + "stage_output_override_path": None, + "output_processor_output_path": None, + "duration_sec": 1.5, + }, + CompletedWorkflowStageMetadata, + ), + ( + {"status": "skipped_empty_upstream", "upstream_stage": "base"}, + SkippedWorkflowStageMetadata, + ), + ], +) +def test_workflow_metadata_supports_stage_statuses( + base_stage_metadata: dict[str, Any], + started_stage_metadata: dict[str, Any], + status_fields: dict[str, Any], + expected_type: type, +) -> None: + stage = base_stage_metadata | status_fields + if status_fields["status"] != "skipped_empty_upstream": + stage |= started_stage_metadata + stage["stage_extension"] = {"value": status_fields["status"]} + + standalone_metadata = WorkflowStageMetadata.model_validate(stage) + metadata = WorkflowMetadata.model_validate({"name": "example", "library_version": "0.9.2", "stages": [stage]}) + + assert isinstance(standalone_metadata.root, expected_type) + assert standalone_metadata.model_dump(mode="json", exclude_unset=True) == stage + assert WorkflowStageMetadata.model_validate_json(standalone_metadata.model_dump_json()) == standalone_metadata + assert isinstance(metadata.stages[0], expected_type) + restored = WorkflowMetadata.model_validate_json(metadata.model_dump_json()) + assert restored == metadata + + +def test_workflow_metadata_supports_legacy_completed_stage( + base_stage_metadata: dict[str, Any], + started_stage_metadata: dict[str, Any], +) -> None: + stage = ( + base_stage_metadata + | started_stage_metadata + | { + "status": "completed", + "num_records_actual": 2, + "output_records": 2, + "output_seed_path": "/tmp/artifacts/example/stage-0-base/parquet-files", + "callback_output_path": None, + "output_processor_output_path": None, + "duration_sec": 1.5, + } + ) + + metadata = WorkflowMetadata.model_validate({"name": "example", "library_version": "0.9.2", "stages": [stage]}) + + assert metadata.stages[0].stage_output_override_path is None + assert "stage_output_override_path" not in metadata.model_dump(mode="json", exclude_unset=True)["stages"][0] + + +def test_workflow_metadata_preserves_extra_fields( + base_stage_metadata: dict[str, Any], + started_stage_metadata: dict[str, Any], +) -> None: + stage = base_stage_metadata | started_stage_metadata | {"status": "running", "stage_extension": {"value": 1}} + metadata = WorkflowMetadata.model_validate( + { + "name": "example", + "library_version": "0.9.2", + "stages": [stage], + "workflow_extension": True, + } + ) + + payload = metadata.model_dump(mode="json", exclude_unset=True) + + assert payload["workflow_extension"] is True + assert payload["stages"][0]["stage_extension"] == {"value": 1} + + +@pytest.mark.parametrize( + "stage_fields", + [ + {"status": "unknown"}, + {"status": "running"}, + {"status": "failed", "duration_sec": "invalid"}, + {"status": "completed"}, + {"status": "skipped_empty_upstream"}, + ], +) +def test_workflow_metadata_rejects_invalid_stage_metadata( + base_stage_metadata: dict[str, Any], + stage_fields: dict[str, Any], +) -> None: + stage = base_stage_metadata | stage_fields + + with pytest.raises(ValidationError): + WorkflowStageMetadata.model_validate(stage) + + with pytest.raises(ValidationError): + WorkflowMetadata.model_validate( + { + "name": "example", + "library_version": "0.9.2", + "stages": [stage], + } + )