Files
server-core/app/api/dependencies.py
T

240 lines
9.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""接口层依赖:数据库句柄与当前登录用户(JWT + RBAC,全异步)。
``Database`` 实例由 ``main.py`` 在启动时创建并挂在 ``app.state.db`` 上,
路由通过 ``Depends(get_db)`` 取用;测试时可替换为临时目录实例。
"""
from __future__ import annotations
import asyncio
import contextlib
from fastapi import Depends, HTTPException, Request
from .. import config
from ..infrastructure.repositories import Database
from ..jwt import decode_access_token
BEARER_PREFIX = "Bearer "
def extract_bearer_token(request: Request) -> str:
"""从 Authorization 头提取 Bearer token,没有则返回空串。"""
auth_header = request.headers.get("Authorization", "")
if auth_header.startswith(BEARER_PREFIX):
return auth_header[len(BEARER_PREFIX):].strip()
return ""
async def get_db():
"""每请求一个独立 Database/AsyncSession(绑定共享引擎)。
之前直接返回 app.state.db(单例 AsyncSession)被并发请求复用,会间歇报
``sqlalchemy.exc.InvalidRequestError: concurrent operations are not permitted``。
改为每请求各建一个 Database 并在请求结束关闭,根除该竞态。
"""
db = Database()
try:
yield db
finally:
try:
await db.close()
except asyncio.CancelledError:
# 请求任务已被取消(客户端断开 / 超时 / 服务关闭):取消上下文中
# 的 await 会立即抛 CancelledErrorSQLAlchemy 的 rollback 会被
# asyncmy 打断并报 "Cancelled during execution",导致 session 关不掉、
# 连接无法归还连接池(随后被 GC 回收时产生 non-checked-in 警告)。
# 用 shield 把 close 移到独立任务执行,确保连接一定归还。
with contextlib.suppress(BaseException):
await asyncio.shield(db.close())
raise
async def _compute_capabilities(db: Database, user: dict) -> list[str]:
"""叠加制能力集合(身份不互斥):
- opc_member:所有账号的基础能力;
- operator:运营端角色;
- carrier:园区端角色 或 已绑定园区管理员(organization_members.role=admin + org.type=carrier,或旧表 park_members.admin / park_tenants.operator_user_id);
- enterprise:已绑定企业管理员(organization_members.role=admin + org.type=enterprise,或旧表 company_members.is_admin)。
优先查统一的 organization_members 表,回退查旧表(park_members/company_members),保持兼容。
"""
from ..infrastructure.models import CompanyMember, Organization, OrganizationMember, ParkMember, ParkTenant
from sqlalchemy import select
role = user.get("role", "")
caps = ["opc_member"]
if role in ("operator", "operator_internal", "op_admin", "op_super_admin", "superadmin", "admin"):
caps.append("operator")
if role in ("carrier", "park", "park_staff", "carrier_staff"):
caps.append("carrier")
else:
# 优先查统一表:organization_members.role=admin + org.type=carrier
admin_orgs = (await db.session.scalars(
select(Organization.id).join(
OrganizationMember, OrganizationMember.org_id == Organization.id
).where(
OrganizationMember.user_id == user.get("id", ""),
OrganizationMember.role == "admin",
OrganizationMember.status == "active",
Organization.type == "carrier",
)
)).all()
if admin_orgs:
caps.append("carrier")
else:
# 回退查旧表
bound = await db.session.scalar(select(ParkMember).where(
ParkMember.user_id == user.get("id", ""), ParkMember.park_id != "",
ParkMember.member_type == "admin", ParkMember.status == "active").limit(1))
if not bound:
bound = await db.session.scalar(select(ParkTenant).where(
ParkTenant.operator_user_id == user.get("id", "")).limit(1))
if bound:
caps.append("carrier")
# enterprise 能力:优先查统一表,回退查旧表
ent_orgs = (await db.session.scalars(
select(Organization.id).join(
OrganizationMember, OrganizationMember.org_id == Organization.id
).where(
OrganizationMember.user_id == user.get("id", ""),
OrganizationMember.role == "admin",
OrganizationMember.status == "active",
Organization.type == "enterprise",
)
)).all()
if ent_orgs:
caps.append("enterprise")
else:
ent = await db.session.scalar(select(CompanyMember).where(
CompanyMember.user_id == user.get("id", ""), CompanyMember.is_admin.is_(True),
CompanyMember.status == "active").limit(1))
if ent:
caps.append("enterprise")
return caps
async def get_user_organizations(db: Database, user_id: str) -> list[dict]:
"""查询用户绑定的所有组织(优先 organization_members,回退 users.org_id)。
返回列表,每项包含:org_id, name, type, role, is_admin
"""
from ..infrastructure.models import Organization, OrganizationMember
from sqlalchemy import select
orgs = []
# 优先查统一表
members = (await db.session.scalars(
select(OrganizationMember).where(
OrganizationMember.user_id == user_id,
OrganizationMember.status == "active",
)
)).all()
for m in members:
org = await db.session.get(Organization, m.org_id)
if org:
orgs.append({
"org_id": org.id, "name": org.name, "type": org.type,
"role": m.role, "is_admin": m.is_admin,
})
# 回退:users.org_id
if not orgs:
user = await db.users.get_by_id(user_id)
if user and user.get("org_id"):
org = await db.session.get(Organization, user["org_id"])
if org:
orgs.append({
"org_id": org.id, "name": org.name, "type": org.type,
"role": "admin", "is_admin": True,
})
return orgs
async def get_user_primary_org(db: Database, user: dict, org_type: str | None = None) -> dict | None:
"""获取用户的主组织(优先 organization_members 中的管理员,可按类型过滤)。"""
orgs = await get_user_organizations(db, user.get("id", ""))
if org_type:
orgs = [o for o in orgs if o["type"] == org_type]
# 优先管理员
admins = [o for o in orgs if o.get("is_admin")]
return admins[0] if admins else (orgs[0] if orgs else None)
async def _resolve_identity(user: dict, db: Database, payload: dict | None = None) -> dict:
"""叠加制能力模型:role 为基础角色,capabilities 为绑定叠加出的能力集合。
capabilities 优先从 JWT payload 读取(登录时已计算并写入),避免每次请求查库;
若 JWT 中无 capabilities(旧 token),则回退到实时计算。
"""
from ..domain.account_types import account_type, permission_role
atype = account_type(user.get("role"))
user["role"] = atype # 归一化为三种账号类型之一
user["permissions"] = await db.roles.permissions_for(permission_role(atype), None)
user["scope_region_ids"] = await db.regions.visible_region_ids(user.get("region_id"))
user["scope_level"] = await db.regions.level(user.get("region_id"))
if payload and payload.get("capabilities"):
user["capabilities"] = payload["capabilities"]
else:
user["capabilities"] = await _compute_capabilities(db, user)
return user
async def get_current_user(
request: Request,
db: Database = Depends(get_db),
) -> dict:
"""校验调用方 JWT,返回当前用户记录(含 role/权限/数据范围)。"""
token = extract_bearer_token(request)
if not token:
raise HTTPException(status_code=401, detail="No token provided")
payload = decode_access_token(token)
if payload is None:
raise HTTPException(status_code=401, detail="Invalid or expired token")
user_id = payload.get("sub")
if not await db.tokens.session_valid(
payload.get("jti", ""),
user_id,
payload.get("ver", 0),
):
raise HTTPException(status_code=401, detail="Invalid or expired token")
user = await db.users.get_by_id(user_id)
if user is None or user.get("status") != "active":
raise HTTPException(status_code=401, detail="Invalid or expired token")
return await _resolve_identity(user, db, payload)
async def optional_current_user(
request: Request,
db: Database = Depends(get_db),
) -> dict | None:
"""可选登录:有有效 JWT 则返回用户,否则返回 None(不报错)。"""
token = extract_bearer_token(request)
if not token:
return None
payload = decode_access_token(token)
if payload is None:
return None
if not await db.tokens.session_valid(
payload.get("jti", ""), payload.get("sub", ""), payload.get("ver", 0),
):
return None
user = await db.users.get_by_id(payload.get("sub"))
if user is None or user.get("status") != "active":
return None
return await _resolve_identity(user, db, payload)
def require_port(user: dict = Depends(get_current_user)) -> dict:
"""按端口隔离的资源(如智能体)—— 单角色后不再按端口隔离,仅要求已登录。
账号唯一角色:智能体不再按端口隔离,统一归到该账号(port="")。为兼容
既有 AgentRepository 的 port 过滤,这里补一个恒定空串。
"""
user.setdefault("port", "")
return user