diff --git a/packages/overture-schema-system/changelog.d/760.bugfix.md b/packages/overture-schema-system/changelog.d/760.bugfix.md new file mode 100644 index 000000000..83ffe1817 --- /dev/null +++ b/packages/overture-schema-system/changelog.d/760.bugfix.md @@ -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`. diff --git a/packages/overture-schema-system/src/overture/schema/system/model_constraint/forbid_if.py b/packages/overture-schema-system/src/overture/schema/system/model_constraint/forbid_if.py index c301edbd4..6f6558235 100644 --- a/packages/overture-schema-system/src/overture/schema/system/model_constraint/forbid_if.py +++ b/packages/overture-schema-system/src/overture/schema/system/model_constraint/forbid_if.py @@ -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, @@ -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. @@ -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 ------- diff --git a/packages/overture-schema-system/src/overture/schema/system/model_constraint/min_fields_set.py b/packages/overture-schema-system/src/overture/schema/system/model_constraint/min_fields_set.py index 7462e6a8b..c3e2de202 100644 --- a/packages/overture-schema-system/src/overture/schema/system/model_constraint/min_fields_set.py +++ b/packages/overture-schema-system/src/overture/schema/system/model_constraint/min_fields_set.py @@ -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. @@ -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 ------- diff --git a/packages/overture-schema-system/src/overture/schema/system/model_constraint/model_constraint.py b/packages/overture-schema-system/src/overture/schema/system/model_constraint/model_constraint.py index ae279da7a..56631359b 100644 --- a/packages/overture-schema-system/src/overture/schema/system/model_constraint/model_constraint.py +++ b/packages/overture-schema-system/src/overture/schema/system/model_constraint/model_constraint.py @@ -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 @@ -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. @@ -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 @@ -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. @@ -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__, @@ -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: """ diff --git a/packages/overture-schema-system/src/overture/schema/system/model_constraint/no_extra_fields.py b/packages/overture-schema-system/src/overture/schema/system/model_constraint/no_extra_fields.py index 4d24f92d7..35c8282c6 100644 --- a/packages/overture-schema-system/src/overture/schema/system/model_constraint/no_extra_fields.py +++ b/packages/overture-schema-system/src/overture/schema/system/model_constraint/no_extra_fields.py @@ -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. @@ -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 diff --git a/packages/overture-schema-system/src/overture/schema/system/model_constraint/radio_group.py b/packages/overture-schema-system/src/overture/schema/system/model_constraint/radio_group.py index 20e7a1202..529efd40d 100644 --- a/packages/overture-schema-system/src/overture/schema/system/model_constraint/radio_group.py +++ b/packages/overture-schema-system/src/overture/schema/system/model_constraint/radio_group.py @@ -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`. @@ -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 ------- diff --git a/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_any_of.py b/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_any_of.py index f05e86043..46091d191 100644 --- a/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_any_of.py +++ b/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_any_of.py @@ -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. @@ -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 ------- diff --git a/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_any_true.py b/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_any_true.py index 88c3d1a61..af139ed37 100644 --- a/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_any_true.py +++ b/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_any_true.py @@ -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`. @@ -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 ------- diff --git a/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_if.py b/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_if.py index fbc354d2f..a10fca785 100644 --- a/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_if.py +++ b/packages/overture-schema-system/src/overture/schema/system/model_constraint/require_if.py @@ -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, @@ -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. @@ -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 ------- diff --git a/packages/overture-schema-system/tests/model_constraint/test_decorator_types.py b/packages/overture-schema-system/tests/model_constraint/test_decorator_types.py new file mode 100644 index 000000000..806bf054d --- /dev/null +++ b/packages/overture-schema-system/tests/model_constraint/test_decorator_types.py @@ -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)