Python/utils/agent.py

399 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# -*- coding: utf-8 -*-
"""
智能体模块
"""
# 列举导入模块
from pathlib import Path
import time
from typing import AsyncGenerator, List, Optional, Dict, Literal
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, Task
from pydantic import Field
from pydantic import BaseModel
from pydantic_ai.messages import (
ModelMessagesTypeAdapter,
ModelResponse,
FinalResultEvent,
ModelMessageEvent,
PartStartEvent,
PartDeltaEvent,
PartEndEvent,
TextPart,
ToolSearchCallPart,
NativeToolSearchCallPart,
LoadCapabilityCallPart,
ThinkingPart,
ToolCallPart,
NativeToolCallPart,
TextPartDelta,
ThinkingPartDelta,
ModelResponseStreamEvent,
ToolReturnPart,
AgentRunResultEvent,
NativeToolReturnPart,
)
from pydantic_ai.models.openai import OpenAIChatModel
from pydantic_ai.output import OutputSpec
from pydantic_ai.providers.openai import OpenAIProvider
from .sqlite import SQLite
class AgentMemory(SQLite):
"""
智能体记忆体,支持:
create新增对话消息
read查询会话历史消息
"""
def __init__(self):
"""
初始化
"""
# 构建数据库路径
super().__init__(database=Path(__file__).parent.resolve() / "agent_memory.db")
try:
with self:
self.execute(
sql="""
CREATE TABLE IF NOT EXISTS new_messages
(
--唯一标识
id TEXT PRIMARY KEY,
--会话唯一标识
session_id TEXT NOT NULL,
--新对话消息
new_messages TEXT NOT NULL,
--时间戳(毫秒)
timestamp INTEGER NOT NULL
)
"""
)
except Exception as exception:
raise RuntimeError(f"初始化数据库发生异常:{str(exception)}") from exception
def create_new_messages(
self, session_id: str, new_messages: List[ModelMessage]
) -> bool:
"""
新增新对话消息
:param session_id: 会话唯一标识
:param new_messages: 新对话消息
:return: 新增是否成功
"""
try:
with self:
return self.execute(
sql="""
INSERT INTO new_messages (id, session_id, new_messages, timestamp) VALUES (?, ?, ?, ?)
""",
parameters=(
uuid4().hex.lower(),
session_id,
ModelMessagesTypeAdapter.dump_json(new_messages),
int(time.time() * 1000),
),
)
except Exception as exception:
raise RuntimeError(f"新增对话消息发生异常:{str(exception)}") from exception
def get_message_history(self, session_id: str) -> List[ModelMessage]:
"""
获取指定会话的消息历史
:param session_id: 会话唯一标识
:return: 消息历史
"""
try:
with self:
result = self.query_all(
sql="""
SELECT new_messages
FROM new_messages
WHERE session_id = ?
ORDER BY timestamp ASC
""",
parameters=(session_id,),
)
message_history = []
for row in result:
message_history.extend(
ModelMessagesTypeAdapter.validate_json(row["new_messages"])
)
return message_history
except Exception as exception:
raise RuntimeError(
f"查询会话历史消息发生异常:{str(exception)}"
) from exception
class Buffer(BaseModel):
"""
模型消息事件缓冲模型
"""
part_type: Literal["thinking", "text"] = Field(..., description="片段类型")
content_delta: str = Field(default="", description="片段内容增量")
task: Optional[Task] = Field(default=None, description="延迟刷新异步协程任务")
class EventsDebouncer:
"""
模型消息事件防抖器
使用 Pydantic-AI 中 run_stream_events 方法时需就其返回的模型消息事件处理,延迟刷新片段类型为思考或文本的片段增量事件以实现防抖
在 Pydantic-AI 中模型流式输出包含若干类型片段,例如思考片段、文本片段等。每类型片段包含若干模型消息事件,例如片段开始事件、片段增量事件、片段结束事件等
"""
def __init__(self, delay: float = 0.2):
"""
初始化
:param delay: 延迟时长(单位为秒),停顿超该时长则刷新
"""
self.delay = delay
# 模型消息事件缓冲字典(数据类型为字典,键为片段索引,值为缓冲模型)
self.buffers: Dict[int, Buffer] = {}
# 待刷新异步队列
self.pending_flush_queue = Queue()
async def handle_event(
self, event: ModelMessageEvent
) -> AsyncGenerator[ModelMessageEvent]:
"""
处理模型消息事件
:param event: 模型消息事件
:yield: AsyncGenerator
"""
# 返回待刷新异步队列中模型消息事件
while True:
try:
yield self.pending_flush_queue.get_nowait()
except QueueEmpty:
break
# 若为思考或文本的片段开始事件则先缓存片段再返回,其它事件直接返回
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
# 新增片段缓存
self.buffers[index] = Buffer(
part_type=part_type,
)
yield event
return
"""
原理:
收到思考或文本片段增量事件时取消未完成的延迟刷新任务、追加增量并重设;若模型持续输出则先缓存,停顿超过阈值再返回,实现防抖
"""
if isinstance(event, PartDeltaEvent):
index, delta = event.index, event.delta
buffer = self.buffers.get(index)
if buffer and isinstance(delta, (ThinkingPartDelta, TextPartDelta)):
# 追加增量
buffer.content_delta += delta.content_delta or ""
# 取消上一轮未完成的延迟刷新任务
self._cancel_task(index=index)
# 创建延迟刷新任务
buffer.task = create_task(coro=self._delay_flush(index=index))
else:
yield event
return
# 若为其它事件则先批量刷新片段增量事件再返回该事件
async for event_ in self._batch_flush():
yield event_
# 若为片段结束事件则删除该片段缓存
if isinstance(event, PartEndEvent):
self.buffers.pop(event.index, None)
yield event
def _cancel_task(self, index: int) -> None:
"""
取消延迟刷新异步协程任务
:param index: 片段索引
:return: None
"""
buffer = self.buffers.get(index)
if not buffer or not buffer.task:
return
if not buffer.task.done():
buffer.task.cancel()
buffer.task = None
async def _delay_flush(self, index: int) -> None:
"""
延迟刷新
:param index: 片段索引
:return: None
"""
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]
"""
buffer = self.buffers.get(index)
if not buffer or not buffer.content_delta:
return
# 构建片段增量事件中增量部分
match buffer.part_type:
case "thinking":
delta = ThinkingPartDelta(content_delta=buffer.content_delta)
case "text":
delta = TextPartDelta(content_delta=buffer.content_delta)
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:
"""
智能体,支持:
1 实例智能体
2 异步运行
"""
def __init__(
self,
session_id: str,
instructions: str,
output_type: OutputSpec = str,
capabilities: Optional[List[AgentCapability]] = None,
):
"""
初始化智能体
:param session_id: 会话唯一标识
:param instructions: 指令
:param skills: 智能体技能列表,默认为不使用技能
:param output_type: 输出类型
:return: 智能体实例
"""
# 会话唯一标识
self.session_id = session_id
# 实例智能体的记忆体
self.agent_memory = AgentMemory()
# 实例智能体
self.agent = PydanticAIAgent(
model=OpenAIChatModel(
model_name="deepseek-v4-flash",
provider=OpenAIProvider(
base_url="https://tokenhub.tencentmaas.com/v1",
api_key="sk-D9Y1mCe8VlvNqLuSC4mAjqEwxJ2nW4C0h8a7EPn8kg9RLsHq",
),
),
instructions=instructions,
capabilities=capabilities,
output_type=output_type,
retries=1,
)
self.agent.to_web()
async def stream_messages_events(
self, user_prompt: str | List[str], delay: float = 0.2
) -> AsyncGenerator[str, None]:
"""
流式输出消息事件
:param user_prompt: 用户提示词(用户输入消息)
:param delay: 延迟时长(单位为秒),停顿超该时长则刷新
:yield: AsyncGenerator
"""
"""定义一次会话session包含若干论对话turn每一轮对话由用户输入消息message和智能体输出消息组成"""
# 获取指定会话的消息历史
message_history = self.agent_memory.get_message_history(
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, 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=new_messages,
)