From 732f635bc0010b19d0b97a77e2b5ee327b3c8adc Mon Sep 17 00:00:00 2001 From: Sampoorn Nagpal Date: Sat, 5 Sep 2026 18:42:04 +0530 Subject: [PATCH 1/2] =?UTF-8?q?fix(storage):=20preserve=20allowed=20Pydant?= =?UTF-8?q?ic=20extras=20in=20snapshots=20=F0=9F=A4=96=F0=9F=A4=96?= =?UTF-8?q?=F0=9F=A4=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Sampoorn Nagpal --- src/nooa/storage/serialization.py | 6 ++++++ tests/storage/test_serialization.py | 28 ++++++++++++++++++++++++++++ 2 files changed, 34 insertions(+) diff --git a/src/nooa/storage/serialization.py b/src/nooa/storage/serialization.py index b3e612a72..e58835712 100644 --- a/src/nooa/storage/serialization.py +++ b/src/nooa/storage/serialization.py @@ -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 @@ -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) diff --git a/tests/storage/test_serialization.py b/tests/storage/test_serialization.py index d6d5dd858..a09ca7a10 100644 --- a/tests/storage/test_serialization.py +++ b/tests/storage/test_serialization.py @@ -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 @@ -170,6 +171,33 @@ def test_SKIP_sentinel_identity(self): class TestPydantic: + def test_model_config_preserves_provider_extras(self): + 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): + 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): + 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): + 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): m = MyModel(name="test", value=42) blob, al = serialize(m) From 575ec6cd64c29446a37aad2c9a1c8e48a46bb031 Mon Sep 17 00:00:00 2001 From: Sampoorn Nagpal Date: Sun, 6 Sep 2026 21:49:50 +0530 Subject: [PATCH 2/2] =?UTF-8?q?docs(storage):=20explain=20Pydantic=20extra?= =?UTF-8?q?s=20regression=20coverage=20=F0=9F=A4=96=F0=9F=A4=96?= =?UTF-8?q?=F0=9F=A4=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Sampoorn Nagpal --- tests/storage/test_serialization.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/tests/storage/test_serialization.py b/tests/storage/test_serialization.py index a09ca7a10..c8d57fbc7 100644 --- a/tests/storage/test_serialization.py +++ b/tests/storage/test_serialization.py @@ -172,6 +172,7 @@ 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) @@ -179,6 +180,7 @@ def test_model_config_preserves_provider_extras(self): 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) @@ -188,17 +190,20 @@ def test_extra_values_use_recursive_serialization(self): 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"