324 lines
12 KiB
Python
324 lines
12 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
Pydantic AI 聊天智能体和相关模块
|
||
"""
|
||
|
||
# 列举导入模块
|
||
from asyncio import Queue, QueueEmpty, Task, Task, create_task, sleep
|
||
from typing import AsyncGenerator, Dict, List, Literal, Optional, Union
|
||
|
||
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.messages import (
|
||
AgentStreamEvent,
|
||
FinalResultEvent,
|
||
LoadCapabilityCallPart,
|
||
ModelMessage,
|
||
NativeToolCallPart,
|
||
NativeToolReturnPart,
|
||
NativeToolSearchCallPart,
|
||
PartDeltaEvent,
|
||
PartEndEvent,
|
||
PartStartEvent,
|
||
TextPart,
|
||
TextPartDelta,
|
||
ThinkingPart,
|
||
ThinkingPartDelta,
|
||
ToolCallPart,
|
||
ToolCallPartDelta,
|
||
ToolReturnPart,
|
||
ToolSearchCallPart,
|
||
)
|
||
from pydantic_ai.models.openai import OpenAIChatModel
|
||
from pydantic_ai.output import OutputSpec
|
||
from pydantic_ai.providers.openai import OpenAIProvider
|
||
from pydantic_ai.run import AgentRunResultEvent
|
||
|
||
|
||
class PartBuffer(BaseModel):
|
||
"""
|
||
模型响应片段缓冲类
|
||
"""
|
||
|
||
part_type: Literal["thinking", "text"] = Field(
|
||
..., description="片段类型,仅支持思考和文本片段"
|
||
)
|
||
part_content_delta: str = Field(default="", description="片段内容增量")
|
||
flush_task: Optional[Task] = Field(
|
||
default=None, description="片段绑定的延迟刷新任务"
|
||
)
|
||
model_config = {"arbitrary_types_allowed": True} # 允许任意类型
|
||
|
||
|
||
class Debouncer:
|
||
"""
|
||
模型响应事件防抖器
|
||
用于平滑 run_stream_events 返回的高频模型响应事件,降低前端渲染频率
|
||
"""
|
||
|
||
def __init__(self, flush_delay: float):
|
||
"""
|
||
初始化
|
||
:param flush_delay: 刷新延迟时长(单位为秒)
|
||
"""
|
||
self.flush_delay = flush_delay
|
||
|
||
# 片段缓冲字典(键为片段索引,值为片段缓冲实例)
|
||
self.part_buffers: Dict[int, PartBuffer] = {}
|
||
# 待推送刷新队列
|
||
self.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.flush_queue.get_nowait()
|
||
except QueueEmpty:
|
||
break
|
||
|
||
# 若为思考或文本的片段开始事件则初始化片段缓冲,其它事件透传
|
||
if isinstance(event, PartStartEvent):
|
||
part_index, part = event.index, event.part # 片段索引、片段
|
||
match part:
|
||
case ThinkingPart():
|
||
part_type = "thinking"
|
||
case TextPart():
|
||
part_type = "text"
|
||
case _:
|
||
yield event
|
||
return
|
||
# 初始化片段缓冲
|
||
self.part_buffers[part_index] = PartBuffer(
|
||
part_type=part_type,
|
||
)
|
||
yield event
|
||
return
|
||
|
||
# 若为思考或文本片段增量事件则先取消延迟刷新任务,更新片段缓冲中的片段内容增量并重新绑定延迟刷新任务,其它事件透传
|
||
if isinstance(event, PartDeltaEvent):
|
||
part_index, delta = event.index, event.delta
|
||
part_buffer = self.part_buffers.get(part_index)
|
||
if part_buffer and isinstance(delta, (ThinkingPartDelta, TextPartDelta)):
|
||
# 取消延迟刷新任务
|
||
self._cancel_task(part_index=part_index)
|
||
# 更新片段缓冲中的片段内容增量
|
||
part_buffer.part_content_delta += delta.content_delta or ""
|
||
# 重新绑定延迟刷新任务
|
||
part_buffer.flush_task = create_task(
|
||
coro=self._flush_after_delay(part_index=part_index)
|
||
)
|
||
else:
|
||
yield event
|
||
return
|
||
|
||
# 若为其它事件则先刷新所有片段缓冲再透传当前事件
|
||
async for flush_event in self._flush_all():
|
||
yield flush_event
|
||
# 若为片段结束事件则删除相应片段缓冲
|
||
if isinstance(event, PartEndEvent):
|
||
self.part_buffers.pop(event.index, None)
|
||
yield event
|
||
|
||
def _cancel_task(self, part_index: int) -> None:
|
||
"""
|
||
取消延迟刷新任务
|
||
:param part_index: 片段索引
|
||
:return: None
|
||
"""
|
||
part_buffer = self.part_buffers.get(part_index)
|
||
if not part_buffer or not part_buffer.flush_task:
|
||
return
|
||
|
||
# 若延迟刷新任务未推送则取消
|
||
if not part_buffer.flush_task.done():
|
||
part_buffer.flush_task.cancel()
|
||
part_buffer.flush_task = None
|
||
|
||
async def _flush_after_delay(self, part_index: int) -> None:
|
||
"""
|
||
延迟刷新
|
||
:param part_index: 片段索引
|
||
:return: None
|
||
"""
|
||
await sleep(delay=self.flush_delay)
|
||
|
||
# 刷新单个片段缓冲
|
||
flush_event = await self._flush(part_index=part_index)
|
||
if flush_event:
|
||
await self.flush_queue.put(item=flush_event)
|
||
|
||
async def _flush(self, part_index: int) -> Optional[PartDeltaEvent]:
|
||
"""
|
||
刷新单个片段缓冲
|
||
:param part_index: 片段索引
|
||
:return: Optional[PartDeltaEvent]
|
||
"""
|
||
part_buffer = self.part_buffers.get(part_index)
|
||
if not part_buffer or not part_buffer.part_content_delta:
|
||
return
|
||
|
||
# 构建片段增量事件中增量
|
||
match part_buffer.part_type:
|
||
case "thinking":
|
||
delta = ThinkingPartDelta(content_delta=part_buffer.part_content_delta)
|
||
case "text":
|
||
delta = TextPartDelta(content_delta=part_buffer.part_content_delta)
|
||
|
||
# 重置片段缓冲中片段内容增量
|
||
part_buffer.part_content_delta = ""
|
||
# 重置片段缓冲中延迟刷新任务
|
||
part_buffer.flush_task = None
|
||
return PartDeltaEvent(index=part_index, delta=delta)
|
||
|
||
async def _flush_all(self) -> AsyncGenerator[PartDeltaEvent]:
|
||
"""
|
||
刷新所有片段缓冲
|
||
:yield: AsyncGenerator
|
||
"""
|
||
for part_index in list(self.part_buffers.keys()):
|
||
flush_event = await self._flush(part_index=part_index)
|
||
if flush_event:
|
||
yield flush_event
|
||
|
||
|
||
DEFAULT_INSTRUCTIONS: str = """
|
||
# 角色
|
||
专业友好AI助手,结构化解答各类问题。
|
||
|
||
# 输出硬性规则
|
||
1. 全文强制标准Markdown,禁止纯文本;不要额外说明排版格式,直接输出内容;
|
||
2. 层级使用 `#/##/###`,列表用 `-` 无序列表或数字有序列表;
|
||
3. 代码块用 ```语言名``` 包裹;
|
||
4. 重点内容标注 **粗体**/*斜体*;
|
||
5. 思考、工具日志仅输出文本,适配前端折叠面板,禁止输出HTML标签;
|
||
6. 内容分点拆分,排版整洁适配前端Markdown渲染。
|
||
|
||
# 行文要求
|
||
语言通俗,逻辑完整简洁,无多余废话。
|
||
"""
|
||
|
||
|
||
class Agent:
|
||
"""
|
||
基于 Pydantic AI 封装的智能体,支持:
|
||
1、流式输出模型响应事件
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
chat_id: str,
|
||
instructions: Optional[str] = None,
|
||
output_type: OutputSpec = str,
|
||
capabilities: Optional[List[AgentCapability]] = None,
|
||
retries: int = 1,
|
||
):
|
||
"""
|
||
初始化
|
||
:param chat_id: 聊天唯一标识
|
||
:param instructions: 指令
|
||
:param capabilities: 技能列表,默认为不使用技能
|
||
:param output_type: 输出类型
|
||
:param retries: 重试次数,默认为1次
|
||
:return: 智能体实例
|
||
"""
|
||
# 聊天唯一标识
|
||
self.chat_id = chat_id
|
||
|
||
# 一次聊天(chat)包含若干论对话(dialog),每一轮对话由用户提示词(user_prompt)和输出(output)组成,两者统称为消息(message)
|
||
|
||
# 本轮对话新增消息列表
|
||
self.new_messages: List[ModelMessage] = []
|
||
|
||
# 若指令为空则使用默认指令
|
||
if not instructions:
|
||
instructions = DEFAULT_INSTRUCTIONS
|
||
|
||
# 初始化智能体
|
||
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,
|
||
)
|
||
|
||
async def stream_events(
|
||
self,
|
||
user_prompt: str | List[str],
|
||
message_history: Optional[List[ModelMessage]] = None,
|
||
flush_delay: float = 0.15,
|
||
) -> AsyncGenerator[str, None]:
|
||
"""
|
||
流式输出模型响应事件
|
||
:param user_prompt: 用户提示词(用户提示词)
|
||
:param flush_delay: 刷新延迟时长(单位为秒)
|
||
:yield: AsyncGenerator
|
||
"""
|
||
# 初始化防抖器
|
||
debouncer = Debouncer(flush_delay=flush_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 PartEndEvent():
|
||
continue
|
||
case FinalResultEvent():
|
||
continue
|
||
case AgentRunResultEvent(result=result):
|
||
self.new_messages = result.new_messages()
|
||
yield "05:"
|
||
case _:
|
||
yield f"06:未知事件类型{type(event).__name__}"
|