# -*- coding: utf-8 -*- """ 生成产品需求文档智能体 """ from dataclasses import dataclass import datetime from enum import Enum from typing import Any, AsyncGenerator, Optional, Literal from pydantic import Any, BaseModel, Enum, Field, Optional from pydantic_ai import ( Agent, ModelMessage, ModelRetry, RunContext, RunUsage, UsageLimits, ) from pydantic_ai.usage import RunUsage from models import DEEPSEEK_V4_FLASH_MODEL #from application.tasks.agent import Agent # 'if-token-present' means nothing will be sent (and the example will work) if you don't have logfire configured class FlightDetails(BaseModel): """Details of the most suitable flight.""" flight_number: str price: int origin: str = Field(description='Three-letter airport code') destination: str = Field(description='Three-letter airport code') date: datetime.date class NoFlightFound(BaseModel): """When no valid flight is found.""" @dataclass class Deps: web_page_text: str req_origin: str req_destination: str req_date: datetime.date # This agent is responsible for controlling the flow of the conversation. search_agent= Agent[Deps, FlightDetails | NoFlightFound]( model=DEEPSEEK_V4_FLASH_MODEL, output_type=FlightDetails | NoFlightFound, deps_type=Deps, retries=4, system_prompt=( 'Your job is to find the cheapest flight for the user on the given date. ' ), ) # This agent is responsible for extracting flight details from web page text. extraction_agent = Agent( output_type=list[FlightDetails], system_prompt='Extract all the flight details from the given text.', ) @search_agent.tool async def extract_flights(ctx: RunContext[Deps]) -> list[FlightDetails]: """Get details of all flights.""" # we pass the usage to the search agent so requests within this agent are counted result = await extraction_agent.run(ctx.deps.web_page_text, usage=ctx.usage) return result.output @search_agent.output_validator async def validate_output( ctx: RunContext[Deps], output: FlightDetails | NoFlightFound ) -> FlightDetails | NoFlightFound: """Procedural validation that the flight meets the constraints.""" if isinstance(output, NoFlightFound): return output errors: list[str] = [] if output.origin != ctx.deps.req_origin: errors.append( f'Flight should have origin {ctx.deps.req_origin}, not {output.origin}' ) if output.destination != ctx.deps.req_destination: errors.append( f'Flight should have destination {ctx.deps.req_destination}, not {output.destination}' ) if output.date != ctx.deps.req_date: errors.append(f'Flight should be on {ctx.deps.req_date}, not {output.date}') if errors: raise ModelRetry('\n'.join(errors)) else: return output class SeatPreference(BaseModel): row: int = Field(ge=1, le=30) seat: Literal['A', 'B', 'C', 'D', 'E', 'F'] class Failed(BaseModel): """Unable to extract a seat selection.""" # This agent is responsible for extracting the user's seat selection seat_preference_agent = Agent[object, SeatPreference | Failed]( 'openai:gpt-5.2', output_type=SeatPreference | Failed, system_prompt=( "Extract the user's seat preference. " 'Seats A and F are window seats. ' 'Row 1 is the front row and has extra leg room. ' 'Rows 14, and 20 also have extra leg room. ' ), ) # in reality this would be downloaded from a booking site, # potentially using another agent to navigate the site flights_web_page = """ 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 """ # restrict how many requests this app can make to the LLM usage_limits = UsageLimits(request_limit=15) async def flight_booking( state: FlightFlowState, user_prompt: Optional[str] ) -> AsyncGenerator: if state.stage == FlowStage.EXEC: prompt = f'Find me a flight from {state.deps.req_origin} to {state.deps.req_destination} on {state.deps.req_date}' run_result = None # 1. 模型流式查询航班 async with search_agent.run_stream_events( user_prompt=prompt, deps=state.deps, message_history=state.message_history, usage_limits=state.usage_limits, ) as stream: async for evt in stream: yield evt run_result = stream.result # 2. 保存本轮模型对话到数据库(核心:上下文持久化,防止断裂) if run_result is not None: # 写入对话历史,下一轮get_message_history可以读到航班内容 await db_state.create_result_record( conversation_id=state.conversation_id, dialog_id=state.dialog_id, new_messages=run_result.new_messages(), ) # 用量回填 state.usage = run_result.usage # 3. 判断业务结果分支 if isinstance(run_result.output, NoFlightFound): # 无航班场景 state.stage = FlowStage.FINISH yield AgentStreamEvent.text_event("未找到符合条件的航班,预订流程结束") return else: # ✅ 查询到航班,推送业务选择提示(前端展示按钮/文字提示) tip_text = "已查询到航班,请回复 buy 购票 / search 重新查询" yield AgentStreamEvent.text_event(tip_text) # 【可选】如果需要让AI记住选择指令,把提示作为系统消息入库 # await db_state.append_system_message(conversation_id=state.conversation_id, content=tip_text) # 阶段挂起,等待用户输入buy/search,不直接结束流程 state.stage = FlowStage.WAIT_USER_INPUT return