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自动授权
This commit is contained in:
Hermes CI Fix
2026-08-11 11:24:52 +08:00
parent 9d0f070c14
commit b31f9b80c4
5 changed files with 336 additions and 20 deletions
+180
View File
@@ -0,0 +1,180 @@
"""用户认证(账套模式:登录前选公司,token绑定entity"""
import hashlib
from fastapi import APIRouter, Depends, HTTPException, Request
from sqlalchemy.orm import Session
from app.database import get_db
from app.models import User, UserEntity, Entity, OperationLog
from app.auth_middleware import (
create_token, require_auth, ROLES,
extract_bearer_token, get_token_entity_id,
)
router = APIRouter(prefix="/api/cma/auth", tags=["认证"])
def _user_dict(user: User, entity: Entity = None) -> dict:
"""用户信息(新增字段设默认值,不破坏现有响应格式)"""
d = {
"id": user.id,
"username": user.username,
"name": user.name,
"role": user.role,
"role_name": ROLES.get(user.role, {}).get("name", user.role),
}
if entity:
d["entity_id"] = entity.id
d["entity_name"] = entity.name
d["entity_short_name"] = entity.short_name
return d
def _check_authorized(db: Session, user_id: int, entity_id: int) -> bool:
"""校验用户是否被授权访问该企业(账套授权表)"""
return db.query(UserEntity).filter(
UserEntity.user_id == user_id,
UserEntity.entity_id == entity_id,
).first() is not None
def _get_active_entity(db: Session, entity_id: int) -> Entity:
ent = db.query(Entity).filter(Entity.id == entity_id).first()
if not ent or ent.status != "active":
raise HTTPException(403, f"企业entity_id={entity_id}不存在或未激活")
return ent
@router.post("/login")
def login(data: dict, db: Session = Depends(get_db)):
username = data.get("username", "")
password = data.get("password", "")
user = db.query(User).filter(User.username == username).first()
if not user or user.password_hash != hashlib.sha256(password.encode()).hexdigest():
raise HTTPException(401, "用户名或密码错误")
# 账套模式:entity_id 必选语义(兼容旧调用:缺省取用户第一个授权企业)
entity_id = data.get("entity_id")
if entity_id is not None:
entity_id = int(entity_id)
if not _check_authorized(db, user.id, entity_id):
raise HTTPException(403, "您没有该企业的访问权限")
else:
first = db.query(UserEntity).filter(UserEntity.user_id == user.id).order_by(UserEntity.entity_id).first()
entity_id = first.entity_id if first else 1
ent = _get_active_entity(db, entity_id)
token = create_token(user.id, entity_id)
return {
"token": token,
"user": _user_dict(user, ent),
}
@router.post("/switch-entity")
def switch_entity(data: dict, request: Request, current_user: User = Depends(require_auth), db: Session = Depends(get_db)):
"""账套切换:校验授权 → 重新签发token(切换留痕)→ 前端替换token+整页刷新"""
entity_id = data.get("entity_id")
if not entity_id:
raise HTTPException(400, "缺少entity_id")
entity_id = int(entity_id)
if not _check_authorized(db, current_user.id, entity_id):
raise HTTPException(403, "您没有该企业的访问权限")
ent = _get_active_entity(db, entity_id)
token = create_token(current_user.id, entity_id)
# 切换留痕
try:
db.add(OperationLog(
user_id=current_user.id,
action="switch_entity",
target_type="entity",
target_id=entity_id,
detail={"entity_name": ent.name, "entity_short_name": ent.short_name},
))
db.commit()
except Exception:
db.rollback()
return {
"token": token,
"entity_id": entity_id,
"entity_name": ent.name,
"entity_short_name": ent.short_name,
"user": _user_dict(current_user, ent),
}
@router.get("/entities")
def get_login_entities(username: str = None, db: Session = Depends(get_db)):
"""登录页公司选择器:按用户名返回授权企业(无鉴权;不暴露用户名是否存在)"""
if not username:
return {"data": []}
user = db.query(User).filter(User.username == username).first()
if not user:
return {"data": []}
rows = db.query(UserEntity).filter(UserEntity.user_id == user.id).all()
ids = [r.entity_id for r in rows]
if not ids:
return {"data": []}
ents = db.query(Entity).filter(Entity.id.in_(ids), Entity.status == "active").order_by(Entity.id).all()
return {"data": [{"id": e.id, "name": e.name, "short_name": e.short_name} for e in ents]}
@router.get("/my-entities")
def get_my_entities(request: Request, current_user: User = Depends(require_auth), db: Session = Depends(get_db)):
"""当前用户授权企业列表(切换器用)"""
rows = db.query(UserEntity).filter(UserEntity.user_id == current_user.id).all()
ids = [r.entity_id for r in rows]
ents = db.query(Entity).filter(Entity.id.in_(ids), Entity.status == "active").order_by(Entity.id).all() if ids else []
token = extract_bearer_token(request)
current_eid = get_token_entity_id(token) if token else None
return {
"data": [{"id": e.id, "name": e.name, "short_name": e.short_name} for e in ents],
"current_entity_id": current_eid,
}
@router.post("/register")
def register(data: dict, db: Session = Depends(get_db)):
exist = db.query(User).filter(User.username == data.get("username")).first()
if exist:
raise HTTPException(400, "用户名已存在")
user = User(
username=data["username"],
password_hash=hashlib.sha256(data["password"].encode()).hexdigest(),
name=data.get("name", data["username"]),
role=data.get("role", "business"),
)
db.add(user)
db.commit()
db.refresh(user)
# 注册默认授权第一个active企业
first_ent = db.query(Entity).filter(Entity.status == "active").order_by(Entity.id).first()
if first_ent:
db.add(UserEntity(user_id=user.id, entity_id=first_ent.id))
db.commit()
return {"message": "注册成功"}
@router.get("/me")
def get_me(request: Request, current_user: User = Depends(require_auth), db: Session = Depends(get_db)):
"""获取当前用户信息(含当前账套)"""
token = extract_bearer_token(request)
eid = get_token_entity_id(token) if token else None
ent = db.query(Entity).filter(Entity.id == eid).first() if eid else None
d = _user_dict(current_user, ent)
d["phone"] = current_user.phone
return d
@router.get("/roles")
def list_roles():
"""返回角色列表(给前端用)"""
return {
"data": [
{"code": k, "name": v["name"], "priority": v["priority"]}
for k, v in ROLES.items()
]
}
+70 -8
View File
@@ -60,7 +60,7 @@ except Exception:
logger.warning("Redis不可用,token存储降级到内存(不支持多worker)")
# 内存 fallback
_token_store: dict[str, int] = {}
_token_store: dict[str, dict] = {}
TOKEN_PREFIX = "cma:token:"
TOKEN_TTL = 86400 # 24小时
@@ -95,22 +95,84 @@ def _load_permissions(db: Session = None):
return DEFAULT_ROUTE_PERMISSIONS, DEFAULT_ACTION_PERMISSIONS
def create_token(user_id: int) -> str:
def create_token(user_id: int, entity_id: int = None) -> str:
"""签发tokenRedis存储 JSON {user_id, entity_id}(账套模式)
兼容旧调用 create_token(user_id) → entity_id=None(切换器会重新签发)
"""
token = secrets.token_hex(32)
payload = json.dumps({"user_id": user_id, "entity_id": entity_id}, ensure_ascii=False)
if _redis_available:
_redis.setex(f"{TOKEN_PREFIX}{token}", TOKEN_TTL, user_id)
_redis.setex(f"{TOKEN_PREFIX}{token}", TOKEN_TTL, payload)
else:
_token_store[token] = user_id
_token_store[token] = payload
return token
def _resolve_user_id(token: str) -> int | None:
def _resolve_token_data(token: str) -> dict | None:
"""解析token → {user_id, entity_id}
- 新格式 JSON → 返回 dict
- 旧格式 int(改造前)→ 返回 None(强制下线,账套模式需重新登录)
- 不存在 → None
"""
if _redis_available:
val = _redis.get(f"{TOKEN_PREFIX}{token}")
if val is not None:
return int(val)
if val is None:
return None
try:
data = json.loads(val)
if isinstance(data, dict) and "user_id" in data:
return data
except (json.JSONDecodeError, ValueError, TypeError):
pass
# 旧格式纯 int → 强制下线
return None
return _token_store.get(token)
val = _token_store.get(token)
if val is None:
return None
if isinstance(val, dict):
return val
try:
data = json.loads(val)
if isinstance(data, dict) and "user_id" in data:
return data
except (json.JSONDecodeError, ValueError, TypeError):
pass
return None
def _resolve_user_id(token: str) -> int | None:
data = _resolve_token_data(token)
return int(data["user_id"]) if data else None
def get_token_entity_id(token: str) -> int | None:
"""从token解析绑定的entity_idtoken不存在/旧格式 → None"""
data = _resolve_token_data(token)
if not data:
return None
eid = data.get("entity_id")
return int(eid) if eid else None
def user_has_entity(db: Session, user_id: int, entity_id: int) -> bool:
"""校验用户是否被授权访问指定企业(账套授权表 user_entities"""
from app.models import UserEntity, Entity
ent = db.query(Entity).filter(Entity.id == entity_id).first()
if not ent or ent.status != "active":
return False
rel = db.query(UserEntity).filter(
UserEntity.user_id == user_id,
UserEntity.entity_id == entity_id,
).first()
return rel is not None
def extract_bearer_token(request) -> str | None:
"""从Request提取Bearer token(无/非Bearer格式 → None"""
auth = request.headers.get("Authorization", "")
if auth.startswith("Bearer "):
return auth[7:].strip()
return None
def require_auth(
+29
View File
@@ -68,6 +68,35 @@ def init_db():
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层级组织示例数据"""
+45 -11
View File
@@ -1,26 +1,60 @@
"""多租户公共依赖 — 从请求头/参数读取当前企业ID"""
from fastapi import Request, Header, Query, Depends
"""多租户公共依赖 — 账套模式:token优先,query/header降级为Bot服务白名单"""
from fastapi import Request, Header, Query, Depends, HTTPException
from typing import Optional
from sqlalchemy.orm import Session
from app.database import get_db
from app.auth_middleware import get_token_entity_id
from app.models import Entity
def get_entity_id(
request: Request,
x_entity_id: Optional[str] = Header(None, alias="X-Entity-Id"),
entity_id: Optional[int] = Query(None, ge=1),
db: Session = Depends(get_db),
) -> int:
"""解析当前企业ID:优先query参数 → header → 默认1(酣客)
前端拦截器已统一附加 X-Entity-Id header entity_id query参数
"""解析当前企业ID(账套模式):token优先 → query/headerBot白名单)→ 默认1
解析链倒置后
1. Authorization Bearer token 优先返回 token.entity_id唯一可信来源
2. query参数 / X-Entity-Id header 仅Bot服务通道使用校验entity状态active
3. body中的 entity_idPOST场景 Bot服务通道兼容
4. 兜底默认 1酣客
安全说明登录用户必须通过token绑定账套query/header传入的entity_id
在token存在时被忽略防止越权传参旧漏洞query优先且无授权校验
"""
# 1. token优先(账套模式唯一来源)
auth = request.headers.get("Authorization", "")
if auth.startswith("Bearer "):
token_entity = get_token_entity_id(auth[7:])
if token_entity is not None:
return token_entity
# token存在但是旧格式/无entity → 账套模式下强制走白名单或默认(由require_auth拦截)
# 这里不抛401:公开接口可能带旧token,交给require_auth统一处理
# 2. query/header → Bot服务通道白名单(校验entity状态active
candidate = None
if entity_id is not None:
return entity_id
if x_entity_id and x_entity_id.isdigit():
return int(x_entity_id)
# 兼容 body 中的 entity_idPOST场景)
if request.method in ("POST", "PUT", "PATCH"):
candidate = entity_id
elif x_entity_id and x_entity_id.isdigit():
candidate = int(x_entity_id)
# 3. body 中的 entity_idPOST场景,Bot通道兼容)
if candidate is None and request.method in ("POST", "PUT", "PATCH"):
try:
body = request.state.body_json or {}
if body.get("entity_id"):
return int(body["entity_id"])
candidate = int(body["entity_id"])
except Exception:
pass
return 1 # 默认酣客
if candidate is not None:
ent = db.query(Entity).filter(Entity.id == candidate).first()
if ent and ent.status == "active":
return candidate
# Bot通道校验不通过:不返回默认,直接拒绝(防越权注入无效entity)
raise HTTPException(403, f"企业 entity_id={candidate} 不存在或未激活")
return 1 # 默认酣客(无token、无参数时兜底,兼容存量公开接口)
+12 -1
View File
@@ -1,5 +1,5 @@
"""管理会计OS 数据模型"""
from sqlalchemy import Column, Integer, String, Text, Float, DateTime, ForeignKey, Boolean, JSON, func
from sqlalchemy import Column, Integer, String, Text, Float, DateTime, ForeignKey, Boolean, JSON, func, UniqueConstraint
from app.database import Base
from app.models.budget_plan import BudgetPlan
@@ -32,6 +32,17 @@ class User(Base):
created_at = Column(DateTime, server_default=func.now())
class UserEntity(Base):
"""用户-企业授权(账套模式多对多)— 新表必须带entity_id(开发规范)"""
__tablename__ = "user_entities"
id = Column(Integer, primary_key=True, index=True)
user_id = Column(Integer, ForeignKey("users.id"), nullable=False, index=True, comment="用户ID")
entity_id = Column(Integer, ForeignKey("entities.id"), nullable=False, index=True, comment="企业ID(账套)")
granted_by = Column(Integer, nullable=True, comment="授权人")
created_at = Column(DateTime, server_default=func.now())
__table_args__ = (UniqueConstraint("user_id", "entity_id", name="uq_user_entity"),)
class StrategicMap(Base):
"""战略地图"""
__tablename__ = "strategic_maps"