398 lines
14 KiB
Python
398 lines
14 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
智能体模块
|
||
"""
|
||
|
||
# 列举导入模块
|
||
from asyncio import Queue, QueueEmpty, Task, Task, create_task, sleep
|
||
from pathlib import Path
|
||
import time
|
||
from typing import AsyncGenerator, Dict, List, Literal, Optional, Union
|
||
from uuid import uuid4
|
||
|
||
from pydantic import Field
|
||
from pydantic import BaseModel
|
||
from pydantic_ai import Agent as PydanticAIAgent, ModelMessage
|
||
from pydantic_ai.capabilities import AgentCapability
|
||
from pydantic_ai.run import AgentRunResultEvent
|
||
from pydantic_ai.messages import (
|
||
ModelMessage,
|
||
FinalResultEvent,
|
||
LoadCapabilityCallPart,
|
||
AgentStreamEvent,
|
||
ModelMessagesTypeAdapter,
|
||
NativeToolCallPart,
|
||
NativeToolReturnPart,
|
||
NativeToolSearchCallPart,
|
||
PartDeltaEvent,
|
||
PartEndEvent,
|
||
PartStartEvent,
|
||
TextPart,
|
||
TextPartDelta,
|
||
ThinkingPart,
|
||
ThinkingPartDelta,
|
||
ToolCallPartDelta,
|
||
ToolCallPart,
|
||
ToolReturnPart,
|
||
ToolSearchCallPart,
|
||
)
|
||
from pydantic_ai.models.openai import OpenAIChatModel
|
||
from pydantic_ai.output import OutputSpec
|
||
from pydantic_ai.providers.openai import OpenAIProvider
|
||
|
||
from sys import path
|
||
from pathlib import Path
|
||
|
||
path.append(Path(__file__).parent.as_posix())
|
||
from sqlite import SQLite
|
||
|
||
|
||
class Memory(SQLite):
|
||
"""
|
||
记忆体
|
||
"""
|
||
|
||
def __init__(self):
|
||
"""
|
||
初始化
|
||
"""
|
||
# 构建数据库路径
|
||
super().__init__(database=Path(__file__).parent.resolve() / "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 save(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="延迟刷新异步协程任务")
|
||
model_config = {"arbitrary_types_allowed": True}
|
||
|
||
|
||
class Debouncer:
|
||
"""
|
||
防抖器
|
||
用于处理 run_stream_events 返回的模型消息事件
|
||
"""
|
||
|
||
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: Union[AgentStreamEvent, AgentRunResultEvent]
|
||
) -> AsyncGenerator[Union[AgentStreamEvent, AgentRunResultEvent]]:
|
||
"""
|
||
处理模型消息事件
|
||
在 Pydantic-AI 中模型流式输出包含若干类型片段,例如思考片段、文本片段等。每类型片段包含若干模型消息事件,例如片段开始事件、片段增量事件、片段结束事件等
|
||
: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,
|
||
retries: int = 1,
|
||
):
|
||
"""
|
||
初始化智能体
|
||
:param session_id: 会话唯一标识
|
||
:param instructions: 指令
|
||
:param capabilities: 技能列表,默认为不使用技能
|
||
:param output_type: 输出类型
|
||
:param retries: 重试次数,默认为1次
|
||
:return: 智能体实例
|
||
"""
|
||
# 会话唯一标识
|
||
self.session_id = session_id
|
||
|
||
# 实例智能体的记忆体
|
||
self.memory = Memory()
|
||
|
||
# 实例智能体
|
||
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=retries,
|
||
)
|
||
|
||
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.memory.get_message_history(session_id=self.session_id)
|
||
|
||
# 实例防抖器
|
||
debouncer = Debouncer(delay=delay)
|
||
|
||
async with self.agent.run_stream_events(
|
||
user_prompt=user_prompt,
|
||
message_history=message_history,
|
||
) as events:
|
||
async for event in events:
|
||
# 处理模型消息事件
|
||
async for event in debouncer.handle_event(event):
|
||
match event:
|
||
# 片段开始事件
|
||
case PartStartEvent(part=part):
|
||
match part:
|
||
case ThinkingPart(content=content):
|
||
yield f"00:{content}"
|
||
case TextPart(content=content):
|
||
yield f"01:{content}"
|
||
case (
|
||
NativeToolSearchCallPart(tool_name=tool_name)
|
||
| NativeToolCallPart(tool_name=tool_name)
|
||
| ToolSearchCallPart(tool_name=tool_name)
|
||
| ToolCallPart(tool_name=tool_name)
|
||
| LoadCapabilityCallPart(tool_name=tool_name)
|
||
):
|
||
yield f"02:正在使用技能:{tool_name}"
|
||
case NativeToolReturnPart(
|
||
content=content
|
||
) | ToolReturnPart(content=content):
|
||
yield f"04:技能返回:{content}"
|
||
case _:
|
||
yield f"06:未知片段类型{type(part).__name__}"
|
||
# 片段增量事件
|
||
case PartDeltaEvent(delta=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 ToolCallPartDelta(args_delta=args_delta):
|
||
yield f"03:技能参数:{args_delta}"
|
||
case _:
|
||
yield f"06:未知片段类型{type(delta).__name__}"
|
||
# 片段结束事件、最终结果事件无需处理
|
||
case PartEndEvent():
|
||
continue
|
||
case FinalResultEvent():
|
||
continue
|
||
case AgentRunResultEvent(result=result):
|
||
self.memory.save(
|
||
session_id=self.session_id,
|
||
new_messages=result.new_messages(),
|
||
)
|
||
yield "05:FinalResultEvent"
|
||
case _:
|
||
yield f"06:未知事件类型{type(event).__name__}"
|