diff --git a/产品需求文档AI生成/application/components/chat.py b/产品需求文档AI生成/application/components/chat.py index 76e9223..1aded1c 100644 --- a/产品需求文档AI生成/application/components/chat.py +++ b/产品需求文档AI生成/application/components/chat.py @@ -42,7 +42,6 @@ def render_part(dialog_id: str, part_id: str, part: Part): PartType.TEXT, rx.markdown( part.content, - font_size="0.90rem", color=rx.color("gray", 12), # 字体颜色 background_color="transparent", # 背景颜色:设置为透明以继承父元素背景颜色 display="block", # 布局模式:铺满 @@ -52,17 +51,35 @@ def render_part(dialog_id: str, part_id: str, part: Part): margin_bottom="12px", # 底部外边距 key=part_id, ), - ), # 文本片段 - (PartType.FINISHED, rx.fragment(key=part_id)), # 结束片段 + ), # 片段类型为文本 + (PartType.FINISHED, rx.fragment(key=part_id)), # 片段类型为结束 rx.box( rx.cond( part.is_open, rx.hstack( - rx.text( - part.content, - font_size="0.90rem", - bold=True, - color=rx.color("gray", 10), + rx.match( + part.part_type, + ( + PartType.THINKING, + rx.text( + "正在思考", + font_size="0.90rem", + bold=True, + color=rx.color("gray", 10), # 字体颜色 + ), + ), + ( + PartType.TOOL_NAME, + rx.text( + "正在调用", + " ", + part.content, + " ", + font_size="0.90rem", + bold=True, + color=rx.color("gray", 10), + ), + ), ), rx.spacer(), rx.icon("chevron_up", size=16, color=rx.color("gray", 6)), @@ -72,7 +89,7 @@ def render_part(dialog_id: str, part_id: str, part: Part): padding_y="0.6em", background_color=rx.color("gray", 2), border_radius="8px", - on_click=ChatState.toggle_collapse_panel(dialog_id, part_id), + on_click=lambda: ChatState.toggle_part_collapse(dialog_id, part_id), ), # 折叠面板打开时标题栏 rx.hstack( rx.hstack( @@ -90,7 +107,7 @@ def render_part(dialog_id: str, part_id: str, part: Part): rx.icon("chevron_down", size=16, color=rx.color("gray", 6)), width="100%", cursor="pointer", - on_click=ChatState.toggle_collapse_panel(dialog_id, part_id), + on_click=lambda: ChatState.toggle_part_collapse(dialog_id, part_id), padding_x="1em", padding_y="0.6em", background_color=rx.color("gray", 2), @@ -103,7 +120,7 @@ def render_part(dialog_id: str, part_id: str, part: Part): margin_bottom="10px", cursor="pointer", key=part_id, - ), + ), # 片段类型为工具相关 ) diff --git a/产品需求文档AI生成/application/models.py b/产品需求文档AI生成/application/models.py index 9f8b81a..6c3cfaf 100644 --- a/产品需求文档AI生成/application/models.py +++ b/产品需求文档AI生成/application/models.py @@ -47,11 +47,11 @@ class MessageHistory(SQLModel, table=True): class PartType(StrEnum): - """片段类型""" + """片段类型(适配前端渲染)""" THINKING = "thinking" TEXT = "text" - CALL = "call" + TOOL_NAME = "tool_name" TOOL_ARGS = "tool_args" TOOL_RETURN = "tool_return" FINISHED = "finished" @@ -59,7 +59,7 @@ class PartType(StrEnum): # 动态生成前缀和片段类型映射表 -PREFIX_MAPING = {f"{i:02d}:": t for i, t in enumerate(PartType)} +PREFIX_MAPING = {f"{i:02d}": t for i, t in enumerate(PartType)} class Part(BaseModel): diff --git a/产品需求文档AI生成/application/state/chat.py b/产品需求文档AI生成/application/state/chat.py index c6744e7..736ff5c 100644 --- a/产品需求文档AI生成/application/state/chat.py +++ b/产品需求文档AI生成/application/state/chat.py @@ -124,9 +124,9 @@ class ChatState(rx.State): self.current_chat_id = next(iter(self.chats)) @rx.event - async def process_input(self, form_data: dict[str, Any]) -> AsyncGenerator: + async def run(self, form_data: dict[str, Any]) -> AsyncGenerator: """ - 处理输入,返回流式输出 + 运行 :param form_data: 表单数据 :return: AsyncGenerator """ @@ -162,7 +162,7 @@ class ChatState(rx.State): ) # 流式输出模型响应事件 - async for event in agent.stream_events( + async for event in agent.run( user_prompt=user_prompt, message_history=message_history ): # 若模型响应事件为空则跳过 @@ -188,8 +188,11 @@ class ChatState(rx.State): # 若当前片段非空则设置片段流式输出状态非正在流式输出 if current_part: current_part.is_streaming = False - new_part = Part(part_type=part_type, is_streaming=True) - current_dialog.output[uuid4().hex] = new_part + + # 新增片段 + current_dialog.output[uuid4().hex] = ( + new_part := Part(part_type=part_type, is_streaming=True) + ) current_part = new_part # 追加片段内容 @@ -208,26 +211,27 @@ class ChatState(rx.State): current_chat.is_streaming = False @rx.event - async def toggle_collapse_panel(self, dialog_id: str, part_id: str) -> None: + def toggle_part_collapse(self, dialog_id: str, part_id: str) -> None: """ - 打开/关闭折叠面板 + 打开/关闭片段折叠面板 + :param dialog_id: 对话唯一标识 + :param part_id: 片段唯一标识 :return: None """ + # 当前聊天 + current_chat = self.chats.get(self.current_chat_id) + if not current_chat: + return + + # 当前对话 + current_dialog = current_chat.dialogs.get(dialog_id) + if not current_dialog: + return + # 当前片段 - current_part = self.chats.get(self.current_chat_id, {}).dialogs.get(dialog_id, {}).get(part_id, None) + current_part = current_dialog.output.get(part_id) if not current_part: return - # 目标对话 - target_dialog = next( - (d for d in current_chat.dialogs if d.id == dialog_id), None - ) - if not target_dialog: - return - # 目标消息 - target_message = next( - (m for m in target_dialog.output if m.id == message_id), None - ) - if not target_message: - return + # 切换折叠面板打开状态 - target_message.is_open = not target_message.is_open + current_part.is_open = not current_part.is_open diff --git a/产品需求文档AI生成/application/utils/agent.py b/产品需求文档AI生成/application/utils/agent.py index fcb1e95..c903605 100644 --- a/产品需求文档AI生成/application/utils/agent.py +++ b/产品需求文档AI生成/application/utils/agent.py @@ -4,191 +4,24 @@ Pydantic AI 聊天智能体和相关模块 """ # 列举导入模块 -from asyncio import Queue, QueueEmpty, Task, Task, create_task, sleep -from typing import AsyncGenerator, Dict, List, Literal, Optional, Union - -from pydantic import Field -from pydantic import BaseModel -from pydantic_ai import Agent as PydanticAIAgent, ModelMessage +from enum import StrEnum +from typing import AsyncGenerator, List, Optional, Union +from uuid import uuid4 +from pydantic_ai import Agent as PydanticAIAgent from pydantic_ai.capabilities import AgentCapability from pydantic_ai.messages import ( AgentStreamEvent, - FinalResultEvent, - LoadCapabilityCallPart, ModelMessage, - NativeToolCallPart, - NativeToolReturnPart, - NativeToolSearchCallPart, - PartDeltaEvent, - PartEndEvent, PartStartEvent, TextPart, - TextPartDelta, ThinkingPart, - ThinkingPartDelta, - ToolCallPart, - ToolCallPartDelta, - ToolReturnPart, - ToolSearchCallPart, + LoadCapabilityCallPart, ) from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.output import OutputSpec from pydantic_ai.providers.openai import OpenAIProvider from pydantic_ai.run import AgentRunResultEvent - - -class PartBuffer(BaseModel): - """ - 模型响应片段缓冲类 - """ - - part_type: Literal["thinking", "text"] = Field( - ..., description="片段类型,仅支持思考和文本片段" - ) - part_content_delta: str = Field(default="", description="片段内容增量") - flush_task: Optional[Task] = Field( - default=None, description="片段绑定的延迟刷新任务" - ) - model_config = {"arbitrary_types_allowed": True} # 允许任意类型 - - -class Debouncer: - """ - 模型响应事件防抖器 - 用于平滑 run_stream_events 返回的高频模型响应事件,降低前端渲染频率 - """ - - def __init__(self, flush_delay: float): - """ - 初始化 - :param flush_delay: 刷新延迟时长(单位为秒) - """ - self.flush_delay = flush_delay - - # 片段缓冲字典(键为片段索引,值为片段缓冲实例) - self.part_buffers: Dict[int, PartBuffer] = {} - # 待推送刷新队列 - self.flush_queue = Queue() - - async def handle_event( - self, event: Union[AgentStreamEvent, AgentRunResultEvent] - ) -> AsyncGenerator[Union[AgentStreamEvent, AgentRunResultEvent]]: - """ - 处理模型响应事件 - 在 Pydantic-AI 中模型流式输出包含若干类型片段,例如思考片段、文本片段等。每类型片段包含若干模型响应事件,例如片段开始事件、片段增量事件、片段结束事件等 - :param event: 模型响应事件 - :yield: AsyncGenerator - """ - # 优先推送队列中积压的刷新任务 - while True: - try: - yield self.flush_queue.get_nowait() - except QueueEmpty: - break - - # 若为思考或文本的片段开始事件则初始化片段缓冲,其它事件透传 - if isinstance(event, PartStartEvent): - part_index, part = event.index, event.part # 片段索引、片段 - match part: - case ThinkingPart(): - part_type = "thinking" - case TextPart(): - part_type = "text" - case _: - yield event - return - # 初始化片段缓冲 - self.part_buffers[part_index] = PartBuffer( - part_type=part_type, - ) - yield event - return - - # 若为思考或文本片段增量事件则先取消延迟刷新任务,更新片段缓冲中的片段内容增量并重新绑定延迟刷新任务,其它事件透传 - if isinstance(event, PartDeltaEvent): - part_index, delta = event.index, event.delta - part_buffer = self.part_buffers.get(part_index) - if part_buffer and isinstance(delta, (ThinkingPartDelta, TextPartDelta)): - # 取消延迟刷新任务 - self._cancel_task(part_index=part_index) - # 更新片段缓冲中的片段内容增量 - part_buffer.part_content_delta += delta.content_delta or "" - # 重新绑定延迟刷新任务 - part_buffer.flush_task = create_task( - coro=self._flush_after_delay(part_index=part_index) - ) - else: - yield event - return - - # 若为其它事件则先刷新所有片段缓冲再透传当前事件 - async for flush_event in self._flush_all(): - yield flush_event - # 若为片段结束事件则删除相应片段缓冲 - if isinstance(event, PartEndEvent): - self.part_buffers.pop(event.index, None) - yield event - - def _cancel_task(self, part_index: int) -> None: - """ - 取消延迟刷新任务 - :param part_index: 片段索引 - :return: None - """ - part_buffer = self.part_buffers.get(part_index) - if not part_buffer or not part_buffer.flush_task: - return - - # 若延迟刷新任务未推送则取消 - if not part_buffer.flush_task.done(): - part_buffer.flush_task.cancel() - part_buffer.flush_task = None - - async def _flush_after_delay(self, part_index: int) -> None: - """ - 延迟刷新 - :param part_index: 片段索引 - :return: None - """ - await sleep(delay=self.flush_delay) - - # 刷新单个片段缓冲 - flush_event = await self._flush(part_index=part_index) - if flush_event: - await self.flush_queue.put(item=flush_event) - - async def _flush(self, part_index: int) -> Optional[PartDeltaEvent]: - """ - 刷新单个片段缓冲 - :param part_index: 片段索引 - :return: Optional[PartDeltaEvent] - """ - part_buffer = self.part_buffers.get(part_index) - if not part_buffer or not part_buffer.part_content_delta: - return - - # 构建片段增量事件中增量 - match part_buffer.part_type: - case "thinking": - delta = ThinkingPartDelta(content_delta=part_buffer.part_content_delta) - case "text": - delta = TextPartDelta(content_delta=part_buffer.part_content_delta) - - # 重置片段缓冲中片段内容增量 - part_buffer.part_content_delta = "" - # 重置片段缓冲中延迟刷新任务 - part_buffer.flush_task = None - return PartDeltaEvent(index=part_index, delta=delta) - - async def _flush_all(self) -> AsyncGenerator[PartDeltaEvent]: - """ - 刷新所有片段缓冲 - :yield: AsyncGenerator - """ - for part_index in list(self.part_buffers.keys()): - flush_event = await self._flush(part_index=part_index) - if flush_event: - yield flush_event +from pydantic import BaseModel, Field DEFAULT_INSTRUCTIONS: str = """ @@ -208,10 +41,34 @@ 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="工具参数") + content: Optional[str] = Field(default=None, description="事件内容") + + class Agent: """ - 基于 Pydantic AI 封装的智能体,支持: - 1、流式输出模型响应事件 + 基于 Pydantic AI 封装的智能体 """ def __init__( @@ -234,7 +91,7 @@ class Agent: # 聊天唯一标识 self.chat_id = chat_id - # 一次聊天(chat)包含若干论对话(dialog),每一轮对话由用户提示词(user_prompt)和输出(output)组成,两者统称为消息(message) + # 一次聊天(chat)包含若干论对话(dialog),每轮对话包含用户提示词(user_prompt)和输出(output)。其中,输出包含若干片段(Part) # 本轮对话新增消息列表 self.new_messages: List[ModelMessage] = [] @@ -258,66 +115,41 @@ class Agent: retries=retries, ) - async def stream_events( + async def run( self, user_prompt: str | List[str], message_history: Optional[List[ModelMessage]] = None, - flush_delay: float = 0.15, - ) -> AsyncGenerator[str, None]: + ) -> AsyncGenerator[Event]: """ - 流式输出模型响应事件 + 运行 :param user_prompt: 用户提示词(用户提示词) :param flush_delay: 刷新延迟时长(单位为秒) - :yield: AsyncGenerator + :yield: AsyncGenerator[Event] """ - # 初始化防抖器 - debouncer = Debouncer(flush_delay=flush_delay) - async with self.agent.run_stream_events( user_prompt=user_prompt, message_history=message_history, ) as events: async for event in events: - # 处理模型响应事件 - async for event in debouncer.handle_event(event): - match event: - # 片段开始事件 - case PartStartEvent(part=part): - match part: - case ThinkingPart(content=content): - yield f"00:{content}" - case TextPart(content=content): - yield f"01:{content}" - case ( - NativeToolSearchCallPart(tool_name=tool_name) - | NativeToolCallPart(tool_name=tool_name) - | ToolSearchCallPart(tool_name=tool_name) - | ToolCallPart(tool_name=tool_name) - | LoadCapabilityCallPart(tool_name=tool_name) - ): - yield f"02:{tool_name}" - case NativeToolReturnPart( - content=content - ) | ToolReturnPart(content=content): - yield f"04:{content}" - case _: - yield f"06:未知片段类型{type(part).__name__}" - # 片段增量事件 - case PartDeltaEvent(delta=delta): - match delta: - case ThinkingPartDelta(content_delta=content_delta): - yield f"00:{content_delta}" - case TextPartDelta(content_delta=content_delta): - yield f"01:{content_delta}" - case ToolCallPartDelta(args_delta=args_delta): - yield f"03:{args_delta}" - # 片段结束事件、最终结果事件无需处理 - case PartEndEvent(): - continue - case FinalResultEvent(): - continue - case AgentRunResultEvent(result=result): - self.new_messages = result.new_messages() - yield "05:" - case _: - yield f"06:未知事件类型{type(event).__name__}" + match event: + case PartStartEvent(): + part = event.part + match part: + case TextPart(content=content): + yield Event( + kind=Kind.TEXTSTART, + part_index=event.index, + content=content, + ) + case ThinkingPart(content=content): + yield Event( + kind=Kind.THINKINGSTART, + part_index=event.index, + content=content, + ) + case LoadCapabilityCallPart(): + yield Event( + kind=Kind.TOOL_NAME, + part_index=event.index, + + )