refactor: 重构数据库访问层,用户/对话/消息都统一迁移至PostgreSQL并添加业务模型支持

This commit is contained in:
Wenjie Zhang 2026-01-21 19:15:52 +08:00
parent a1812f2e97
commit 3da0d57c1c
21 changed files with 1992 additions and 395 deletions

View File

@ -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"]

View 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())

View File

@ -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()

View File

@ -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" # 禁止登录

View File

@ -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 + INTERVALSQLite 使用 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)
)

View File

@ -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(

View File

@ -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

View File

@ -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

View File

@ -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():

View File

@ -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)

View 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

View 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

View 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

View 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

View 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

View 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

View File

@ -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

View File

@ -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()

View 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="服务器 URLsse/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

View File

@ -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)

View 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