Merge branch 'main' into vb/update-get-state

This commit is contained in:
vbarda
2024-08-22 15:34:14 -04:00
32 changed files with 1265 additions and 540 deletions
@@ -22,6 +22,7 @@ from langgraph.checkpoint.base.id import uuid6
from langgraph.checkpoint.serde.base import SerializerProtocol, maybe_add_typed_methods
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
from langgraph.checkpoint.serde.types import (
ERROR,
ChannelProtocol,
SendProtocol,
)
@@ -98,6 +99,7 @@ class Checkpoint(TypedDict):
Cleared by the next checkpoint."""
current_tasks: Dict[str, TaskInfo]
"""Map from task ID to task info."""
# TODO remove this
def empty_checkpoint() -> Checkpoint:
@@ -140,6 +142,8 @@ def create_checkpoint(
else:
values: dict[str, Any] = {}
for k, v in channels.items():
if k not in checkpoint["channel_versions"]:
continue
try:
values[k] = v.checkpoint()
except EmptyChannelError:
@@ -437,3 +441,13 @@ def get_checkpoint_id(config: RunnableConfig) -> Optional[str]:
return config["configurable"].get(
"checkpoint_id", config["configurable"].get("thread_ts")
)
"""
Mapping from error type to error index.
Regular writes just map to their index in the list of writes being saved.
Special writes (e.g. errors) map to negative indices, to avoid those writes from
saving regular writes.
Each Checkpointer implementation should use this mapping in put_writes.
"""
WRITES_IDX_MAP = {ERROR: -1}
@@ -8,6 +8,7 @@ from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, Tuple
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
WRITES_IDX_MAP,
BaseCheckpointSaver,
ChannelVersions,
Checkpoint,
@@ -52,6 +53,9 @@ class MemorySaver(
# thread ID -> checkpoint NS -> checkpoint ID -> checkpoint mapping
storage: defaultdict[str, dict[str, dict[str, tuple[bytes, bytes, Optional[str]]]]]
writes: defaultdict[
tuple[str, str, str], dict[tuple[str, int], tuple[str, str, bytes]]
]
def __init__(
self,
@@ -60,7 +64,7 @@ class MemorySaver(
) -> None:
super().__init__(serde=serde)
self.storage = defaultdict(lambda: defaultdict(dict))
self.writes = defaultdict(list)
self.writes = defaultdict(dict)
def __enter__(self) -> "MemorySaver":
return self
@@ -103,7 +107,7 @@ class MemorySaver(
if checkpoint_id := get_checkpoint_id(config):
if saved := self.storage[thread_id][checkpoint_ns].get(checkpoint_id):
checkpoint, metadata, parent_checkpoint_id = saved
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)]
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()
return CheckpointTuple(
config=config,
checkpoint=self.serde.loads_typed(checkpoint),
@@ -125,7 +129,7 @@ class MemorySaver(
if checkpoints := self.storage[thread_id][checkpoint_ns]:
checkpoint_id = max(checkpoints.keys())
checkpoint, metadata, parent_checkpoint_id = checkpoints[checkpoint_id]
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)]
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()
return CheckpointTuple(
config={
"configurable": {
@@ -206,7 +210,8 @@ class MemorySaver(
elif limit is not None:
limit -= 1
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)]
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()
yield CheckpointTuple(
config={
"configurable": {
@@ -216,9 +221,6 @@ class MemorySaver(
}
},
checkpoint=self.serde.loads_typed(checkpoint),
pending_writes=[
(id, c, self.serde.loads_typed(v)) for id, c, v in writes
],
metadata=metadata,
parent_config={
"configurable": {
@@ -229,6 +231,9 @@ class MemorySaver(
}
if parent_checkpoint_id
else None,
pending_writes=[
(id, c, self.serde.loads_typed(v)) for id, c, v in writes
],
)
def put(
@@ -293,11 +298,10 @@ class MemorySaver(
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"]["checkpoint_ns"]
checkpoint_id = config["configurable"]["checkpoint_id"]
key = (thread_id, checkpoint_ns, checkpoint_id)
self.writes[key] = [w for w in self.writes[key] if w[0] != task_id]
self.writes[key].extend(
[(task_id, c, self.serde.dumps_typed(v)) for c, v in writes]
)
outer_key = (thread_id, checkpoint_ns, checkpoint_id)
for idx, (c, v) in enumerate(writes):
inner_key = (task_id, WRITES_IDX_MAP.get(c, idx))
self.writes[outer_key][inner_key] = (task_id, c, self.serde.dumps_typed(v))
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
"""Asynchronous version of get_tuple.
@@ -12,6 +12,8 @@ from typing import (
from langchain_core.runnables import RunnableConfig
from typing_extensions import Self
ERROR = "__error__"
Value = TypeVar("Value")
Update = TypeVar("Update")
C = TypeVar("C")