365 lines
13 KiB
Python
365 lines
13 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
预定航班工作流(范式)
|
||
"""
|
||
import datetime
|
||
from typing import AsyncGenerator, cast
|
||
|
||
from pydantic import BaseModel, Field, field_validator
|
||
from pydantic_ai import (
|
||
Agent,
|
||
ApprovalRequired,
|
||
DeferredToolRequests,
|
||
DeferredToolResults,
|
||
ModelMessage,
|
||
ModelRetry,
|
||
RunContext,
|
||
RunUsage,
|
||
UsageLimits,
|
||
)
|
||
from pydantic_ai.messages import (
|
||
AgentStreamEvent,
|
||
FunctionToolResultEvent,
|
||
PartStartEvent,
|
||
ToolCallPart,
|
||
ToolReturnPart,
|
||
)
|
||
from pydantic_ai.run import AgentRunResult, AgentRunResultEvent
|
||
|
||
from application.states.models import AgentRunWarmUpEvent
|
||
from application.workshop.models import (
|
||
DEEPSEEK_V4_FLASH_MODEL,
|
||
DEEPSEEK_V4_FLASH_MODEL_DISABLED_THINKING,
|
||
)
|
||
|
||
|
||
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 NoResult(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="查询到的航班"
|
||
)
|
||
extracted_flight_number: str | None = Field(
|
||
default=None, description="提取到的航班号"
|
||
)
|
||
booked_flight: Flight | None = Field(default=None, description="预定到的航班")
|
||
|
||
|
||
# 主智能体
|
||
agent = Agent[Deps, Flight | NoResult | DeferredToolRequests](
|
||
model=DEEPSEEK_V4_FLASH_MODEL,
|
||
model_settings=DEEPSEEK_V4_FLASH_MODEL_DISABLED_THINKING,
|
||
deps_type=Deps,
|
||
output_type=Flight | NoResult | DeferredToolRequests,
|
||
system_prompt=(
|
||
"你的任务是帮助用户查询并预定航班,**必须按照下述步骤执行**:",
|
||
"1. 使用 search_flights 查询满足用户需求的航班;",
|
||
"2. 使用 book_flight 预定航班。若未选择到航班则返回 NoSelectedFlight;",
|
||
),
|
||
)
|
||
|
||
|
||
# 提取所有航班信息智能体
|
||
extraction_flights_agent = Agent[Deps, list[Flight]](
|
||
model=DEEPSEEK_V4_FLASH_MODEL,
|
||
model_settings=DEEPSEEK_V4_FLASH_MODEL_DISABLED_THINKING,
|
||
deps_type=Deps,
|
||
output_type=list[Flight],
|
||
system_prompt=(
|
||
"你的任务是在来源资料中提取所有航班信息,包括航班号、日期、出发机场代码、到达机场代码和机票价格。",
|
||
"其中,出发机场代码和到达机场代码须为英文、大写、3位字符串,例如 San Francisco International Airport (SFO) 的机场代码为 SFO。",
|
||
"若某航班信息不全则跳过该航班,若没有任何航班信息则返回空列表。",
|
||
"禁止编造。",
|
||
),
|
||
)
|
||
|
||
# 提取航班号智能体
|
||
extraction_flight_number_agent = Agent[Deps, str | None](
|
||
model=DEEPSEEK_V4_FLASH_MODEL,
|
||
model_settings=DEEPSEEK_V4_FLASH_MODEL_DISABLED_THINKING,
|
||
deps_type=Deps,
|
||
output_type=str | None,
|
||
system_prompt=(
|
||
"你的任务是从用户提示词中提取航班号且须在查询到的航班中。",
|
||
"其中,航班号须由英文、数字和短横线组成、大写,例如 SFO-AK123。",
|
||
"若用户提示词中没有航班号则返回 None。",
|
||
"禁止编造。",
|
||
),
|
||
)
|
||
|
||
|
||
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(name="查询航班")
|
||
async def search_flights(ctx: RunContext[Deps]) -> list[Flight]:
|
||
"""
|
||
查询航班
|
||
**如何提高提取准确率**
|
||
1. 系统提示词规则约束
|
||
2. 降低模型温度
|
||
3. 输出类型约定和输出模型校验
|
||
4. 业务规则约束(例如,输出不可为空列表)
|
||
|
||
**设计工具需先设计业务流程**
|
||
本示例按照 查询航班 -> 推荐航班 -> 预定航班
|
||
"""
|
||
# 提取到的所有航班信息
|
||
extracted_flights = (
|
||
await extraction_flights_agent.run(
|
||
f"来源资料:\n{source_material}",
|
||
deps=ctx.deps,
|
||
usage=ctx.usage,
|
||
usage_limits=UsageLimits(request_limit=10),
|
||
)
|
||
).output
|
||
if not extracted_flights: # 模拟业务规则约束:提取到的所有航班信息不可能为空
|
||
raise ModelRetry(f"提取到的所有航班信息不可能为空")
|
||
|
||
# 查询到的航班
|
||
ctx.deps.searched_flights = sorted(
|
||
[
|
||
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
|
||
],
|
||
key=lambda flight: flight.airfare,
|
||
)
|
||
return ctx.deps.searched_flights
|
||
|
||
|
||
@agent.tool(name="预定航班")
|
||
async def book_flight(ctx: RunContext[Deps]) -> Flight | NoResult:
|
||
"""
|
||
预定航班
|
||
"""
|
||
if (searched_flights := ctx.deps.searched_flights) is None:
|
||
raise ModelRetry("必须先使用 search_flights 查询航班")
|
||
|
||
if not ctx.tool_call_approved:
|
||
raise ApprovalRequired()
|
||
|
||
# 提取到的航班号
|
||
ctx.deps.extracted_flight_number = (
|
||
await extraction_flight_number_agent.run(
|
||
f"用户提示词:\n{ctx.prompt}\n查询到的航班:\n{(flight_numbers:= [flight.number for flight in searched_flights])}",
|
||
deps=ctx.deps,
|
||
usage=ctx.usage,
|
||
usage_limits=UsageLimits(request_limit=10),
|
||
)
|
||
).output
|
||
if not ctx.deps.extracted_flight_number:
|
||
return NoResult()
|
||
if ctx.deps.extracted_flight_number not in flight_numbers:
|
||
raise ModelRetry(
|
||
f"提取到的航班号 {ctx.deps.extracted_flight_number} 不在查询到的航班中"
|
||
)
|
||
# 预定到的航班
|
||
ctx.deps.booked_flight = next(
|
||
flight
|
||
for flight in searched_flights
|
||
if flight.number == ctx.deps.extracted_flight_number
|
||
)
|
||
return ctx.deps.booked_flight
|
||
|
||
|
||
@agent.output_validator
|
||
async def validate_output(
|
||
ctx: RunContext[Deps],
|
||
output: Flight | NoResult | DeferredToolRequests,
|
||
) -> Flight | NoResult | DeferredToolRequests:
|
||
"""
|
||
输出校验
|
||
"""
|
||
if isinstance(output, (NoResult, 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
|
||
|
||
|
||
deps = Deps(
|
||
date=datetime.date(2025, 1, 10),
|
||
origin_airport_code="SFO",
|
||
destination_airport_code="ANC",
|
||
)
|
||
|
||
|
||
async def run_stream_events(
|
||
deps: Deps,
|
||
user_prompt: str,
|
||
message_history: list[ModelMessage],
|
||
usage: RunUsage,
|
||
usage_limits: UsageLimits | None = None,
|
||
deferred_tool_results: DeferredToolResults | None = None,
|
||
) -> AsyncGenerator[
|
||
AgentStreamEvent
|
||
| AgentRunResultEvent[str | DeferredToolRequests]
|
||
| AgentRunWarmUpEvent
|
||
| None,
|
||
]:
|
||
"""
|
||
运行并流式输出事件
|
||
"""
|
||
# 构建智能体运行预热事件
|
||
yield AgentRunWarmUpEvent(
|
||
content=f"正在预定 {deps.date.strftime('%Y-%m-%d')} 从 {deps.origin_airport_code} 到 {deps.destination_airport_code} 的航班"
|
||
)
|
||
tool_names = set()
|
||
async with agent.run_stream_events(
|
||
user_prompt=user_prompt,
|
||
deps=deps,
|
||
message_history=message_history,
|
||
usage=usage,
|
||
usage_limits=usage_limits,
|
||
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):
|
||
match tool_name:
|
||
case "查询航班" | "预定航班":
|
||
if tool_name not in tool_names:
|
||
tool_names.add(tool_name)
|
||
yield event
|
||
|
||
case FunctionToolResultEvent(part):
|
||
match part:
|
||
case ToolReturnPart(
|
||
tool_name=tool_name,
|
||
content=content,
|
||
):
|
||
match tool_name:
|
||
case "查询航班":
|
||
# 查询到的航班
|
||
searched_flights = cast(list[Flight], content)
|
||
event.part.content = f"已查询到 **{len(searched_flights)} 班航班**:\n"
|
||
for searched_flight in searched_flights:
|
||
event.part.content += f"{searched_flight.number} {searched_flight.airfare} 于 {searched_flight.date.strftime('%Y-%m-%d')} 从 {searched_flight.origin_airport_code} 到 {searched_flight.destination_airport_code}\n"
|
||
yield event
|
||
|
||
case "预定航班":
|
||
if isinstance(content, NoResult):
|
||
event.part.content = (
|
||
"未提取到航班号,请检查后重试"
|
||
)
|
||
|
||
if isinstance(content, Flight):
|
||
event.part.content = f"已预定航班:\n{content.number} {content.airfare} 于 {content.date.strftime('%Y-%m-%d')} 从 {content.origin_airport_code} 到 {content.destination_airport_code}\n"
|
||
yield event
|
||
|
||
case AgentRunResultEvent(
|
||
result=result,
|
||
):
|
||
match result:
|
||
case AgentRunResult(output=output):
|
||
match output:
|
||
case Flight():
|
||
yield AgentRunResultEvent(
|
||
result=AgentRunResult(output="预定成功")
|
||
)
|
||
case DeferredToolRequests(approvals=approvals):
|
||
yield AgentRunResultEvent(
|
||
result=AgentRunResult(
|
||
output=DeferredToolRequests(
|
||
approvals=approvals
|
||
)
|
||
)
|
||
)
|