Python/agent/application/tasks/generating_prd_agent.py

197 lines
7.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 -*-
"""
生成产品需求文档智能体
"""
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