This commit is contained in:
liubiren 2026-06-26 20:01:54 +08:00
parent 249a38e6b9
commit 97a8bb09fe
1 changed files with 146 additions and 35 deletions

View File

@ -7,16 +7,27 @@ Pydantic AI 聊天智能体和相关模块
from enum import StrEnum from enum import StrEnum
from typing import AsyncGenerator, List, Optional, Union from typing import AsyncGenerator, List, Optional, Union
from uuid import uuid4 from uuid import uuid4
from pydantic_ai import Agent as PydanticAIAgent from numpy._core.numeric import str_
from pydantic_ai import Agent
from pydantic_ai.capabilities import AgentCapability from pydantic_ai.capabilities import AgentCapability
from pydantic_ai.messages import ( from pydantic_ai.messages import (
AgentStreamEvent, AgentStreamEvent,
ModelMessage, ModelMessage,
PartStartEvent, PartStartEvent,
TextPart,
ThinkingPart, ThinkingPart,
ToolSearchCallPart,
ThinkingPartDelta,
ToolCallPart,
LoadCapabilityCallPart, LoadCapabilityCallPart,
TextPartDelta,
ToolCallPartDelta,
FunctionToolResultEvent,
PartDeltaEvent,
TextPart,
PartEndEvent,
) )
from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.models.openai import OpenAIChatModel
from pydantic_ai.output import OutputSpec from pydantic_ai.output import OutputSpec
from pydantic_ai.providers.openai import OpenAIProvider from pydantic_ai.providers.openai import OpenAIProvider
@ -41,32 +52,19 @@ DEFAULT_INSTRUCTIONS: str = """
""" """
class Kind(StrEnum):
"""种类"""
TEXTSTART = "text_start"
THINKINGSTART = "thinking_start"
TOOL_NAME = "tool_name"
TOOL_ARGS = "tool_args"
TOOL_RETURN = "tool_return"
FINISHED = "finished"
ERROR = "error"
class Event(BaseModel): class Event(BaseModel):
""" """
片段 事件类
""" """
part_index: Optional[int] = Field(default=None, description="片段索引") event_kind: str = Field(..., description="事件种类")
kind: Kind = Field(..., description="事件种类")
tool_name: Optional[str] = Field(default=None, description="工具名称")
args: Optional[LoadCapabilityArgs] = Field(default=None, description="工具参数")
content: Optional[str] = Field(default=None, description="事件内容") content: Optional[str] = Field(default=None, description="事件内容")
part_index: Optional[int] = Field(default=None, description="分片索引")
part_kind: str = Field(..., description="分片种类")
tool_name: Optional[str] = Field(default=None, description="工具名称")
class Agent: class AIAgent:
""" """
基于 Pydantic AI 封装的智能体 基于 Pydantic AI 封装的智能体
""" """
@ -91,7 +89,10 @@ class Agent:
# 聊天唯一标识 # 聊天唯一标识
self.chat_id = chat_id self.chat_id = chat_id
# 一次聊天chat包含若干论对话dialog每轮对话包含用户提示词user_prompt和输出output。其中输出包含若干片段Part # 一次聊天chat包含若干论对话dialog每轮对话包含用户提示词user_prompt和输出output。其中输出包含若干分片Part
# 本轮对话工具调用唯一标识与片段索引映射表
self.tool_call_ids: dict[str, int] = {}
# 本轮对话新增消息列表 # 本轮对话新增消息列表
self.new_messages: List[ModelMessage] = [] self.new_messages: List[ModelMessage] = []
@ -101,7 +102,7 @@ class Agent:
instructions = DEFAULT_INSTRUCTIONS instructions = DEFAULT_INSTRUCTIONS
# 初始化智能体 # 初始化智能体
self.agent = PydanticAIAgent( self.agent = Agent(
model=OpenAIChatModel( model=OpenAIChatModel(
model_name="deepseek-v4-flash", model_name="deepseek-v4-flash",
provider=OpenAIProvider( provider=OpenAIProvider(
@ -132,24 +133,134 @@ class Agent:
) as events: ) as events:
async for event in events: async for event in events:
match event: match event:
case PartStartEvent(): # ========== 分片开始事件 ==========
part = event.part case PartStartEvent(event_kind=event_kind, index=index, part=part):
match part: match part:
case TextPart(content=content): # 思考分片开始事件
case ThinkingPart(part_kind=part_kind, content=content):
yield Event( yield Event(
kind=Kind.TEXTSTART, event_kind=event_kind,
part_index=event.index,
content=content, content=content,
part_index=index,
part_kind=part_kind,
) )
case ThinkingPart(content=content): # 检索工具分片开始事件
case ToolSearchCallPart(
part_kind=part_kind, tool_call_id=tool_call_id, tool_name=tool_name
):
# 记录工具调用唯一标识与片段索引映射
self.tool_call_ids[tool_call_id] = index
yield Event( yield Event(
kind=Kind.THINKINGSTART, event_kind=event_kind,
part_index=event.index, part_index=index,
part_kind=part_kind,
tool_name=tool_name,
)
# 加载能力分片开始事件
case LoadCapabilityCallPart(
part_kind=part_kind, tool_call_id=tool_call_id, tool_name=tool_name
):
# 记录工具调用唯一标识与片段索引映射
self.tool_call_ids[tool_call_id] = index
yield Event(
event_kind=event_kind,
part_index=index,
part_kind=part_kind,
tool_name=tool_name,
)
# 调用工具分片开始事件
case ToolCallPart(
part_kind=part_kind,
tool_call_id=tool_call_id,
tool_name=tool_name,
):
# 记录工具调用唯一标识与片段索引映射
self.tool_call_ids[tool_call_id] = index
yield Event(
event_kind=event_kind,
part_index=index,
part_kind=part_kind,
tool_name=tool_name,
)
# 文本分片开始事件
case TextPart(part_kind=part_kind, content=content):
yield Event(
event_kind=event_kind,
content=content, content=content,
part_index=index,
part_kind=part_kind,
) )
case LoadCapabilityCallPart():
yield Event(
kind=Kind.TOOL_NAME,
part_index=event.index,
# ========== 分片增量事件(前端仅就文本片段实现打字机效果) ==========
case PartDeltaEvent(
event_kind=event_kind, index=index, delta=delta
):
match delta:
# 文本分片增量事件
case TextPartDelta(
part_delta_kind=part_delta_kind,
content_delta=content_delta,
):
yield Event(
event_kind=event_kind,
content=content_delta, # 增量
part_index=index,
part_kind=part_delta_kind,
) )
# ========== 分片结束事件 ==========
case PartEndEvent(event_kind=event_kind, index=index, part=part):
match part:
# 思考分片结束事件
case ThinkingPart(part_kind=part_kind, content=content):
yield Event(
event_kind=event_kind,
content=content,
part_index=index,
part_kind=part_kind,
)
# 检索工具分片结束事件
case ToolSearchCallPart(part_kind=part_kind, tool_name=tool_name):
yield Event(
event_kind=event_kind,
part_index=index,
part_kind=part_kind,
tool_name=tool_name,
)
# 加载能力分片结束事件
case LoadCapabilityCallPart(part_kind=part_kind, tool_name=tool_name):
yield Event(
event_kind=event_kind,
part_index=index,
part_kind=part_kind,
tool_name=tool_name,
)
# 调用工具分片结束事件
case ToolCallPart(part_kind=part_kind, tool_name=tool_name):
yield Event(
event_kind=event_kind,
part_index=index,
part_kind=part_kind,
tool_name=tool_name,
)
# 文本分片结束事件
case TextPart(part_kind=part_kind):
yield Event(
event_kind=event_kind,
part_index=index,
part_kind=part_kind,
)
# ========== 函数工具结果回调事件 ==========
case FunctionToolResultEvent(
event_kind=event_kind,
content=content,
part=part,
):
yield Event(
event_kind=event_kind,
content=content,
part_index=index,
part_kind=part_kind,
tool_name=tool_name,
)