Python/agent/application/workshop/book_flight.py

339 lines
12 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 datetime
from typing import Any, 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 (
NoResult,
InteractionType,
Approval,
)
from application.workshop.models import (
DEEPSEEK_V4_FLASH_MODEL,
DEEPSEEK_V4_FLASH_MODEL_SETTINGS,
)
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 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="查询到的航班"
)
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_SETTINGS,
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_SETTINGS,
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_SETTINGS,
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"提取到的所有航班信息不可能为空")
# 查询到的航班
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,
)
if not searched_flights: # 模拟业务规则约束:查询到的航班不能为空
raise ModelRetry("查询到的航班不能为空")
ctx.deps.searched_flights = searched_flights
return searched_flights
@agent.tool(
name="预定航班",
requires_approval=True,
)
async def book_flight(ctx: RunContext[Deps]) -> Flight | NoResult:
"""
预定航班
"""
if (searched_flights := ctx.deps.searched_flights) is None:
raise ModelRetry("必须先使用 search_flights 查询航班")
# 提取到的航班号
extracted_flight_number = (
await extraction_flight_number_agent.run(
f"用户提示词:\n{cast(str, ctx.prompt).upper()}\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 extracted_flight_number:
return NoResult()
if extracted_flight_number not in flight_numbers:
raise ModelRetry(f"提取到的航班号 {extracted_flight_number} 不在查询到的航班中")
# 预定到的航班
ctx.deps.booked_flight = next(
flight
for flight in searched_flights
if flight.number == 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: Any,
user_prompt: str,
message_history: list[ModelMessage],
usage: RunUsage,
usage_limits: UsageLimits | None = None,
deferred_tool_results: DeferredToolResults | None = None,
) -> AsyncGenerator[
AgentStreamEvent
| AgentRunResultEvent[Flight | NoResult | DeferredToolRequests]
| None,
]:
"""
运行并流式输出事件
"""
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\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\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:
yield event