374 lines
17 KiB
Python
374 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", "租户列表接口尚未配置", 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
|