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
13 changes: 11 additions & 2 deletions src/unitxt/dataclass.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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):
Expand Down Expand Up @@ -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
}
Expand Down
34 changes: 34 additions & 0 deletions tests/library/test_dataclass.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,13 @@
import sys
import unittest
from dataclasses import field
from typing import Callable

from unitxt.dataclass import (
AbstractField,
AbstractFieldError,
Dataclass,
DataclassMeta,
FinalField,
FinalFieldError,
MissingDefaultError,
Expand All @@ -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
Expand Down