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自动授权
This commit is contained in:
@@ -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()
|
||||
]
|
||||
}
|
||||
@@ -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(
|
||||
|
||||
@@ -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层级组织示例数据"""
|
||||
|
||||
+45
-11
@@ -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、无参数时兜底,兼容存量公开接口)
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user