Skip to content

Latest commit

 

History

History
786 lines (621 loc) · 21.5 KB

File metadata and controls

786 lines (621 loc) · 21.5 KB

教程 02:Tool Calling - 工具调用

学习目标: 让 LLM 能够调用真实的业务函数,实现 Agent 的核心能力

📖 本章内容

  • Tool Calling / Function Calling 原理
  • 工具定义与参数校验
  • 安全控制与人工确认
  • 实战:客服 Agent(多工具协作)

🎯 核心概念

什么是 Tool Calling?

Tool Calling(工具调用)是 LLM 的核心能力之一,让 LLM 能够:

  1. 识别需求 - 理解用户何时需要调用工具
  2. 选择工具 - 从多个工具中选择合适的
  3. 生成参数 - 根据上下文生成正确的参数
  4. 处理结果 - 基于工具返回结果给出最终答案

Java vs Python 对比

特性 Java Python
工具定义 Interface + 实现类 函数 + 装饰器
参数校验 Bean Validation Pydantic
函数引用 Method Reference 函数对象
反射调用 Method.invoke() func(**kwargs)
异步执行 CompletableFuture asyncio

工作流程

用户提问
   ↓
LLM 分析(是否需要工具?)
   ↓
[需要] → LLM 生成工具调用请求
   ↓         - function_name
   ↓         - arguments (JSON)
   ↓
执行工具函数
   ↓
返回结果给 LLM
   ↓
LLM 基于结果生成最终答案
   ↓
返回给用户

📝 Part 1:基础工具调用

1.1 定义工具函数

def get_weather(city: str, unit: str = "celsius") -> str:
    """
    获取指定城市的天气信息
    
    Args:
        city: 城市名称
        unit: 温度单位,celsius(摄氏度)或 fahrenheit(华氏度)
    
    Returns:
        天气描述字符串
    """
    # 模拟天气 API 调用
    weather_data = {
        "北京": {"temp": 25, "condition": "晴朗"},
        "上海": {"temp": 28, "condition": "多云"},
    }
    
    data = weather_data.get(city, {"temp": 20, "condition": "未知"})
    temp = data["temp"]
    
    if unit == "fahrenheit":
        temp = temp * 9/5 + 32
        unit_symbol = "°F"
    else:
        unit_symbol = "°C"
    
    return f"{city}的天气:{data['condition']},温度 {temp}{unit_symbol}"

重点:

  1. 类型注解 - city: str 会被解析为参数类型
  2. Docstring - 会作为工具描述传递给 LLM
  3. 默认值 - unit: str = "celsius" 表示可选参数

1.2 转换为 OpenAI Tools 格式

tools = [
    {
        "type": "function",
        "function": {
            "name": "get_weather",  # 函数名
            "description": "获取指定城市的天气信息",  # 描述(帮助 LLM 选择)
            "parameters": {  # JSON Schema 格式
                "type": "object",
                "properties": {
                    "city": {
                        "type": "string",
                        "description": "城市名称,例如:北京、上海"
                    },
                    "unit": {
                        "type": "string",
                        "enum": ["celsius", "fahrenheit"],
                        "description": "温度单位"
                    }
                },
                "required": ["city"]  # city 必填,unit 可选
            }
        }
    }
]

Java 对比:

// Java 需要手动定义接口和实现
@FunctionalInterface
public interface WeatherTool {
    String execute(String city, String unit);
}

// 使用注解描述
@ToolDefinition(
    name = "get_weather",
    description = "获取指定城市的天气信息"
)
public class WeatherToolImpl implements WeatherTool {
    @Override
    public String execute(
        @NotBlank String city, 
        String unit
    ) {
        // ...
    }
}

1.3 Agent 执行循环

def run_agent(user_query: str, max_iterations: int = 5) -> str:
    """执行 Agent 循环"""
    messages = [{"role": "user", "content": user_query}]
    
    for i in range(max_iterations):
        # 1. 调用 LLM
        response = client.chat.completions.create(
            model="gpt-4o-mini",
            messages=messages,
            tools=tools,  # 传递工具列表
            tool_choice="auto"  # auto | none | required
        )
        
        response_message = response.choices[0].message
        messages.append(response_message)  # 加入历史
        
        # 2. 检查是否有工具调用
        tool_calls = response_message.tool_calls
        if not tool_calls:
            # 没有工具调用,返回最终答案
            return response_message.content
        
        # 3. 执行工具函数
        for tool_call in tool_calls:
            function_name = tool_call.function.name
            function_args = eval(tool_call.function.arguments)  # JSON → dict
            
            # 调用函数
            function_to_call = available_functions[function_name]
            function_response = function_to_call(**function_args)  # ** 解包参数
            
            # 4. 将结果加入历史
            messages.append({
                "role": "tool",
                "tool_call_id": tool_call.id,
                "name": function_name,
                "content": function_response
            })
    
    return "达到最大迭代次数"

关键点:

  1. 迭代循环 - Agent 可能需要多次调用工具
  2. 消息历史 - 每次交互都加入 messages
  3. 工具响应 - 必须包含 tool_call_id 关联请求

📝 Part 2:参数校验

2.1 为什么需要参数校验?

  1. LLM 可能生成不合规参数 - 格式错误、超出范围
  2. 防止安全问题 - SQL 注入、命令注入
  3. 提供清晰错误信息 - 帮助 LLM 重试

2.2 使用 Pydantic 进行校验

from pydantic import BaseModel, Field, field_validator

class QueryOrderParams(BaseModel):
    """订单查询参数"""
    order_id: str = Field(
        ...,  # 必填
        description="订单号",
        pattern=r"^ORD\d{10}$",  # 正则:ORD + 10位数字
        examples=["ORD1234567890"]
    )
    
    @field_validator("order_id")
    @classmethod
    def validate_order_id(cls, v: str) -> str:
        """自定义校验器"""
        # 防止 SQL 注入
        if "'" in v or '"' in v or ';' in v:
            raise ValueError("订单号包含非法字符")
        return v

Java 对比:

public class QueryOrderParams {
    @NotBlank
    @Pattern(regexp = "^ORD\\d{10}$")
    private String orderId;
    
    @AssertTrue(message = "订单号包含非法字符")
    public boolean isOrderIdValid() {
        return !orderId.contains("'") 
            && !orderId.contains("\"");
    }
}

2.3 在工具函数中使用校验

def query_order_validated(order_id: str) -> str:
    """查询订单(参数已校验)"""
    try:
        # 使用 Pydantic 校验
        params = QueryOrderParams(order_id=order_id)
        
        # 查询数据库
        order = orders_db.get(params.order_id)
        if not order:
            return f"❌ 订单 {params.order_id} 不存在"
        
        return f"订单 {params.order_id}:状态 {order.status},金额 ¥{order.amount}"
        
    except Exception as e:
        return f"❌ 参数错误:{str(e)}"

2.4 直接使用 Pydantic Schema

tools = [
    {
        "type": "function",
        "function": {
            "name": "query_order_validated",
            "description": "查询订单详情",
            # 直接使用 Pydantic 的 JSON Schema
            "parameters": QueryOrderParams.model_json_schema()
        }
    }
]

📝 Part 3:安全控制与人工确认

3.1 操作风险等级

from enum import Enum

class RiskLevel(str, Enum):
    LOW = "low"       # 查询操作
    MEDIUM = "medium" # 修改操作
    HIGH = "high"     # 删除、转账、发布

3.2 人工确认装饰器

def require_approval(risk_level: RiskLevel):
    """人工确认装饰器"""
    def decorator(func):
        def wrapper(*args, **kwargs):
            # 构建确认提示
            print(f"\n⚠️ 需要人工确认 [{risk_level.value.upper()} RISK]")
            print(f"操作:{func.__name__}")
            print(f"参数:{kwargs}")
            
            # 请求用户确认
            user_input = input("是否执行?(yes/no): ").strip().lower()
            if user_input in ["yes", "y"]:
                return func(*args, **kwargs)
            else:
                return f"操作被用户拒绝:{func.__name__}"
        
        return wrapper
    return decorator

Java 对比:

@RequireApproval(level = RiskLevel.HIGH)
public String transferMoney(String from, String to, double amount) {
    // ...
}

// 通过 AOP 拦截
@Aspect
public class ApprovalAspect {
    @Around("@annotation(requireApproval)")
    public Object requireApproval(ProceedingJoinPoint pjp, 
                                  RequireApproval requireApproval) {
        // 弹出确认对话框
        if (showConfirmDialog(pjp)) {
            return pjp.proceed();
        }
        return "操作被拒绝";
    }
}

3.3 使用装饰器

@require_approval(RiskLevel.HIGH)
def transfer_money(from_account: str, to_account: str, amount: float) -> str:
    """转账 - 高风险操作"""
    # 执行转账
    return f"✅ 转账成功:{from_account}{to_account},金额 ¥{amount}"

@require_approval(RiskLevel.MEDIUM)
def send_email(to: str, subject: str, body: str) -> str:
    """发送邮件 - 中风险操作"""
    # 发送邮件
    return f"✅ 邮件已发送到 {to}"

def query_balance(account: str) -> str:
    """查询余额 - 低风险,无需确认"""
    return f"账户 {account} 余额:¥10000"

3.4 生产环境的人工确认

实际生产环境中,人工确认可能通过:

  1. Web 界面 - 用户在 UI 上点击确认按钮
  2. 短信验证码 - 发送验证码到用户手机
  3. 工作流审批 - 提交到审批系统
  4. 双因素认证 - 指纹、人脸识别
# 异步确认示例(伪代码)
async def require_approval_async(operation_id: str) -> bool:
    """异步等待用户确认"""
    # 1. 生成确认 token
    token = generate_token(operation_id)
    
    # 2. 发送通知(短信/邮件)
    await send_notification(user_phone, token)
    
    # 3. 等待用户确认(轮询或 WebSocket)
    result = await wait_for_confirmation(operation_id, timeout=300)  # 5分钟
    
    return result

📝 Part 4:实战 - 客服 Agent

4.1 场景设计

功能需求:

  • 查询订单状态
  • 查询物流信息
  • 申请退款
  • 查询退款政策
  • 知识库问答(退货、发票、会员等)

技术要点:

  • 多工具协作
  • 条件判断(是否超过退款期限)
  • 知识库检索(简化版)
  • 友好的用户体验

4.2 数据模型

class Order:
    """订单模型"""
    def __init__(self, order_id: str, status: str, amount: float, 
                 items: List[str], create_time: str):
        self.order_id = order_id
        self.status = status
        self.amount = amount
        self.items = items
        self.create_time = create_time

class Logistics:
    """物流信息模型"""
    def __init__(self, order_id: str, carrier: str, tracking_no: str,
                 status: str, location: str):
        self.order_id = order_id
        self.carrier = carrier
        self.tracking_no = tracking_no
        self.status = status
        self.location = location

4.3 核心工具函数

def query_order(order_id: str) -> str:
    """查询订单详情"""
    order = ORDERS_DB.get(order_id)
    if not order:
        return f"❌ 订单 {order_id} 不存在"
    
    return (
        f"📦 订单详情:\n"
        f"  订单号:{order.order_id}\n"
        f"  状态:{order.status}\n"
        f"  金额:¥{order.amount:.2f}\n"
        f"  商品:{', '.join(order.items)}\n"
        f"  下单时间:{order.create_time}"
    )

def query_logistics(order_id: str) -> str:
    """查询物流信息"""
    logistics = LOGISTICS_DB.get(order_id)
    if not logistics:
        return "❌ 暂无物流信息"
    
    return (
        f"🚚 物流详情:\n"
        f"  快递公司:{logistics.carrier}\n"
        f"  运单号:{logistics.tracking_no}\n"
        f"  状态:{logistics.status}\n"
        f"  位置:{logistics.location}"
    )

def apply_refund(order_id: str, reason: str) -> str:
    """申请退款"""
    order = ORDERS_DB.get(order_id)
    if not order:
        return f"❌ 订单 {order_id} 不存在"
    
    # 检查是否超过7天
    order_time = datetime.strptime(order.create_time, "%Y-%m-%d %H:%M:%S")
    days_passed = (datetime.now() - order_time).days
    
    if days_passed > 7:
        return f"❌ 订单已超过7天无理由退货期限(已过{days_passed}天)"
    
    return (
        f"✅ 退款申请已提交:\n"
        f"  订单号:{order.order_id}\n"
        f"  退款金额:¥{order.amount:.2f}\n"
        f"  退款原因:{reason}\n"
        f"  预计到账时间:3-5个工作日"
    )

def search_product_info(keyword: str) -> str:
    """查询商品相关信息(知识库)"""
    PRODUCT_KB = {
        "退货政策": "7天无理由退货,商品需保持原包装完整",
        "退款时间": "退货签收后3-5个工作日原路退回",
        "发票": "所有订单默认开具电子发票",
        "会员": "消费满1000元自动升级为银卡会员",
    }
    
    results = [f"【{k}{v}" for k, v in PRODUCT_KB.items() 
               if keyword in k or keyword in v]
    
    if not results:
        return f"❌ 未找到关于 '{keyword}' 的信息"
    
    return "📚 相关信息:\n\n" + "\n\n".join(results)

4.4 系统提示词(System Prompt)

system_prompt = """你是一个专业的电商客服 Agent。

你的职责:
1. 热情友好地回答用户问题
2. 合理使用工具查询订单、物流、退款等信息
3. 根据退款政策帮助用户判断是否可以退款
4. 对于不确定的问题,查询知识库或建议联系人工客服

注意事项:
- 订单号格式为 ORD + 10位数字
- 退款需要在7天内申请
- 物流信息只有在订单发货后才能查询
- 态度要友好、专业、耐心
"""

系统提示词的重要性:

  1. 定义 Agent 的角色和行为
  2. 提供业务规则和约束
  3. 引导 LLM 正确使用工具
  4. 设定回复的语气和风格

4.5 完整执行流程

def run_customer_service_agent(user_query: str) -> str:
    """客服 Agent"""
    messages = [
        {"role": "system", "content": system_prompt},  # 系统提示词
        {"role": "user", "content": user_query}
    ]
    
    for i in range(10):  # 最多10次迭代
        response = client.chat.completions.create(
            model="gpt-4o-mini",
            messages=messages,
            tools=tools,
            tool_choice="auto"
        )
        
        response_message = response.choices[0].message
        messages.append(response_message)
        
        tool_calls = response_message.tool_calls
        if not tool_calls:
            return response_message.content
        
        # 执行工具
        for tool_call in tool_calls:
            function_name = tool_call.function.name
            function_args = eval(tool_call.function.arguments)
            
            function_to_call = available_functions[function_name]
            function_response = function_to_call(**function_args)
            
            messages.append({
                "role": "tool",
                "tool_call_id": tool_call.id,
                "name": function_name,
                "content": function_response
            })
    
    return "达到最大迭代次数"

🎯 最佳实践

1. 工具设计原则

单一职责

# ✅ 好:每个工具只做一件事
def query_order(order_id: str) -> str: ...
def query_logistics(order_id: str) -> str: ...

# ❌ 差:一个工具做太多事
def query_order_and_logistics_and_maybe_refund(order_id: str, ...) -> str: ...

清晰的描述

# ✅ 好:描述清晰,LLM 容易选择
"description": "查询订单详情,包括订单状态、金额、商品等信息"

# ❌ 差:描述模糊
"description": "查询订单"

严格的参数校验

# ✅ 好:使用 Pydantic 校验
class Params(BaseModel):
    order_id: str = Field(..., pattern=r"^ORD\d{10}$")

# ❌ 差:直接信任 LLM 生成的参数
def query_order(order_id: str):
    result = db.query(f"SELECT * FROM orders WHERE id = '{order_id}'")  # SQL 注入风险

2. 错误处理

def safe_tool_call(func, **kwargs):
    """安全的工具调用包装器"""
    try:
        # 1. 参数校验
        # 2. 权限检查
        # 3. 执行函数
        result = func(**kwargs)
        # 4. 记录日志
        return result
    except ValidationError as e:
        return f"❌ 参数错误:{str(e)}"
    except PermissionError:
        return "❌ 权限不足"
    except Exception as e:
        logger.error(f"Tool error: {func.__name__}", exc_info=True)
        return "❌ 系统错误,请稍后重试"

3. Token 成本控制

# ❌ 差:返回大量数据
def search_products(keyword: str) -> str:
    products = db.query(f"SELECT * FROM products WHERE name LIKE '%{keyword}%'")
    return json.dumps(products)  # 可能有数千条记录

# ✅ 好:返回摘要信息
def search_products(keyword: str, limit: int = 10) -> str:
    products = db.query(f"... LIMIT {limit}")
    # 只返回关键字段
    summary = [f"{p.id}: {p.name} - ¥{p.price}" for p in products]
    return "\n".join(summary)

4. 日志与可观测性

import logging

logger = logging.getLogger(__name__)

def query_order(order_id: str) -> str:
    """带日志的工具函数"""
    logger.info(f"[Tool] query_order called: order_id={order_id}")
    
    try:
        order = ORDERS_DB.get(order_id)
        
        if order:
            logger.info(f"[Tool] query_order success: {order_id}")
        else:
            logger.warning(f"[Tool] query_order not found: {order_id}")
        
        return format_order(order)
        
    except Exception as e:
        logger.error(f"[Tool] query_order error: {order_id}", exc_info=True)
        raise

🧪 练习题

练习 1:计算器工具(入门)

实现一个简单的计算器 Agent:

  • 工具:add, subtract, multiply, divide
  • 要求:能处理"计算 (3 + 5) * 2"这样的复合问题
💡 提示
  1. 定义4个独立的工具函数
  2. LLM 会自动拆解计算步骤
  3. 注意除零错误处理

练习 2:数据库查询 Agent(进阶)

实现一个数据库查询 Agent:

  • 工具:list_tables, describe_table, query_data
  • 要求:用户用自然语言查询,Agent 转换为 SQL
💡 提示
  1. list_tables 返回所有表名
  2. describe_table 返回表结构
  3. query_data 执行 SELECT 语句(禁止 UPDATE/DELETE)
  4. 使用 Pydantic 校验 SQL 语句

练习 3:邮件助手 Agent(高级)

实现一个邮件助手 Agent:

  • 工具:list_emails, read_email, send_email, search_contacts
  • 要求:
    • 发送邮件需要人工确认
    • 自动根据上下文搜索联系人
    • 支持"回复最新邮件"这样的指令
💡 提示
  1. 维护对话上下文("最新邮件"需要记录)
  2. 发送邮件使用 @require_approval 装饰器
  3. 提供详细的工具描述,帮助 LLM 理解"最新"等概念

📚 扩展阅读

OpenAI Function Calling

Anthropic Tool Use

Pydantic 数据校验


🎓 本章小结

核心知识点

  1. ✅ Tool Calling 的工作原理和流程
  2. ✅ 如何定义工具函数和 JSON Schema
  3. ✅ 使用 Pydantic 进行参数校验
  4. ✅ 实现人工确认(Human-in-the-Loop)
  5. ✅ 多工具协作的 Agent 设计

关键技能

  • 能定义清晰、安全的工具函数
  • 能处理 LLM 的工具调用请求
  • 能设计多工具协作的业务流程
  • 能实现安全控制和人工确认

下一步

在教程 03 中,我们将学习 RAG(检索增强生成),让 Agent 能够访问外部知识库,解决"知识过时"和"领域知识不足"的问题。


💬 常见问题

Q: LLM 会不会调用错误的工具?
A: 可能会。解决方法:

  1. 提供清晰的工具描述
  2. 在 System Prompt 中说明每个工具的使用场景
  3. 实现参数校验,拒绝不合法的调用

Q: 如何限制工具调用次数?
A: 设置 max_iterations 参数,防止死循环。

Q: 工具执行时间过长怎么办?
A:

  1. 为工具函数设置超时
  2. 使用异步执行(见教程 07)
  3. 返回"任务已提交",稍后轮询结果

Q: 如何处理工具执行失败?
A: 将错误信息返回给 LLM,LLM 会自动重试或给出解释。

Q: 能否让 LLM 生成并执行代码?
A: 可以,但有巨大安全风险!生产环境需要:

  1. 沙箱环境执行
  2. 代码审计
  3. 资源限制(CPU、内存、网络)
  4. 人工审批

下一课:Tutorial 03 - RAG 检索增强生成 📖