deerflow-code/skills-data/xdfx-step5-audience-strategy/scripts/xdfx_auth.py
2026-09-07 18:24:55 +08:00

145 lines
4.4 KiB
Python

#!/usr/bin/env python3
"""Authentication helper for xdfx skill scripts.
The caller should not pass tokens through command-line arguments. This module
owns the fixed admin login flow and returns a Bearer token before business
interfaces are called.
"""
from __future__ import annotations
import json
import re
import ssl
import urllib.error
import urllib.request
from dataclasses import dataclass
from typing import Any
@dataclass
class AuthResult:
token: str
url: str
status: int | None
raw: str
class AuthError(RuntimeError):
def __init__(
self,
message: str,
*,
url: str,
status: int | None = None,
response: str = "",
exception: BaseException | None = None,
) -> None:
super().__init__(message)
self.url = url
self.status = status
self.response = response
self.exception = exception
def auth_header_value(token: str) -> str:
token = token.strip()
return token if token.lower().startswith("bearer ") else f"Bearer {token}"
def redact_url(url: str) -> str:
return re.sub(r"([?&](?:password|token|authToken|access_token)=)[^&]+", r"\1***", url, flags=re.IGNORECASE)
def ssl_context(*, verify_ssl: bool) -> ssl.SSLContext | None:
if verify_ssl:
return None
return ssl._create_unverified_context()
def build_login_url(api_base: str, path: str) -> str:
return path if re.match(r"^https?://", path, flags=re.IGNORECASE) else api_base.rstrip("/") + "/" + path.lstrip("/")
def extract_access_token(data: Any) -> str | None:
if isinstance(data, dict):
for key in ("access_token", "accessToken", "token", "jwt", "id_token", "authorization", "Authorization"):
value = data.get(key)
if isinstance(value, str) and value.strip():
return value.strip()
for value in data.values():
nested = extract_access_token(value)
if nested:
return nested
elif isinstance(data, list):
for value in data:
nested = extract_access_token(value)
if nested:
return nested
return None
def parse_json_loose(text: str) -> Any:
try:
return json.loads(text)
except Exception:
pass
match = re.search(r"[\[{]", text or "")
if not match:
return {}
start = match.start()
for end in range(len(text), start, -1):
try:
return json.loads(text[start:end])
except Exception:
continue
return {}
def fetch_access_token(
*,
api_base: str,
username: str = "admin",
password: str = "ch@user",
login_path: str = "/api/auth/login",
timeout: float = 90,
verify_ssl: bool = False,
) -> AuthResult:
"""Login once and return the extracted access token."""
url = build_login_url(api_base, login_path)
payload = {"username": username, "password": password}
data = json.dumps(payload, ensure_ascii=False).encode("utf-8")
headers = {"Content-Type": "application/json", "Accept": "application/json"}
req = urllib.request.Request(url, data=data, headers=headers, method="POST")
try:
with urllib.request.urlopen(req, timeout=timeout, context=ssl_context(verify_ssl=verify_ssl)) as resp:
raw = resp.read().decode("utf-8", errors="replace")
parsed = parse_json_loose(raw)
token = extract_access_token(parsed)
if not token:
raise AuthError(
"登录接口已响应但未找到 access_token",
url=redact_url(url),
status=getattr(resp, "status", None),
response=raw,
)
return AuthResult(token=token, url=redact_url(url), status=getattr(resp, "status", None), raw=raw)
except urllib.error.HTTPError as exc:
body = exc.read().decode("utf-8", errors="replace") if exc.fp else ""
raise AuthError(
"登录接口返回错误",
url=redact_url(url),
status=exc.code,
response=body,
exception=exc,
) from exc
except AuthError:
raise
except Exception as exc: # noqa: BLE001 - auth errors must preserve detail for user feedback.
raise AuthError(
"登录接口调用异常",
url=redact_url(url),
response=str(exc),
exception=exc,
) from exc