This commit is contained in:
parent
406c3877b1
commit
107c7ed03f
|
|
@ -28,6 +28,8 @@ instruction = """
|
|||
# 行文要求
|
||||
语言通俗,逻辑完整简洁,无多余废话。
|
||||
"""
|
||||
|
||||
|
||||
async def run_stream_events(
|
||||
task_type: TaskType,
|
||||
user_prompt: str,
|
||||
|
|
@ -45,8 +47,8 @@ async def run_stream_events(
|
|||
async with agent.run_stream_events(
|
||||
user_prompt=user_prompt,
|
||||
message_history=message_history,
|
||||
) as event_stream:
|
||||
async for event in event_stream:
|
||||
) as events:
|
||||
async for event in events:
|
||||
yield event
|
||||
|
||||
case "flight":
|
||||
|
|
|
|||
|
|
@ -4,10 +4,9 @@
|
|||
"""
|
||||
from dataclasses import dataclass
|
||||
import datetime
|
||||
from enum import Enum
|
||||
from typing import Any, AsyncGenerator, Optional, Literal
|
||||
from typing import AsyncGenerator, Optional, Literal, List
|
||||
|
||||
from pydantic import Any, BaseModel, Enum, Field, Optional
|
||||
from pydantic import BaseModel, Field, Optional
|
||||
from pydantic_ai import (
|
||||
Agent,
|
||||
ModelMessage,
|
||||
|
|
@ -17,52 +16,62 @@ from pydantic_ai import (
|
|||
UsageLimits,
|
||||
)
|
||||
from pydantic_ai.usage import RunUsage
|
||||
|
||||
from pydantic_ai.run import AgentRunResultEvent
|
||||
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:
|
||||
class Deps(BaseModel):
|
||||
web_page_text: str
|
||||
req_origin: str
|
||||
req_destination: str
|
||||
req_date: datetime.date
|
||||
|
||||
|
||||
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
|
||||
"""
|
||||
|
||||
|
||||
# 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,
|
||||
output_type=FlightDetails | NoFlightFound,
|
||||
retries=3,
|
||||
system_prompt=(
|
||||
'Your job is to find the cheapest flight for the user on the given date. '
|
||||
"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.',
|
||||
model=DEEPSEEK_V4_FLASH_MODEL,
|
||||
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]:
|
||||
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
|
||||
|
|
@ -73,32 +82,38 @@ async def validate_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}'
|
||||
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}'
|
||||
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}')
|
||||
errors.append(f"Flight should be on {ctx.deps.req_date}, not {output.date}")
|
||||
if errors:
|
||||
raise ModelRetry('\n'.join(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']
|
||||
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',
|
||||
model=DEEPSEEK_V4_FLASH_MODEL,
|
||||
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. '
|
||||
"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,
|
||||
|
|
@ -146,26 +161,27 @@ flights_web_page = """
|
|||
- Date: January 10, 2025
|
||||
"""
|
||||
# restrict how many requests this app can make to the LLM
|
||||
usage_limits = UsageLimits(request_limit=15)
|
||||
usage_limits = UsageLimits(request_limit=5)
|
||||
|
||||
|
||||
async def flight_booking(
|
||||
state: FlightFlowState,
|
||||
user_prompt: Optional[str]
|
||||
deps: Deps, user_prompt: Optional[str] = None
|
||||
) -> 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}'
|
||||
if deps.stage == FlowStage.EXEC:
|
||||
prompt = f"Find me a flight from {deps.req_origin} to {deps.req_destination} on {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
|
||||
deps=deps,
|
||||
message_history=message_history,
|
||||
usage_limits=usage_limits,
|
||||
) as events:
|
||||
async for event in events:
|
||||
if isinstance(event, AgentRunResultEvent):
|
||||
run_result = event
|
||||
yield event
|
||||
|
||||
# 2. 保存本轮模型对话到数据库(核心:上下文持久化,防止断裂)
|
||||
if run_result is not None:
|
||||
|
|
@ -189,9 +205,6 @@ async def flight_booking(
|
|||
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
|
||||
Loading…
Reference in New Issue