11
This commit is contained in:
+400
-161
@@ -1,191 +1,430 @@
|
||||
# core/request.py
|
||||
import json
|
||||
import aiohttp
|
||||
import asyncio
|
||||
from typing import Optional, Dict, Any, Union, List
|
||||
from aiohttp import ClientTimeout, ClientSession
|
||||
import json
|
||||
import random
|
||||
from collections.abc import Mapping
|
||||
from email.utils import parsedate_to_datetime
|
||||
from time import monotonic, time
|
||||
from typing import Any
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
import aiohttp
|
||||
from aiohttp import ClientSession, ClientTimeout
|
||||
|
||||
from astrbot.api import logger
|
||||
|
||||
class APIClient:
|
||||
"""
|
||||
API客户端类
|
||||
|
||||
优化说明:
|
||||
1. 复用 aiohttp.ClientSession 以提高性能。
|
||||
2. 增加类型提示 (Type Hints)。
|
||||
3. 支持异步上下文管理器 (Async Context Manager)。
|
||||
"""
|
||||
SENSITIVE_KEYS = {
|
||||
"access_token",
|
||||
"api_key",
|
||||
"authorization",
|
||||
"cookie",
|
||||
"jx3api_ticket",
|
||||
"jx3api_token",
|
||||
"password",
|
||||
"refresh_token",
|
||||
"secret",
|
||||
"ticket",
|
||||
"token",
|
||||
}
|
||||
|
||||
def __init__(self, base_timeout: int = 10, ssl_verify: bool = False):
|
||||
self.base_timeout = base_timeout
|
||||
RETRYABLE_STATUS_CODES = {429, 500, 502, 503, 504}
|
||||
|
||||
|
||||
class APIClientError(Exception):
|
||||
"""HTTP 客户端的结构化异常。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
*,
|
||||
method: str = "",
|
||||
url: str = "",
|
||||
status: int | None = None,
|
||||
retryable: bool = False,
|
||||
):
|
||||
super().__init__(message)
|
||||
self.method = method
|
||||
self.url = url
|
||||
self.status = status
|
||||
self.retryable = retryable
|
||||
|
||||
|
||||
class APIResponseError(APIClientError):
|
||||
"""服务端返回了错误状态或无法解析的响应。"""
|
||||
|
||||
|
||||
class APIResponseTooLargeError(APIResponseError):
|
||||
"""响应体超过客户端允许的大小。"""
|
||||
|
||||
|
||||
class APIResponseFormatError(APIResponseError):
|
||||
"""响应结构与调用方约定不一致。"""
|
||||
|
||||
|
||||
def redact_mapping(data: Mapping[str, Any] | None) -> dict[str, Any] | None:
|
||||
"""递归清理日志中的敏感字段。"""
|
||||
if data is None:
|
||||
return None
|
||||
|
||||
redacted: dict[str, Any] = {}
|
||||
for key, value in data.items():
|
||||
key_text = str(key)
|
||||
if key_text.lower() in SENSITIVE_KEYS:
|
||||
redacted[key_text] = "***"
|
||||
elif isinstance(value, Mapping):
|
||||
redacted[key_text] = redact_mapping(value)
|
||||
elif isinstance(value, list):
|
||||
redacted[key_text] = [
|
||||
redact_mapping(item) if isinstance(item, Mapping) else item
|
||||
for item in value
|
||||
]
|
||||
else:
|
||||
redacted[key_text] = value
|
||||
return redacted
|
||||
|
||||
|
||||
def safe_url(url: str) -> str:
|
||||
"""移除 URL 中的用户名、密码、查询参数和片段,仅用于日志。"""
|
||||
parsed = urlsplit(url)
|
||||
host = parsed.hostname or ""
|
||||
if parsed.port:
|
||||
host = f"{host}:{parsed.port}"
|
||||
return urlunsplit((parsed.scheme, host, parsed.path, "", ""))
|
||||
|
||||
|
||||
class APIClient:
|
||||
"""可复用、带安全与可靠性保护的 aiohttp 客户端。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_timeout: float = 15,
|
||||
*,
|
||||
connect_timeout: float = 5,
|
||||
read_timeout: float = 10,
|
||||
ssl_verify: bool = True,
|
||||
max_retries: int = 2,
|
||||
max_response_bytes: int = 10 * 1024 * 1024,
|
||||
max_binary_response_bytes: int = 20 * 1024 * 1024,
|
||||
connection_limit: int = 50,
|
||||
connection_limit_per_host: int = 10,
|
||||
log_payloads: bool = False,
|
||||
):
|
||||
if base_timeout <= 0 or connect_timeout <= 0 or read_timeout <= 0:
|
||||
raise ValueError("请求超时时间必须大于 0")
|
||||
if max_retries < 0:
|
||||
raise ValueError("max_retries 不能小于 0")
|
||||
if max_response_bytes <= 0 or max_binary_response_bytes <= 0:
|
||||
raise ValueError("响应大小限制必须大于 0")
|
||||
|
||||
self.timeout = ClientTimeout(
|
||||
total=base_timeout,
|
||||
connect=connect_timeout,
|
||||
sock_connect=connect_timeout,
|
||||
sock_read=read_timeout,
|
||||
)
|
||||
self.ssl_verify = ssl_verify
|
||||
self._session: Optional[ClientSession] = None
|
||||
self.max_retries = max_retries
|
||||
self.max_response_bytes = max_response_bytes
|
||||
self.max_binary_response_bytes = max_binary_response_bytes
|
||||
self.connection_limit = connection_limit
|
||||
self.connection_limit_per_host = connection_limit_per_host
|
||||
self.log_payloads = log_payloads
|
||||
|
||||
self._session: ClientSession | None = None
|
||||
self._session_lock = asyncio.Lock()
|
||||
|
||||
async def start(self) -> None:
|
||||
"""显式初始化连接池;重复调用是安全的。"""
|
||||
await self.get_session()
|
||||
|
||||
async def get_session(self) -> ClientSession:
|
||||
"""获取或创建单例 Session"""
|
||||
if self._session is None or self._session.closed:
|
||||
timeout = ClientTimeout(total=self.base_timeout)
|
||||
self._session = ClientSession(timeout=timeout)
|
||||
return self._session
|
||||
"""获取或并发安全地创建共享 Session。"""
|
||||
if self._session is not None and not self._session.closed:
|
||||
return self._session
|
||||
|
||||
async def close(self):
|
||||
"""关闭 Session"""
|
||||
if self._session and not self._session.closed:
|
||||
await self._session.close()
|
||||
async with self._session_lock:
|
||||
if self._session is None or self._session.closed:
|
||||
connector = aiohttp.TCPConnector(
|
||||
ssl=self.ssl_verify,
|
||||
limit=self.connection_limit,
|
||||
limit_per_host=self.connection_limit_per_host,
|
||||
ttl_dns_cache=300,
|
||||
)
|
||||
self._session = ClientSession(
|
||||
timeout=self.timeout,
|
||||
connector=connector,
|
||||
trust_env=True,
|
||||
)
|
||||
return self._session
|
||||
|
||||
async def close(self) -> None:
|
||||
"""关闭连接池;重复调用是安全的。"""
|
||||
async with self._session_lock:
|
||||
session = self._session
|
||||
self._session = None
|
||||
|
||||
if session is not None and not session.closed:
|
||||
await session.close()
|
||||
|
||||
async def __aenter__(self):
|
||||
await self.get_session()
|
||||
await self.start()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
await self.close()
|
||||
|
||||
async def _request(self, method: str, url: str, params: Optional[Dict] = None, json_data: Optional[Dict] = None) -> Any:
|
||||
"""
|
||||
统一的内部请求处理方法
|
||||
"""
|
||||
session = await self.get_session()
|
||||
async def _request(
|
||||
self,
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
params: Mapping[str, Any] | None = None,
|
||||
json_data: Any = None,
|
||||
form_data: Any = None,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
timeout: ClientTimeout | None = None,
|
||||
retry: bool | None = None,
|
||||
) -> Any:
|
||||
method = method.upper()
|
||||
|
||||
# 记录日志
|
||||
logger.debug(f"发起 {method} 请求: {url}")
|
||||
if params: logger.debug(f"Query参数: {params}")
|
||||
if json_data: logger.debug(f"Body数据: {json_data}")
|
||||
request_url = safe_url(url)
|
||||
started_at = monotonic()
|
||||
retry_enabled = method in {"GET", "HEAD", "OPTIONS"} if retry is None else retry
|
||||
retry_count = self.max_retries if retry_enabled else 0
|
||||
|
||||
try:
|
||||
# aiohttp 会自动处理 json=json_data 时的 Content-Type
|
||||
async with session.request(
|
||||
method=method,
|
||||
url=url,
|
||||
params=params,
|
||||
json=json_data,
|
||||
ssl=self.ssl_verify
|
||||
) as response:
|
||||
return await self._handle_response(response)
|
||||
|
||||
except aiohttp.ClientError as e:
|
||||
logger.error(f"网络请求出错 ({method} {url}): {e}")
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.error(f"未知错误 ({method} {url}): {e}")
|
||||
return None
|
||||
if self.log_payloads:
|
||||
params_log: Any = redact_mapping(params)
|
||||
json_log: Any = (
|
||||
redact_mapping(json_data)
|
||||
if isinstance(json_data, Mapping)
|
||||
else "<omitted>"
|
||||
if json_data is not None
|
||||
else None
|
||||
)
|
||||
else:
|
||||
params_log = sorted(str(key) for key in (params or {}))
|
||||
json_log = (
|
||||
sorted(str(key) for key in json_data)
|
||||
if isinstance(json_data, Mapping)
|
||||
else "<omitted>"
|
||||
if json_data is not None
|
||||
else None
|
||||
)
|
||||
logger.debug(
|
||||
f"发起 {method} 请求: {request_url}, "
|
||||
f"param_keys={params_log}, json_keys={json_log}"
|
||||
)
|
||||
|
||||
async def _handle_response(self, response: aiohttp.ClientResponse) -> Any:
|
||||
"""处理响应:自动识别二进制或JSON"""
|
||||
try:
|
||||
logger.debug(f"响应状态: {response.status}")
|
||||
response.raise_for_status()
|
||||
|
||||
content_type = response.headers.get('Content-Type', '').lower()
|
||||
|
||||
if 'image' in content_type or 'octet-stream' in content_type:
|
||||
return await response.read()
|
||||
retryable_exceptions = (
|
||||
aiohttp.ClientConnectionError,
|
||||
aiohttp.ServerTimeoutError,
|
||||
asyncio.TimeoutError,
|
||||
)
|
||||
|
||||
for attempt in range(retry_count + 1):
|
||||
session = await self.get_session()
|
||||
try:
|
||||
data = await response.json()
|
||||
except Exception:
|
||||
text = await response.text()
|
||||
async with session.request(
|
||||
method=method,
|
||||
url=url,
|
||||
params=params,
|
||||
json=json_data,
|
||||
data=form_data,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
) as response:
|
||||
if (
|
||||
response.status in RETRYABLE_STATUS_CODES
|
||||
and attempt < retry_count
|
||||
):
|
||||
delay = self._retry_delay(attempt, response)
|
||||
logger.warning(
|
||||
f"{method} {request_url} 返回 HTTP {response.status},"
|
||||
f"{delay:.2f}s 后重试 ({attempt + 1}/{retry_count})"
|
||||
)
|
||||
response.release()
|
||||
await asyncio.sleep(delay)
|
||||
continue
|
||||
|
||||
result = await self._handle_response(response, method, request_url)
|
||||
elapsed_ms = (monotonic() - started_at) * 1000
|
||||
logger.debug(
|
||||
f"{method} {request_url} 请求完成: "
|
||||
f"status={response.status}, elapsed={elapsed_ms:.0f}ms"
|
||||
)
|
||||
return result
|
||||
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except retryable_exceptions as exc:
|
||||
if attempt < retry_count:
|
||||
delay = self._retry_delay(attempt)
|
||||
logger.warning(
|
||||
f"{method} {request_url} 网络异常,{delay:.2f}s 后重试 "
|
||||
f"({attempt + 1}/{retry_count}): {type(exc).__name__}"
|
||||
)
|
||||
await asyncio.sleep(delay)
|
||||
continue
|
||||
raise APIClientError(
|
||||
f"网络请求失败: {type(exc).__name__}",
|
||||
method=method,
|
||||
url=request_url,
|
||||
retryable=True,
|
||||
) from exc
|
||||
except APIClientError:
|
||||
raise
|
||||
except aiohttp.ClientError as exc:
|
||||
raise APIClientError(
|
||||
f"HTTP 客户端错误: {type(exc).__name__}",
|
||||
method=method,
|
||||
url=request_url,
|
||||
retryable=False,
|
||||
) from exc
|
||||
|
||||
raise APIClientError(
|
||||
"请求在重试后仍然失败",
|
||||
method=method,
|
||||
url=request_url,
|
||||
retryable=True,
|
||||
)
|
||||
|
||||
def _retry_delay(
|
||||
self,
|
||||
attempt: int,
|
||||
response: aiohttp.ClientResponse | None = None,
|
||||
) -> float:
|
||||
if response is not None:
|
||||
retry_after = response.headers.get("Retry-After")
|
||||
if retry_after:
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
data = await loop.run_in_executor(None, json.loads, text)
|
||||
except json.JSONDecodeError:
|
||||
logger.error(f"无法解析响应为 JSON。原始内容: {text[:100]}...")
|
||||
return None
|
||||
return min(max(float(retry_after), 0.0), 60.0)
|
||||
except ValueError:
|
||||
try:
|
||||
retry_at = parsedate_to_datetime(retry_after)
|
||||
seconds = retry_at.timestamp() - time()
|
||||
return min(max(seconds, 0.0), 60.0)
|
||||
except (TypeError, ValueError, OverflowError):
|
||||
pass
|
||||
|
||||
logger.debug(f"响应数据: {data}")
|
||||
return self._validate_api_payload(data)
|
||||
base = min(0.5 * (2**attempt), 8.0)
|
||||
return base + random.uniform(0, base * 0.25)
|
||||
|
||||
except aiohttp.ClientError as e:
|
||||
logger.error(f"HTTP响应错误: {e}")
|
||||
async def _handle_response(
|
||||
self,
|
||||
response: aiohttp.ClientResponse,
|
||||
method: str,
|
||||
request_url: str,
|
||||
) -> Any:
|
||||
content_type = response.headers.get("Content-Type", "").lower()
|
||||
is_binary = "image/" in content_type or "octet-stream" in content_type
|
||||
size_limit = (
|
||||
self.max_binary_response_bytes if is_binary else self.max_response_bytes
|
||||
)
|
||||
|
||||
if response.status >= 400:
|
||||
await self._read_limited(response, self.max_response_bytes)
|
||||
raise APIResponseError(
|
||||
f"HTTP {response.status}",
|
||||
method=method,
|
||||
url=request_url,
|
||||
status=response.status,
|
||||
retryable=response.status in RETRYABLE_STATUS_CODES,
|
||||
)
|
||||
|
||||
body = await self._read_limited(response, size_limit)
|
||||
if is_binary:
|
||||
return body
|
||||
|
||||
if not body:
|
||||
return None
|
||||
|
||||
def _validate_api_payload(self, data: Any) -> Any:
|
||||
"""校验业务层面的 JSON 数据结构"""
|
||||
if not data:
|
||||
logger.error("API返回空数据")
|
||||
return None
|
||||
encoding = response.charset or "utf-8"
|
||||
try:
|
||||
text = body.decode(encoding)
|
||||
except (LookupError, UnicodeDecodeError) as exc:
|
||||
raise APIResponseFormatError(
|
||||
f"响应字符编码无效: {encoding}",
|
||||
method=method,
|
||||
url=request_url,
|
||||
status=response.status,
|
||||
) from exc
|
||||
|
||||
# 如果返回的是 JSON 字符串而非对象,再次解析
|
||||
if isinstance(data, str):
|
||||
try:
|
||||
data = json.loads(data)
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
if isinstance(data, dict) and 'code' in data:
|
||||
# 兼容多种成功状态码:200, "0", 0, 1
|
||||
code = data.get('code')
|
||||
if code not in [200, "0", 0, 1]:
|
||||
msg = data.get('msg') or data.get('message', '未知错误')
|
||||
logger.error(f"API业务报错: code={code}, msg={msg}")
|
||||
return None
|
||||
|
||||
try:
|
||||
return json.loads(text)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise APIResponseFormatError(
|
||||
"响应不是合法 JSON",
|
||||
method=method,
|
||||
url=request_url,
|
||||
status=response.status,
|
||||
) from exc
|
||||
|
||||
async def _read_limited(
|
||||
self,
|
||||
response: aiohttp.ClientResponse,
|
||||
limit: int,
|
||||
) -> bytes:
|
||||
content_length = response.content_length
|
||||
if content_length is not None and content_length > limit:
|
||||
raise APIResponseTooLargeError(
|
||||
f"响应体过大: {content_length} bytes,限制为 {limit} bytes",
|
||||
method=response.method,
|
||||
url=safe_url(str(response.url)),
|
||||
status=response.status,
|
||||
)
|
||||
|
||||
chunks: list[bytes] = []
|
||||
received = 0
|
||||
async for chunk in response.content.iter_chunked(64 * 1024):
|
||||
received += len(chunk)
|
||||
if received > limit:
|
||||
raise APIResponseTooLargeError(
|
||||
f"响应体超过限制: {limit} bytes",
|
||||
method=response.method,
|
||||
url=safe_url(str(response.url)),
|
||||
status=response.status,
|
||||
)
|
||||
chunks.append(chunk)
|
||||
return b"".join(chunks)
|
||||
|
||||
async def get(
|
||||
self,
|
||||
url: str,
|
||||
params: Mapping[str, Any] | None = None,
|
||||
*,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
timeout: ClientTimeout | None = None,
|
||||
retry: bool = True,
|
||||
) -> Any:
|
||||
data = await self._request(
|
||||
"GET",
|
||||
url,
|
||||
params=params,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
retry=retry,
|
||||
)
|
||||
return data
|
||||
|
||||
async def get(self, url: str, params: Optional[Dict] = None, out_key: Optional[str] = None) -> Any:
|
||||
"""GET 请求封装"""
|
||||
data = await self._request('GET', url, params=params)
|
||||
return self._extract_data(data, out_key)
|
||||
|
||||
async def post(self, url: str, data: Optional[Dict] = None, out_key: Optional[str] = None) -> Any:
|
||||
"""POST 请求封装 (默认发送 JSON)"""
|
||||
data = await self._request('POST', url, json_data=data)
|
||||
return self._extract_data(data, out_key)
|
||||
|
||||
def _extract_data(self, data: Any, key: Optional[str]) -> Any:
|
||||
"""辅助方法:从结果中提取指定字段"""
|
||||
if data is None:
|
||||
return None
|
||||
if isinstance(data, bytes):
|
||||
return data
|
||||
if key and isinstance(data, dict):
|
||||
return data.get(key, {})
|
||||
return data
|
||||
|
||||
async def all_pages(
|
||||
self,
|
||||
method: str,
|
||||
url: str,
|
||||
params_data: Optional[Dict] = None,
|
||||
out_key: str = "",
|
||||
list_key: str = "list",
|
||||
max_pages: int = 10
|
||||
) -> List[Any]:
|
||||
"""
|
||||
分页获取所有数据
|
||||
:param method: GET 或 POST
|
||||
:param list_key: 列表数据在 JSON 中的字段名,如 'data' 或 'list'
|
||||
"""
|
||||
all_data = []
|
||||
current_page = 1
|
||||
params = params_data.copy() if params_data else {}
|
||||
|
||||
while True:
|
||||
params["page"] = str(current_page)
|
||||
|
||||
if method.upper() == "POST":
|
||||
data = await self.post(url, data=params, out_key=out_key)
|
||||
else:
|
||||
data = await self.get(url, params=params, out_key=out_key)
|
||||
|
||||
# 终止条件判断
|
||||
if not data or isinstance(data, bytes):
|
||||
break
|
||||
|
||||
# 如果 data 是列表本身(有些API直接返回列表)
|
||||
page_items = data if isinstance(data, list) else data.get(list_key)
|
||||
|
||||
if not page_items:
|
||||
break
|
||||
|
||||
all_data.extend(page_items)
|
||||
|
||||
if current_page >= max_pages:
|
||||
break
|
||||
|
||||
current_page += 1
|
||||
logger.info(f"已获取第 {current_page} 页数据")
|
||||
|
||||
return all_data
|
||||
async def post(
|
||||
self,
|
||||
url: str,
|
||||
data: Any = None,
|
||||
*,
|
||||
params: Mapping[str, Any] | None = None,
|
||||
form_data: Any = None,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
timeout: ClientTimeout | None = None,
|
||||
retry: bool = False,
|
||||
) -> Any:
|
||||
if data is not None and form_data is not None:
|
||||
raise ValueError("JSON data 和 form_data 不能同时传入")
|
||||
response_data = await self._request(
|
||||
"POST",
|
||||
url,
|
||||
params=params,
|
||||
json_data=data,
|
||||
form_data=form_data,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
retry=retry,
|
||||
)
|
||||
return response_data
|
||||
|
||||
Reference in New Issue
Block a user