267 lines
9.2 KiB
Python
267 lines
9.2 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
预定航班
|
||
"""
|
||
import datetime
|
||
from typing import AsyncGenerator, Literal
|
||
|
||
from pydantic import BaseModel, Field
|
||
from pydantic_ai import Agent, ModelMessage, ModelRetry, RunContext, UsageLimits
|
||
from pydantic_ai.run import AgentRunResultEvent
|
||
|
||
from application.states.models import (
|
||
Task,
|
||
TaskNodeResultEvent,
|
||
TaskType,
|
||
usage_limits_to_object,
|
||
usage_to_dict,
|
||
usage_to_object,
|
||
)
|
||
from application.workshop.models import DEEPSEEK_V4_FLASH_MODEL, MODEL_SETTINGS
|
||
|
||
|
||
class Deps(BaseModel):
|
||
"""
|
||
依赖项类
|
||
"""
|
||
|
||
flight_info: str = Field(..., description="所有航班信息")
|
||
date: datetime.date = Field(..., description="航班日期")
|
||
origin: str = Field(..., description="出发机场")
|
||
destination: str = Field(..., description="到达机场")
|
||
|
||
|
||
class FlightDetail(BaseModel):
|
||
"""
|
||
航班详情类
|
||
"""
|
||
|
||
number: str = Field(description="航班号")
|
||
date: datetime.date = Field(description="航班日期")
|
||
origin: str = Field(description="出发机场")
|
||
destination: str = Field(description="到达机场")
|
||
price: int = Field(description="机票价格")
|
||
|
||
|
||
class NoFlightFound(BaseModel):
|
||
"""
|
||
未查询到航班类
|
||
"""
|
||
|
||
|
||
# 航班查询智能体
|
||
flight_search_agent = Agent[Deps, FlightDetail | NoFlightFound](
|
||
model=DEEPSEEK_V4_FLASH_MODEL,
|
||
model_settings=MODEL_SETTINGS,
|
||
deps_type=Deps,
|
||
output_type=FlightDetail | NoFlightFound,
|
||
system_prompt="你的工作是在给定日期、出发机场和到达机场为用户找到最便宜的航班",
|
||
)
|
||
|
||
# 所有航班详情提取智能体
|
||
flight_details_extraction_agent = Agent(
|
||
model=DEEPSEEK_V4_FLASH_MODEL,
|
||
model_settings=MODEL_SETTINGS,
|
||
output_type=list[FlightDetail],
|
||
system_prompt="从给定文本中提取所有航班详细,包括航班号、航班日期、出发机场、到达机场和机票价格",
|
||
)
|
||
|
||
|
||
@flight_search_agent.tool
|
||
async def extract_flight_details(ctx: RunContext[Deps]) -> list[FlightDetail]:
|
||
"""
|
||
工具:提取所有航班详情
|
||
"""
|
||
result = await flight_details_extraction_agent.run(
|
||
ctx.deps.flight_info, usage=ctx.usage
|
||
)
|
||
return result.output
|
||
|
||
|
||
@flight_search_agent.output_validator
|
||
async def validate_output(
|
||
ctx: RunContext[Deps], output: FlightDetail | NoFlightFound
|
||
) -> FlightDetail | NoFlightFound:
|
||
"""
|
||
输出校验:航班详情
|
||
"""
|
||
if isinstance(output, NoFlightFound):
|
||
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
|
||
|
||
|
||
class SeatPreference(BaseModel):
|
||
"""
|
||
座位偏好类
|
||
"""
|
||
|
||
row: int = Field(ge=1, le=30, description="座位行")
|
||
column: Literal["A", "B", "C", "D", "E", "F"] = Field(description="座位列")
|
||
|
||
|
||
class NoSeatExtracted(BaseModel):
|
||
"""
|
||
未提取到座位偏好类
|
||
"""
|
||
|
||
|
||
# 提取座位偏好智能体(无依赖项)
|
||
seat_preference_extraction_agent = Agent[object, SeatPreference | NoSeatExtracted](
|
||
model=DEEPSEEK_V4_FLASH_MODEL,
|
||
model_settings=MODEL_SETTINGS,
|
||
output_type=SeatPreference | NoSeatExtracted,
|
||
system_prompt="提取用户的座位偏好。座位规则说明:A 座、F 座为靠窗座位;第 1 排是前排座位,腿部空间更大;14 排、20 排同样拥有加宽腿部空间",
|
||
)
|
||
|
||
flight_info = """
|
||
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_task() -> Task:
|
||
return WorkFlow(
|
||
type=WorkType.BOOK_FLIGHT,
|
||
tools={"extract_flight_details": {"": ""}},
|
||
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],
|
||
work: WorkType | None = None,
|
||
|
||
) -> 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_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)
|
||
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
|