# -*- 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