docs: update redis how-to (#2727)

Fixes #2712
This commit is contained in:
Vadym Barda
2024-12-11 23:22:47 +00:00
committed by GitHub
parent 67f96063e2
commit 0400c5236e
+100 -55
View File
@@ -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,