Skip to content
Draft
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
11 changes: 10 additions & 1 deletion python/packages/core/agent_framework/_sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -271,7 +271,16 @@ def decode_from_dict(payload: Mapping[str, Any]) -> Any:
if issubclass(cls, BaseModel):

def decode_pydantic(payload: Mapping[str, Any]) -> Any:
return cls.model_validate({key: value for key, value in payload.items() if key != "type"})
# The encoder's payload items went through _serialize_value, so a
# nested registered value (e.g. Message) sits in the payload as its
# tagged dict. Restore nested values the same way before validating,
# otherwise Pydantic receives the tagged dict where the field
# expects the reconstructed instance.
return cls.model_validate({
key: _deserialize_value(value, path=f"{cls.__name__}.{key}")
for key, value in payload.items()
if key != "type"
})

return decode_pydantic

Expand Down
42 changes: 42 additions & 0 deletions python/packages/core/tests/core/test_sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -901,6 +901,26 @@ class PydanticState(BaseModel):
assert isinstance(restored.state["value"], PydanticState)
assert restored.state["value"].value == "ok"

def test_registered_pydantic_state_restores_nested_registered_values(self) -> None:
from pydantic import BaseModel, ConfigDict

class PydanticMessageState(BaseModel):
model_config = ConfigDict(arbitrary_types_allowed=True)

message: Message

register_state_type(PydanticMessageState, type_id="pydantic_message_state_test")
session = AgentSession(session_id="nested-message")
session.state["value"] = PydanticMessageState(
message=Message(role="user", contents=["persisted"])
)

restored = AgentSession.from_dict(session.to_dict())

assert isinstance(restored.state["value"], PydanticMessageState)
assert isinstance(restored.state["value"].message, Message)
assert restored.state["value"].message.text == "persisted"

def test_conflicting_type_identifier_is_rejected(self) -> None:
class FirstState:
def to_dict(self) -> dict[str, Any]:
Expand Down Expand Up @@ -1146,6 +1166,28 @@ async def test_round_trips_session_across_store_instances(self, tmp_path: Path)
assert files[0].suffix == ".json"
assert json.loads(files[0].read_bytes())["version"] == "1.0"

async def test_round_trips_registered_pydantic_state_with_nested_message(self, tmp_path: Path) -> None:
from pydantic import BaseModel, ConfigDict

class PydanticFileMessageState(BaseModel):
model_config = ConfigDict(arbitrary_types_allowed=True)

message: Message

register_state_type(PydanticFileMessageState, type_id="pydantic_file_message_state_test")
session = AgentSession(session_id="file-nested-message")
session.state["value"] = PydanticFileMessageState(
message=Message(role="user", contents=["persisted"])
)

await FileSessionStore(tmp_path).set("tenant_file-nested-message", session)
restored = await FileSessionStore(tmp_path).get("tenant_file-nested-message")

assert restored is not None
assert isinstance(restored.state["value"], PydanticFileMessageState)
assert isinstance(restored.state["value"].message, Message)
assert restored.state["value"].message.text == "persisted"

async def test_round_trips_binary_messagepack_session(self, tmp_path: Path) -> None:
store = FileSessionStore(tmp_path, serialization_format="msgpack")
session = AgentSession(session_id="binary-session")
Expand Down
Loading