From 3da0d57c1cc9b8922c2909478287b13c748f13a5 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Wed, 21 Jan 2026 19:15:52 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E9=87=8D=E6=9E=84=E6=95=B0?= =?UTF-8?q?=E6=8D=AE=E5=BA=93=E8=AE=BF=E9=97=AE=E5=B1=82=EF=BC=8C=E7=94=A8?= =?UTF-8?q?=E6=88=B7/=E5=AF=B9=E8=AF=9D/=E6=B6=88=E6=81=AF=E9=83=BD?= =?UTF-8?q?=E7=BB=9F=E4=B8=80=E8=BF=81=E7=A7=BB=E8=87=B3PostgreSQL?= =?UTF-8?q?=E5=B9=B6=E6=B7=BB=E5=8A=A0=E4=B8=9A=E5=8A=A1=E6=A8=A1=E5=9E=8B?= =?UTF-8?q?=E6=94=AF=E6=8C=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docker-compose.yml | 2 +- scripts/migrate_business_from_sqlite.py | 601 ++++++++++++++++++ scripts/migrate_kb_metadata_to_db.py | 7 + server/routers/auth_router.py | 116 ++-- server/routers/dashboard_router.py | 92 +-- server/routers/department_router.py | 67 +- server/routers/graph_router.py | 6 +- server/utils/auth_middleware.py | 8 +- server/utils/common_utils.py | 2 +- src/knowledge/adapters/factory.py | 12 +- src/repositories/conversation_repository.py | 227 +++++++ src/repositories/department_repository.py | 95 +++ src/repositories/mcp_server_repository.py | 81 +++ .../message_feedback_repository.py | 44 ++ src/repositories/operation_log_repository.py | 49 ++ src/repositories/user_repository.py | 146 +++++ src/storage/conversation/manager.py | 313 +++------ src/storage/postgres/manager.py | 69 +- src/storage/postgres/models_business.py | 387 +++++++++++ test/api/test_dashboard_router.py | 33 + test/api/test_graph_router_list.py | 30 + 21 files changed, 1992 insertions(+), 395 deletions(-) create mode 100644 scripts/migrate_business_from_sqlite.py create mode 100644 src/repositories/conversation_repository.py create mode 100644 src/repositories/department_repository.py create mode 100644 src/repositories/mcp_server_repository.py create mode 100644 src/repositories/message_feedback_repository.py create mode 100644 src/repositories/operation_log_repository.py create mode 100644 src/repositories/user_repository.py create mode 100644 src/storage/postgres/models_business.py create mode 100644 test/api/test_graph_router_list.py diff --git a/docker-compose.yml b/docker-compose.yml index 9b67cbe0..d0b2dc45 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -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"] diff --git a/scripts/migrate_business_from_sqlite.py b/scripts/migrate_business_from_sqlite.py new file mode 100644 index 00000000..84c64291 --- /dev/null +++ b/scripts/migrate_business_from_sqlite.py @@ -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()) diff --git a/scripts/migrate_kb_metadata_to_db.py b/scripts/migrate_kb_metadata_to_db.py index 7cc14f9a..0498fd71 100644 --- a/scripts/migrate_kb_metadata_to_db.py +++ b/scripts/migrate_kb_metadata_to_db.py @@ -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() diff --git a/server/routers/auth_router.py b/server/routers/auth_router.py index c200962c..6c3a200a 100644 --- a/server/routers/auth_router.py +++ b/server/routers/auth_router.py @@ -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" # 禁止登录 diff --git a/server/routers/dashboard_router.py b/server/routers/dashboard_router.py index bb25b3c4..98762845 100644 --- a/server/routers/dashboard_router.py +++ b/server/routers/dashboard_router.py @@ -8,16 +8,18 @@ Provides centralized dashboard APIs for monitoring system-wide statistics. import traceback from datetime import datetime, timedelta +from typing import Any from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel -from sqlalchemy import String, cast, distinct, func, or_, select +from sqlalchemy import Integer, String, cast, distinct, func, or_, select, text from sqlalchemy.ext.asyncio import AsyncSession from server.routers.auth_router import get_admin_user from server.utils.auth_middleware import get_db from src.storage.conversation import ConversationManager from src.storage.db.models import User +from src.storage.postgres.manager import pg_manager from src.utils.datetime_utils import UTC, ensure_shanghai, shanghai_now, utc_now from src.utils.logging_config import logger @@ -25,6 +27,25 @@ from src.utils.logging_config import logger dashboard = APIRouter(prefix="/dashboard", tags=["Dashboard"]) +def _get_time_group_format(column, time_range: str) -> Any: + """ + 根据数据库类型生成时间分组格式化表达式。 + PostgreSQL 使用 to_char + INTERVAL,SQLite 使用 datetime + strftime。 + """ + # 检查是否是 PostgreSQL(通过检测 engine 或使用方言) + # 这里直接使用 PostgreSQL 语法,因为所有业务数据现在都在 PostgreSQL 上 + if time_range == "14hours": + # 每小时: YYYY-MM-DD HH:00 + time_expr = func.to_char(column + text("INTERVAL '8 hours'"), "YYYY-MM-DD HH24:00") + elif time_range == "14weeks": + # 每周: YYYY-WW + time_expr = func.to_char(column + text("INTERVAL '8 hours'"), "YYYY-IW") + else: # 14days + # 每天: YYYY-MM-DD + time_expr = func.to_char(column + text("INTERVAL '8 hours'"), "YYYY-MM-DD") + return time_expr + + # ============================================================================= # Response Models # ============================================================================= @@ -236,6 +257,8 @@ async def get_user_activity_stats( from src.storage.db.models import User, Conversation now = utc_now() + # PostgreSQL with asyncpg requires naive datetime for naive DateTime columns + naive_now = now.replace(tzinfo=None) # Conversations may store either the numeric user primary key or the login user_id string. # Join condition accounts for both representations. @@ -253,7 +276,7 @@ async def get_user_activity_stats( select(func.count(distinct(User.id))) .select_from(Conversation) .join(User, user_join_condition) - .filter(Conversation.updated_at >= now - timedelta(days=1), User.is_deleted == 0) + .filter(Conversation.updated_at >= naive_now - timedelta(days=1), User.is_deleted == 0) ) active_users_24h = active_users_24h_result.scalar() or 0 @@ -261,14 +284,14 @@ async def get_user_activity_stats( select(func.count(distinct(User.id))) .select_from(Conversation) .join(User, user_join_condition) - .filter(Conversation.updated_at >= now - timedelta(days=30), User.is_deleted == 0) + .filter(Conversation.updated_at >= naive_now - timedelta(days=30), User.is_deleted == 0) ) active_users_30d = active_users_30d_result.scalar() or 0 # 最近7天每日活跃用户(排除已删除用户) daily_active_users = [] for i in range(7): - day_start = now - timedelta(days=i + 1) - day_end = now - timedelta(days=i) + day_start = naive_now - timedelta(days=i + 1) + day_end = naive_now - timedelta(days=i) active_count_result = await db.execute( select(func.count(distinct(User.id))) @@ -308,6 +331,8 @@ async def get_tool_call_stats( from src.storage.db.models import ToolCall now = utc_now() + # PostgreSQL with asyncpg requires naive datetime for naive DateTime columns + naive_now = now.replace(tzinfo=None) # 基础工具调用统计 total_calls_result = await db.execute(select(func.count(ToolCall.id))) @@ -340,8 +365,8 @@ async def get_tool_call_stats( # 最近7天每日工具调用数 daily_tool_calls = [] for i in range(7): - day_start = now - timedelta(days=i + 1) - day_end = now - timedelta(days=i) + day_start = naive_now - timedelta(days=i + 1) + day_end = naive_now - timedelta(days=i) daily_count_result = await db.execute( select(func.count(ToolCall.id)).filter(ToolCall.created_at >= day_start, ToolCall.created_at < day_end) @@ -647,7 +672,7 @@ async def get_all_feedbacks( .join(Conversation, Message.conversation_id == Conversation.id) .outerjoin( User, - (MessageFeedback.user_id == User.id) | (MessageFeedback.user_id == User.user_id), + (MessageFeedback.user_id == cast(User.id, String)) | (MessageFeedback.user_id == User.user_id), ) ) @@ -724,7 +749,7 @@ async def get_call_timeseries_stats( intervals = 14 # 包含当前小时:从13小时前开始 start_time = now - timedelta(hours=intervals - 1) - group_format = func.strftime("%Y-%m-%d %H:00", func.datetime(Message.created_at, "+8 hours")) + group_format = _get_time_group_format(Message.created_at, time_range) base_local_time = ensure_shanghai(start_time) elif time_range == "14weeks": intervals = 14 @@ -733,40 +758,40 @@ async def get_call_timeseries_stats( local_start = local_start - timedelta(days=local_start.weekday()) local_start = local_start.replace(hour=0, minute=0, second=0, microsecond=0) start_time = local_start.astimezone(UTC) - group_format = func.strftime("%Y-%W", func.datetime(Message.created_at, "+8 hours")) + group_format = _get_time_group_format(Message.created_at, time_range) base_local_time = local_start else: # 14days (default) intervals = 14 # 包含当前天:从13天前开始 start_time = now - timedelta(days=intervals - 1) - group_format = func.strftime("%Y-%m-%d", func.datetime(Message.created_at, "+8 hours")) + group_format = _get_time_group_format(Message.created_at, time_range) base_local_time = ensure_shanghai(start_time) + # Convert start_time to naive UTC datetime for PostgreSQL query + # PostgreSQL with asyncpg and naive DateTime columns requires naive datetime objects + query_start_time = start_time.replace(tzinfo=None) + # 根据类型查询数据 if type == "models": # 模型调用统计(基于消息数量,按模型分组) # 从message的extra_metadata中提取模型信息 + category_expr = cast(Message.extra_metadata["response_metadata"]["model_name"], String) query_result = await db.execute( select( group_format.label("date"), func.count(Message.id).label("count"), - func.json_extract(Message.extra_metadata, "$.response_metadata.model_name").label("category"), + category_expr.label("category"), ) - .filter(Message.role == "assistant", Message.created_at >= start_time) + .filter(Message.role == "assistant", Message.created_at >= query_start_time) .filter(Message.extra_metadata.isnot(None)) - .group_by(group_format, func.json_extract(Message.extra_metadata, "$.response_metadata.model_name")) + .group_by(group_format, category_expr) .order_by(group_format) ) query = query_result.all() elif type == "agents": # 智能体调用统计(基于对话更新时间,按智能体分组) - # 为对话创建独立的时间格式化器 - if time_range == "14hours": - conv_group_format = func.strftime("%Y-%m-%d %H:00", func.datetime(Conversation.updated_at, "+8 hours")) - elif time_range == "14weeks": - conv_group_format = func.strftime("%Y-%W", func.datetime(Conversation.updated_at, "+8 hours")) - else: # 14days - conv_group_format = func.strftime("%Y-%m-%d", func.datetime(Conversation.updated_at, "+8 hours")) + # 为对话创建独立的时间格式化器(使用 PostgreSQL 兼容的 to_char + INTERVAL) + conv_group_format = _get_time_group_format(Conversation.updated_at, time_range) query_result = await db.execute( select( @@ -775,7 +800,7 @@ async def get_call_timeseries_stats( Conversation.agent_id.label("category"), ) .filter(Conversation.updated_at.isnot(None)) - .filter(Conversation.updated_at >= start_time) + .filter(Conversation.updated_at >= query_start_time) .group_by(conv_group_format, Conversation.agent_id) .order_by(conv_group_format) ) @@ -789,14 +814,14 @@ async def get_call_timeseries_stats( select( group_format.label("date"), func.sum( - func.coalesce(func.json_extract(Message.extra_metadata, "$.usage_metadata.input_tokens"), 0) + func.coalesce(cast(cast(Message.extra_metadata["usage_metadata"]["input_tokens"], String), Integer), 0) ).label("count"), literal("input_tokens").label("category"), ) .filter( - Message.created_at >= start_time, + Message.created_at >= query_start_time, Message.extra_metadata.isnot(None), - func.json_extract(Message.extra_metadata, "$.usage_metadata").isnot(None), + Message.extra_metadata["usage_metadata"].isnot(None), ) .group_by(group_format) .order_by(group_format) @@ -808,14 +833,14 @@ async def get_call_timeseries_stats( select( group_format.label("date"), func.sum( - func.coalesce(func.json_extract(Message.extra_metadata, "$.usage_metadata.output_tokens"), 0) + func.coalesce(cast(cast(Message.extra_metadata["usage_metadata"]["output_tokens"], String), Integer), 0) ).label("count"), literal("output_tokens").label("category"), ) .filter( - Message.created_at >= start_time, + Message.created_at >= query_start_time, Message.extra_metadata.isnot(None), - func.json_extract(Message.extra_metadata, "$.usage_metadata").isnot(None), + Message.extra_metadata["usage_metadata"].isnot(None), ) .group_by(group_format) .order_by(group_format) @@ -828,13 +853,8 @@ async def get_call_timeseries_stats( results = input_results + output_results elif type == "tools": # 工具调用统计(按工具名称分组) - # 为工具调用创建独立的时间格式化器 - if time_range == "14hours": - tool_group_format = func.strftime("%Y-%m-%d %H:00", func.datetime(ToolCall.created_at, "+8 hours")) - elif time_range == "14weeks": - tool_group_format = func.strftime("%Y-%W", func.datetime(ToolCall.created_at, "+8 hours")) - else: # 14days - tool_group_format = func.strftime("%Y-%m-%d", func.datetime(ToolCall.created_at, "+8 hours")) + # 为工具调用创建独立的时间格式化器(使用 PostgreSQL 兼容的 to_char + INTERVAL) + tool_group_format = _get_time_group_format(ToolCall.created_at, time_range) query_result = await db.execute( select( @@ -842,7 +862,7 @@ async def get_call_timeseries_stats( func.count(ToolCall.id).label("count"), ToolCall.tool_name.label("category"), ) - .filter(ToolCall.created_at >= start_time) + .filter(ToolCall.created_at >= query_start_time) .group_by(tool_group_format, ToolCall.tool_name) .order_by(tool_group_format) ) diff --git a/server/routers/department_router.py b/server/routers/department_router.py index c4eca3b7..9424c177 100644 --- a/server/routers/department_router.py +++ b/server/routers/department_router.py @@ -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( diff --git a/server/routers/graph_router.py b/server/routers/graph_router.py index 7cca6ea4..4496a752 100644 --- a/server/routers/graph_router.py +++ b/server/routers/graph_router.py @@ -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 diff --git a/server/utils/auth_middleware.py b/server/utils/auth_middleware.py index 2be55704..04a2bd31 100644 --- a/server/utils/auth_middleware.py +++ b/server/utils/auth_middleware.py @@ -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 diff --git a/server/utils/common_utils.py b/server/utils/common_utils.py index 02d68576..0adc252b 100644 --- a/server/utils/common_utils.py +++ b/server/utils/common_utils.py @@ -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(): diff --git a/src/knowledge/adapters/factory.py b/src/knowledge/adapters/factory.py index 06c98eab..4804d18a 100644 --- a/src/knowledge/adapters/factory.py +++ b/src/knowledge/adapters/factory.py @@ -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) diff --git a/src/repositories/conversation_repository.py b/src/repositories/conversation_repository.py new file mode 100644 index 00000000..d82ce913 --- /dev/null +++ b/src/repositories/conversation_repository.py @@ -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 diff --git a/src/repositories/department_repository.py b/src/repositories/department_repository.py new file mode 100644 index 00000000..df8ff80c --- /dev/null +++ b/src/repositories/department_repository.py @@ -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 diff --git a/src/repositories/mcp_server_repository.py b/src/repositories/mcp_server_repository.py new file mode 100644 index 00000000..04feae11 --- /dev/null +++ b/src/repositories/mcp_server_repository.py @@ -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 diff --git a/src/repositories/message_feedback_repository.py b/src/repositories/message_feedback_repository.py new file mode 100644 index 00000000..9e9d68ad --- /dev/null +++ b/src/repositories/message_feedback_repository.py @@ -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 diff --git a/src/repositories/operation_log_repository.py b/src/repositories/operation_log_repository.py new file mode 100644 index 00000000..f365763a --- /dev/null +++ b/src/repositories/operation_log_repository.py @@ -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 diff --git a/src/repositories/user_repository.py b/src/repositories/user_repository.py new file mode 100644 index 00000000..b820eee2 --- /dev/null +++ b/src/repositories/user_repository.py @@ -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 diff --git a/src/storage/conversation/manager.py b/src/storage/conversation/manager.py index a3e0e30f..207a4422 100644 --- a/src/storage/conversation/manager.py +++ b/src/storage/conversation/manager.py @@ -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 diff --git a/src/storage/postgres/manager.py b/src/storage/postgres/manager.py index e3245641..f6b137b3 100644 --- a/src/storage/postgres/manager.py +++ b/src/storage/postgres/manager.py @@ -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() diff --git a/src/storage/postgres/models_business.py b/src/storage/postgres/models_business.py new file mode 100644 index 00000000..031d8539 --- /dev/null +++ b/src/storage/postgres/models_business.py @@ -0,0 +1,387 @@ +"""PostgreSQL 业务数据模型 - 用户、部门、对话等相关表""" + +from datetime import datetime as dt, timezone +from typing import Any + +from sqlalchemy import JSON, Boolean, Column, DateTime, ForeignKey, Integer, String, Text +from sqlalchemy.ext.declarative import declarative_base +from sqlalchemy.orm import relationship + +Base = declarative_base() + + +def _utc_now_naive(): + """返回 naive UTC datetime 以兼容 PostgreSQL TIMESTAMP WITHOUT TIME ZONE 列""" + return dt.now(timezone.utc).replace(tzinfo=None) + + +def _format_utc_datetime(dt_value) -> str | None: + """Helper to format datetime to UTC ISO string, assuming naive datetimes are UTC.""" + if dt_value is None: + return None + if isinstance(dt_value, dt): + if dt_value.tzinfo is None: + dt_value = dt_value.replace(tzinfo=timezone.utc) + return dt_value.isoformat() + return str(dt_value) + + +class Department(Base): + """部门模型""" + + __tablename__ = "departments" + + id = Column(Integer, primary_key=True, autoincrement=True) + name = Column(String(50), nullable=False, unique=True, index=True) + description = Column(String(255), nullable=True) + created_at = Column(DateTime, default=_utc_now_naive) + + # 关联关系 + users = relationship("User", back_populates="department", cascade="all, delete-orphan") + + def to_dict(self) -> dict[str, Any]: + return { + "id": self.id, + "name": self.name, + "description": self.description, + "created_at": _format_utc_datetime(self.created_at), + } + + +class User(Base): + """用户模型""" + + __tablename__ = "users" + + id = Column(Integer, primary_key=True, autoincrement=True) + username = Column(String, nullable=False, unique=True, index=True) # 显示名称 + user_id = Column(String, nullable=False, unique=True, index=True) # 登录ID + phone_number = Column(String, nullable=True, unique=True, index=True) # 手机号 + avatar = Column(String, nullable=True) # 头像URL + password_hash = Column(String, nullable=False) + role = Column(String, nullable=False, default="user") # 角色: superadmin, admin, user + department_id = Column(Integer, ForeignKey("departments.id"), nullable=True) # 部门ID + created_at = Column(DateTime, default=_utc_now_naive) + last_login = Column(DateTime, nullable=True) + + # 登录失败限制相关字段 + login_failed_count = Column(Integer, nullable=False, default=0) # 登录失败次数 + last_failed_login = Column(DateTime, nullable=True) # 最后一次登录失败时间 + login_locked_until = Column(DateTime, nullable=True) # 锁定到什么时候 + + # 软删除相关字段 + is_deleted = Column(Integer, nullable=False, default=0, index=True) # 是否已删除:0=否,1=是 + deleted_at = Column(DateTime, nullable=True) # 删除时间 + + # 关联操作日志 + operation_logs = relationship("OperationLog", back_populates="user", cascade="all, delete-orphan") + + # 关联部门 + department = relationship("Department", back_populates="users") + + def to_dict(self, include_password: bool = False) -> dict[str, Any]: + result = { + "id": self.id, + "username": self.username, + "user_id": self.user_id, + "phone_number": self.phone_number, + "avatar": self.avatar, + "role": self.role, + "department_id": self.department_id, + "created_at": _format_utc_datetime(self.created_at), + "last_login": _format_utc_datetime(self.last_login), + "login_failed_count": self.login_failed_count, + "last_failed_login": _format_utc_datetime(self.last_failed_login), + "login_locked_until": _format_utc_datetime(self.login_locked_until), + "is_deleted": self.is_deleted, + "deleted_at": _format_utc_datetime(self.deleted_at), + } + if include_password: + result["password_hash"] = self.password_hash + return result + + def is_login_locked(self) -> bool: + """检查用户是否处于登录锁定状态""" + if self.login_locked_until is None: + return False + return _utc_now_naive() < self.login_locked_until + + def get_remaining_lock_time(self) -> int: + """获取剩余锁定时间(秒)""" + if self.login_locked_until is None: + return 0 + remaining = int((self.login_locked_until - _utc_now_naive()).total_seconds()) + return max(0, remaining) + + def reset_failed_login(self): + """重置登录失败相关字段""" + self.login_failed_count = 0 + self.last_failed_login = None + self.login_locked_until = None + + +class Conversation(Base): + """Conversation table - 对话表""" + + __tablename__ = "conversations" + + id = Column(Integer, primary_key=True, autoincrement=True, comment="Primary key") + thread_id = Column(String(64), unique=True, index=True, nullable=False, comment="Thread ID (UUID)") + user_id = Column(String(64), index=True, nullable=False, comment="User ID") + agent_id = Column(String(64), index=True, nullable=False, comment="Agent ID") + title = Column(String(255), nullable=True, comment="Conversation title") + status = Column(String(20), default="active", comment="Status: active/archived/deleted") + created_at = Column(DateTime, default=_utc_now_naive, comment="Creation time") + updated_at = Column(DateTime, default=_utc_now_naive, onupdate=_utc_now_naive, comment="Update time") + extra_metadata = Column(JSON, nullable=True, comment="Additional metadata") + + # Relationships + messages = relationship("Message", back_populates="conversation", cascade="all, delete-orphan") + stats = relationship("ConversationStats", back_populates="conversation", uselist=False, cascade="all, delete-orphan") + + def to_dict(self) -> dict[str, Any]: + return { + "id": self.id, + "thread_id": self.thread_id, + "user_id": self.user_id, + "agent_id": self.agent_id, + "title": self.title, + "status": self.status, + "created_at": _format_utc_datetime(self.created_at), + "updated_at": _format_utc_datetime(self.updated_at), + "metadata": self.extra_metadata or {}, + } + + +class Message(Base): + """Message table - 消息表""" + + __tablename__ = "messages" + + id = Column(Integer, primary_key=True, autoincrement=True, comment="Primary key") + conversation_id = Column( + Integer, ForeignKey("conversations.id"), nullable=False, index=True, comment="Conversation ID" + ) + role = Column(String(20), nullable=False, comment="Message role: user/assistant/system/tool") + content = Column(Text, nullable=False, comment="Message content") + message_type = Column(String(30), default="text", comment="Message type: text/tool_call/tool_result") + created_at = Column(DateTime, default=_utc_now_naive, comment="Creation time") + token_count = Column(Integer, nullable=True, comment="Token count (optional)") + extra_metadata = Column(JSON, nullable=True, comment="Additional metadata (complete message dump)") + image_content = Column(Text, nullable=True, comment="Base64 encoded image content for multimodal messages") + + # Relationships + conversation = relationship("Conversation", back_populates="messages") + tool_calls = relationship("ToolCall", back_populates="message", cascade="all, delete-orphan") + feedbacks = relationship("MessageFeedback", back_populates="message", cascade="all, delete-orphan") + + def to_dict(self) -> dict[str, Any]: + return { + "id": self.id, + "conversation_id": self.conversation_id, + "role": self.role, + "content": self.content, + "message_type": self.message_type, + "created_at": _format_utc_datetime(self.created_at), + "token_count": self.token_count, + "metadata": self.extra_metadata or {}, + "image_content": self.image_content, + "tool_calls": [tc.to_dict() for tc in self.tool_calls] if self.tool_calls else [], + } + + def to_simple_dict(self) -> dict[str, Any]: + return { + "role": self.role, + "content": self.content, + } + + +class ToolCall(Base): + """ToolCall table - 工具调用表""" + + __tablename__ = "tool_calls" + + id = Column(Integer, primary_key=True, autoincrement=True, comment="Primary key") + message_id = Column(Integer, ForeignKey("messages.id"), nullable=False, index=True, comment="Message ID") + langgraph_tool_call_id = Column(String(100), nullable=True, index=True, comment="LangGraph tool_call_id") + tool_name = Column(String(100), nullable=False, comment="Tool name") + tool_input = Column(JSON, nullable=True, comment="Tool input parameters") + tool_output = Column(Text, nullable=True, comment="Tool execution result") + status = Column(String(20), default="pending", comment="Status: pending/success/error") + error_message = Column(Text, nullable=True, comment="Error message if failed") + created_at = Column(DateTime, default=_utc_now_naive, comment="Creation time") + + # Relationships + message = relationship("Message", back_populates="tool_calls") + + def to_dict(self) -> dict[str, Any]: + return { + "id": self.id, + "message_id": self.message_id, + "langgraph_tool_call_id": self.langgraph_tool_call_id, + "tool_name": self.tool_name, + "tool_input": self.tool_input or {}, + "tool_output": self.tool_output, + "status": self.status, + "error_message": self.error_message, + "created_at": _format_utc_datetime(self.created_at), + } + + +class ConversationStats(Base): + """ConversationStats table - 对话统计表""" + + __tablename__ = "conversation_stats" + + id = Column(Integer, primary_key=True, autoincrement=True, comment="Primary key") + conversation_id = Column( + Integer, ForeignKey("conversations.id"), unique=True, nullable=False, comment="Conversation ID" + ) + message_count = Column(Integer, default=0, comment="Total message count") + total_tokens = Column(Integer, default=0, comment="Total tokens used") + model_used = Column(String(100), nullable=True, comment="Model used") + user_feedback = Column(JSON, nullable=True, comment="User feedback") + created_at = Column(DateTime, default=_utc_now_naive, comment="Creation time") + updated_at = Column(DateTime, default=_utc_now_naive, onupdate=_utc_now_naive, comment="Update time") + + # Relationships + conversation = relationship("Conversation", back_populates="stats") + + def to_dict(self) -> dict[str, Any]: + return { + "id": self.id, + "conversation_id": self.conversation_id, + "message_count": self.message_count, + "total_tokens": self.total_tokens, + "model_used": self.model_used, + "user_feedback": self.user_feedback or {}, + "created_at": _format_utc_datetime(self.created_at), + "updated_at": _format_utc_datetime(self.updated_at), + } + + +class OperationLog(Base): + """操作日志模型""" + + __tablename__ = "operation_logs" + + id = Column(Integer, primary_key=True, autoincrement=True) + user_id = Column(Integer, ForeignKey("users.id"), nullable=False) + operation = Column(String, nullable=False) + details = Column(Text, nullable=True) + ip_address = Column(String, nullable=True) + timestamp = Column(DateTime, default=_utc_now_naive) + + # 关联用户 + user = relationship("User", back_populates="operation_logs") + + def to_dict(self) -> dict[str, Any]: + return { + "id": self.id, + "user_id": self.user_id, + "operation": self.operation, + "details": self.details, + "ip_address": self.ip_address, + "timestamp": _format_utc_datetime(self.timestamp), + } + + +class MessageFeedback(Base): + """Message feedback table - 消息反馈表""" + + __tablename__ = "message_feedbacks" + + id = Column(Integer, primary_key=True, autoincrement=True, comment="Primary key") + message_id = Column(Integer, ForeignKey("messages.id"), nullable=False, index=True, comment="Message ID being rated") + user_id = Column(String(64), nullable=False, index=True, comment="User ID who provided feedback") + rating = Column(String(10), nullable=False, comment="Feedback rating: like or dislike") + reason = Column(Text, nullable=True, comment="Optional reason for dislike feedback") + created_at = Column(DateTime, default=_utc_now_naive, comment="Feedback creation time") + + # Relationships + message = relationship("Message", back_populates="feedbacks") + + def to_dict(self) -> dict[str, Any]: + return { + "id": self.id, + "message_id": self.message_id, + "user_id": self.user_id, + "rating": self.rating, + "reason": self.reason, + "created_at": _format_utc_datetime(self.created_at), + } + + +class MCPServer(Base): + """MCP 服务器配置模型""" + + __tablename__ = "mcp_servers" + + # 核心字段 - name 作为主键 + name = Column(String(100), primary_key=True, comment="服务器名称(唯一标识)") + description = Column(String(500), nullable=True, comment="描述") + + # 连接配置 + transport = Column(String(20), nullable=False, comment="传输类型:sse/streamable_http/stdio") + url = Column(String(500), nullable=True, comment="服务器 URL(sse/streamable_http)") + command = Column(String(500), nullable=True, comment="命令(stdio)") + args = Column(JSON, nullable=True, comment="命令参数数组(stdio)") + headers = Column(JSON, nullable=True, comment="HTTP 请求头") + timeout = Column(Integer, nullable=True, comment="HTTP 超时时间(秒)") + sse_read_timeout = Column(Integer, nullable=True, comment="SSE 读取超时(秒)") + + # UI 增强字段 + tags = Column(JSON, nullable=True, comment="标签数组") + icon = Column(String(50), nullable=True, comment="图标(emoji)") + + # 状态字段 + enabled = Column(Integer, nullable=False, default=1, comment="是否启用:1=是,0=否") + disabled_tools = Column(JSON, nullable=True, comment="禁用的工具名称列表") + + # 用户追踪 + created_by = Column(String(100), nullable=False, comment="创建人用户名") + updated_by = Column(String(100), nullable=False, comment="修改人用户名") + + # 时间戳 + created_at = Column(DateTime, default=_utc_now_naive, comment="创建时间") + updated_at = Column(DateTime, default=_utc_now_naive, onupdate=_utc_now_naive, comment="更新时间") + + def to_dict(self) -> dict[str, Any]: + return { + "name": self.name, + "description": self.description, + "transport": self.transport, + "url": self.url, + "command": self.command, + "args": self.args or [], + "headers": self.headers or {}, + "timeout": self.timeout, + "sse_read_timeout": self.sse_read_timeout, + "tags": self.tags or [], + "icon": self.icon, + "enabled": bool(self.enabled), + "disabled_tools": self.disabled_tools or [], + "created_by": self.created_by, + "updated_by": self.updated_by, + "created_at": _format_utc_datetime(self.created_at), + "updated_at": _format_utc_datetime(self.updated_at), + } + + def to_mcp_config(self) -> dict[str, Any]: + """转换为 MCP 配置格式(用于加载到 MCP_SERVERS 缓存)""" + config = {"transport": self.transport} + if self.url: + config["url"] = self.url + if self.command: + config["command"] = self.command + if self.args: + config["args"] = self.args + if self.headers: + config["headers"] = self.headers + if self.timeout is not None: + config["timeout"] = self.timeout + if self.sse_read_timeout is not None: + config["sse_read_timeout"] = self.sse_read_timeout + if self.disabled_tools: + config["disabled_tools"] = self.disabled_tools + return config diff --git a/test/api/test_dashboard_router.py b/test/api/test_dashboard_router.py index d0557862..eadef753 100644 --- a/test/api/test_dashboard_router.py +++ b/test/api/test_dashboard_router.py @@ -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) diff --git a/test/api/test_graph_router_list.py b/test/api/test_graph_router_list.py new file mode 100644 index 00000000..5c70cadc --- /dev/null +++ b/test/api/test_graph_router_list.py @@ -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