Python/agent/application/workshop/book_flight.py

602 lines
22 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
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"
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="到达机场代码")
matched_flights: list[Flight] | None = Field(
default=None, description="匹配到的航班"
)
selected_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 match_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.matched_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.matched_flights
@agent.tool
async def select_flight(ctx: RunContext[Deps]) -> Flight | NoSelectedFlight:
"""
选择航班
**如何提高提取准确率**
1. 系统提示词规则约束
2. 降低模型温度
3. 输出类型约定和输出模型校验
4. 业务规则约束(例如,输出不可为空列表)
"""
if (matched_flights := ctx.deps.matched_flights) is None:
raise ModelRetry("必须先使用 match_flights 匹配满足用户需求的航班")
ctx.deps.selected_flight = (
min(matched_flights, key=lambda flight: flight.airfare)
if matched_flights
else NoSelectedFlight()
)
return ctx.deps.selected_flight
@agent.tool
async def book_flight(ctx: RunContext[Deps]) -> Flight | NoSelectedFlight:
"""
预定航班
"""
if (selected_flight := ctx.deps.selected_flight) is None:
raise ModelRetry("必须先使用 select_flight 选择航班")
if isinstance(selected_flight, NoSelectedFlight):
return selected_flight
else:
if not ctx.tool_call_approved:
raise ApprovalRequired(
metadata={
"reason": f"需要用户确认是否预定 {selected_flight.number}{selected_flight.date.strftime('%Y-%m-%d')}{selected_flight.origin_airport_code}{selected_flight.destination_airport_code} 的航班"
}
)
return selected_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:
match event:
case PartStartEvent(
part=part,
):
match part:
case ToolCallPart(
tool_name=tool_name, tool_call_id=tool_call_id
):
match tool_name:
case "match_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 "select_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 "match_flights":
matched_flights = cast(list[Flight], content)
task = message.tasks[tool_name]
task.title = f"已查询到 {len(matched_flights)} 班航班"
task.content = "\n".join(
[
f"{matched_flight.number} - {matched_flight.date.strftime('%Y-%m-%d')} - {matched_flight.origin_airport_code} ~ {matched_flight.destination_airport_code} ${matched_flight.airfare}"
for matched_flight in matched_flights
]
)
task.status = TaskStatus.DONE
yield message
case "select_flight":
selected_flight = cast(Flight, content)
task = message.tasks[tool_name]
task.title = f"已选择航班"
task.content = f"{selected_flight.number} - {selected_flight.date.strftime('%Y-%m-%d')} - {selected_flight.origin_airport_code} ~ {selected_flight.destination_airport_code} ${selected_flight.airfare}"
task.status = TaskStatus.DONE
yield message
async def main():
# 实例化依赖
deps = Deps(
date=datetime.date(2025, 1, 10),
origin_airport_code="SFO",
destination_airport_code="ANC",
selected_flight=None,
)
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(event)
if isinstance(event, AgentRunResultEvent):
# 保存本次运行新增消息
message_history.extend(event.result.new_messages())
asyncio.run(main())
"""
@selection_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
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,
)
"""