Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 12 additions & 3 deletions afp/schemas.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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(
Expand Down
30 changes: 15 additions & 15 deletions afp/validators.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"
)


Expand Down
32 changes: 21 additions & 11 deletions tests/test_validators.py
Original file line number Diff line number Diff line change
Expand Up @@ -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": {
Expand All @@ -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": {
Expand All @@ -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(
Expand Down