This commit is contained in:
liubiren 2026-06-16 20:00:46 +08:00
parent 214b7789d6
commit ad71e4b16d
1 changed files with 66 additions and 68 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 from asyncio import Task, create_task, sleep, Queue, QueueEmpty
from pydantic import Field from pydantic import Field
from pydantic import BaseModel from pydantic import BaseModel
@ -133,33 +133,32 @@ class AgentMemory(SQLite):
) from exception ) from exception
class PartCache(BaseModel): class PendingFlushPart(BaseModel):
""" """
片段缓存模型 待刷新模型消息片段
""" """
type_: Literal["thinking", "text"] = Field(..., description="片段类型") type_: Literal["thinking", "text"] = Field(..., description="片段类型")
content: str = 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 中模型流式输出包含若干片段例如思考片段文本片段等每个片段包含若干模型消息事件例如片段开始事件片段增量事件片段结束事件等
""" """
def __init__(self, delay: float = 0.2, content_delta_limit: int = 50): def __init__(self, flush_delay: float = 0.2):
""" """
初始化 初始化
:param delay: 异步协程任务延迟时间 :param flush_delay: 刷新延迟单位为秒
:param content_delta_limit: 增量片段内容的字数上限
""" """
self.delay = delay self.flush_delay = flush_delay
self.content_delta_limit = content_delta_limit
# 所有片段缓存 # 待刷新模型响应片段
self.part_caches: Dict[int, PartCache] = {} self.pending_flush_parts: Dict[int, PendingFlushPart] = {}
# 待刷新列队 # 待刷新列队
self.pending_flush_queue = Queue() self.pending_flush_queue = Queue()
@ -167,17 +166,20 @@ class EventsDebouncer:
self, event: ModelMessageEvent self, event: ModelMessageEvent
) -> AsyncGenerator[ModelMessageEvent, None]: ) -> AsyncGenerator[ModelMessageEvent, None]:
""" """
处理模型消息事件合并片段增量事件 处理模型消息事件
:param event: 模型消息事件 :param event: 模型消息事件
:yield: AsyncGenerator :yield: AsyncGenerator
""" """
# 从待刷新列队中返回刷新事件 # 返回待刷新列队中模型消息事件
async for flush_event in self._yield_pending_flush_events(): while True:
yield flush_event try:
yield self.pending_flush_queue.get_nowait()
except QueueEmpty:
break
# 若为片段开始事件且片段类型为思考或纯文本则创建片段缓存,否则直接返回 # 若为思考或文本的片段开始事件则先缓存片段再返回,其它事件直接返回
if isinstance(event, PartStartEvent): if isinstance(event, PartStartEvent):
index, part = event.index, event.part index, part = event.index, event.part # 片段索引、片段
match part: match part:
case ThinkingPart(): case ThinkingPart():
type_ = "thinking" type_ = "thinking"
@ -186,97 +188,93 @@ class EventsDebouncer:
case _: case _:
yield event yield event
return return
# 创建片段缓存 # 新增片段缓存
self.part_caches[index] = PartCache( self.pending_flush_parts[index] = PendingFlushPart(
type_=type_, type_=type_,
content=part.content or "",
) )
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
part_cache = self.part_caches.get(index) pending_flush_part = self.pending_flush_parts.get(index)
if part_cache and isinstance(delta, (ThinkingPartDelta, TextPartDelta)): if pending_flush_part and isinstance(
part_cache.content += delta.content_delta or "" delta, (ThinkingPartDelta, TextPartDelta)
part_cache.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)
# 创建异步协程任务 # 创建异步协程任务
part_cache.task = create_task(coro=self._delay_flush(index=index)) pending_flush_part.task = create_task(
coro=self._delay_flush(index=index)
)
else: else:
yield event yield event
return return
# 若为其它事件则先刷新所有片段缓存再直接返回 # 若为其它事件则先刷新所有片段缓存再直接返回
async for flush_event in self._batch_flush(): for index in list(self.pending_flush_parts.keys()):
flush_event = await self._flush(index=index)
if flush_event:
yield flush_event yield flush_event
yield event
# 若为片段结束事件则删除该片段缓存 # 若为片段结束事件则删除该片段缓存
if isinstance(event, PartEndEvent): if isinstance(event, PartEndEvent):
self.part_caches.pop(event.index, None) self.part_caches.pop(event.index, None)
return return
async def _yield_pending_flush_events(
self,
) -> AsyncGenerator[ModelMessageEvent, None]:
"""
从待刷新列队中返回刷新事件
:return: AsyncGenerator
"""
while not self.pending_flush_queue.empty():
flush_event = await self.pending_flush_queue.get_nowait()
yield flush_event
def _cancel_task(self, index: int) -> None: def _cancel_task(self, index: int) -> None:
""" """
取消异步协程任务 取消异步协程任务
:param index: 片段索引 :param index: 片段索引
:yield: None :yield: None
""" """
part_cache = self.part_caches.get(index) pending_flush_part = self.pending_flush_parts.get(index)
if part_cache and part_cache.task and not part_cache.task.done(): if not pending_flush_part or not pending_flush_part.task:
part_cache.task.cancel() return
part_cache.task = None
async def _flush(self, index: int) -> Optional[PartDeltaEvent]: # 若异步协程任务未完成则取消
"""刷新片段缓存""" if not pending_flush_part.task.done():
part_cache = self.part_caches.get(index) pending_flush_part.task.cancel()
if not part_cache or not part_cache.content_delta:
return None
match part_cache.type_: pending_flush_part.task = None
case "thinking":
delta = ThinkingPartDelta(content_delta=part_cache.content_delta)
case "text":
delta = TextPartDelta(content_delta=part_cache.content_delta)
part_cache.content_delta = ""
part_cache.task = None
return PartDeltaEvent(index=index, delta=delta)
async def _delay_flush(self, index: int): async def _delay_flush(self, index: int):
""" """
延迟刷新片段缓存 延迟刷新
:param index: 片段索引 :param index: 片段索引
:return: None :return: None
""" """
await sleep(delay=self.delay) await sleep(delay=self.flush_delay)
flush_event = await self._flush(index=index) flush_event = await self._flush(index=index)
if flush_event: if flush_event:
await self.pending_flush_queue.put(item=flush_event) await self.pending_flush_queue.put(item=flush_event)
async def _batch_flush(self): async def _flush(self, index: int) -> Optional[PartDeltaEvent]:
""" """
批量刷新片段缓存 刷新
:param index: 片段索引
:return: Optional[PartDeltaEvent]
""" """
for index in list(self.part_caches.keys()): pending_flush_part = self.pending_flush_parts.get(index)
# 刷新片段缓存 if not pending_flush_part or not pending_flush_part.content_delta:
flush_event = await self._flush(index) return None
if flush_event:
yield flush_event match pending_flush_part.type_:
case "thinking":
delta = ThinkingPartDelta(content_delta=pending_flush_part.content_delta)
case "text":
delta = TextPartDelta(content_delta=pending_flush_part.content_delta)
pending_flush_part.content_delta = ""
pending_flush_part.task = None
return PartDeltaEvent(index=index, delta=delta)
class Agent: class Agent: