This commit is contained in:
parent
491ea47a9d
commit
214b7789d6
200
utils/agent.py
200
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
|
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:
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue