feat: auth.py 账套模式 — login强制entity_id + my-entities授权查询 + switch-entity重签发token

This commit is contained in:
Hermes CI Fix
2026-08-11 11:25:18 +08:00
parent b31f9b80c4
commit 35839dd38e
+110 -4
View File
@@ -1,14 +1,31 @@
"""用户认证""" """用户认证"""
import hashlib import hashlib
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException, Request
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.database import get_db from app.database import get_db
from app.models import User from app.models import User, UserEntity, Entity, OperationLog
from app.auth_middleware import create_token, require_auth, ROLES from app.auth_middleware import create_token, require_auth, ROLES
router = APIRouter(prefix="/api/cma/auth", tags=["认证"]) router = APIRouter(prefix="/api/cma/auth", tags=["认证"])
def _user_entity_ids(db: Session, user_id: int) -> list[int]:
"""返回用户被授权的企业ID列表"""
rows = db.query(UserEntity.entity_id).filter(UserEntity.user_id == user_id).all()
return [r[0] for r in rows]
def _ensure_default_grants(db: Session, user_id: int) -> None:
"""存量用户兼容:若用户没有任何授权记录,则授予所有active企业(不锁死老账号)"""
cnt = db.query(UserEntity).filter(UserEntity.user_id == user_id).count()
if cnt > 0:
return
entities = db.query(Entity).filter(Entity.status == "active").all()
for e in entities:
db.add(UserEntity(user_id=user_id, entity_id=e.id, granted_by=1))
db.commit()
@router.post("/login") @router.post("/login")
def login(data: dict, db: Session = Depends(get_db)): def login(data: dict, db: Session = Depends(get_db)):
username = data.get("username", "") username = data.get("username", "")
@@ -17,9 +34,28 @@ def login(data: dict, db: Session = Depends(get_db)):
if not user or user.password_hash != hashlib.sha256(password.encode()).hexdigest(): if not user or user.password_hash != hashlib.sha256(password.encode()).hexdigest():
raise HTTPException(401, "用户名或密码错误") raise HTTPException(401, "用户名或密码错误")
token = create_token(user.id) # 账套模式:登录必须指定企业(entity_id)
entity_id = data.get("entity_id")
if entity_id is None:
raise HTTPException(400, "账套模式:请选择登录企业(entity_id)")
# 存量兼容:无授权记录时自动授予active企业
_ensure_default_grants(db, user.id)
allowed = _user_entity_ids(db, user.id)
if int(entity_id) not in allowed:
raise HTTPException(403, f"该账号未被授权访问企业 entity_id={entity_id}")
entity = db.query(Entity).filter(Entity.id == int(entity_id)).first()
if not entity or entity.status != "active":
raise HTTPException(403, f"企业 entity_id={entity_id} 不存在或未激活")
token = create_token(user.id, int(entity_id))
return { return {
"token": token, "token": token,
"entity_id": int(entity_id),
"entity_name": entity.name,
"entity_short_name": entity.short_name,
"user": { "user": {
"id": user.id, "id": user.id,
"username": user.username, "username": user.username,
@@ -43,12 +79,14 @@ def register(data: dict, db: Session = Depends(get_db)):
) )
db.add(user) db.add(user)
db.commit() db.commit()
# 新用户默认授予所有active企业
_ensure_default_grants(db, user.id)
return {"message": "注册成功"} return {"message": "注册成功"}
@router.get("/me") @router.get("/me")
def get_me(current_user: User = Depends(require_auth)): def get_me(current_user: User = Depends(require_auth)):
"""获取当前用户信息""" """获取当前用户信息(含当前账套)"""
return { return {
"id": current_user.id, "id": current_user.id,
"username": current_user.username, "username": current_user.username,
@@ -68,3 +106,71 @@ def list_roles():
for k, v in ROLES.items() for k, v in ROLES.items()
] ]
} }
@router.get("/entities")
def my_entities(current_user: User = Depends(require_auth), db: Session = Depends(get_db)):
"""当前用户被授权的企业列表(登录页下拉/切换器数据源)"""
_ensure_default_grants(db, current_user.id)
ids = _user_entity_ids(db, current_user.id)
entities = 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, "industry": e.industry}
for e in entities
]
}
@router.post("/switch-entity")
def switch_entity(
data: dict,
request: Request,
current_user: User = Depends(require_auth),
db: Session = Depends(get_db),
):
"""账套切换:校验授权 → 重新签发token → 写操作日志(方案B:无缝刷新token)"""
entity_id = data.get("entity_id")
if not entity_id:
raise HTTPException(400, "缺少entity_id")
_ensure_default_grants(db, current_user.id)
allowed = _user_entity_ids(db, current_user.id)
if int(entity_id) not in allowed:
raise HTTPException(403, f"该账号未被授权访问企业 entity_id={entity_id}")
entity = db.query(Entity).filter(Entity.id == int(entity_id)).first()
if not entity or entity.status != "active":
raise HTTPException(403, f"企业 entity_id={entity_id} 不存在或未激活")
# 旧token失效:删除当前请求的旧token(切换即下线旧账套凭据)
auth = request.headers.get("Authorization", "")
if auth.startswith("Bearer "):
try:
import redis as redis_lib
r = redis_lib.Redis(host="127.0.0.1", port=6379, db=1, decode_responses=True)
r.delete(f"cma:token:{auth[7:]}")
except Exception:
pass
token = create_token(current_user.id, int(entity_id))
# 切换留痕
db.add(OperationLog(
user_id=current_user.id,
action="switch_entity",
target_type="entity",
target_id=int(entity_id),
detail={"entity_id": int(entity_id), "entity_name": entity.name, "from": "account-switch"},
))
db.commit()
return {
"token": token,
"entity_id": int(entity_id),
"entity_name": entity.name,
"entity_short_name": entity.short_name,
}