# -*- coding: utf-8 -*- """ 对话状态 """ from typing import Any, AsyncGenerator, Dict, List, Optional, Tuple, cast from pydantic_ai._uuid import uuid7 import reflex as rx from application.models import Conversation, EventKind, PartKind, Run from application.state.create_conversation_modal import CreateConversationModalState from application.state.database import DatabaseState from application.utils.agent import AIAgent class ConversationState(rx.State): """ 对话状态 """ # 当前对话唯一标识 conversation_id: str = str(uuid7()) # 对话字典 conversations: Dict[str, Conversation] = {conversation_id: Conversation()} # 当前运行唯一标识 run_id: Optional[str] = None # 初始化智能体 agent: AIAgent = AIAgent() @rx.var def get_conversations(self) -> Dict[str, Conversation]: """ 获取对话字典,用于前端按照创建倒序渲染对话历史 :return: 对话字典 """ return dict(reversed(self.conversations.items())) @rx.event async def create_conversation(self, form_data: Dict[str, Any]) -> None: """ 创建对话 :param form_data: 表单数据 :return: None """ # 获取描述 description = form_data["description"].strip() if not description: description = "新对话" # 创建对话唯一标识 self.conversation_id = str(uuid7()) self.conversations[self.conversation_id] = Conversation(description=description) # 获取创建对话模态窗状态 create_conversation_modal_state = await self.get_state( CreateConversationModalState ) # 关闭创建对话模态窗 create_conversation_modal_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)) @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_runs(self) -> Dict[str, Run]: """ 获取运行字典,用于前端渲染运行历史 :return: 运行字典 """ # 当前对话 conversation = self.conversations.get(self.conversation_id) return conversation.runs if conversation else {} @rx.var def get_run_streaming_status(self) -> bool: """ 获取当前运行流式输出状态 :return: 当前运行流式输出状态(True 表示正在流式输出,False 表示非正在流式输出) """ # 当前对话 conversation = self.conversations.get(self.conversation_id) if not conversation or not conversation.runs: return False # 末位运行 run = next(reversed(conversation.runs.values())) return run.is_streaming @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 ) # 运行,处理分片生命周期事件 async with self.agent.run( conversation_id=self.conversation_id, user_prompt=user_prompt, message_history=message_history, ) as async_event: if ( not event or not self.current_run_id or self.current_run_id != event.run_id ): continue # ========== 运行开始 ========== if event.event_kind == EventKind.RUN_START: # 将事件运行唯一标识设置为当前运行唯一标识 self.current_run_id = event.run_id # 创建运行 conversation.runs[self.current_run_id] = Run( user_prompt=user_prompt, is_streaming=True ) yield # 通知前端渲染 continue # 当前运行 current_run = current_conversation.runs[self.current_run_id] # ========== 推理 ========== if ( event.event_kind == EventKind.PART_START and event.part_kind == PartKind.THINKING ): current_run.thinking += event.event_content yield # 通知前端渲染 continue if part_index := event.part_index: if part_index not in current_run.reasonings: current_run.reasonings[part_index] += event.event_content yield # 通知前端渲染 continue # ========== 回答 ========== if ( event.part_kind == PartKind.TEXT and event.event_content != EventKind.PART_END ): current_run.answer += event.event_content yield # 通知前端渲染 continue # ========== 运行结束 ========== if event.event_kind == EventKind.RUN_END: # 保存运行新增消息 await database_state.save_new_message( conversation_id=self.current_conversation_id, run_id=self.current_run_id, run_new_message=event.run_new_messages, ) # 设置当前运行流式输出状态非正在流式输出 current_run.is_streaming = False yield # 通知前端渲染 continue