5.4 Wrap-style hooks函数用法

支持两种用法

装饰器是函数式挂载,把一个hook快速挂载到Agent的某个节点。

类写法是对象化中间件,把中间件封装为一个可配置、可复用、可扩展的组件。

1. wrap_model_call

① 基于装饰器实现

我们可以同时在模型调用前后做事,所以命名为 wrap_model_call ,wrap意为 包裹 。

from typing import Callable

from langchain.agents import create_agent
from langchain.agents.middleware import wrap_model_call
from langchain.agents.middleware import ModelRequest, ModelResponse
from langchain_core.messages import HumanMessage


@wrap_model_call
def wrap_model_call_middleware(
    request: ModelRequest,
    handle: Callable[[ModelRequest], ModelResponse]
) -> ModelResponse:

    # 模型调用之前
    request.messages[-1].content += "---> wrap_model_call_before <----"

    # 调用模型
    response = handle(request)

    # 模型调用之后
    response.result[0].content += "---> wrap_model_call_middleware after <----"

    return response


agent = create_agent(
    model=model,
    middleware=[
        wrap_model_call_middleware
    ]
)

response = agent.invoke({
    "messages": [
        HumanMessage(content="你好")
    ]
})

for msg in response["messages"]:
    msg.pretty_print()

模型调用前消息列表的最后一条是HumanMessage,调用后最后一条是AIMessage,可以看到,模型调 用前后的更改都生效了。

② 基于类实现

from langchain.agents import create_agent
from langchain.agents.middleware import (
    AgentMiddleware,
    ModelRequest,
    ModelResponse,
)
from langchain_core.messages import HumanMessage


class WrapModelCallMiddleware(AgentMiddleware):

    def wrap_model_call(
        self,
        request: ModelRequest,
        handler
    ) -> ModelResponse:

        # 模型调用之前
        request.messages[-1].content += "---> wrap_model_call_before <----"

        # 调用模型
        response = handler(request)

        # 模型调用之后
        response.result[0].content += "---> wrap_model_call_middleware after <----"

        return response


agent = create_agent(
    model=model,
    middleware=[
        WrapModelCallMiddleware()
    ]
)

response = agent.invoke({
    "messages": [
        HumanMessage(content="你好")
    ]
})

for msg in response["messages"]:
    msg.pretty_print()

使用场景:用于拦截、重试、缓存模型调用。

场景1:重试逻辑

from typing import Callable

import time

from langchain.agents import create_agent
from langchain.agents.middleware import (
    ModelRequest,
    ModelResponse,
    wrap_model_call,
)
from langchain.messages import HumanMessage


@wrap_model_call
def retry_model(
    request: ModelRequest,
    handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
    """
    自动重试失败的模型调用
    """
    max_retries = 3

    for attempt in range(max_retries):
        try:
            print(
                f"🔄 尝试调用模型"
                f"(第 {attempt + 1}/{max_retries} 次)"
            )

            return handler(request)

        except Exception as e:
            # 最后一次重试仍然失败
            if attempt == max_retries - 1:
                print(f"❌ 所有重试失败:{e}")
                raise

            # 指数退避
            wait_time = 2 ** attempt

            print(
                f"⚠ 调用失败:{e},"
                f"{wait_time} 秒后重试"
            )

            time.sleep(wait_time)


agent = create_agent(
    model=model,
    middleware=[
        retry_model,
    ],
)

response = agent.invoke(
    {
        "messages": [
            HumanMessage(content="你好")
        ]
    }
)

for msg in response["messages"]:
    msg.pretty_print()

场景2:响应缓存

import hashlib

from typing import Callable

from langchain.agents import create_agent
from langchain.agents.middleware import (
    ModelRequest,
    ModelResponse,
    wrap_model_call,
)


class ModelCache:
    """模型响应缓存"""

    def __init__(self):
        self.cache = {}

    def create_hook(self):
        @wrap_model_call
        def cache_model(
            request: ModelRequest,
            handler: Callable[[ModelRequest], ModelResponse],
        ) -> ModelResponse:

            # 获取最后一条消息
            last_message = request.messages[-1]

            # 使用最后一条消息生成缓存 Key
            cache_key = hashlib.md5(
                str(last_message.content).encode("utf-8")
            ).hexdigest()

            print("cache_key:", cache_key)

            # 查询缓存
            if cache_key in self.cache:
                print("💾 缓存命中!")
                return self.cache[cache_key]

            # 缓存未命中
            print("🔍 缓存未命中,调用模型")

            response = handler(request)

            # 保存模型响应
            self.cache[cache_key] = response

            return response

        return cache_model


cache = ModelCache()

agent = create_agent(
    model=model,
    middleware=[
        cache.create_hook(),
    ],
)


# 第一次调用
response1 = agent.invoke(
    {
        "messages": [
            HumanMessage(content="1+1")
        ]
    }
)

# 第二次调用
response2 = agent.invoke(
    {
        "messages": [
            HumanMessage(content="1+1")
        ]
    }
)

场景3:修改系统提示

from datetime import datetime
from typing import Callable

from langchain.agents import create_agent
from langchain.agents.middleware import (
    wrap_model_call,
    ModelRequest,
    ModelResponse,
)
from langchain_core.messages import (
    HumanMessage,
    SystemMessage,
)


@wrap_model_call
def add_context(
    request: ModelRequest,
    handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
    """动态添加上下文信息到系统提示"""

    # 获取当前时间
    current_time = datetime.now().strftime("%Y-%m-%d %H:%M:%S")

    # 获取原来的系统提示
    original_content = (
        request.system_message.content
        if request.system_message
        else ""
    )

    # 构建新的系统提示
    new_content = f"""
{original_content}
当前时间:{current_time}
用户位置:中国
语言偏好:中文
"""

    # 创建新的系统消息
    new_system_message = SystemMessage(
        content=new_content
    )

    # 修改请求
    modified_request = request.override(
        system_message=new_system_message
    )

    # 调用模型
    return handler(modified_request)


# 创建 Agent
agent = create_agent(
    model=model,
    middleware=[
        add_context,
    ],
)


# 调用 Agent
response = agent.invoke({
    "messages": [
        HumanMessage(content="你好")
    ]
})


# 输出消息
for msg in response["messages"]:
    msg.pretty_print()

2. wrap_tool_call

我们可以同时在工具调用前后做事,所以命名为 wrap_tool_call 。

① 基于装饰器实现

from typing import Callable

from langchain.agents import create_agent
from langchain.agents.middleware import wrap_tool_call
from langchain.messages import HumanMessage, ToolMessage
from langchain.tools import tool
from langchain.tools.tool_node import ToolCallRequest
from langgraph.types import Command


@tool
def get_weather(city: str, is_forcast: bool) -> str:
    """
    获取当日特定城市的天气

    Args:
        city: 城市名称
        is_forcast: 是否包含明天的天气预报
    """
    res = f"{city}今天天气不错"

    if is_forcast:
        res += "\n明天天气也很好"

    return res


@wrap_tool_call
def wrap_tool_call_middleware(
    request: ToolCallRequest,
    handler: Callable[[ToolCallRequest], ToolMessage | Command],
) -> ToolMessage | Command:
    # 第一次:使用模型原始参数执行工具
    result = handler(request)

    print(f"原始参数:{request.tool_call['args']}")
    print(f"原始参数调用结果:{result}")

    # 修改工具参数
    request.tool_call["args"]["is_forcast"] = True

    # 第二次:使用修改后的参数执行工具
    result = handler(request)

    print(f"更新后的参数:{request.tool_call['args']}")
    print(f"更新参数调用结果:{result}")

    return result


agent = create_agent(
    model=model,
    tools=[get_weather],
    middleware=[wrap_tool_call_middleware],
)

response = agent.invoke(
    {
        "messages": [
            HumanMessage("你好啊,今天杭州的天气怎么样")
        ]
    }
)

for msg in response["messages"]:
    msg.pretty_print()

在@wrap_tool_call装饰的函数中两次调用函数并更改参数。

② 基于类实现

from typing import Callable

from langchain.agents import create_agent
from langchain.agents.middleware import AgentMiddleware
from langchain.messages import HumanMessage, ToolMessage
from langchain.tools import tool
from langchain.tools.tool_node import ToolCallRequest
from langgraph.types import Command


@tool
def get_weather(city: str, is_forcast: bool) -> str:
    """
    获取当日特定城市的天气

    Args:
        city: 城市名称
        is_forcast: 是否包含明天的天气预报
    """
    res = f"{city}今天天气不错"

    if is_forcast:
        res += "\n明天天气也很好"

    return res


class WrapToolCallMiddleware(AgentMiddleware):

    def wrap_tool_call(
        self,
        request: ToolCallRequest,
        handler: Callable[[ToolCallRequest], ToolMessage | Command],
    ) -> ToolMessage | Command:

        # 原始参数
        print(f"原始参数:{request.tool_call['args']}")

        # 修改工具调用参数
        request.tool_call["args"]["is_forcast"] = True

        print(f"更新后的参数:{request.tool_call['args']}")

        # 执行工具
        result = handler(request)

        print(f"调用结果:{result}")

        return result


agent = create_agent(
    model=model,
    tools=[get_weather],
    middleware=[WrapToolCallMiddleware()],
)

response = agent.invoke(
    {
        "messages": [
            HumanMessage("你好啊,今天杭州的天气怎么样")
        ],
    }
)

for msg in response["messages"]:
    msg.pretty_print()

使用场景:用于监控、重试、修改工具执行。

import time
from typing import Callable

from langchain.agents import create_agent
from langchain.agents.middleware import wrap_tool_call
from langchain.tools import tool
from langchain.tools.tool_node import ToolCallRequest
from langchain_core.messages import HumanMessage, ToolMessage
from langgraph.types import Command


@tool
def get_weather(city: str) -> str:
    """获取指定城市的天气"""
    time.sleep(1)

    return f"{city}今天晴天,温度 25℃"


@wrap_tool_call
def monitor_tool(
    request: ToolCallRequest,
    handler: Callable[
        [ToolCallRequest],
        ToolMessage | Command,
    ],
) -> ToolMessage | Command:
    """监控工具执行时间和状态"""

    # 获取工具名称
    tool_name = request.tool_call["name"]

    # 获取工具参数
    tool_args = request.tool_call.get("args", {})

    print(f"🔧 开始执行工具:{tool_name}")
    print(f"   参数:{tool_args}")

    # 开始计时
    start_time = time.time()

    try:
        # 执行工具
        result = handler(request)

        # 计算执行时间
        elapsed = time.time() - start_time

        print(f"✅ 工具执行成功,耗时:{elapsed:.2f} 秒")

        return result

    except Exception as e:
        # 计算执行时间
        elapsed = time.time() - start_time

        print(f"❌ 工具执行失败,耗时:{elapsed:.2f} 秒")
        print(f"   错误信息:{e}")

        raise


# 创建 Agent
agent = create_agent(
    model=model,
    tools=[
        get_weather,
    ],
    middleware=[
        monitor_tool,
    ],
)


# 调用 Agent
response = agent.invoke({
    "messages": [
        HumanMessage(content="北京今天天气怎么样?")
    ]
})


# 输出结果
for message in response["messages"]:
    message.pretty_print()

两种方法的统一

同上,装饰器方法底层也会创建一个AgentMiddleware的实例。

参数说明

request:被封装的请求对象,可以是模型或工具调用请

handler:处理器,用于处理请求并返回调用结果。

Logo

一座年轻的奋斗人之城,一个温馨的开发者之家。在这里,代码改变人生,开发创造未来!

更多推荐