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