# -*- coding: utf-8 -*- """ 数据库状态 """ from datetime import datetime, timedelta from random import choices from typing import Any from pydantic_ai import ModelMessage, ModelMessagesTypeAdapter from pydantic_ai._uuid import uuid7 import reflex as rx from sqlalchemy import desc from sqlmodel import Field, JSON, SQLModel, select, update from application.states.models import ( Conversation, TaskType, TaskStatus, RunStatus, MessageType, Task, Run, Message, deps_to_object, usage_to_object, usage_limits_to_object, ) class VerificationCodeTable(SQLModel, table=True): """ 验证码表 """ email: str = Field(..., primary_key=True, description="邮箱") verification_code: str = Field( default_factory=lambda: "".join(choices("0123456789", k=6)), primary_key=True, description="验证码", ) is_valid: bool = Field( default=True, index=True, description="验证码有效:True 表示有效,False 表示无效", ) is_verified: bool = Field( default=False, index=True, description="验证码已核验:True 表示已核验,False 表示未核验", ) expired_at: datetime = Field( default_factory=lambda: datetime.now() + timedelta(minutes=30), primary_key=True, description="过期时间", ) class UserTable(SQLModel, table=True): """ 用户表 """ id: str = Field( default_factory=lambda: str(uuid7()), primary_key=True, description="用户唯一标识", ) email: str = Field(..., index=True, description="邮箱") class ConversationTable(SQLModel, table=True): """ 会话表 """ id: str = Field( default_factory=lambda: str(uuid7()), primary_key=True, description="会话唯一标识", ) user_id: str = Field(..., index=True, description="用户唯一标识") description: str = Field(default="新会话", description="会话描述") is_deleted: bool = Field( default=False, index=True, description="会话已删除:True 表示已删除,False 表示未删除", ) created_at: datetime = Field( default_factory=datetime.now, description="会话创建时间" ) class TaskTable(SQLModel, table=True): """ 任务表 """ id: str = Field( default_factory=lambda: str(uuid7()), primary_key=True, description="任务唯一标识", ) conversation_id: str = Field(..., index=True, description="会话唯一标识") type: TaskType = Field(default=TaskType.CHAT, description="任务类型") status: TaskStatus = Field(default=TaskStatus.NONE, description="任务状态") deps: dict[str, Any] = Field( default_factory=dict, sa_type=JSON, description="任务依赖项" ) usage: dict[str, Any] = Field( default_factory=dict, sa_type=JSON, description="任务使用量" ) usage_limits: dict[str, Any] = Field( default_factory=dict, sa_type=JSON, description="任务使用量限制" ) class RunTable(SQLModel, table=True): """ 运行表 """ id: str = Field( default_factory=lambda: str(uuid7()), primary_key=True, description="运行唯一标识", ) task_id: str = Field(..., index=True, description="任务唯一标识") status: RunStatus = Field(default=RunStatus.RUNNING, description="运行状态") usage: dict[str, Any] = Field( default_factory=dict, sa_type=JSON, description="运行使用量" ) usage_limits: dict[str, Any] = Field( default_factory=dict, sa_type=JSON, description="运行使用量限制" ) class MessageTable(SQLModel, table=True): """ 消息表 """ id: str = Field( default_factory=lambda: str(uuid7()), primary_key=True, description="消息唯一标识", ) run_id: str = Field(..., index=True, description="运行唯一标识") type: MessageType = Field(default=MessageType.USER_PROMPT, description="消息类型") title: str = Field(default="", description="消息标题") content: str = Field(default="", description="消息内容") class DatabaseState(rx.State): """ 数据库状态 """ async def create_verification_code_record(self, email: str) -> str: """ 创建验证码记录 :param email: 邮箱 :return: 验证码 """ async with rx.asession() as session: # 先将该邮箱的有效、未核验的验证码记录设置为无效 await session.exec( update(VerificationCodeTable) .where( VerificationCodeTable.email == email, # type: ignore VerificationCodeTable.is_valid == True, # type: ignore VerificationCodeTable.is_verified == False, # type: ignore ) .values(is_valid=False) ) await session.flush() # 创建验证码记录 record = VerificationCodeTable( email=email, ) session.add(record) await session.commit() await session.refresh(record) return record.verification_code async def verify_verification_code( self, email: str, verification_code: str ) -> bool: """ 核验验证码 :param email: 邮箱 :param verification_code: 验证码 :return: 是否核验成功,True 表示核验成功,False 表示核验失败(根据邮箱和验证码未查询到有效且未核验的记录,或失效) """ async with rx.asession() as session: # 查询该邮箱验证码且有效、未核验、过期时间大于当前时间的验证码记录 result = await session.exec( select(VerificationCodeTable).where( VerificationCodeTable.email == email, VerificationCodeTable.verification_code == verification_code, VerificationCodeTable.is_valid == True, VerificationCodeTable.is_verified == False, VerificationCodeTable.expired_at > datetime.now(), ) ) record = result.first() if not record: return False record.is_verified = True await session.commit() await session.refresh(record) return True async def create_user_record(self, email: str) -> str: """ 创建用户记录 :param email: 邮箱 :return: 用户唯一标识 """ async with rx.asession() as session: result = await session.exec( select(UserTable).where(UserTable.email == email) ) record = result.first() # 若用户记录已存在则返回用户唯一标识,否则先创建用户记录再返回用户唯一标识 if record: return record.id record = UserTable(email=email) session.add(record) await session.commit() await session.refresh(record) return record.id async def get_conversations(self, user_id: str) -> dict[str, Conversation]: """ 获取当前用户的会话字典 :param user_id: 用户唯一标识 :return: 当前用户的会话字典 """ records: dict[str, Conversation] = {} async with rx.asession() as session: result = await session.exec( select(ConversationTable, RunTable, MessageTable) .outerjoin(RunTable, RunTable.conversation_id == ConversationTable.id) # type: ignore .outerjoin(MessageTable, MessageTable.run_id == RunTable.id) # type: ignore .where( ConversationTable.user_id == user_id, ConversationTable.is_deleted == False, ) .order_by( ConversationTable.id, TaskTable.id, RunTable.id, MessageTable.id ) ) for ( conversation_result, run_result, message_result, ) in result.all(): conversation_record = records.setdefault( conversation_result.id, Conversation( id=conversation_result.id, description=conversation_result.description, created_at=conversation_result.created_at, ), ) if not run_result: continue run_record = conversation_record.runs.setdefault( run_result.id, Run( id=run_result.id, status=run_result.status, usage=usage_to_object(run_result.usage), usage_limits=usage_limits_to_object(run_result.usage_limits), ), ) if not message_result: continue run_record.messages.setdefault( message_result.id, Message( id=message_result.id, type=message_result.type, title=message_result.title, content=message_result.content, ), ) return records async def create_conversation_record(self, user_id: str) -> dict[str, Conversation]: """ 创建会话记录 :param user_id: 用户唯一标识 :return: 会话实例 """ async with rx.asession() as session: record = ConversationTable(user_id=user_id) session.add(record) await session.commit() await session.refresh(record) return { record.id: Conversation( id=record.id, description=record.description, created_at=record.created_at, ) } async def delete_conversation_record(self, conversation_id: str) -> None: """ 删除会话记录(逻辑删除) :param conversation_id: 指定会话唯一标识 :return: None """ async with rx.asession() as session: record = await session.get(ConversationTable, conversation_id) if not record: return record.is_deleted = True await session.commit() async def create_dialog_record( self, conversation_id: str, id: str, user_prompt: str, thoughts: dict[int, Any], result_output: str, usage: dict[str, Any], ) -> dict[str, Dialog]: """ 创建对话记录 :param conversation_id: 会话唯一标识 :param user_prompt: 用户提示词 :param result_output: 结果输出 :return: 创建对话记录的唯一标识 """ async with rx.asession() as session: record = DialogRecord( id=id, conversation_id=conversation_id, user_prompt=user_prompt, thoughts=thoughts, result_output=result_output, usage=usage, ) session.add(record) await session.commit() await session.refresh(record) return { record.id: Dialog( id=record.id, user_prompt=record.user_prompt, thoughts=record.thoughts, result_output=record.result_output, usage=record.usage, ) } async def create_result_record( self, conversation_id: str, dialog_id: str, new_messages: list[ModelMessage], ) -> None: """ 创建结果记录 :param conversation_id: 会话唯一标识 :param dialog_id: 对话唯一标识 :param new_messages: 新增消息 :return: None """ async with rx.asession() as session: session.add( ResultRecord( conversation_id=conversation_id, dialog_id=dialog_id, new_messages=ModelMessagesTypeAdapter.dump_json( new_messages ).decode( "utf-8" ), # 序列化为 JSON 字符串 ) ) await session.commit() async def get_message_history(self, conversation_id: str) -> list[ModelMessage]: """ 获取消息历史列表 :param conversation_id: 会话唯一标识 :return: 消息历史 """ records: list[ModelMessage] = [] async with rx.asession() as session: result = await session.exec( select(ResultRecord) .where(ResultRecord.conversation_id == conversation_id) .order_by(desc(ResultRecord.dialog_id)) ) for record in result.all(): records.extend( ModelMessagesTypeAdapter.validate_json(record.new_messages) ) return records