init
This commit is contained in:
@@ -3,7 +3,6 @@ import sys
|
|||||||
import os
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
# 添加项目根目录到Python路径
|
|
||||||
project_root = Path(__file__).parent.parent.parent
|
project_root = Path(__file__).parent.parent.parent
|
||||||
sys.path.insert(0, str(project_root))
|
sys.path.insert(0, str(project_root))
|
||||||
|
|
||||||
@@ -12,6 +11,7 @@ from sqlalchemy import text
|
|||||||
from sqlalchemy.orm import sessionmaker
|
from sqlalchemy.orm import sessionmaker
|
||||||
from config.settings import settings
|
from config.settings import settings
|
||||||
import asyncio
|
import asyncio
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
from utils.logger import get_logger
|
from utils.logger import get_logger
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
@@ -67,13 +67,17 @@ class DatabaseManager:
|
|||||||
self.is_connected = False
|
self.is_connected = False
|
||||||
logger.info("数据库连接已断开")
|
logger.info("数据库连接已断开")
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
async def session(self):
|
async def session(self):
|
||||||
"""获取数据库会话的异步上下文管理器"""
|
"""获取数据库会话的异步上下文管理器"""
|
||||||
if not self.is_connected:
|
if not self.is_connected:
|
||||||
await self.connect()
|
await self.connect()
|
||||||
|
|
||||||
async with self.async_session() as session:
|
session = self.async_session()
|
||||||
|
try:
|
||||||
yield session
|
yield session
|
||||||
|
finally:
|
||||||
|
await session.close()
|
||||||
|
|
||||||
async def get_session(self) -> AsyncSession:
|
async def get_session(self) -> AsyncSession:
|
||||||
"""获取数据库会话"""
|
"""获取数据库会话"""
|
||||||
|
|||||||
@@ -119,7 +119,7 @@ async def init_database():
|
|||||||
await db_manager.connect()
|
await db_manager.connect()
|
||||||
await db_manager.create_tables()
|
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)
|
perm_map = await init_permissions(session)
|
||||||
if perm_map is None:
|
if perm_map is None:
|
||||||
# Permissions already existed, fetch them from database
|
# Permissions already existed, fetch them from database
|
||||||
|
|||||||
+29
-12
@@ -1,8 +1,17 @@
|
|||||||
import asyncio
|
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 sqlalchemy import select
|
||||||
from database.database import db_manager
|
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 services.auth_service import get_password_hash
|
||||||
|
from config.settings import settings
|
||||||
from utils.logger import get_logger
|
from utils.logger import get_logger
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
@@ -15,35 +24,43 @@ async def create_admin_user():
|
|||||||
|
|
||||||
async with db_manager.session() as session:
|
async with db_manager.session() as session:
|
||||||
result = await session.execute(
|
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()
|
existing_admin = result.scalar_one_or_none()
|
||||||
|
|
||||||
if existing_admin:
|
if existing_admin:
|
||||||
logger.info("管理员账户已存在")
|
logger.info("管理员账户已存在")
|
||||||
print("管理员账户已存在")
|
print("管理员账户已存在")
|
||||||
print("用户名: admin")
|
print(f"用户名: {settings.ADMIN_USERNAME}")
|
||||||
return
|
return
|
||||||
|
|
||||||
admin = User(
|
admin = User(
|
||||||
username="admin",
|
username=settings.ADMIN_USERNAME,
|
||||||
email="admin@gemold.com",
|
email=settings.ADMIN_EMAIL,
|
||||||
hashed_password=get_password_hash("admin123"),
|
hashed_password=get_password_hash(settings.ADMIN_PASSWORD),
|
||||||
full_name="系统管理员",
|
full_name=settings.ADMIN_FULL_NAME,
|
||||||
is_active=True,
|
is_active=True
|
||||||
is_superuser=True
|
|
||||||
)
|
)
|
||||||
|
|
||||||
session.add(admin)
|
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()
|
await session.commit()
|
||||||
|
|
||||||
logger.info("管理员账户创建成功")
|
logger.info("管理员账户创建成功")
|
||||||
print("=" * 50)
|
print("=" * 50)
|
||||||
print("管理员账户创建成功!")
|
print("管理员账户创建成功!")
|
||||||
print("=" * 50)
|
print("=" * 50)
|
||||||
print("用户名: admin")
|
print(f"用户名: {settings.ADMIN_USERNAME}")
|
||||||
print("密码: admin123")
|
print(f"密码: {settings.ADMIN_PASSWORD}")
|
||||||
print("邮箱: admin@gemold.com")
|
print(f"邮箱: {settings.ADMIN_EMAIL}")
|
||||||
print("=" * 50)
|
print("=" * 50)
|
||||||
print("⚠️ 请登录后立即修改密码!")
|
print("⚠️ 请登录后立即修改密码!")
|
||||||
print("=" * 50)
|
print("=" * 50)
|
||||||
|
|||||||
Reference in New Issue
Block a user