From 92674c434b14e238871d046722fc639808bda43c Mon Sep 17 00:00:00 2001 From: liubiren Date: Thu, 13 Aug 2026 20:09:34 +0800 Subject: [PATCH] 1 --- agent/application/workshop/book_flight.py | 185 ++++++++++++++-------- 1 file changed, 120 insertions(+), 65 deletions(-) diff --git a/agent/application/workshop/book_flight.py b/agent/application/workshop/book_flight.py index 7c432b4..023e21e 100644 --- a/agent/application/workshop/book_flight.py +++ b/agent/application/workshop/book_flight.py @@ -7,7 +7,7 @@ import datetime from typing import AsyncGenerator from logfire_api.variables import ValueDoesNotEqual -from pydantic import BaseModel, Field, field_validator +from pydantic import BaseModel, Field, field_validator, TypeAdapter from pydantic_ai import ( Agent, ApprovalRequired, @@ -23,6 +23,7 @@ from pydantic_ai import ( from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.providers.openai import OpenAIProvider from sqlalchemy.sql.dml import ReturningDelete +from pydantic_ai.run import AgentRunResultEvent DEEPSEEK_V4_FLASH_MODEL = OpenAIChatModel( @@ -38,17 +39,6 @@ MODEL_SETTINGS = ModelSettings( ) # 禁用思考模式 -class Deps(BaseModel): - """ - 依赖类 - """ - - date: datetime.date = Field(..., description="日期") - origin_airport_code: str = Field(..., description="出发机场代码") - destination_airport_code: str = Field(..., description="到达机场代码") - source_material: str = Field(..., description="来源资料") - - class Flight(BaseModel): """ 航班类 @@ -79,18 +69,35 @@ class FlightNoFound(BaseModel): """ -system_prompt = """你的任务是根据用户提供的日期、出发机场和到达机场选择航班。 -1. 使用 extract_flights 从来源资料提取所有航班信息; -2. 使用 select_flight 从所有航班信息中选择匹配航班; -""" +class Deps(BaseModel): + """ + 依赖类 + """ -# 主智能体 -agent = Agent[Deps, Flight | FlightNoFound]( + date: datetime.date = Field(..., description="日期") + origin_airport_code: str = Field(..., description="出发机场代码") + destination_airport_code: str = Field(..., description="到达机场代码") + source_material: str = Field(..., description="来源资料") + extracted_flights: list[Flight] | None = Field( + default=None, description="提取到的所有航班信息" + ) + selected_flight: Flight | FlightNoFound | None = Field( + default=None, description="查询到的航班" + ) + + +# 航班查询与预定智能体 +agent = Agent[Deps, Flight | FlightNoFound | DeferredToolRequests]( model=DEEPSEEK_V4_FLASH_MODEL, model_settings=MODEL_SETTINGS, deps_type=Deps, - output_type=Flight | FlightNoFound, - system_prompt=system_prompt, + output_type=Flight | FlightNoFound | DeferredToolRequests, + system_prompt=( + "你的任务是帮助用户查询并预定航班,**必须按照下述步骤执行**:", + "1. 使用 extract_flights 提取所有航班信息;", + "2. 使用 select_flight 查询航班;", + "3. 若查询到航班则使用 book_flight 预定航班;若未查询到航班则返回 FlightNoFound。", + ), ) # 航班信息提取智能体 @@ -99,64 +106,78 @@ extraction_agent = Agent( model_settings=MODEL_SETTINGS, output_type=list[Flight], system_prompt=( - "你的任务是根据给定的来源资料提取出所有航班信息,包括航班号、日期、出发机场代码、到达机场代码和机票价格。", + "你的任务是根据提供的来源资料提取出所有航班信息,包括航班号、日期、出发机场代码、到达机场代码和机票价格。", "其中,出发机场代码和到达机场代码须为英文、大写、3位字符串,例如 San Francisco International Airport (SFO) 的机场代码为 SFO。", "若某航班信息不全则跳过该航班,若没有任何航班信息则返回空列表。", "禁止编造。", ), -) # 通过提示词约束输出 - - -# 航班选择智能体 -selection_agent = Agent( - model=DEEPSEEK_V4_FLASH_MODEL, - model_settings=MODEL_SETTINGS, - output_type=Flight | FlightNoFound, - system_prompt="你的任务是从所有航班信息中选择匹配用户提供的日期、出发机场和到达机场的航班:若有多个航班则选择最便宜的航班,选择完成需用户确认是否预定;若没有则返回 FlightNoFound。", ) -""" -如何提高提取准确率 -1. 系统提示词规则约束 -2. 降低模型温度 -3. 输出类型约定和输出模型校验 -4. 业务规则约束(例如,输出不可为空列表) -""" - @agent.tool async def extract_flights(ctx: RunContext[Deps]) -> list[Flight]: """ 提取所有航班信息 + **如何提高提取准确率** + 1. 系统提示词规则约束 + 2. 降低模型温度 + 3. 输出类型约定和输出模型校验 + 4. 业务规则约束(例如,输出不可为空列表) """ result = await extraction_agent.run( ctx.deps.source_material, - usage_limits=UsageLimits(request_limit=3), usage=ctx.usage, + usage_limits=UsageLimits(request_limit=10), ) - if len(output := result.output) != 8: - raise ModelRetry("提取到的所有航班信息应为 8 条") + + if (flight_counts := len(output := result.output)) != 8: # 模拟业务规则约束 + raise ModelRetry(f"所有航班信息数量应为 8 ,不是 {flight_counts}") + # 更新提取到的航班信息 + ctx.deps.extracted_flights = output return output +# 航班查询智能体 +selection_agent = Agent[Deps, Flight | FlightNoFound]( + model=DEEPSEEK_V4_FLASH_MODEL, + model_settings=MODEL_SETTINGS, + deps_type=Deps, + output_type=Flight | FlightNoFound, + system_prompt=( + "你的任务是在提取到的所有航班信息中查询匹配用户提供的日期、出发机场代码和到达机场代码的机票价格最低的航班。", + "若没有符合要求的航班则返回 FlightNoFound。", + ), +) + + @agent.tool async def select_flight(ctx: RunContext[Deps]) -> Flight | FlightNoFound: """ - 选择航班 + 查询航班 """ - result = await selection_agent.run(ctx.deps.source_material, usage=ctx.usage) - if isinstance(result.output, Flight): - if not ctx.tool_call_approved: - raise ApprovalRequired(metadata={"reason": "protected"}) - return result.output + extracted_flights = ctx.deps.extracted_flights + if not extracted_flights: + raise ModelRetry("须先使用 extract_flights 提取所有航班信息") + + result = await selection_agent.run( + ( + f"{ctx.prompt}", + f"所有航班信息:\n{TypeAdapter(list[Flight]).dump_json(extracted_flights).decode('utf-8')}", + ), + deps=ctx.deps, + usage=ctx.usage, + usage_limits=UsageLimits(request_limit=10), + ) + ctx.deps.selected_flight = (output := result.output) + return output -@agent.output_validator +@selection_agent.output_validator async def validate_output( ctx: RunContext[Deps], output: Flight | FlightNoFound ) -> Flight | FlightNoFound: """ - 输出校验 + 航班查询智能体输出校验 """ # 不校验未查询到航班 if isinstance(output, FlightNoFound): @@ -165,17 +186,39 @@ async def validate_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: + if output.origin_airport_code != ctx.deps.origin_airport_code: errors.append( - f"航班到达机场应为 {ctx.deps.destination}, 不是 {output.destination}" + 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 +@agent.tool +async def book_flight(ctx: RunContext[Deps]) -> Flight | FlightNoFound: + """ + 预定航班 + """ + if not (selected_flight := ctx.deps.selected_flight): + raise ModelRetry("须先使用 select_flight 查询航班") + + if isinstance(selected_flight, FlightNoFound): + return selected_flight + else: + if not ctx.tool_call_approved: + raise ApprovalRequired( + metadata={ + "reason": f"需要用户确认是否预定 {selected_flight.number} 在 {selected_flight.date.strftime('%Y-%m-%d')} 从 {selected_flight.origin_airport_code} 到 {selected_flight.destination_airport_code} 的航班" + } + ) + return selected_flight + + source_material = """ 1. Flight SFO-AK123 - Price: $350 @@ -221,8 +264,8 @@ source_material = """ async def run_stream_events( - user_prompt: str, deps: Deps, + user_prompt: str, message_history: list[ModelMessage], usage: RunUsage, deferred_tool_results: DeferredToolResults | None = None, @@ -239,30 +282,42 @@ async def run_stream_events( async def main(): + # 实例化依赖 deps = Deps( source_material=source_material, date=datetime.date(2025, 1, 10), - origin="SFO", - destination="ANC", + origin_airport_code="SFO", + destination_airport_code="ANC", + selected_flight=None, ) - # message_history 传空列表,不要传 None - msg_history: list[ModelMessage] = [] - prompt = f"Find me a flight from {deps.origin} to {deps.destination} on {deps.date}" - - gen = run_stream_events( - user_prompt=prompt, + user_prompt = f"帮我查询并预定在 {deps.date.strftime('%Y-%m-%d')} 从 {deps.origin_airport_code} 到 {deps.destination_airport_code} 的航班" + message_history: list[ModelMessage] = [] + events = run_stream_events( deps=deps, - message_history=msg_history, + user_prompt=user_prompt, + message_history=message_history, usage=RunUsage(), ) # 消费异步生成器 - async for ev in gen: - print(ev) + async for event in events: + print(event) + if isinstance(event, AgentRunResultEvent): + # 保存本次运行新增消息 + message_history.extend(event.result.new_messages()) asyncio.run(main()) """ +@selection_agent.tool +async def select_flight(ctx: RunContext[Deps]) -> Flight | FlightNoFound: + + result = await selection_agent.run(ctx.deps.source_material, usage=ctx.usage) + if isinstance(result.output, Flight): + if not ctx.tool_call_approved: + raise ApprovalRequired(metadata={"reason": "protected"}) + return result.output + if not isinstance(event, AgentRunResultEvent): yield event else: