aps-agent/server/auth/providers.py

396 lines
17 KiB
Python
Raw Normal View History

from __future__ import annotations
import base64
import hashlib
import hmac
import json
import os
2026-07-29 23:22:40 +08:00
import re
import time
from abc import ABC, abstractmethod
from typing import Any
from fastapi import Request
from server.auth.context import IdentityContext
2026-07-29 23:22:40 +08:00
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
2026-07-29 23:22:40 +08:00
async def tenants(self) -> list[dict[str, str]]:
raise AuthError("AUTH_NOT_CONFIGURED", "已取消租户列表接口,请直接输入企业名称登录", 410)
2026-07-29 23:22:40 +08:00
async def precheck(self, username: str = "", tenant_code: str = "", tenant_name: str = "") -> dict[str, Any]:
2026-07-29 23:22:40 +08:00
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]:
2026-07-29 23:22:40 +08:00
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]]:
2026-07-29 23:22:40 +08:00
raise AuthError("AUTH_USER_DIRECTORY_NOT_CONFIGURED", "租户用户查询接口尚未接入", 503)
class UnconfiguredAuthProvider(AuthProvider):
async def login(self, payload: dict[str, Any]) -> tuple[IdentityContext, str]:
2026-07-29 23:22:40 +08:00
raise AuthError("AUTH_NOT_CONFIGURED", "JMS 用户认证尚未配置", 503)
async def authenticate(self, request: Request) -> IdentityContext:
2026-07-29 23:22:40 +08:00
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))
2026-07-29 23:22:40 +08:00
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",
}
2026-07-29 23:22:40 +08:00
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()
2026-07-29 23:22:40 +08:00
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)
2026-07-29 23:22:40 +08:00
@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 ""
2026-07-29 23:22:40 +08:00
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)
2026-07-29 23:22:40 +08:00
signature = _b64encode(hmac.new(self.session_secret, body.encode("ascii"), hashlib.sha256).digest())
return f"{body}.{signature}"
2026-07-29 23:22:40 +08:00
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)
2026-07-29 23:22:40 +08:00
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))
2026-07-29 23:22:40 +08:00
if int(payload.get("expiresAt") or 0) <= int(time.time()):
raise AuthError("AUTH_EXPIRED", "登录状态已过期", 401)
2026-07-29 23:22:40 +08:00
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
2026-07-29 23:22:40 +08:00
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,)
2026-07-29 23:22:40 +08:00
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)
2026-07-29 23:22:40 +08:00
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)
2026-07-29 23:22:40 +08:00
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]:
2026-07-29 23:22:40 +08:00
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 ""),
)
2026-07-29 23:22:40 +08:00
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
2026-07-29 23:22:40 +08:00
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)
2026-07-29 23:22:40 +08:00
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)
2026-07-29 23:22:40 +08:00
identity = self._identity(
user, remote_tenant, expected_tenant_code=remote_code,
2026-07-29 23:22:40 +08:00
)
return identity, self._issue_session(
token_name, token_value, remote_code,
)
2026-07-29 23:22:40 +08:00
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)
2026-07-29 23:22:40 +08:00
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:
2026-07-29 23:22:40 +08:00
_provider = JmsAuthProvider() if mode == "jms" else UnconfiguredAuthProvider()
_provider_mode = mode
return _provider