From 0a46a12765875a1bce34dd14aa7c6cf3b92c64b4 Mon Sep 17 00:00:00 2001 From: Plamen Neykov Date: Wed, 16 Sep 2026 19:14:59 +0100 Subject: [PATCH] Fix reference binding for inherited fields Use Pydantic model_fields to retrieve inherited field annotations, preventing TypeError while preserving reference target type checks. Add regression coverage for multilevel inheritance, incompatible targets, and deserialization with validation enabled and disabled. Fixes #39 --- src/rune/runtime/metadata.py | 55 +++++++++++++-------------- test/test_basic_types_with_meta.py | 2 +- test/test_keys_and_references.py | 60 ++++++++++++++++++++++++++++-- 3 files changed, 86 insertions(+), 31 deletions(-) diff --git a/src/rune/runtime/metadata.py b/src/rune/runtime/metadata.py index 58ffc11..1f4067a 100644 --- a/src/rune/runtime/metadata.py +++ b/src/rune/runtime/metadata.py @@ -1,15 +1,16 @@ '''Classes representing annotated basic Rune types''' -import uuid import datetime import importlib -from enum import Enum -from functools import partial, lru_cache +import uuid +from collections.abc import Iterable from decimal import Decimal -from typing import Any, Never, get_args, Iterable -from typing_extensions import Self, Tuple -from pydantic import (PlainSerializer, PlainValidator, WrapValidator, - WrapSerializer) +from enum import Enum +from functools import lru_cache, partial +from typing import Any, Never, Self, get_args + +from pydantic import PlainSerializer, PlainValidator, WrapSerializer, WrapValidator from pydantic_core import PydanticCustomError + # from rune.runtime.object_registry import get_object DEFAULT_META = '_ALLOWED_METADATA' @@ -135,7 +136,7 @@ def get_reference(self, _): class UnresolvedReference(BaseReference): '''used by the deserialization to hold temporarily unresolved references''' def __init__(self, key): - rune_type, self.key = list(key.items())[0] + rune_type, self.key = next(iter(key.items())) self.key_type = KeyType.from_rune(rune_type) def get_reference(self, parent): @@ -219,7 +220,7 @@ def get_or_create_key(self) -> str: self.set_meta(key=key) try: self._get_object_map(KeyType.INTERNAL)[key] = self - except: # noqa + except: self.set_meta(key=None) raise return key @@ -238,7 +239,7 @@ def set_external_key(self, self.set_meta(check_allowed=True, **{key_type.key_tag: key}) try: self._get_object_map(key_type)[key] = self - except: # noqa + except: self.set_meta(check_allowed=True, **{key_type.key_tag: None}) raise return self @@ -298,13 +299,14 @@ def _bind_property_to(self, property_nm: str, ref: Reference): f'not allowed for {property_nm}. Allowed types ' f'are: {allowed_ref_types.get(property_nm, {})}') - field_type = self.__class__.__annotations__.get(property_nm) + # Pydantic includes inherited fields and resolves their annotations. + field_type = type(self).model_fields[property_nm].annotation # type: ignore allowed_type = _get_basic_type(field_type) if not (isinstance(allowed_type, str) or isinstance(ref.target, allowed_type)): - raise ValueError("Can't set reference. Incompatible types: " - f"expected {allowed_type}, " - f"got {ref.target.__class__}") + raise TypeError("Can't set reference. Incompatible types: " + f"expected {allowed_type}, " + f"got {ref.target.__class__}") refs = self.__dict__.setdefault(REFS_CONTAINER, {}) if property_nm not in refs: @@ -416,9 +418,8 @@ def _type_to_cls(cls, prefix = namespace_prefix if prefix is None: prefix = cls._get_rune_namespace_prefix() - if prefix: - if not rune_type.startswith(prefix + '.'): - import_path = prefix + '.' + rune_type + if prefix and not rune_type.startswith(prefix + '.'): + import_path = prefix + '.' + rune_type rune_module = importlib.import_module(import_path) return getattr(rune_module, rune_class_name) return cls # support for legacy json @@ -453,14 +454,14 @@ def deserialize(cls, obj, allowed_meta: set[str]): 'Expected either {my_type} or dict but ' 'got {type}.', {'type': type(obj), 'my_type': cls}) - metadata = {k: obj[k] for k in obj.keys() if k.startswith('@')} + metadata = {k: obj[k] for k in obj if k.startswith('@')} # References deserialization treatment if aux := cls._create_unresolved_ref(metadata): return aux # Model creation - for k in metadata.keys(): + for k in metadata: obj.pop(k) rune_cls = cls._type_to_cls(metadata) @@ -483,7 +484,7 @@ def serializer(cls): @classmethod @lru_cache - def validator(cls, allowed_meta: tuple[str] | tuple[Never, ...] = tuple()): + def validator(cls, allowed_meta: tuple[str] | tuple[Never, ...] = ()): '''default validator for the specific class''' allowed = set(allowed_meta) return PlainValidator(partial(cls.deserialize, allowed_meta=allowed), @@ -492,16 +493,16 @@ def validator(cls, allowed_meta: tuple[str] | tuple[Never, ...] = tuple()): class BasicTypeMetaDataMixin(BaseMetaDataMixin): '''holds the metadata associated with an instance''' - _INPUT_TYPES: Any | Tuple[Any, ...] = str # to be overridden by subclasses + _INPUT_TYPES: Any | tuple[Any, ...] = str # to be overridden by subclasses _OUTPUT_TYPE: Any = str # to be overridden by subclasses _JSON_OUTPUT = str | dict @classmethod def _check_type(cls, value): if not isinstance(value, cls._INPUT_TYPES): - raise ValueError(f'{cls.__name__} can be instantiated only with ' - f'one of the following type(s): {cls._INPUT_TYPES},' - f' however the value is of type {type(value)}') + raise TypeError(f'{cls.__name__} can be instantiated only with ' + f'one of the following type(s): {cls._INPUT_TYPES},' + f' however the value is of type {type(value)}') @classmethod def serialise(cls, obj, base_type) -> dict: @@ -514,7 +515,7 @@ def serialise(cls, obj, base_type) -> dict: def deserialize(cls, obj, handler, base_types, allowed_meta: set[str]): '''method used as pydantic `validator`''' if isinstance(obj, list): - identity = lambda x: x # noqa: E731 + identity = lambda x: x processed = [cls.deserialize(item, identity, base_types, allowed_meta) for item in obj] return handler(processed) model = obj @@ -649,7 +650,7 @@ class _EnumWrapper(BaseMetaDataMixin): '''wrapper for enums with metadata''' def __init__(self, enum_instance=_EnumWrapperDefaultVal.NOT_SET): if not isinstance(enum_instance, Enum): - raise ValueError("enum_instance must be an instance of an Enum") + raise TypeError("enum_instance must be an instance of an Enum") self._enum_instance = enum_instance @property @@ -723,7 +724,7 @@ def serializer(cls): @classmethod @lru_cache - def validator(cls, allowed_meta: tuple[str] | tuple[Never, ...] = tuple()): + def validator(cls, allowed_meta: tuple[str] | tuple[Never, ...] = ()): '''default validator for the specific class''' allowed = set(allowed_meta) return PlainValidator(partial(cls.deserialize, allowed_meta=allowed), diff --git a/test/test_basic_types_with_meta.py b/test/test_basic_types_with_meta.py index 166e040..c33123f 100644 --- a/test/test_basic_types_with_meta.py +++ b/test/test_basic_types_with_meta.py @@ -235,7 +235,7 @@ def test_annotated_date_fail(): def test_date_with_meta_fail(): '''test instantiation failure with an incorrect type''' - with pytest.raises(ValueError): + with pytest.raises(TypeError): DateWithMeta(10) diff --git a/test/test_keys_and_references.py b/test/test_keys_and_references.py index 7ccfa1c..322231c 100644 --- a/test/test_keys_and_references.py +++ b/test/test_keys_and_references.py @@ -1,4 +1,5 @@ '''test key generation/retrieval runtime functions''' +import json from decimal import Decimal from typing_extensions import Annotated import pytest @@ -62,6 +63,23 @@ class DummyLoan2(BaseDataClass): 'repayment': {'@ref', '@ref:external'} } + +class NamedLoan(DummyLoan2): + '''Inherit both reference fields without redeclaring their annotations.''' + label: str = 'loan' + + +class ExtendedNamedLoan(NamedLoan): + '''Reference fields can be inherited through multiple levels.''' + pass + + +class LoanBook(BaseDataClass): + '''Key targets and inherited references in separate list fields.''' + cashflows: list[CashFlow] + loans: list[NamedLoan] + + class DummyLoan3(BaseDataClass): '''number test class''' loan: Annotated[NumberWithMeta, @@ -220,6 +238,38 @@ def test_ref_assign(): assert id(model.loan) == id(model.repayment) +@pytest.mark.parametrize('model_type', [NamedLoan, ExtendedNamedLoan]) +def test_ref_assign_to_inherited_field(model_type): + '''Inherited reference fields retain their declared target type.''' + model = model_type(loan=CashFlow(currency='EUR', amount=100), + repayment=CashFlow(currency='EUR', amount=101)) + + model.repayment = Reference(model.loan) + + assert model.repayment is model.loan + assert model.resolve_ref_key('repayment') == model.loan.get_meta('@key') + + +@pytest.mark.parametrize('validate_model', [False, True]) +def test_deserialize_inherited_references_in_lists(validate_model): + '''Inherited references resolve before optional model validation.''' + data = json.dumps({ + 'cashflows': [{ + '@key:external': 'cashflow1', 'currency': 'EUR', 'amount': '100', + }], + 'loans': [{ + 'loan': {'@ref:external': 'cashflow1'}, + 'repayment': {'@ref:external': 'cashflow1'}, + }], + }) + + book = LoanBook.rune_deserialize(data, validate_model=validate_model) + + assert book.loans[0].loan is book.cashflows[0] + assert book.loans[0].repayment is book.cashflows[0] + assert book.loans[0].resolve_ref_key('repayment') == 'cashflow1' + + def test_ref_assign_from_cow_wrapped_object(): '''test use a ref from a COW-wrapped object''' model = DummyLoan2(loan=CashFlow(currency='EUR', amount=100), @@ -445,13 +495,17 @@ def test_reference_non_metadata_target_with_ext_key_raises(): Reference('not_an_object', 'ext_key') -def test_bind_property_rejects_wrong_type(): +@pytest.mark.parametrize('model_type', [DummyLoan2, NamedLoan, ExtendedNamedLoan]) +def test_bind_property_rejects_wrong_type(model_type): '''reject refs that don't match field type''' - model = DummyLoan2(loan=CashFlow(currency='EUR', amount=100), + model = model_type(loan=CashFlow(currency='EUR', amount=100), repayment=CashFlow(currency='EUR', amount=101)) + original_repayment = model.repayment other = OtherThing(name='other') - with pytest.raises(ValueError): + with pytest.raises(TypeError, match='Incompatible types'): model.repayment = Reference(other) + assert model.repayment is original_repayment + assert model.resolve_ref_key('repayment') is None def test_bind_property_rejects_non_replaceable():