396 lines
17 KiB
Python
396 lines
17 KiB
Python
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
|