refactor: 重构数据库访问层,用户/对话/消息都统一迁移至PostgreSQL并添加业务模型支持
This commit is contained in:
parent
a1812f2e97
commit
3da0d57c1c
@ -45,7 +45,7 @@ services:
|
||||
- NO_PROXY=localhost,127.0.0.1,milvus,graph,milvus-minio,milvus-etcd-dev,etcd,minio,mineru,paddlex,api.siliconflow.cn
|
||||
- no_proxy=localhost,127.0.0.1,milvus,graph,milvus-minio,milvus-etcd-dev,etcd,minio,mineru,paddlex,api.siliconflow.cn
|
||||
# endregion api_envs
|
||||
command: uv run --no-dev uvicorn server.main:app --host 0.0.0.0 --port 5050 --reload
|
||||
command: uv run --no-dev uvicorn server.main:app --host 0.0.0.0 --port 5050 # --reload
|
||||
restart: unless-stopped
|
||||
healthcheck:
|
||||
test: ["CMD-SHELL", "curl -f http://localhost:5050/api/system/health || exit 1"]
|
||||
|
||||
601
scripts/migrate_business_from_sqlite.py
Normal file
601
scripts/migrate_business_from_sqlite.py
Normal file
@ -0,0 +1,601 @@
|
||||
"""
|
||||
SQLite 到 PostgreSQL 业务数据迁移脚本
|
||||
|
||||
将用户、部门、对话等业务数据从 SQLite 迁移到 PostgreSQL。
|
||||
迁移顺序(按外键依赖):
|
||||
1. departments (无依赖)
|
||||
2. users (依赖 departments)
|
||||
3. conversations (依赖 users)
|
||||
4. messages (依赖 conversations)
|
||||
5. tool_calls (依赖 messages)
|
||||
6. conversation_stats (依赖 conversations)
|
||||
7. operation_logs (依赖 users)
|
||||
8. message_feedbacks (依赖 messages)
|
||||
9. mcp_servers (无依赖)
|
||||
|
||||
用法:
|
||||
python scripts/migrate_business_from_sqlite.py --dry-run # 预览迁移
|
||||
python scripts/migrate_business_from_sqlite.py --execute # 执行迁移
|
||||
python scripts/migrate_business_from_sqlite.py --verify # 验证数据
|
||||
python scripts/migrate_business_from_sqlite.py --rollback # 回滚迁移
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, UTC
|
||||
from typing import Any
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))
|
||||
os.environ.setdefault("YUXI_SKIP_APP_INIT", "1")
|
||||
|
||||
from sqlalchemy import create_engine, select, text
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from src import config
|
||||
from src.storage.db.models import (
|
||||
Base as SqliteBase,
|
||||
Department as SqliteDepartment,
|
||||
User as SqliteUser,
|
||||
Conversation as SqliteConversation,
|
||||
Message as SqliteMessage,
|
||||
ToolCall as SqliteToolCall,
|
||||
ConversationStats as SqliteConversationStats,
|
||||
OperationLog as SqliteOperationLog,
|
||||
MessageFeedback as SqliteMessageFeedback,
|
||||
MCPServer as SqliteMCPServer,
|
||||
)
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import (
|
||||
Department,
|
||||
User,
|
||||
Conversation,
|
||||
Message,
|
||||
ToolCall,
|
||||
ConversationStats,
|
||||
OperationLog,
|
||||
MessageFeedback,
|
||||
MCPServer,
|
||||
)
|
||||
from src.utils import logger
|
||||
|
||||
|
||||
def _utc_dt(value: Any) -> datetime | None:
|
||||
"""Convert various datetime formats to naive UTC datetime."""
|
||||
if not value:
|
||||
return None
|
||||
if isinstance(value, datetime):
|
||||
if value.tzinfo is None:
|
||||
return value
|
||||
return value.astimezone(UTC).replace(tzinfo=None)
|
||||
if isinstance(value, (int, float)):
|
||||
return datetime.fromtimestamp(value, tz=UTC).replace(tzinfo=None)
|
||||
if isinstance(value, str):
|
||||
v = value.strip()
|
||||
if not v:
|
||||
return None
|
||||
try:
|
||||
dt_val = datetime.fromisoformat(v.replace("Z", "+00:00"))
|
||||
if dt_val.tzinfo is None:
|
||||
return dt_val
|
||||
return dt_val.astimezone(UTC).replace(tzinfo=None)
|
||||
except ValueError:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
class SQLiteReader:
|
||||
"""SQLite 数据读取器"""
|
||||
|
||||
def __init__(self):
|
||||
db_path = os.path.join(config.save_dir, "database", "server.db")
|
||||
self.engine = create_engine(f"sqlite:///{db_path}")
|
||||
self.Session = sessionmaker(bind=self.engine)
|
||||
|
||||
def get_session(self):
|
||||
return self.Session()
|
||||
|
||||
def read_departments(self) -> list[SqliteDepartment]:
|
||||
with self.get_session() as session:
|
||||
return session.execute(select(SqliteDepartment)).scalars().all()
|
||||
|
||||
def read_users(self) -> list[SqliteUser]:
|
||||
with self.get_session() as session:
|
||||
return session.execute(select(SqliteUser)).scalars().all()
|
||||
|
||||
def read_conversations(self) -> list[SqliteConversation]:
|
||||
with self.get_session() as session:
|
||||
return session.execute(select(SqliteConversation)).scalars().all()
|
||||
|
||||
def read_messages(self) -> list[SqliteMessage]:
|
||||
with self.get_session() as session:
|
||||
return session.execute(select(SqliteMessage)).scalars().all()
|
||||
|
||||
def read_tool_calls(self) -> list[SqliteToolCall]:
|
||||
with self.get_session() as session:
|
||||
return session.execute(select(SqliteToolCall)).scalars().all()
|
||||
|
||||
def read_conversation_stats(self) -> list[SqliteConversationStats]:
|
||||
with self.get_session() as session:
|
||||
return session.execute(select(SqliteConversationStats)).scalars().all()
|
||||
|
||||
def read_operation_logs(self) -> list[SqliteOperationLog]:
|
||||
with self.get_session() as session:
|
||||
return session.execute(select(SqliteOperationLog)).scalars().all()
|
||||
|
||||
def read_message_feedbacks(self) -> list[SqliteMessageFeedback]:
|
||||
with self.get_session() as session:
|
||||
return session.execute(select(SqliteMessageFeedback)).scalars().all()
|
||||
|
||||
def read_mcp_servers(self) -> list[SqliteMCPServer]:
|
||||
with self.get_session() as session:
|
||||
return session.execute(select(SqliteMCPServer)).scalars().all()
|
||||
|
||||
def count_table(self, table_name: str) -> int:
|
||||
with self.get_session() as session:
|
||||
result = session.execute(text(f"SELECT COUNT(*) FROM {table_name}"))
|
||||
return result.scalar() or 0
|
||||
|
||||
|
||||
async def migrate_departments(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
|
||||
"""迁移部门数据"""
|
||||
sqlite_depts = sqlite_reader.read_departments()
|
||||
logger.info(f"准备迁移 {len(sqlite_depts)} 个部门")
|
||||
|
||||
created = 0
|
||||
if dry_run:
|
||||
for sqlite_dept in sqlite_depts:
|
||||
logger.info(f"[DRY-RUN] 将创建部门: {sqlite_dept.name}")
|
||||
elif execute:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
for sqlite_dept in sqlite_depts:
|
||||
# 检查是否已存在
|
||||
existing = await session.execute(
|
||||
select(Department).where(Department.id == sqlite_dept.id)
|
||||
)
|
||||
if existing.scalar_one_or_none() is None:
|
||||
dept = Department(
|
||||
id=sqlite_dept.id,
|
||||
name=sqlite_dept.name,
|
||||
description=sqlite_dept.description,
|
||||
created_at=_utc_dt(sqlite_dept.created_at),
|
||||
)
|
||||
session.add(dept)
|
||||
created += 1
|
||||
|
||||
return {"total": len(sqlite_depts), "created": created}
|
||||
|
||||
|
||||
async def migrate_users(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
|
||||
"""迁移用户数据"""
|
||||
sqlite_users = sqlite_reader.read_users()
|
||||
logger.info(f"准备迁移 {len(sqlite_users)} 个用户")
|
||||
|
||||
created = 0
|
||||
if dry_run:
|
||||
for sqlite_user in sqlite_users:
|
||||
logger.info(f"[DRY-RUN] 将创建用户: {sqlite_user.username} ({sqlite_user.user_id})")
|
||||
elif execute:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
for sqlite_user in sqlite_users:
|
||||
existing = await session.execute(
|
||||
select(User).where(User.id == sqlite_user.id)
|
||||
)
|
||||
if existing.scalar_one_or_none() is None:
|
||||
user = User(
|
||||
id=sqlite_user.id,
|
||||
username=sqlite_user.username,
|
||||
user_id=sqlite_user.user_id,
|
||||
phone_number=sqlite_user.phone_number,
|
||||
avatar=sqlite_user.avatar,
|
||||
password_hash=sqlite_user.password_hash,
|
||||
role=sqlite_user.role,
|
||||
department_id=sqlite_user.department_id,
|
||||
created_at=_utc_dt(sqlite_user.created_at),
|
||||
last_login=_utc_dt(sqlite_user.last_login),
|
||||
login_failed_count=sqlite_user.login_failed_count,
|
||||
last_failed_login=_utc_dt(sqlite_user.last_failed_login),
|
||||
login_locked_until=_utc_dt(sqlite_user.login_locked_until),
|
||||
is_deleted=sqlite_user.is_deleted,
|
||||
deleted_at=_utc_dt(sqlite_user.deleted_at),
|
||||
)
|
||||
session.add(user)
|
||||
created += 1
|
||||
|
||||
return {"total": len(sqlite_users), "created": created}
|
||||
|
||||
|
||||
async def migrate_conversations(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
|
||||
"""迁移对话数据"""
|
||||
sqlite_convs = sqlite_reader.read_conversations()
|
||||
logger.info(f"准备迁移 {len(sqlite_convs)} 个对话")
|
||||
|
||||
created = 0
|
||||
if dry_run:
|
||||
for sqlite_conv in sqlite_convs:
|
||||
logger.info(f"[DRY-RUN] 将创建对话: {sqlite_conv.thread_id}")
|
||||
elif execute:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
for sqlite_conv in sqlite_convs:
|
||||
existing = await session.execute(
|
||||
select(Conversation).where(Conversation.id == sqlite_conv.id)
|
||||
)
|
||||
if existing.scalar_one_or_none() is None:
|
||||
# 截断过长的 title
|
||||
title = sqlite_conv.title
|
||||
if title and len(title) > 255:
|
||||
title = title[:255]
|
||||
logger.warning(f"截断对话标题 (id={sqlite_conv.id}): 原始长度={len(sqlite_conv.title)}")
|
||||
conv = Conversation(
|
||||
id=sqlite_conv.id,
|
||||
thread_id=sqlite_conv.thread_id,
|
||||
user_id=sqlite_conv.user_id,
|
||||
agent_id=sqlite_conv.agent_id,
|
||||
title=title,
|
||||
status=sqlite_conv.status,
|
||||
created_at=_utc_dt(sqlite_conv.created_at),
|
||||
updated_at=_utc_dt(sqlite_conv.updated_at),
|
||||
extra_metadata=sqlite_conv.extra_metadata,
|
||||
)
|
||||
session.add(conv)
|
||||
created += 1
|
||||
|
||||
return {"total": len(sqlite_convs), "created": created}
|
||||
|
||||
|
||||
async def migrate_messages(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
|
||||
"""迁移消息数据"""
|
||||
sqlite_messages = sqlite_reader.read_messages()
|
||||
logger.info(f"准备迁移 {len(sqlite_messages)} 条消息")
|
||||
|
||||
created = 0
|
||||
if dry_run:
|
||||
for sqlite_msg in sqlite_messages:
|
||||
logger.info(f"[DRY-RUN] 将创建消息: id={sqlite_msg.id}, conversation={sqlite_msg.conversation_id}")
|
||||
elif execute:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
for sqlite_msg in sqlite_messages:
|
||||
existing = await session.execute(
|
||||
select(Message).where(Message.id == sqlite_msg.id)
|
||||
)
|
||||
if existing.scalar_one_or_none() is None:
|
||||
msg = Message(
|
||||
id=sqlite_msg.id,
|
||||
conversation_id=sqlite_msg.conversation_id,
|
||||
role=sqlite_msg.role,
|
||||
content=sqlite_msg.content,
|
||||
message_type=sqlite_msg.message_type,
|
||||
created_at=_utc_dt(sqlite_msg.created_at),
|
||||
token_count=sqlite_msg.token_count,
|
||||
extra_metadata=sqlite_msg.extra_metadata,
|
||||
image_content=sqlite_msg.image_content,
|
||||
)
|
||||
session.add(msg)
|
||||
created += 1
|
||||
|
||||
return {"total": len(sqlite_messages), "created": created}
|
||||
|
||||
|
||||
async def migrate_tool_calls(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
|
||||
"""迁移工具调用数据"""
|
||||
sqlite_calls = sqlite_reader.read_tool_calls()
|
||||
logger.info(f"准备迁移 {len(sqlite_calls)} 个工具调用")
|
||||
|
||||
created = 0
|
||||
if dry_run:
|
||||
for sqlite_call in sqlite_calls:
|
||||
logger.info(f"[DRY-RUN] 将创建工具调用: id={sqlite_call.id}, tool={sqlite_call.tool_name}")
|
||||
elif execute:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
for sqlite_call in sqlite_calls:
|
||||
existing = await session.execute(
|
||||
select(ToolCall).where(ToolCall.id == sqlite_call.id)
|
||||
)
|
||||
if existing.scalar_one_or_none() is None:
|
||||
call = ToolCall(
|
||||
id=sqlite_call.id,
|
||||
message_id=sqlite_call.message_id,
|
||||
langgraph_tool_call_id=sqlite_call.langgraph_tool_call_id,
|
||||
tool_name=sqlite_call.tool_name,
|
||||
tool_input=sqlite_call.tool_input,
|
||||
tool_output=sqlite_call.tool_output,
|
||||
status=sqlite_call.status,
|
||||
error_message=sqlite_call.error_message,
|
||||
created_at=_utc_dt(sqlite_call.created_at),
|
||||
)
|
||||
session.add(call)
|
||||
created += 1
|
||||
|
||||
return {"total": len(sqlite_calls), "created": created}
|
||||
|
||||
|
||||
async def migrate_conversation_stats(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
|
||||
"""迁移对话统计数据"""
|
||||
sqlite_stats = sqlite_reader.read_conversation_stats()
|
||||
logger.info(f"准备迁移 {len(sqlite_stats)} 条对话统计")
|
||||
|
||||
created = 0
|
||||
if dry_run:
|
||||
for sqlite_stat in sqlite_stats:
|
||||
logger.info(f"[DRY-RUN] 将创建对话统计: conversation_id={sqlite_stat.conversation_id}")
|
||||
elif execute:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
for sqlite_stat in sqlite_stats:
|
||||
existing = await session.execute(
|
||||
select(ConversationStats).where(ConversationStats.id == sqlite_stat.id)
|
||||
)
|
||||
if existing.scalar_one_or_none() is None:
|
||||
stat = ConversationStats(
|
||||
id=sqlite_stat.id,
|
||||
conversation_id=sqlite_stat.conversation_id,
|
||||
message_count=sqlite_stat.message_count,
|
||||
total_tokens=sqlite_stat.total_tokens,
|
||||
model_used=sqlite_stat.model_used,
|
||||
user_feedback=sqlite_stat.user_feedback,
|
||||
created_at=_utc_dt(sqlite_stat.created_at),
|
||||
updated_at=_utc_dt(sqlite_stat.updated_at),
|
||||
)
|
||||
session.add(stat)
|
||||
created += 1
|
||||
|
||||
return {"total": len(sqlite_stats), "created": created}
|
||||
|
||||
|
||||
async def migrate_operation_logs(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
|
||||
"""迁移操作日志数据"""
|
||||
sqlite_logs = sqlite_reader.read_operation_logs()
|
||||
logger.info(f"准备迁移 {len(sqlite_logs)} 条操作日志")
|
||||
|
||||
created = 0
|
||||
if dry_run:
|
||||
for sqlite_log in sqlite_logs:
|
||||
logger.info(f"[DRY-RUN] 将创建操作日志: id={sqlite_log.id}, operation={sqlite_log.operation}")
|
||||
elif execute:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
for sqlite_log in sqlite_logs:
|
||||
existing = await session.execute(
|
||||
select(OperationLog).where(OperationLog.id == sqlite_log.id)
|
||||
)
|
||||
if existing.scalar_one_or_none() is None:
|
||||
log = OperationLog(
|
||||
id=sqlite_log.id,
|
||||
user_id=sqlite_log.user_id,
|
||||
operation=sqlite_log.operation,
|
||||
details=sqlite_log.details,
|
||||
ip_address=sqlite_log.ip_address,
|
||||
timestamp=_utc_dt(sqlite_log.timestamp),
|
||||
)
|
||||
session.add(log)
|
||||
created += 1
|
||||
|
||||
return {"total": len(sqlite_logs), "created": created}
|
||||
|
||||
|
||||
async def migrate_message_feedbacks(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
|
||||
"""迁移消息反馈数据"""
|
||||
sqlite_feedbacks = sqlite_reader.read_message_feedbacks()
|
||||
logger.info(f"准备迁移 {len(sqlite_feedbacks)} 条消息反馈")
|
||||
|
||||
created = 0
|
||||
if dry_run:
|
||||
for sqlite_fb in sqlite_feedbacks:
|
||||
logger.info(f"[DRY-RUN] 将创建消息反馈: id={sqlite_fb.id}, rating={sqlite_fb.rating}")
|
||||
elif execute:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
for sqlite_fb in sqlite_feedbacks:
|
||||
existing = await session.execute(
|
||||
select(MessageFeedback).where(MessageFeedback.id == sqlite_fb.id)
|
||||
)
|
||||
if existing.scalar_one_or_none() is None:
|
||||
fb = MessageFeedback(
|
||||
id=sqlite_fb.id,
|
||||
message_id=sqlite_fb.message_id,
|
||||
user_id=sqlite_fb.user_id,
|
||||
rating=sqlite_fb.rating,
|
||||
reason=sqlite_fb.reason,
|
||||
created_at=_utc_dt(sqlite_fb.created_at),
|
||||
)
|
||||
session.add(fb)
|
||||
created += 1
|
||||
|
||||
return {"total": len(sqlite_feedbacks), "created": created}
|
||||
|
||||
|
||||
async def migrate_mcp_servers(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
|
||||
"""迁移 MCP 服务器数据"""
|
||||
sqlite_servers = sqlite_reader.read_mcp_servers()
|
||||
logger.info(f"准备迁移 {len(sqlite_servers)} 个 MCP 服务器")
|
||||
|
||||
created = 0
|
||||
if dry_run:
|
||||
for sqlite_server in sqlite_servers:
|
||||
logger.info(f"[DRY-RUN] 将创建 MCP 服务器: {sqlite_server.name}")
|
||||
elif execute:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
for sqlite_server in sqlite_servers:
|
||||
existing = await session.execute(
|
||||
select(MCPServer).where(MCPServer.name == sqlite_server.name)
|
||||
)
|
||||
if existing.scalar_one_or_none() is None:
|
||||
server = MCPServer(
|
||||
name=sqlite_server.name,
|
||||
description=sqlite_server.description,
|
||||
transport=sqlite_server.transport,
|
||||
url=sqlite_server.url,
|
||||
command=sqlite_server.command,
|
||||
args=sqlite_server.args,
|
||||
headers=sqlite_server.headers,
|
||||
timeout=sqlite_server.timeout,
|
||||
sse_read_timeout=sqlite_server.sse_read_timeout,
|
||||
tags=sqlite_server.tags,
|
||||
icon=sqlite_server.icon,
|
||||
enabled=sqlite_server.enabled,
|
||||
disabled_tools=sqlite_server.disabled_tools,
|
||||
created_by=sqlite_server.created_by,
|
||||
updated_by=sqlite_server.updated_by,
|
||||
created_at=_utc_dt(sqlite_server.created_at),
|
||||
updated_at=_utc_dt(sqlite_server.updated_at),
|
||||
)
|
||||
session.add(server)
|
||||
created += 1
|
||||
|
||||
return {"total": len(sqlite_servers), "created": created}
|
||||
|
||||
|
||||
async def verify_migration(sqlite_reader: SQLiteReader) -> dict[str, dict]:
|
||||
"""验证迁移结果"""
|
||||
# 使用 (模型, 主键列名) 格式,支持不同表使用不同的主键
|
||||
tables = [
|
||||
("departments", Department, "id"),
|
||||
("users", User, "id"),
|
||||
("conversations", Conversation, "id"),
|
||||
("messages", Message, "id"),
|
||||
("tool_calls", ToolCall, "id"),
|
||||
("conversation_stats", ConversationStats, "id"),
|
||||
("operation_logs", OperationLog, "id"),
|
||||
("message_feedbacks", MessageFeedback, "id"),
|
||||
("mcp_servers", MCPServer, "name"), # MCPServer 使用 name 作为主键
|
||||
]
|
||||
|
||||
results = {}
|
||||
for table_name, model, pk_column in tables:
|
||||
sqlite_count = sqlite_reader.count_table(table_name)
|
||||
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
from sqlalchemy import func
|
||||
pk_attr = getattr(model, pk_column)
|
||||
result = await session.execute(select(func.count(pk_attr)))
|
||||
pg_count = result.scalar() or 0
|
||||
|
||||
results[table_name] = {
|
||||
"sqlite": sqlite_count,
|
||||
"postgresql": pg_count,
|
||||
"match": sqlite_count == pg_count,
|
||||
}
|
||||
|
||||
return results
|
||||
|
||||
|
||||
async def rollback_migration() -> None:
|
||||
"""回滚迁移 - 删除所有业务数据表"""
|
||||
logger.warning("开始回滚迁移...")
|
||||
|
||||
# 按外键依赖顺序删除
|
||||
tables_to_delete = [
|
||||
MessageFeedback,
|
||||
OperationLog,
|
||||
ConversationStats,
|
||||
ToolCall,
|
||||
Message,
|
||||
Conversation,
|
||||
User,
|
||||
Department,
|
||||
MCPServer,
|
||||
]
|
||||
|
||||
for model in tables_to_delete:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(model))
|
||||
records = result.scalars().all()
|
||||
for record in records:
|
||||
await session.delete(record)
|
||||
|
||||
logger.warning("回滚完成 - 已删除所有迁移的业务数据")
|
||||
|
||||
|
||||
async def migrate_all(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, Any]:
|
||||
"""执行所有迁移"""
|
||||
results = {}
|
||||
|
||||
# 按外键依赖顺序迁移
|
||||
results["departments"] = await migrate_departments(sqlite_reader, dry_run, execute)
|
||||
results["users"] = await migrate_users(sqlite_reader, dry_run, execute)
|
||||
results["conversations"] = await migrate_conversations(sqlite_reader, dry_run, execute)
|
||||
results["messages"] = await migrate_messages(sqlite_reader, dry_run, execute)
|
||||
results["tool_calls"] = await migrate_tool_calls(sqlite_reader, dry_run, execute)
|
||||
results["conversation_stats"] = await migrate_conversation_stats(sqlite_reader, dry_run, execute)
|
||||
results["operation_logs"] = await migrate_operation_logs(sqlite_reader, dry_run, execute)
|
||||
results["message_feedbacks"] = await migrate_message_feedbacks(sqlite_reader, dry_run, execute)
|
||||
results["mcp_servers"] = await migrate_mcp_servers(sqlite_reader, dry_run, execute)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="SQLite 到 PostgreSQL 业务数据迁移")
|
||||
parser.add_argument("--dry-run", action="store_true", help="预览迁移,不执行")
|
||||
parser.add_argument("--execute", action="store_true", help="执行迁移")
|
||||
parser.add_argument("--verify", action="store_true", help="验证迁移结果")
|
||||
parser.add_argument("--rollback", action="store_true", help="回滚迁移")
|
||||
parser.add_argument("--migrate-all", action="store_true", help="迁移所有业务数据")
|
||||
parser.add_argument("--init-tables", action="store_true", help="仅初始化业务表结构")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if not any([args.dry_run, args.execute, args.verify, args.rollback, args.migrate_all, args.init_tables]):
|
||||
args.dry_run = True
|
||||
|
||||
# 初始化 PostgreSQL 管理器
|
||||
pg_manager.initialize()
|
||||
logger.info("PostgreSQL manager initialized")
|
||||
|
||||
if args.init_tables:
|
||||
# 仅初始化表结构
|
||||
await pg_manager.create_business_tables()
|
||||
logger.info("业务表结构初始化完成")
|
||||
return
|
||||
|
||||
if args.verify:
|
||||
# 验证模式
|
||||
sqlite_reader = SQLiteReader()
|
||||
results = await verify_migration(sqlite_reader)
|
||||
|
||||
logger.info("=" * 60)
|
||||
logger.info("迁移验证结果:")
|
||||
logger.info("=" * 60)
|
||||
all_match = True
|
||||
for table_name, counts in results.items():
|
||||
status = "✓" if counts["match"] else "✗"
|
||||
logger.info(
|
||||
f"{status} {table_name}: SQLite={counts['sqlite']}, PostgreSQL={counts['postgresql']}"
|
||||
)
|
||||
if not counts["match"]:
|
||||
all_match = False
|
||||
logger.info("=" * 60)
|
||||
logger.info(f"全部匹配: {'是' if all_match else '否'}")
|
||||
return
|
||||
|
||||
if args.rollback:
|
||||
# 回滚模式
|
||||
if args.dry_run:
|
||||
logger.info("[DRY-RUN] 将回滚所有迁移的业务数据")
|
||||
else:
|
||||
await rollback_migration()
|
||||
return
|
||||
|
||||
# 迁移模式
|
||||
sqlite_reader = SQLiteReader()
|
||||
|
||||
if args.migrate_all:
|
||||
# 检查是否需要初始化表结构
|
||||
logger.info("检查业务表结构...")
|
||||
await pg_manager.create_business_tables()
|
||||
logger.info("业务表结构就绪")
|
||||
|
||||
results = await migrate_all(sqlite_reader, args.dry_run, args.execute)
|
||||
|
||||
logger.info("=" * 60)
|
||||
logger.info("迁移完成:")
|
||||
for table_name, counts in results.items():
|
||||
logger.info(f" {table_name}: {counts['created']}/{counts['total']}")
|
||||
logger.info("=" * 60)
|
||||
|
||||
if not args.dry_run:
|
||||
logger.info("建议运行 --verify 验证数据完整性")
|
||||
else:
|
||||
logger.info("使用 --migrate-all 执行迁移,或使用 --verify 验证数据")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
@ -74,6 +74,8 @@ async def rollback_all() -> None:
|
||||
|
||||
|
||||
async def migrate(dry_run: bool, execute: bool, rollback: bool) -> None:
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
|
||||
base_dir = os.path.join(config.save_dir, "knowledge_base_data")
|
||||
global_meta_path = os.path.join(base_dir, "global_metadata.json")
|
||||
global_meta = _load_json(global_meta_path).get("databases", {})
|
||||
@ -86,6 +88,11 @@ async def migrate(dry_run: bool, execute: bool, rollback: bool) -> None:
|
||||
logger.info("Rollback completed")
|
||||
return
|
||||
|
||||
# 初始化表结构
|
||||
pg_manager.initialize()
|
||||
await pg_manager.create_tables()
|
||||
logger.info("知识库表结构初始化完成")
|
||||
|
||||
kb_repo = KnowledgeBaseRepository()
|
||||
file_repo = KnowledgeFileRepository()
|
||||
eval_repo = EvaluationRepository()
|
||||
|
||||
@ -8,8 +8,10 @@ from pydantic import BaseModel
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.storage.db.manager import db_manager
|
||||
from src.storage.db.models import User, Department
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import User, Department
|
||||
from src.repositories.user_repository import UserRepository
|
||||
from src.repositories.department_repository import DepartmentRepository
|
||||
from server.utils.auth_middleware import (
|
||||
get_admin_user,
|
||||
get_superadmin_user,
|
||||
@ -21,7 +23,10 @@ from server.utils.auth_utils import AuthUtils
|
||||
from server.utils.user_utils import generate_unique_user_id, validate_username, is_valid_phone_number
|
||||
from server.utils.common_utils import log_operation
|
||||
from src.storage.minio import aupload_file_to_minio
|
||||
from src.utils.datetime_utils import utc_now
|
||||
from datetime import datetime as dt, timezone
|
||||
|
||||
# 使用 naive datetime 以兼容 PostgreSQL TIMESTAMP WITHOUT TIME ZONE 列
|
||||
_utc_now = lambda: dt.now(timezone.utc).replace(tzinfo=None)
|
||||
|
||||
# 创建路由器
|
||||
auth = APIRouter(prefix="/auth", tags=["authentication"])
|
||||
@ -175,7 +180,7 @@ async def login_for_access_token(form_data: OAuth2PasswordRequestForm = Depends(
|
||||
|
||||
# 登录成功,重置失败计数器
|
||||
user.reset_failed_login()
|
||||
user.last_login = utc_now()
|
||||
user.last_login = _utc_now()
|
||||
await db.commit()
|
||||
|
||||
# 生成访问令牌
|
||||
@ -208,7 +213,7 @@ async def login_for_access_token(form_data: OAuth2PasswordRequestForm = Depends(
|
||||
# 路由:校验是否需要初始化管理员
|
||||
@auth.get("/check-first-run")
|
||||
async def check_first_run():
|
||||
is_first_run = await db_manager.async_check_first_run()
|
||||
is_first_run = await pg_manager.async_check_first_run()
|
||||
return {"first_run": is_first_run}
|
||||
|
||||
|
||||
@ -216,7 +221,7 @@ async def check_first_run():
|
||||
@auth.post("/initialize", response_model=Token)
|
||||
async def initialize_admin(admin_data: InitializeAdmin, db: AsyncSession = Depends(get_db)):
|
||||
# 检查是否是首次运行
|
||||
if not await db_manager.async_check_first_run():
|
||||
if not await pg_manager.async_check_first_run():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="系统已经初始化,无法再次创建初始管理员",
|
||||
@ -246,24 +251,24 @@ async def initialize_admin(admin_data: InitializeAdmin, db: AsyncSession = Depen
|
||||
user_id = admin_data.user_id
|
||||
|
||||
# 创建默认部门
|
||||
default_department = Department(name="默认部门", description="系统初始化时创建的默认部门")
|
||||
db.add(default_department)
|
||||
await db.flush() # 获取部门ID
|
||||
dept_repo = DepartmentRepository()
|
||||
default_department = await dept_repo.create({
|
||||
"name": "默认部门",
|
||||
"description": "系统初始化时创建的默认部门",
|
||||
})
|
||||
|
||||
new_admin = User(
|
||||
username=admin_data.user_id, # username和user_id设置为相同值
|
||||
user_id=user_id,
|
||||
phone_number=admin_data.phone_number,
|
||||
avatar=None, # 初始化时头像为空
|
||||
password_hash=hashed_password,
|
||||
role="superadmin",
|
||||
department_id=default_department.id,
|
||||
last_login=utc_now(),
|
||||
)
|
||||
|
||||
db.add(new_admin)
|
||||
await db.commit()
|
||||
await db.refresh(new_admin)
|
||||
# 创建管理员用户
|
||||
user_repo = UserRepository()
|
||||
new_admin = await user_repo.create({
|
||||
"username": admin_data.user_id,
|
||||
"user_id": user_id,
|
||||
"phone_number": admin_data.phone_number,
|
||||
"avatar": None,
|
||||
"password_hash": hashed_password,
|
||||
"role": "superadmin",
|
||||
"department_id": default_department.id,
|
||||
"last_login": _utc_now(),
|
||||
})
|
||||
|
||||
# 生成访问令牌
|
||||
token_data = {"sub": str(new_admin.id)}
|
||||
@ -378,6 +383,7 @@ async def create_user(
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""创建新用户(管理员权限)"""
|
||||
user_repo = UserRepository()
|
||||
|
||||
# 验证用户名
|
||||
is_valid, error_msg = validate_username(user_data.username)
|
||||
@ -388,9 +394,8 @@ async def create_user(
|
||||
)
|
||||
|
||||
# 检查用户名是否已存在
|
||||
result = await db.execute(select(User).filter(User.username == user_data.username))
|
||||
existing_user = result.scalar_one_or_none()
|
||||
if existing_user:
|
||||
users = await user_repo.list_users()
|
||||
if any(u.username == user_data.username for u in users):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="用户名已存在",
|
||||
@ -398,17 +403,14 @@ async def create_user(
|
||||
|
||||
# 检查手机号是否已存在(如果提供了)
|
||||
if user_data.phone_number:
|
||||
result = await db.execute(select(User).filter(User.phone_number == user_data.phone_number))
|
||||
existing_phone = result.scalar_one_or_none()
|
||||
if existing_phone:
|
||||
if await user_repo.exists_by_phone(user_data.phone_number):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="手机号已存在",
|
||||
)
|
||||
|
||||
# 生成唯一的user_id
|
||||
result = await db.execute(select(User.user_id))
|
||||
existing_user_ids = [user_id for (user_id,) in result.all()]
|
||||
existing_user_ids = await user_repo.get_all_user_ids()
|
||||
user_id = generate_unique_user_id(user_data.username, existing_user_ids)
|
||||
|
||||
# 创建新用户
|
||||
@ -434,7 +436,11 @@ async def create_user(
|
||||
# 超级管理员创建用户时,使用指定的部门或默认部门
|
||||
department_id = user_data.department_id
|
||||
if department_id is None:
|
||||
department_id = await get_default_department_id(db)
|
||||
# 获取默认部门
|
||||
dept_repo = DepartmentRepository()
|
||||
departments = await dept_repo.list_departments()
|
||||
default_dept = next((d for d in departments if d.name == "默认部门"), None)
|
||||
department_id = default_dept.id if default_dept else None
|
||||
else:
|
||||
# 普通管理员创建用户时,自动继承该管理员的部门
|
||||
department_id = current_user.department_id
|
||||
@ -445,18 +451,14 @@ async def create_user(
|
||||
detail="普通管理员不能指定部门",
|
||||
)
|
||||
|
||||
new_user = User(
|
||||
username=user_data.username,
|
||||
user_id=user_id,
|
||||
phone_number=user_data.phone_number,
|
||||
password_hash=hashed_password,
|
||||
role=user_data.role,
|
||||
department_id=department_id,
|
||||
)
|
||||
|
||||
db.add(new_user)
|
||||
await db.commit()
|
||||
await db.refresh(new_user)
|
||||
new_user = await user_repo.create({
|
||||
"username": user_data.username,
|
||||
"user_id": user_id,
|
||||
"phone_number": user_data.phone_number,
|
||||
"password_hash": hashed_password,
|
||||
"role": user_data.role,
|
||||
"department_id": department_id,
|
||||
})
|
||||
|
||||
# 记录操作
|
||||
await log_operation(
|
||||
@ -471,28 +473,20 @@ async def create_user(
|
||||
async def read_users(
|
||||
skip: int = 0, limit: int = 100, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db)
|
||||
):
|
||||
user_repo = UserRepository()
|
||||
|
||||
# 部门隔离逻辑
|
||||
if current_user.role == "superadmin":
|
||||
# 超级管理员可以看到所有用户,使用 JOIN 获取部门名称
|
||||
result = await db.execute(
|
||||
select(User, Department.name.label("department_name"))
|
||||
.outerjoin(Department, User.department_id == Department.id)
|
||||
.filter(User.is_deleted == 0)
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
)
|
||||
# 超级管理员可以看到所有用户
|
||||
users_with_dept = await user_repo.list_with_department(skip=skip, limit=limit)
|
||||
else:
|
||||
# 普通管理员只能看到本部门用户
|
||||
result = await db.execute(
|
||||
select(User, Department.name.label("department_name"))
|
||||
.outerjoin(Department, User.department_id == Department.id)
|
||||
.filter(User.is_deleted == 0, User.department_id == current_user.department_id)
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
users_with_dept = await user_repo.list_with_department(
|
||||
skip=skip, limit=limit, department_id=current_user.department_id
|
||||
)
|
||||
rows = result.all()
|
||||
|
||||
users = []
|
||||
for user, dept_name in rows:
|
||||
for user, dept_name in users_with_dept:
|
||||
user_dict = user.to_dict()
|
||||
user_dict["department_name"] = dept_name
|
||||
users.append(user_dict)
|
||||
@ -679,7 +673,7 @@ async def delete_user(
|
||||
hash_suffix = hashlib.md5(user.user_id.encode()).hexdigest()[:4]
|
||||
|
||||
user.is_deleted = 1
|
||||
user.deleted_at = utc_now()
|
||||
user.deleted_at = _utc_now()
|
||||
user.username = f"已注销用户-{hash_suffix}"
|
||||
user.phone_number = None # 清空手机号,释放该手机号供其他用户使用
|
||||
user.password_hash = "DELETED" # 禁止登录
|
||||
|
||||
@ -8,16 +8,18 @@ Provides centralized dashboard APIs for monitoring system-wide statistics.
|
||||
|
||||
import traceback
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy import String, cast, distinct, func, or_, select
|
||||
from sqlalchemy import Integer, String, cast, distinct, func, or_, select, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from server.routers.auth_router import get_admin_user
|
||||
from server.utils.auth_middleware import get_db
|
||||
from src.storage.conversation import ConversationManager
|
||||
from src.storage.db.models import User
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.utils.datetime_utils import UTC, ensure_shanghai, shanghai_now, utc_now
|
||||
from src.utils.logging_config import logger
|
||||
|
||||
@ -25,6 +27,25 @@ from src.utils.logging_config import logger
|
||||
dashboard = APIRouter(prefix="/dashboard", tags=["Dashboard"])
|
||||
|
||||
|
||||
def _get_time_group_format(column, time_range: str) -> Any:
|
||||
"""
|
||||
根据数据库类型生成时间分组格式化表达式。
|
||||
PostgreSQL 使用 to_char + INTERVAL,SQLite 使用 datetime + strftime。
|
||||
"""
|
||||
# 检查是否是 PostgreSQL(通过检测 engine 或使用方言)
|
||||
# 这里直接使用 PostgreSQL 语法,因为所有业务数据现在都在 PostgreSQL 上
|
||||
if time_range == "14hours":
|
||||
# 每小时: YYYY-MM-DD HH:00
|
||||
time_expr = func.to_char(column + text("INTERVAL '8 hours'"), "YYYY-MM-DD HH24:00")
|
||||
elif time_range == "14weeks":
|
||||
# 每周: YYYY-WW
|
||||
time_expr = func.to_char(column + text("INTERVAL '8 hours'"), "YYYY-IW")
|
||||
else: # 14days
|
||||
# 每天: YYYY-MM-DD
|
||||
time_expr = func.to_char(column + text("INTERVAL '8 hours'"), "YYYY-MM-DD")
|
||||
return time_expr
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Response Models
|
||||
# =============================================================================
|
||||
@ -236,6 +257,8 @@ async def get_user_activity_stats(
|
||||
from src.storage.db.models import User, Conversation
|
||||
|
||||
now = utc_now()
|
||||
# PostgreSQL with asyncpg requires naive datetime for naive DateTime columns
|
||||
naive_now = now.replace(tzinfo=None)
|
||||
|
||||
# Conversations may store either the numeric user primary key or the login user_id string.
|
||||
# Join condition accounts for both representations.
|
||||
@ -253,7 +276,7 @@ async def get_user_activity_stats(
|
||||
select(func.count(distinct(User.id)))
|
||||
.select_from(Conversation)
|
||||
.join(User, user_join_condition)
|
||||
.filter(Conversation.updated_at >= now - timedelta(days=1), User.is_deleted == 0)
|
||||
.filter(Conversation.updated_at >= naive_now - timedelta(days=1), User.is_deleted == 0)
|
||||
)
|
||||
active_users_24h = active_users_24h_result.scalar() or 0
|
||||
|
||||
@ -261,14 +284,14 @@ async def get_user_activity_stats(
|
||||
select(func.count(distinct(User.id)))
|
||||
.select_from(Conversation)
|
||||
.join(User, user_join_condition)
|
||||
.filter(Conversation.updated_at >= now - timedelta(days=30), User.is_deleted == 0)
|
||||
.filter(Conversation.updated_at >= naive_now - timedelta(days=30), User.is_deleted == 0)
|
||||
)
|
||||
active_users_30d = active_users_30d_result.scalar() or 0
|
||||
# 最近7天每日活跃用户(排除已删除用户)
|
||||
daily_active_users = []
|
||||
for i in range(7):
|
||||
day_start = now - timedelta(days=i + 1)
|
||||
day_end = now - timedelta(days=i)
|
||||
day_start = naive_now - timedelta(days=i + 1)
|
||||
day_end = naive_now - timedelta(days=i)
|
||||
|
||||
active_count_result = await db.execute(
|
||||
select(func.count(distinct(User.id)))
|
||||
@ -308,6 +331,8 @@ async def get_tool_call_stats(
|
||||
from src.storage.db.models import ToolCall
|
||||
|
||||
now = utc_now()
|
||||
# PostgreSQL with asyncpg requires naive datetime for naive DateTime columns
|
||||
naive_now = now.replace(tzinfo=None)
|
||||
|
||||
# 基础工具调用统计
|
||||
total_calls_result = await db.execute(select(func.count(ToolCall.id)))
|
||||
@ -340,8 +365,8 @@ async def get_tool_call_stats(
|
||||
# 最近7天每日工具调用数
|
||||
daily_tool_calls = []
|
||||
for i in range(7):
|
||||
day_start = now - timedelta(days=i + 1)
|
||||
day_end = now - timedelta(days=i)
|
||||
day_start = naive_now - timedelta(days=i + 1)
|
||||
day_end = naive_now - timedelta(days=i)
|
||||
|
||||
daily_count_result = await db.execute(
|
||||
select(func.count(ToolCall.id)).filter(ToolCall.created_at >= day_start, ToolCall.created_at < day_end)
|
||||
@ -647,7 +672,7 @@ async def get_all_feedbacks(
|
||||
.join(Conversation, Message.conversation_id == Conversation.id)
|
||||
.outerjoin(
|
||||
User,
|
||||
(MessageFeedback.user_id == User.id) | (MessageFeedback.user_id == User.user_id),
|
||||
(MessageFeedback.user_id == cast(User.id, String)) | (MessageFeedback.user_id == User.user_id),
|
||||
)
|
||||
)
|
||||
|
||||
@ -724,7 +749,7 @@ async def get_call_timeseries_stats(
|
||||
intervals = 14
|
||||
# 包含当前小时:从13小时前开始
|
||||
start_time = now - timedelta(hours=intervals - 1)
|
||||
group_format = func.strftime("%Y-%m-%d %H:00", func.datetime(Message.created_at, "+8 hours"))
|
||||
group_format = _get_time_group_format(Message.created_at, time_range)
|
||||
base_local_time = ensure_shanghai(start_time)
|
||||
elif time_range == "14weeks":
|
||||
intervals = 14
|
||||
@ -733,40 +758,40 @@ async def get_call_timeseries_stats(
|
||||
local_start = local_start - timedelta(days=local_start.weekday())
|
||||
local_start = local_start.replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
start_time = local_start.astimezone(UTC)
|
||||
group_format = func.strftime("%Y-%W", func.datetime(Message.created_at, "+8 hours"))
|
||||
group_format = _get_time_group_format(Message.created_at, time_range)
|
||||
base_local_time = local_start
|
||||
else: # 14days (default)
|
||||
intervals = 14
|
||||
# 包含当前天:从13天前开始
|
||||
start_time = now - timedelta(days=intervals - 1)
|
||||
group_format = func.strftime("%Y-%m-%d", func.datetime(Message.created_at, "+8 hours"))
|
||||
group_format = _get_time_group_format(Message.created_at, time_range)
|
||||
base_local_time = ensure_shanghai(start_time)
|
||||
|
||||
# Convert start_time to naive UTC datetime for PostgreSQL query
|
||||
# PostgreSQL with asyncpg and naive DateTime columns requires naive datetime objects
|
||||
query_start_time = start_time.replace(tzinfo=None)
|
||||
|
||||
# 根据类型查询数据
|
||||
if type == "models":
|
||||
# 模型调用统计(基于消息数量,按模型分组)
|
||||
# 从message的extra_metadata中提取模型信息
|
||||
category_expr = cast(Message.extra_metadata["response_metadata"]["model_name"], String)
|
||||
query_result = await db.execute(
|
||||
select(
|
||||
group_format.label("date"),
|
||||
func.count(Message.id).label("count"),
|
||||
func.json_extract(Message.extra_metadata, "$.response_metadata.model_name").label("category"),
|
||||
category_expr.label("category"),
|
||||
)
|
||||
.filter(Message.role == "assistant", Message.created_at >= start_time)
|
||||
.filter(Message.role == "assistant", Message.created_at >= query_start_time)
|
||||
.filter(Message.extra_metadata.isnot(None))
|
||||
.group_by(group_format, func.json_extract(Message.extra_metadata, "$.response_metadata.model_name"))
|
||||
.group_by(group_format, category_expr)
|
||||
.order_by(group_format)
|
||||
)
|
||||
query = query_result.all()
|
||||
elif type == "agents":
|
||||
# 智能体调用统计(基于对话更新时间,按智能体分组)
|
||||
# 为对话创建独立的时间格式化器
|
||||
if time_range == "14hours":
|
||||
conv_group_format = func.strftime("%Y-%m-%d %H:00", func.datetime(Conversation.updated_at, "+8 hours"))
|
||||
elif time_range == "14weeks":
|
||||
conv_group_format = func.strftime("%Y-%W", func.datetime(Conversation.updated_at, "+8 hours"))
|
||||
else: # 14days
|
||||
conv_group_format = func.strftime("%Y-%m-%d", func.datetime(Conversation.updated_at, "+8 hours"))
|
||||
# 为对话创建独立的时间格式化器(使用 PostgreSQL 兼容的 to_char + INTERVAL)
|
||||
conv_group_format = _get_time_group_format(Conversation.updated_at, time_range)
|
||||
|
||||
query_result = await db.execute(
|
||||
select(
|
||||
@ -775,7 +800,7 @@ async def get_call_timeseries_stats(
|
||||
Conversation.agent_id.label("category"),
|
||||
)
|
||||
.filter(Conversation.updated_at.isnot(None))
|
||||
.filter(Conversation.updated_at >= start_time)
|
||||
.filter(Conversation.updated_at >= query_start_time)
|
||||
.group_by(conv_group_format, Conversation.agent_id)
|
||||
.order_by(conv_group_format)
|
||||
)
|
||||
@ -789,14 +814,14 @@ async def get_call_timeseries_stats(
|
||||
select(
|
||||
group_format.label("date"),
|
||||
func.sum(
|
||||
func.coalesce(func.json_extract(Message.extra_metadata, "$.usage_metadata.input_tokens"), 0)
|
||||
func.coalesce(cast(cast(Message.extra_metadata["usage_metadata"]["input_tokens"], String), Integer), 0)
|
||||
).label("count"),
|
||||
literal("input_tokens").label("category"),
|
||||
)
|
||||
.filter(
|
||||
Message.created_at >= start_time,
|
||||
Message.created_at >= query_start_time,
|
||||
Message.extra_metadata.isnot(None),
|
||||
func.json_extract(Message.extra_metadata, "$.usage_metadata").isnot(None),
|
||||
Message.extra_metadata["usage_metadata"].isnot(None),
|
||||
)
|
||||
.group_by(group_format)
|
||||
.order_by(group_format)
|
||||
@ -808,14 +833,14 @@ async def get_call_timeseries_stats(
|
||||
select(
|
||||
group_format.label("date"),
|
||||
func.sum(
|
||||
func.coalesce(func.json_extract(Message.extra_metadata, "$.usage_metadata.output_tokens"), 0)
|
||||
func.coalesce(cast(cast(Message.extra_metadata["usage_metadata"]["output_tokens"], String), Integer), 0)
|
||||
).label("count"),
|
||||
literal("output_tokens").label("category"),
|
||||
)
|
||||
.filter(
|
||||
Message.created_at >= start_time,
|
||||
Message.created_at >= query_start_time,
|
||||
Message.extra_metadata.isnot(None),
|
||||
func.json_extract(Message.extra_metadata, "$.usage_metadata").isnot(None),
|
||||
Message.extra_metadata["usage_metadata"].isnot(None),
|
||||
)
|
||||
.group_by(group_format)
|
||||
.order_by(group_format)
|
||||
@ -828,13 +853,8 @@ async def get_call_timeseries_stats(
|
||||
results = input_results + output_results
|
||||
elif type == "tools":
|
||||
# 工具调用统计(按工具名称分组)
|
||||
# 为工具调用创建独立的时间格式化器
|
||||
if time_range == "14hours":
|
||||
tool_group_format = func.strftime("%Y-%m-%d %H:00", func.datetime(ToolCall.created_at, "+8 hours"))
|
||||
elif time_range == "14weeks":
|
||||
tool_group_format = func.strftime("%Y-%W", func.datetime(ToolCall.created_at, "+8 hours"))
|
||||
else: # 14days
|
||||
tool_group_format = func.strftime("%Y-%m-%d", func.datetime(ToolCall.created_at, "+8 hours"))
|
||||
# 为工具调用创建独立的时间格式化器(使用 PostgreSQL 兼容的 to_char + INTERVAL)
|
||||
tool_group_format = _get_time_group_format(ToolCall.created_at, time_range)
|
||||
|
||||
query_result = await db.execute(
|
||||
select(
|
||||
@ -842,7 +862,7 @@ async def get_call_timeseries_stats(
|
||||
func.count(ToolCall.id).label("count"),
|
||||
ToolCall.tool_name.label("category"),
|
||||
)
|
||||
.filter(ToolCall.created_at >= start_time)
|
||||
.filter(ToolCall.created_at >= query_start_time)
|
||||
.group_by(tool_group_format, ToolCall.tool_name)
|
||||
.order_by(tool_group_format)
|
||||
)
|
||||
|
||||
@ -10,7 +10,9 @@ from pydantic import BaseModel
|
||||
from sqlalchemy import select, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.storage.db.models import Department, User
|
||||
from src.storage.postgres.models_business import Department, User
|
||||
from src.repositories.department_repository import DepartmentRepository
|
||||
from src.repositories.user_repository import UserRepository
|
||||
from server.utils.auth_middleware import get_superadmin_user, get_admin_user, get_db
|
||||
from server.utils.auth_utils import AuthUtils
|
||||
from server.utils.common_utils import log_operation
|
||||
@ -65,26 +67,11 @@ class DepartmentResponse(BaseModel):
|
||||
# =============================================================================
|
||||
|
||||
|
||||
async def _get_departments_with_user_count(db: AsyncSession) -> list[dict]:
|
||||
"""获取所有部门列表,包含用户数量(内部辅助函数)"""
|
||||
result = await db.execute(select(Department).order_by(Department.created_at.desc()))
|
||||
departments = result.scalars().all()
|
||||
|
||||
department_list = []
|
||||
for dep in departments:
|
||||
user_count_result = await db.execute(
|
||||
select(func.count(User.id)).filter(User.department_id == dep.id, User.is_deleted == 0)
|
||||
)
|
||||
user_count = user_count_result.scalar()
|
||||
department_list.append({**dep.to_dict(), "user_count": user_count})
|
||||
|
||||
return department_list
|
||||
|
||||
|
||||
@department.get("", response_model=list[DepartmentResponse])
|
||||
async def get_departments(current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db)):
|
||||
"""获取所有部门列表(管理员可访问)"""
|
||||
return await _get_departments_with_user_count(db)
|
||||
dept_repo = DepartmentRepository()
|
||||
return await dept_repo.list_with_user_count()
|
||||
|
||||
|
||||
@department.get("/{department_id}", response_model=DepartmentResponse)
|
||||
@ -115,10 +102,11 @@ async def create_department(
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""创建新部门,同时创建该部门的管理员"""
|
||||
dept_repo = DepartmentRepository()
|
||||
user_repo = UserRepository()
|
||||
|
||||
# 检查部门名称是否已存在
|
||||
result = await db.execute(select(Department).filter(Department.name == department_data.name))
|
||||
existing = result.scalar_one_or_none()
|
||||
if existing:
|
||||
if await dept_repo.exists_by_name(department_data.name):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="部门名称已存在")
|
||||
|
||||
# 验证管理员 user_id 格式
|
||||
@ -136,9 +124,7 @@ async def create_department(
|
||||
)
|
||||
|
||||
# 检查 user_id 是否已存在
|
||||
result = await db.execute(select(User).filter(User.user_id == admin_user_id))
|
||||
existing_user = result.scalar_one_or_none()
|
||||
if existing_user:
|
||||
if await user_repo.exists_by_user_id(admin_user_id):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="用户ID已存在",
|
||||
@ -149,33 +135,28 @@ async def create_department(
|
||||
if admin_phone:
|
||||
if not is_valid_phone_number(admin_phone):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="手机号格式不正确")
|
||||
result = await db.execute(select(User).filter(User.phone_number == admin_phone))
|
||||
existing_phone = result.scalar_one_or_none()
|
||||
if existing_phone:
|
||||
if await user_repo.exists_by_phone(admin_phone):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="手机号已存在",
|
||||
)
|
||||
|
||||
new_department = Department(name=department_data.name, description=department_data.description)
|
||||
|
||||
db.add(new_department)
|
||||
await db.flush() # 获取部门ID
|
||||
# 创建部门
|
||||
new_department = await dept_repo.create({
|
||||
"name": department_data.name,
|
||||
"description": department_data.description,
|
||||
})
|
||||
|
||||
# 创建管理员用户
|
||||
hashed_password = AuthUtils.hash_password(department_data.admin_password)
|
||||
new_admin = User(
|
||||
username=admin_user_id, # username 和 user_id 设置为相同值
|
||||
user_id=admin_user_id,
|
||||
phone_number=admin_phone,
|
||||
password_hash=hashed_password,
|
||||
role="admin",
|
||||
department_id=new_department.id,
|
||||
)
|
||||
db.add(new_admin)
|
||||
|
||||
await db.commit()
|
||||
await db.refresh(new_department)
|
||||
await user_repo.create({
|
||||
"username": admin_user_id,
|
||||
"user_id": admin_user_id,
|
||||
"phone_number": admin_phone,
|
||||
"password_hash": hashed_password,
|
||||
"role": "admin",
|
||||
"department_id": new_department.id,
|
||||
})
|
||||
|
||||
# 记录操作
|
||||
await log_operation(
|
||||
|
||||
@ -31,12 +31,12 @@ async def _get_graph_adapter(db_id: str) -> GraphAdapter:
|
||||
# 检查图数据库服务状态 (仅对 Upload 类型需要)
|
||||
if not graph_base.is_running():
|
||||
# 先尝试检测图谱类型,如果是不需要 graph_base 的类型则允许
|
||||
graph_type = GraphAdapterFactory.detect_graph_type(db_id, knowledge_base)
|
||||
graph_type = await GraphAdapterFactory.detect_graph_type(db_id, knowledge_base)
|
||||
if graph_type == "upload":
|
||||
raise HTTPException(status_code=503, detail="Graph database service is not running")
|
||||
|
||||
# 使用工厂方法自动创建适配器
|
||||
return GraphAdapterFactory.create_adapter_by_db_id(
|
||||
return await GraphAdapterFactory.create_adapter_by_db_id(
|
||||
db_id=db_id, knowledge_base_manager=knowledge_base, graph_db_instance=graph_base
|
||||
)
|
||||
|
||||
@ -83,7 +83,7 @@ async def get_graphs(current_user: User = Depends(get_admin_user)):
|
||||
)
|
||||
|
||||
# 2. 获取 LightRAG 数据库信息
|
||||
lightrag_dbs = knowledge_base.get_lightrag_databases()
|
||||
lightrag_dbs = await knowledge_base.get_lightrag_databases()
|
||||
# 直接使用 LightRAG 适配器的默认 metadata
|
||||
from src.knowledge.adapters.lightrag import LightRAGGraphAdapter
|
||||
|
||||
|
||||
@ -5,8 +5,8 @@ from fastapi.security import OAuth2PasswordBearer
|
||||
from jose import JWTError
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.storage.db.manager import db_manager
|
||||
from src.storage.db.models import User
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import User
|
||||
from server.utils.auth_utils import AuthUtils
|
||||
|
||||
# 定义OAuth2密码承载器,指定token URL
|
||||
@ -25,7 +25,7 @@ PUBLIC_PATHS = [
|
||||
|
||||
# 获取数据库会话(异步版本)
|
||||
async def get_db():
|
||||
async with db_manager.get_async_session_context() as db:
|
||||
async with pg_manager.get_async_session_context() as db:
|
||||
yield db
|
||||
|
||||
|
||||
@ -61,7 +61,7 @@ async def get_current_user(token: str | None = Depends(oauth2_scheme), db: Async
|
||||
# 查找用户(异步版本)
|
||||
from sqlalchemy import select
|
||||
|
||||
result = await db.execute(select(User).filter(User.id == user_id))
|
||||
result = await db.execute(select(User).filter(User.id == int(user_id)))
|
||||
user = result.scalar_one_or_none()
|
||||
if user is None:
|
||||
raise credentials_exception
|
||||
|
||||
@ -5,7 +5,7 @@ import logging
|
||||
from fastapi import Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.storage.db.models import OperationLog, User
|
||||
from src.storage.postgres.models_business import OperationLog, User
|
||||
|
||||
|
||||
def setup_logging():
|
||||
|
||||
@ -34,7 +34,7 @@ class GraphAdapterFactory:
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def detect_graph_type(cls, db_id: str, knowledge_base_manager=None) -> str:
|
||||
async def detect_graph_type(cls, db_id: str, knowledge_base_manager=None) -> str:
|
||||
"""
|
||||
自动检测图谱类型
|
||||
|
||||
@ -47,7 +47,7 @@ class GraphAdapterFactory:
|
||||
"""
|
||||
# 1. 首先检查是否是 LightRAG 数据库 (通过知识库管理器)
|
||||
if knowledge_base_manager:
|
||||
db_info = knowledge_base_manager.get_database_info(db_id)
|
||||
db_info = await knowledge_base_manager.get_database_info(db_id)
|
||||
if db_info: # 有信息表示是 LightRAG 数据库
|
||||
return "lightrag"
|
||||
|
||||
@ -59,7 +59,7 @@ class GraphAdapterFactory:
|
||||
return "upload"
|
||||
|
||||
@classmethod
|
||||
def create_adapter_by_db_id(cls, db_id: str, knowledge_base_manager=None, graph_db_instance=None) -> GraphAdapter:
|
||||
async def create_adapter_by_db_id(cls, db_id: str, knowledge_base_manager=None, graph_db_instance=None) -> GraphAdapter:
|
||||
"""
|
||||
根据数据库ID自动创建对应的适配器
|
||||
|
||||
@ -71,7 +71,7 @@ class GraphAdapterFactory:
|
||||
Returns:
|
||||
对应的图谱适配器
|
||||
"""
|
||||
graph_type = cls.detect_graph_type(db_id, knowledge_base_manager)
|
||||
graph_type = await cls.detect_graph_type(db_id, knowledge_base_manager)
|
||||
|
||||
if graph_type == "lightrag":
|
||||
# LightRAG 类型,使用 kb_id 作为配置
|
||||
@ -81,8 +81,8 @@ class GraphAdapterFactory:
|
||||
return cls.create_adapter("upload", graph_db_instance=graph_db_instance, config={"kgdb_name": db_id})
|
||||
|
||||
@classmethod
|
||||
def create_adapter_for_db_id(cls, db_id: str, knowledge_base_manager=None, graph_db_instance=None) -> GraphAdapter:
|
||||
async def create_adapter_for_db_id(cls, db_id: str, knowledge_base_manager=None, graph_db_instance=None) -> GraphAdapter:
|
||||
"""
|
||||
兼容性方法,调用 create_adapter_by_db_id
|
||||
"""
|
||||
return cls.create_adapter_by_db_id(db_id, knowledge_base_manager, graph_db_instance)
|
||||
return await cls.create_adapter_by_db_id(db_id, knowledge_base_manager, graph_db_instance)
|
||||
|
||||
227
src/repositories/conversation_repository.py
Normal file
227
src/repositories/conversation_repository.py
Normal file
@ -0,0 +1,227 @@
|
||||
"""对话数据访问层 - Repository"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import Conversation, ConversationStats, Message, ToolCall
|
||||
|
||||
|
||||
class ConversationRepository:
|
||||
"""对话数据访问层"""
|
||||
|
||||
async def get_by_thread_id(self, thread_id: str) -> Conversation | None:
|
||||
"""根据 thread_id 获取对话"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Conversation).where(Conversation.thread_id == thread_id))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def get_by_id(self, id: int) -> Conversation | None:
|
||||
"""根据 ID 获取对话"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Conversation).where(Conversation.id == id))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def list_by_user(
|
||||
self, user_id: str, skip: int = 0, limit: int = 100, agent_id: str | None = None, status: str = "active"
|
||||
) -> list[Conversation]:
|
||||
"""获取用户的对话列表"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
query = select(Conversation).where(Conversation.user_id == str(user_id), Conversation.status == status)
|
||||
if agent_id:
|
||||
query = query.where(Conversation.agent_id == agent_id)
|
||||
query = query.order_by(Conversation.updated_at.desc()).offset(skip).limit(limit)
|
||||
result = await session.execute(query)
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def create(self, data: dict[str, Any]) -> Conversation:
|
||||
"""创建对话"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
conversation = Conversation(**data)
|
||||
session.add(conversation)
|
||||
await session.flush()
|
||||
|
||||
# 创建关联的 stats 记录
|
||||
stats = ConversationStats(conversation_id=conversation.id)
|
||||
session.add(stats)
|
||||
return conversation
|
||||
|
||||
async def update(self, id: int, data: dict[str, Any]) -> Conversation | None:
|
||||
"""更新对话"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Conversation).where(Conversation.id == id))
|
||||
conversation = result.scalar_one_or_none()
|
||||
if conversation is None:
|
||||
return None
|
||||
for key, value in data.items():
|
||||
if key not in ("id", "thread_id", "user_id", "agent_id"):
|
||||
setattr(conversation, key, value)
|
||||
return conversation
|
||||
|
||||
async def update_status(self, id: int, status: str) -> bool:
|
||||
"""更新对话状态"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Conversation).where(Conversation.id == id))
|
||||
conversation = result.scalar_one_or_none()
|
||||
if conversation is None:
|
||||
return False
|
||||
conversation.status = status
|
||||
return True
|
||||
|
||||
async def delete(self, id: int, soft_delete: bool = True) -> bool:
|
||||
"""删除对话"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Conversation).where(Conversation.id == id))
|
||||
conversation = result.scalar_one_or_none()
|
||||
if conversation is None:
|
||||
return False
|
||||
if soft_delete:
|
||||
conversation.status = "deleted"
|
||||
else:
|
||||
await session.delete(conversation)
|
||||
return True
|
||||
|
||||
async def count_by_user(self, user_id: str) -> int:
|
||||
"""统计用户对话数量"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(
|
||||
select(func.count(Conversation.id)).where(Conversation.user_id == str(user_id))
|
||||
)
|
||||
return result.scalar() or 0
|
||||
|
||||
|
||||
class MessageRepository:
|
||||
"""消息数据访问层"""
|
||||
|
||||
async def get_by_id(self, id: int) -> Message | None:
|
||||
"""根据 ID 获取消息"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Message).where(Message.id == id))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def list_by_conversation(self, conversation_id: int, limit: int | None = None, offset: int = 0) -> list[Message]:
|
||||
"""获取对话的消息列表"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
query = (
|
||||
select(Message)
|
||||
.options(selectinload(Message.tool_calls), selectinload(Message.feedbacks))
|
||||
.where(Message.conversation_id == conversation_id)
|
||||
.order_by(Message.created_at.asc())
|
||||
)
|
||||
query = query.offset(offset)
|
||||
if limit:
|
||||
query = query.limit(limit)
|
||||
result = await session.execute(query)
|
||||
return list(result.scalars().unique().all())
|
||||
|
||||
async def create(self, data: dict[str, Any]) -> Message:
|
||||
"""创建消息"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
message = Message(**data)
|
||||
session.add(message)
|
||||
return message
|
||||
|
||||
async def count_by_conversation(self, conversation_id: int) -> int:
|
||||
"""统计对话消息数量"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(
|
||||
select(func.count(Message.id)).where(Message.conversation_id == conversation_id)
|
||||
)
|
||||
return result.scalar() or 0
|
||||
|
||||
|
||||
class ToolCallRepository:
|
||||
"""工具调用数据访问层"""
|
||||
|
||||
async def get_by_id(self, id: int) -> ToolCall | None:
|
||||
"""根据 ID 获取工具调用"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(ToolCall).where(ToolCall.id == id))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def get_by_langgraph_id(self, langgraph_tool_call_id: str) -> ToolCall | None:
|
||||
"""根据 LangGraph tool_call_id 获取工具调用"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(
|
||||
select(ToolCall).where(ToolCall.langgraph_tool_call_id == langgraph_tool_call_id)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def list_by_message(self, message_id: int) -> list[ToolCall]:
|
||||
"""获取消息的工具调用列表"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(ToolCall).where(ToolCall.message_id == message_id))
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def create(self, data: dict[str, Any]) -> ToolCall:
|
||||
"""创建工具调用"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
tool_call = ToolCall(**data)
|
||||
session.add(tool_call)
|
||||
return tool_call
|
||||
|
||||
async def update_output(
|
||||
self, langgraph_tool_call_id: str, tool_output: str, status: str = "success", error_message: str | None = None
|
||||
) -> ToolCall | None:
|
||||
"""更新工具调用输出"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(
|
||||
select(ToolCall).where(ToolCall.langgraph_tool_call_id == langgraph_tool_call_id)
|
||||
)
|
||||
tool_call = result.scalar_one_or_none()
|
||||
if tool_call is None:
|
||||
return None
|
||||
tool_call.tool_output = tool_output
|
||||
tool_call.status = status
|
||||
if error_message:
|
||||
tool_call.error_message = error_message
|
||||
return tool_call
|
||||
|
||||
|
||||
class ConversationStatsRepository:
|
||||
"""对话统计数据访问层"""
|
||||
|
||||
async def get_by_conversation_id(self, conversation_id: int) -> ConversationStats | None:
|
||||
"""根据对话 ID 获取统计信息"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(
|
||||
select(ConversationStats).where(ConversationStats.conversation_id == conversation_id)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def create(self, data: dict[str, Any]) -> ConversationStats:
|
||||
"""创建统计信息"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
stats = ConversationStats(**data)
|
||||
session.add(stats)
|
||||
return stats
|
||||
|
||||
async def update(
|
||||
self, conversation_id: int, data: dict[str, Any]
|
||||
) -> ConversationStats | None:
|
||||
"""更新统计信息"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(
|
||||
select(ConversationStats).where(ConversationStats.conversation_id == conversation_id)
|
||||
)
|
||||
stats = result.scalar_one_or_none()
|
||||
if stats is None:
|
||||
return None
|
||||
for key, value in data.items():
|
||||
if key != "conversation_id":
|
||||
setattr(stats, key, value)
|
||||
return stats
|
||||
|
||||
async def update_message_count(self, conversation_id: int, message_count: int) -> bool:
|
||||
"""更新消息计数"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(
|
||||
select(ConversationStats).where(ConversationStats.conversation_id == conversation_id)
|
||||
)
|
||||
stats = result.scalar_one_or_none()
|
||||
if stats is None:
|
||||
return False
|
||||
stats.message_count = message_count
|
||||
return True
|
||||
95
src/repositories/department_repository.py
Normal file
95
src/repositories/department_repository.py
Normal file
@ -0,0 +1,95 @@
|
||||
"""部门数据访问层 - Repository"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import Department
|
||||
|
||||
|
||||
class DepartmentRepository:
|
||||
"""部门数据访问层"""
|
||||
|
||||
async def get_by_id(self, id: int) -> Department | None:
|
||||
"""根据 ID 获取部门"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Department).where(Department.id == id))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def get_by_name(self, name: str) -> Department | None:
|
||||
"""根据名称获取部门"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Department).where(Department.name == name))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def list_departments(self) -> list[Department]:
|
||||
"""获取所有部门列表"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Department).order_by(Department.created_at.desc()))
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def list_with_user_count(self) -> list[dict[str, Any]]:
|
||||
"""获取所有部门列表,包含用户数量"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
from src.storage.postgres.models_business import User
|
||||
|
||||
result = await session.execute(select(Department).order_by(Department.created_at.desc()))
|
||||
departments = result.scalars().all()
|
||||
|
||||
department_list = []
|
||||
for dep in departments:
|
||||
user_count_result = await session.execute(
|
||||
select(func.count(User.id)).where(User.department_id == dep.id, User.is_deleted == 0)
|
||||
)
|
||||
user_count = user_count_result.scalar()
|
||||
dep_dict = dep.to_dict()
|
||||
dep_dict["user_count"] = user_count
|
||||
department_list.append(dep_dict)
|
||||
|
||||
return department_list
|
||||
|
||||
async def create(self, data: dict[str, Any]) -> Department:
|
||||
"""创建部门"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
department = Department(**data)
|
||||
session.add(department)
|
||||
return department
|
||||
|
||||
async def update(self, id: int, data: dict[str, Any]) -> Department | None:
|
||||
"""更新部门"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Department).where(Department.id == id))
|
||||
department = result.scalar_one_or_none()
|
||||
if department is None:
|
||||
return None
|
||||
for key, value in data.items():
|
||||
if key != "id":
|
||||
setattr(department, key, value)
|
||||
return department
|
||||
|
||||
async def delete(self, id: int) -> bool:
|
||||
"""删除部门"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Department).where(Department.id == id))
|
||||
department = result.scalar_one_or_none()
|
||||
if department is None:
|
||||
return False
|
||||
await session.delete(department)
|
||||
return True
|
||||
|
||||
async def count_users(self, id: int) -> int:
|
||||
"""统计部门用户数量"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
from src.storage.postgres.models_business import User
|
||||
|
||||
result = await session.execute(
|
||||
select(func.count(User.id)).where(User.department_id == id, User.is_deleted == 0)
|
||||
)
|
||||
return result.scalar() or 0
|
||||
|
||||
async def exists_by_name(self, name: str) -> bool:
|
||||
"""检查部门名称是否存在"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(Department.id).where(Department.name == name))
|
||||
return result.scalar_one_or_none() is not None
|
||||
81
src/repositories/mcp_server_repository.py
Normal file
81
src/repositories/mcp_server_repository.py
Normal file
@ -0,0 +1,81 @@
|
||||
"""MCP 服务器数据访问层 - Repository"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import MCPServer
|
||||
|
||||
|
||||
class MCPServerRepository:
|
||||
"""MCP 服务器数据访问层"""
|
||||
|
||||
async def get_by_name(self, name: str) -> MCPServer | None:
|
||||
"""根据名称获取 MCP 服务器"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(MCPServer).where(MCPServer.name == name))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def list(self) -> list[MCPServer]:
|
||||
"""获取所有 MCP 服务器"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(MCPServer))
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def list_enabled(self) -> list[MCPServer]:
|
||||
"""获取所有启用的 MCP 服务器"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(MCPServer).where(MCPServer.enabled == 1))
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def create(self, data: dict[str, Any]) -> MCPServer:
|
||||
"""创建 MCP 服务器"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
server = MCPServer(**data)
|
||||
session.add(server)
|
||||
return server
|
||||
|
||||
async def update(self, name: str, data: dict[str, Any]) -> MCPServer | None:
|
||||
"""更新 MCP 服务器"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(MCPServer).where(MCPServer.name == name))
|
||||
server = result.scalar_one_or_none()
|
||||
if server is None:
|
||||
return None
|
||||
for key, value in data.items():
|
||||
if key != "name":
|
||||
setattr(server, key, value)
|
||||
return server
|
||||
|
||||
async def delete(self, name: str) -> bool:
|
||||
"""删除 MCP 服务器"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(MCPServer).where(MCPServer.name == name))
|
||||
server = result.scalar_one_or_none()
|
||||
if server is None:
|
||||
return False
|
||||
await session.delete(server)
|
||||
return True
|
||||
|
||||
async def upsert(self, data: dict[str, Any]) -> MCPServer:
|
||||
"""插入或更新 MCP 服务器"""
|
||||
name = data.get("name")
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(MCPServer).where(MCPServer.name == name))
|
||||
existing = result.scalar_one_or_none()
|
||||
if existing is None:
|
||||
server = MCPServer(**data)
|
||||
session.add(server)
|
||||
else:
|
||||
for key, value in data.items():
|
||||
if key != "name":
|
||||
setattr(existing, key, value)
|
||||
server = existing
|
||||
return server
|
||||
|
||||
async def exists_by_name(self, name: str) -> bool:
|
||||
"""检查 MCP 服务器是否存在"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(MCPServer.id).where(MCPServer.name == name))
|
||||
return result.scalar_one_or_none() is not None
|
||||
44
src/repositories/message_feedback_repository.py
Normal file
44
src/repositories/message_feedback_repository.py
Normal file
@ -0,0 +1,44 @@
|
||||
"""消息反馈数据访问层 - Repository"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import MessageFeedback
|
||||
|
||||
|
||||
class MessageFeedbackRepository:
|
||||
"""消息反馈数据访问层"""
|
||||
|
||||
async def get_by_id(self, id: int) -> MessageFeedback | None:
|
||||
"""根据 ID 获取消息反馈"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(MessageFeedback).where(MessageFeedback.id == id))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def list_by_message(self, message_id: int) -> list[MessageFeedback]:
|
||||
"""获取消息的反馈列表"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(
|
||||
select(MessageFeedback).where(MessageFeedback.message_id == message_id)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def create(self, data: dict[str, Any]) -> MessageFeedback:
|
||||
"""创建消息反馈"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
feedback = MessageFeedback(**data)
|
||||
session.add(feedback)
|
||||
return feedback
|
||||
|
||||
async def exists_by_message_and_user(self, message_id: int, user_id: str) -> bool:
|
||||
"""检查用户是否已对消息反馈"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(
|
||||
select(MessageFeedback.id).where(
|
||||
MessageFeedback.message_id == message_id,
|
||||
MessageFeedback.user_id == user_id
|
||||
)
|
||||
)
|
||||
return result.scalar_one_or_none() is not None
|
||||
49
src/repositories/operation_log_repository.py
Normal file
49
src/repositories/operation_log_repository.py
Normal file
@ -0,0 +1,49 @@
|
||||
"""操作日志数据访问层 - Repository"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import OperationLog
|
||||
|
||||
|
||||
class OperationLogRepository:
|
||||
"""操作日志数据访问层"""
|
||||
|
||||
async def get_by_id(self, id: int) -> OperationLog | None:
|
||||
"""根据 ID 获取操作日志"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(OperationLog).where(OperationLog.id == id))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def list_by_user(
|
||||
self, user_id: int, skip: int = 0, limit: int = 100
|
||||
) -> list[OperationLog]:
|
||||
"""获取用户的操作日志列表"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(
|
||||
select(OperationLog)
|
||||
.where(OperationLog.user_id == user_id)
|
||||
.order_by(OperationLog.timestamp.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def create(self, data: dict[str, Any]) -> OperationLog:
|
||||
"""创建操作日志"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
log = OperationLog(**data)
|
||||
session.add(log)
|
||||
return log
|
||||
|
||||
async def count_by_user(self, user_id: int) -> int:
|
||||
"""统计用户操作日志数量"""
|
||||
from sqlalchemy import func
|
||||
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(
|
||||
select(func.count(OperationLog.id)).where(OperationLog.user_id == user_id)
|
||||
)
|
||||
return result.scalar() or 0
|
||||
146
src/repositories/user_repository.py
Normal file
146
src/repositories/user_repository.py
Normal file
@ -0,0 +1,146 @@
|
||||
"""用户数据访问层 - Repository"""
|
||||
|
||||
from datetime import datetime as dt, timezone
|
||||
from typing import Any, Annotated
|
||||
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import User
|
||||
|
||||
# 使用 naive datetime 以兼容 PostgreSQL TIMESTAMP WITHOUT TIME ZONE 列
|
||||
_utc_now = lambda: dt.now(timezone.utc).replace(tzinfo=None)
|
||||
|
||||
|
||||
class UserRepository:
|
||||
"""用户数据访问层"""
|
||||
|
||||
async def get_by_id(self, id: int) -> User | None:
|
||||
"""根据 ID 获取用户"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(User).where(User.id == id))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def get_by_user_id(self, user_id: str) -> User | None:
|
||||
"""根据 user_id 获取用户"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(User).where(User.user_id == user_id))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def get_by_phone(self, phone: str) -> User | None:
|
||||
"""根据手机号获取用户"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(User).where(User.phone_number == phone))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def list_users(
|
||||
self, skip: int = 0, limit: int = 100, department_id: int | None = None, role: str | None = None
|
||||
) -> list[User]:
|
||||
"""获取用户列表"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
query = select(User).where(User.is_deleted == 0)
|
||||
if department_id is not None:
|
||||
query = query.where(User.department_id == department_id)
|
||||
if role is not None:
|
||||
query = query.where(User.role == role)
|
||||
query = query.offset(skip).limit(limit)
|
||||
result = await session.execute(query)
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def list_with_department(
|
||||
self, skip: int = 0, limit: int = 100, department_id: int | None = None, role: str | None = None
|
||||
) -> Annotated[list[tuple[User, str | None]], "用户列表,包含部门名称"]:
|
||||
"""获取用户列表,包含部门名称"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
from src.storage.postgres.models_business import Department
|
||||
|
||||
query = (
|
||||
select(User, Department.name.label("department_name"))
|
||||
.outerjoin(Department, User.department_id == Department.id)
|
||||
.where(User.is_deleted == 0)
|
||||
)
|
||||
if department_id is not None:
|
||||
query = query.where(User.department_id == department_id)
|
||||
if role is not None:
|
||||
query = query.where(User.role == role)
|
||||
query = query.offset(skip).limit(limit)
|
||||
result = await session.execute(query)
|
||||
return list(result.all())
|
||||
|
||||
async def create(self, data: dict[str, Any]) -> User:
|
||||
"""创建用户"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
user = User(**data)
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
return user
|
||||
|
||||
async def update(self, id: int, data: dict[str, Any]) -> User | None:
|
||||
"""更新用户"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(User).where(User.id == id, User.is_deleted == 0))
|
||||
user = result.scalar_one_or_none()
|
||||
if user is None:
|
||||
return None
|
||||
for key, value in data.items():
|
||||
if key != "id":
|
||||
setattr(user, key, value)
|
||||
return user
|
||||
|
||||
async def soft_delete(self, id: int, username: str | None = None, phone_number: str | None = None) -> bool:
|
||||
"""软删除用户"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(User).where(User.id == id, User.is_deleted == 0))
|
||||
user = result.scalar_one_or_none()
|
||||
if user is None:
|
||||
return False
|
||||
user.is_deleted = 1
|
||||
|
||||
user.deleted_at = _utc_now()
|
||||
if username:
|
||||
import hashlib
|
||||
|
||||
hash_suffix = hashlib.md5(user.user_id.encode()).hexdigest()[:4]
|
||||
user.username = f"已注销用户-{hash_suffix}"
|
||||
if phone_number:
|
||||
user.phone_number = None
|
||||
return True
|
||||
|
||||
async def exists_by_user_id(self, user_id: str) -> bool:
|
||||
"""检查 user_id 是否存在"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(User.id).where(User.user_id == user_id))
|
||||
return result.scalar_one_or_none() is not None
|
||||
|
||||
async def exists_by_phone(self, phone: str) -> bool:
|
||||
"""检查手机号是否存在"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(User.id).where(User.phone_number == phone))
|
||||
return result.scalar_one_or_none() is not None
|
||||
|
||||
async def count(self, department_id: int | None = None) -> int:
|
||||
"""统计用户数量"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
query = select(func.count(User.id)).where(User.is_deleted == 0)
|
||||
if department_id is not None:
|
||||
query = query.where(User.department_id == department_id)
|
||||
result = await session.execute(query)
|
||||
return result.scalar() or 0
|
||||
|
||||
async def get_all_user_ids(self) -> list[str]:
|
||||
"""获取所有用户 ID"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(User.user_id))
|
||||
return [uid for (uid,) in result.all()]
|
||||
|
||||
async def get_admin_count_in_department(self, department_id: int, exclude_user_id: int | None = None) -> int:
|
||||
"""统计部门中管理员数量"""
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
query = select(func.count(User.id)).where(
|
||||
User.department_id == department_id, User.role == "admin", User.is_deleted == 0
|
||||
)
|
||||
if exclude_user_id is not None:
|
||||
query = query.where(User.id != exclude_user_id)
|
||||
result = await session.execute(query)
|
||||
return result.scalar() or 0
|
||||
@ -5,13 +5,14 @@ Manages conversation data storage including messages, tool calls, and statistics
|
||||
All database operations are now asynchronous for improved performance.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
import uuid as uuid_lib
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
from src.storage.db.models import Conversation, ConversationStats, Message, ToolCall
|
||||
from src.storage.postgres.models_business import Conversation, ConversationStats, Message, ToolCall
|
||||
from src.utils import logger
|
||||
from src.utils.datetime_utils import utc_now
|
||||
|
||||
@ -20,31 +21,24 @@ class ConversationManager:
|
||||
"""Async Manager for conversation storage operations"""
|
||||
|
||||
def __init__(self, db_session: AsyncSession):
|
||||
"""初始化 ConversationManager
|
||||
|
||||
Args:
|
||||
db_session: 异步数据库会话
|
||||
"""
|
||||
self.db = db_session
|
||||
|
||||
async def create_conversation(
|
||||
self,
|
||||
user_id: str,
|
||||
agent_id: str,
|
||||
title: str | None = None,
|
||||
thread_id: str | None = None,
|
||||
metadata: dict | None = None,
|
||||
title: Optional[str] = None,
|
||||
thread_id: Optional[str] = None,
|
||||
metadata: Optional[dict] = None,
|
||||
) -> Conversation:
|
||||
"""
|
||||
Create a new conversation
|
||||
|
||||
Args:
|
||||
user_id: User ID
|
||||
agent_id: Agent ID
|
||||
title: Conversation title
|
||||
thread_id: Optional thread ID (will generate UUID if not provided)
|
||||
metadata: Optional additional metadata
|
||||
|
||||
Returns:
|
||||
Created Conversation object
|
||||
"""
|
||||
"""创建新对话"""
|
||||
if not thread_id:
|
||||
thread_id = str(uuid.uuid4())
|
||||
thread_id = str(uuid_lib.uuid4())
|
||||
|
||||
metadata = (metadata or {}).copy()
|
||||
metadata.setdefault("attachments", [])
|
||||
@ -59,10 +53,9 @@ class ConversationManager:
|
||||
)
|
||||
|
||||
self.db.add(conversation)
|
||||
# Flush to assign primary key without committing
|
||||
await self.db.flush()
|
||||
|
||||
# Create associated stats record and commit once
|
||||
# 创建关联的 stats 记录
|
||||
stats = ConversationStats(conversation_id=conversation.id)
|
||||
self.db.add(stats)
|
||||
await self.db.commit()
|
||||
@ -71,30 +64,17 @@ class ConversationManager:
|
||||
logger.info(f"Created conversation: {conversation.thread_id} for user {user_id}")
|
||||
return conversation
|
||||
|
||||
async def get_conversation_by_thread_id(self, thread_id: str) -> Conversation | None:
|
||||
"""
|
||||
Get conversation by thread ID
|
||||
|
||||
Args:
|
||||
thread_id: Thread ID
|
||||
|
||||
Returns:
|
||||
Conversation object or None if not found
|
||||
"""
|
||||
result = await self.db.execute(select(Conversation).filter(Conversation.thread_id == thread_id))
|
||||
async def get_conversation_by_thread_id(self, thread_id: str) -> Optional[Conversation]:
|
||||
"""根据 thread_id 获取对话"""
|
||||
result = await self.db.execute(select(Conversation).where(Conversation.thread_id == thread_id))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def _get_conversation_by_id(self, conversation_id: int) -> Conversation | None:
|
||||
result = await self.db.execute(select(Conversation).filter(Conversation.id == conversation_id))
|
||||
async def _get_conversation_by_id(self, conversation_id: int) -> Optional[Conversation]:
|
||||
result = await self.db.execute(select(Conversation).where(Conversation.id == conversation_id))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
def _ensure_metadata(self, conversation: Conversation) -> dict:
|
||||
"""
|
||||
Return a shallow copy of conversation metadata with a standalone attachments list.
|
||||
|
||||
We copy here because SQLAlchemy's JSON type does not automatically detect in-place
|
||||
mutations. By assigning a fresh dict/list back we ensure the ORM marks the row dirty.
|
||||
"""
|
||||
"""确保元数据是独立副本"""
|
||||
metadata = dict(conversation.extra_metadata or {})
|
||||
metadata["attachments"] = list(metadata.get("attachments", []))
|
||||
return metadata
|
||||
@ -111,23 +91,10 @@ class ConversationManager:
|
||||
role: str,
|
||||
content: str,
|
||||
message_type: str = "text",
|
||||
extra_metadata: dict | None = None,
|
||||
image_content: str | None = None,
|
||||
extra_metadata: Optional[dict] = None,
|
||||
image_content: Optional[str] = None,
|
||||
) -> Message:
|
||||
"""
|
||||
Add a message to a conversation
|
||||
|
||||
Args:
|
||||
conversation_id: Conversation ID
|
||||
role: Message role (user/assistant/system/tool)
|
||||
content: Message content
|
||||
message_type: Message type (text/tool_call/tool_result/multimodal_image)
|
||||
extra_metadata: Additional metadata (complete message dump)
|
||||
image_content: Base64 encoded image content for multimodal messages
|
||||
|
||||
Returns:
|
||||
Created Message object
|
||||
"""
|
||||
"""添加消息到对话"""
|
||||
message = Message(
|
||||
conversation_id=conversation_id,
|
||||
role=role,
|
||||
@ -138,7 +105,7 @@ class ConversationManager:
|
||||
)
|
||||
|
||||
self.db.add(message)
|
||||
# Mark the parent conversation as active for sorting/analytics
|
||||
# 更新父对话的更新时间
|
||||
conversation = await self._get_conversation_by_id(conversation_id)
|
||||
if conversation:
|
||||
conversation.updated_at = utc_now()
|
||||
@ -146,7 +113,7 @@ class ConversationManager:
|
||||
await self.db.commit()
|
||||
await self.db.refresh(message)
|
||||
|
||||
# Update conversation stats
|
||||
# 更新对话统计
|
||||
await self._update_message_count(conversation_id)
|
||||
|
||||
logger.debug(f"Added {role} message to conversation {conversation_id}")
|
||||
@ -158,23 +125,10 @@ class ConversationManager:
|
||||
role: str,
|
||||
content: str,
|
||||
message_type: str = "text",
|
||||
extra_metadata: dict | None = None,
|
||||
image_content: str | None = None,
|
||||
) -> Message | None:
|
||||
"""
|
||||
Add a message to a conversation by thread ID
|
||||
|
||||
Args:
|
||||
thread_id: Thread ID
|
||||
role: Message role (user/assistant/system/tool)
|
||||
content: Message content
|
||||
message_type: Message type (text/tool_call/tool_result/multimodal_image)
|
||||
extra_metadata: Additional metadata (complete message dump)
|
||||
image_content: Base64 encoded image content for multimodal messages
|
||||
|
||||
Returns:
|
||||
Created Message object or None if conversation not found
|
||||
"""
|
||||
extra_metadata: Optional[dict] = None,
|
||||
image_content: Optional[str] = None,
|
||||
) -> Optional[Message]:
|
||||
"""根据 thread_id 添加消息到对话"""
|
||||
conversation = await self.get_conversation_by_thread_id(thread_id)
|
||||
if not conversation:
|
||||
logger.warning(f"Conversation not found for thread_id: {thread_id}")
|
||||
@ -193,27 +147,13 @@ class ConversationManager:
|
||||
self,
|
||||
message_id: int,
|
||||
tool_name: str,
|
||||
tool_input: dict | None = None,
|
||||
tool_output: str | None = None,
|
||||
tool_input: Optional[dict] = None,
|
||||
tool_output: Optional[str] = None,
|
||||
status: str = "pending",
|
||||
error_message: str | None = None,
|
||||
langgraph_tool_call_id: str | None = None,
|
||||
error_message: Optional[str] = None,
|
||||
langgraph_tool_call_id: Optional[str] = None,
|
||||
) -> ToolCall:
|
||||
"""
|
||||
Add a tool call record
|
||||
|
||||
Args:
|
||||
message_id: Message ID
|
||||
tool_name: Tool name
|
||||
tool_input: Tool input parameters
|
||||
tool_output: Tool execution result
|
||||
status: Status (pending/success/error)
|
||||
error_message: Error message if failed
|
||||
langgraph_tool_call_id: LangGraph tool_call_id for precise matching
|
||||
|
||||
Returns:
|
||||
Created ToolCall object
|
||||
"""
|
||||
"""添加工具调用记录"""
|
||||
tool_call = ToolCall(
|
||||
message_id=message_id,
|
||||
tool_name=tool_name,
|
||||
@ -231,25 +171,17 @@ class ConversationManager:
|
||||
logger.debug(f"Added tool call {tool_name} to message {message_id}")
|
||||
return tool_call
|
||||
|
||||
async def get_messages(self, conversation_id: int, limit: int | None = None, offset: int = 0) -> list[Message]:
|
||||
"""
|
||||
Get messages for a conversation
|
||||
|
||||
Args:
|
||||
conversation_id: Conversation ID
|
||||
limit: Maximum number of messages to return
|
||||
offset: Number of messages to skip
|
||||
|
||||
Returns:
|
||||
List of Message objects with preloaded tool_calls and feedbacks
|
||||
"""
|
||||
async def get_messages(
|
||||
self, conversation_id: int, limit: Optional[int] = None, offset: int = 0
|
||||
) -> list[Message]:
|
||||
"""获取对话的消息列表"""
|
||||
query = (
|
||||
select(Message)
|
||||
.options(
|
||||
selectinload(Message.tool_calls), # Preload tool calls
|
||||
selectinload(Message.feedbacks), # Preload feedbacks for UI state
|
||||
selectinload(Message.tool_calls),
|
||||
selectinload(Message.feedbacks),
|
||||
)
|
||||
.filter(Message.conversation_id == conversation_id)
|
||||
.where(Message.conversation_id == conversation_id)
|
||||
.order_by(Message.created_at.asc())
|
||||
)
|
||||
|
||||
@ -257,22 +189,12 @@ class ConversationManager:
|
||||
query = query.limit(limit).offset(offset)
|
||||
|
||||
result = await self.db.execute(query)
|
||||
return result.scalars().unique().all()
|
||||
return list(result.scalars().unique().all())
|
||||
|
||||
async def get_messages_by_thread_id(
|
||||
self, thread_id: str, limit: int | None = None, offset: int = 0
|
||||
self, thread_id: str, limit: Optional[int] = None, offset: int = 0
|
||||
) -> list[Message]:
|
||||
"""
|
||||
Get messages for a conversation by thread ID
|
||||
|
||||
Args:
|
||||
thread_id: Thread ID
|
||||
limit: Maximum number of messages to return
|
||||
offset: Number of messages to skip
|
||||
|
||||
Returns:
|
||||
List of Message objects
|
||||
"""
|
||||
"""根据 thread_id 获取对话的消息"""
|
||||
conversation = await self.get_conversation_by_thread_id(thread_id)
|
||||
if not conversation:
|
||||
logger.warning(f"Conversation not found for thread_id: {thread_id}")
|
||||
@ -281,51 +203,28 @@ class ConversationManager:
|
||||
return await self.get_messages(conversation.id, limit, offset)
|
||||
|
||||
async def list_conversations(
|
||||
self, user_id: str | None = None, agent_id: str | None = None, status: str = "active"
|
||||
self, user_id: Optional[str] = None, agent_id: Optional[str] = None, status: str = "active"
|
||||
) -> list[Conversation]:
|
||||
"""
|
||||
List conversations for a user or all users
|
||||
"""列出对话"""
|
||||
query = select(Conversation).where(Conversation.status == status)
|
||||
|
||||
Args:
|
||||
user_id: User ID (optional, if None or empty string, returns all users' conversations)
|
||||
agent_id: Optional agent ID filter
|
||||
status: Conversation status filter
|
||||
|
||||
Returns:
|
||||
List of Conversation objects
|
||||
"""
|
||||
query = select(Conversation).filter(Conversation.status == status)
|
||||
|
||||
# Only filter by user_id if it's provided and not empty
|
||||
if user_id:
|
||||
query = query.filter(Conversation.user_id == str(user_id))
|
||||
|
||||
query = query.where(Conversation.user_id == str(user_id))
|
||||
if agent_id:
|
||||
query = query.filter(Conversation.agent_id == agent_id)
|
||||
query = query.where(Conversation.agent_id == agent_id)
|
||||
|
||||
query = query.order_by(Conversation.updated_at.desc())
|
||||
result = await self.db.execute(query)
|
||||
return result.scalars().all()
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def update_conversation(
|
||||
self,
|
||||
thread_id: str,
|
||||
title: str | None = None,
|
||||
status: str | None = None,
|
||||
metadata: dict | None = None,
|
||||
) -> Conversation | None:
|
||||
"""
|
||||
Update conversation information
|
||||
|
||||
Args:
|
||||
thread_id: Thread ID
|
||||
title: New title
|
||||
status: New status
|
||||
metadata: Additional metadata to merge
|
||||
|
||||
Returns:
|
||||
Updated Conversation object or None if not found
|
||||
"""
|
||||
title: Optional[str] = None,
|
||||
status: Optional[str] = None,
|
||||
metadata: Optional[dict] = None,
|
||||
) -> Optional[Conversation]:
|
||||
"""更新对话信息"""
|
||||
conversation = await self.get_conversation_by_thread_id(thread_id)
|
||||
if not conversation:
|
||||
return None
|
||||
@ -335,7 +234,6 @@ class ConversationManager:
|
||||
if status is not None:
|
||||
conversation.status = status
|
||||
|
||||
# Handle metadata updates
|
||||
if metadata is not None:
|
||||
current_metadata = conversation.extra_metadata or {}
|
||||
current_metadata.update(metadata)
|
||||
@ -349,16 +247,7 @@ class ConversationManager:
|
||||
return conversation
|
||||
|
||||
async def delete_conversation(self, thread_id: str, soft_delete: bool = True) -> bool:
|
||||
"""
|
||||
Delete a conversation
|
||||
|
||||
Args:
|
||||
thread_id: Thread ID
|
||||
soft_delete: If True, mark as deleted; if False, permanently delete
|
||||
|
||||
Returns:
|
||||
True if successful, False otherwise
|
||||
"""
|
||||
"""删除对话"""
|
||||
conversation = await self.get_conversation_by_thread_id(thread_id)
|
||||
if not conversation:
|
||||
return False
|
||||
@ -374,40 +263,21 @@ class ConversationManager:
|
||||
|
||||
return True
|
||||
|
||||
async def get_stats(self, conversation_id: int) -> ConversationStats | None:
|
||||
"""
|
||||
Get conversation statistics
|
||||
|
||||
Args:
|
||||
conversation_id: Conversation ID
|
||||
|
||||
Returns:
|
||||
ConversationStats object or None if not found
|
||||
"""
|
||||
async def get_stats(self, conversation_id: int) -> Optional[ConversationStats]:
|
||||
"""获取对话统计"""
|
||||
result = await self.db.execute(
|
||||
select(ConversationStats).filter(ConversationStats.conversation_id == conversation_id)
|
||||
select(ConversationStats).where(ConversationStats.conversation_id == conversation_id)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def update_stats(
|
||||
self,
|
||||
conversation_id: int,
|
||||
tokens_used: int | None = None,
|
||||
model_used: str | None = None,
|
||||
user_feedback: dict | None = None,
|
||||
) -> ConversationStats | None:
|
||||
"""
|
||||
Update conversation statistics
|
||||
|
||||
Args:
|
||||
conversation_id: Conversation ID
|
||||
tokens_used: Number of tokens to add
|
||||
model_used: Model name
|
||||
user_feedback: User feedback data
|
||||
|
||||
Returns:
|
||||
Updated ConversationStats object or None if not found
|
||||
"""
|
||||
tokens_used: Optional[int] = None,
|
||||
model_used: Optional[str] = None,
|
||||
user_feedback: Optional[dict] = None,
|
||||
) -> Optional[ConversationStats]:
|
||||
"""更新对话统计"""
|
||||
stats = await self.get_stats(conversation_id)
|
||||
if not stats:
|
||||
return None
|
||||
@ -425,18 +295,10 @@ class ConversationManager:
|
||||
|
||||
return stats
|
||||
|
||||
async def get_tool_call_by_langgraph_id(self, langgraph_tool_call_id: str) -> ToolCall | None:
|
||||
"""
|
||||
Get tool call by LangGraph tool_call_id
|
||||
|
||||
Args:
|
||||
langgraph_tool_call_id: LangGraph tool_call_id
|
||||
|
||||
Returns:
|
||||
ToolCall object or None if not found
|
||||
"""
|
||||
async def get_tool_call_by_langgraph_id(self, langgraph_tool_call_id: str) -> Optional[ToolCall]:
|
||||
"""根据 LangGraph tool_call_id 获取工具调用"""
|
||||
result = await self.db.execute(
|
||||
select(ToolCall).filter(ToolCall.langgraph_tool_call_id == langgraph_tool_call_id)
|
||||
select(ToolCall).where(ToolCall.langgraph_tool_call_id == langgraph_tool_call_id)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
@ -445,20 +307,9 @@ class ConversationManager:
|
||||
langgraph_tool_call_id: str,
|
||||
tool_output: str,
|
||||
status: str = "success",
|
||||
error_message: str | None = None,
|
||||
) -> ToolCall | None:
|
||||
"""
|
||||
Update tool call output by LangGraph tool_call_id
|
||||
|
||||
Args:
|
||||
langgraph_tool_call_id: LangGraph tool_call_id
|
||||
tool_output: Tool execution result
|
||||
status: Status (success/error)
|
||||
error_message: Error message if failed
|
||||
|
||||
Returns:
|
||||
Updated ToolCall object or None if not found
|
||||
"""
|
||||
error_message: Optional[str] = None,
|
||||
) -> Optional[ToolCall]:
|
||||
"""根据 LangGraph tool_call_id 更新工具调用输出"""
|
||||
tool_call = await self.get_tool_call_by_langgraph_id(langgraph_tool_call_id)
|
||||
if not tool_call:
|
||||
logger.warning(f"Tool call not found for langgraph_tool_call_id: {langgraph_tool_call_id}")
|
||||
@ -476,26 +327,24 @@ class ConversationManager:
|
||||
return tool_call
|
||||
|
||||
async def _update_message_count(self, conversation_id: int) -> None:
|
||||
"""
|
||||
Update message count in conversation stats
|
||||
|
||||
Args:
|
||||
conversation_id: Conversation ID
|
||||
"""
|
||||
"""更新对话统计中的消息计数"""
|
||||
from sqlalchemy import func
|
||||
|
||||
stats = await self.get_stats(conversation_id)
|
||||
if stats:
|
||||
result = await self.db.execute(select(func.count()).filter(Message.conversation_id == conversation_id))
|
||||
result = await self.db.execute(
|
||||
select(func.count()).where(Message.conversation_id == conversation_id)
|
||||
)
|
||||
message_count = result.scalar()
|
||||
stats.message_count = message_count
|
||||
await self.db.commit()
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# Attachment helpers
|
||||
# 附件辅助方法
|
||||
# -------------------------------------------------------------------------
|
||||
|
||||
async def get_attachments(self, conversation_id: int) -> list[dict]:
|
||||
"""获取对话的附件列表"""
|
||||
conversation = await self._get_conversation_by_id(conversation_id)
|
||||
if not conversation:
|
||||
return []
|
||||
@ -503,12 +352,14 @@ class ConversationManager:
|
||||
return list(metadata.get("attachments", []))
|
||||
|
||||
async def get_attachments_by_thread_id(self, thread_id: str) -> list[dict]:
|
||||
"""根据 thread_id 获取附件列表"""
|
||||
conversation = await self.get_conversation_by_thread_id(thread_id)
|
||||
if not conversation:
|
||||
return []
|
||||
return await self.get_attachments(conversation.id)
|
||||
|
||||
async def add_attachment(self, conversation_id: int, attachment_info: dict) -> dict | None:
|
||||
async def add_attachment(self, conversation_id: int, attachment_info: dict) -> Optional[dict]:
|
||||
"""添加附件到对话"""
|
||||
conversation = await self._get_conversation_by_id(conversation_id)
|
||||
if not conversation:
|
||||
return None
|
||||
@ -522,8 +373,9 @@ class ConversationManager:
|
||||
return attachment_info
|
||||
|
||||
async def update_attachment_status(
|
||||
self, conversation_id: int, file_id: str, status: str, update_fields: dict | None = None
|
||||
) -> dict | None:
|
||||
self, conversation_id: int, file_id: str, status: str, update_fields: Optional[dict] = None
|
||||
) -> Optional[dict]:
|
||||
"""更新附件状态"""
|
||||
conversation = await self._get_conversation_by_id(conversation_id)
|
||||
if not conversation:
|
||||
return None
|
||||
@ -545,6 +397,7 @@ class ConversationManager:
|
||||
return target
|
||||
|
||||
async def remove_attachment(self, conversation_id: int, file_id: str) -> bool:
|
||||
"""从对话中移除附件"""
|
||||
conversation = await self._get_conversation_by_id(conversation_id)
|
||||
if not conversation:
|
||||
return False
|
||||
|
||||
@ -1,4 +1,4 @@
|
||||
"""PostgreSQL 数据库管理器 - 专门用于知识库数据"""
|
||||
"""PostgreSQL 数据库管理器 - 支持知识库和业务数据"""
|
||||
|
||||
import json
|
||||
import os
|
||||
@ -6,16 +6,26 @@ from contextlib import asynccontextmanager
|
||||
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
from sqlalchemy.orm import declarative_base
|
||||
|
||||
from server.utils.singleton import SingletonMeta
|
||||
from src.storage.db.models_knowledge import (
|
||||
Base,
|
||||
)
|
||||
from src.storage.db.models_knowledge import Base as KnowledgeBase
|
||||
from src.storage.postgres.models_business import Base as BusinessBase
|
||||
from src.utils import logger
|
||||
|
||||
# 合并两个 Base
|
||||
CombinedBase = declarative_base()
|
||||
|
||||
# 继承所有表
|
||||
for module in [KnowledgeBase, BusinessBase]:
|
||||
for table_name in dir(module):
|
||||
table = getattr(module, table_name)
|
||||
if isinstance(table, type) and hasattr(table, '__tablename__'):
|
||||
setattr(CombinedBase, table_name, table)
|
||||
|
||||
|
||||
class PostgresManager(metaclass=SingletonMeta):
|
||||
"""PostgreSQL 数据库管理器 - 专门用于知识库元数据"""
|
||||
"""PostgreSQL 数据库管理器 - 支持知识库和业务数据"""
|
||||
|
||||
# 知识库 PostgreSQL URL 环境变量名
|
||||
KB_DATABASE_URL_ENV = "YUXI_KNOWLEDGE_DATABASE_URL"
|
||||
@ -65,17 +75,26 @@ class PostgresManager(metaclass=SingletonMeta):
|
||||
raise RuntimeError("PostgreSQL manager not initialized. Please check configuration.")
|
||||
|
||||
async def create_tables(self):
|
||||
"""创建所有知识库相关表"""
|
||||
"""创建所有表(知识库和业务表)"""
|
||||
self._check_initialized()
|
||||
async with self.async_engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
logger.info("PostgreSQL tables created/checked")
|
||||
await conn.run_sync(KnowledgeBase.metadata.create_all)
|
||||
await conn.run_sync(BusinessBase.metadata.create_all)
|
||||
logger.info("PostgreSQL tables created/checked (knowledge + business)")
|
||||
|
||||
async def create_business_tables(self):
|
||||
"""创建所有业务数据表"""
|
||||
self._check_initialized()
|
||||
async with self.async_engine.begin() as conn:
|
||||
await conn.run_sync(BusinessBase.metadata.create_all)
|
||||
logger.info("PostgreSQL business tables created/checked")
|
||||
|
||||
async def drop_tables(self):
|
||||
"""删除所有知识库相关表(慎用!)"""
|
||||
"""删除所有表(慎用!)"""
|
||||
self._check_initialized()
|
||||
async with self.async_engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.drop_all)
|
||||
await conn.run_sync(BusinessBase.metadata.drop_all)
|
||||
await conn.run_sync(KnowledgeBase.metadata.drop_all)
|
||||
logger.info("PostgreSQL tables dropped")
|
||||
|
||||
async def ensure_knowledge_schema(self):
|
||||
@ -170,6 +189,36 @@ class PostgresManager(metaclass=SingletonMeta):
|
||||
if self.async_engine:
|
||||
await self.async_engine.dispose()
|
||||
|
||||
async def async_check_first_run(self):
|
||||
"""检查是否首次运行(异步版本)- 检查用户表是否有数据"""
|
||||
from sqlalchemy import func, select
|
||||
|
||||
self._check_initialized()
|
||||
async with self.get_async_session_context() as session:
|
||||
from src.storage.postgres.models_business import User
|
||||
|
||||
result = await session.execute(select(func.count(User.id)))
|
||||
count = result.scalar()
|
||||
return count == 0
|
||||
|
||||
async def execute(self, statement):
|
||||
"""直接执行 SQL 语句(用于迁移脚本)"""
|
||||
self._check_initialized()
|
||||
async with self.get_async_session_context() as session:
|
||||
return await session.execute(statement)
|
||||
|
||||
async def add(self, instance):
|
||||
"""添加实例到会话(用于迁移脚本)"""
|
||||
self._check_initialized()
|
||||
async with self.get_async_session_context() as session:
|
||||
session.add(instance)
|
||||
|
||||
async def commit(self):
|
||||
"""提交当前会话"""
|
||||
self._check_initialized()
|
||||
async with self.get_async_session_context() as session:
|
||||
pass # commit is automatic in context manager
|
||||
|
||||
|
||||
# 创建全局 PostgreSQL 管理器实例
|
||||
pg_manager = PostgresManager()
|
||||
|
||||
387
src/storage/postgres/models_business.py
Normal file
387
src/storage/postgres/models_business.py
Normal file
@ -0,0 +1,387 @@
|
||||
"""PostgreSQL 业务数据模型 - 用户、部门、对话等相关表"""
|
||||
|
||||
from datetime import datetime as dt, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import JSON, Boolean, Column, DateTime, ForeignKey, Integer, String, Text
|
||||
from sqlalchemy.ext.declarative import declarative_base
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
Base = declarative_base()
|
||||
|
||||
|
||||
def _utc_now_naive():
|
||||
"""返回 naive UTC datetime 以兼容 PostgreSQL TIMESTAMP WITHOUT TIME ZONE 列"""
|
||||
return dt.now(timezone.utc).replace(tzinfo=None)
|
||||
|
||||
|
||||
def _format_utc_datetime(dt_value) -> str | None:
|
||||
"""Helper to format datetime to UTC ISO string, assuming naive datetimes are UTC."""
|
||||
if dt_value is None:
|
||||
return None
|
||||
if isinstance(dt_value, dt):
|
||||
if dt_value.tzinfo is None:
|
||||
dt_value = dt_value.replace(tzinfo=timezone.utc)
|
||||
return dt_value.isoformat()
|
||||
return str(dt_value)
|
||||
|
||||
|
||||
class Department(Base):
|
||||
"""部门模型"""
|
||||
|
||||
__tablename__ = "departments"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
name = Column(String(50), nullable=False, unique=True, index=True)
|
||||
description = Column(String(255), nullable=True)
|
||||
created_at = Column(DateTime, default=_utc_now_naive)
|
||||
|
||||
# 关联关系
|
||||
users = relationship("User", back_populates="department", cascade="all, delete-orphan")
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
}
|
||||
|
||||
|
||||
class User(Base):
|
||||
"""用户模型"""
|
||||
|
||||
__tablename__ = "users"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
username = Column(String, nullable=False, unique=True, index=True) # 显示名称
|
||||
user_id = Column(String, nullable=False, unique=True, index=True) # 登录ID
|
||||
phone_number = Column(String, nullable=True, unique=True, index=True) # 手机号
|
||||
avatar = Column(String, nullable=True) # 头像URL
|
||||
password_hash = Column(String, nullable=False)
|
||||
role = Column(String, nullable=False, default="user") # 角色: superadmin, admin, user
|
||||
department_id = Column(Integer, ForeignKey("departments.id"), nullable=True) # 部门ID
|
||||
created_at = Column(DateTime, default=_utc_now_naive)
|
||||
last_login = Column(DateTime, nullable=True)
|
||||
|
||||
# 登录失败限制相关字段
|
||||
login_failed_count = Column(Integer, nullable=False, default=0) # 登录失败次数
|
||||
last_failed_login = Column(DateTime, nullable=True) # 最后一次登录失败时间
|
||||
login_locked_until = Column(DateTime, nullable=True) # 锁定到什么时候
|
||||
|
||||
# 软删除相关字段
|
||||
is_deleted = Column(Integer, nullable=False, default=0, index=True) # 是否已删除:0=否,1=是
|
||||
deleted_at = Column(DateTime, nullable=True) # 删除时间
|
||||
|
||||
# 关联操作日志
|
||||
operation_logs = relationship("OperationLog", back_populates="user", cascade="all, delete-orphan")
|
||||
|
||||
# 关联部门
|
||||
department = relationship("Department", back_populates="users")
|
||||
|
||||
def to_dict(self, include_password: bool = False) -> dict[str, Any]:
|
||||
result = {
|
||||
"id": self.id,
|
||||
"username": self.username,
|
||||
"user_id": self.user_id,
|
||||
"phone_number": self.phone_number,
|
||||
"avatar": self.avatar,
|
||||
"role": self.role,
|
||||
"department_id": self.department_id,
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
"last_login": _format_utc_datetime(self.last_login),
|
||||
"login_failed_count": self.login_failed_count,
|
||||
"last_failed_login": _format_utc_datetime(self.last_failed_login),
|
||||
"login_locked_until": _format_utc_datetime(self.login_locked_until),
|
||||
"is_deleted": self.is_deleted,
|
||||
"deleted_at": _format_utc_datetime(self.deleted_at),
|
||||
}
|
||||
if include_password:
|
||||
result["password_hash"] = self.password_hash
|
||||
return result
|
||||
|
||||
def is_login_locked(self) -> bool:
|
||||
"""检查用户是否处于登录锁定状态"""
|
||||
if self.login_locked_until is None:
|
||||
return False
|
||||
return _utc_now_naive() < self.login_locked_until
|
||||
|
||||
def get_remaining_lock_time(self) -> int:
|
||||
"""获取剩余锁定时间(秒)"""
|
||||
if self.login_locked_until is None:
|
||||
return 0
|
||||
remaining = int((self.login_locked_until - _utc_now_naive()).total_seconds())
|
||||
return max(0, remaining)
|
||||
|
||||
def reset_failed_login(self):
|
||||
"""重置登录失败相关字段"""
|
||||
self.login_failed_count = 0
|
||||
self.last_failed_login = None
|
||||
self.login_locked_until = None
|
||||
|
||||
|
||||
class Conversation(Base):
|
||||
"""Conversation table - 对话表"""
|
||||
|
||||
__tablename__ = "conversations"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True, comment="Primary key")
|
||||
thread_id = Column(String(64), unique=True, index=True, nullable=False, comment="Thread ID (UUID)")
|
||||
user_id = Column(String(64), index=True, nullable=False, comment="User ID")
|
||||
agent_id = Column(String(64), index=True, nullable=False, comment="Agent ID")
|
||||
title = Column(String(255), nullable=True, comment="Conversation title")
|
||||
status = Column(String(20), default="active", comment="Status: active/archived/deleted")
|
||||
created_at = Column(DateTime, default=_utc_now_naive, comment="Creation time")
|
||||
updated_at = Column(DateTime, default=_utc_now_naive, onupdate=_utc_now_naive, comment="Update time")
|
||||
extra_metadata = Column(JSON, nullable=True, comment="Additional metadata")
|
||||
|
||||
# Relationships
|
||||
messages = relationship("Message", back_populates="conversation", cascade="all, delete-orphan")
|
||||
stats = relationship("ConversationStats", back_populates="conversation", uselist=False, cascade="all, delete-orphan")
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"thread_id": self.thread_id,
|
||||
"user_id": self.user_id,
|
||||
"agent_id": self.agent_id,
|
||||
"title": self.title,
|
||||
"status": self.status,
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
"updated_at": _format_utc_datetime(self.updated_at),
|
||||
"metadata": self.extra_metadata or {},
|
||||
}
|
||||
|
||||
|
||||
class Message(Base):
|
||||
"""Message table - 消息表"""
|
||||
|
||||
__tablename__ = "messages"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True, comment="Primary key")
|
||||
conversation_id = Column(
|
||||
Integer, ForeignKey("conversations.id"), nullable=False, index=True, comment="Conversation ID"
|
||||
)
|
||||
role = Column(String(20), nullable=False, comment="Message role: user/assistant/system/tool")
|
||||
content = Column(Text, nullable=False, comment="Message content")
|
||||
message_type = Column(String(30), default="text", comment="Message type: text/tool_call/tool_result")
|
||||
created_at = Column(DateTime, default=_utc_now_naive, comment="Creation time")
|
||||
token_count = Column(Integer, nullable=True, comment="Token count (optional)")
|
||||
extra_metadata = Column(JSON, nullable=True, comment="Additional metadata (complete message dump)")
|
||||
image_content = Column(Text, nullable=True, comment="Base64 encoded image content for multimodal messages")
|
||||
|
||||
# Relationships
|
||||
conversation = relationship("Conversation", back_populates="messages")
|
||||
tool_calls = relationship("ToolCall", back_populates="message", cascade="all, delete-orphan")
|
||||
feedbacks = relationship("MessageFeedback", back_populates="message", cascade="all, delete-orphan")
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"conversation_id": self.conversation_id,
|
||||
"role": self.role,
|
||||
"content": self.content,
|
||||
"message_type": self.message_type,
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
"token_count": self.token_count,
|
||||
"metadata": self.extra_metadata or {},
|
||||
"image_content": self.image_content,
|
||||
"tool_calls": [tc.to_dict() for tc in self.tool_calls] if self.tool_calls else [],
|
||||
}
|
||||
|
||||
def to_simple_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"role": self.role,
|
||||
"content": self.content,
|
||||
}
|
||||
|
||||
|
||||
class ToolCall(Base):
|
||||
"""ToolCall table - 工具调用表"""
|
||||
|
||||
__tablename__ = "tool_calls"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True, comment="Primary key")
|
||||
message_id = Column(Integer, ForeignKey("messages.id"), nullable=False, index=True, comment="Message ID")
|
||||
langgraph_tool_call_id = Column(String(100), nullable=True, index=True, comment="LangGraph tool_call_id")
|
||||
tool_name = Column(String(100), nullable=False, comment="Tool name")
|
||||
tool_input = Column(JSON, nullable=True, comment="Tool input parameters")
|
||||
tool_output = Column(Text, nullable=True, comment="Tool execution result")
|
||||
status = Column(String(20), default="pending", comment="Status: pending/success/error")
|
||||
error_message = Column(Text, nullable=True, comment="Error message if failed")
|
||||
created_at = Column(DateTime, default=_utc_now_naive, comment="Creation time")
|
||||
|
||||
# Relationships
|
||||
message = relationship("Message", back_populates="tool_calls")
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"message_id": self.message_id,
|
||||
"langgraph_tool_call_id": self.langgraph_tool_call_id,
|
||||
"tool_name": self.tool_name,
|
||||
"tool_input": self.tool_input or {},
|
||||
"tool_output": self.tool_output,
|
||||
"status": self.status,
|
||||
"error_message": self.error_message,
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
}
|
||||
|
||||
|
||||
class ConversationStats(Base):
|
||||
"""ConversationStats table - 对话统计表"""
|
||||
|
||||
__tablename__ = "conversation_stats"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True, comment="Primary key")
|
||||
conversation_id = Column(
|
||||
Integer, ForeignKey("conversations.id"), unique=True, nullable=False, comment="Conversation ID"
|
||||
)
|
||||
message_count = Column(Integer, default=0, comment="Total message count")
|
||||
total_tokens = Column(Integer, default=0, comment="Total tokens used")
|
||||
model_used = Column(String(100), nullable=True, comment="Model used")
|
||||
user_feedback = Column(JSON, nullable=True, comment="User feedback")
|
||||
created_at = Column(DateTime, default=_utc_now_naive, comment="Creation time")
|
||||
updated_at = Column(DateTime, default=_utc_now_naive, onupdate=_utc_now_naive, comment="Update time")
|
||||
|
||||
# Relationships
|
||||
conversation = relationship("Conversation", back_populates="stats")
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"conversation_id": self.conversation_id,
|
||||
"message_count": self.message_count,
|
||||
"total_tokens": self.total_tokens,
|
||||
"model_used": self.model_used,
|
||||
"user_feedback": self.user_feedback or {},
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
"updated_at": _format_utc_datetime(self.updated_at),
|
||||
}
|
||||
|
||||
|
||||
class OperationLog(Base):
|
||||
"""操作日志模型"""
|
||||
|
||||
__tablename__ = "operation_logs"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
user_id = Column(Integer, ForeignKey("users.id"), nullable=False)
|
||||
operation = Column(String, nullable=False)
|
||||
details = Column(Text, nullable=True)
|
||||
ip_address = Column(String, nullable=True)
|
||||
timestamp = Column(DateTime, default=_utc_now_naive)
|
||||
|
||||
# 关联用户
|
||||
user = relationship("User", back_populates="operation_logs")
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"user_id": self.user_id,
|
||||
"operation": self.operation,
|
||||
"details": self.details,
|
||||
"ip_address": self.ip_address,
|
||||
"timestamp": _format_utc_datetime(self.timestamp),
|
||||
}
|
||||
|
||||
|
||||
class MessageFeedback(Base):
|
||||
"""Message feedback table - 消息反馈表"""
|
||||
|
||||
__tablename__ = "message_feedbacks"
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True, comment="Primary key")
|
||||
message_id = Column(Integer, ForeignKey("messages.id"), nullable=False, index=True, comment="Message ID being rated")
|
||||
user_id = Column(String(64), nullable=False, index=True, comment="User ID who provided feedback")
|
||||
rating = Column(String(10), nullable=False, comment="Feedback rating: like or dislike")
|
||||
reason = Column(Text, nullable=True, comment="Optional reason for dislike feedback")
|
||||
created_at = Column(DateTime, default=_utc_now_naive, comment="Feedback creation time")
|
||||
|
||||
# Relationships
|
||||
message = relationship("Message", back_populates="feedbacks")
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"message_id": self.message_id,
|
||||
"user_id": self.user_id,
|
||||
"rating": self.rating,
|
||||
"reason": self.reason,
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
}
|
||||
|
||||
|
||||
class MCPServer(Base):
|
||||
"""MCP 服务器配置模型"""
|
||||
|
||||
__tablename__ = "mcp_servers"
|
||||
|
||||
# 核心字段 - name 作为主键
|
||||
name = Column(String(100), primary_key=True, comment="服务器名称(唯一标识)")
|
||||
description = Column(String(500), nullable=True, comment="描述")
|
||||
|
||||
# 连接配置
|
||||
transport = Column(String(20), nullable=False, comment="传输类型:sse/streamable_http/stdio")
|
||||
url = Column(String(500), nullable=True, comment="服务器 URL(sse/streamable_http)")
|
||||
command = Column(String(500), nullable=True, comment="命令(stdio)")
|
||||
args = Column(JSON, nullable=True, comment="命令参数数组(stdio)")
|
||||
headers = Column(JSON, nullable=True, comment="HTTP 请求头")
|
||||
timeout = Column(Integer, nullable=True, comment="HTTP 超时时间(秒)")
|
||||
sse_read_timeout = Column(Integer, nullable=True, comment="SSE 读取超时(秒)")
|
||||
|
||||
# UI 增强字段
|
||||
tags = Column(JSON, nullable=True, comment="标签数组")
|
||||
icon = Column(String(50), nullable=True, comment="图标(emoji)")
|
||||
|
||||
# 状态字段
|
||||
enabled = Column(Integer, nullable=False, default=1, comment="是否启用:1=是,0=否")
|
||||
disabled_tools = Column(JSON, nullable=True, comment="禁用的工具名称列表")
|
||||
|
||||
# 用户追踪
|
||||
created_by = Column(String(100), nullable=False, comment="创建人用户名")
|
||||
updated_by = Column(String(100), nullable=False, comment="修改人用户名")
|
||||
|
||||
# 时间戳
|
||||
created_at = Column(DateTime, default=_utc_now_naive, comment="创建时间")
|
||||
updated_at = Column(DateTime, default=_utc_now_naive, onupdate=_utc_now_naive, comment="更新时间")
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"transport": self.transport,
|
||||
"url": self.url,
|
||||
"command": self.command,
|
||||
"args": self.args or [],
|
||||
"headers": self.headers or {},
|
||||
"timeout": self.timeout,
|
||||
"sse_read_timeout": self.sse_read_timeout,
|
||||
"tags": self.tags or [],
|
||||
"icon": self.icon,
|
||||
"enabled": bool(self.enabled),
|
||||
"disabled_tools": self.disabled_tools or [],
|
||||
"created_by": self.created_by,
|
||||
"updated_by": self.updated_by,
|
||||
"created_at": _format_utc_datetime(self.created_at),
|
||||
"updated_at": _format_utc_datetime(self.updated_at),
|
||||
}
|
||||
|
||||
def to_mcp_config(self) -> dict[str, Any]:
|
||||
"""转换为 MCP 配置格式(用于加载到 MCP_SERVERS 缓存)"""
|
||||
config = {"transport": self.transport}
|
||||
if self.url:
|
||||
config["url"] = self.url
|
||||
if self.command:
|
||||
config["command"] = self.command
|
||||
if self.args:
|
||||
config["args"] = self.args
|
||||
if self.headers:
|
||||
config["headers"] = self.headers
|
||||
if self.timeout is not None:
|
||||
config["timeout"] = self.timeout
|
||||
if self.sse_read_timeout is not None:
|
||||
config["sse_read_timeout"] = self.sse_read_timeout
|
||||
if self.disabled_tools:
|
||||
config["disabled_tools"] = self.disabled_tools
|
||||
return config
|
||||
@ -23,3 +23,36 @@ async def test_admin_can_fetch_conversations(test_client, admin_headers):
|
||||
response = await test_client.get("/api/dashboard/conversations", headers=admin_headers)
|
||||
assert response.status_code == 200, response.text
|
||||
assert isinstance(response.json(), list)
|
||||
|
||||
|
||||
async def test_admin_can_fetch_stats(test_client, admin_headers):
|
||||
"""Test that all stats endpoints return 200 and don't crash on DB queries."""
|
||||
|
||||
# Test call timeseries stats for all types
|
||||
types = ["models", "agents", "tokens", "tools"]
|
||||
for stats_type in types:
|
||||
response = await test_client.get(
|
||||
f"/api/dashboard/stats/calls/timeseries?type={stats_type}&time_range=14days",
|
||||
headers=admin_headers
|
||||
)
|
||||
assert response.status_code == 200, f"{stats_type} stats failed: {response.text}"
|
||||
data = response.json()
|
||||
assert "data" in data
|
||||
assert "categories" in data
|
||||
|
||||
# Test user activity stats
|
||||
response = await test_client.get("/api/dashboard/stats/users", headers=admin_headers)
|
||||
assert response.status_code == 200, f"user stats failed: {response.text}"
|
||||
assert "total_users" in response.json()
|
||||
|
||||
# Test tool call stats
|
||||
response = await test_client.get("/api/dashboard/stats/tools", headers=admin_headers)
|
||||
assert response.status_code == 200, f"tool stats failed: {response.text}"
|
||||
assert "total_calls" in response.json()
|
||||
|
||||
|
||||
async def test_admin_can_fetch_feedbacks(test_client, admin_headers):
|
||||
"""Test that feedback endpoint returns 200 and handles the User join correctly."""
|
||||
response = await test_client.get("/api/dashboard/feedbacks", headers=admin_headers)
|
||||
assert response.status_code == 200, f"feedbacks failed: {response.text}"
|
||||
assert isinstance(response.json(), list)
|
||||
|
||||
30
test/api/test_graph_router_list.py
Normal file
30
test/api/test_graph_router_list.py
Normal file
@ -0,0 +1,30 @@
|
||||
"""
|
||||
Integration tests for graph router list endpoint.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
pytestmark = [pytest.mark.asyncio, pytest.mark.integration]
|
||||
|
||||
|
||||
async def test_admin_can_list_graphs(test_client, admin_headers):
|
||||
"""Test that listing graphs returns 200 and a list of graphs."""
|
||||
response = await test_client.get("/api/graph/list", headers=admin_headers)
|
||||
assert response.status_code == 200, f"Failed to list graphs: {response.text}"
|
||||
data = response.json()
|
||||
|
||||
# Check if response is wrapped
|
||||
if isinstance(data, dict) and "data" in data:
|
||||
graphs = data["data"]
|
||||
else:
|
||||
graphs = data
|
||||
|
||||
assert isinstance(graphs, list)
|
||||
# Check structure of returned items if list is not empty
|
||||
if graphs:
|
||||
item = graphs[0]
|
||||
assert "id" in item
|
||||
assert "name" in item
|
||||
assert "type" in item
|
||||
Loading…
Reference in New Issue
Block a user