Skip to content
Open
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
1 change: 1 addition & 0 deletions packages/overture-schema-system/changelog.d/760.bugfix.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Typed the model-constraint decorators to return the class they decorate, so pyright and ty keep a decorated model's fields instead of reducing it to `BaseModel`.
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from typing_extensions import override

from .._json_schema import get_static_json_schema_extra, put_if, required_non_null
from ..create_model import ModelT
from .model_constraint import (
Condition,
OptionalFieldGroupConstraint,
Expand All @@ -19,7 +20,7 @@
def forbid_if(
field_names: list[str] | tuple[str, ...],
condition: Condition,
) -> Callable[[type[BaseModel]], type[BaseModel]]:
) -> Callable[[type[ModelT]], type[ModelT]]:
"""
Decorate a Pydantic model class with a constraint forbidding any of the named fields from
holding a non-`None` value, but only if a field value condition is true.
Expand All @@ -37,8 +38,8 @@ def forbid_if(

Returns
-------
Callable
Decorator factory
Callable[[type[ModelT]], type[ModelT]]
Decorator that applies the constraint to a Pydantic model class

Example
-------
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,11 @@
from typing_extensions import override

from .._json_schema import get_static_json_schema_extra
from ..create_model import ModelT
from .model_constraint import ModelConstraint


def min_fields_set(count: int) -> Callable[[type[BaseModel]], type[BaseModel]]:
def min_fields_set(count: int) -> Callable[[type[ModelT]], type[ModelT]]:
"""
Decorate a Pydantic model class with a constraint that requires a minimum number of fields in
the model to be set to a non-`None` value.
Expand All @@ -26,8 +27,8 @@ def min_fields_set(count: int) -> Callable[[type[BaseModel]], type[BaseModel]]:

Returns
-------
type[BaseModel]
Decorated Pydantic model class
Callable[[type[ModelT]], type[ModelT]]
Decorator that applies the constraint to a Pydantic model class

Example
-------
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@
from pydantic.json_schema import JsonDict, to_jsonable_python
from typing_extensions import override

from ..create_model import create_model
from ..create_model import ModelT, create_model
from ..metadata import Key, Metadata


Expand Down Expand Up @@ -62,7 +62,7 @@ def name(self) -> str:
return self.__name

@final
def decorate(self, model_class: type[BaseModel]) -> type[BaseModel]:
def decorate(self, model_class: type[ModelT]) -> type[ModelT]:
"""
Decorate a Pydantic model, returning a new version of the model that has this constraint
applied to it.
Expand All @@ -71,13 +71,13 @@ def decorate(self, model_class: type[BaseModel]) -> type[BaseModel]:

Parameters
----------
model_class : type[BaseModel]
model_class : type[ModelT]
Pydantic model to decorate. It is not decorated in-place, rather a new version of the
model class is returned with this constraint attached to it.

Returns
-------
type[BaseModel]
type[ModelT]
New version of `model_class` with this constraint applied to it

Example
Expand All @@ -94,7 +94,7 @@ def decorate(self, model_class: type[BaseModel]) -> type[BaseModel]:
... raise ValueError('the `foo` field must equal "bar"')
...
>>> # Define a decorator.
>>> def foo(model_class: type[BaseModel]) -> type[BaseModel]:
>>> def foo(model_class: type[ModelT]) -> type[ModelT]:
... return FooConstraint().decorate(model_class)
...
>>> # Apply the decorator.
Expand Down Expand Up @@ -131,7 +131,7 @@ def _after_validator(model_instance: BaseModel) -> BaseModel:
constraint.validate_instance(model_instance)
return model_instance

new_model_class = create_model(
return create_model(
model_class.__name__,
__config__=config,
__doc__=model_class.__doc__,
Expand All @@ -145,7 +145,6 @@ def _after_validator(model_instance: BaseModel) -> BaseModel:
},
__metadata__=metadata,
)
return new_model_class

def validate_class(self, model_class: type[BaseModel]) -> None:
"""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,10 +5,11 @@
from pydantic import BaseModel, ConfigDict
from typing_extensions import override

from ..create_model import ModelT
from .model_constraint import ModelConstraint


def no_extra_fields(model_class: type[BaseModel]) -> type[BaseModel]:
def no_extra_fields(model_class: type[ModelT]) -> type[ModelT]:
"""
Decorate a Pydantic model class with a constraint that forbids extra fields that aren't
explicitly part of the model.
Expand All @@ -18,12 +19,12 @@ def no_extra_fields(model_class: type[BaseModel]) -> type[BaseModel]:

Parameters
----------
model_class: type[BaseModel]
model_class : type[ModelT]
Pydantic model class being decorated

Returns
-------
type[BaseModel]
type[ModelT]
Decorated Pydantic model class

Example
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -11,10 +11,11 @@
from typing_extensions import override

from .._json_schema import get_static_json_schema_extra, put_one_of
from ..create_model import ModelT
from .model_constraint import FieldGroupConstraint, apply_alias


def radio_group(*field_names: str) -> Callable[[type[BaseModel]], type[BaseModel]]:
def radio_group(*field_names: str) -> Callable[[type[ModelT]], type[ModelT]]:
"""
Decorate a Pydantic model class with a constraint requiring that exactly one field in a group of
`bool` fields has the value `True`.
Expand All @@ -34,8 +35,8 @@ def radio_group(*field_names: str) -> Callable[[type[BaseModel]], type[BaseModel

Returns
-------
Callable
Decorator factory
Callable[[type[ModelT]], type[ModelT]]
Decorator that applies the constraint to a Pydantic model class

Example
-------
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,11 @@
from typing_extensions import override

from .._json_schema import get_static_json_schema_extra, put_any_of, required_non_null
from ..create_model import ModelT
from .model_constraint import OptionalFieldGroupConstraint, apply_alias


def require_any_of(*field_names: str) -> Callable[[type[BaseModel]], type[BaseModel]]:
def require_any_of(*field_names: str) -> Callable[[type[ModelT]], type[ModelT]]:
"""
Decorate a Pydantic model class with a constraint requiring that at least one of the named
fields has a non-`None` value.
Expand All @@ -29,8 +30,8 @@ def require_any_of(*field_names: str) -> Callable[[type[BaseModel]], type[BaseMo

Returns
-------
Callable
Decorator factory
Callable[[type[ModelT]], type[ModelT]]
Decorator that applies the constraint to a Pydantic model class

Example
-------
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,12 +9,18 @@
from typing_extensions import override

from .._json_schema import get_static_json_schema_extra, put_any_of
from .model_constraint import Condition, FieldEqCondition, ModelConstraint, apply_alias
from ..create_model import ModelT
from .model_constraint import (
Condition,
FieldEqCondition,
ModelConstraint,
apply_alias,
)


def require_any_true(
*conditions: Condition,
) -> Callable[[type[BaseModel]], type[BaseModel]]:
) -> Callable[[type[ModelT]], type[ModelT]]:
"""
Decorate a Pydantic model class with a constraint requiring at least one condition in a group
of conditions to evaluate to `True`.
Expand All @@ -28,8 +34,8 @@ def require_any_true(

Returns
-------
Callable
Decorator factory
Callable[[type[ModelT]], type[ModelT]]
Decorator that applies the constraint to a Pydantic model class

Example
-------
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from typing_extensions import override

from .._json_schema import get_static_json_schema_extra, put_if, required_non_null
from ..create_model import ModelT
from .model_constraint import (
Condition,
OptionalFieldGroupConstraint,
Expand All @@ -19,7 +20,7 @@
def require_if(
field_names: list[str] | tuple[str, ...],
condition: Condition,
) -> Callable[[type[BaseModel]], type[BaseModel]]:
) -> Callable[[type[ModelT]], type[ModelT]]:
"""
Decorate a Pydantic model class with a constraint requiring all of the named fields to have a
non-`None` value, but only if a condition is true.
Expand All @@ -38,8 +39,8 @@ def require_if(

Returns
-------
Callable
Decorator factory
Callable[[type[ModelT]], type[ModelT]]
Decorator that applies the constraint to a Pydantic model class

Example
-------
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
"""
Each model-constraint decorator returns the decorated class typed as that class, not as `BaseModel`.

mypy ignores class-decorator return types, so a check written with `@decorator` syntax passes
even when the decorator erases the class (as pyright and ty then do). Calling each decorator as a
function puts its return type in front of mypy.
"""

from pydantic import BaseModel
from typing_extensions import assert_type

from overture.schema.system.model_constraint import (
FieldEqCondition,
NoExtraFieldsConstraint,
forbid_if,
min_fields_set,
no_extra_fields,
radio_group,
require_any_of,
require_any_true,
require_if,
)


class Model(BaseModel):
f: int | None = None
g: int | None = None
b: bool | None = None
c: bool | None = None


def test_decorators_preserve_the_decorated_class_type() -> None:
decorated = [
assert_type(NoExtraFieldsConstraint().decorate(Model), type[Model]),
assert_type(no_extra_fields(Model), type[Model]),
assert_type(require_if(["f"], FieldEqCondition("g", 1))(Model), type[Model]),
assert_type(forbid_if(["f"], FieldEqCondition("g", 1))(Model), type[Model]),
assert_type(require_any_of("f", "g")(Model), type[Model]),
assert_type(require_any_true(FieldEqCondition("b", True))(Model), type[Model]),
assert_type(radio_group("b", "c")(Model), type[Model]),
assert_type(min_fields_set(1)(Model), type[Model]),
]

for model_class in decorated:
assert issubclass(model_class, Model)
Loading