Files
cma-management/backend/app/database.py
T
Hermes CI Fix b31f9b80c4 feat: 账套模式后端 — token绑定entity_id + user_entities授权表 + 解析链倒置
- auth_middleware: create_token存JSON{user_id,entity_id},旧int格式token强制下线
- models: 新增UserEntity授权表(user_id↔entity_id多对多,唯一约束)
- database: init_db自动建表+存量用户×active企业默认授权(平滑迁移)
- deps: get_entity_id解析链倒置 token优先 → Bot白名单(query/header校验entity active) → 默认1
- auth: login加entity_id+授权校验; 新增switch-entity/my-entities/登录页entities接口; register自动授权
2026-08-11 11:24:52 +08:00

152 lines
5.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""数据库配置"""
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker, declarative_base
from sqlalchemy import inspect
import os
import logging
logger = logging.getLogger("cma")
DB_USER = os.getenv("CMA_DB_USER", "cma_user")
DB_PASS = os.getenv("CMA_DB_PASS", "cma_pass_2026")
DB_HOST = os.getenv("CMA_DB_HOST", "127.0.0.1")
DB_PORT = os.getenv("CMA_DB_PORT", "3306")
DB_NAME = os.getenv("CMA_DB_NAME", "cma")
DATABASE_URL = "mysql+pymysql://%(user)s:%(password)s@%(host)s:%(port)s/%(name)s?charset=utf8mb4" % {
"user": DB_USER,
"password": DB_PASS,
"host": DB_HOST,
"port": DB_PORT,
"name": DB_NAME,
}
_engine = None
_SessionLocal = None
Base = declarative_base()
def get_engine():
global _engine
if _engine is None:
_engine = create_engine(DATABASE_URL, echo=False, pool_size=5, max_overflow=10, pool_pre_ping=True)
return _engine
def get_session_local():
global _SessionLocal
if _SessionLocal is None:
_SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=get_engine())
return _SessionLocal
def get_db():
db = get_session_local()()
try:
yield db
finally:
db.close()
def init_db():
import app.models
Base.metadata.create_all(bind=get_engine())
logger.info("CMA数据库已初始化")
# ── 初始化组织层级示例数据 ──
try:
inspector = inspect(get_engine())
if "org_nodes" in inspector.get_table_names():
Session = get_session_local()
session = Session()
try:
cnt = session.query(app.models.OrgNode).count()
if cnt == 0:
_seed_org_data(session)
finally:
session.close()
except Exception as e:
logger.warning(f"组织数据初始化跳过: {e}")
# ── user_entities 授权表初始化(账套模式平滑迁移)──
# 首次建表(表空)时,给存量用户默认授权所有 active 企业,保证现有登录不丢权限
try:
inspector = inspect(get_engine())
if "user_entities" in inspector.get_table_names():
Session = get_session_local()
session = Session()
try:
from app.models import UserEntity, User, Entity
cnt = session.query(UserEntity).count()
if cnt == 0:
users = session.query(User).all()
entities = session.query(Entity).filter(Entity.status == "active").all()
if users and entities:
for u in users:
for e in entities:
exists = session.query(UserEntity).filter(
UserEntity.user_id == u.id,
UserEntity.entity_id == e.id,
).first()
if not exists:
session.add(UserEntity(user_id=u.id, entity_id=e.id, granted_by=None))
session.commit()
logger.info(f"user_entities初始化: {len(users)}用户 × {len(entities)}企业")
finally:
session.close()
except Exception as e:
logger.warning(f"user_entities初始化跳过: {e}")
def _seed_org_data(db_session):
"""插入5层级组织示例数据"""
from app.models import OrgNode
# 1. 集团
g = OrgNode(id=1, parent_id=None, name="博海网络科技", code="BH", level=1, sort_order=1, enabled=1)
db_session.add(g)
db_session.flush()
# 2. 事业部
depts = [
OrgNode(parent_id=1, name="技术事业部", code="TECH", level=2, sort_order=1, enabled=1),
OrgNode(parent_id=1, name="销售事业部", code="SALES", level=2, sort_order=2, enabled=1),
OrgNode(parent_id=1, name="财务事业部", code="FIN", level=2, sort_order=3, enabled=1),
]
db_session.add_all(depts)
db_session.flush()
# 3. 区域/部门级
regions = [
OrgNode(parent_id=2, name="华南区域", code="SC", level=3, sort_order=1, enabled=1),
OrgNode(parent_id=2, name="华东区域", code="EC", level=3, sort_order=2, enabled=1),
OrgNode(parent_id=3, name="销售一部", code="S1", level=3, sort_order=1, enabled=1),
OrgNode(parent_id=3, name="销售二部", code="S2", level=3, sort_order=2, enabled=1),
]
db_session.add_all(regions)
db_session.flush()
# 4. 部门
departs = [
OrgNode(parent_id=5, name="研发部", code="RD", level=4, sort_order=1, enabled=1),
OrgNode(parent_id=5, name="实施部", code="IMP", level=4, sort_order=2, enabled=1),
OrgNode(parent_id=5, name="运维部", code="OPS", level=4, sort_order=3, enabled=1),
OrgNode(parent_id=6, name="前端研发", code="FE", level=4, sort_order=1, enabled=1),
OrgNode(parent_id=6, name="后端研发", code="BE", level=4, sort_order=2, enabled=1),
OrgNode(parent_id=8, name="KA客户部", code="KA", level=4, sort_order=1, enabled=1),
]
db_session.add_all(departs)
db_session.flush()
# 5. 班组
teams = [
OrgNode(parent_id=12, name="前端组", code="FE-TEAM", level=5, sort_order=1, enabled=1),
OrgNode(parent_id=12, name="后端组", code="BE-TEAM", level=5, sort_order=2, enabled=1),
OrgNode(parent_id=12, name="测试组", code="QA-TEAM", level=5, sort_order=3, enabled=1),
OrgNode(parent_id=13, name="实施一组", code="IMP1", level=5, sort_order=1, enabled=1),
OrgNode(parent_id=13, name="实施二组", code="IMP2", level=5, sort_order=2, enabled=1),
]
db_session.add_all(teams)
db_session.commit()
logger.info("组织层级示例数据已初始化")