Python/产品需求文档AI生成/application/utils/agent.py

324 lines
12 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 -*-
"""
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__}"