Python/agent/application/tasks/book_flight.py

264 lines
9.1 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
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.tasks.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 Task(
type=TaskType.BOOK_FLIGHT,
node="flight_search",
deps=Deps(
flight_info=flight_info,
date=datetime.date(2025, 1, 10),
origin="SFO",
destination="ANC",
),
)
async def run_stream_events(
task: Task | None,
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_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