Files
cma-management/backend/app/auth.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

181 lines
6.6 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.
"""用户认证(账套模式:登录前选公司,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()
]
}