This commit is contained in:
parent
483b7dc7ea
commit
72dae054f1
|
|
@ -22,15 +22,14 @@ from application.states.models import (
|
|||
Run,
|
||||
Message,
|
||||
deps_to_object,
|
||||
usage_to_object,
|
||||
usage_validate_python,
|
||||
usage_limits_to_object,
|
||||
|
||||
)
|
||||
|
||||
|
||||
class VerificationCodeTable(SQLModel, table=True):
|
||||
class VerificationCodeRecord(SQLModel, table=True, table_name="verification_code"):
|
||||
"""
|
||||
验证码表
|
||||
验证码记录
|
||||
"""
|
||||
|
||||
email: str = Field(..., primary_key=True, description="邮箱")
|
||||
|
|
@ -56,9 +55,9 @@ class VerificationCodeTable(SQLModel, table=True):
|
|||
)
|
||||
|
||||
|
||||
class UserTable(SQLModel, table=True):
|
||||
class UserRecord(SQLModel, table=True, table_name="user"):
|
||||
"""
|
||||
用户表
|
||||
用户记录
|
||||
"""
|
||||
|
||||
id: str = Field(
|
||||
|
|
@ -69,9 +68,9 @@ class UserTable(SQLModel, table=True):
|
|||
email: str = Field(..., index=True, description="邮箱")
|
||||
|
||||
|
||||
class ConversationTable(SQLModel, table=True):
|
||||
class ConversationRecord(SQLModel, table=True, table_name="conversation"):
|
||||
"""
|
||||
会话表
|
||||
会话记录
|
||||
"""
|
||||
|
||||
id: str = Field(
|
||||
|
|
@ -81,6 +80,9 @@ class ConversationTable(SQLModel, table=True):
|
|||
)
|
||||
user_id: str = Field(..., index=True, description="用户唯一标识")
|
||||
description: str = Field(default="新会话", description="会话描述")
|
||||
usage: dict[str, Any] = Field(
|
||||
default_factory=dict, sa_type=JSON, description="会话使用量"
|
||||
)
|
||||
is_deleted: bool = Field(
|
||||
default=False,
|
||||
index=True,
|
||||
|
|
@ -91,53 +93,9 @@ class ConversationTable(SQLModel, table=True):
|
|||
)
|
||||
|
||||
|
||||
class TaskTable(SQLModel, table=True):
|
||||
class MessageRecord(SQLModel, table=True, table_name="message"):
|
||||
"""
|
||||
任务表
|
||||
"""
|
||||
|
||||
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(
|
||||
|
|
@ -145,15 +103,12 @@ class MessageTable(SQLModel, table=True):
|
|||
primary_key=True,
|
||||
description="消息唯一标识",
|
||||
)
|
||||
run_id: str = Field(..., index=True, description="运行唯一标识")
|
||||
type: MessageType = Field(default=MessageType.USER_PROMPT, description="消息类型")
|
||||
conversation_id: str = Field(..., index=True, description="会话唯一标识")
|
||||
type: MessageType = Field(..., description="消息类型")
|
||||
title: str = Field(default="", description="消息标题")
|
||||
content: str = Field(default="", description="消息内容")
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class DatabaseState(rx.State):
|
||||
"""
|
||||
数据库状态
|
||||
|
|
@ -166,19 +121,19 @@ class DatabaseState(rx.State):
|
|||
:return: 验证码
|
||||
"""
|
||||
async with rx.asession() as session:
|
||||
# 先将该邮箱的有效、未核验的验证码记录设置为无效
|
||||
# 先将该邮箱有效、未核验的验证码记录设置为无效
|
||||
await session.exec(
|
||||
update(VerificationCodeTable)
|
||||
update(VerificationCodeRecord)
|
||||
.where(
|
||||
VerificationCodeTable.email == email, # type: ignore
|
||||
VerificationCodeTable.is_valid == True, # type: ignore
|
||||
VerificationCodeTable.is_verified == False, # type: ignore
|
||||
VerificationCodeRecord.email == email, # type: ignore
|
||||
VerificationCodeRecord.is_valid == True, # type: ignore
|
||||
VerificationCodeRecord.is_verified == False, # type: ignore
|
||||
)
|
||||
.values(is_valid=False)
|
||||
)
|
||||
await session.flush()
|
||||
# 创建验证码记录
|
||||
record = VerificationCodeTable(
|
||||
record = VerificationCodeRecord(
|
||||
email=email,
|
||||
)
|
||||
session.add(record)
|
||||
|
|
@ -196,14 +151,13 @@ class DatabaseState(rx.State):
|
|||
: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(),
|
||||
select(VerificationCodeRecord).where(
|
||||
VerificationCodeRecord.email == email,
|
||||
VerificationCodeRecord.verification_code == verification_code,
|
||||
VerificationCodeRecord.is_valid == True,
|
||||
VerificationCodeRecord.is_verified == False,
|
||||
VerificationCodeRecord.expired_at > datetime.now(),
|
||||
)
|
||||
)
|
||||
record = result.first()
|
||||
|
|
@ -222,13 +176,13 @@ class DatabaseState(rx.State):
|
|||
"""
|
||||
async with rx.asession() as session:
|
||||
result = await session.exec(
|
||||
select(UserTable).where(UserTable.email == email)
|
||||
select(UserRecord).where(UserRecord.email == email)
|
||||
)
|
||||
record = result.first()
|
||||
# 若用户记录已存在则返回用户唯一标识,否则先创建用户记录再返回用户唯一标识
|
||||
if record:
|
||||
return record.id
|
||||
record = UserTable(email=email)
|
||||
record = UserRecord(email=email)
|
||||
session.add(record)
|
||||
await session.commit()
|
||||
await session.refresh(record)
|
||||
|
|
@ -240,56 +194,42 @@ class DatabaseState(rx.State):
|
|||
:param user_id: 用户唯一标识
|
||||
:return: 当前用户的会话字典
|
||||
"""
|
||||
records: dict[str, Conversation] = {}
|
||||
conversations: 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
|
||||
select(ConversationRecord, MessageRecord)
|
||||
.outerjoin(MessageRecord, MessageRecord.conversation_id == ConversationRecord.id) # type: ignore
|
||||
.where(
|
||||
ConversationTable.user_id == user_id,
|
||||
ConversationTable.is_deleted == False,
|
||||
)
|
||||
.order_by(
|
||||
ConversationTable.id, TaskTable.id, RunTable.id, MessageTable.id
|
||||
ConversationRecord.user_id == user_id,
|
||||
ConversationRecord.is_deleted == False,
|
||||
)
|
||||
.order_by(ConversationRecord.id, MessageRecord.id)
|
||||
)
|
||||
for (
|
||||
conversation_result,
|
||||
run_result,
|
||||
message_result,
|
||||
conversation_record,
|
||||
message_record,
|
||||
) in result.all():
|
||||
conversation_record = records.setdefault(
|
||||
conversation_result.id,
|
||||
conversation = conversations.setdefault(
|
||||
conversation_record.id,
|
||||
Conversation(
|
||||
id=conversation_result.id,
|
||||
description=conversation_result.description,
|
||||
created_at=conversation_result.created_at,
|
||||
id=conversation_record.id,
|
||||
description=conversation_record.description,
|
||||
usage=usage_validate_python(conversation_record.usage),
|
||||
created_at=conversation_record.created_at,
|
||||
),
|
||||
)
|
||||
if not run_result:
|
||||
if not message_record:
|
||||
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,
|
||||
conversation.messages.setdefault(
|
||||
message_record.id,
|
||||
Message(
|
||||
id=message_result.id,
|
||||
type=message_result.type,
|
||||
title=message_result.title,
|
||||
content=message_result.content,
|
||||
id=message_record.id,
|
||||
type=message_record.type,
|
||||
title=message_record.title,
|
||||
content=message_record.content,
|
||||
),
|
||||
)
|
||||
return records
|
||||
return conversations
|
||||
|
||||
async def create_conversation_record(self, user_id: str) -> dict[str, Conversation]:
|
||||
"""
|
||||
|
|
@ -298,7 +238,7 @@ class DatabaseState(rx.State):
|
|||
:return: 会话实例
|
||||
"""
|
||||
async with rx.asession() as session:
|
||||
record = ConversationTable(user_id=user_id)
|
||||
record = ConversationRecord(user_id=user_id)
|
||||
session.add(record)
|
||||
await session.commit()
|
||||
await session.refresh(record)
|
||||
|
|
|
|||
|
|
@ -31,8 +31,8 @@ class Message(BaseModel):
|
|||
type: MessageType = Field(..., description="消息类型")
|
||||
title: str = Field(default="", description="消息标题")
|
||||
content: str = Field(default="", description="消息内容")
|
||||
is_component_expanded: bool = Field(
|
||||
default=True, description="展开组件,True 表示展开,False 表示折叠"
|
||||
is_expanded: bool = Field(
|
||||
default=True, description="展开,True 表示展开,False 表示折叠"
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -106,8 +106,9 @@ class Conversation(BaseModel):
|
|||
"""
|
||||
|
||||
id: str = Field(default_factory=lambda: str(uuid7()), description="会话唯一标识")
|
||||
runs: dict[str, Run] = Field(default_factory=dict, description="运行字典")
|
||||
description: str = Field(default="新会话", description="会话描述")
|
||||
usage: RunUsage = Field(default=RunUsage(), description="会话使用量")
|
||||
messages: dict[str, Message] = Field(default_factory=dict, description="运行字典")
|
||||
created_at: datetime = Field(..., description="会话创建日期时间")
|
||||
|
||||
|
||||
|
|
@ -146,7 +147,7 @@ def deps_to_dict(deps: Deps | None) -> dict[str, Any]:
|
|||
UsageAdapter = TypeAdapter(RunUsage)
|
||||
|
||||
|
||||
def usage_to_object(usage: dict[str, Any]) -> RunUsage:
|
||||
def usage_validate_python(usage: dict[str, Any]) -> RunUsage:
|
||||
"""
|
||||
Usage 转为对象
|
||||
"""
|
||||
|
|
|
|||
Loading…
Reference in New Issue