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

297 lines
11 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="片段类型")
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.25):
"""
初始化
:param delay: 延迟时长(单位为秒),停顿超该时长则刷新
"""
self.delay = delay
# 片段缓冲字典(键为片段索引,值为片段缓冲实例)
self.part_buffers: Dict[int, PartBuffer] = {}
# 待刷新异步队列
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:
"""
Pydantic AI 智能体
"""
def __init__(
self,
chat_id: str,
instructions: str,
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每一轮对话由输入消息message和输出消息message组成
self.new_messages: List[ModelMessage] = []
# 实例智能体
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_messages_events(
self,
user_prompt: str | List[str],
message_history: Optional[List[ModelMessage]] = None,
delay: float = 0.2,
) -> AsyncGenerator[str, None]:
"""
流式输出消息事件
:param user_prompt: 用户提示词(用户输入消息)
:param delay: 延迟时长(单位为秒),停顿超该时长则刷新
:yield: AsyncGenerator
"""
# 实例防抖器
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.new_messages = result.new_messages()
yield "05:FinalResultEvent"
case _:
yield f"06:未知事件类型{type(event).__name__}"