Python/agent/application/workshop/book_flight.py

417 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# -*- coding: utf-8 -*-
"""
预定航班(范式)
"""
import asyncio
import datetime
from typing import AsyncGenerator
from logfire_api.variables import ValueDoesNotEqual
from pydantic import BaseModel, Field, field_validator
from pydantic_ai import (
Agent,
ApprovalRequired,
DeferredToolRequests,
DeferredToolResults,
ModelMessage,
ModelRetry,
ModelSettings,
RunContext,
RunUsage,
UsageLimits,
)
from pydantic_ai.models.openai import OpenAIChatModel
from pydantic_ai.providers.openai import OpenAIProvider
from sqlalchemy.sql.dml import ReturningDelete
DEEPSEEK_V4_FLASH_MODEL = OpenAIChatModel(
model_name="deepseek-v4-flash",
provider=OpenAIProvider(
base_url="https://tokenhub.tencentmaas.com/v1",
api_key="sk-D9Y1mCe8VlvNqLuSC4mAjqEwxJ2nW4C0h8a7EPn8kg9RLsHq",
),
)
MODEL_SETTINGS = ModelSettings(
temperature=0, extra_body={"thinking": {"type": "disabled"}} # 温度控制
) # 禁用思考模式
class Deps(BaseModel):
"""
依赖类
"""
date: datetime.date = Field(..., description="日期")
origin_airport_code: str = Field(..., description="出发机场代码")
destination_airport_code: str = Field(..., description="到达机场代码")
source_material: str = Field(..., description="来源资料")
class Flight(BaseModel):
"""
航班类
"""
number: str = Field(..., description="航班号")
date: datetime.date = Field(..., description="日期")
origin_airport_code: str = Field(..., description="出发机场代码")
destination_airport_code: str = Field(..., description="到达机场代码")
airfare: int = Field(..., description="机票价格")
@field_validator("origin_airport_code", "destination_airport_code")
@classmethod
def validate_airport_code(cls, airport_code: str) -> str:
"""校验机场代码"""
if (
not airport_code.isalpha()
or not airport_code.isupper()
or len(airport_code) != 3
):
raise ValueError(f"机场代码须为英文、大写、3位字符串")
return airport_code
class FlightNoFound(BaseModel):
"""
未查询到航班类
"""
system_prompt = """你的任务是根据用户提供的日期、出发机场和到达机场选择航班。
1. 使用 extract_flights 从来源资料提取所有航班信息;
2. 使用 select_flight 从所有航班信息中选择匹配航班;
"""
# 主智能体
agent = Agent[Deps, Flight | FlightNoFound](
model=DEEPSEEK_V4_FLASH_MODEL,
model_settings=MODEL_SETTINGS,
deps_type=Deps,
output_type=Flight | FlightNoFound,
system_prompt=system_prompt,
)
# 航班信息提取智能体
extraction_agent = Agent(
model=DEEPSEEK_V4_FLASH_MODEL,
model_settings=MODEL_SETTINGS,
output_type=list[Flight],
system_prompt=(
"你的任务是根据给定的来源资料提取出所有航班信息,包括航班号、日期、出发机场代码、到达机场代码和机票价格。",
"其中出发机场代码和到达机场代码须为英文、大写、3位字符串例如 San Francisco International Airport (SFO) 的机场代码为 SFO。",
"若某航班信息不全则跳过该航班,若没有任何航班信息则返回空列表。",
"禁止编造。",
),
) # 通过提示词约束输出
# 航班选择智能体
selection_agent = Agent(
model=DEEPSEEK_V4_FLASH_MODEL,
model_settings=MODEL_SETTINGS,
output_type=Flight | FlightNoFound,
system_prompt="你的任务是从所有航班信息中选择匹配用户提供的日期、出发机场和到达机场的航班:若有多个航班则选择最便宜的航班,选择完成需用户确认是否预定;若没有则返回 FlightNoFound。",
)
"""
如何提高提取准确率
1. 系统提示词规则约束
2. 降低模型温度
3. 输出类型约定和输出模型校验
4. 业务规则约束(例如,输出不可为空列表)
"""
@agent.tool
async def extract_flights(ctx: RunContext[Deps]) -> list[Flight]:
"""
提取所有航班信息
"""
result = await extraction_agent.run(
ctx.deps.source_material,
usage_limits=UsageLimits(request_limit=3),
usage=ctx.usage,
)
if len(output := result.output) != 8:
raise ModelRetry("提取到的所有航班信息应为 8 条")
return output
@agent.tool
async def select_flight(ctx: RunContext[Deps]) -> Flight | FlightNoFound:
"""
选择航班
"""
result = await selection_agent.run(ctx.deps.source_material, usage=ctx.usage)
if isinstance(result.output, Flight):
if not ctx.tool_call_approved:
raise ApprovalRequired(metadata={"reason": "protected"})
return result.output
@agent.output_validator
async def validate_output(
ctx: RunContext[Deps], output: Flight | FlightNoFound
) -> Flight | FlightNoFound:
"""
输出校验
"""
# 不校验未查询到航班
if isinstance(output, FlightNoFound):
return output
errors = []
if output.date != ctx.deps.date:
errors.append(f"航班日期应为 {ctx.deps.date}, 不是 {output.date}")
if output.origin != ctx.deps.origin:
errors.append(f"航班出发机场应为 {ctx.deps.origin}, 不是 {output.origin}")
if output.destination != ctx.deps.destination:
errors.append(
f"航班到达机场应为 {ctx.deps.destination}, 不是 {output.destination}"
)
if errors:
raise ModelRetry("\n".join(errors))
return output
source_material = """
1. Flight SFO-AK123
- Price: $350
- Origin: San Francisco International Airport (SFO)
- Destination: Ted Stevens Anchorage International Airport (ANC)
- Date: January 10, 2025
2. Flight SFO-AK456
- Price: $370
- Origin: San Francisco International Airport (SFO)
- Destination: Fairbanks International Airport (FAI)
- Date: January 10, 2025
3. Flight SFO-AK789
- Price: $400
- Origin: San Francisco International Airport (SFO)
- Destination: Juneau International Airport (JNU)
- Date: January 20, 2025
4. Flight NYC-LA101
- Price: $250
- Origin: San Francisco International Airport (SFO)
- Destination: Ted Stevens Anchorage International Airport (ANC)
- Date: January 10, 2025
5. Flight CHI-MIA202
- Price: $200
- Origin: Chicago O'Hare International Airport (ORD)
- Destination: Miami International Airport (MIA)
- Date: January 12, 2025
6. Flight BOS-SEA303
- Price: $120
- Origin: Boston Logan International Airport (BOS)
- Destination: Ted Stevens Anchorage International Airport (ANC)
- Date: January 12, 2025
7. Flight DFW-DEN404
- Price: $150
- Origin: Dallas/Fort Worth International Airport (DFW)
- Destination: Denver International Airport (DEN)
- Date: January 10, 2025
8. Flight ATL-HOU505
- Price: $180
- Origin: Hartsfield-Jackson Atlanta International Airport (ATL)
- Destination: George Bush Intercontinental Airport (IAH)
- Date: January 10, 2025
"""
async def run_stream_events(
user_prompt: str,
deps: Deps,
message_history: list[ModelMessage],
usage: RunUsage,
deferred_tool_results: DeferredToolResults | None = None,
) -> AsyncGenerator:
async with agent.run_stream_events(
user_prompt=user_prompt,
deps=deps,
message_history=message_history,
usage=usage,
deferred_tool_results=deferred_tool_results,
) as events:
async for event in events:
yield event
async def main():
deps = Deps(
source_material=source_material,
date=datetime.date(2025, 1, 10),
origin="SFO",
destination="ANC",
)
# message_history 传空列表,不要传 None
msg_history: list[ModelMessage] = []
prompt = f"Find me a flight from {deps.origin} to {deps.destination} on {deps.date}"
gen = run_stream_events(
user_prompt=prompt,
deps=deps,
message_history=msg_history,
usage=RunUsage(),
)
# 消费异步生成器
async for ev in gen:
print(ev)
asyncio.run(main())
"""
if not isinstance(event, AgentRunResultEvent):
yield event
else:
result = event.result
# 更新任务使用量
task.usage = usage_to_dict(result.usage)
if isinstance(result.output, FlightDetail):
content = "\n---\n已查询到航班,请回复 buy 购票 或 search 重新查询\n"
# 更新任务节点为提取座位偏好
task.node = "seat_preference_extraction"
else:
content = (
"\n---\n未查询到满足您需求的航班,流程结束!\n"
)
# 更新任务为空
task = None
# 返回任务节点结果事件
yield TaskNodeResultEvent(
task=task,
content=content,
)
yield event
return
# 提取座位偏好
case "seat_preference_extraction":
if user_prompt == "buy":
# 返回任务节点结果事件
yield TaskNodeResultEvent(
task=task,
content="请和我说下您的座位偏好:\nA、F 座位是靠窗位1 排、14 排、20 排腿部空间更大、更舒展,你更想要靠窗座位,宽敞大空间座位",
)
return
elif user_prompt == "search":
# 更新任务节点为航班查询
task.node = "flight_search"
else:
async with seat_preference_extraction_agent.run_stream_events(
user_prompt=user_prompt,
message_history=message_history,
usage=usage_to_object(dict(task.usage)),
usage_limits=usage_limits_to_object(dict(task.usage_limits)),
) as events:
async for event in events:
if not isinstance(event, AgentRunResultEvent):
yield event
else:
result = event.result
# 更新任务使用量
task.usage = usage_to_dict(result.usage)
# 更新任务为空
task = None
# 返回任务节点结果事件
yield TaskNodeResultEvent(
task=task,
content="已帮您定好座位,流程结束!",
)
yield event
return
# 工具调用分片开始事件
case ToolCallPart(
tool_name=tool_name, tool_call_id=tool_call_id
):
# 构建消息实例
message = Message(
type=MessageType.TOOL_CALL,
title=tool_name,
content=",",
is_running=True,
)
tool_call_ids.add(tool_call_id)
# 添加至消息字典
dialog.thoughts[index] = Thought(
type="tool_call",
content="正在生成调用参数",
)
# ========== 函数工具调用事件 ==========
case FunctionToolCallEvent(tool_call_id=tool_call_id, part=part):
# 获取分片索引
index = tool_call_ids[tool_call_id]
match dialog.thoughts[index].type:
# 工具检索
case "tool_search":
dialog.thoughts[index].content = (
f"正在检索 {part.args_as_json_str()}"
)
# 能力加载
case "capability_load":
dialog.thoughts[index].content = (
f"正在加载能力 {part.tool_name}"
)
# 工具调用
case "tool_call":
dialog.thoughts[index].content = (
f"正在调用工具 {part.tool_name}"
)
case _:
continue
# ========== 函数工具结果事件 ==========
case FunctionToolResultEvent(
tool_call_id=tool_call_id,
content=content,
):
index = tool_call_ids[tool_call_id]
match dialog.thoughts[index].type:
# 工具检索
case "tool_search":
dialog.thoughts[index].content = (
content if isinstance(content, str) else ""
) # 暂仅考虑文本内容
# 能力加载
case "capability_load":
dialog.thoughts[index].content = f"已加载 {content}"
# 工具调用
case "tool_call":
dialog.thoughts[index].content = f"已调用 {content}"
case _:
continue
def init_work_flow() -> WorkFlow:
return WorkFlow(
type=WorkFlowType.BOOK_FLIGHT,
deps=Dependences(
source_material=source_material,
date=datetime.date(2025, 1, 10),
origin="SFO",
destination="ANC",
),
)
from application.states.models import (
WorkFlow,
WorkFlowType,
)
from application.workshop.models import (
DEEPSEEK_V4_FLASH_MODEL,
MODEL_SETTINGS_DISABLED_THINKING,
)
"""