mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-13 13:17:52 +02:00
@@ -151,6 +151,7 @@
|
||||
"from langchain_core.runnables import RunnableConfig\n",
|
||||
"\n",
|
||||
"from langgraph.checkpoint.base import (\n",
|
||||
" WRITES_IDX_MAP,\n",
|
||||
" BaseCheckpointSaver,\n",
|
||||
" ChannelVersions,\n",
|
||||
" Checkpoint,\n",
|
||||
@@ -163,7 +164,7 @@
|
||||
"from redis import Redis\n",
|
||||
"from redis.asyncio import Redis as AsyncRedis\n",
|
||||
"\n",
|
||||
"REDIS_KEY_SEPARATOR = \":\"\n",
|
||||
"REDIS_KEY_SEPARATOR = \"$\"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# Utilities shared by both RedisSaver and AsyncRedisSaver\n",
|
||||
@@ -246,17 +247,6 @@
|
||||
" return keys\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def _dump_writes(serde: SerializerProtocol, writes: tuple[str, Any]) -> list[dict]:\n",
|
||||
" \"\"\"Serialize pending writes.\"\"\"\n",
|
||||
" serialized_writes = []\n",
|
||||
" for channel, value in writes:\n",
|
||||
" type_, serialized_value = serde.dumps_typed(value)\n",
|
||||
" serialized_writes.append(\n",
|
||||
" {\"channel\": channel, \"type\": type_, \"value\": serialized_value}\n",
|
||||
" )\n",
|
||||
" return serialized_writes\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def _load_writes(\n",
|
||||
" serde: SerializerProtocol, task_id_to_data: dict[tuple[str, str], dict]\n",
|
||||
") -> list[PendingWrite]:\n",
|
||||
@@ -413,7 +403,7 @@
|
||||
" config: RunnableConfig,\n",
|
||||
" writes: List[Tuple[str, Any]],\n",
|
||||
" task_id: str,\n",
|
||||
" ) -> RunnableConfig:\n",
|
||||
" ) -> None:\n",
|
||||
" \"\"\"Store intermediate writes linked to a checkpoint.\n",
|
||||
"\n",
|
||||
" Args:\n",
|
||||
@@ -425,12 +415,23 @@
|
||||
" checkpoint_ns = config[\"configurable\"][\"checkpoint_ns\"]\n",
|
||||
" checkpoint_id = config[\"configurable\"][\"checkpoint_id\"]\n",
|
||||
"\n",
|
||||
" for idx, data in enumerate(_dump_writes(self.serde, writes)):\n",
|
||||
" for idx, (channel, value) in enumerate(writes):\n",
|
||||
" key = _make_redis_checkpoint_writes_key(\n",
|
||||
" thread_id, checkpoint_ns, checkpoint_id, task_id, idx\n",
|
||||
" thread_id,\n",
|
||||
" checkpoint_ns,\n",
|
||||
" checkpoint_id,\n",
|
||||
" task_id,\n",
|
||||
" WRITES_IDX_MAP.get(channel, idx),\n",
|
||||
" )\n",
|
||||
" self.conn.hset(key, mapping=data)\n",
|
||||
" return config\n",
|
||||
" type_, serialized_value = self.serde.dumps_typed(value)\n",
|
||||
" data = {\"channel\": channel, \"type\": type_, \"value\": serialized_value}\n",
|
||||
" if all(w[0] in WRITES_IDX_MAP for w in writes):\n",
|
||||
" # Use HSET which will overwrite existing values\n",
|
||||
" self.conn.hset(key, mapping=data)\n",
|
||||
" else:\n",
|
||||
" # Use HSETNX which will not overwrite existing values\n",
|
||||
" for field, value in data.items():\n",
|
||||
" self.conn.hsetnx(key, field, value)\n",
|
||||
"\n",
|
||||
" def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:\n",
|
||||
" \"\"\"Get a checkpoint tuple from Redis.\n",
|
||||
@@ -463,21 +464,8 @@
|
||||
" checkpoint_id\n",
|
||||
" or _parse_redis_checkpoint_key(checkpoint_key)[\"checkpoint_id\"]\n",
|
||||
" )\n",
|
||||
" writes_key = _make_redis_checkpoint_writes_key(\n",
|
||||
" thread_id, checkpoint_ns, checkpoint_id, \"*\", None\n",
|
||||
" )\n",
|
||||
" matching_keys = self.conn.keys(pattern=writes_key)\n",
|
||||
" parsed_keys = [\n",
|
||||
" _parse_redis_checkpoint_writes_key(key.decode()) for key in matching_keys\n",
|
||||
" ]\n",
|
||||
" pending_writes = _load_writes(\n",
|
||||
" self.serde,\n",
|
||||
" {\n",
|
||||
" (parsed_key[\"task_id\"], parsed_key[\"idx\"]): self.conn.hgetall(key)\n",
|
||||
" for key, parsed_key in sorted(\n",
|
||||
" zip(matching_keys, parsed_keys), key=lambda x: x[1][\"idx\"]\n",
|
||||
" )\n",
|
||||
" },\n",
|
||||
" pending_writes = self._load_pending_writes(\n",
|
||||
" thread_id, checkpoint_ns, checkpoint_id\n",
|
||||
" )\n",
|
||||
" return _parse_redis_checkpoint_data(\n",
|
||||
" self.serde, checkpoint_key, checkpoint_data, pending_writes=pending_writes\n",
|
||||
@@ -514,7 +502,37 @@
|
||||
" for key in keys:\n",
|
||||
" data = self.conn.hgetall(key)\n",
|
||||
" if data and b\"checkpoint\" in data and b\"metadata\" in data:\n",
|
||||
" yield _parse_redis_checkpoint_data(self.serde, key.decode(), data)\n",
|
||||
" # load pending writes\n",
|
||||
" checkpoint_id = _parse_redis_checkpoint_key(key.decode())[\n",
|
||||
" \"checkpoint_id\"\n",
|
||||
" ]\n",
|
||||
" pending_writes = self._load_pending_writes(\n",
|
||||
" thread_id, checkpoint_ns, checkpoint_id\n",
|
||||
" )\n",
|
||||
" yield _parse_redis_checkpoint_data(\n",
|
||||
" self.serde, key.decode(), data, pending_writes=pending_writes\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" def _load_pending_writes(\n",
|
||||
" self, thread_id: str, checkpoint_ns: str, checkpoint_id: str\n",
|
||||
" ) -> List[PendingWrite]:\n",
|
||||
" writes_key = _make_redis_checkpoint_writes_key(\n",
|
||||
" thread_id, checkpoint_ns, checkpoint_id, \"*\", None\n",
|
||||
" )\n",
|
||||
" matching_keys = self.conn.keys(pattern=writes_key)\n",
|
||||
" parsed_keys = [\n",
|
||||
" _parse_redis_checkpoint_writes_key(key.decode()) for key in matching_keys\n",
|
||||
" ]\n",
|
||||
" pending_writes = _load_writes(\n",
|
||||
" self.serde,\n",
|
||||
" {\n",
|
||||
" (parsed_key[\"task_id\"], parsed_key[\"idx\"]): self.conn.hgetall(key)\n",
|
||||
" for key, parsed_key in sorted(\n",
|
||||
" zip(matching_keys, parsed_keys), key=lambda x: x[1][\"idx\"]\n",
|
||||
" )\n",
|
||||
" },\n",
|
||||
" )\n",
|
||||
" return pending_writes\n",
|
||||
"\n",
|
||||
" def _get_checkpoint_key(\n",
|
||||
" self, conn, thread_id: str, checkpoint_ns: str, checkpoint_id: Optional[str]\n",
|
||||
@@ -637,7 +655,7 @@
|
||||
" config: RunnableConfig,\n",
|
||||
" writes: List[Tuple[str, Any]],\n",
|
||||
" task_id: str,\n",
|
||||
" ) -> RunnableConfig:\n",
|
||||
" ) -> None:\n",
|
||||
" \"\"\"Store intermediate writes linked to a checkpoint asynchronously.\n",
|
||||
"\n",
|
||||
" This method saves intermediate writes associated with a checkpoint to the database.\n",
|
||||
@@ -651,12 +669,23 @@
|
||||
" checkpoint_ns = config[\"configurable\"][\"checkpoint_ns\"]\n",
|
||||
" checkpoint_id = config[\"configurable\"][\"checkpoint_id\"]\n",
|
||||
"\n",
|
||||
" for idx, data in enumerate(_dump_writes(self.serde, writes)):\n",
|
||||
" for idx, (channel, value) in enumerate(writes):\n",
|
||||
" key = _make_redis_checkpoint_writes_key(\n",
|
||||
" thread_id, checkpoint_ns, checkpoint_id, task_id, idx\n",
|
||||
" thread_id,\n",
|
||||
" checkpoint_ns,\n",
|
||||
" checkpoint_id,\n",
|
||||
" task_id,\n",
|
||||
" WRITES_IDX_MAP.get(channel, idx),\n",
|
||||
" )\n",
|
||||
" await self.conn.hset(key, mapping=data)\n",
|
||||
" return config\n",
|
||||
" type_, serialized_value = self.serde.dumps_typed(value)\n",
|
||||
" data = {\"channel\": channel, \"type\": type_, \"value\": serialized_value}\n",
|
||||
" if all(w[0] in WRITES_IDX_MAP for w in writes):\n",
|
||||
" # Use HSET which will overwrite existing values\n",
|
||||
" await self.conn.hset(key, mapping=data)\n",
|
||||
" else:\n",
|
||||
" # Use HSETNX which will not overwrite existing values\n",
|
||||
" for field, value in data.items():\n",
|
||||
" await self.conn.hsetnx(key, field, value)\n",
|
||||
"\n",
|
||||
" async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:\n",
|
||||
" \"\"\"Get a checkpoint tuple from Redis asynchronously.\n",
|
||||
@@ -688,21 +717,8 @@
|
||||
" checkpoint_id\n",
|
||||
" or _parse_redis_checkpoint_key(checkpoint_key)[\"checkpoint_id\"]\n",
|
||||
" )\n",
|
||||
" writes_key = _make_redis_checkpoint_writes_key(\n",
|
||||
" thread_id, checkpoint_ns, checkpoint_id, \"*\", None\n",
|
||||
" )\n",
|
||||
" matching_keys = await self.conn.keys(pattern=writes_key)\n",
|
||||
" parsed_keys = [\n",
|
||||
" _parse_redis_checkpoint_writes_key(key.decode()) for key in matching_keys\n",
|
||||
" ]\n",
|
||||
" pending_writes = _load_writes(\n",
|
||||
" self.serde,\n",
|
||||
" {\n",
|
||||
" (parsed_key[\"task_id\"], parsed_key[\"idx\"]): await self.conn.hgetall(key)\n",
|
||||
" for key, parsed_key in sorted(\n",
|
||||
" zip(matching_keys, parsed_keys), key=lambda x: x[1][\"idx\"]\n",
|
||||
" )\n",
|
||||
" },\n",
|
||||
" pending_writes = await self._aload_pending_writes(\n",
|
||||
" thread_id, checkpoint_ns, checkpoint_id\n",
|
||||
" )\n",
|
||||
" return _parse_redis_checkpoint_data(\n",
|
||||
" self.serde, checkpoint_key, checkpoint_data, pending_writes=pending_writes\n",
|
||||
@@ -738,7 +754,36 @@
|
||||
" for key in keys:\n",
|
||||
" data = await self.conn.hgetall(key)\n",
|
||||
" if data and b\"checkpoint\" in data and b\"metadata\" in data:\n",
|
||||
" yield _parse_redis_checkpoint_data(self.serde, key.decode(), data)\n",
|
||||
" checkpoint_id = _parse_redis_checkpoint_key(key.decode())[\n",
|
||||
" \"checkpoint_id\"\n",
|
||||
" ]\n",
|
||||
" pending_writes = await self._aload_pending_writes(\n",
|
||||
" thread_id, checkpoint_ns, checkpoint_id\n",
|
||||
" )\n",
|
||||
" yield _parse_redis_checkpoint_data(\n",
|
||||
" self.serde, key.decode(), data, pending_writes=pending_writes\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" async def _aload_pending_writes(\n",
|
||||
" self, thread_id: str, checkpoint_ns: str, checkpoint_id: str\n",
|
||||
" ) -> List[PendingWrite]:\n",
|
||||
" writes_key = _make_redis_checkpoint_writes_key(\n",
|
||||
" thread_id, checkpoint_ns, checkpoint_id, \"*\", None\n",
|
||||
" )\n",
|
||||
" matching_keys = await self.conn.keys(pattern=writes_key)\n",
|
||||
" parsed_keys = [\n",
|
||||
" _parse_redis_checkpoint_writes_key(key.decode()) for key in matching_keys\n",
|
||||
" ]\n",
|
||||
" pending_writes = _load_writes(\n",
|
||||
" self.serde,\n",
|
||||
" {\n",
|
||||
" (parsed_key[\"task_id\"], parsed_key[\"idx\"]): await self.conn.hgetall(key)\n",
|
||||
" for key, parsed_key in sorted(\n",
|
||||
" zip(matching_keys, parsed_keys), key=lambda x: x[1][\"idx\"]\n",
|
||||
" )\n",
|
||||
" },\n",
|
||||
" )\n",
|
||||
" return pending_writes\n",
|
||||
"\n",
|
||||
" async def _aget_checkpoint_key(\n",
|
||||
" self, conn, thread_id: str, checkpoint_ns: str, checkpoint_id: Optional[str]\n",
|
||||
@@ -1042,7 +1087,7 @@
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.11.4"
|
||||
"version": "3.12.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
|
||||
Reference in New Issue
Block a user