mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-05 09:17:47 +02:00
feat(checkpoint/postgres): diff chain reconstruction in async saver
Add `_load_diff_chains_async` to `AsyncPostgresSaver` and override `_load_checkpoint_tuple` to inline blob-parsing and diff-chain resolution via async point-lookup traversal, mirroring the sync `PostgresSaver._load_diff_chains` implementation. Add integration test `test_diff_channel_chain_reconstruction` that skips gracefully when `langgraph` core is not installed. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Sonnet 4.6
parent
f5b27e8cda
commit
afc2e12201
@@ -391,6 +391,45 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
async with conn.cursor(binary=True, row_factory=dict_row) as cur:
|
||||
yield cur
|
||||
|
||||
async def _load_diff_chains_async(
|
||||
self,
|
||||
thread_id: str,
|
||||
checkpoint_ns: str,
|
||||
diff_channel_payloads: dict[str, dict[str, Any]],
|
||||
) -> dict[str, Any]:
|
||||
from langgraph.checkpoint.base import DiffChainValue
|
||||
|
||||
result: dict[str, Any] = {}
|
||||
async with self._cursor() as cur:
|
||||
for channel, current_payload in diff_channel_payloads.items():
|
||||
payloads: list[dict[str, Any]] = [current_payload]
|
||||
version_cursor: str | None = current_payload["p"]
|
||||
base: list[Any] | None = None
|
||||
|
||||
while version_cursor is not None:
|
||||
await cur.execute(
|
||||
"SELECT type, blob FROM checkpoint_blobs "
|
||||
"WHERE thread_id = %s AND checkpoint_ns = %s "
|
||||
"AND channel = %s AND version = %s",
|
||||
(thread_id, checkpoint_ns, channel, version_cursor),
|
||||
)
|
||||
row = await cur.fetchone()
|
||||
if row is None:
|
||||
break
|
||||
if row["type"] == "diff":
|
||||
payload = self.serde.loads_typed(("diff", row["blob"]))
|
||||
payloads.append(payload)
|
||||
version_cursor = payload["p"]
|
||||
else:
|
||||
base = self.serde.loads_typed((row["type"], row["blob"]))
|
||||
break
|
||||
|
||||
payloads.reverse()
|
||||
result[channel] = DiffChainValue(
|
||||
base=base, deltas=[p["d"] for p in payloads]
|
||||
)
|
||||
return result
|
||||
|
||||
async def _load_checkpoint_tuple(self, value: DictRow) -> CheckpointTuple:
|
||||
"""
|
||||
Convert a database row into a CheckpointTuple object.
|
||||
@@ -403,11 +442,32 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
including its configuration, metadata, parent checkpoint (if any),
|
||||
and pending writes.
|
||||
"""
|
||||
thread_id = value["thread_id"]
|
||||
checkpoint_ns = value["checkpoint_ns"]
|
||||
blob_values = value["channel_values"]
|
||||
|
||||
non_diff: dict[str, Any] = {}
|
||||
diff_payloads: dict[str, dict[str, Any]] = {}
|
||||
if blob_values:
|
||||
for k, t, v in blob_values:
|
||||
channel = k.decode()
|
||||
type_tag = t.decode()
|
||||
if type_tag == "diff":
|
||||
diff_payloads[channel] = self.serde.loads_typed((type_tag, v))
|
||||
elif type_tag != "empty":
|
||||
non_diff[channel] = self.serde.loads_typed((type_tag, v))
|
||||
|
||||
diff_values = (
|
||||
await self._load_diff_chains_async(thread_id, checkpoint_ns, diff_payloads)
|
||||
if diff_payloads
|
||||
else {}
|
||||
)
|
||||
|
||||
return CheckpointTuple(
|
||||
{
|
||||
"configurable": {
|
||||
"thread_id": value["thread_id"],
|
||||
"checkpoint_ns": value["checkpoint_ns"],
|
||||
"thread_id": thread_id,
|
||||
"checkpoint_ns": checkpoint_ns,
|
||||
"checkpoint_id": value["checkpoint_id"],
|
||||
}
|
||||
},
|
||||
@@ -415,15 +475,16 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
**value["checkpoint"],
|
||||
"channel_values": {
|
||||
**(value["checkpoint"].get("channel_values") or {}),
|
||||
**self._load_blobs(value["channel_values"]),
|
||||
**non_diff,
|
||||
**diff_values,
|
||||
},
|
||||
},
|
||||
value["metadata"],
|
||||
(
|
||||
{
|
||||
"configurable": {
|
||||
"thread_id": value["thread_id"],
|
||||
"checkpoint_ns": value["checkpoint_ns"],
|
||||
"thread_id": thread_id,
|
||||
"checkpoint_ns": checkpoint_ns,
|
||||
"checkpoint_id": value["parent_checkpoint_id"],
|
||||
}
|
||||
}
|
||||
|
||||
@@ -371,3 +371,47 @@ async def test_get_checkpoint_no_channel_values(
|
||||
|
||||
checkpoint = await saver.aget_tuple(config)
|
||||
assert checkpoint.checkpoint["channel_values"] == {}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"])
|
||||
async def test_diff_channel_chain_reconstruction(saver_name: str) -> None:
|
||||
"""AsyncPostgresSaver reconstructs DiffChannel chain via point-lookup traversal."""
|
||||
pytest.importorskip(
|
||||
"langgraph.channels.diff", reason="langgraph core not installed"
|
||||
)
|
||||
|
||||
from typing import Annotated
|
||||
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
from langgraph.channels.diff import DiffChannel
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.graph.message import add_messages
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
class State(TypedDict):
|
||||
messages: Annotated[list, DiffChannel(add_messages)]
|
||||
|
||||
def respond(state: State) -> dict:
|
||||
n = len(state["messages"])
|
||||
return {"messages": [AIMessage(content=f"reply-{n}", id=f"ai-{n}")]}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("respond", respond)
|
||||
builder.add_edge(START, "respond")
|
||||
|
||||
async with _saver(saver_name) as saver:
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
config = {"configurable": {"thread_id": "diff-channel-test-1"}}
|
||||
|
||||
await graph.ainvoke({"messages": [HumanMessage(content="hi", id="h1")]}, config)
|
||||
await graph.ainvoke(
|
||||
{"messages": [HumanMessage(content="there", id="h2")]}, config
|
||||
)
|
||||
|
||||
state = await graph.aget_state(config)
|
||||
msgs = state.values["messages"]
|
||||
assert len(msgs) == 4, f"expected 4, got {len(msgs)}: {msgs}"
|
||||
assert msgs[0].content == "hi"
|
||||
assert msgs[1].content == "reply-1"
|
||||
assert msgs[2].content == "there"
|
||||
assert msgs[3].content == "reply-3"
|
||||
|
||||
Reference in New Issue
Block a user