# -*- coding: utf-8 -*- """ 预定航班工作流(范式) """ import datetime from re import I from typing import AsyncGenerator, cast, Any from dataclasses import replace 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 AgentRunEvent, NoResult 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="预定航班") async def book_flight(ctx: RunContext[Deps]) -> Flight | NoResult: """ 预定航班 """ if (searched_flights := ctx.deps.searched_flights) is None: raise ModelRetry("必须先使用 search_flights 查询航班") if not ctx.tool_call_approved: raise ApprovalRequired(metadata={"content": "您要预定哪班航班?"}) # 提取到的航班号 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] | AgentRunEvent | 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( result=result, ): match result: case AgentRunResult(output=output): match output: case Flight(): yield AgentRunEvent(content="预定成功") case DeferredToolRequests(approvals=approvals): yield AgentRunResultEvent( result=AgentRunResult( output=DeferredToolRequests( approvals=approvals ) ) ) yield event