from __future__ import annotations import base64 import hashlib import hmac import json import os import re import time from abc import ABC, abstractmethod from typing import Any from fastapi import Request from server.auth.context import IdentityContext from server.integrations.jms_auth_client import JmsAuthClient, JmsAuthClientError class AuthError(Exception): def __init__(self, code: str, message: str, status_code: int = 401) -> None: super().__init__(message) self.code = code self.message = message self.status_code = status_code class AuthProvider(ABC): mode = "unconfigured" @abstractmethod async def login(self, payload: dict[str, Any]) -> tuple[IdentityContext, str]: raise NotImplementedError @abstractmethod async def authenticate(self, request: Request) -> IdentityContext: raise NotImplementedError async def tenants(self) -> list[dict[str, str]]: raise AuthError("AUTH_NOT_CONFIGURED", "已取消租户列表接口,请直接输入企业名称登录", 410) async def precheck(self, username: str = "", tenant_code: str = "", tenant_name: str = "") -> dict[str, Any]: raise AuthError("AUTH_NOT_CONFIGURED", "登录预检接口尚未配置", 503) async def captcha(self) -> dict[str, str]: raise AuthError("AUTH_NOT_CONFIGURED", "验证码接口尚未配置", 503) async def refresh(self, request: Request) -> tuple[IdentityContext, str]: await self.authenticate(request) raise AuthError("AUTH_NOT_CONFIGURED", "刷新接口尚未配置", 503) async def logout(self, request: Request) -> None: return None async def search_users(self, query: str, identity: IdentityContext) -> list[dict[str, Any]]: raise AuthError("AUTH_USER_DIRECTORY_NOT_CONFIGURED", "租户用户查询接口尚未接入", 503) class UnconfiguredAuthProvider(AuthProvider): async def login(self, payload: dict[str, Any]) -> tuple[IdentityContext, str]: raise AuthError("AUTH_NOT_CONFIGURED", "JMS 用户认证尚未配置", 503) async def authenticate(self, request: Request) -> IdentityContext: raise AuthError("AUTH_NOT_CONFIGURED", "JMS 用户认证尚未配置", 503) def _b64encode(raw: bytes) -> str: return base64.urlsafe_b64encode(raw).decode("ascii").rstrip("=") def _b64decode(raw: str) -> bytes: return base64.urlsafe_b64decode(raw + "=" * (-len(raw) % 4)) class JmsAuthProvider(AuthProvider): mode = "jms" _TOKEN_NAME_RE = re.compile(r"^[A-Za-z0-9-]{1,64}$") _PROHIBITED_TOKEN_NAMES = { "connection", "content-length", "cookie", "host", "proxy-authorization", "transfer-encoding", } def __init__(self, client: JmsAuthClient | None = None) -> None: self.client: JmsAuthClient | None = client self.client_config_error: AuthError | None = None if self.client is None: try: self.client = JmsAuthClient() except JmsAuthClientError as exc: self.client_config_error = AuthError(exc.code, exc.message, exc.status_code) self.tenant_code = (os.environ.get("JMS_AUTH_TENANT_CODE") or "").strip() self.tenant_name = (os.environ.get("JMS_AUTH_TENANT_NAME") or "").strip() self.expected_token_name = (os.environ.get("JMS_AUTH_TOKEN_NAME") or "").strip() self.session_secret = ( os.environ.get("JMS_AUTH_SESSION_SECRET") or os.environ.get("APS_AUTH_SESSION_SECRET") or "" ).encode("utf-8") try: self.ttl_seconds = int(os.environ.get("APS_AUTH_TTL_SECONDS") or "28800") except ValueError: self.ttl_seconds = 0 def _get_client(self) -> JmsAuthClient: if self.client_config_error is not None: raise self.client_config_error if self.client is None: raise AuthError("AUTH_CONFIG_INVALID", "JMS 认证客户端未正确配置", 503) return self.client def _require_login_config(self) -> None: missing = [] if len(self.session_secret) < 32: missing.append("JMS_AUTH_SESSION_SECRET(至少 32 个字符)") if self.ttl_seconds <= 0: missing.append("APS_AUTH_TTL_SECONDS(必须为正整数)") if missing: raise AuthError( "AUTH_CONFIG_INVALID", f"JMS 认证配置缺失:{', '.join(missing)}", 503, ) self._get_client() def _validate_token_name(self, token_name: str) -> None: if not self._TOKEN_NAME_RE.fullmatch(token_name): raise AuthError("AUTH_UPSTREAM_INVALID", "JMS 登录凭证名称无效", 502) if token_name.lower() in self._PROHIBITED_TOKEN_NAMES: raise AuthError("AUTH_UPSTREAM_INVALID", "JMS 登录凭证名称不安全", 502) if self.expected_token_name and token_name.lower() != self.expected_token_name.lower(): raise AuthError("AUTH_UPSTREAM_INVALID", "JMS 登录凭证名称与本地配置不一致", 502) @staticmethod def _request_token(request: Request) -> str: authorization = request.headers.get("authorization") or "" if authorization.lower().startswith("bearer "): return authorization[7:].strip() return request.cookies.get("aps_session") or "" def _issue_session(self, token_name: str, token_value: str, tenant_code: str) -> str: self._validate_token_name(token_name) payload = { "tokenName": token_name, "tokenValue": token_value, "tenantCode": tenant_code, "expiresAt": int(time.time()) + self.ttl_seconds, } raw = json.dumps(payload, separators=(",", ":"), ensure_ascii=True).encode("utf-8") body = _b64encode(raw) signature = _b64encode(hmac.new(self.session_secret, body.encode("ascii"), hashlib.sha256).digest()) return f"{body}.{signature}" def _read_session(self, request: Request) -> tuple[str, str, str]: token = self._request_token(request) if not token: raise AuthError("AUTH_REQUIRED", "请先登录", 401) if len(self.session_secret) < 32: raise AuthError("AUTH_CONFIG_INVALID", "JMS 会话密钥未正确配置", 503) try: body, signature = token.split(".", 1) expected = _b64encode( hmac.new(self.session_secret, body.encode("ascii"), hashlib.sha256).digest() ) if not hmac.compare_digest(signature, expected): raise ValueError("signature") payload = json.loads(_b64decode(body)) if int(payload.get("expiresAt") or 0) <= int(time.time()): raise AuthError("AUTH_EXPIRED", "登录状态已过期", 401) token_name = str(payload["tokenName"]) token_value = str(payload["tokenValue"]) tenant_code = str(payload["tenantCode"]) if not token_value or not tenant_code: raise ValueError("token") self._validate_token_name(token_name) return token_name, token_value, tenant_code except AuthError: raise except Exception as exc: raise AuthError("AUTH_INVALID", "登录凭证无效", 401) from exc @staticmethod def _client_error(exc: JmsAuthClientError, *, login: bool = False) -> AuthError: if login and exc.status_code in {400, 401, 403}: return AuthError("AUTH_LOGIN_FAILED", exc.message, 401) if exc.status_code in {400, 401, 403}: return AuthError("AUTH_INVALID", exc.message, 401) return AuthError(exc.code, exc.message, exc.status_code) def _identity( self, user: dict[str, Any], tenant: dict[str, Any], *, expected_tenant_code: str, ) -> IdentityContext: if str(user.get("state", "0")) != "0": raise AuthError("AUTH_USER_DISABLED", "当前用户已被停用", 403) if str(tenant.get("state", "0")) != "0": raise AuthError("AUTH_TENANT_DISABLED", "当前租户已被停用", 403) raw_user_id = str(user.get("id") or "") if not raw_user_id.isdigit() or int(raw_user_id) <= 0 or int(raw_user_id) > 2**63 - 1: raise AuthError("AUTH_UPSTREAM_INVALID", "JMS 用户标识格式无效", 502) tenant_uuid = str(tenant.get("uuid") or "").strip() username = str(user.get("username") or "").strip() if not tenant_uuid or len(tenant_uuid) > 32 or not username: raise AuthError("AUTH_UPSTREAM_INVALID", "JMS 用户或租户信息不完整", 502) returned_tenant_code = str(tenant.get("code") or "").strip() if returned_tenant_code and returned_tenant_code != expected_tenant_code: raise AuthError("AUTH_TENANT_MISMATCH", "JMS 登录租户与本地配置不一致", 403) roles = tuple( str(role.get("code")).strip() for role in (user.get("roles") or []) if isinstance(role, dict) and str(role.get("code") or "").strip() ) if not roles: # JMS profile 不返回 roles 字段(真实平台现场已验证)→ 默认给计划员 # 角色,保证登录用户可发起/审批 P2 门禁(导入确认/发布/写主数据); # 可用 JMS_AUTH_DEFAULT_ROLE 覆盖或留空关闭。 default_role = (os.environ.get("JMS_AUTH_DEFAULT_ROLE") or "planner").strip() if default_role: roles = (default_role,) return IdentityContext( user_id=int(raw_user_id), username=username, fullname=str(user.get("fullname") or username), tenant_uuid=tenant_uuid, roles=roles, ) # 常见展示名 → 技术编码(不再依赖公开租户列表) _TENANT_NAME_ALIASES = { "平台": "platform", "platform": "platform", } @classmethod def _alias_tenant_code(cls, name_or_code: str) -> str: key = name_or_code.strip() if not key: return "" return cls._TENANT_NAME_ALIASES.get(key) or cls._TENANT_NAME_ALIASES.get(key.lower()) or key def _select_tenant(self, tenant_code: str = "", tenant_name: str = "") -> dict[str, str]: """用手填企业名/编码解析登录租户,不调用公开租户列表。""" raw_name = (tenant_name or "").strip() raw_code = (tenant_code or "").strip() if self.tenant_code: if raw_code and raw_code != self.tenant_code: raise AuthError("AUTH_TENANT_INVALID", "所选企业与系统配置不一致", 400) if raw_name: alias = self._alias_tenant_code(raw_name) if ( self.tenant_name and raw_name != self.tenant_name and alias != self.tenant_code ): raise AuthError("AUTH_TENANT_INVALID", "所选企业与系统配置不一致", 400) return { "code": self.tenant_code, "name": self.tenant_name or raw_name or self.tenant_code, } if raw_name: return { "code": raw_code or self._alias_tenant_code(raw_name), "name": raw_name, } if raw_code: display = "平台" if raw_code == "platform" else raw_code return {"code": raw_code, "name": display} raise AuthError("AUTH_TENANT_REQUIRED", "请输入所属企业名称", 400) async def tenants(self) -> list[dict[str, str]]: raise AuthError( "AUTH_TENANT_LIST_REMOVED", "已取消租户列表接口,请直接输入企业名称登录", 410, ) async def precheck( self, username: str = "", tenant_code: str = "", tenant_name: str = "", ) -> dict[str, Any]: tenant = self._select_tenant(tenant_code, tenant_name) try: data = await self._get_client().precheck( tenant_code=tenant["code"], username=username.strip(), ) except JmsAuthClientError as exc: raise self._client_error(exc) from exc return {"requireCaptcha": bool(data.get("requireCaptcha"))} async def captcha(self) -> dict[str, str]: try: data = await self._get_client().captcha() except JmsAuthClientError as exc: raise self._client_error(exc) from exc captcha_id = str(data.get("captchaId") or "") image = str(data.get("image") or "") if not captcha_id or not image.startswith("data:image/png;base64,"): raise AuthError("AUTH_UPSTREAM_INVALID", "JMS 验证码响应不完整", 502) return {"captchaId": captcha_id, "image": image} async def login(self, payload: dict[str, Any]) -> tuple[IdentityContext, str]: self._require_login_config() if str(payload.get("method") or "password") != "password": raise AuthError("AUTH_METHOD_UNSUPPORTED", "当前仅支持账号密码登录", 400) username = str(payload.get("username") or "").strip() password = str(payload.get("password") or "") selected_tenant = self._select_tenant( str(payload.get("tenantCode") or ""), str(payload.get("tenantName") or ""), ) captcha_id = str(payload.get("captchaId") or "").strip() captcha_code = str(payload.get("captchaCode") or "").strip() if not username or not password: raise AuthError("AUTH_CREDENTIALS_REQUIRED", "请输入账号和密码", 400) if bool(captcha_id) != bool(captcha_code): raise AuthError("AUTH_CAPTCHA_REQUIRED", "请输入完整的图形验证码", 400) try: result = await self._get_client().login( tenant_code=selected_tenant["code"], tenant_name=selected_tenant["name"], username=username, password=password, captcha_id=captcha_id, captcha_code=captcha_code, ) except JmsAuthClientError as exc: raise self._client_error(exc, login=True) from exc token_name = str(result.get("tokenName") or "").strip() token_value = str(result.get("tokenValue") or "").strip() if not token_name or not token_value: raise AuthError("AUTH_UPSTREAM_INVALID", "JMS 登录响应缺少凭证", 502) self._validate_token_name(token_name) action = str(result.get("action") or "").strip() if action: try: await self._get_client().logout(token_name=token_name, token_value=token_value) except JmsAuthClientError: pass raise AuthError("AUTH_ACTION_REQUIRED", f"账号需要先完成安全操作:{action}", 409) try: client = self._get_client() user = await client.profile(token_name=token_name, token_value=token_value) remote_tenant = await client.tenant(token_name=token_name, token_value=token_value) except JmsAuthClientError as exc: raise self._client_error(exc) from exc # 登录成功后以 JMS 返回的租户编码为准(手填企业名可能只是展示名) remote_code = str(remote_tenant.get("code") or "").strip() or selected_tenant["code"] if self.tenant_code and remote_code != self.tenant_code: raise AuthError("AUTH_TENANT_MISMATCH", "JMS 登录租户与本地配置不一致", 403) identity = self._identity( user, remote_tenant, expected_tenant_code=remote_code, ) return identity, self._issue_session( token_name, token_value, remote_code, ) async def authenticate(self, request: Request) -> IdentityContext: token_name, token_value, tenant_code = self._read_session(request) try: client = self._get_client() user = await client.profile(token_name=token_name, token_value=token_value) tenant = await client.tenant(token_name=token_name, token_value=token_value) except JmsAuthClientError as exc: raise self._client_error(exc) from exc return self._identity(user, tenant, expected_tenant_code=tenant_code) async def refresh(self, request: Request) -> tuple[IdentityContext, str]: token_name, token_value, tenant_code = self._read_session(request) identity = await self.authenticate(request) return identity, self._issue_session(token_name, token_value, tenant_code) async def logout(self, request: Request) -> None: try: token_name, token_value, _ = self._read_session(request) await self._get_client().logout(token_name=token_name, token_value=token_value) except (AuthError, JmsAuthClientError): # Local logout must always succeed even when the remote session has expired. return None _provider: AuthProvider | None = None _provider_mode: str | None = None def get_auth_provider() -> AuthProvider: global _provider, _provider_mode mode = (os.environ.get("APS_AUTH_PROVIDER") or "unconfigured").strip().lower() if _provider is None or _provider_mode != mode: _provider = JmsAuthProvider() if mode == "jms" else UnconfiguredAuthProvider() _provider_mode = mode return _provider