一、核心应用场景

1. 会话状态持久化(最核心用途)

  • 保存多轮对话历史
  • 维护 Agent 执行状态
  • 支持断点续传和故障恢复

2. 短期记忆管理

  • 存储最近 N 条对话记录
  • 滑动窗口记忆机制
  • 跨会话共享上下文

3. 缓存层

  • LLM 响应缓存(降低 API 成本)
  • Embedding 结果缓存
  • 工具调用结果缓存

4. 分布式锁与并发控制

  • 防止同一 Agent 实例重复执行
  • 限制并发工具调用
  • 协调多节点部署

5. RAG 向量存储

  • 文档向量存储与检索
  • 语义缓存
  • 混合搜索

6. 工作流状态管理

  • LangGraph 节点状态存储
  • 任务队列管理
  • 失败重试跟踪

二、详细实现方案与代码

环境准备

# requirements.txt
langchain>=0.1.0
langchain-openai>=0.0.5
langchain-redis>=0.0.1
langgraph>=0.0.20
redis>=5.0.0
pydantic>=2.0.0

场景 1:会话状态持久化(LangGraph Checkpointer)

思路:使用 Redis 作为 LangGraph 的 checkpointer,实现 Agent 状态的持久化和恢复。

from langgraph.checkpoint.redis import RedisSaver
from langgraph.graph import StateGraph, END
from langchain_core.messages import HumanMessage, AIMessage
from typing import TypedDict, List, Optional
import os

# 配置 Redis
REDIS_URL = "redis://localhost:6379/0"

class AgentState(TypedDict):
    messages: List[HumanMessage | AIMessage]
    step_count: int
    current_task: Optional[str]
    tool_results: dict

def create_agent_graph():
    """创建带 Redis 持久化的 Agent"""
    
    # 初始化 Redis Checkpointer
    with RedisSaver.from_conn_string(REDIS_URL) as checkpointer:
        
        # 定义节点函数
        def process_message(state: AgentState):
            """处理用户消息"""
            last_message = state["messages"][-1].content
            
            # 模拟 Agent 处理逻辑
            response = f"处理了消息: {last_message}"
            
            return {
                "messages": [AIMessage(content=response)],
                "step_count": state.get("step_count", 0) + 1,
                "current_task": None
            }
        
        def should_continue(state: AgentState):
            """决定是否继续处理"""
            if state.get("step_count", 0) > 5:
                return END
            return "process"
        
        # 构建图
        workflow = StateGraph(AgentState)
        workflow.add_node("process", process_message)
        workflow.set_entry_point("process")
        workflow.add_conditional_edges(
            "process",
            should_continue,
            {"process": "process", END: END}
        )
        
        # 编译图并绑定 Redis checkpointer
        app = workflow.compile(checkpointer=checkpointer)
        return app

# 使用示例
if __name__ == "__main__":
    app = create_agent_graph()
    
    config = {
        "configurable": {
            "thread_id": "user_session_123"  # 关键:通过 thread_id 区分不同会话
        }
    }
    
    # 第一次调用
    result1 = app.invoke(
        {"messages": [HumanMessage(content="你好,我是用户")]},
        config=config
    )
    print(f"Step 1: {result1['messages'][-1].content}")
    
    # 第二次调用(自动从 Redis 恢复状态)
    result2 = app.invoke(
        {"messages": [HumanMessage(content="继续我们的对话")]},
        config=config
    )
    print(f"Step 2: {result2['messages'][-1].content}")
    print(f"总步数: {result2['step_count']}")

场景 2:短期记忆管理(滑动窗口)

思路:使用 Redis List 维护固定长度的对话历史,实现滑动窗口记忆。

import json
import redis
from langchain.memory import ConversationBufferWindowMemory
from langchain.schema import BaseChatMessageHistory
from typing import List
from langchain_core.messages import BaseMessage, HumanMessage, AIMessage

class RedisChatMessageHistory(BaseChatMessageHistory):
    """自定义 Redis 聊天历史存储"""
    
    def __init__(self, session_id: str, redis_url: str = "redis://localhost:6379/0", ttl: int = 3600):
        self.redis_client = redis.from_url(redis_url)
        self.session_id = session_id
        self.ttl = ttl  # 会话过期时间(秒)
        self.key_prefix = "chat_history:"
    
    @property
    def messages(self) -> List[BaseMessage]:
        """从 Redis 获取所有消息"""
        key = f"{self.key_prefix}{self.session_id}"
        messages_json = self.redis_client.lrange(key, 0, -1)
        
        messages = []
        for msg_json in messages_json:
            msg_dict = json.loads(msg_json)
            if msg_dict["type"] == "human":
                messages.append(HumanMessage(content=msg_dict["content"]))
            elif msg_dict["type"] == "ai":
                messages.append(AIMessage(content=msg_dict["content"]))
        return messages
    
    def add_message(self, message: BaseMessage) -> None:
        """添加消息到 Redis"""
        key = f"{self.key_prefix}{self.session_id}"
        msg_dict = {
            "type": "human" if isinstance(message, HumanMessage) else "ai",
            "content": message.content
        }
        
        # 添加到列表末尾
        self.redis_client.rpush(key, json.dumps(msg_dict))
        # 设置过期时间
        self.redis_client.expire(key, self.ttl)
    
    def clear(self) -> None:
        """清空历史"""
        key = f"{self.key_prefix}{self.session_id}"
        self.redis_client.delete(key)

# 集成到 LangChain Agent
from langchain_openai import ChatOpenAI
from langchain.agents import AgentExecutor, create_react_agent
from langchain.tools import Tool
from langchain import hub

def create_agent_with_redis_memory(session_id: str):
    """创建带 Redis 记忆的 Agent"""
    
    # 创建 Redis 消息历史
    message_history = RedisChatMessageHistory(
        session_id=session_id,
        ttl=7200  # 2小时过期
    )
    
    # 创建滑动窗口记忆
    memory = ConversationBufferWindowMemory(
        k=10,  # 保留最近10轮对话
        chat_memory=message_history,
        return_messages=True
    )
    
    # 定义工具
    tools = [
        Tool(
            name="Calculator",
            func=lambda x: str(eval(x)),
            description="用于数学计算"
        )
    ]
    
    # 创建 LLM
    llm = ChatOpenAI(model="gpt-4", temperature=0)
    
    # 获取 ReAct 提示模板
    prompt = hub.pull("hwchase17/react-chat")
    
    # 创建 Agent
    agent = create_react_agent(llm, tools, prompt)
    
    # 创建 Agent Executor,集成记忆
    agent_executor = AgentExecutor(
        agent=agent,
        tools=tools,
        memory=memory,
        verbose=True,
        handle_parsing_errors=True
    )
    
    return agent_executor

# 使用示例
if __name__ == "__main__":
    agent = create_agent_with_redis_memory("user_123")
    
    # 第一次对话
    response1 = agent.invoke({
        "input": "我的名字是张三,请计算 15 * 8"
    })
    print(f"Response 1: {response1['output']}")
    
    # 第二次对话(会记住上下文)
    response2 = agent.invoke({
        "input": "我叫什么名字?刚才的计算结果是多少?"
    })
    print(f"Response 2: {response2['output']}")

场景 3:LLM 响应缓存

思路:使用 Redis 缓存 LLM 的响应,避免重复调用相同提示词。

import hashlib
import pickle
from langchain.cache import RedisCache
from langchain.globals import set_llm_cache
from langchain_openai import ChatOpenAI

def setup_llm_cache():
    """配置 Redis LLM 缓存"""
    
    # 方式1:使用 LangChain 内置的 RedisCache
    redis_cache = RedisCache(redis_=redis.Redis(host='localhost', port=6379, db=1))
    set_llm_cache(redis_cache)
    
    # 方式2:自定义高级缓存策略
    class AdvancedRedisLLMCache:
        def __init__(self, redis_client, ttl=86400):
            self.redis = redis_client
            self.ttl = ttl  # 默认缓存1天
            self.key_prefix = "llm_cache:"
        
        def _generate_key(self, prompt: str, model: str, **kwargs) -> str:
            """生成缓存键"""
            # 包含模型名称和参数,确保不同配置的缓存不冲突
            cache_data = {
                "prompt": prompt,
                "model": model,
                "temperature": kwargs.get("temperature", 0),
                "max_tokens": kwargs.get("max_tokens", 1000)
            }
            data_str = json.dumps(cache_data, sort_keys=True)
            key_hash = hashlib.sha256(data_str.encode()).hexdigest()
            return f"{self.key_prefix}{key_hash}"
        
        def get(self, prompt: str, model: str, **kwargs):
            """获取缓存"""
            key = self._generate_key(prompt, model, **kwargs)
            cached = self.redis.get(key)
            
            if cached:
                try:
                    return pickle.loads(cached)
                except:
                    return None
            return None
        
        def set(self, prompt: str, model: str, value, **kwargs):
            """设置缓存"""
            key = self._generate_key(prompt, model, **kwargs)
            try:
                serialized = pickle.dumps(value)
                self.redis.setex(key, self.ttl, serialized)
            except Exception as e:
                print(f"缓存写入失败: {e}")
        
        def clear_cache_for_prompt(self, prompt_pattern: str):
            """清除匹配模式的缓存"""
            pattern = f"{self.key_prefix}*{hashlib.sha256(prompt_pattern.encode()).hexdigest()[:8]}*"
            keys = self.redis.keys(pattern)
            if keys:
                self.redis.delete(*keys)
    
    # 使用自定义缓存
    redis_client = redis.Redis(host='localhost', port=6379, db=2)
    custom_cache = AdvancedRedisLLMCache(redis_client, ttl=3600)
    
    # 包装 LLM
    class CachedChatOpenAI(ChatOpenAI):
        def __init__(self, cache, *args, **kwargs):
            super().__init__(*args, **kwargs)
            self.cache = cache
        
        def _generate(self, messages, stop=None, run_manager=None, **kwargs):
            # 将消息转换为字符串作为缓存键的一部分
            prompt = "\n".join([f"{msg.type}: {msg.content}" for msg in messages])
            
            # 尝试从缓存获取
            cached_result = self.cache.get(prompt, self.model_name, **kwargs)
            if cached_result:
                print("✅ 从缓存返回结果")
                return cached_result
            
            # 调用原始方法
            result = super()._generate(messages, stop, run_manager, **kwargs)
            
            # 存入缓存
            self.cache.set(prompt, self.model_name, result, **kwargs)
            print("💾 结果已缓存")
            
            return result
    
    # 创建带缓存的 LLM
    cached_llm = CachedChatOpenAI(
        cache=custom_cache,
        model="gpt-4",
        temperature=0
    )
    
    return cached_llm

# 使用示例
if __name__ == "__main__":
    llm = setup_llm_cache()
    
    # 第一次调用(会调用 API)
    response1 = llm.invoke("解释什么是机器学习?")
    print(f"Response 1: {response1.content[:100]}...")
    
    # 第二次调用相同问题(直接从 Redis 缓存返回)
    response2 = llm.invoke("解释什么是机器学习?")
    print(f"Response 2: {response2.content[:100]}...")

场景 4:分布式锁与并发控制

思路:使用 Redis 分布式锁防止 Agent 重复执行或资源竞争。

import redis
import time
import uuid
from contextlib import contextmanager
from typing import Optional

class RedisDistributedLock:
    """Redis 分布式锁"""
    
    def __init__(self, redis_client, lock_name: str, expire_time: int = 30):
        self.redis = redis_client
        self.lock_name = f"distributed_lock:{lock_name}"
        self.expire_time = expire_time  # 锁过期时间(秒)
        self.identifier = str(uuid.uuid4())  # 锁持有者标识
    
    @contextmanager
    def acquire(self, block=True, timeout=10):
        """获取锁的上下文管理器"""
        locked = False
        start_time = time.time()
        
        try:
            while not locked:
                # 尝试获取锁
                locked = self.redis.set(
                    self.lock_name, 
                    self.identifier, 
                    nx=True, 
                    ex=self.expire_time
                )
                
                if locked:
                    yield True
                    break
                elif not block:
                    yield False
                    break
                elif time.time() - start_time > timeout:
                    yield False
                    break
                else:
                    time.sleep(0.1)
        finally:
            # 释放锁(只有锁的持有者才能释放)
            if locked:
                script = """
                if redis.call("get", KEYS[1]) == ARGV[1] then
                    return redis.call("del", KEYS[1])
                else
                    return 0
                end
                """
                self.redis.eval(script, 1, self.lock_name, self.identifier)

class AgentConcurrencyManager:
    """Agent 并发管理器"""
    
    def __init__(self, redis_url: str = "redis://localhost:6379/0"):
        self.redis = redis.from_url(redis_url)
        self.active_agents_key = "active_agents"
        self.max_concurrent = 5  # 最大并发数
    
    def register_agent(self, agent_id: str, task_type: str = "default") -> bool:
        """注册 Agent 执行"""
        key = f"agent_registry:{task_type}"
        
        # 使用 Lua 脚本确保原子性操作
        script = """
        local current = tonumber(redis.call('get', KEYS[1]) or '0')
        if current >= tonumber(ARGV[1]) then
            return 0
        end
        redis.call('incr', KEYS[1])
        redis.call('expire', KEYS[1], 300)  -- 5分钟过期
        redis.call('hset', KEYS[2], ARGV[2], tostring(current + 1))
        redis.call('expire', KEYS[2], 300)
        return 1
        """
        
        result = self.redis.eval(
            script, 
            2,  # 两个键
            key, 
            f"{key}:details",
            self.max_concurrent,
            agent_id
        )
        
        return bool(result)
    
    def unregister_agent(self, agent_id: str, task_type: str = "default"):
        """注销 Agent 执行"""
        key = f"agent_registry:{task_type}"
        self.redis.hdel(f"{key}:details", agent_id)
        self.redis.decr(key)

# 在 LangGraph Agent 中使用分布式锁
from langgraph.graph import StateGraph, END
from typing import TypedDict, Literal

class LockProtectedState(TypedDict):
    agent_id: str
    task_status: Literal["pending", "running", "completed", "failed"]
    lock_acquired: bool

def create_lock_protected_agent():
    """创建带分布式锁保护的 Agent"""
    
    redis_client = redis.from_url("redis://localhost:6379/0")
    concurrency_manager = AgentConcurrencyManager()
    
    def acquire_lock_node(state: LockProtectedState):
        """获取分布式锁"""
        agent_id = state["agent_id"]
        
        # 检查并发限制
        if not concurrency_manager.register_agent(agent_id):
            return {
                **state,
                "task_status": "failed",
                "lock_acquired": False
            }
        
        # 获取具体任务锁
        lock = RedisDistributedLock(redis_client, f"task_lock:{agent_id}")
        
        with lock.acquire(block=False) as acquired:
            if acquired:
                return {
                    **state,
                    "task_status": "running",
                    "lock_acquired": True
                }
            else:
                concurrency_manager.unregister_agent(agent_id)
                return {
                    **state,
                    "task_status": "failed",
                    "lock_acquired": False
                }
    
    def execute_task_node(state: LockProtectedState):
        """执行任务(只有获取到锁的实例才能执行)"""
        if not state["lock_acquired"]:
            return state
        
        # 模拟长时间运行的任务
        time.sleep(2)
        
        return {
            **state,
            "task_status": "completed"
        }
    
    def release_resources_node(state: LockProtectedState):
        """释放资源"""
        if state["lock_acquired"]:
            concurrency_manager.unregister_agent(state["agent_id"])
        
        return state
    
    # 构建工作流
    workflow = StateGraph(LockProtectedState)
    workflow.add_node("acquire_lock", acquire_lock_node)
    workflow.add_node("execute_task", execute_task_node)
    workflow.add_node("release_resources", release_resources_node)
    
    workflow.set_entry_point("acquire_lock")
    workflow.add_edge("acquire_lock", "execute_task")
    workflow.add_edge("execute_task", "release_resources")
    workflow.add_edge("release_resources", END)
    
    return workflow.compile()

# 使用示例
if __name__ == "__main__":
    app = create_lock_protected_agent()
    
    # 模拟多个并发请求
    import threading
    
    def run_agent(agent_id):
        result = app.invoke({
            "agent_id": agent_id,
            "task_status": "pending",
            "lock_acquired": False
        })
        print(f"Agent {agent_id}: {result['task_status']}")
    
    threads = []
    for i in range(10):  # 尝试启动10个并发Agent
        t = threading.Thread(target=run_agent, args=(f"agent_{i}",))
        threads.append(t)
        t.start()
    
    for t in threads:
        t.join()

场景 5:RAG 向量存储

思路:使用 Redis Stack 的向量搜索功能作为 RAG 的向量数据库。

from langchain_community.vectorstores import Redis
from langchain_openai import OpenAIEmbeddings
from langchain.text_splitter import RecursiveCharacterTextSplitter
from langchain.document_loaders import TextLoader
from langchain.chains import RetrievalQA
from langchain_openai import ChatOpenAI

class RedisRAGSystem:
    """基于 Redis 的 RAG 系统"""
    
    def __init__(self, redis_url: str = "redis://localhost:6379"):
        self.redis_url = redis_url
        self.embeddings = OpenAIEmbeddings()
        self.index_name = "document_index"
        self.index_schema = {
            "tag": [{"name": "source"}],
            "text": [{"name": "content"}],
            "vector": [{
                "name": "embedding",
                "algorithm": "HNSW",
                "datatype": "float32",
                "dims": 1536,  # OpenAI embedding 维度
                "distance_metric": "cosine"
            }]
        }
    
    def ingest_documents(self, file_path: str, chunk_size: int = 1000):
        """文档摄入流程"""
        # 加载文档
        loader = TextLoader(file_path)
        documents = loader.load()
        
        # 分割文档
        text_splitter = RecursiveCharacterTextSplitter(
            chunk_size=chunk_size,
            chunk_overlap=200
        )
        splits = text_splitter.split_documents(documents)
        
        # 创建 Redis 向量存储
        vectorstore = Redis.from_documents(
            documents=splits,
            embedding=self.embeddings,
            redis_url=self.redis_url,
            index_name=self.index_name,
            index_schema=self.index_schema
        )
        
        return vectorstore
    
    def create_qa_chain(self):
        """创建 QA 链"""
        # 加载现有向量存储
        vectorstore = Redis(
            redis_url=self.redis_url,
            index_name=self.index_name,
            embedding_function=self.embeddings
        )
        
        # 创建检索器
        retriever = vectorstore.as_retriever(
            search_type="similarity",
            search_kwargs={"k": 4}
        )
        
        # 创建 LLM
        llm = ChatOpenAI(model="gpt-4", temperature=0)
        
        # 创建 QA 链
        qa_chain = RetrievalQA.from_chain_type(
            llm=llm,
            chain_type="stuff",
            retriever=retriever,
            return_source_documents=True,
            verbose=True
        )
        
        return qa_chain
    
    def hybrid_search(self, query: str, filters: dict = None):
        """混合搜索(向量+标签过滤)"""
        vectorstore = Redis(
            redis_url=self.redis_url,
            index_name=self.index_name,
            embedding_function=self.embeddings
        )
        
        # 构建过滤器
        filter_expression = ""
        if filters:
            filter_parts = []
            for key, value in filters.items():
                filter_parts.append(f"@{key}:{{{value}}}")
            filter_expression = " ".join(filter_parts)
        
        # 执行搜索
        results = vectorstore.similarity_search(
            query=query,
            k=5,
            filter=filter_expression if filter_expression else None
        )
        
        return results

# 集成到 Agent
from langchain.agents import AgentExecutor, create_react_agent
from langchain.tools import Tool

def create_rag_enabled_agent():
    """创建带 RAG 能力的 Agent"""
    
    rag_system = RedisRAGSystem()
    
    # 如果还没有摄入文档,先摄入
    try:
        qa_chain = rag_system.create_qa_chain()
    except:
        print("索引不存在,正在摄入文档...")
        rag_system.ingest_documents("./knowledge_base.txt")
        qa_chain = rag_system.create_qa_chain()
    
    # 定义 RAG 工具
    def rag_search(query: str) -> str:
        """基于知识库的搜索"""
        result = qa_chain.invoke({"query": query})
        sources = [doc.metadata.get("source", "") for doc in result["source_documents"]]
        return f"答案: {result['result']}\n来源: {', '.join(set(sources))}"
    
    tools = [
        Tool(
            name="KnowledgeBaseSearch",
            func=rag_search,
            description="当需要查询公司内部文档、技术手册或特定领域知识时使用此工具"
        )
    ]
    
    # 创建 Agent
    llm = ChatOpenAI(model="gpt-4", temperature=0)
    prompt = hub.pull("hwchase17/react")
    agent = create_react_agent(llm, tools, prompt)
    
    agent_executor = AgentExecutor(
        agent=agent,
        tools=tools,
        verbose=True
    )
    
    return agent_executor

# 使用示例
if __name__ == "__main__":
    agent = create_rag_enabled_agent()
    
    response = agent.invoke({
        "input": "我们的产品退货政策是什么?"
    })
    print(response["output"])

场景 6:工作流状态管理与任务队列

思路:使用 Redis Streams 或 Lists 管理工作流任务和状态。

import redis
import json
import uuid
from datetime import datetime
from typing import Dict, Any, List, Optional
from enum import Enum

class TaskStatus(Enum):
    PENDING = "pending"
    PROCESSING = "processing"
    COMPLETED = "completed"
    FAILED = "failed"
    RETRYING = "retrying"

class WorkflowTaskManager:
    """工作流任务管理器"""
    
    def __init__(self, redis_url: str = "redis://localhost:6379/0"):
        self.redis = redis.from_url(redis_url)
        self.task_stream = "workflow_tasks"
        self.task_status_prefix = "task_status:"
        self.task_results_prefix = "task_results:"
        self.consumer_group = "workflow_workers"
        self.max_retries = 3
    
    def initialize_consumer_group(self):
        """初始化消费者组"""
        try:
            self.redis.xgroup_create(
                self.task_stream, 
                self.consumer_group, 
                id='0',
                mkstream=True
            )
        except redis.ResponseError as e:
            if "BUSYGROUP" not in str(e):
                raise
    
    def submit_workflow(self, workflow_id: str, tasks: List[Dict[str, Any]]) -> str:
        """提交工作流任务"""
        execution_id = f"{workflow_id}:{uuid.uuid4()}"
        
        # 为每个任务创建记录
        for idx, task in enumerate(tasks):
            task_id = f"{execution_id}:task_{idx}"
            task_data = {
                "task_id": task_id,
                "workflow_id": workflow_id,
                "execution_id": execution_id,
                "task_type": task["type"],
                "payload": json.dumps(task["payload"]),
                "dependencies": json.dumps(task.get("dependencies", [])),
                "status": TaskStatus.PENDING.value,
                "retry_count": 0,
                "created_at": datetime.now().isoformat()
            }
            
            # 存储任务状态
            self.redis.hset(
                f"{self.task_status_prefix}{task_id}",
                mapping=task_data
            )
            
            # 添加到任务队列
            self.redis.xadd(
                self.task_stream,
                {"task_id": task_id, "data": json.dumps(task_data)}
            )
        
        # 设置工作流元数据
        self.redis.hset(
            f"workflow:{execution_id}",
            mapping={
                "workflow_id": workflow_id,
                "total_tasks": len(tasks),
                "completed_tasks": 0,
                "failed_tasks": 0,
                "status": "running",
                "started_at": datetime.now().isoformat()
            }
        )
        
        return execution_id
    
    def claim_pending_task(self, consumer_id: str, block_ms: int = 5000) -> Optional[Dict]:
        """领取待处理任务"""
        try:
            # 从流中读取任务
            streams = self.redis.xreadgroup(
                self.consumer_group,
                consumer_id,
                {self.task_stream: ">"},
                count=1,
                block=block_ms
            )
            
            if not streams:
                return None
            
            stream_name, messages = streams[0]
            
            for message_id, data in messages:
                task_id = data[b"task_id"].decode()
                
                # 检查任务状态
                task_status = self.redis.hgetall(f"{self.task_status_prefix}{task_id}")
                
                if not task_status:
                    continue
                
                # 更新为处理中
                self.redis.hset(
                    f"{self.task_status_prefix}{task_id}",
                    mapping={
                        "status": TaskStatus.PROCESSING.value,
                        "consumer_id": consumer_id,
                        "message_id": message_id.decode(),
                        "started_at": datetime.now().isoformat()
                    }
                )
                
                return {
                    "task_id": task_id,
                    "message_id": message_id.decode(),
                    "task_data": json.loads(task_status[b"payload"].decode()),
                    "task_type": task_status[b"task_type"].decode()
                }
        
        except Exception as e:
            print(f"领取任务失败: {e}")
        
        return None
    
    def complete_task(self, task_id: str, result: Any = None, error: str = None):
        """完成任务"""
        status = TaskStatus.COMPLETED.value if not error else TaskStatus.FAILED.value
        
        update_data = {
            "status": status,
            "completed_at": datetime.now().isoformat()
        }
        
        if result is not None:
            # 存储结果
            self.redis.set(
                f"{self.task_results_prefix}{task_id}",
                json.dumps(result),
                ex=86400  # 24小时过期
            )
            update_data["result_available"] = "true"
        
        if error:
            update_data["error"] = error
            # 检查是否需要重试
            retry_count = int(self.redis.hget(f"{self.task_status_prefix}{task_id}", "retry_count") or 0)
            if retry_count < self.max_retries:
                update_data["status"] = TaskStatus.RETRYING.value
                update_data["retry_count"] = retry_count + 1
        
        self.redis.hset(f"{self.task_status_prefix}{task_id}", mapping=update_data)
        
        # 更新工作流进度
        execution_id = task_id.split(":")[0] + ":" + task_id.split(":")[1]
        if status == TaskStatus.COMPLETED.value:
            self.redis.hincrby(f"workflow:{execution_id}", "completed_tasks", 1)
        else:
            self.redis.hincrby(f"workflow:{execution_id}", "failed_tasks", 1)
        
        # 检查工作流是否完成
        workflow_info = self.redis.hgetall(f"workflow:{execution_id}")
        total = int(workflow_info.get(b"total_tasks", 0))
        completed = int(workflow_info.get(b"completed_tasks", 0))
        failed = int(workflow_info.get(b"failed_tasks", 0))
        
        if completed + failed >= total:
            final_status = "completed" if failed == 0 else "partially_failed"
            self.redis.hset(
                f"workflow:{execution_id}",
                mapping={
                    "status": final_status,
                    "finished_at": datetime.now().isoformat()
                }
            )
    
    def get_workflow_status(self, execution_id: str) -> Dict:
        """获取工作流状态"""
        workflow_info = self.redis.hgetall(f"workflow:{execution_id}")
        
        # 解码字节键
        decoded_info = {}
        for k, v in workflow_info.items():
            decoded_info[k.decode()] = v.decode() if isinstance(v, bytes) else v
        
        # 获取所有任务状态
        task_pattern = f"{self.task_status_prefix}{execution_id}*"
        task_keys = self.redis.keys(task_pattern)
        
        tasks = []
        for key in task_keys:
            task_data = self.redis.hgetall(key)
            decoded_task = {}
            for k, v in task_data.items():
                decoded_task[k.decode()] = v.decode() if isinstance(v, bytes) else v
            tasks.append(decoded_task)
        
        decoded_info["tasks"] = tasks
        return decoded_info

# 在 LangGraph 中集成任务队列
from langgraph.graph import StateGraph, END
from typing import TypedDict, List

class TaskQueueState(TypedDict):
    execution_id: str
    pending_tasks: List[str]
    completed_tasks: List[str]
    failed_tasks: List[str]

def create_queued_workflow_agent():
    """创建带任务队列的 Agent"""
    
    task_manager = WorkflowTaskManager()
    task_manager.initialize_consumer_group()
    
    def submit_tasks_node(state: TaskQueueState):
        """提交任务到队列"""
        execution_id = state.get("execution_id") or f"exec_{uuid.uuid4()}"
        
        # 定义工作流任务
        tasks = [
            {
                "type": "data_extraction",
                "payload": {"source": "database", "table": "users"}
            },
            {
                "type": "data_transformation",
                "payload": {"operation": "normalize"},
                "dependencies": ["data_extraction"]
            },
            {
                "type": "report_generation",
                "payload": {"format": "pdf"},
                "dependencies": ["data_transformation"]
            }
        ]
        
        # 提交到任务管理器
        exec_id = task_manager.submit_workflow("etl_workflow", tasks)
        
        return {
            **state,
            "execution_id": exec_id,
            "pending_tasks": [t["type"] for t in tasks]
        }
    
    def monitor_progress_node(state: TaskQueueState):
        """监控任务进度"""
        execution_id = state["execution_id"]
        status = task_manager.get_workflow_status(execution_id)
        
        completed = []
        failed = []
        
        for task in status.get("tasks", []):
            if task.get("status") == "completed":
                completed.append(task["task_type"])
            elif task.get("status") == "failed":
                failed.append(task["task_type"])
        
        return {
            **state,
            "completed_tasks": completed,
            "failed_tasks": failed
        }
    
    def should_continue(state: TaskQueueState):
        """判断是否继续执行"""
        total_pending = len(state.get("pending_tasks", []))
        total_completed = len(state.get("completed_tasks", []))
        total_failed = len(state.get("failed_tasks", []))
        
        if total_completed + total_failed >= total_pending:
            if total_failed > 0:
                return "handle_failures"
            return END
        
        return "monitor"
    
    # 构建工作流
    workflow = StateGraph(TaskQueueState)
    workflow.add_node("submit", submit_tasks_node)
    workflow.add_node("monitor", monitor_progress_node)
    workflow.add_node("handle_failures", lambda s: {**s, "status": "failed"})
    
    workflow.set_entry_point("submit")
    workflow.add_edge("submit", "monitor")
    workflow.add_conditional_edges(
        "monitor",
        should_continue,
        {
            "monitor": "monitor",
            "handle_failures": "handle_failures",
            END: END
        }
    )
    
    return workflow.compile()

# 任务处理器(Worker)
def task_processor_worker(worker_id: str):
    """任务处理工作线程"""
    task_manager = WorkflowTaskManager()
    
    while True:
        task = task_manager.claim_pending_task(worker_id, block_ms=5000)
        
        if not task:
            time.sleep(1)
            continue
        
        print(f"Worker {worker_id} 处理任务: {task['task_type']}")
        
        try:
            # 根据任务类型执行不同的处理逻辑
            if task["task_type"] == "data_extraction":
                result = {"records": 1000, "status": "success"}
            elif task["task_type"] == "data_transformation":
                result = {"normalized": True, "count": 950}
            elif task["task_type"] == "report_generation":
                result = {"file": "report.pdf", "pages": 25}
            else:
                raise ValueError(f"未知任务类型: {task['task_type']}")
            
            task_manager.complete_task(task["task_id"], result=result)
        
        except Exception as e:
            print(f"任务失败: {e}")
            task_manager.complete_task(task["task_id"], error=str(e))

# 使用示例
if __name__ == "__main__":
    # 启动工作流
    app = create_queued_workflow_agent()
    result = app.invoke({})
    
    print(f"工作流已启动: {result['execution_id']}")
    
    # 启动工作线程(实际部署时会是独立的进程)
    import threading
    workers = []
    for i in range(3):
        worker = threading.Thread(
            target=task_processor_worker,
            args=(f"worker_{i}",),
            daemon=True
        )
        worker.start()
        workers.append(worker)
    
    # 等待工作流完成
    time.sleep(10)
    
    # 检查最终状态
    task_manager = WorkflowTaskManager()
    final_status = task_manager.get_workflow_status(result['execution_id'])
    print(f"最终状态: {final_status['status']}")

三、综合最佳实践示例

"""
完整的生产级 Agent 系统示例
集成 Redis 的所有主要功能
"""
import os
from typing import TypedDict, List, Optional, Dict, Any
from langgraph.graph import StateGraph, END
from langgraph.checkpoint.redis import RedisSaver
from langchain_openai import ChatOpenAI
from langchain_core.messages import HumanMessage, AIMessage
from langchain.memory import ConversationBufferWindowMemory
from langchain.cache import RedisCache
from langchain.globals import set_llm_cache
import redis
import json

# 配置
REDIS_URL = "redis://localhost:6379/0"
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")

class ProductionAgentState(TypedDict):
    """生产级 Agent 状态定义"""
    messages: List[HumanMessage | AIMessage]
    session_id: str
    user_id: str
    current_step: str
    context: Dict[str, Any]
    tool_calls: List[Dict]
    error_count: int
    metadata: Dict[str, Any]

class ProductionAgentSystem:
    """生产级 Agent 系统"""
    
    def __init__(self):
        # Redis 客户端
        self.redis_client = redis.from_url(REDIS_URL)
        
        # 配置 LLM 缓存
        set_llm_cache(RedisCache(redis_=self.redis_client))
        
        # LLM
        self.llm = ChatOpenAI(
            model="gpt-4",
            temperature=0.7,
            max_tokens=2000
        )
        
        # 初始化系统
        self._setup_components()
        self.graph = self._build_graph()
    
    def _setup_components(self):
        """初始化各个组件"""
        
        # 1. 分布式锁
        self.locks = {
            "session": lambda sid: redis.lock.Lock(
                self.redis_client, 
                f"session_lock:{sid}",
                timeout=30
            ),
            "tool": lambda tid: redis.lock.Lock(
                self.redis_client,
                f"tool_lock:{tid}",
                timeout=60
            )
        }
        
        # 2. 记忆管理器
        self.memory_store = {}
        
        # 3. 速率限制器
        self.rate_limiter_script = """
        local key = KEYS[1]
        local limit = tonumber(ARGV[1])
        local window = tonumber(ARGV[2])
        
        local current = redis.call('INCR', key)
        if current == 1 then
            redis.call('EXPIRE', key, window)
        end
        
        if current > limit then
            return 0
        end
        
        return 1
        """
    
    def _check_rate_limit(self, user_id: str, limit: int = 100, window: int = 3600) -> bool:
        """检查用户速率限制"""
        key = f"rate_limit:{user_id}"
        return bool(self.redis_client.eval(
            self.rate_limiter_script,
            1,
            key,
            limit,
            window
        ))
    
    def _get_or_create_memory(self, session_id: str):
        """获取或创建会话记忆"""
        if session_id not in self.memory_store:
            # 从 Redis 恢复记忆
            history_key = f"memory:{session_id}"
            messages_json = self.redis_client.lrange(history_key, 0, -1)
            
            memory = ConversationBufferWindowMemory(
                k=20,
                return_messages=True
            )
            
            # 恢复历史消息
            for msg_json in messages_json:
                msg_dict = json.loads(msg_json)
                if msg_dict["role"] == "user":
                    memory.chat_memory.add_user_message(msg_dict["content"])
                else:
                    memory.chat_memory.add_ai_message(msg_dict["content"])
            
            self.memory_store[session_id] = memory
        
        return self.memory_store[session_id]
    
    def _save_memory(self, session_id: str, memory):
        """保存记忆到 Redis"""
        history_key = f"memory:{session_id}"
        
        # 清空现有历史
        self.redis_client.delete(history_key)
        
        # 保存消息
        for msg in memory.chat_memory.messages:
            msg_dict = {
                "role": "user" if isinstance(msg, HumanMessage) else "ai",
                "content": msg.content,
                "timestamp": json.dumps({"$date": msg.additional_kwargs.get("timestamp")})
            }
            self.redis_client.rpush(history_key, json.dumps(msg_dict))
        
        # 设置过期时间(7天)
        self.redis_client.expire(history_key, 604800)
    
    def _build_graph(self):
        """构建 LangGraph 工作流"""
        
        def validate_input(state: ProductionAgentState):
            """验证输入"""
            session_id = state["session_id"]
            
            # 速率限制检查
            if not self._check_rate_limit(state["user_id"]):
                return {
                    **state,
                    "current_step": "rate_limited",
                    "error_count": state.get("error_count", 0) + 1
                }
            
            # 会话锁检查
            with self.lockssession_id:
                # 获取或创建记忆
                memory = self._get_or_create_memory(session_id)
                
                return {
                    **state,
                    "current_step": "process",
                    "context": {
                        **state.get("context", {}),
                        "memory": memory
                    }
                }
        
        def process_request(state: ProductionAgentState):
            """处理用户请求"""
            memory = state["context"]["memory"]
            
            # 构建提示词(包含历史记忆)
            messages = memory.chat_memory.messages + [state["messages"][-1]]
            
            # 调用 LLM
            response = self.llm.invoke(messages)
            
            # 更新记忆
            memory.chat_memory.add_user_message(state["messages"][-1].content)
            memory.chat_memory.add_ai_message(response.content)
            
            # 保存到 Redis
            self._save_memory(state["session_id"], memory)
            
            return {
                **state,
                "messages": state["messages"] + [response],
                "current_step": "complete"
            }
        
        def handle_error(state: ProductionAgentState):
            """错误处理"""
            error_count = state.get("error_count", 0)
            
            # 指数退避重试
            if error_count < 3:
                backoff = 2 ** error_count
                time.sleep(backoff)
                return {
                    **state,
                    "current_step": "validate",
                    "error_count": error_count + 1
                }
            
            # 超过重试次数,记录错误并返回友好消息
            error_msg = AIMessage(
                content="抱歉,系统暂时无法处理您的请求,请稍后再试。"
            )
            
            return {
                **state,
                "messages": state["messages"] + [error_msg],
                "current_step": "failed"
            }
        
        # 构建图
        workflow = StateGraph(ProductionAgentState)
        
        workflow.add_node("validate", validate_input)
        workflow.add_node("process", process_request)
        workflow.add_node("error_handler", handle_error)
        
        workflow.set_entry_point("validate")
        
        workflow.add_conditional_edges(
            "validate",
            lambda s: s["current_step"],
            {
                "process": "process",
                "rate_limited": "error_handler"
            }
        )
        
        workflow.add_conditional_edges(
            "process",
            lambda s: s["current_step"],
            {
                "complete": END,
                "failed": "error_handler"
            }
        )
        
        workflow.add_edge("error_handler", END)
        
        # 使用 Redis Checkpointer
        with RedisSaver.from_conn_string(REDIS_URL) as checkpointer:
            return workflow.compile(checkpointer=checkpointer)
    
    def invoke(self, input_data: Dict[str, Any], config: Optional[Dict] = None):
        """调用 Agent"""
        default_config = {
            "configurable": {
                "thread_id": input_data.get("session_id", "default")
            }
        }
        
        return self.graph.invoke(
            input_data,
            config=config or default_config
        )

# 使用示例
if __name__ == "__main__":
    # 初始化系统
    agent_system = ProductionAgentSystem()
    
    # 模拟用户请求
    requests = [
        {
            "messages": [HumanMessage(content="你好,我想了解你们的AI服务")],
            "session_id": "user_001_session",
            "user_id": "user_001",
            "current_step": "start",
            "context": {},
            "tool_calls": [],
            "error_count": 0,
            "metadata": {"source": "web"}
        },
        {
            "messages": [HumanMessage(content="价格是多少?有什么套餐?")],
            "session_id": "user_001_session",
            "user_id": "user_001",
            "current_step": "start",
            "context": {},
            "tool_calls": [],
            "error_count": 0,
            "metadata": {"source": "web"}
        }
    ]
    
    # 处理请求
    for req in requests:
        print(f"\n用户: {req['messages'][-1].content}")
        result = agent_system.invoke(req)
        print(f"Agent: {result['messages'][-1].content}")

四、总结

Redis 在 LangChain/LangGraph Agent 开发中的核心价值

场景 Redis 作用 关键优势
会话状态持久化 LangGraph Checkpointer 支持断点续传、多节点部署、故障恢复
短期记忆管理 滑动窗口存储 跨重启保持上下文、自动过期清理
LLM 响应缓存 KV 缓存层 降低 API 成本、提高响应速度
分布式锁 并发控制 防止重复执行、资源竞争协调
RAG 向量存储 向量数据库 低延迟检索、混合搜索、实时更新
任务队列 Streams/Lists 异步处理、负载均衡、失败重试

最佳实践建议

  1. 分层设计

    # 推荐架构
    ┌─────────────────────────────────────┐
    │         LangGraph Workflow          │
    │  ┌─────────┐  ┌─────────────────┐  │
    │  │  Nodes  │  │  Redis Check-   │  │
    │  │         │  │  pointer        │  │
    │  └─────────┘  └─────────────────┘  │
    ├─────────────────────────────────────┤
    │  Redis Layer                        │
    │  ┌──────┐ ┌──────┐ ┌────────────┐  │
    │  │Cache │ │State │ │Vector Store│  │
    │  └──────┘ └──────┘ └────────────┘  │
    ├─────────────────────────────────────┤
    │  Infrastructure                     │
    │  ┌───────────────────────────────┐  │
    │  │  Redis Cluster / Sentinel     │  │
    │  └───────────────────────────────┘  │
    └─────────────────────────────────────┘
    
  2. 性能优化技巧

    • 使用 Pipeline 批量操作
    • 合理设置 TTL 避免内存溢出
    • 使用 Hash 而非 String 存储结构化数据
    • 启用 Redis 压缩(LZF/ZSTD)
  3. 生产环境注意事项

    • 配置 Redis 持久化(RDB + AOF)
    • 设置内存上限和淘汰策略
    • 监控内存使用和命中率
    • 使用连接池避免连接泄漏
  4. 安全考虑

    • 启用 Redis AUTH
    • 使用 TLS 加密传输
    • 网络隔离(VPC/安全组)
    • 定期备份关键数据

技术选型决策树

需要状态持久化?
├── 是 → LangGraph + Redis Checkpointer
└── 否 → 继续判断

需要记忆管理?
├── 是 → RedisChatMessageHistory
└── 否 → 继续判断

需要缓存?
├── 是 → RedisCache / 自定义缓存
└── 否 → 继续判断

需要向量搜索?
├── 是 → Redis Stack (RediSearch)
└── 否 → 继续判断

需要任务队列?
├── 是 → Redis Streams
└── 否 → 基础 KV 存储即可

结论

Redis 是构建生产级 LangChain/LangGraph Agent 系统的基础设施级组件,它解决了 Agent 开发中最关键的几个问题:

  • 状态一致性:确保分布式环境下的状态同步
  • 性能瓶颈:通过缓存大幅降低延迟和成本
  • 可靠性:提供持久化和故障恢复能力
  • 可扩展性:支持水平扩展和负载均衡

在实际项目中,建议从会话状态持久化LLM 响应缓存这两个最高价值的场景开始引入 Redis,然后逐步扩展到其他场景。随着 Agent 复杂度的提升,Redis 的价值会越来越明显。

Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐