From b31f9b80c46b62a53aab1febfc6eda517129d036 Mon Sep 17 00:00:00 2001 From: Hermes CI Fix Date: Tue, 11 Aug 2026 11:24:52 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E8=B4=A6=E5=A5=97=E6=A8=A1=E5=BC=8F?= =?UTF-8?q?=E5=90=8E=E7=AB=AF=20=E2=80=94=20token=E7=BB=91=E5=AE=9Aentity?= =?UTF-8?q?=5Fid=20+=20user=5Fentities=E6=8E=88=E6=9D=83=E8=A1=A8=20+=20?= =?UTF-8?q?=E8=A7=A3=E6=9E=90=E9=93=BE=E5=80=92=E7=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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自动授权 --- backend/app/auth.py | 180 +++++++++++++++++++++++++++++++++ backend/app/auth_middleware.py | 78 ++++++++++++-- backend/app/database.py | 29 ++++++ backend/app/deps.py | 56 ++++++++-- backend/app/models/__init__.py | 13 ++- 5 files changed, 336 insertions(+), 20 deletions(-) create mode 100644 backend/app/auth.py diff --git a/backend/app/auth.py b/backend/app/auth.py new file mode 100644 index 00000000..717bb8f7 --- /dev/null +++ b/backend/app/auth.py @@ -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() + ] + } diff --git a/backend/app/auth_middleware.py b/backend/app/auth_middleware.py index 6454f7d8..7e25bf6a 100644 --- a/backend/app/auth_middleware.py +++ b/backend/app/auth_middleware.py @@ -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: + """签发token:Redis存储 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_id(token不存在/旧格式 → 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( diff --git a/backend/app/database.py b/backend/app/database.py index a55438ca..dfb65580 100644 --- a/backend/app/database.py +++ b/backend/app/database.py @@ -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层级组织示例数据""" diff --git a/backend/app/deps.py b/backend/app/deps.py index 210de261..3d8bfab3 100644 --- a/backend/app/deps.py +++ b/backend/app/deps.py @@ -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/header(Bot白名单)→ 默认1 + + 解析链(倒置后): + 1. Authorization Bearer token → 优先返回 token.entity_id(唯一可信来源) + 2. query参数 / X-Entity-Id header → 仅Bot服务通道使用(校验entity状态active) + 3. body中的 entity_id(POST场景)→ 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_id(POST场景) - 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_id(POST场景,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、无参数时兜底,兼容存量公开接口) diff --git a/backend/app/models/__init__.py b/backend/app/models/__init__.py index 0d6e125f..a72c4099 100644 --- a/backend/app/models/__init__.py +++ b/backend/app/models/__init__.py @@ -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"