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
6 changes: 6 additions & 0 deletions src/nooa/storage/serialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -154,6 +154,10 @@ def _serialize(value: Any, allowlist: set[str]) -> Any:
serialized = _serialize(getattr(value, field_name), allowlist)
if serialized is not SKIP:
data[field_name] = serialized
for field_name, extra in (value.model_extra or {}).items():
serialized = _serialize(extra, allowlist)
if serialized is not SKIP:
data[field_name] = serialized
return {"__type__": _PYDANTIC, "__class__": fqn, "data": data}

# 6. Dataclasses
Expand Down Expand Up @@ -235,6 +239,8 @@ def _deserialize_envelope(blob: dict[str, Any], allowlist: set[str]) -> Any:
# Deserialize nested envelopes first (e.g. dataclass/snapshotable
# fields), then let Pydantic validate the reconstructed objects.
deserialized_data = {k: _deserialize(v, allowlist) for k, v in data.items()}
if cls.model_config.get("extra") == "allow":
return cls.model_validate(deserialized_data)
# Filter to known fields for lenient restoration (handles extra="forbid"
# models that would reject extra fields from schema drift).
known_fields = set(cls.model_fields)
Expand Down
33 changes: 33 additions & 0 deletions tests/storage/test_serialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
import pytest
from pydantic import BaseModel, ConfigDict

from nooa.config.model_config import ModelConfig
from nooa.errors.storage import DeserializationError, SerializationError
from nooa.storage.markers import snapshotable
from nooa.storage.serialization import SKIP, deserialize, serialize
Expand Down Expand Up @@ -170,7 +171,39 @@ def test_SKIP_sentinel_identity(self):


class TestPydantic:
def test_model_config_preserves_provider_extras(self):
"""Snapshot round trips retain provider options stored as allowed model extras."""
model = ModelConfig(model_name="test", num_retries=7, custom={"stops": ("a", "b")})
blob, allowlist = serialize(model)
restored = deserialize(blob, allowlist)
assert restored == model
assert restored.model_extra == model.model_extra

def test_extra_values_use_recursive_serialization(self):
"""Nested dataclasses and snapshotable values in extras retain their types."""
model = ModelConfig(custom={"point": Point(1, 2), "config": Config("localhost")})
blob, allowlist = serialize(model)
restored = deserialize(blob, allowlist)
assert restored.model_extra["custom"]["point"] == Point(1, 2)
config = restored.model_extra["custom"]["config"]
assert isinstance(config, Config)
assert config.host == "localhost"

def test_nested_extra_class_requires_allowlist(self):
"""Allowing extra fields does not bypass the nested-class restoration allowlist."""
model = ModelConfig(custom=Point(1, 2))
blob, allowlist = serialize(model)
with pytest.raises(DeserializationError, match="not in the allowlist"):
deserialize(blob, allowlist - {f"{Point.__module__}.Point"})

def test_nosnapshot_extra_is_skipped(self):
"""Opted-out extra values are skipped without losing serializable sibling options."""
model = ModelConfig(custom=_NoSnapshotThing(), num_retries=7)
blob, allowlist = serialize(model)
assert deserialize(blob, allowlist).model_extra == {"num_retries": 7}

def test_simple_model_roundtrip(self):
"""Ordinary declared Pydantic fields still survive a snapshot round trip."""
m = MyModel(name="test", value=42)
blob, al = serialize(m)
assert blob["__type__"] == "pydantic"
Expand Down