This commit is contained in:
liubiren 2026-06-25 20:04:06 +08:00
parent d0eb2932dc
commit 249a38e6b9
4 changed files with 115 additions and 262 deletions

View File

@ -42,7 +42,6 @@ def render_part(dialog_id: str, part_id: str, part: Part):
PartType.TEXT, PartType.TEXT,
rx.markdown( rx.markdown(
part.content, part.content,
font_size="0.90rem",
color=rx.color("gray", 12), # 字体颜色 color=rx.color("gray", 12), # 字体颜色
background_color="transparent", # 背景颜色:设置为透明以继承父元素背景颜色 background_color="transparent", # 背景颜色:设置为透明以继承父元素背景颜色
display="block", # 布局模式:铺满 display="block", # 布局模式:铺满
@ -52,18 +51,36 @@ def render_part(dialog_id: str, part_id: str, part: Part):
margin_bottom="12px", # 底部外边距 margin_bottom="12px", # 底部外边距
key=part_id, key=part_id,
), ),
), # 文本片段 ), # 片段类型为文本
(PartType.FINISHED, rx.fragment(key=part_id)), # 结束片段 (PartType.FINISHED, rx.fragment(key=part_id)), # 片段类型为结束
rx.box( rx.box(
rx.cond( rx.cond(
part.is_open, part.is_open,
rx.hstack( rx.hstack(
rx.match(
part.part_type,
(
PartType.THINKING,
rx.text( rx.text(
"正在思考",
font_size="0.90rem",
bold=True,
color=rx.color("gray", 10), # 字体颜色
),
),
(
PartType.TOOL_NAME,
rx.text(
"正在调用",
" ",
part.content, part.content,
" ",
font_size="0.90rem", font_size="0.90rem",
bold=True, bold=True,
color=rx.color("gray", 10), color=rx.color("gray", 10),
), ),
),
),
rx.spacer(), rx.spacer(),
rx.icon("chevron_up", size=16, color=rx.color("gray", 6)), rx.icon("chevron_up", size=16, color=rx.color("gray", 6)),
width="100%", width="100%",
@ -72,7 +89,7 @@ def render_part(dialog_id: str, part_id: str, part: Part):
padding_y="0.6em", padding_y="0.6em",
background_color=rx.color("gray", 2), background_color=rx.color("gray", 2),
border_radius="8px", 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(
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)), rx.icon("chevron_down", size=16, color=rx.color("gray", 6)),
width="100%", width="100%",
cursor="pointer", 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_x="1em",
padding_y="0.6em", padding_y="0.6em",
background_color=rx.color("gray", 2), background_color=rx.color("gray", 2),
@ -103,7 +120,7 @@ def render_part(dialog_id: str, part_id: str, part: Part):
margin_bottom="10px", margin_bottom="10px",
cursor="pointer", cursor="pointer",
key=part_id, key=part_id,
), ), # 片段类型为工具相关
) )

View File

@ -47,11 +47,11 @@ class MessageHistory(SQLModel, table=True):
class PartType(StrEnum): class PartType(StrEnum):
"""片段类型""" """片段类型(适配前端渲染)"""
THINKING = "thinking" THINKING = "thinking"
TEXT = "text" TEXT = "text"
CALL = "call" TOOL_NAME = "tool_name"
TOOL_ARGS = "tool_args" TOOL_ARGS = "tool_args"
TOOL_RETURN = "tool_return" TOOL_RETURN = "tool_return"
FINISHED = "finished" 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): class Part(BaseModel):

View File

@ -124,9 +124,9 @@ class ChatState(rx.State):
self.current_chat_id = next(iter(self.chats)) self.current_chat_id = next(iter(self.chats))
@rx.event @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: 表单数据 :param form_data: 表单数据
:return: AsyncGenerator :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 user_prompt=user_prompt, message_history=message_history
): ):
# 若模型响应事件为空则跳过 # 若模型响应事件为空则跳过
@ -188,8 +188,11 @@ class ChatState(rx.State):
# 若当前片段非空则设置片段流式输出状态非正在流式输出 # 若当前片段非空则设置片段流式输出状态非正在流式输出
if current_part: if current_part:
current_part.is_streaming = False 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 current_part = new_part
# 追加片段内容 # 追加片段内容
@ -208,26 +211,27 @@ class ChatState(rx.State):
current_chat.is_streaming = False current_chat.is_streaming = False
@rx.event @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 :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: if not current_part:
return 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

View File

@ -4,191 +4,24 @@ Pydantic AI 聊天智能体和相关模块
""" """
# 列举导入模块 # 列举导入模块
from asyncio import Queue, QueueEmpty, Task, Task, create_task, sleep from enum import StrEnum
from typing import AsyncGenerator, Dict, List, Literal, Optional, Union from typing import AsyncGenerator, List, Optional, Union
from uuid import uuid4
from pydantic import Field from pydantic_ai import Agent as PydanticAIAgent
from pydantic import BaseModel
from pydantic_ai import Agent as PydanticAIAgent, ModelMessage
from pydantic_ai.capabilities import AgentCapability from pydantic_ai.capabilities import AgentCapability
from pydantic_ai.messages import ( from pydantic_ai.messages import (
AgentStreamEvent, AgentStreamEvent,
FinalResultEvent,
LoadCapabilityCallPart,
ModelMessage, ModelMessage,
NativeToolCallPart,
NativeToolReturnPart,
NativeToolSearchCallPart,
PartDeltaEvent,
PartEndEvent,
PartStartEvent, PartStartEvent,
TextPart, TextPart,
TextPartDelta,
ThinkingPart, ThinkingPart,
ThinkingPartDelta, LoadCapabilityCallPart,
ToolCallPart,
ToolCallPartDelta,
ToolReturnPart,
ToolSearchCallPart,
) )
from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.models.openai import OpenAIChatModel
from pydantic_ai.output import OutputSpec from pydantic_ai.output import OutputSpec
from pydantic_ai.providers.openai import OpenAIProvider from pydantic_ai.providers.openai import OpenAIProvider
from pydantic_ai.run import AgentRunResultEvent from pydantic_ai.run import AgentRunResultEvent
from pydantic import BaseModel, Field
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
DEFAULT_INSTRUCTIONS: str = """ 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: class Agent:
""" """
基于 Pydantic AI 封装的智能体支持 基于 Pydantic AI 封装的智能体
1流式输出模型响应事件
""" """
def __init__( def __init__(
@ -234,7 +91,7 @@ class Agent:
# 聊天唯一标识 # 聊天唯一标识
self.chat_id = chat_id self.chat_id = chat_id
# 一次聊天chat包含若干论对话dialog一轮对话由用户提示词user_prompt和输出output组成两者统称为消息message # 一次聊天chat包含若干论对话dialog轮对话包含用户提示词user_prompt和输出output。其中输出包含若干片段Part
# 本轮对话新增消息列表 # 本轮对话新增消息列表
self.new_messages: List[ModelMessage] = [] self.new_messages: List[ModelMessage] = []
@ -258,66 +115,41 @@ class Agent:
retries=retries, retries=retries,
) )
async def stream_events( async def run(
self, self,
user_prompt: str | List[str], user_prompt: str | List[str],
message_history: Optional[List[ModelMessage]] = None, message_history: Optional[List[ModelMessage]] = None,
flush_delay: float = 0.15, ) -> AsyncGenerator[Event]:
) -> AsyncGenerator[str, None]:
""" """
流式输出模型响应事件 运行
:param user_prompt: 用户提示词用户提示词 :param user_prompt: 用户提示词用户提示词
:param flush_delay: 刷新延迟时长单位为秒 :param flush_delay: 刷新延迟时长单位为秒
:yield: AsyncGenerator :yield: AsyncGenerator[Event]
""" """
# 初始化防抖器
debouncer = Debouncer(flush_delay=flush_delay)
async with self.agent.run_stream_events( async with self.agent.run_stream_events(
user_prompt=user_prompt, user_prompt=user_prompt,
message_history=message_history, message_history=message_history,
) as events: ) as events:
async for event in events: async for event in events:
# 处理模型响应事件
async for event in debouncer.handle_event(event):
match event: match event:
# 片段开始事件 case PartStartEvent():
case PartStartEvent(part=part): part = event.part
match part: match part:
case ThinkingPart(content=content):
yield f"00:{content}"
case TextPart(content=content): case TextPart(content=content):
yield f"01:{content}" yield Event(
case ( kind=Kind.TEXTSTART,
NativeToolSearchCallPart(tool_name=tool_name) part_index=event.index,
| NativeToolCallPart(tool_name=tool_name) content=content,
| ToolSearchCallPart(tool_name=tool_name) )
| ToolCallPart(tool_name=tool_name) case ThinkingPart(content=content):
| LoadCapabilityCallPart(tool_name=tool_name) yield Event(
): kind=Kind.THINKINGSTART,
yield f"02:{tool_name}" part_index=event.index,
case NativeToolReturnPart( content=content,
content=content )
) | ToolReturnPart(content=content): case LoadCapabilityCallPart():
yield f"04:{content}" yield Event(
case _: kind=Kind.TOOL_NAME,
yield f"06:未知片段类型{type(part).__name__}" part_index=event.index,
# 片段增量事件
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__}"