1. Agent模型调用的拦截需求与Middleware解决方案
在开发AI Agent时,我们经常遇到这样的场景:Agent运行过程中需要插入额外的处理逻辑,比如记录日志、安全检查、性能监控等。传统做法是直接修改Agent核心代码,但这会导致代码臃肿且难以维护。Middleware模式提供了一种优雅的解决方案,它就像给Agent安装了一个"拦截器",可以在不修改核心逻辑的情况下,灵活地插入各种处理逻辑。
Middleware的核心思想是"面向切面编程"(AOP),它允许我们在模型调用的关键节点(调用前和调用后)插入自定义逻辑。这种设计模式在Web开发中很常见(如Express.js的中间件),现在也被广泛应用于AI Agent开发领域。
2. Middleware的核心机制与实现原理
2.1 Middleware的基本结构
一个标准的Middleware通常需要实现两个关键方法:
from langchain.agents.middleware import AgentMiddleware from langchain.agents import AgentState from langgraph.runtime import Runtime class CustomMiddleware(AgentMiddleware): def before_model(self, state: AgentState, runtime: Runtime) -> dict | None: """模型调用前执行的逻辑""" # 可以访问和修改state中的消息、工具等 # 返回None表示继续流程,返回dict可以改变流程走向 pass def after_model(self, state: AgentState, runtime: Runtime) -> None: """模型调用后执行的逻辑""" # 可以访问模型响应结果,但不能改变流程 pass2.2 执行流程控制
Middleware最强大的特性在于它可以控制执行流程:
- 被动观察:返回None表示不干预流程,仅执行附加逻辑(如日志记录)
- 主动干预:返回dict可以改变流程,比如:
- 跳过模型调用直接返回结果
- 修改输入消息内容
- 终止当前会话
{ "jump_to": "end", # 跳过模型调用直接结束 "messages": [AIMessage(content="自定义响应")] # 添加的消息 }3. 实战:构建实用的Middleware组件
3.1 日志记录Middleware
一个完整的日志记录Middleware应该包含以下功能:
import logging from datetime import datetime class EnhancedLoggingMiddleware(AgentMiddleware): def __init__(self, log_level=logging.INFO): self.logger = logging.getLogger("AgentLogger") self.logger.setLevel(log_level) def before_model(self, state: AgentState, runtime: Runtime) -> None: context = { "timestamp": datetime.now().isoformat(), "message_count": len(state['messages']), "last_user_input": next( (msg.content for msg in reversed(state['messages']) if msg.type == "human"), None) } self.logger.info(f"Pre-model call: {context}") def after_model(self, state: AgentState, runtime: Runtime) -> None: response = state['messages'][-1].content metrics = { "response_length": len(response), "response_time": runtime.get("model_time", 0) } self.logger.info(f"Post-model call: {metrics}")提示:在生产环境中,建议将日志输出到文件或日志系统,而非直接打印到控制台。
3.2 安全拦截Middleware
增强版的安全Middleware可以支持:
- 关键词黑名单
- 正则表达式模式匹配
- 敏感操作检测
import re from typing import List class AdvancedSafetyMiddleware(AgentMiddleware): def __init__(self, blacklist: List[str] = None, dangerous_patterns: List[str] = None): self.blacklist = blacklist or ["删除", "危险", "密码", "root"] self.patterns = [re.compile(p) for p in dangerous_patterns or []] def before_model(self, state: AgentState, runtime: Runtime) -> dict | None: last_msg = state['messages'][-1].content # 检查黑名单关键词 if any(keyword in last_msg for keyword in self.blacklist): return self._block_action("检测到禁用关键词") # 检查危险模式 if any(pattern.search(last_msg) for pattern in self.patterns): return self._block_action("检测到危险操作模式") return None def _block_action(self, reason: str) -> dict: return { "jump_to": "end", "messages": [AIMessage( content=f"{reason},操作已终止。如需帮助请联系管理员。" )] }4. Middleware的高级应用场景
4.1 上下文增强Middleware
可以在模型调用前注入相关上下文信息:
class ContextEnhancementMiddleware(AgentMiddleware): def __init__(self, knowledge_base): self.knowledge = knowledge_base def before_model(self, state: AgentState, runtime: Runtime) -> None: user_query = state['messages'][-1].content related_info = self.knowledge.search(user_query) if related_info: state['context'] = related_info # 注入上下文4.2 限流Middleware
控制模型调用频率,防止滥用:
from collections import deque import time class RateLimitMiddleware(AgentMiddleware): def __init__(self, max_calls=5, period=60): self.max_calls = max_calls self.period = period self.call_times = deque() def before_model(self, state: AgentState, runtime: Runtime) -> dict | None: now = time.time() # 移除过期的调用记录 while self.call_times and now - self.call_times[0] > self.period: self.call_times.popleft() if len(self.call_times) >= self.max_calls: return { "jump_to": "end", "messages": [AIMessage( content="请求过于频繁,请稍后再试" )] } self.call_times.append(now) return None4.3 缓存Middleware
对重复请求返回缓存结果:
import hashlib class CacheMiddleware(AgentMiddleware): def __init__(self, cache_size=100): self.cache = {} self.cache_size = cache_size def before_model(self, state: AgentState, runtime: Runtime) -> dict | None: query = state['messages'][-1].content query_hash = hashlib.md5(query.encode()).hexdigest() if query_hash in self.cache: return { "jump_to": "end", "messages": [AIMessage( content=self.cache[query_hash] )] } return None def after_model(self, state: AgentState, runtime: Runtime) -> None: if len(self.cache) >= self.cache_size: self.cache.popitem() # 简单LRU策略 query = state['messages'][-2].content # 用户的上一条消息 response = state['messages'][-1].content query_hash = hashlib.md5(query.encode()).hexdigest() self.cache[query_hash] = response5. Middleware的链式调用与执行顺序
当多个Middleware组合使用时,它们的执行顺序很重要:
- before_model调用顺序:按照Middleware列表顺序依次执行
- after_model调用顺序:与before_model相反(栈式结构)
- 流程中断:任一Middleware返回非None值都会中断后续Middleware执行
# 推荐的Middleware顺序: middlewares = [ RateLimitMiddleware(), # 最先执行限流检查 SafetyMiddleware(), # 然后安全检查 LoggingMiddleware(), # 记录原始请求 ContextEnhancementMiddleware(kb), # 上下文增强 CacheMiddleware() # 最后检查缓存 ]6. 性能优化与调试技巧
6.1 Middleware性能监控
可以添加专门的性能监控Middleware:
class PerformanceMonitorMiddleware(AgentMiddleware): def before_model(self, state: AgentState, runtime: Runtime) -> None: runtime['start_time'] = time.time() def after_model(self, state: AgentState, runtime: Runtime) -> None: duration = time.time() - runtime['start_time'] print(f"模型调用耗时: {duration:.3f}秒") if duration > 1.0: # 慢请求警告 print(f"慢请求警告: {state['messages'][-2].content[:50]}...")6.2 调试技巧
- 隔离测试:逐个启用Middleware,确认各自功能正常
- 状态检查:在before_model中打印state完整内容
- 错误处理:Middleware内部应该捕获自己的异常,避免影响主流程
class SafeMiddleware(AgentMiddleware): def before_model(self, state: AgentState, runtime: Runtime) -> dict | None: try: # 业务逻辑 return None except Exception as e: print(f"Middleware错误: {str(e)}") return None # 即使出错也不中断流程7. 生产环境最佳实践
- 配置化:通过配置文件管理Middleware开关和参数
- 依赖注入:避免Middleware直接实例化外部依赖
- 单元测试:为每个Middleware编写独立测试用例
- 性能考量:IO密集型操作(如网络请求)应该异步化
# 配置示例 MIDDLEWARE_CONFIG = { "safety": { "enable": True, "blacklist": ["删除", "格式化", "关机"], "patterns": [r"rm -rf", r"DROP TABLE"] }, "logging": { "enable": True, "level": "INFO" } } # 根据配置动态创建Middleware链 def setup_middlewares(config): middlewares = [] if config["safety"]["enable"]: middlewares.append(SafetyMiddleware( blacklist=config["safety"]["blacklist"], patterns=config["safety"]["patterns"] )) if config["logging"]["enable"]: middlewares.append(LoggingMiddleware( log_level=config["logging"]["level"] )) return middlewares8. 常见问题与解决方案
8.1 Middleware执行顺序问题
问题:多个Middleware相互影响,顺序不当导致功能异常
解决方案:
- 按照"安全→日志→业务→缓存"的通用顺序排列
- 为Middleware添加优先级属性,动态排序
8.2 状态污染问题
问题:Middleware意外修改了state导致后续流程异常
解决方案:
- 在修改state前创建深拷贝
- 使用不可变数据结构
import copy class SafeStateMiddleware(AgentMiddleware): def before_model(self, state: AgentState, runtime: Runtime) -> dict | None: original_state = copy.deepcopy(state) # 保存原始状态 try: # 修改state return None except Exception: return {"restore_state": original_state} # 出错时恢复状态8.3 性能瓶颈问题
问题:Middleware引入过多计算或IO导致延迟增加
解决方案:
- 异步化处理
- 采样记录而非全量记录
- 使用轻量级检查(如布隆过滤器)
import asyncio class AsyncLoggingMiddleware(AgentMiddleware): async def _async_log(self, message): # 异步写入日志系统 pass def after_model(self, state: AgentState, runtime: Runtime) -> None: loop = asyncio.get_event_loop() message = state['messages'][-1].content loop.create_task(self._async_log(message)) # 异步记录9. 扩展思考:Middleware设计模式的应用
Middleware模式不仅适用于模型调用拦截,还可以应用于:
- 工具调用拦截:在工具执行前后插入逻辑
- 消息处理管道:对输入/输出消息进行统一处理
- Agent生命周期钩子:在Agent启动/停止时执行操作
class ToolMiddleware: def before_tool(self, tool_name: str, input: dict) -> dict | None: """工具调用前执行""" pass def after_tool(self, tool_name: str, output: str) -> str | None: """工具调用后执行""" pass class LifecycleMiddleware: def on_agent_start(self): """Agent启动时执行""" pass def on_agent_stop(self): """Agent停止时执行""" passMiddleware模式为Agent开发提供了极大的灵活性和可扩展性。通过合理设计和组合Middleware,可以实现各种横切关注点,而无需修改核心业务逻辑。这种设计模式遵循了开闭原则(对扩展开放,对修改关闭),是构建可维护、可扩展AI系统的重要实践。