Discriminate TestOutputAssertions union on class key

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) <noreply@anthropic.com>
This commit is contained in:
John Chilton
2026-04-22 09:53:59 -04:00
co-authored by Claude Opus 4.7
parent aef25d3cb0
commit 7b764303a2
5 changed files with 143 additions and 4 deletions
+45 -4
View File
@@ -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]]
@@ -98,6 +98,7 @@
collection_type: list:paired_or_unpaired
elements:
el1:
class: Collection
elements:
forward:
asserts:
@@ -90,6 +90,7 @@
collection_type: list:list
elements:
test_level_3:
class: Collection
elements:
test_level_2:
asserts:
@@ -81,6 +81,7 @@
collection_type: list:list
elements:
test_level_3:
class: Collection
elements:
test_level_2:
asserts:
@@ -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