This commit is contained in:
parent
214b7789d6
commit
ad71e4b16d
134
utils/agent.py
134
utils/agent.py
|
|
@ -11,7 +11,7 @@ from uuid import uuid4
|
|||
from pydantic_ai.messages import ModelMessage
|
||||
from pydantic_ai import Agent as PydanticAIAgent, ModelMessage
|
||||
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 BaseModel
|
||||
|
|
@ -133,33 +133,32 @@ class AgentMemory(SQLite):
|
|||
) from exception
|
||||
|
||||
|
||||
class PartCache(BaseModel):
|
||||
class PendingFlushPart(BaseModel):
|
||||
"""
|
||||
片段缓存模型
|
||||
待刷新模型消息片段
|
||||
"""
|
||||
|
||||
type_: Literal["thinking", "text"] = Field(..., description="片段类型")
|
||||
content: str = Field(..., description="片段内容")
|
||||
content_delta: str = Field(default="", description="片段内容增量")
|
||||
task: Optional[Task] = Field(default=None, description="片段绑定的异步协程任务")
|
||||
task: Optional[Task] = Field(default=None, description="异步协程任务")
|
||||
|
||||
|
||||
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 content_delta_limit: 增量片段内容的字数上限
|
||||
:param flush_delay: 刷新延迟(单位为秒)
|
||||
"""
|
||||
self.delay = delay
|
||||
self.content_delta_limit = content_delta_limit
|
||||
self.flush_delay = flush_delay
|
||||
|
||||
# 所有片段缓存
|
||||
self.part_caches: Dict[int, PartCache] = {}
|
||||
# 待刷新模型响应片段
|
||||
self.pending_flush_parts: Dict[int, PendingFlushPart] = {}
|
||||
# 待刷新列队
|
||||
self.pending_flush_queue = Queue()
|
||||
|
||||
|
|
@ -167,17 +166,20 @@ class EventsDebouncer:
|
|||
self, event: ModelMessageEvent
|
||||
) -> AsyncGenerator[ModelMessageEvent, None]:
|
||||
"""
|
||||
处理模型消息事件:合并片段增量事件
|
||||
处理模型消息事件
|
||||
:param event: 模型消息事件
|
||||
:yield: AsyncGenerator
|
||||
"""
|
||||
# 从待刷新列队中返回刷新事件
|
||||
async for flush_event in self._yield_pending_flush_events():
|
||||
yield flush_event
|
||||
# 返回待刷新列队中模型消息事件
|
||||
while True:
|
||||
try:
|
||||
yield self.pending_flush_queue.get_nowait()
|
||||
except QueueEmpty:
|
||||
break
|
||||
|
||||
# 若为片段开始事件且片段类型为思考或纯文本则创建片段缓存,否则直接返回
|
||||
# 若为思考或文本的片段开始事件则先缓存片段再返回,其它事件直接返回
|
||||
if isinstance(event, PartStartEvent):
|
||||
index, part = event.index, event.part
|
||||
index, part = event.index, event.part # 片段索引、片段
|
||||
match part:
|
||||
case ThinkingPart():
|
||||
type_ = "thinking"
|
||||
|
|
@ -186,97 +188,93 @@ class EventsDebouncer:
|
|||
case _:
|
||||
yield event
|
||||
return
|
||||
# 创建片段缓存
|
||||
self.part_caches[index] = PartCache(
|
||||
# 新增片段缓存
|
||||
self.pending_flush_parts[index] = PendingFlushPart(
|
||||
type_=type_,
|
||||
content=part.content or "",
|
||||
)
|
||||
yield event
|
||||
return
|
||||
|
||||
# 若为片段增量事件且增量类型为思考或纯文本则更新片段内容和增量片段内容并刷新片段缓存,否则直接返回
|
||||
# 若为思考或文本的片段增量事件则更新缓存,其它事件直接返回
|
||||
if isinstance(event, PartDeltaEvent):
|
||||
index, delta = event.index, event.delta
|
||||
part_cache = self.part_caches.get(index)
|
||||
if part_cache and isinstance(delta, (ThinkingPartDelta, TextPartDelta)):
|
||||
part_cache.content += delta.content_delta or ""
|
||||
part_cache.content_delta += delta.content_delta or ""
|
||||
pending_flush_part = self.pending_flush_parts.get(index)
|
||||
if pending_flush_part and isinstance(
|
||||
delta, (ThinkingPartDelta, TextPartDelta)
|
||||
):
|
||||
pending_flush_part.content_delta += delta.content_delta or ""
|
||||
|
||||
# 原理:若持续输出片段增量事件则拦截并重新创建异步协程任务,至输出间隔超过刷新延迟则返回
|
||||
# 取消异步协程任务
|
||||
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:
|
||||
yield event
|
||||
return
|
||||
|
||||
# 若为其它事件则先刷新所有片段缓存再直接返回
|
||||
async for flush_event in self._batch_flush():
|
||||
yield flush_event
|
||||
yield event
|
||||
for index in list(self.pending_flush_parts.keys()):
|
||||
flush_event = await self._flush(index=index)
|
||||
if flush_event:
|
||||
yield flush_event
|
||||
|
||||
# 若为片段结束事件则删除该片段缓存
|
||||
if isinstance(event, PartEndEvent):
|
||||
self.part_caches.pop(event.index, None)
|
||||
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:
|
||||
"""
|
||||
取消异步协程任务
|
||||
:param index: 片段索引
|
||||
:yield: None
|
||||
"""
|
||||
part_cache = self.part_caches.get(index)
|
||||
if part_cache and part_cache.task and not part_cache.task.done():
|
||||
part_cache.task.cancel()
|
||||
part_cache.task = None
|
||||
pending_flush_part = self.pending_flush_parts.get(index)
|
||||
if not pending_flush_part or not pending_flush_part.task:
|
||||
return
|
||||
|
||||
async def _flush(self, index: int) -> Optional[PartDeltaEvent]:
|
||||
"""刷新片段缓存"""
|
||||
part_cache = self.part_caches.get(index)
|
||||
if not part_cache or not part_cache.content_delta:
|
||||
return None
|
||||
# 若异步协程任务未完成则取消
|
||||
if not pending_flush_part.task.done():
|
||||
pending_flush_part.task.cancel()
|
||||
|
||||
match part_cache.type_:
|
||||
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)
|
||||
pending_flush_part.task = None
|
||||
|
||||
async def _delay_flush(self, index: int):
|
||||
"""
|
||||
延迟刷新片段缓存
|
||||
延迟刷新
|
||||
:param index: 片段索引
|
||||
:return: None
|
||||
"""
|
||||
await sleep(delay=self.delay)
|
||||
await sleep(delay=self.flush_delay)
|
||||
flush_event = await self._flush(index=index)
|
||||
if 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()):
|
||||
# 刷新片段缓存
|
||||
flush_event = await self._flush(index)
|
||||
if flush_event:
|
||||
yield flush_event
|
||||
pending_flush_part = self.pending_flush_parts.get(index)
|
||||
if not pending_flush_part or not pending_flush_part.content_delta:
|
||||
return None
|
||||
|
||||
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:
|
||||
|
|
|
|||
Loading…
Reference in New Issue