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 typing import AsyncGenerator, List, Optional, Union
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.messages import (
AgentStreamEvent,
ModelMessage,
PartStartEvent,
TextPart,
ThinkingPart,
ToolSearchCallPart,
ThinkingPartDelta,
ToolCallPart,
LoadCapabilityCallPart,
TextPartDelta,
ToolCallPartDelta,
FunctionToolResultEvent,
PartDeltaEvent,
TextPart,
PartEndEvent,
)
from pydantic_ai.models.openai import OpenAIChatModel
from pydantic_ai.output import OutputSpec
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):
"""
片段
事件类
"""
part_index: Optional[int] = Field(default=None, description="片段索引")
kind: Kind = Field(..., description="事件种类")
tool_name: Optional[str] = Field(default=None, description="工具名称")
args: Optional[LoadCapabilityArgs] = Field(default=None, description="工具参数")
event_kind: str = Field(..., 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 封装的智能体
"""
@ -91,7 +89,10 @@ class Agent:
# 聊天唯一标识
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] = []
@ -101,7 +102,7 @@ class Agent:
instructions = DEFAULT_INSTRUCTIONS
# 初始化智能体
self.agent = PydanticAIAgent(
self.agent = Agent(
model=OpenAIChatModel(
model_name="deepseek-v4-flash",
provider=OpenAIProvider(
@ -132,24 +133,134 @@ class Agent:
) as events:
async for event in events:
match event:
case PartStartEvent():
part = event.part
# ========== 分片开始事件 ==========
case PartStartEvent(event_kind=event_kind, index=index, part=part):
match part:
case TextPart(content=content):
# 思考分片开始事件
case ThinkingPart(part_kind=part_kind, content=content):
yield Event(
kind=Kind.TEXTSTART,
part_index=event.index,
event_kind=event_kind,
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(
kind=Kind.THINKINGSTART,
part_index=event.index,
event_kind=event_kind,
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,
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,
)