checkpoint-postgres: use correct metadata serialization (#1255)

This commit is contained in:
Vadym Barda
2024-08-07 12:13:44 -04:00
committed by GitHub
parent fb8390e138
commit ca4b65e1a6
3 changed files with 21 additions and 6 deletions
@@ -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
@@ -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
@@ -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