This commit is contained in:
parent
ad71e4b16d
commit
4e65dc783f
202
utils/agent.py
202
utils/agent.py
|
|
@ -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,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue