Python/agent/application/workshop/book_flight.py

331 lines
13 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 AsyncGenerator, Literal
from pydantic import BaseModel, Field
from pydantic_ai import Agent, ModelMessage, ModelRetry, RunContext, UsageLimits, RunUsage, DeferredToolRequests
from pydantic_ai.run import AgentRunResultEvent
from application.states.models import (
WorkFlow,
WorkFlowType,
)
from application.workshop.models import DEEPSEEK_V4_FLASH_MODEL, MODEL_SETTINGS_DISABLED_THINKING
class Dependences(BaseModel):
"""
依赖类
"""
date: datetime.date = Field(..., description="航班日期")
origin: str = Field(..., description="航班出发机场")
destination: str = Field(..., description="航班到达机场")
source_material: str = Field(..., description="航班来源资料")
class Flight(BaseModel):
"""
航班类
"""
number: str = Field(..., description="航班号")
date: datetime.date = Field(..., description="航班日期")
origin: str = Field(..., description="航班出发机场")
destination: str = Field(..., description="航班到达机场")
airfare: int = Field(..., description="航班机票价格")
class FlightNoFound(BaseModel):
"""
未查询到航班类
"""
class SeatPreference(BaseModel):
"""
座位偏好类
"""
row: int = Field(ge=1, le=30, description="座位行")
column: Literal["A", "B", "C", "D", "E", "F"] = Field(description="座位列")
class SeatPreferenceNoExtracted(BaseModel):
"""
未提取到座位偏好类
"""
# 航班查询智能体
flight_inquiry_agent = Agent[Dependences, Flight | FlightNoFound | DeferredToolRequests](
model=DEEPSEEK_V4_FLASH_MODEL,
model_settings=MODEL_SETTINGS_DISABLED_THINKING,
deps_type=Dependences,
output_type=Flight | FlightNoFound | DeferredToolRequests,
system_prompt="你的任务是根据给定的日期、出发机场和到达机场帮助用户找到最便宜的航班。",
)
# 航班信息提取智能体
flights_extraction_agent = Agent(
model=DEEPSEEK_V4_FLASH_MODEL,
model_settings=MODEL_SETTINGS_DISABLED_THINKING,
output_type=list[Flight],
system_prompt="你的任务是根据给定的航班来源资料提取出所有航班信息,包括航班号、航班日期、航班出发机场、航班到达机场和航班机票价格。禁止编造。")
@flight_inquiry_agent.tool
async def extract_flight_details(ctx: RunContext[Dependences]) -> list[Flight]:
"""
提取所有航班信息
"""
result = await flights_extraction_agent.run(
ctx.deps.source_material, usage=ctx.usage
)
return result.output
@flight_inquiry_agent.output_validator
async def validate_output(
ctx: RunContext[Dependences], output: Flight | FlightNoFound | DeferredToolRequests
) -> Flight | FlightNoFound | DeferredToolRequests:
"""
输出校验:航班信息
"""
# 不校验未查询到航班
if isinstance(output, FlightNoFound):
return output
# 不校验延迟工具请求
if isinstance(output, DeferredToolRequests):
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
# 座位偏好提取智能体
seat_preference_extraction_agent = Agent[object, SeatPreference | SeatPreferenceNoExtracted](
model=DEEPSEEK_V4_FLASH_MODEL,
model_settings=MODEL_SETTINGS_DISABLED_THINKING,
output_type=SeatPreference | SeatPreferenceNoExtracted,
system_prompt="你的任务是根据用户回答提取座位偏好。座位规则说明A 座、F 座为靠窗座位;第 1 排是前排座位腿部空间更大14 排、20 排同样拥有加宽腿部空间",
)
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
"""
def init_work_flow() -> WorkFlow:
return WorkFlow(
type=WorkFlowType.BOOK_FLIGHT,
deps=Deps_(
flight_info=flight_info,
date=datetime.date(2025, 1, 10),
origin="SFO",
destination="ANC",
),
usage_limits={},
)
async def run_stream_events(
usage: RunUsage,
user_prompt: str,
message_history: list[ModelMessage],
) -> AsyncGenerator:
result = None
while True:
if not task:
return
match task.node:
# 航班查询
case "flight_search":
async with flight_search_agent.run_stream_events(
user_prompt=user_prompt,
deps=task.deps,
message_history=message_history,
usage=usage,
) 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)
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