diff --git a/src/database/database.py b/src/database/database.py index fbdc4b6..38daad0 100644 --- a/src/database/database.py +++ b/src/database/database.py @@ -3,7 +3,6 @@ import sys import os from pathlib import Path -# 添加项目根目录到Python路径 project_root = Path(__file__).parent.parent.parent sys.path.insert(0, str(project_root)) @@ -12,6 +11,7 @@ from sqlalchemy import text from sqlalchemy.orm import sessionmaker from config.settings import settings import asyncio +from contextlib import asynccontextmanager from utils.logger import get_logger logger = get_logger(__name__) @@ -67,13 +67,17 @@ class DatabaseManager: self.is_connected = False logger.info("数据库连接已断开") + @asynccontextmanager async def session(self): """获取数据库会话的异步上下文管理器""" if not self.is_connected: await self.connect() - async with self.async_session() as session: + session = self.async_session() + try: yield session + finally: + await session.close() async def get_session(self) -> AsyncSession: """获取数据库会话""" diff --git a/src/database/init_db.py b/src/database/init_db.py index 120aa67..c5a657d 100644 --- a/src/database/init_db.py +++ b/src/database/init_db.py @@ -119,7 +119,7 @@ async def init_database(): await db_manager.connect() await db_manager.create_tables() - async with db_manager.get_session() as session: + async with db_manager.session() as session: perm_map = await init_permissions(session) if perm_map is None: # Permissions already existed, fetch them from database diff --git a/src/scripts/create_admin.py b/src/scripts/create_admin.py index daa1616..028a9ae 100644 --- a/src/scripts/create_admin.py +++ b/src/scripts/create_admin.py @@ -1,8 +1,17 @@ import asyncio +import sys +from pathlib import Path + +project_root = Path(__file__).parent.parent.parent +src_root = Path(__file__).parent.parent +sys.path.insert(0, str(project_root)) +sys.path.insert(0, str(src_root)) + from sqlalchemy import select from database.database import db_manager -from models.database import User +from models.database import User, Role, UserRole from services.auth_service import get_password_hash +from config.settings import settings from utils.logger import get_logger logger = get_logger(__name__) @@ -15,35 +24,43 @@ async def create_admin_user(): async with db_manager.session() as session: result = await session.execute( - select(User).where(User.username == "admin") + select(User).where(User.username == settings.ADMIN_USERNAME) ) existing_admin = result.scalar_one_or_none() if existing_admin: logger.info("管理员账户已存在") print("管理员账户已存在") - print("用户名: admin") + print(f"用户名: {settings.ADMIN_USERNAME}") return admin = User( - username="admin", - email="admin@gemold.com", - hashed_password=get_password_hash("admin123"), - full_name="系统管理员", - is_active=True, - is_superuser=True + username=settings.ADMIN_USERNAME, + email=settings.ADMIN_EMAIL, + hashed_password=get_password_hash(settings.ADMIN_PASSWORD), + full_name=settings.ADMIN_FULL_NAME, + is_active=True ) session.add(admin) + await session.flush() + + result = await session.execute(select(Role).where(Role.code == "admin")) + admin_role = result.scalar_one_or_none() + + if admin_role: + user_role = UserRole(user_id=admin.id, role_id=admin_role.id) + session.add(user_role) + await session.commit() logger.info("管理员账户创建成功") print("=" * 50) print("管理员账户创建成功!") print("=" * 50) - print("用户名: admin") - print("密码: admin123") - print("邮箱: admin@gemold.com") + print(f"用户名: {settings.ADMIN_USERNAME}") + print(f"密码: {settings.ADMIN_PASSWORD}") + print(f"邮箱: {settings.ADMIN_EMAIL}") print("=" * 50) print("⚠️ 请登录后立即修改密码!") print("=" * 50)