mirror of
https://github.com/galaxyproject/galaxy.git
synced 2026-09-24 16:30:27 +08:00
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:
co-authored by
Claude Opus 4.7
parent
aef25d3cb0
commit
7b764303a2
@@ -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
|
||||
Reference in New Issue
Block a user