This commit is contained in:
parent
d0eb2932dc
commit
249a38e6b9
|
|
@ -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,18 +51,36 @@ 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.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)),
|
||||
width="100%",
|
||||
|
|
@ -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,
|
||||
),
|
||||
), # 片段类型为工具相关
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
case PartStartEvent():
|
||||
part = event.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__}"
|
||||
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,
|
||||
|
||||
)
|
||||
|
|
|
|||
Loading…
Reference in New Issue