Python/agent/application/states/conversation.py

427 lines
16 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 -*-
"""
智能体状态
"""
from datetime import datetime
from enum import StrEnum
from typing import Any, AsyncGenerator, Dict
from typing import Dict, List
from pydantic import BaseModel, Field
from pydantic_ai import Agent, ThinkingPartDelta
from pydantic_ai._uuid import uuid7
from pydantic_ai._uuid import uuid7
from pydantic_ai.messages import (
FunctionToolCallEvent,
FunctionToolResultEvent,
LoadCapabilityCallPart,
PartDeltaEvent,
PartEndEvent,
PartStartEvent,
TextPart,
TextPartDelta,
ThinkingPart,
ToolCallPart,
ToolSearchCallPart,
)
from pydantic_ai.messages import ModelMessage, ModelMessagesTypeAdapter
from pydantic_ai.models.openai import OpenAIChatModel
from pydantic_ai.providers.openai import OpenAIProvider
from pydantic_ai.run import AgentRunResultEvent
import reflex as rx
from sqlmodel import Field as SqlField, SQLModel
from application.states import CreateConversationState, DatabaseState
class ReasoningPartKind(StrEnum):
"""
推理分片种类
"""
THINKING = "thinking"
TOOL_SEARCH = "tool-search"
CAPABILITY_LOAD = "load-capability"
TOOL_CALL = "tool-call"
TEXT = "text"
TOOL_RETURN = "tool-return"
RETRY_PROMPT = "retry-prompt"
RUN_RETURN = "run-return"
class ReasoningPart(BaseModel):
"""
推理分片内存模型
"""
kind: ReasoningPartKind = Field(..., description="推理分片种类")
content: str = Field(default="", description="推理分片内容")
class Dialog(BaseModel):
"""
对话内存模型
"""
user_prompt: str = Field(..., description="用户提示词")
is_reasoning: bool = Field(
default=False,
description="推理状态True 表示正在推理False 表示推理完成",
)
is_collapse_open: bool = Field(
default=False,
description="推理面板展开状态True 表示推理面板展开False 表示推理面板折叠",
)
reasoning_parts: Dict[int, ReasoningPart] = Field(
default_factory=dict, description="推理分片字典"
)
assistant_content: str = Field(default="", description="回复正文")
is_running: bool = Field(
default=False,
description="运行状态True 表示正在运行False 表示运行完成",
)
class Conversation(BaseModel):
"""
会话内存模型
"""
description: str = Field(default="新会话", description="会话描述")
dialogs: Dict[str, Dialog] = Field(default_factory=dict, description="对话字典")
instructions: str = """
# 角色
专业友好AI助手结构化解答各类问题。
# 输出硬性规则
1. 全文强制标准Markdown禁止纯文本不要额外说明排版格式直接输出内容
2. 层级使用 `#/##/###`,列表用 `-` 无序列表或数字有序列表;
3. 代码块用 ```语言名``` 包裹;
4. 重点内容标注 **粗体**/*斜体*
5. 思考、工具日志仅输出文本适配前端折叠面板禁止输出HTML标签
6. 内容分点拆分排版整洁适配前端Markdown渲染。
# 行文要求
语言通俗,逻辑完整简洁,无多余废话。
"""
# 实例化智能体(因无法序列化故剥离出状态管理)
agent: Agent = Agent(
model=OpenAIChatModel(
model_name="deepseek-v4-flash",
provider=OpenAIProvider(
base_url="https://tokenhub.tencentmaas.com/v1",
api_key="sk-D9Y1mCe8VlvNqLuSC4mAjqEwxJ2nW4C0h8a7EPn8kg9RLsHq",
),
),
instructions=instructions,
capabilities=None,
output_type=str,
retries=1,
)
class ConversationState(rx.State):
"""
会话状态
"""
# 初始化会话字典
conversations: Dict[str, Conversation] = {str(uuid7()): Conversation()}
# 获取当前会话唯一标识
conversation_id: str = next(reversed(conversations.keys()))
@rx.var
def get_chats(self) -> Dict[str, Conversation]:
"""
获取会话字典,用于前端渲染会话历史
:return: 会话字典
"""
return dict(reversed(self.conversations.items()))
@rx.event
async def create_chat(self, form_data: Dict[str, Any]) -> None:
"""
新建会话
:param form_data: 表单数据
:return: None
"""
# 获取会话描述
description = form_data["description"].strip()
if not description:
description = "新会话"
# 新建会话
self.conversations[str(uuid7())] = Conversation(description=description)
# 将末位会话唯一标识设置为当前会话唯一标识
self.conversation_id = next(reversed(self.conversations.keys()))
# 获取新建会话状态
create_chat_state = await self.get_state(CreateConversationState)
# 关闭新建会话模态窗
create_chat_state.is_open = False
@rx.event
def delete_conversation(self, conversation_id: str) -> None:
"""
删除会话
:param conversation_id: 会话唯一标识
:return: None
"""
if conversation_id not in self.conversations:
return
del self.conversations[conversation_id]
# 删除后,若会话字典为空则创建会话
if not self.conversations:
self.conversations[str(uuid7())] = Conversation()
# 删除后,若当前会话唯一标识不存在则将末位会话唯一标识设置为当前会话唯一标识
if self.conversation_id not in self.conversations:
self.conversation_id = next(reversed(self.conversations.keys()))
@rx.event
def switch_conversation(self, conversation_id: str) -> None:
"""
将指定会话唯一标识设置为当前会话唯一标识
:param conversation_id: 指定会话唯一标识
:return: None
"""
if conversation_id not in self.conversations:
return
self.conversation_id = conversation_id
@rx.var
def get_conversation_description(self) -> str:
"""
获取当前会话描述,用于前端渲染会话标题
:return: 当前会话描述
"""
# 当前会话
conversation = self.conversations.get(self.conversation_id)
return conversation.description if conversation else "新会话"
@rx.var
def get_dialogs(self) -> Dict[str, Dialog]:
"""
获取对话字典,用于前端渲染对话历史
:return: 对话字典
"""
# 当前会话
conversation = self.conversations.get(self.conversation_id)
return conversation.dialogs if conversation else {}
@rx.var
def get_running_status(self) -> bool:
"""
获取当前运行状态
:return: 当前运行状态True 表示正在运行False 表示运行完成)
"""
# 当前会话
conversation = self.conversations.get(self.conversation_id)
if not conversation or not conversation.dialogs:
return False
# 末位对话
dialog = next(reversed(conversation.dialogs.values()))
return dialog.is_running
@rx.event
async def run(self, form_data: dict[str, Any]) -> AsyncGenerator[None]:
"""
运行
:param form_data: 表单数据
:return: AsyncGenerator
"""
# 当前会话
conversation = self.conversations.get(self.conversation_id)
if not conversation:
return
# 获取用户提示词
user_prompt = form_data.get("user_prompt", "").strip()
if not user_prompt:
return
# 获取数据库状态
database_state = await self.get_state(DatabaseState)
# 获取消息历史
message_history = await database_state.get_message_history(
conversation_id=self.conversation_id
)
# 初始化工具调用唯一标识和片段索引映射字典
tool_call_ids: Dict[str, int] = {}
# 创建对话
conversation.dialogs[str(uuid7())] = Dialog(user_prompt=user_prompt)
# 将末位对话唯一标识、对话设置为当前对话唯一标识、对话
dialog_id, dialog = next(reversed(conversation.dialogs.items()))
# 将运行状态设置为正在运行
dialog.is_running = True
yield # 通知前端渲染
async with agent.run_stream_events(
conversation_id=self.conversation_id,
user_prompt=user_prompt,
message_history=message_history,
) as events:
async for event in events:
match event:
# ========== 开始事件 ==========
case PartStartEvent(
index=index,
part=part,
previous_part_kind=previous_part_kind,
):
match part:
# 思考分片开始事件
case ThinkingPart(content=content):
# 若上一分片种类为空则将推理状态设置为正在推理、推理面板展开状态设置为展开
if not previous_part_kind:
dialog.is_reasoning = True
dialog.is_collapse_open = True
dialog.reasoning_parts[index] = ReasoningPart(
kind=ReasoningPartKind.THINKING, content=content
)
yield
# 工具检索分片开始事件
case ToolSearchCallPart(tool_call_id=tool_call_id):
# 创建工具调用唯一标识与片段索引映射
tool_call_ids[tool_call_id] = index
dialog.reasoning_parts[index] = ReasoningPart(
kind=ReasoningPartKind.TOOL_SEARCH,
content="正在生成检索关键词",
)
yield
# 能力加载分片开始事件
case LoadCapabilityCallPart(tool_call_id=tool_call_id):
tool_call_ids[tool_call_id] = index
dialog.reasoning_parts[index] = ReasoningPart(
kind=ReasoningPartKind.CAPABILITY_LOAD,
content="正在生成加载参数",
)
yield
# 工具调用分片开始事件
case ToolCallPart(tool_call_id=tool_call_id):
tool_call_ids[tool_call_id] = index
dialog.reasoning_parts[index] = ReasoningPart(
kind=ReasoningPartKind.TOOL_CALL,
content="正在生成调用参数",
)
yield
# 文本分片开始事件
case TextPart(content=content):
dialog.assistant_content = content
yield
# ========== 增量事件 ==========
case PartDeltaEvent(index=index, delta=delta):
match delta:
# 思考分片增量事件
case ThinkingPartDelta(
content_delta=content_delta,
):
dialog.reasoning_parts[index].content += content_delta or ""
yield
# 文本分片增量事件
case TextPartDelta(
content_delta=content_delta,
):
dialog.assistant_content += content_delta
yield
# ========== 结束事件 ==========
case PartEndEvent(
index=index,
part=part,
next_part_kind=next_part_kind,
):
match part:
# 思考分片结束事件
case ThinkingPart(part_kind=part_kind, content=content):
# 若下一分片种类为文本则将推理状态设置为推理完成、推理面板展开状态设置为折叠
if next_part_kind == ReasoningPartKind.TEXT:
dialog.is_reasoning = False
dialog.is_collapse_open = False
yield
# ========== 函数工具调用事件 ==========
case FunctionToolCallEvent(tool_call_id=tool_call_id, part=part):
# 获取分片索引
index = tool_call_ids[tool_call_id]
match dialog.reasoning_parts[index].kind:
# 工具检索
case ReasoningPartKind.TOOL_SEARCH:
dialog.reasoning_parts[index].content = "正在检索"
yield
# 能力加载
case ReasoningPartKind.CAPABILITY_LOAD:
dialog.reasoning_parts[index].content = (
f"正在加载能力 {part.tool_name}"
)
yield
# 工具调用
case ReasoningPartKind.TOOL_CALL:
dialog.reasoning_parts[index].content = (
f"正在调用工具 {part.tool_name}"
)
yield
# ========== 函数工具结果事件 ==========
case FunctionToolResultEvent(
tool_call_id=tool_call_id,
content=content,
):
index = tool_call_ids[tool_call_id]
match dialog.reasoning_parts[index].kind:
# 工具检索
case ReasoningPartKind.TOOL_SEARCH:
dialog.reasoning_parts[index].content = (
content if isinstance(content, str) else ""
) # 暂仅考虑文本内容
yield
# 能力加载
case ReasoningPartKind.CAPABILITY_LOAD:
dialog.reasoning_parts[index].content = "已加载"
yield
# 工具调用
case ReasoningPartKind.TOOL_CALL:
dialog.reasoning_parts[index].content = f"已调用"
yield
# ========== 智能体运行结果事件 ==========
case AgentRunResultEvent(result=result):
# 保存新增消息
await database_state.save_new_messages(
conversation_id=self.conversation_id,
dialog_id=dialog_id,
new_messages=result.new_messages(),
)
# 将运行状态设置为运行完成
dialog.is_running = False
yield
@rx.event
def toggle_reasoning_panel(self, dialog_id: str) -> None:
"""
展开/折叠指定运行唯一标识的推理面板
"""
# 指定运行
dialog = self.conversations[self.conversation_id].dialogs[dialog_id]
dialog.is_collapse_open = not dialog.is_collapse_open