diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py index 734ccd7e6..2734cfebe 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py @@ -111,7 +111,7 @@ class PostgresSaver(BasePostgresSaver): **self._load_checkpoint(value["checkpoint"]), "channel_values": self._load_blobs(value["channel_values"]), }, - value["metadata"], + self._load_metadata(value["metadata"]), { "configurable": { "thread_id": value["thread_id"], @@ -155,7 +155,7 @@ class PostgresSaver(BasePostgresSaver): **self._load_checkpoint(value["checkpoint"]), "channel_values": self._load_blobs(value["channel_values"]), }, - value["metadata"], + self._load_metadata(value["metadata"]), { "configurable": { "thread_id": thread_id, @@ -210,7 +210,7 @@ class PostgresSaver(BasePostgresSaver): checkpoint["id"], checkpoint_id, Jsonb(self._dump_checkpoint(copy)), - Jsonb(metadata), + self._dump_metadata(metadata), ), ) return next_config diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py index 686f0b305..5a56eec97 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py @@ -111,7 +111,7 @@ class AsyncPostgresSaver(BasePostgresSaver): self._load_blobs, value["channel_values"] ), }, - value["metadata"], + self._load_metadata(value["metadata"]), { "configurable": { "thread_id": value["thread_id"], @@ -157,7 +157,7 @@ class AsyncPostgresSaver(BasePostgresSaver): self._load_blobs, value["channel_values"] ), }, - value["metadata"], + self._load_metadata(value["metadata"]), { "configurable": { "thread_id": thread_id, @@ -214,7 +214,7 @@ class AsyncPostgresSaver(BasePostgresSaver): checkpoint["id"], checkpoint_id, Jsonb(self._dump_checkpoint(copy)), - Jsonb(metadata), + self._dump_metadata(metadata), ), ) return next_config diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py index 5a88f61ac..b4c56e781 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py @@ -11,6 +11,7 @@ from langgraph.checkpoint.base import ( EmptyChannelError, get_checkpoint_id, ) +from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer from langgraph.checkpoint.serde.types import ChannelProtocol MetadataInput = Optional[dict[str, Any]] @@ -109,6 +110,7 @@ class BasePostgresSaver(BaseCheckpointSaver): UPSERT_CHECKPOINT_BLOBS_SQL = UPSERT_CHECKPOINT_BLOBS_SQL UPSERT_CHECKPOINTS_SQL = UPSERT_CHECKPOINTS_SQL UPSERT_CHECKPOINT_WRITES_SQL = UPSERT_CHECKPOINT_WRITES_SQL + jsonplus_serde = JsonPlusSerializer() def _load_checkpoint(self, checkpoint: dict[str, Any]) -> Checkpoint: if len(checkpoint["pending_sends"]) == 2 and all( @@ -202,6 +204,19 @@ class BasePostgresSaver(BaseCheckpointSaver): for idx, (channel, value) in enumerate(writes) ] + def _load_metadata(self, metadata: dict[str, Any]) -> dict[str, Any]: + return self.jsonplus_serde.loads(self.jsonplus_serde.dumps(metadata)) + + def _dump_metadata(self, metadata) -> str: + serialized_metadata_type, serialized_metadata = self.jsonplus_serde.dumps_typed( + metadata + ) + if serialized_metadata_type != "json": + raise TypeError( + f"Failed to properly serialize metadata -- expected 'json', got '{serialized_metadata_type}'" + ) + return serialized_metadata.decode() + def get_next_version(self, current: Optional[str], channel: ChannelProtocol) -> str: if current is None: current_v = 0