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
55 changes: 28 additions & 27 deletions src/rune/runtime/metadata.py
Original file line number Diff line number Diff line change
@@ -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'
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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),
Expand All @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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),
Expand Down
2 changes: 1 addition & 1 deletion test/test_basic_types_with_meta.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)


Expand Down
60 changes: 57 additions & 3 deletions test/test_keys_and_references.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
'''test key generation/retrieval runtime functions'''
import json
from decimal import Decimal
from typing_extensions import Annotated
import pytest
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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),
Expand Down Expand Up @@ -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():
Expand Down
Loading