feat(memory): implement get_channel_blob; remove diff handling from _load_blobs

This commit is contained in:
Sydney Runkle
2026-04-30 14:44:39 -04:00
parent fe783a53f0
commit 6df84d7421
2 changed files with 52 additions and 88 deletions
@@ -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.
+18 -44
View File
@@ -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