# -*- coding: utf-8 -*- """ 数据库状态 """ from typing import List import reflex as rx import sqlalchemy as sa from sqlmodel import select as sql_select from application.models import MessageHistory, ModelMessage, ModelMessagesTypeAdapter class DatabaseState(rx.State): """数据库状态""" async def get_message_history(self, conversation_id: str) -> List[ModelMessage]: """ 获取会话消息历史 :param conversation_id: 会话唯一标识 :return: 会话消息历史 """ message_history = [] async with rx.asession() as session: records = await session.exec( sql_select(MessageHistory) .where(MessageHistory.conversation_id == conversation_id) .order_by(sa.desc("run_id")) ) for record in records.all(): message_history.extend( ModelMessagesTypeAdapter.validate_json(record.run_new_message) ) return message_history async def save_new_message( self, conversation_id: str, run_id: str, run_new_message: List[ModelMessage] ) -> None: """ 保存运行新增消息 :param conversation_id: 会话唯一标识 :param run_id: 运行唯一标识 :param run_new_message: 运行新增消息 :return: None """ async with rx.asession() as session: record = MessageHistory.adapt( conversation_id=conversation_id, run_id=run_id, run_new_message=run_new_message, ) session.add(record) await session.commit()