From 6df84d7421b54051b350a5f5cee33bd42d5d6940 Mon Sep 17 00:00:00 2001 From: Sydney Runkle Date: Tue, 21 Apr 2026 10:05:08 -0400 Subject: [PATCH] feat(memory): implement get_channel_blob; remove diff handling from _load_blobs --- .../langgraph/checkpoint/memory/__init__.py | 78 ++++++++----------- libs/checkpoint/tests/test_memory.py | 62 +++++---------- 2 files changed, 52 insertions(+), 88 deletions(-) diff --git a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py index 4cc98f9db..7ca2ccf0c 100644 --- a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py @@ -123,59 +123,49 @@ class InMemorySaver( def _load_blobs( self, thread_id: str, checkpoint_ns: str, versions: ChannelVersions ) -> dict[str, Any]: - from langgraph.checkpoint.base import DeltaChainValue - channel_values: dict[str, Any] = {} - diff_channels: dict[str, str] = {} - for k, v in versions.items(): kk = (thread_id, checkpoint_ns, k, v) if kk not in self.blobs: continue vv = self.blobs[kk] - if vv[0] == "diff": - diff_channels[k] = str(v) - elif vv[0] != "empty": + if vv[0] != "empty": channel_values[k] = self.serde.loads_typed(vv) - - for k, current_version in diff_channels.items(): - chain_deltas: list[list[Any]] = [] - base: list[Any] | None = None - version: str | None = current_version - visited: set[str] = set() - while version is not None: - if version in visited: - logger.warning( - "DeltaChannel chain cycle detected at version %r for channel %r; breaking", - version, - k, - ) - break - visited.add(version) - kk = (thread_id, checkpoint_ns, k, version) - if kk not in self.blobs: - logger.warning( - "DeltaChannel chain is broken: blob for channel %r version %r not found; " - "partial history will be returned", - k, - version, - ) - break - vv = self.blobs[kk] - if vv[0] == "diff": - payload = self.serde.loads_typed( - vv - ) # {"d": [...], "p": version|None} - chain_deltas.append(payload["d"]) - version = payload["p"] - else: - base = self.serde.loads_typed(vv) - break - chain_deltas.reverse() - channel_values[k] = DeltaChainValue(base=base, deltas=chain_deltas) - return channel_values + def get_channel_blob( + self, + thread_id: str, + checkpoint_ns: str, + checkpoint_id: str, + channel: str, + ) -> Any: + """Fast-path blob lookup: checkpoint → channel version → blob.""" + ns_storage = self.storage.get((thread_id, checkpoint_ns), {}) + entry = ns_storage.get(checkpoint_id) + if entry is None: + return NotImplemented + checkpoint = entry[1] + version = checkpoint["channel_versions"].get(channel) + if version is None: + return NotImplemented + kk = (thread_id, checkpoint_ns, channel, version) + if kk not in self.blobs: + return NotImplemented + vv = self.blobs[kk] + if vv[0] == "empty": + return NotImplemented + return self.serde.loads_typed(vv) + + async def aget_channel_blob( + self, + thread_id: str, + checkpoint_ns: str, + checkpoint_id: str, + channel: str, + ) -> Any: + return self.get_channel_blob(thread_id, checkpoint_ns, checkpoint_id, channel) + def get_tuple(self, config: RunnableConfig) -> CheckpointTuple | None: """Get a checkpoint tuple from the in-memory storage. diff --git a/libs/checkpoint/tests/test_memory.py b/libs/checkpoint/tests/test_memory.py index 0c2204db1..61aa7f965 100644 --- a/libs/checkpoint/tests/test_memory.py +++ b/libs/checkpoint/tests/test_memory.py @@ -323,55 +323,29 @@ def test_memory_saver_with_allowlist_proxy_isolated() -> None: class TestInMemorySaverDeltaChannel: - def test_delta_channel_chain_reconstruction(self) -> None: - """_load_blobs follows the diff chain and returns DeltaChainValue.""" - from langgraph.checkpoint.base import DeltaChainValue, DeltaValue + def test_get_channel_blob(self) -> None: + """get_channel_blob returns the deserialized blob for a checkpoint+channel.""" + from langgraph.checkpoint.base import DeltaValue, empty_checkpoint saver = InMemorySaver() serde = JsonPlusSerializer() - thread_id = "t1" - ns = "" + thread_id, ns, channel = "t1", "", "messages" + version = "00000000000000000000000000000001.0000000000000000" + delta = DeltaValue(delta=[{"content": "hi"}], prev_checkpoint_id=None) + saver.blobs[(thread_id, ns, channel, version)] = serde.dumps_typed(delta) - # Simulate two steps: v1 (root) and v2 (chained to v1) - v1 = "00000000000000000000000000000001.1234567890000000" - v2 = "00000000000000000000000000000002.1234567890000000" + cp = empty_checkpoint() + cp["id"] = "cp1" + cp["channel_versions"][channel] = version + saver.storage[(thread_id, ns)] = {"cp1": ({}, cp, {})} - delta1 = DeltaValue(delta=["msg1"], prev_version=None) - delta2 = DeltaValue(delta=["msg2"], prev_version=v1) - - saver.blobs[(thread_id, ns, "messages", v1)] = serde.dumps_typed(delta1) - saver.blobs[(thread_id, ns, "messages", v2)] = serde.dumps_typed(delta2) - - channel_values = saver._load_blobs(thread_id, ns, {"messages": v2}) - - assert "messages" in channel_values - result = channel_values["messages"] - assert isinstance(result, DeltaChainValue) - assert result.base is None - assert result.deltas == [["msg1"], ["msg2"]] - - def test_delta_channel_mixed_old_and_new_blobs(self) -> None: - """When chain hits an old non-diff blob, it becomes base.""" - from langgraph.checkpoint.base import DeltaChainValue, DeltaValue + result = saver.get_channel_blob(thread_id, ns, "cp1", channel) + assert isinstance(result, DeltaValue) + assert result.delta == [{"content": "hi"}] + assert result.prev_checkpoint_id is None + def test_get_channel_blob_missing(self) -> None: + """get_channel_blob returns NotImplemented when checkpoint or channel not found.""" saver = InMemorySaver() - serde = JsonPlusSerializer() - - thread_id = "t2" - ns = "" - - v_old = "00000000000000000000000000000001.0000000000000000" - v_new = "00000000000000000000000000000002.0000000000000000" - - # Old-style full-list blob - saver.blobs[(thread_id, ns, "messages", v_old)] = serde.dumps_typed(["old_msg"]) - # New diff blob chained to old - delta = DeltaValue(delta=["new_msg"], prev_version=v_old) - saver.blobs[(thread_id, ns, "messages", v_new)] = serde.dumps_typed(delta) - - channel_values = saver._load_blobs(thread_id, ns, {"messages": v_new}) - result = channel_values["messages"] - assert isinstance(result, DeltaChainValue) - assert result.base == ["old_msg"] - assert result.deltas == [["new_msg"]] + assert saver.get_channel_blob("t1", "", "no-such-cp", "messages") is NotImplemented