Python/agent/application/workshop/book_flight.py

472 lines
17 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, cast
from pydantic import BaseModel, Field, field_validator, TypeAdapter
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 pydantic_ai.run import AgentRunResultEvent, AgentRunResult
from enum import StrEnum
from pydantic_ai._uuid import uuid7
from pydantic_ai.messages import (
FunctionToolCallEvent,
FunctionToolResultEvent,
LoadCapabilityCallPart,
PartDeltaEvent,
PartEndEvent,
PartStartEvent,
TextPart,
TextPartDelta,
ThinkingPart,
ThinkingPartDelta,
ToolCallPart,
ToolSearchCallPart,
ToolReturnPart,
)
class TaskStatus(StrEnum):
"""
任务状态枚举
"""
RUNNING = "running"
DONE = "done"
ERROR = "error"
PENDING_APPROVAL = "pending_approval"
class Task(BaseModel):
"""
任务类
"""
tool_name: str = Field(..., description="工具名称")
tool_call_id: str = Field(..., description="工具调用唯一标识")
status: TaskStatus = Field(default=TaskStatus.RUNNING, description="任务状态")
title: str = Field(default="", description="任务标题")
content: str = Field(default="", description="任务内容")
class MessageType(StrEnum):
"""
消息类型枚举
"""
USER_PROMPT = "user_prompt"
THINKING = "thinking"
WORK_OUTPUT = "work_output"
RESULT_OUTPUT = "result_output"
class Message(BaseModel):
"""
消息类
"""
id: str = Field(default_factory=lambda: str(uuid7()), description="消息唯一标识")
type: MessageType = Field(..., description="消息类型")
title: str = Field(default="", description="消息标题")
content: str = Field(default="", description="消息内容")
tasks: dict[str, Task] = Field(
default_factory=dict, description="任务字典,键为工具调用唯一标识"
)
is_running: bool = Field(
default=False, description="正在运行True 表示正在运行False 表示运行完成"
)
is_shown: bool = Field(
default=False, description="展示组件True 表示展示False 表示隐藏"
)
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 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 NoSelectedFlight(BaseModel):
"""
未选择到航班
"""
class Deps(BaseModel):
"""
依赖类
"""
date: datetime.date = Field(..., description="日期")
origin_airport_code: str = Field(..., description="出发机场代码")
destination_airport_code: str = Field(..., description="到达机场代码")
searched_flights: list[Flight] | None = Field(
default=None, description="查询到的航班"
)
recommended_flight: Flight | NoSelectedFlight | None = Field(
default=None, description="推荐航班"
)
# 主智能体
agent = Agent[Deps, Flight | NoSelectedFlight | DeferredToolRequests](
model=DEEPSEEK_V4_FLASH_MODEL,
model_settings=MODEL_SETTINGS,
deps_type=Deps,
output_type=Flight | NoSelectedFlight | DeferredToolRequests,
system_prompt=(
"你的任务是帮助用户查询并预定航班,**必须按照下述步骤执行**",
"1. 使用 match_flight 匹配满足用户需求的航班;",
"2. 使用 select_flight 选择航班。若未选择到航班则返回 NoSelectedFlight",
"3. 若选择到航班则使用 book_flight 预定航班。由用户确认后预定。",
),
)
# 提取智能体
extraction_agent = Agent[Deps, list[Flight]](
model=DEEPSEEK_V4_FLASH_MODEL,
model_settings=MODEL_SETTINGS,
deps_type=Deps,
output_type=list[Flight],
system_prompt=(
"你的任务是在来源资料中提取所有航班信息,包括航班号、日期、出发机场代码、到达机场代码和机票价格。",
"其中出发机场代码和到达机场代码须为英文、大写、3位字符串例如 San Francisco International Airport (SFO) 的机场代码为 SFO。",
"若某航班信息不全则跳过该航班,若没有任何航班信息则返回空列表。",
"禁止编造。",
),
)
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
"""
@agent.tool
async def search_flights(ctx: RunContext[Deps]) -> list[Flight]:
"""
查询航班
**如何提高提取准确率**
1. 系统提示词规则约束
2. 降低模型温度
3. 输出类型约定和输出模型校验
4. 业务规则约束(例如,输出不可为空列表)
**设计工具需先设计业务流程**
本示例按照 查询航班 -> 推荐航班 -> 预定航班
"""
extraction_agent_result = await extraction_agent.run(
f"来源资料:\n{source_material}",
deps=ctx.deps,
usage=ctx.usage,
usage_limits=UsageLimits(request_limit=10),
)
if (
extracted_flight_counts := len(
extracted_flights := extraction_agent_result.output
)
) != 8: # 模拟业务规则约束
raise ModelRetry(
f"提取到的所有航班信息数量应为 8 ,不是 {extracted_flight_counts}"
)
# 查询到的航班
ctx.deps.searched_flights = [
flight
for flight in extracted_flights
if flight.date == ctx.deps.date
and flight.origin_airport_code == ctx.deps.origin_airport_code
and flight.destination_airport_code == ctx.deps.destination_airport_code
]
return ctx.deps.searched_flights
@agent.tool
async def recommend_flight(ctx: RunContext[Deps]) -> Flight | NoSelectedFlight:
"""
推荐航班
**如何提高提取准确率**
1. 系统提示词规则约束
2. 降低模型温度
3. 输出类型约定和输出模型校验
4. 业务规则约束(例如,输出不可为空列表)
"""
if (searched_flights := ctx.deps.searched_flights) is None:
raise ModelRetry("必须先使用 search_flights 查询航班")
ctx.deps.recommended_flight = (
min(searched_flights, key=lambda flight: flight.airfare)
if searched_flights
else NoSelectedFlight()
)
return ctx.deps.recommended_flight
@agent.tool
async def book_flight(ctx: RunContext[Deps]) -> Flight | NoSelectedFlight:
"""
预定航班
"""
if (recommended_flight := ctx.deps.recommended_flight) is None:
raise ModelRetry("必须先使用 recommend_flight 推荐航班")
if isinstance(recommended_flight, NoSelectedFlight):
return recommended_flight
else:
if not ctx.tool_call_approved:
raise ApprovalRequired()
return recommended_flight
@agent.output_validator
async def validate_output(
ctx: RunContext[Deps], output: Flight | NoSelectedFlight | DeferredToolRequests
) -> Flight | NoSelectedFlight | DeferredToolRequests:
"""
输出校验
"""
# 不校验未选择到航班或延迟工具请求
if isinstance(output, NoSelectedFlight | DeferredToolRequests):
return output
errors = []
if output.date != ctx.deps.date:
errors.append(f"航班日期应为 {ctx.deps.date}, 不是 {output.date}")
if output.origin_airport_code != ctx.deps.origin_airport_code:
errors.append(
f"航班出发机场应为 {ctx.deps.origin_airport_code}, 不是 {output.origin_airport_code}"
)
if output.destination_airport_code != ctx.deps.destination_airport_code:
errors.append(
f"航班到达机场应为 {ctx.deps.destination_airport_code}, 不是 {output.destination_airport_code}"
)
if errors:
raise ModelRetry("\n".join(errors))
return output
async def run_stream_events(
deps: Deps,
user_prompt: str,
message_history: list[ModelMessage],
usage: RunUsage,
deferred_tool_results: DeferredToolResults | None = None,
) -> AsyncGenerator:
# 构建消息
message = Message(
type=MessageType.WORK_OUTPUT,
title="预定航班",
is_running=True,
)
yield message
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:
print(event)
match event:
case PartStartEvent(
part=part,
):
match part:
case ToolCallPart(
tool_name=tool_name, tool_call_id=tool_call_id
):
match tool_name:
case "search_flights":
if tool_name not in message.tasks:
message.tasks[tool_name] = Task(
tool_name=tool_name,
tool_call_id=tool_call_id,
title="正在查询航班",
)
yield message
case "recommend_flight":
if tool_name not in message.tasks:
message.tasks[tool_name] = Task(
tool_name=tool_name,
tool_call_id=tool_call_id,
title="正在推荐航班",
)
yield message
case "book_flight":
if tool_name not in message.tasks:
message.tasks[tool_name] = Task(
tool_name=tool_name,
tool_call_id=tool_call_id,
title="正在预定航班",
)
yield message
case FunctionToolResultEvent(part):
match part:
case ToolReturnPart(
tool_name=tool_name,
tool_call_id=tool_call_id,
content=content,
):
match tool_name:
case "search_flights":
task = message.tasks[tool_name]
searched_flights = cast(list[Flight], content)
task.title = (
f"已查询到 {len(searched_flights)} 班航班"
)
task.content = "\n".join(
[
f"{searched_flight.number} - {searched_flight.date.strftime('%Y-%m-%d')}{searched_flight.origin_airport_code} 飞往 {searched_flight.destination_airport_code} - {searched_flight.airfare}"
for searched_flight in searched_flights
]
)
task.status = TaskStatus.DONE
yield message
case "recommend_flight":
task = message.tasks[tool_name]
recommended_flight = cast(
Flight | NoSelectedFlight, content
)
if isinstance(recommended_flight, Flight):
task.title = "为您推荐"
task.content = f"{recommended_flight.number} - {recommended_flight.date.strftime('%Y-%m-%d')}{recommended_flight.origin_airport_code} 飞往 {recommended_flight.destination_airport_code} - {recommended_flight.airfare}"
else:
task.title = "为您推荐"
task.status = TaskStatus.DONE
yield message
case AgentRunResultEvent(
result=result,
):
match result:
case AgentRunResult(output=output):
match output:
case DeferredToolRequests(approvals=approvals):
for approval in approvals:
match approval:
case ToolCallPart(tool_name=tool_name):
task = message.tasks[tool_name]
task.content = "请确认是否预定 "
task.status = (
TaskStatus.PENDING_APPROVAL
)
yield message
async def main():
# 实例化依赖
deps = Deps(
date=datetime.date(2025, 1, 10),
origin_airport_code="SFO",
destination_airport_code="ANC",
)
user_prompt = f"帮我查询并预定在 {deps.date.strftime('%Y-%m-%d')}{deps.origin_airport_code}{deps.destination_airport_code} 的航班"
message_history: list[ModelMessage] = []
events = run_stream_events(
deps=deps,
user_prompt=user_prompt,
message_history=message_history,
usage=RunUsage(),
)
# 消费异步生成器
async for event in events:
print()
if isinstance(event, AgentRunResultEvent):
# 保存本次运行新增消息
message_history.extend(event.result.new_messages())
if isinstance(event, Message):
print(event)
asyncio.run(main())