Python/agent/application/states/database.py

269 lines
8.7 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 typing import Dict, List, cast
from pydantic import TypeAdapter
from pydantic_ai import ModelMessage, ModelMessagesTypeAdapter
import reflex as rx
from sqlalchemy import ColumnElement, desc
from sqlmodel import select, update
from application.models import (
Captchas,
Conversation,
Conversations,
Dialog,
Dialogs,
RunResults,
ThoughtNode,
Users,
)
# 思考节点内存模型类型适配器
ThoughtNodeTypeAdapter = TypeAdapter(dict[int, ThoughtNode])
class DatabaseState(rx.State):
"""
数据库状态
"""
async def create_captchas_record(self, email: str) -> str:
"""
创建验证码记录
:param email: 邮箱
:return: 验证码
"""
async with rx.asession() as session:
# 前置处理:将该邮箱有效且未使用的验证码记录的是否有效设置为无效
await session.exec(
update(Captchas)
.where(
cast(ColumnElement[bool], Captchas.email == email),
cast(
ColumnElement[bool],
Captchas.is_valid == True,
),
cast(
ColumnElement[bool],
Captchas.is_unused == True,
),
)
.values(is_valid=False)
)
await session.flush()
# 创建验证码记录
record = Captchas(
email=email,
)
session.add(record)
await session.commit()
await session.refresh(record)
return record.captcha
async def update_captcha(self, email: str, captcha: str) -> bool:
"""
更新验证码记录
:param email: 邮箱
:param captcha: 验证码
:return: 是否更新成功True 表示更新成功False 表示更新失败
"""
async with rx.asession() as session:
result = await session.exec(
select(Captchas).where(
Captchas.email == email,
Captchas.captcha == captcha,
Captchas.is_valid == True,
Captchas.is_unused == True,
Captchas.expired_at > datetime.now(),
)
)
record = result.first()
if not record:
return False
record.is_unused = False
await session.commit()
await session.refresh(record)
return True
async def create_users_record(self, email: str) -> str:
"""
创建用户记录
:param email: 邮箱
:return: 创建用户记录的唯一标识
"""
async with rx.asession() as session:
result = await session.exec(select(Users).where(Users.email == email))
record = result.first()
# 若用户记录存在则返回用户唯一标识,否则创建新用户记录并返回所创建的用户唯一标识
if record:
return record.id
new_user = Users(email=email)
session.add(new_user)
await session.commit()
await session.refresh(new_user)
return new_user.id
async def retrieve_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(Conversations, Dialogs)
.outerjoin(
Dialogs,
cast(
ColumnElement[bool],
Conversations.id == Dialogs.conversation_id,
),
)
.where(
Conversations.user_id == user_id,
Conversations.is_deleted == False,
)
.order_by(Conversations.id, Dialogs.id)
)
for conversation, dialog in result.all():
record = records.setdefault(
conversation.id,
Conversation(
description=conversation.description,
created_at=conversation.created_at,
),
)
if not dialog:
continue
record.dialogs[dialog.id] = Dialog(
question=dialog.question,
thought_nodes=ThoughtNodeTypeAdapter.validate_python(
dialog.thought_nodes
),
answer=dialog.answer,
)
return records
async def create_conversations_record(
self, user_id: str, description: str = "新会话"
) -> tuple[str, datetime]:
"""
创建会话记录
:param user_id: 用户唯一标识
:param description: 会话描述
:return: 创建会话记录的唯一标识和创建时间
"""
async with rx.asession() as session:
record = Conversations(user_id=user_id, description=description)
session.add(record)
await session.commit()
await session.refresh(record)
return record.id, record.created_at
async def delete_conversations_record(self, conversation_id: str) -> None:
"""
逻辑删除会话记录
:param conversation_id: 指定会话记录唯一标识
:return: None
"""
async with rx.asession() as session:
record = await session.get(Conversations, conversation_id)
if not record:
return
record.is_deleted = True
await session.commit()
async def create_dialogs_record(self, conversation_id: str, question: str) -> str:
"""
创建对话记录
:param conversation_id: 会话唯一标识
:param question: 问题
:return: 创建对话记录的唯一标识
"""
async with rx.asession() as session:
record = Dialogs(
conversation_id=conversation_id,
question=question,
)
session.add(record)
await session.commit()
await session.refresh(record)
return record.id
async def update_dialog_record(
self,
dialog_id: str,
thought_nodes: dict[int, ThoughtNode],
answer: str,
) -> None:
"""
更新对话记录
:param dialog_id: 对话唯一标识
:param thought_nodes: 思考节点列表
:param answer: 回答
:return: None
"""
async with rx.asession() as session:
record = await session.get(Dialogs, dialog_id)
if not record:
return
record.thought_nodes = ThoughtNodeTypeAdapter.dump_python(
{k: v for k, v in thought_nodes.items()}
)
record.answer = answer
await session.commit()
async def create_run_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(
RunResults(
conversation_id=conversation_id,
dialog_id=dialog_id,
new_messages=ModelMessagesTypeAdapter.dump_json(
new_messages
).decode(
"utf-8"
), # 序列化为 JSON 字符串
)
)
await session.commit()
async def retrieve_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(RunResults)
.where(RunResults.conversation_id == conversation_id)
.order_by(desc(RunResults.dialog_id))
)
for record in result.all():
records.extend(
ModelMessagesTypeAdapter.validate_json(record.new_messages)
)
return records