From 8eb7892dc761ebade6db840449591bffebdb7032 Mon Sep 17 00:00:00 2001 From: Excelius <57819425+Excelius-Wang@users.noreply.github.com> Date: Thu, 20 Aug 2026 13:38:19 +0800 Subject: [PATCH] fix: support deferred class annotations on Python 3.14 Signed-off-by: Excelius <57819425+Excelius-Wang@users.noreply.github.com> --- src/unitxt/dataclass.py | 13 +++++++++++-- tests/library/test_dataclass.py | 34 +++++++++++++++++++++++++++++++++ 2 files changed, 45 insertions(+), 2 deletions(-) diff --git a/src/unitxt/dataclass.py b/src/unitxt/dataclass.py index afb92bcf0e..b483ba85d0 100644 --- a/src/unitxt/dataclass.py +++ b/src/unitxt/dataclass.py @@ -152,6 +152,15 @@ def is_possible_field(field_name, field_value): return True +def _get_class_annotations(obj): + cls = obj if isinstance(obj, type) else type(obj) + if hasattr(inspect, "Format"): + return inspect.get_annotations(cls, format=inspect.Format.FORWARDREF) + if hasattr(inspect, "get_annotations"): + return inspect.get_annotations(cls) + return dict(vars(cls).get("__annotations__", {})) + + def get_fields(cls, attrs): """Get the fields for a class based on its attributes. @@ -165,7 +174,7 @@ def get_fields(cls, attrs): fields = {} for base in cls.__bases__: fields = {**getattr(base, _FIELDS, {}), **fields} - annotations = {**attrs.get("__annotations__", {})} + annotations = _get_class_annotations(cls) for attr_name, attr_value in attrs.items(): if attr_name not in annotations and is_possible_field(attr_name, attr_value): @@ -593,7 +602,7 @@ def to_dict(self, classes: Optional[List] = None, keep_empty: bool = True): else: attributes = [] for cls in classes: - attributes += list(cls.__annotations__.keys()) + attributes += list(_get_class_annotations(cls).keys()) attributes_dict = { attribute: getattr(self, attribute) for attribute in attributes } diff --git a/tests/library/test_dataclass.py b/tests/library/test_dataclass.py index fcbc373198..3dc4daa105 100644 --- a/tests/library/test_dataclass.py +++ b/tests/library/test_dataclass.py @@ -1,3 +1,5 @@ +import sys +import unittest from dataclasses import field from typing import Callable @@ -5,6 +7,7 @@ AbstractField, AbstractFieldError, Dataclass, + DataclassMeta, FinalField, FinalFieldError, MissingDefaultError, @@ -25,6 +28,37 @@ class TestDataclass(UnitxtTestCase): + @unittest.skipUnless(sys.version_info >= (3, 14), "requires deferred annotations") + def test_self_referencing_annotation(self): + class SelfReferencingDataclass(Dataclass): + child: SelfReferencingDataclass | None = None # noqa: F821 + + instance = SelfReferencingDataclass() + + self.assertIsNone(instance.child) + self.assertListEqual(fields_names(SelfReferencingDataclass), ["child"]) + + def test_annotations_not_stored_in_class_namespace(self): + class DeferredAnnotationsMeta(DataclassMeta): + def __init__(self, name, bases, attrs): + attrs.pop("__annotations__", None) + super().__init__(name, bases, attrs) + + class DeferredAnnotationsDataclass( + Dataclass, metaclass=DeferredAnnotationsMeta + ): + pass + + class Dummy(DeferredAnnotationsDataclass): + name: str + count: int = 0 + + dummy = Dummy(name="example") + + self.assertEqual(dummy.name, "example") + self.assertEqual(dummy.count, 0) + self.assertListEqual(fields_names(Dummy), ["name", "count"]) + def test_dataclass(self): class GrandParent(Dataclass): a: int