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