From 97a8bb09fe532b063249843cd160e58b99877e6a Mon Sep 17 00:00:00 2001 From: liubiren Date: Fri, 26 Jun 2026 20:01:54 +0800 Subject: [PATCH] 1 --- 产品需求文档AI生成/application/utils/agent.py | 181 ++++++++++++++---- 1 file changed, 146 insertions(+), 35 deletions(-) diff --git a/产品需求文档AI生成/application/utils/agent.py b/产品需求文档AI生成/application/utils/agent.py index c903605..9fe0662 100644 --- a/产品需求文档AI生成/application/utils/agent.py +++ b/产品需求文档AI生成/application/utils/agent.py @@ -7,16 +7,27 @@ Pydantic AI 聊天智能体和相关模块 from enum import StrEnum from typing import AsyncGenerator, List, Optional, Union from uuid import uuid4 -from pydantic_ai import Agent as PydanticAIAgent +from numpy._core.numeric import str_ +from pydantic_ai import Agent from pydantic_ai.capabilities import AgentCapability from pydantic_ai.messages import ( AgentStreamEvent, ModelMessage, PartStartEvent, - TextPart, ThinkingPart, + ToolSearchCallPart, + ThinkingPartDelta, + ToolCallPart, LoadCapabilityCallPart, + TextPartDelta, + ToolCallPartDelta, + FunctionToolResultEvent, + PartDeltaEvent, + TextPart, + PartEndEvent, ) + + from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.output import OutputSpec from pydantic_ai.providers.openai import OpenAIProvider @@ -41,32 +52,19 @@ DEFAULT_INSTRUCTIONS: str = """ """ -class Kind(StrEnum): - """种类""" - - TEXTSTART = "text_start" - THINKINGSTART = "thinking_start" - - TOOL_NAME = "tool_name" - TOOL_ARGS = "tool_args" - TOOL_RETURN = "tool_return" - FINISHED = "finished" - ERROR = "error" - - class Event(BaseModel): """ - 片段类 + 事件类 """ - part_index: Optional[int] = Field(default=None, description="片段索引") - kind: Kind = Field(..., description="事件种类") - tool_name: Optional[str] = Field(default=None, description="工具名称") - args: Optional[LoadCapabilityArgs] = Field(default=None, description="工具参数") + event_kind: str = Field(..., description="事件种类") content: Optional[str] = Field(default=None, description="事件内容") + part_index: Optional[int] = Field(default=None, description="分片索引") + part_kind: str = Field(..., description="分片种类") + tool_name: Optional[str] = Field(default=None, description="工具名称") -class Agent: +class AIAgent: """ 基于 Pydantic AI 封装的智能体 """ @@ -91,7 +89,10 @@ class Agent: # 聊天唯一标识 self.chat_id = chat_id - # 一次聊天(chat)包含若干论对话(dialog),每轮对话包含用户提示词(user_prompt)和输出(output)。其中,输出包含若干片段(Part) + # 一次聊天(chat)包含若干论对话(dialog),每轮对话包含用户提示词(user_prompt)和输出(output)。其中,输出包含若干分片(Part) + + # 本轮对话工具调用唯一标识与片段索引映射表 + self.tool_call_ids: dict[str, int] = {} # 本轮对话新增消息列表 self.new_messages: List[ModelMessage] = [] @@ -101,7 +102,7 @@ class Agent: instructions = DEFAULT_INSTRUCTIONS # 初始化智能体 - self.agent = PydanticAIAgent( + self.agent = Agent( model=OpenAIChatModel( model_name="deepseek-v4-flash", provider=OpenAIProvider( @@ -132,24 +133,134 @@ class Agent: ) as events: async for event in events: match event: - case PartStartEvent(): - part = event.part + # ========== 分片开始事件 ========== + case PartStartEvent(event_kind=event_kind, index=index, part=part): match part: - case TextPart(content=content): + # 思考分片开始事件 + case ThinkingPart(part_kind=part_kind, content=content): yield Event( - kind=Kind.TEXTSTART, - part_index=event.index, + event_kind=event_kind, content=content, + part_index=index, + part_kind=part_kind, ) - case ThinkingPart(content=content): + # 检索工具分片开始事件 + case ToolSearchCallPart( + part_kind=part_kind, tool_call_id=tool_call_id, tool_name=tool_name + ): + # 记录工具调用唯一标识与片段索引映射 + self.tool_call_ids[tool_call_id] = index yield Event( - kind=Kind.THINKINGSTART, - part_index=event.index, + event_kind=event_kind, + part_index=index, + part_kind=part_kind, + tool_name=tool_name, + ) + # 加载能力分片开始事件 + case LoadCapabilityCallPart( + part_kind=part_kind, tool_call_id=tool_call_id, tool_name=tool_name + ): + # 记录工具调用唯一标识与片段索引映射 + self.tool_call_ids[tool_call_id] = index + yield Event( + event_kind=event_kind, + part_index=index, + part_kind=part_kind, + tool_name=tool_name, + ) + # 调用工具分片开始事件 + case ToolCallPart( + part_kind=part_kind, + tool_call_id=tool_call_id, + tool_name=tool_name, + ): + # 记录工具调用唯一标识与片段索引映射 + self.tool_call_ids[tool_call_id] = index + yield Event( + event_kind=event_kind, + part_index=index, + part_kind=part_kind, + tool_name=tool_name, + ) + # 文本分片开始事件 + case TextPart(part_kind=part_kind, content=content): + yield Event( + event_kind=event_kind, content=content, + part_index=index, + part_kind=part_kind, ) - case LoadCapabilityCallPart(): - yield Event( - kind=Kind.TOOL_NAME, - part_index=event.index, + # ========== 分片增量事件(前端仅就文本片段实现打字机效果) ========== + case PartDeltaEvent( + event_kind=event_kind, index=index, delta=delta + ): + match delta: + # 文本分片增量事件 + case TextPartDelta( + part_delta_kind=part_delta_kind, + content_delta=content_delta, + ): + yield Event( + event_kind=event_kind, + content=content_delta, # 增量 + part_index=index, + part_kind=part_delta_kind, ) + + # ========== 分片结束事件 ========== + case PartEndEvent(event_kind=event_kind, index=index, part=part): + match part: + # 思考分片结束事件 + case ThinkingPart(part_kind=part_kind, content=content): + yield Event( + event_kind=event_kind, + content=content, + part_index=index, + part_kind=part_kind, + ) + # 检索工具分片结束事件 + case ToolSearchCallPart(part_kind=part_kind, tool_name=tool_name): + yield Event( + event_kind=event_kind, + part_index=index, + part_kind=part_kind, + tool_name=tool_name, + ) + # 加载能力分片结束事件 + case LoadCapabilityCallPart(part_kind=part_kind, tool_name=tool_name): + yield Event( + event_kind=event_kind, + part_index=index, + part_kind=part_kind, + tool_name=tool_name, + ) + # 调用工具分片结束事件 + case ToolCallPart(part_kind=part_kind, tool_name=tool_name): + yield Event( + event_kind=event_kind, + part_index=index, + part_kind=part_kind, + tool_name=tool_name, + ) + # 文本分片结束事件 + case TextPart(part_kind=part_kind): + yield Event( + event_kind=event_kind, + part_index=index, + part_kind=part_kind, + ) + # ========== 函数工具结果回调事件 ========== + case FunctionToolResultEvent( + event_kind=event_kind, + content=content, + part=part, + ): + yield Event( + event_kind=event_kind, + content=content, + part_index=index, + part_kind=part_kind, + tool_name=tool_name, + ) +