aps-agent/server/auth/providers.py

374 lines
17 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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