This commit is contained in:
liubiren 2026-06-17 20:12:56 +08:00
parent ad71e4b16d
commit 4e65dc783f
1 changed files with 118 additions and 88 deletions

View File

@ -11,7 +11,7 @@ from uuid import uuid4
from pydantic_ai.messages import ModelMessage from pydantic_ai.messages import ModelMessage
from pydantic_ai import Agent as PydanticAIAgent, ModelMessage from pydantic_ai import Agent as PydanticAIAgent, ModelMessage
from pydantic_ai.capabilities import AgentCapability from pydantic_ai.capabilities import AgentCapability
from asyncio import Task, create_task, sleep, Queue, QueueEmpty from asyncio import Task, create_task, sleep, Queue, QueueEmpty, Task
from pydantic import Field from pydantic import Field
from pydantic import BaseModel from pydantic import BaseModel
@ -20,6 +20,7 @@ from pydantic import BaseModel
from pydantic_ai.messages import ( from pydantic_ai.messages import (
ModelMessagesTypeAdapter, ModelMessagesTypeAdapter,
ModelResponse, ModelResponse,
FinalResultEvent,
ModelMessageEvent, ModelMessageEvent,
PartStartEvent, PartStartEvent,
PartDeltaEvent, PartDeltaEvent,
@ -133,44 +134,44 @@ class AgentMemory(SQLite):
) from exception ) from exception
class PendingFlushPart(BaseModel): class Buffer(BaseModel):
""" """
待刷新模型消息片段 模型消息事件缓冲模型
""" """
type_: Literal["thinking", "text"] = Field(..., description="片段类型") part_type: Literal["thinking", "text"] = Field(..., description="片段类型")
content_delta: str = Field(default="", description="片段内容增量") content_delta: str = Field(default="", description="片段内容增量")
task: Optional[Task] = Field(default=None, description="异步协程任务") task: Optional[Task] = Field(default=None, description="延迟刷新异步协程任务")
class EventsDebouncer: class EventsDebouncer:
""" """
模型消息事件防抖器 模型消息事件防抖器
使用 Pydantic-AI run_stream_events 方法时需就其返回的模型消息事件处理延迟刷新片段类型为思考或文本的片段增量事件以实现防抖打字机效果 使用 Pydantic-AI run_stream_events 方法时需就其返回的模型消息事件处理延迟刷新片段类型为思考或文本的片段增量事件以实现防抖
Pydantic-AI 中模型流式输出包含若干片段例如思考片段文本片段等片段包含若干模型消息事件例如片段开始事件片段增量事件片段结束事件等 Pydantic-AI 中模型流式输出包含若干类型片段例如思考片段文本片段等类型片段包含若干模型消息事件例如片段开始事件片段增量事件片段结束事件等
""" """
def __init__(self, flush_delay: float = 0.2): def __init__(self, delay: float = 0.2):
""" """
初始化 初始化
:param flush_delay: 刷新延迟单位为秒 :param delay: 延迟时长单位为秒停顿超该时长则刷新
""" """
self.flush_delay = flush_delay self.delay = delay
# 待刷新模型响应片段 # 模型消息事件缓冲字典(数据类型为字典,键为片段索引,值为缓冲模型)
self.pending_flush_parts: Dict[int, PendingFlushPart] = {} self.buffers: Dict[int, Buffer] = {}
# 待刷新列队 # 待刷新异步队列
self.pending_flush_queue = Queue() self.pending_flush_queue = Queue()
async def handle_event( async def handle_event(
self, event: ModelMessageEvent self, event: ModelMessageEvent
) -> AsyncGenerator[ModelMessageEvent, None]: ) -> AsyncGenerator[ModelMessageEvent]:
""" """
处理模型消息事件 处理模型消息事件
:param event: 模型消息事件 :param event: 模型消息事件
:yield: AsyncGenerator :yield: AsyncGenerator
""" """
# 返回待刷新列队中模型消息事件 # 返回待刷新异步队列中模型消息事件
while True: while True:
try: try:
yield self.pending_flush_queue.get_nowait() yield self.pending_flush_queue.get_nowait()
@ -182,99 +183,102 @@ class EventsDebouncer:
index, part = event.index, event.part # 片段索引、片段 index, part = event.index, event.part # 片段索引、片段
match part: match part:
case ThinkingPart(): case ThinkingPart():
type_ = "thinking" part_type = "thinking"
case TextPart(): case TextPart():
type_ = "text" part_type = "text"
case _: case _:
yield event yield event
return return
# 新增片段缓存 # 新增片段缓存
self.pending_flush_parts[index] = PendingFlushPart( self.buffers[index] = Buffer(
type_=type_, part_type=part_type,
) )
yield event yield event
return return
# 若为思考或文本的片段增量事件则更新缓存,其它事件直接返回 """
原理
收到思考或文本片段增量事件时取消未完成的延迟刷新任务追加增量并重设若模型持续输出则先缓存停顿超过阈值再返回实现防抖
"""
if isinstance(event, PartDeltaEvent): if isinstance(event, PartDeltaEvent):
index, delta = event.index, event.delta index, delta = event.index, event.delta
pending_flush_part = self.pending_flush_parts.get(index) buffer = self.buffers.get(index)
if pending_flush_part and isinstance( if buffer and isinstance(delta, (ThinkingPartDelta, TextPartDelta)):
delta, (ThinkingPartDelta, TextPartDelta) # 追加增量
): buffer.content_delta += delta.content_delta or ""
pending_flush_part.content_delta += delta.content_delta or ""
# 原理:若持续输出片段增量事件则拦截并重新创建异步协程任务,至输出间隔超过刷新延迟则返回 # 取消上一轮未完成的延迟刷新任务
# 取消异步协程任务
self._cancel_task(index=index) self._cancel_task(index=index)
# 创建异步协程任务 # 创建延迟刷新任务
pending_flush_part.task = create_task( buffer.task = create_task(coro=self._delay_flush(index=index))
coro=self._delay_flush(index=index)
)
else: else:
yield event yield event
return return
# 若为其它事件则先刷新所有片段缓存再直接返回 # 若为其它事件则先批量刷新片段增量事件再返回该事件
for index in list(self.pending_flush_parts.keys()): async for event_ in self._batch_flush():
flush_event = await self._flush(index=index) yield event_
if flush_event:
yield flush_event
# 若为片段结束事件则删除该片段缓存 # 若为片段结束事件则删除该片段缓存
if isinstance(event, PartEndEvent): if isinstance(event, PartEndEvent):
self.part_caches.pop(event.index, None) self.buffers.pop(event.index, None)
return yield event
def _cancel_task(self, index: int) -> None: def _cancel_task(self, index: int) -> None:
""" """
取消异步协程任务 取消延迟刷新异步协程任务
:param index: 片段索引 :param index: 片段索引
:yield: None :return: None
""" """
pending_flush_part = self.pending_flush_parts.get(index) buffer = self.buffers.get(index)
if not pending_flush_part or not pending_flush_part.task: if not buffer or not buffer.task:
return return
if not buffer.task.done():
buffer.task.cancel()
buffer.task = None
# 若异步协程任务未完成则取消 async def _delay_flush(self, index: int) -> None:
if not pending_flush_part.task.done():
pending_flush_part.task.cancel()
pending_flush_part.task = None
async def _delay_flush(self, index: int):
""" """
延迟刷新 延迟刷新
:param index: 片段索引 :param index: 片段索引
:return: None :return: None
""" """
await sleep(delay=self.flush_delay) await sleep(delay=self.delay)
flush_event = await self._flush(index=index)
if flush_event: # 刷新片段增量事件
await self.pending_flush_queue.put(item=flush_event) event = await self._flush(index=index)
if event:
await self.pending_flush_queue.put(item=event)
async def _flush(self, index: int) -> Optional[PartDeltaEvent]: async def _flush(self, index: int) -> Optional[PartDeltaEvent]:
""" """
刷新 刷新片段增量事件
:param index: 片段索引 :param index: 片段索引
:return: Optional[PartDeltaEvent] :return: Optional[PartDeltaEvent]
""" """
pending_flush_part = self.pending_flush_parts.get(index) buffer = self.buffers.get(index)
if not pending_flush_part or not pending_flush_part.content_delta: if not buffer or not buffer.content_delta:
return None return
match pending_flush_part.type_: # 构建片段增量事件中增量部分
match buffer.part_type:
case "thinking": case "thinking":
delta = ThinkingPartDelta(content_delta=pending_flush_part.content_delta) delta = ThinkingPartDelta(content_delta=buffer.content_delta)
case "text": case "text":
delta = TextPartDelta(content_delta=pending_flush_part.content_delta) delta = TextPartDelta(content_delta=buffer.content_delta)
pending_flush_part.content_delta = "" buffer.content_delta = ""
pending_flush_part.task = None buffer.task = None
return PartDeltaEvent(index=index, delta=delta) return PartDeltaEvent(index=index, delta=delta)
async def _batch_flush(self) -> AsyncGenerator[PartDeltaEvent, None]:
"""
批量刷新片段增量事件
:yield: AsyncGenerator
"""
for index in list(self.buffers.keys()):
event = await self._flush(index=index)
if event:
yield event
class Agent: class Agent:
@ -323,12 +327,13 @@ class Agent:
self.agent.to_web() self.agent.to_web()
async def stream_messages_events( async def stream_messages_events(
self, user_prompt: str | List[str] self, user_prompt: str | List[str], delay: float = 0.2
) -> AsyncGenerator[str, None]: ) -> AsyncGenerator[str, None]:
""" """
流式输出消息事件 流式输出消息事件
:param user_prompt: 用户提示词用户输入消息 :param user_prompt: 用户提示词用户输入消息
:return: 消息事件 :param delay: 延迟时长单位为秒停顿超该时长则刷新
:yield: AsyncGenerator
""" """
"""定义一次会话session包含若干论对话turn每一轮对话由用户输入消息message和智能体输出消息组成""" """定义一次会话session包含若干论对话turn每一轮对话由用户输入消息message和智能体输出消息组成"""
# 获取指定会话的消息历史 # 获取指定会话的消息历史
@ -336,33 +341,58 @@ class Agent:
session_id=self.session_id session_id=self.session_id
) )
# 实例模型消息事件防抖器
event_debouncer = EventsDebouncer(delay=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:
# 处理片段开始事件 if isinstance(event, AgentRunResultEvent):
if isinstance(event, PartStartEvent): evnet : AgentRunResultEvent = event
part = event.part new_messages = evnet.result
if isinstance(part, ThinkingPart): # 处理思考片段
if content := part.content: # 处理模型消息事件
async for event_ in event_debouncer.handle_event(event):
match event_: # event_ 为处理后模型消息事件
case PartStartEvent():
part = event_.part
match part:
case ThinkingPart(content=content):
yield f"00:{content}" yield f"00:{content}"
elif isinstance(part, TextPart): # 处理文本片段 case TextPart(content=content):
if content := part.content:
yield f"01:{content}" yield f"01:{content}"
elif isinstance( case (
part, NativeToolSearchCallPart()
( | NativeToolCallPart()
NativeToolSearchCallPart, | ToolSearchCallPart()
NativeToolCallPart, | ToolCallPart()
ToolSearchCallPart, | LoadCapabilityCallPart()
ToolCallPart, ) as tool:
LoadCapabilityCallPart, yield f"02:使用技能:{tool.tool_name}"
), case (
): # 处理原生搜索、原生工具、自封搜索、自封工具和加载能力调用片段 NativeToolReturnPart() | ToolReturnPart()
yield f"02:技能名称:{part.tool_name}" ) as tool_return:
yield f"03:技能返回:{tool_return.content}"
case _:
yield f"99:未知片段类型{type(part)}"
case PartDeltaEvent():
delta = event_.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 _:
yield f"99:未知片段类型{type(delta)}"
# 片段结束事件无需处理
case PartEndEvent() | AgentRunResultEvent():
pass
case _:
yield f"99:未知事件类型{type(event_)}"
self.agent_memory.create_new_messages( self.agent_memory.create_new_messages(
session_id=self.session_id, session_id=self.session_id,
new_messages=events.new_messages(), new_messages=new_messages,
) )