This commit is contained in:
liubiren 2026-06-15 19:59:49 +08:00
parent 491ea47a9d
commit 214b7789d6
10 changed files with 131 additions and 69 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 from asyncio import Task, create_task, sleep, Queue
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,
ModelMessageEvent,
PartStartEvent, PartStartEvent,
PartDeltaEvent, PartDeltaEvent,
PartEndEvent, PartEndEvent,
@ -34,6 +35,8 @@ from pydantic_ai.messages import (
ThinkingPartDelta, ThinkingPartDelta,
ModelResponseStreamEvent, ModelResponseStreamEvent,
ToolReturnPart, ToolReturnPart,
AgentRunResultEvent,
NativeToolReturnPart,
) )
from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.models.openai import OpenAIChatModel
from pydantic_ai.output import OutputSpec from pydantic_ai.output import OutputSpec
@ -44,16 +47,16 @@ from .sqlite import SQLite
class AgentMemory(SQLite): class AgentMemory(SQLite):
""" """
智能体记忆体支持 智能体记忆体支持
create新增对话消息 create新增对话消息
read查询会话历史消息 read查询会话历史消息
""" """
def __init__(self): def __init__(self):
""" """
初始化智能体的记忆体 初始化
""" """
# 构建智能体的记忆体的数据库路径 # 构建数据库路径
super().__init__(database=Path(__file__).parent.resolve() / "agent_memory.db") super().__init__(database=Path(__file__).parent.resolve() / "agent_memory.db")
try: try:
@ -74,9 +77,7 @@ class AgentMemory(SQLite):
""" """
) )
except Exception as exception: except Exception as exception:
raise RuntimeError( raise RuntimeError(f"初始化数据库发生异常:{str(exception)}") from exception
f"初始化智能体的记忆体发生异常:{str(exception)}"
) from exception
def create_new_messages( def create_new_messages(
self, session_id: str, new_messages: List[ModelMessage] self, session_id: str, new_messages: List[ModelMessage]
@ -132,89 +133,150 @@ class AgentMemory(SQLite):
) from exception ) from exception
class Part(BaseModel): class PartCache(BaseModel):
""" """
片段模型 片段缓存模型
""" """
type: Literal["thinking", "text"] = Field(..., description="片段类型") type_: Literal["thinking", "text"] = Field(..., description="片段类型")
content: str = Field(..., description="片段内容") content: str = Field(..., description="片段内容")
task: Optional[Task] = Field(default=None, description="片段绑定的异步任务") content_delta: str = Field(default="", description="片段内容增量")
task: Optional[Task] = Field(default=None, description="片段绑定的异步协程任务")
class StreamEventsDebouncer: class EventsDebouncer:
""" """
流式事件防抖器 模型消息事件防抖器
""" """
def __init__(self, debounce_by: float = 0.2): def __init__(self, delay: float = 0.2, content_delta_limit: int = 50):
""" """
初始化 初始化
:param debounce_by: 防抖间隔 :param delay: 异步协程任务延迟时间
:param content_delta_limit: 增量片段内容的字数上限
""" """
# 片段缓存 self.delay = delay
self.parts: Dict[int, Part] = {} self.content_delta_limit = content_delta_limit
async def handle_event(self, event): # 所有片段缓存
""" self.part_caches: Dict[int, PartCache] = {}
处理流式事件 # 待刷新列队
:param event: 流式事件 self.pending_flush_queue = Queue()
:yield: 处理后流式事件
"""
# 若为片段开始事件且片段类型为思考或纯文本则创建片段缓存,否则直接返回
if isinstance(event, PartStartEvent):
index, part = event.index, event.part
match part:
case ThinkingPart():
part_type = "thinking"
case TextPart():
part_type = "text"
case _:
yield event
return
# 缓存片段 async def handle_event(
self.parts[index] = Part( self, event: ModelMessageEvent
type=part_type, ) -> AsyncGenerator[ModelMessageEvent, None]:
content=part.content if part.content else "", """
) 处理模型消息事件合并片段增量事件
# 片段首次触发直接推送原始事件,保证首条内容立刻展示 :param event: 模型消息事件
yield event :yield: AsyncGenerator
return """
# 从待刷新列队中返回刷新事件
async for flush_event in self._yield_pending_flush_events():
yield flush_event
# 若为片段增量事件且增量类型为思考或纯文本则在片段缓存的片段内容追加 # 若为片段开始事件且片段类型为思考或纯文本则创建片段缓存,否则直接返回
if isinstance(event, PartDeltaEvent): if isinstance(event, PartStartEvent):
index, delta = event.index, event.delta index, part = event.index, event.part
part = self.parts.get(index) match part:
if isinstance(delta, (ThinkingPartDelta, TextPartDelta)): case ThinkingPart():
# 增量内容追加至缓冲区 type_ = "thinking"
part.content += delta.content_delta case TextPart():
# 重置防抖计时器,实现滑动窗口效果 type_ = "text"
self._cancel_timer(index) case _:
part.timer = create_task(self._delay_flush(index))
else:
# 非文本类增量事件直接透传
yield event yield event
return return
# 创建片段缓存
# 命中强制刷新规则:先推送所有缓存聚合内容,再透传当前触发事件 self.part_caches[index] = PartCache(
if isinstance(event, FLUSH_TRIGGER_TYPES): type_=type_,
async for flush_event in self.flush_all(): content=part.content or "",
yield flush_event )
yield event
return
# 其余未匹配的所有事件原样透传
yield event yield event
return
def _cancel_task(self, index: int): # 若为片段增量事件且增量类型为思考或纯文本则更新片段内容和增量片段内容并刷新片段缓存,否则直接返回
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 ""
# 取消异步协程任务
self._cancel_task(index=index)
# 创建异步协程任务
part_cache.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
# 若为片段结束事件则删除该片段缓存
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: 片段索引 :param index: 片段索引
:yield: None
""" """
part = self.parts.get(index) part_cache = self.part_caches.get(index)
if part and part.task is not None and not part.task.done(): if part_cache and part_cache.task and not part_cache.task.done():
part.task.cancel() part_cache.task.cancel()
part_cache.task = None
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
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)
async def _delay_flush(self, index: int):
"""
延迟刷新片段缓存
:param index: 片段索引
:return: None
"""
await sleep(delay=self.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):
"""
批量刷新片段缓存
"""
for index in list(self.part_caches.keys()):
# 刷新片段缓存
flush_event = await self._flush(index)
if flush_event:
yield flush_event
class Agent: class Agent: