234 lines
7.1 KiB
Python
234 lines
7.1 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
应用状态管理模块
|
||
"""
|
||
|
||
from enum import StrEnum
|
||
from uuid import uuid4
|
||
from typing import Any, AsyncGenerator, Dict, List
|
||
|
||
from pydantic import BaseModel, Field
|
||
import reflex
|
||
|
||
from sys import path
|
||
from pathlib import Path
|
||
|
||
path.append(Path(__file__).parent.parent.parent.parent.as_posix())
|
||
from utils.agent import Agent
|
||
|
||
# 所有会话绑定的智能体
|
||
agents: Dict[str, Agent] = {}
|
||
|
||
|
||
def get_current_session_agent(state) -> Agent:
|
||
"""
|
||
获取当前会话绑定的智能体
|
||
:return: 当前会话绑定的智能体
|
||
"""
|
||
current_session_name = state.current_session_name
|
||
if current_session_name not in agents:
|
||
agents[current_session_name] = Agent(
|
||
session_id=uuid4().hex,
|
||
instructions="You are a friendly chatbot",
|
||
)
|
||
return agents[current_session_name]
|
||
|
||
|
||
class MessageType(StrEnum):
|
||
"""消息类型类"""
|
||
|
||
THINKING = "thinking"
|
||
TEXT = "text"
|
||
CALL = "call"
|
||
TOOL_ARGS = "tool_args"
|
||
TOOL_RETURN = "tool_return"
|
||
RESULT = "result"
|
||
ERROR = "error"
|
||
|
||
|
||
# 消息类型前缀映射
|
||
MESSAGE_TYPE_PREFIX_MAP = {f"{i:02d}:": mt for i, mt in enumerate(MessageType)}
|
||
|
||
|
||
# 会话、对话、消息关系:一次会话包含若干轮对话,每轮对话包含若干条消息
|
||
class Message(BaseModel):
|
||
"""消息类"""
|
||
|
||
id: str = Field(default_factory=lambda: uuid4().hex, description="消息唯一标识")
|
||
type_: MessageType = Field(..., description="类型")
|
||
content: str = Field(default="", description="内容")
|
||
|
||
|
||
class Turn(BaseModel):
|
||
"""对话类"""
|
||
|
||
id: str = Field(default_factory=lambda: uuid4().hex, description="对话唯一标识")
|
||
input_: str = Field(..., description="输入消息")
|
||
output: List[Message] = Field(default_factory=list, description="输出消息")
|
||
|
||
|
||
class Session(BaseModel):
|
||
"""会话类"""
|
||
|
||
is_processing: bool = Field(
|
||
default=False,
|
||
description="会话状态:True 表示正在处理中,False 表示未处理或处理完成",
|
||
)
|
||
turns: List[Turn] = Field(default_factory=list, description="对话列表")
|
||
|
||
|
||
class State(reflex.State):
|
||
"""统一管理应用数据与功能状态,作为前后端交互枢纽,借助响应式特性实现页面自动更新"""
|
||
|
||
# 当前会话名称
|
||
current_session_name: str = "新会话"
|
||
# 会话列表(会话名称作为会话对象的唯一标识,不允许重复)
|
||
sessions: Dict[str, Session] = {current_session_name: Session()}
|
||
|
||
# 新建会话模态窗是否打开
|
||
create_session_modal_is_open: bool = False
|
||
|
||
@reflex.var
|
||
def get_session_names(self) -> List[str]:
|
||
"""
|
||
获取会话名称列表
|
||
:return: 会话名称列表
|
||
"""
|
||
return list(self.sessions)
|
||
|
||
@reflex.event
|
||
def switch_session(self, session_name: str) -> None:
|
||
"""
|
||
切换会话
|
||
:param session_name: 会话名称
|
||
:return: None
|
||
"""
|
||
self.current_session_name = session_name
|
||
|
||
@reflex.var
|
||
def get_current_session_status(self) -> bool:
|
||
"""
|
||
获取当前会话状态
|
||
:return: 当前会话状态
|
||
"""
|
||
if self.current_session_name not in self.sessions:
|
||
return False
|
||
return self.sessions[self.current_session_name].is_processing
|
||
|
||
@reflex.var
|
||
def get_current_session_turns(self) -> List[Turn]:
|
||
"""
|
||
获取当前会话对话列表
|
||
:return: 对话列表
|
||
"""
|
||
if self.current_session_name not in self.sessions:
|
||
return []
|
||
return self.sessions[self.current_session_name].turns
|
||
|
||
@reflex.event
|
||
def create_session(self, form_data: Dict[str, Any]) -> None:
|
||
"""
|
||
新建会话
|
||
:param form_data: 新建会话表单数据
|
||
:return: None
|
||
"""
|
||
session_name = form_data["session_name"].strip()
|
||
|
||
# 若新建会话名称为空则默认使用"新会话"作为会话名称
|
||
if not session_name:
|
||
session_name = "新会话"
|
||
|
||
original_session_name = session_name
|
||
counter = 1
|
||
# 若会话名称重复则在会话名称后面添加标号至不重复
|
||
while session_name in self.sessions:
|
||
session_name = f"{original_session_name}({counter})"
|
||
counter += 1
|
||
|
||
self.current_session_name = session_name
|
||
self.sessions[session_name] = Session()
|
||
|
||
# 关闭新建会话模态窗
|
||
self.create_session_modal_is_open = False
|
||
|
||
@reflex.event
|
||
def delete_session(self, session_name: str) -> None:
|
||
"""
|
||
删除会话
|
||
:param session_name: 会话名称
|
||
:return: None
|
||
"""
|
||
if session_name not in self.sessions:
|
||
return
|
||
del self.sessions[session_name]
|
||
|
||
# 若会话列表为空则新建会话(区别于 create_session 方法,此处为后台新建会话)
|
||
if not self.sessions:
|
||
self.sessions["新会话"] = Session()
|
||
|
||
# 若当前会话名称不存在则默认使用第一个会话名称
|
||
if self.current_session_name not in self.sessions:
|
||
self.current_session_name = next(iter(self.sessions))
|
||
|
||
@reflex.event
|
||
def toggle_create_session_modal(self, is_open: bool) -> None:
|
||
"""
|
||
打开 / 关闭新建会话模态窗
|
||
:param is_open: 打开或关闭新建会话模态窗
|
||
:return: None
|
||
"""
|
||
self.create_session_modal_is_open = is_open
|
||
|
||
@reflex.event
|
||
async def adapt_input(self, form_data: dict[str, Any]) -> AsyncGenerator:
|
||
"""
|
||
适配输入
|
||
:param form_data: 输入栏组件的表单数据
|
||
:return: AsyncGenerator
|
||
"""
|
||
input_ = form_data["input"].strip()
|
||
if not input_:
|
||
return
|
||
|
||
# 当前会话
|
||
current_session = self.sessions[self.current_session_name]
|
||
# 当前会话正在处理
|
||
current_session.is_processing = True
|
||
# 将输入添加到当前会话对话列表
|
||
current_session.turns.append(
|
||
Turn(
|
||
input_=input_,
|
||
)
|
||
)
|
||
yield # 通知前端渲染输入消息
|
||
|
||
# 当前对话
|
||
current_turn = current_session.turns[-1]
|
||
|
||
# 获取当前会话绑定的智能体
|
||
agent = get_current_session_agent(self)
|
||
async for event in agent.stream_messages_events(user_prompt=input_):
|
||
# 跳过空事件
|
||
if not event:
|
||
continue
|
||
|
||
# 匹配消息类型
|
||
prefix = next(
|
||
(t for t in MESSAGE_TYPE_PREFIX_MAP if event.startswith(t)), None
|
||
)
|
||
# 跳过未匹配事件
|
||
if not prefix:
|
||
continue
|
||
|
||
# 消息类型
|
||
type_ = MESSAGE_TYPE_PREFIX_MAP[prefix]
|
||
# 若当前对话输出为空或当前消息类型和上一个消息类型不一致则创建消息
|
||
if not current_turn.output or current_turn.output[-1].type_ != type_:
|
||
current_turn.output.append(Message(type_=type_))
|
||
current_turn.output[-1].content += event.removeprefix(prefix)
|
||
|
||
yield # 通知前端渲染输出消息
|
||
|
||
# 当前会话处理完成
|
||
current_session.is_processing = False
|