diff --git a/py/src/braintrust/parameters.py b/py/src/braintrust/parameters.py index c8b71ddc..ca85ef8d 100644 --- a/py/src/braintrust/parameters.py +++ b/py/src/braintrust/parameters.py @@ -221,7 +221,7 @@ def _prompt_data_to_dict( if prompt_data is None: return None if isinstance(prompt_data, PromptData): - return prompt_data.as_dict() + return prompt_data.as_dict(exclude_unset=True) return dict(prompt_data) diff --git a/py/src/braintrust/serializable_data_class.py b/py/src/braintrust/serializable_data_class.py index ba32ecca..af7c971e 100644 --- a/py/src/braintrust/serializable_data_class.py +++ b/py/src/braintrust/serializable_data_class.py @@ -1,16 +1,105 @@ +import copy import dataclasses import json import types -from typing import Union, get_origin +from typing import Any, Union, get_origin + + +_EXPLICITLY_SET_FIELDS_ATTR = "_braintrust_explicitly_set_fields" +_INIT_FIELDS_REMAINING_ATTR = "_braintrust_init_fields_remaining" + + +def _dataclass_fields(cls_or_instance: Any) -> tuple[dataclasses.Field, ...]: + try: + return dataclasses.fields(cls_or_instance) + except TypeError: + return () + + +def _explicit_constructor_fields(cls: type, args: tuple[Any, ...], kwargs: dict[str, Any]) -> set[str]: + fields = _dataclass_fields(cls) + positional_fields = [f.name for f in fields if f.init and not getattr(f, "kw_only", False)] + init_field_names = {f.name for f in fields if f.init} + + explicit_fields = set(positional_fields[: len(args)]) + explicit_fields.update(k for k in kwargs if k in init_field_names) + return explicit_fields + + +def _field_names(cls_or_instance: Any) -> set[str]: + return {f.name for f in _dataclass_fields(cls_or_instance)} + + +def _init_field_names(cls_or_instance: Any) -> set[str]: + return {f.name for f in _dataclass_fields(cls_or_instance) if f.init} + + +def _clear_init_tracking(obj: Any) -> None: + if hasattr(obj, _INIT_FIELDS_REMAINING_ATTR): + object.__delattr__(obj, _INIT_FIELDS_REMAINING_ATTR) class SerializableDataClass: - def as_dict(self): + def __new__(cls, *args, **kwargs): + instance = super().__new__(cls) + init_fields = _init_field_names(cls) + if init_fields: + object.__setattr__(instance, _INIT_FIELDS_REMAINING_ATTR, init_fields) + object.__setattr__(instance, _EXPLICITLY_SET_FIELDS_ATTR, _explicit_constructor_fields(cls, args, kwargs)) + return instance + + def __setattr__(self, name: str, value: Any) -> None: + object.__setattr__(self, name, value) + + if name not in _field_names(type(self)): + return + + init_fields_remaining = getattr(self, _INIT_FIELDS_REMAINING_ATTR, None) + if init_fields_remaining is not None: + init_fields_remaining.discard(name) + if not init_fields_remaining: + _clear_init_tracking(self) + return + + getattr(self, _EXPLICITLY_SET_FIELDS_ATTR).add(name) + + def __post_init__(self) -> None: + _clear_init_tracking(self) + + def __getstate__(self) -> dict[str, Any]: + state = self.__dict__.copy() + # copy/deepcopy and pickle rebuild via __new__ without running dataclass __init__, + # so do not persist the temporary constructor-assignment tracking marker. + state.pop(_INIT_FIELDS_REMAINING_ATTR, None) + return state + + def __setstate__(self, state: dict[str, Any]) -> None: + self.__dict__.update(state) + explicitly_set_fields = getattr(self, _EXPLICITLY_SET_FIELDS_ATTR, None) + if explicitly_set_fields is not None: + object.__setattr__(self, _EXPLICITLY_SET_FIELDS_ATTR, set(explicitly_set_fields)) + _clear_init_tracking(self) + + def as_dict(self, exclude_unset: bool = False): """Serialize the object to a dictionary.""" - return dataclasses.asdict(self) + if not exclude_unset: + return dataclasses.asdict(self) + + explicitly_set_fields = getattr(self, _EXPLICITLY_SET_FIELDS_ATTR, None) + if explicitly_set_fields is None: + return dataclasses.asdict(self) + explicitly_set_fields = set(explicitly_set_fields) - def as_json(self, **kwargs): + return { + f.name: _as_dict_value(getattr(self, f.name), exclude_unset=exclude_unset) + for f in dataclasses.fields(self) + if f.name in explicitly_set_fields or getattr(self, f.name) is not None + } + + def as_json(self, exclude_unset: bool = False, **kwargs): """Serialize the object to JSON.""" + if exclude_unset: + return json.dumps(self.as_dict(exclude_unset=True), **kwargs) return json.dumps(self.as_dict(), **kwargs) def __getitem__(self, item: str): @@ -64,3 +153,17 @@ def from_dict_deep(cls, d: dict): else: filtered[k] = v return cls(**filtered) + + +def _as_dict_value(value: Any, exclude_unset: bool) -> Any: + if isinstance(value, SerializableDataClass): + return value.as_dict(exclude_unset=exclude_unset) + if dataclasses.is_dataclass(value) and not isinstance(value, type): + return dataclasses.asdict(value) + if isinstance(value, list): + return [_as_dict_value(v, exclude_unset=exclude_unset) for v in value] + if isinstance(value, tuple): + return tuple(_as_dict_value(v, exclude_unset=exclude_unset) for v in value) + if isinstance(value, dict): + return {copy.deepcopy(k): _as_dict_value(v, exclude_unset=exclude_unset) for k, v in value.items()} + return copy.deepcopy(value) diff --git a/py/src/braintrust/test_serializable_data_class.py b/py/src/braintrust/test_serializable_data_class.py index e31e2078..3e233ed1 100644 --- a/py/src/braintrust/test_serializable_data_class.py +++ b/py/src/braintrust/test_serializable_data_class.py @@ -1,7 +1,9 @@ +import copy +import pickle import unittest -from dataclasses import dataclass +from dataclasses import dataclass, field -from .serializable_data_class import SerializableDataClass +from .serializable_data_class import _EXPLICITLY_SET_FIELDS_ATTR, _INIT_FIELDS_REMAINING_ATTR, SerializableDataClass @dataclass @@ -22,6 +24,24 @@ class PromptSchema(SerializableDataClass): tags: list[str] | None +@dataclass +class Child(SerializableDataClass): + value: str | None = None + label: str = "child" + + +@dataclass +class Parent(SerializableDataClass): + child: Child | None = None + children: list[Child] | None = None + metadata: dict | None = None + + +@dataclass +class WithFactory(SerializableDataClass): + items: list[str] = field(default_factory=list) + + class TestSerializableDataClass(unittest.TestCase): def test_from_dict_deep_with_none_values(self): """Test that from_dict_deep correctly handles None values in nested objects.""" @@ -56,6 +76,107 @@ def test_from_dict_deep_with_none_values(self): round_trip = PromptSchema.from_dict_deep(prompt.as_dict()) self.assertEqual(round_trip.as_dict(), test_dict) + def test_as_dict_exclude_unset_omits_defaults(self): + prompt_data = PromptData() + + self.assertEqual(prompt_data.as_dict(), {"prompt": None, "options": None}) + self.assertEqual(prompt_data.as_dict(exclude_unset=True), {}) + + def test_as_dict_exclude_unset_keeps_explicit_none(self): + keyword_prompt_data = PromptData(prompt=None) + positional_prompt_data = PromptData(None) + + self.assertEqual(keyword_prompt_data.as_dict(exclude_unset=True), {"prompt": None}) + self.assertEqual(positional_prompt_data.as_dict(exclude_unset=True), {"prompt": None}) + + def test_as_dict_exclude_unset_keeps_default_factory_values(self): + default_factory = WithFactory() + explicit_factory_value = WithFactory(items=[]) + + self.assertEqual(default_factory.as_dict(), {"items": []}) + self.assertEqual(default_factory.as_dict(exclude_unset=True), {"items": []}) + self.assertEqual(explicit_factory_value.as_dict(exclude_unset=True), {"items": []}) + + def test_as_dict_exclude_unset_tracks_assignments(self): + prompt_data = PromptData() + + prompt_data.prompt = None + + self.assertEqual(prompt_data.as_dict(exclude_unset=True), {"prompt": None}) + + def test_as_dict_exclude_unset_tracks_assignments_after_copy_or_pickle(self): + reconstructors = ( + ("copy", copy.copy), + ("deepcopy", copy.deepcopy), + ("pickle", lambda value: pickle.loads(pickle.dumps(value))), + ) + + for name, reconstruct in reconstructors: + with self.subTest(name=name): + original = PromptData() + prompt_data = reconstruct(original) + + self.assertFalse(hasattr(prompt_data, _INIT_FIELDS_REMAINING_ATTR)) + prompt_data.prompt = None + + self.assertEqual(original.as_dict(exclude_unset=True), {}) + self.assertEqual(prompt_data.as_dict(exclude_unset=True), {"prompt": None}) + + for name, reconstruct in reconstructors: + with self.subTest(name=f"{name}_explicit"): + original = PromptData(prompt=None) + prompt_data = reconstruct(original) + + self.assertFalse(hasattr(prompt_data, _INIT_FIELDS_REMAINING_ATTR)) + self.assertEqual(prompt_data.as_dict(exclude_unset=True), {"prompt": None}) + prompt_data.options = None + + self.assertEqual(original.as_dict(exclude_unset=True), {"prompt": None}) + self.assertEqual(prompt_data.as_dict(exclude_unset=True), {"prompt": None, "options": None}) + + def test_as_dict_exclude_unset_recurses_into_nested_values(self): + parent = Parent( + child=Child(value=None), + children=[Child(label="child")], + metadata={"nested": Child(value="set")}, + ) + + self.assertEqual( + parent.as_dict(exclude_unset=True), + { + "child": {"value": None, "label": "child"}, + "children": [{"label": "child"}], + "metadata": {"nested": {"value": "set", "label": "child"}}, + }, + ) + + def test_from_dict_deep_tracks_explicit_none_without_marking_missing_defaults(self): + test_dict = { + "id": "456", + "project_id": "123", + "_xact_id": "789", + "name": "test-prompt", + "slug": "test-prompt", + "description": None, + "prompt_data": {"prompt": None}, + "tags": None, + } + + prompt = PromptSchema.from_dict_deep(test_dict) + + self.assertEqual(prompt.as_dict(exclude_unset=True), test_dict) + + def test_as_json_supports_exclude_unset(self): + prompt_data = PromptData(prompt=None) + + self.assertEqual(prompt_data.as_json(exclude_unset=True, sort_keys=True), '{"prompt": null}') + + def test_as_dict_exclude_unset_without_tracking_falls_back_to_full_serialization(self): + prompt_data = PromptData() + object.__delattr__(prompt_data, _EXPLICITLY_SET_FIELDS_ATTR) + + self.assertEqual(prompt_data.as_dict(exclude_unset=True), prompt_data.as_dict()) + if __name__ == "__main__": unittest.main()