# -*- 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__}"