From 7b764303a2ddb6081d9785968aed3e2afdaa64c7 Mon Sep 17 00:00:00 2001 From: John Chilton Date: Mon, 20 Apr 2026 18:48:01 -0400 Subject: [PATCH] Discriminate TestOutputAssertions union on class key MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Collapses un-tagged Union into tagged File/Collection/scalar so validation errors point at one branch, not three. Adds class_ field to nested element types for symmetry. Fixes three gxwf-tests YAMLs that were missing class: Collection on nested collection outputs — the un-discriminated Union had been hiding the bug. Co-Authored-By: Claude Opus 4.7 (1M context) --- lib/galaxy/tool_util_models/__init__.py | 49 +++++++++- .../collection_semantics_cat.gxwf-tests.yml | 1 + .../subcollection_rank_sorting.gxwf-tests.yml | 1 + ...lection_rank_sorting_paired.gxwf-tests.yml | 1 + .../test_output_assertions.py | 95 +++++++++++++++++++ 5 files changed, 143 insertions(+), 4 deletions(-) create mode 100644 test/unit/tool_util_models/test_output_assertions.py diff --git a/lib/galaxy/tool_util_models/__init__.py b/lib/galaxy/tool_util_models/__init__.py index 7565e001c9d..adaf69413b4 100644 --- a/lib/galaxy/tool_util_models/__init__.py +++ b/lib/galaxy/tool_util_models/__init__.py @@ -16,9 +16,11 @@ from pydantic import ( AnyUrl, BaseModel, ConfigDict, + Discriminator, Field, model_validator, RootModel, + Tag, ) from typing_extensions import ( Annotated, @@ -190,16 +192,33 @@ class TestDataOutputAssertions(BaseTestOutputModel): class TestCollectionCollectionElementAssertions(StrictModel): + class_: Optional[Literal["Collection"]] = Field("Collection", alias="class") elements: Optional[Dict[str, "TestCollectionElementAssertion"]] = None element_tests: Optional[Dict[str, "TestCollectionElementAssertion"]] = None class TestCollectionDatasetElementAssertions(BaseTestOutputModel): - pass + class_: Optional[Literal["File"]] = Field("File", alias="class") -TestCollectionElementAssertion = Union[ - TestCollectionDatasetElementAssertions, TestCollectionCollectionElementAssertions +def _discriminate_collection_element(v): + if isinstance(v, dict): + if v.get("class") == "Collection": + return "Collection" + return "File" + if isinstance(v, TestCollectionCollectionElementAssertions): + return "Collection" + if isinstance(v, TestCollectionDatasetElementAssertions): + return "File" + return None + + +TestCollectionElementAssertion = Annotated[ + Union[ + Annotated[TestCollectionDatasetElementAssertions, Tag("File")], + Annotated[TestCollectionCollectionElementAssertions, Tag("Collection")], + ], + Discriminator(_discriminate_collection_element), ] TestCollectionCollectionElementAssertions.model_rebuild() @@ -219,7 +238,29 @@ class TestCollectionOutputAssertions(StrictModel): TestOutputLiteral = Union[bool, int, float, str] -TestOutputAssertions = Union[TestCollectionOutputAssertions, TestDataOutputAssertions, TestOutputLiteral] + +def _discriminate_output(v): + if isinstance(v, dict): + if v.get("class") == "Collection": + return "Collection" + return "File" + if isinstance(v, TestCollectionOutputAssertions): + return "Collection" + if isinstance(v, TestDataOutputAssertions): + return "File" + if isinstance(v, (bool, int, float, str)): + return "scalar" + return None + + +TestOutputAssertions = Annotated[ + Union[ + Annotated[TestCollectionOutputAssertions, Tag("Collection")], + Annotated[TestDataOutputAssertions, Tag("File")], + Annotated[TestOutputLiteral, Tag("scalar")], + ], + Discriminator(_discriminate_output), +] TestInputValue = Union[bool, int, float, str, List[Any], Dict[str, Any]] diff --git a/lib/galaxy_test/workflow/collection_semantics_cat.gxwf-tests.yml b/lib/galaxy_test/workflow/collection_semantics_cat.gxwf-tests.yml index 1cfd3adf027..e4a9a4137d1 100644 --- a/lib/galaxy_test/workflow/collection_semantics_cat.gxwf-tests.yml +++ b/lib/galaxy_test/workflow/collection_semantics_cat.gxwf-tests.yml @@ -98,6 +98,7 @@ collection_type: list:paired_or_unpaired elements: el1: + class: Collection elements: forward: asserts: diff --git a/lib/galaxy_test/workflow/subcollection_rank_sorting.gxwf-tests.yml b/lib/galaxy_test/workflow/subcollection_rank_sorting.gxwf-tests.yml index 1de9fc991a5..7555515538a 100644 --- a/lib/galaxy_test/workflow/subcollection_rank_sorting.gxwf-tests.yml +++ b/lib/galaxy_test/workflow/subcollection_rank_sorting.gxwf-tests.yml @@ -90,6 +90,7 @@ collection_type: list:list elements: test_level_3: + class: Collection elements: test_level_2: asserts: diff --git a/lib/galaxy_test/workflow/subcollection_rank_sorting_paired.gxwf-tests.yml b/lib/galaxy_test/workflow/subcollection_rank_sorting_paired.gxwf-tests.yml index 68c4faae75f..3e62371b1ac 100644 --- a/lib/galaxy_test/workflow/subcollection_rank_sorting_paired.gxwf-tests.yml +++ b/lib/galaxy_test/workflow/subcollection_rank_sorting_paired.gxwf-tests.yml @@ -81,6 +81,7 @@ collection_type: list:list elements: test_level_3: + class: Collection elements: test_level_2: asserts: diff --git a/test/unit/tool_util_models/test_output_assertions.py b/test/unit/tool_util_models/test_output_assertions.py new file mode 100644 index 00000000000..2fa8823ac26 --- /dev/null +++ b/test/unit/tool_util_models/test_output_assertions.py @@ -0,0 +1,95 @@ +"""Tests for the discriminated `TestOutputAssertions` Union.""" + +import json + +import pytest +from pydantic import ValidationError + +from galaxy.tool_util_models import ( + Tests, + TestCollectionCollectionElementAssertions, + TestCollectionOutputAssertions, + TestDataOutputAssertions, +) + + +def _one_test(outputs): + return [{"doc": "t", "job": {}, "outputs": outputs}] + + +def test_implicit_file_output_validates_as_data(): + tests = Tests.model_validate(_one_test({"out": {"asserts": [{"that": "has_text", "text": "x"}]}})) + out = tests.root[0].outputs["out"] + assert isinstance(out, TestDataOutputAssertions) + + +def test_explicit_collection_output_validates_as_collection(): + tests = Tests.model_validate(_one_test({"out": {"class": "Collection", "elements": {"a": {"asserts": []}}}})) + out = tests.root[0].outputs["out"] + assert isinstance(out, TestCollectionOutputAssertions) + + +@pytest.mark.parametrize("value", [True, 42, 3.14, "hello"]) +def test_scalar_literal_output_validates(value): + tests = Tests.model_validate(_one_test({"out": value})) + assert tests.root[0].outputs["out"] == value + + +def test_nested_collection_element_with_class_collection_validates(): + tests = Tests.model_validate( + _one_test( + { + "out": { + "class": "Collection", + "elements": { + "inner": { + "class": "Collection", + "elements": {"leaf": {"asserts": []}}, + } + }, + } + } + ) + ) + out = tests.root[0].outputs["out"] + assert isinstance(out, TestCollectionOutputAssertions) + assert out.elements is not None + inner = out.elements["inner"] + assert isinstance(inner, TestCollectionCollectionElementAssertions) + + +def test_unknown_field_on_file_output_yields_single_error(): + with pytest.raises(ValidationError) as exc: + Tests.model_validate(_one_test({"out": {"asserts": [], "garbage_key": 1}})) + errs = exc.value.errors() + assert len(errs) == 1 + err = errs[0] + assert err["loc"][:4] == (0, "outputs", "out", "File") + assert err["loc"][-1] == "garbage_key" + assert err["type"] == "extra_forbidden" + + +def test_unknown_class_value_yields_single_class_literal_error(): + with pytest.raises(ValidationError) as exc: + Tests.model_validate(_one_test({"out": {"class": "Banana"}})) + errs = exc.value.errors() + assert len(errs) == 1 + err = errs[0] + assert err["loc"][:4] == (0, "outputs", "out", "File") + assert err["type"] == "literal_error" + + +def test_bad_asserts_error_scoped_to_file_branch(): + with pytest.raises(ValidationError) as exc: + Tests.model_validate(_one_test({"out": {"asserts": [{"that": "not_a_real_assert"}]}})) + errs = exc.value.errors() + for err in errs: + assert "Collection" not in err["loc"] + assert "scalar" not in err["loc"] + assert err["loc"][:4] == (0, "outputs", "out", "File") + + +def test_json_schema_emits_discriminator_for_outputs(): + schema = Tests.model_json_schema() + dumped = json.dumps(schema) + assert "discriminator" in dumped or '"oneOf"' in dumped