Python/agent/application/states/database.py

409 lines
13 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, 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