# -*- 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 会立即抛 CancelledError,SQLAlchemy 的 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