diff --git a/langfuse/_utils/serializer.py b/langfuse/_utils/serializer.py index 53889f29a..d284b5055 100644 --- a/langfuse/_utils/serializer.py +++ b/langfuse/_utils/serializer.py @@ -5,7 +5,7 @@ import math from asyncio import Queue from collections.abc import Sequence -from dataclasses import asdict, is_dataclass +from dataclasses import fields, is_dataclass from datetime import date, datetime from json import JSONEncoder from logging import getLogger @@ -127,7 +127,19 @@ def _default_inner(self, obj: Any) -> Any: return f"<{type(obj).__name__}>" if is_dataclass(obj): - return asdict(obj) # type: ignore + obj_id = id(obj) + + if obj_id in self.seen: + return type(obj).__name__ + + self.seen.add(obj_id) + try: + return { + field.name: self.default(getattr(obj, field.name)) + for field in fields(obj) + } + finally: + self.seen.remove(obj_id) if isinstance(obj, BaseModel): obj.model_rebuild() diff --git a/tests/unit/test_serializer.py b/tests/unit/test_serializer.py index ce5798f67..2d65bc8e5 100644 --- a/tests/unit/test_serializer.py +++ b/tests/unit/test_serializer.py @@ -165,6 +165,23 @@ def __init__(self): assert result == {"next": {"next": "Node"}} +def test_circular_dataclass_reference(): + @dataclass + class Node: + name: str + next: "Node | None" = None + + node1 = Node("first") + node2 = Node("second") + node1.next = node2 + node2.next = node1 + + serializer = EventSerializer() + result = json.loads(serializer.encode(node1)) + + assert result == {"name": "first", "next": {"name": "second", "next": "Node"}} + + def test_not_serializable(): class NotSerializable: def __init__(self):