diff --git a/afp/schemas.py b/afp/schemas.py index d2acd7b..260c8af 100644 --- a/afp/schemas.py +++ b/afp/schemas.py @@ -1,6 +1,7 @@ """AFP data structures.""" from decimal import Decimal +from itertools import chain from typing import Annotated, Any, ClassVar, Literal, Self from pydantic import AfterValidator, BeforeValidator, Field, model_validator @@ -321,9 +322,17 @@ def _cross_validate(self) -> Self: self.product.min_price, self.product.max_price, ) - validators.validate_outcome_space_conditions( - self.outcome_space.base_case.condition, - [case.condition for case in self.outcome_space.edge_cases], + validators.validate_outcome_space_template_variables( + [ + self.outcome_space.base_case.condition, + self.outcome_space.base_case.fsp_resolution, + ] + + list( + chain.from_iterable( + [edge_case.condition, edge_case.fsp_resolution] + for edge_case in self.outcome_space.edge_cases + ) + ), self.outcome_point.model_dump(), ) if isinstance(self.outcome_space, OutcomeSpaceTimeSeries) and isinstance( diff --git a/afp/validators.py b/afp/validators.py index 8b28844..86a162c 100644 --- a/afp/validators.py +++ b/afp/validators.py @@ -3,7 +3,7 @@ from decimal import Decimal from functools import reduce from operator import getitem -from typing import Any +from typing import Any, Iterable import requests from binascii import Error @@ -178,23 +178,23 @@ def validate_oracle_fallback_fsp( ) -def validate_outcome_space_conditions( - base_case_condition: str, - edge_case_conditions: list[str], - outcome_point_dict: dict[Any, Any], +def validate_outcome_space_template_variables( + values: Iterable[str], outcome_point_dict: dict[Any, Any] ) -> None: - conditions = [base_case_condition] + edge_case_conditions - schemas = ["BaseCaseResolution"] + [ - f"EdgeCase[{i}]" for i in range(len(edge_case_conditions)) - ] - for condition, schema in zip(conditions, schemas): - for variable in re.findall(r"{(.+?)}", condition): - parts = variable.split(".") + for value in values: + for variable in re.findall(r"{(.*?)}", value): try: - reduce(getitem, parts, outcome_point_dict) - except KeyError: + referred_value = reduce( + getitem, variable.split("."), outcome_point_dict + ) + except (TypeError, KeyError): + raise ValueError( + f"OutcomeSpace: Invalid template variable '{variable}'" + ) + if isinstance(referred_value, dict) or isinstance(referred_value, list): # type: ignore raise ValueError( - f"{schema}: condition: Invalid template variable '{variable}'" + f"OutcomeSpace: Template variable '{variable}' " + "should not refer to a nested object or list" ) diff --git a/tests/test_validators.py b/tests/test_validators.py index 4a14b46..c6e2e8f 100644 --- a/tests/test_validators.py +++ b/tests/test_validators.py @@ -110,7 +110,7 @@ def test_validate_price_limits__error(): validators.validate_price_limits(Decimal("0.11"), Decimal("0.10")) -def test_validate_outcome_space_conditions__pass(): +def test_validate_outcome_space_template_variables__pass(): dct = { "a": { "b": { @@ -119,15 +119,29 @@ def test_validate_outcome_space_conditions__pass(): }, }, } - base_case_condition = "Reference to {a.b.c}" - edge_case_conditions = ["And to {a.b.d} as well"] - validators.validate_outcome_space_conditions( - base_case_condition, edge_case_conditions, dct + validators.validate_outcome_space_template_variables( + [ + "Reference to {a.b.c} should pass", + "And to {a.b.d} as well", + "So as having no template variable", + ], + dct, ) -def test_validate_outcome_space_conditions__error(): +@pytest.mark.parametrize( + "value", + [ + "{a.b.c} and {a.b.e}", + "{c}", + "{a.b.c.d}", + "{a.b}", + "{}", + ], + ids=str, +) +def test_validate_outcome_space_tempate_variables__error(value): dct = { "a": { "b": { @@ -136,13 +150,9 @@ def test_validate_outcome_space_conditions__error(): }, }, } - base_case_condition = "Reference to {a.b.e}" - edge_case_conditions = [] with pytest.raises(ValueError): - validators.validate_outcome_space_conditions( - base_case_condition, edge_case_conditions, dct - ) + validators.validate_outcome_space_template_variables([value], dct) @pytest.mark.parametrize(