diff --git a/docs/docs/how-tos/persistence_redis.ipynb b/docs/docs/how-tos/persistence_redis.ipynb index 869751ca1..8fab22b3c 100644 --- a/docs/docs/how-tos/persistence_redis.ipynb +++ b/docs/docs/how-tos/persistence_redis.ipynb @@ -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,