aps-agent/server/auth/providers.py

396 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", "已取消租户列表接口,请直接输入企业名称登录", 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