mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-10 19:57:52 +02:00
feat(memory): implement get_channel_blob; remove diff handling from _load_blobs
This commit is contained in:
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user