This commit is contained in:
liubiren 2026-08-07 01:16:53 +08:00
parent ae05009acd
commit 0280261508
5 changed files with 166 additions and 177 deletions

View File

@ -25,9 +25,7 @@ from application.states.database import DatabaseState
from application.states.models import ( from application.states.models import (
Conversation, Conversation,
Dialog, Dialog,
Task, TaskNodeResultEvent,
TaskResultEvent,
TaskType,
Thought, Thought,
thoughts_to_dict, thoughts_to_dict,
usage_to_dict, usage_to_dict,
@ -52,8 +50,6 @@ class ConversationState(rx.State):
# 显示更多的会话唯一标识 # 显示更多的会话唯一标识
shown_more_conversation_id: str = "" shown_more_conversation_id: str = ""
# 任务类型
task_type: TaskType = TaskType.NONE
# 用户提示词 # 用户提示词
user_prompt: str = "" user_prompt: str = ""
@ -237,7 +233,7 @@ class ConversationState(rx.State):
tool_call_ids: dict[str, int] = {} tool_call_ids: dict[str, int] = {}
# 获取运行流式输出事件 # 获取运行流式输出事件
async for event in run_stream_events( async for event in run_stream_events(
task_type=self.task_type, task=conversation.task,
user_prompt=dialog.user_prompt, user_prompt=dialog.user_prompt,
message_history=message_history, message_history=message_history,
): ):
@ -362,7 +358,7 @@ class ConversationState(rx.State):
case "tool_call": case "tool_call":
dialog.thoughts[index].content = f"已调用 {content}" dialog.thoughts[index].content = f"已调用 {content}"
case TaskResultEvent(task=task, content=content): case TaskNodeResultEvent(task=task, content=content):
# 更新任务 # 更新任务
conversation.task = task conversation.task = task
dialog.result_output += content dialog.result_output += content
@ -373,7 +369,7 @@ class ConversationState(rx.State):
await db_state.create_dialog_record( await db_state.create_dialog_record(
conversation_id=self.conversation_id, conversation_id=self.conversation_id,
id=dialog.id, id=dialog.id,
user_prompt=self.user_prompt, user_prompt=dialog.user_prompt,
thoughts=thoughts_to_dict(dialog.thoughts), thoughts=thoughts_to_dict(dialog.thoughts),
result_output=dialog.result_output, result_output=dialog.result_output,
usage=usage_to_dict(result.usage), usage=usage_to_dict(result.usage),
@ -405,7 +401,6 @@ class ConversationState(rx.State):
conversation = self.conversations[self.conversation_id] conversation = self.conversations[self.conversation_id]
# 初始化预定航班任务 # 初始化预定航班任务
conversation.task = init_task() conversation.task = init_task()
self.user_prompt = f"帮我找一班从 {conversation.task.deps.origin}{conversation.task.deps.destination}{conversation.task.deps.date} 的航班" self.user_prompt = f"帮我找一班从 {conversation.task.deps.origin}{conversation.task.deps.destination}{conversation.task.deps.date} 的航班"
@rx.event @rx.event

View File

@ -2,12 +2,12 @@
""" """
面向 reflex.state 的类 面向 reflex.state 的类
""" """
from datetime import datetime from enum import StrEnum
from typing import Any from typing import Any
from pydantic import BaseModel, Field, TypeAdapter from pydantic import BaseModel, Field, TypeAdapter
from pydantic_ai import RunUsage from pydantic_ai import RunUsage
from pydantic_ai._uuid import uuid7 from pydantic_ai._uuid import uuid7
from enum import StrEnum
from pydantic_ai.usage import UsageLimits from pydantic_ai.usage import UsageLimits
@ -53,14 +53,14 @@ class Dialog(BaseModel):
UsageAdapter = TypeAdapter(RunUsage) UsageAdapter = TypeAdapter(RunUsage)
def usage_to_object(usage: dict) -> RunUsage: def usage_to_object(usage: dict[str, Any]) -> RunUsage:
""" """
Usage 转为对象 Usage 转为对象
""" """
return UsageAdapter.validate_python(usage) if usage else RunUsage() return UsageAdapter.validate_python(usage) if usage else RunUsage()
def usage_to_dict(usage: RunUsage) -> dict: def usage_to_dict(usage: RunUsage) -> dict[str, Any]:
""" """
Usage 转为字典 Usage 转为字典
""" """
@ -71,7 +71,7 @@ def usage_to_dict(usage: RunUsage) -> dict:
UsageLimitsAdapter = TypeAdapter(UsageLimits) UsageLimitsAdapter = TypeAdapter(UsageLimits)
def usage_limits_to_object(usage_limits: dict) -> UsageLimits: def usage_limits_to_object(usage_limits: dict[str, Any]) -> UsageLimits:
""" """
UsageLimits 转为对象 UsageLimits 转为对象
""" """
@ -82,7 +82,7 @@ def usage_limits_to_object(usage_limits: dict) -> UsageLimits:
) )
def usage_limits_to_dict(usage_limits: UsageLimits) -> dict: def usage_limits_to_dict(usage_limits: UsageLimits) -> dict[str, Any]:
""" """
UsageLimits 转为字典 UsageLimits 转为字典
""" """
@ -94,20 +94,9 @@ class TaskType(StrEnum):
任务类型枚举 任务类型枚举
""" """
NONE = "none"
BOOK_FLIGHT = "book_flight" BOOK_FLIGHT = "book_flight"
class TaskNode(StrEnum):
"""
任务节点枚举
"""
EXECUTION = "execution"
PENDING_INPUT = "pending_input"
FINISH = "finish"
class Task(BaseModel): class Task(BaseModel):
""" """
任务类 任务类
@ -115,7 +104,7 @@ class Task(BaseModel):
id: str = Field(default_factory=lambda: str(uuid7()), description="任务唯一标识") id: str = Field(default_factory=lambda: str(uuid7()), description="任务唯一标识")
type: TaskType = Field(..., description="任务类型") type: TaskType = Field(..., description="任务类型")
node: TaskNode = Field(default=TaskNode.EXECUTION, description="任务节点") node: str | None = Field(default="", description="任务节点")
deps: Any = Field(default=None, description="任务依赖项") deps: Any = Field(default=None, description="任务依赖项")
usage: dict[str, Any] = Field(default_factory=dict, description="任务使用量") usage: dict[str, Any] = Field(default_factory=dict, description="任务使用量")
usage_limits: dict[str, Any] = Field( usage_limits: dict[str, Any] = Field(
@ -123,15 +112,6 @@ class Task(BaseModel):
) )
class TaskResultEvent(BaseModel):
"""
任务结果事件类
"""
task: Task = Field(..., description="任务实例")
content: str = Field(default="", description="任务结果内容")
class Conversation(BaseModel): class Conversation(BaseModel):
""" """
会话类 会话类
@ -149,3 +129,12 @@ class Conversation(BaseModel):
..., ...,
description="创建时间", description="创建时间",
) )
class TaskNodeResultEvent(BaseModel):
"""
任务节点结果事件类
"""
task: Task | None = Field(..., description="任务实例")
content: str = Field(default="", description="任务结果内容")

View File

@ -2,17 +2,15 @@
""" """
任务模块 任务模块
""" """
from typing import Any, AsyncGenerator, Callable, Dict, Optional
from typing import AsyncGenerator, List from typing import AsyncGenerator, List
from pydantic import BaseModel
from pydantic_ai import Agent, ModelMessage from pydantic_ai import Agent, ModelMessage
from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.messages import AgentStreamEvent
from pydantic_ai.providers.openai import OpenAIProvider from pydantic_ai.run import AgentRunResultEvent
from application.states.models import Dialog, TaskType from application.states.models import TaskNodeResultEvent, TaskType, Task
from application.tasks.models import DEEPSEEK_V4_FLASH_MODEL from application.tasks.models import DEEPSEEK_V4_FLASH_MODEL
from application.states.models import TaskNode, Task
instruction = """ instruction = """
# 角色 # 角色
@ -32,15 +30,26 @@ instruction = """
async def run_stream_events( async def run_stream_events(
task_type: TaskType, task: Task | None,
user_prompt: str, user_prompt: str,
message_history: List[ModelMessage], message_history: List[ModelMessage],
) -> AsyncGenerator: ) -> AsyncGenerator[AgentStreamEvent | AgentRunResultEvent | TaskNodeResultEvent, None]:
""" """
以流式事件模式运行 以流式事件模式运行
""" """
match task_type: if task:
case TaskType.NONE: match task.type:
case TaskType.BOOK_FLIGHT:
from application.tasks.book_flight import run_stream_events
async for event in run_stream_events(
task=task,
user_prompt=user_prompt,
message_history=message_history,
):
yield event
else:
agent = Agent( agent = Agent(
model=DEEPSEEK_V4_FLASH_MODEL, model=DEEPSEEK_V4_FLASH_MODEL,
instructions=instruction, instructions=instruction,
@ -51,13 +60,3 @@ async def run_stream_events(
) as events: ) as events:
async for event in events: async for event in events:
yield event yield event
case TaskType.BOOK_FLIGHT:
from application.tasks.book_flight import run_stream_events, init_task
async for event in run_stream_events(
task=init_task(),
user_prompt=user_prompt,
message_history=message_history,
):
yield event

View File

@ -1,42 +1,40 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
""" """
生成产品需求文档智能体 预定航班任务
""" """
import datetime import datetime
from typing import AsyncGenerator, Literal from typing import AsyncGenerator, Literal
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from pydantic_ai import ( from pydantic_ai import Agent, ModelMessage, ModelRetry, RunContext, UsageLimits
Agent,
ModelMessage,
ModelRetry,
RunContext,
UsageLimits,
ModelMessage,
)
from pydantic_ai.run import AgentRunResultEvent from pydantic_ai.run import AgentRunResultEvent
from application.tasks.models import DEEPSEEK_V4_FLASH_MODEL
from application.states.models import ( from application.states.models import (
TaskType,
Task, Task,
TaskNode, TaskNodeResultEvent,
usage_to_dict, TaskType,
usage_to_object,
TaskResultEvent,
usage_limits_to_dict, usage_limits_to_dict,
usage_limits_to_object, usage_limits_to_object,
usage_to_dict,
usage_to_object,
) )
from application.tasks.models import DEEPSEEK_V4_FLASH_MODEL
class Deps(BaseModel): class Deps(BaseModel):
flights: str = Field(..., description="航班信息") """
依赖项类
"""
flight_info: str = Field(..., description="所有航班信息")
date: datetime.date = Field(..., description="航班日期") date: datetime.date = Field(..., description="航班日期")
origin: str = Field(..., description="出发机场") origin: str = Field(..., description="出发机场")
destination: str = Field(..., description="到达机场") destination: str = Field(..., description="到达机场")
class FlightDetails(BaseModel): class FlightDetail(BaseModel):
""" """
航班详情 航班详情
""" """
number: str = Field(description="航班号") number: str = Field(description="航班号")
@ -48,83 +46,82 @@ class FlightDetails(BaseModel):
class NoFlightFound(BaseModel): class NoFlightFound(BaseModel):
""" """
未查询到航班 未查询到航班
""" """
# 主控智能体
master_agent = Agent[Deps, FlightDetails | NoFlightFound](
model=DEEPSEEK_V4_FLASH_MODEL,
deps_type=Deps,
output_type=FlightDetails | NoFlightFound,
retries=2,
system_prompt=("你的工作是在给定日期为用户找到最便宜的航班"),
)
# 航班查询智能体 # 航班查询智能体
search_agent = Agent( flight_search_agent = Agent[Deps, FlightDetail | NoFlightFound](
model=DEEPSEEK_V4_FLASH_MODEL, model=DEEPSEEK_V4_FLASH_MODEL,
output_type=list[FlightDetails], deps_type=Deps,
system_prompt="从给定文本中提取所有航班详细信息", output_type=FlightDetail | NoFlightFound,
system_prompt="你的工作是在给定日期、出发机场和到达机场为用户找到最便宜的航班",
)
# 所有航班详情提取智能体
flight_details_extraction_agent = Agent(
model=DEEPSEEK_V4_FLASH_MODEL,
output_type=list[FlightDetail],
system_prompt="从给定文本中提取所有航班详细,包括航班号、航班日期、出发机场、到达机场和机票价格",
) )
@master_agent.tool @flight_search_agent.tool
async def search_flights(ctx: RunContext[Deps]) -> list[FlightDetails]: async def extract_flight_details(ctx: RunContext[Deps]) -> list[FlightDetail]:
""" """
查询并返回航班详情列表 工具提取所有航班详情
""" """
result = await search_agent.run(ctx.deps.flights, usage=ctx.usage) result = await flight_details_extraction_agent.run(
ctx.deps.flight_info, usage=ctx.usage
)
return result.output return result.output
@master_agent.output_validator @flight_search_agent.output_validator
async def validate_output( async def validate_output(
ctx: RunContext[Deps], output: FlightDetails | NoFlightFound ctx: RunContext[Deps], output: FlightDetail | NoFlightFound
) -> FlightDetails | NoFlightFound: ) -> FlightDetail | NoFlightFound:
""" """
校验主控智能体输出 输出校验航班详情
""" """
if isinstance(output, NoFlightFound): if isinstance(output, NoFlightFound):
return output return output
errors = "" errors = []
if output.date != ctx.deps.date: if output.date != ctx.deps.date:
errors += f"航班日期应为 {ctx.deps.date}, 不是 {output.date}\n" errors.append(f"航班日期应为 {ctx.deps.date}, 不是 {output.date}")
if output.origin != ctx.deps.origin: if output.origin != ctx.deps.origin:
errors += f"航班出发机场应为 {ctx.deps.origin}, 不是 {output.origin}\n" errors.append(f"出发机场应为 {ctx.deps.origin}, 不是 {output.origin}")
if output.destination != ctx.deps.destination: if output.destination != ctx.deps.destination:
errors += ( errors.append(f"到达机场应为 {ctx.deps.destination}, 不是 {output.destination}")
f"航班到达机场应为 {ctx.deps.destination}, 不是 {output.destination}\n"
)
if errors: if errors:
raise ModelRetry(errors) raise ModelRetry("\n".join(errors))
return output return output
class SeatPreference(BaseModel): class SeatPreference(BaseModel):
row: int = Field(ge=1, le=30) """
seat: Literal["A", "B", "C", "D", "E", "F"] 座位偏好类
"""
row: int = Field(ge=1, le=30, description="座位行")
column: Literal["A", "B", "C", "D", "E", "F"] = Field(description="座位列")
class Failed(BaseModel): class NoSeatExtracted(BaseModel):
"""Unable to extract a seat selection.""" """
未提取到座位偏好类
"""
# 选座智能体 # 提取座位偏好智能体(无依赖项)
seat_selection_agent = Agent[object, SeatPreference | Failed]( seat_preference_extraction_agent = Agent[object, SeatPreference | NoSeatExtracted](
model=DEEPSEEK_V4_FLASH_MODEL, model=DEEPSEEK_V4_FLASH_MODEL,
output_type=SeatPreference | Failed, output_type=SeatPreference | NoSeatExtracted,
system_prompt=( system_prompt="提取用户的座位偏好。座位规则说明A 座、F 座为靠窗座位;第 1 排是前排座位腿部空间更大14 排、20 排同样拥有加宽腿部空间",
"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. "
),
) )
flights = """ flight_info = """
1. Flight SFO-AK123 1. Flight SFO-AK123
- Price: $350 - Price: $350
- Origin: San Francisco International Airport (SFO) - Origin: San Francisco International Airport (SFO)
@ -171,9 +168,9 @@ flights = """
def init_task() -> Task: def init_task() -> Task:
return Task( return Task(
type=TaskType.BOOK_FLIGHT, type=TaskType.BOOK_FLIGHT,
node=TaskNode.EXECUTION, node="flight_search",
deps=Deps( deps=Deps(
flights=flights, flight_info=flight_info,
date=datetime.date(2025, 1, 10), date=datetime.date(2025, 1, 10),
origin="SFO", origin="SFO",
destination="ANC", destination="ANC",
@ -183,15 +180,19 @@ def init_task() -> Task:
async def run_stream_events( async def run_stream_events(
task: Task, task: Task | None,
user_prompt: str | None = None, user_prompt: str,
message_history: list[ModelMessage] | None = None, message_history: list[ModelMessage],
) -> AsyncGenerator: ) -> AsyncGenerator:
result = None result = None
while True: while True:
if task.node == TaskNode.EXECUTION: if not task:
async with master_agent.run_stream_events( return
match task.node:
# 航班查询
case "flight_search":
async with flight_search_agent.run_stream_events(
user_prompt=user_prompt, user_prompt=user_prompt,
deps=task.deps, deps=task.deps,
message_history=message_history, message_history=message_history,
@ -205,35 +206,39 @@ async def run_stream_events(
result = event.result result = event.result
# 更新任务使用量 # 更新任务使用量
task.usage = usage_to_dict(result.usage) task.usage = usage_to_dict(result.usage)
if isinstance(result.output, FlightDetails): if isinstance(result.output, FlightDetail):
content = "\n---\n已查询到航班,请回复 buy 购票 / search 重新查询\n" content = "\n---\n已查询到航班,请回复 buy 购票 search 重新查询\n"
# 更新任务节点为待用户输入 # 更新任务节点为提取座位偏好
task.node = TaskNode.PENDING_INPUT task.node = "seat_preference_extraction"
else: else:
content = "\n---\n未找到符合条件的航班,流程结束\n" content = (
# 更新任务节点为结束 "\n---\n未查询到满足您需求的航班,流程结束!\n"
task.node = TaskNode.FINISH )
yield TaskResultEvent( # 更新任务为空
task = None
# 返回任务节点结果事件
yield TaskNodeResultEvent(
task=task, task=task,
content=content, content=content,
) )
yield event yield event
return return
# 提取座位偏好
if task.node == TaskNode.PENDING_INPUT: case "seat_preference_extraction":
if user_prompt == "buy": if user_prompt == "buy":
yield TaskResultEvent( # 返回任务节点结果事件
yield TaskNodeResultEvent(
task=task, task=task,
content="请和我说下你的座位偏好吧\nA、F 座位是靠窗位1 排、14 排、20 排腿部空间更大、更舒展,你更想要靠窗座位,宽敞大空间座位", content="请和我说下您的座位偏好\nA、F 座位是靠窗位1 排、14 排、20 排腿部空间更大、更舒展,你更想要靠窗座位,宽敞大空间座位",
) )
return return
elif user_prompt == "search": elif user_prompt == "search":
# 更新任务节点为执行 # 更新任务节点为航班查询
task.node = TaskNode.EXECUTION task.node = "flight_search"
else: else:
async with seat_selection_agent.run_stream_events( async with seat_preference_extraction_agent.run_stream_events(
user_prompt=user_prompt, user_prompt=user_prompt,
message_history=message_history, message_history=message_history,
usage=usage_to_object(task.usage), usage=usage_to_object(task.usage),
@ -246,11 +251,12 @@ async def run_stream_events(
result = event.result result = event.result
# 更新任务使用量 # 更新任务使用量
task.usage = usage_to_dict(result.usage) task.usage = usage_to_dict(result.usage)
# 更新任务节点为结束 # 更新任务为空
task.node = TaskNode.FINISH task = None
yield TaskResultEvent( # 返回任务节点结果事件
yield TaskNodeResultEvent(
task=task, task=task,
content="为您预定好座位,流程结束", content="帮您定好座位,流程结束!",
) )
yield event yield event
return return

Binary file not shown.