52 lines
1.7 KiB
Python
52 lines
1.7 KiB
Python
import pytest
|
|
|
|
from infrasynth.workflows.validators import (
|
|
DataValidatorProtocol,
|
|
DataValidatorRegistry,
|
|
)
|
|
|
|
|
|
class _RequiredFieldValidator:
|
|
def validate(self, node, data, context):
|
|
if "resolution_note" not in data:
|
|
raise ValueError("resolution_note is required")
|
|
return data
|
|
|
|
|
|
class TestDataValidatorProtocol:
|
|
def test_protocol_runtime_check(self):
|
|
assert isinstance(_RequiredFieldValidator(), DataValidatorProtocol)
|
|
|
|
def test_non_matching_class_not_compatible(self):
|
|
class NotAValidator:
|
|
pass
|
|
|
|
assert not isinstance(NotAValidator(), DataValidatorProtocol)
|
|
|
|
|
|
class TestDataValidatorRegistry:
|
|
def teardown_method(self):
|
|
DataValidatorRegistry._validators.clear()
|
|
|
|
def test_register_and_get(self):
|
|
validator = _RequiredFieldValidator()
|
|
DataValidatorRegistry.register("ticket_approval", validator)
|
|
assert DataValidatorRegistry.get("ticket_approval") is validator
|
|
|
|
def test_get_unknown_returns_none(self):
|
|
assert DataValidatorRegistry.get("unknown") is None
|
|
|
|
def test_register_overwrites(self):
|
|
v1 = _RequiredFieldValidator()
|
|
v2 = _RequiredFieldValidator()
|
|
DataValidatorRegistry.register("wf", v1)
|
|
DataValidatorRegistry.register("wf", v2)
|
|
assert DataValidatorRegistry.get("wf") is v2
|
|
|
|
def test_validator_raises_on_missing_field(self):
|
|
with pytest.raises(ValueError, match="resolution_note"):
|
|
_RequiredFieldValidator().validate(None, {}, {})
|
|
|
|
def test_validator_passes_with_field(self):
|
|
data = {"resolution_note": "Fixed"}
|
|
assert _RequiredFieldValidator().validate(None, data, {}) == data
|