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