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", "租户列表接口尚未配置", 503) async def precheck(self, username: str = "", tenant_code: 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 self._tenants_cache: tuple[float, list[dict[str, str]]] | None = None 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() ) return IdentityContext( user_id=int(raw_user_id), username=username, fullname=str(user.get("fullname") or username), tenant_uuid=tenant_uuid, roles=roles, ) async def tenants(self) -> list[dict[str, str]]: now = time.monotonic() if self._tenants_cache and self._tenants_cache[0] > now: return [dict(row) for row in self._tenants_cache[1]] try: data = await self._get_client().tenants() except JmsAuthClientError as exc: raise self._client_error(exc) from exc rows: list[dict[str, str]] = [] seen: set[str] = set() for tenant in data: code = str(tenant.get("code") or "").strip() name = str(tenant.get("name") or "").strip() if ( not code or not name or code in seen or str(tenant.get("state", "0")) != "0" or str(tenant.get("deleted", "0")) != "0" ): continue if self.tenant_code and code != self.tenant_code: continue seen.add(code) display_name = self.tenant_name if self.tenant_code == code and self.tenant_name else name rows.append({"code": code, "name": display_name}) if (not self.tenant_code or self.tenant_code == "platform") and "platform" not in seen: platform_name = self.tenant_name if self.tenant_code == "platform" and self.tenant_name else "平台" rows.append({"code": "platform", "name": platform_name}) if self.tenant_code and self.tenant_name and not rows: rows.append({"code": self.tenant_code, "name": self.tenant_name}) rows.sort(key=lambda row: (row["name"].lower(), row["code"].lower())) self._tenants_cache = (now + 300, rows) return [dict(row) for row in rows] async def _resolve_tenant(self, tenant_code: str) -> dict[str, str]: selected_code = (self.tenant_code or tenant_code).strip() if not selected_code: raise AuthError("AUTH_TENANT_REQUIRED", "请选择所属企业", 400) if self.tenant_code and tenant_code and tenant_code != self.tenant_code: raise AuthError("AUTH_TENANT_INVALID", "所选企业与系统配置不一致", 400) tenant = next((row for row in await self.tenants() if row["code"] == selected_code), None) if tenant is None: raise AuthError("AUTH_TENANT_INVALID", "所选企业不可用,请重新选择", 400) return tenant async def precheck(self, username: str = "", tenant_code: str = "") -> dict[str, Any]: tenant = await self._resolve_tenant(tenant_code) 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 = await self._resolve_tenant(str(payload.get("tenantCode") 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 identity = self._identity( user, remote_tenant, expected_tenant_code=selected_tenant["code"], ) return identity, self._issue_session( token_name, token_value, selected_tenant["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