568 lines
24 KiB
Python
568 lines
24 KiB
Python
from __future__ import annotations
|
||
|
||
import logging
|
||
import sys
|
||
import time
|
||
from decimal import Decimal, InvalidOperation, ROUND_HALF_UP
|
||
from typing import Any
|
||
from urllib.parse import quote, urlparse
|
||
|
||
import httpx
|
||
|
||
from app.utils.number import fixed_decimal_number, fixed_decimal_string
|
||
from app.services.external_api_logger import log_external_api_call
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
class WebsiteError(RuntimeError):
|
||
pass
|
||
|
||
|
||
def _friendly_http_error(exc: httpx.HTTPStatusError) -> str:
|
||
"""将常见 HTTP 错误转换为中文友好提示,原始信息保留在日志中。"""
|
||
status = exc.response.status_code
|
||
url = exc.request.url if exc.request else "?"
|
||
logger.warning("website_client HTTP %s from %s: %s", status, url, exc)
|
||
if status == 401:
|
||
return "目标网站认证失败,请检查 Admin API Key / JWT 是否正确"
|
||
if status == 403:
|
||
return "目标网站权限不足,请检查当前凭证是否有分组管理权限"
|
||
if status == 404:
|
||
return f"目标网站接口不存在,请检查 API Prefix 和分组接口路径({exc.response.url.path})"
|
||
if 500 <= status < 600:
|
||
return "目标网站服务异常,请稍后重试"
|
||
return f"目标网站返回错误(HTTP {status})"
|
||
|
||
|
||
def _friendly_connection_error(exc: Exception) -> str:
|
||
"""将网络/超时异常转换为中文友好提示。"""
|
||
logger.warning("website_client connection error: %s", exc)
|
||
if isinstance(exc, httpx.TimeoutException):
|
||
return "目标网站请求超时,请检查网络连接和 API 地址是否正确"
|
||
if isinstance(exc, httpx.ConnectError):
|
||
return "无法连接目标网站,请检查 API 地址和网络连通性"
|
||
return f"目标网站通信异常:{exc}"
|
||
|
||
|
||
def parse_positive_decimal(value: Any) -> Decimal | None:
|
||
if value is None or value == "":
|
||
return None
|
||
try:
|
||
d = Decimal(str(value))
|
||
except (InvalidOperation, ValueError):
|
||
return None
|
||
return d if d > 0 else None
|
||
|
||
|
||
def calculate_target_rate(values: list[Any], percent: Any = 0, algorithm: str = "max_plus_percent") -> Decimal:
|
||
rates = [rate for rate in (parse_positive_decimal(v) for v in values) if rate is not None]
|
||
if not rates:
|
||
raise WebsiteError("没有可用的正数上游倍率")
|
||
if algorithm == "average_plus_percent":
|
||
base = sum(rates, Decimal("0")) / Decimal(len(rates))
|
||
elif algorithm == "min_plus_percent":
|
||
base = min(rates)
|
||
elif algorithm == "max_plus_percent":
|
||
base = max(rates)
|
||
else:
|
||
raise WebsiteError(f"不支持的算法:{algorithm}")
|
||
pct = Decimal(str(percent or 0))
|
||
if pct < 0:
|
||
raise WebsiteError("百分比不能为负数")
|
||
return (base * (Decimal("1") + pct / Decimal("100"))).quantize(Decimal("0.01"), rounding=ROUND_HALF_UP)
|
||
|
||
|
||
def _unwrap_data(value: Any) -> Any:
|
||
if isinstance(value, dict):
|
||
data = value.get("data")
|
||
if "data" in value and (
|
||
"code" in value
|
||
or "message" in value
|
||
or isinstance(data, list)
|
||
or (isinstance(data, dict) and any(key in data for key in ("items", "groups")))
|
||
):
|
||
value = data
|
||
if not isinstance(value, dict):
|
||
return value
|
||
for key in ("items", "groups"):
|
||
if key in value:
|
||
return value.get(key)
|
||
return value
|
||
|
||
|
||
def _extract_id(value: Any) -> str:
|
||
if isinstance(value, dict):
|
||
for key in ("id", "account_id", "accountId", "group_id", "groupId"):
|
||
candidate = value.get(key)
|
||
if candidate is not None:
|
||
return str(candidate)
|
||
for key in ("data", "result", "account", "group"):
|
||
found = _extract_id(value.get(key))
|
||
if found:
|
||
return found
|
||
return ""
|
||
|
||
|
||
def normalize_groups(value: Any) -> list[dict[str, Any]]:
|
||
raw = _unwrap_data(value)
|
||
if isinstance(raw, dict):
|
||
raw = list(raw.values())
|
||
if not isinstance(raw, list):
|
||
raise WebsiteError("分组接口没有返回列表")
|
||
groups: list[dict[str, Any]] = []
|
||
for item in raw:
|
||
if isinstance(item, str):
|
||
groups.append({
|
||
"id": item,
|
||
"name": item,
|
||
"rate_multiplier": None,
|
||
"description": None,
|
||
"raw": {"id": item, "name": item}
|
||
})
|
||
continue
|
||
if not isinstance(item, dict):
|
||
continue
|
||
gid = item.get("id") or item.get("group_id") or item.get("groupId") or item.get("name") or item.get("group_name")
|
||
if gid is None:
|
||
continue
|
||
name = item.get("name") or item.get("group_name") or str(gid)
|
||
rate = item.get("rate_multiplier") or item.get("rateMultiplier") or item.get("ratio")
|
||
desc = item.get("description") or item.get("desc") or item.get("remark")
|
||
groups.append({
|
||
"id": str(gid),
|
||
"name": str(name),
|
||
"rate_multiplier": fixed_decimal_string(rate, 2) if rate is not None else None,
|
||
"description": str(desc) if desc is not None else None,
|
||
"raw": item,
|
||
})
|
||
return groups
|
||
|
||
|
||
class Sub2ApiWebsiteClient:
|
||
def __init__(
|
||
self,
|
||
base_url: str,
|
||
api_prefix: str,
|
||
auth_type: str,
|
||
auth_config: dict[str, Any],
|
||
timeout: float = 30.0,
|
||
target_id: int | None = None,
|
||
target_name: str | None = None,
|
||
) -> None:
|
||
self.base_url = base_url.rstrip("/")
|
||
self.api_prefix = api_prefix.strip("/")
|
||
self.auth_type = auth_type
|
||
self.auth_config = auth_config
|
||
self.timeout = timeout
|
||
self.target_id = target_id
|
||
self.target_name = target_name
|
||
self._client = httpx.Client(timeout=timeout)
|
||
|
||
def close(self) -> None:
|
||
self._client.close()
|
||
|
||
def __enter__(self) -> Sub2ApiWebsiteClient:
|
||
return self
|
||
|
||
def __exit__(self, *args: Any) -> None:
|
||
self.close()
|
||
|
||
def _url(self, path: str) -> str:
|
||
prefix = f"/{self.api_prefix}" if self.api_prefix else ""
|
||
return f"{self.base_url}{prefix}/{path.lstrip('/')}"
|
||
|
||
def _headers(self) -> dict[str, str]:
|
||
headers = {"Accept": "application/json", "User-Agent": "SmartUp/1.0"}
|
||
if self.auth_type == "api_key":
|
||
key = self.auth_config.get("key") or self.auth_config.get("api_key") or ""
|
||
header = self.auth_config.get("header") or "x-api-key"
|
||
if key:
|
||
headers[header] = key
|
||
elif self.auth_type == "bearer":
|
||
token = self.auth_config.get("token") or ""
|
||
if token:
|
||
headers["Authorization"] = f"Bearer {token}"
|
||
return headers
|
||
|
||
def _request(self, method: str, path: str, body: Any = None) -> Any:
|
||
url = self._url(path)
|
||
started = time.monotonic()
|
||
status_code: int | None = None
|
||
error_type: str | None = None
|
||
error_msg: str | None = None
|
||
try:
|
||
try:
|
||
resp = self._client.request(method, url, json=body, headers=self._headers())
|
||
except httpx.TimeoutException as exc:
|
||
error_type, error_msg = type(exc).__name__, str(exc)[:500]
|
||
raise WebsiteError(_friendly_connection_error(exc)) from exc
|
||
except httpx.ConnectError as exc:
|
||
error_type, error_msg = type(exc).__name__, str(exc)[:500]
|
||
raise WebsiteError(_friendly_connection_error(exc)) from exc
|
||
except httpx.HTTPStatusError as exc:
|
||
error_type, error_msg = type(exc).__name__, str(exc)[:500]
|
||
status_code = exc.response.status_code
|
||
raise WebsiteError(_friendly_http_error(exc)) from exc
|
||
status_code = getattr(resp, "status_code", None)
|
||
try:
|
||
resp.raise_for_status()
|
||
except httpx.HTTPStatusError as exc:
|
||
error_type, error_msg = type(exc).__name__, str(exc)[:500]
|
||
status_code = exc.response.status_code
|
||
raise WebsiteError(_friendly_http_error(exc)) from exc
|
||
if not resp.content:
|
||
return None
|
||
text = resp.text
|
||
if "application/json" not in resp.headers.get("content-type", "") and text.lstrip().startswith("<"):
|
||
error_type = "WebsiteError"
|
||
error_msg = f"{method} {path} returned HTML"
|
||
raise WebsiteError(f"{method} {path} 返回了 HTML,请检查接口地址是否正确")
|
||
return resp.json()
|
||
except WebsiteError:
|
||
if error_type is None:
|
||
error_type = "WebsiteError"
|
||
e = sys.exc_info()[1]
|
||
if e:
|
||
error_msg = str(e)[:500]
|
||
raise
|
||
finally:
|
||
elapsed = int((time.monotonic() - started) * 1000)
|
||
parsed = urlparse(url)
|
||
log_external_api_call(
|
||
direction="website",
|
||
target_type="website",
|
||
method=method,
|
||
path=parsed.path or path,
|
||
url_host=parsed.hostname or "",
|
||
duration_ms=elapsed,
|
||
status_code=status_code,
|
||
success=error_type is None,
|
||
target_id=self.target_id,
|
||
target_name=self.target_name,
|
||
error_type=error_type,
|
||
error_message=error_msg,
|
||
)
|
||
|
||
def get_groups(self, endpoint: str = "/groups") -> list[dict[str, Any]]:
|
||
"""拉取分组列表,尝试 endpoint 和 fallback /groups/all。"""
|
||
last_error: Exception | None = None
|
||
tried_paths: list[str] = []
|
||
for path in [endpoint, "/groups/all"]:
|
||
tried_paths.append(path)
|
||
try:
|
||
return normalize_groups(self._request("GET", path))
|
||
except WebsiteError as exc:
|
||
msg = str(exc)
|
||
# 认证/权限类错误:直接抛出,不需要尝试 fallback
|
||
if "认证失败" in msg or "权限不足" in msg:
|
||
raise
|
||
# 404/5xx 等路径相关错误,试试另一个路径
|
||
last_error = exc
|
||
except Exception as exc:
|
||
last_error = exc
|
||
logger.info("get_groups fallback %s failed: %s", path, exc)
|
||
|
||
msg = str(last_error) if last_error else "拉取分组失败"
|
||
raise WebsiteError(f"{msg}(尝试接口:{'、'.join(tried_paths)})")
|
||
|
||
def update_group_rate(self, endpoint_template: str, group_id: str, rate: Any) -> Any:
|
||
path = endpoint_template.replace("{id}", quote(group_id, safe=""))
|
||
return self._request("PUT", path, {"rate_multiplier": fixed_decimal_number(rate, 2)})
|
||
|
||
def update_group(self, endpoint_template: str, group_id: str, body: dict[str, Any]) -> dict[str, Any]:
|
||
path = endpoint_template.replace("{id}", quote(group_id, safe=""))
|
||
resp = self._request("PUT", path, body)
|
||
data = _unwrap_data(resp)
|
||
return data if isinstance(data, dict) else {"value": data}
|
||
|
||
def create_group(self, body: dict[str, Any], endpoint: str = "/groups") -> dict[str, Any]:
|
||
resp = self._request("POST", endpoint, body)
|
||
data = _unwrap_data(resp)
|
||
return data if isinstance(data, dict) else {"value": data}
|
||
|
||
def create_account(self, body: dict[str, Any], endpoint: str = "/accounts") -> dict[str, Any]:
|
||
resp = self._request("POST", endpoint, body)
|
||
data = _unwrap_data(resp)
|
||
return data if isinstance(data, dict) else {"value": data}
|
||
|
||
def update_account(self, account_id: str, body: dict[str, Any], endpoint: str = "/accounts") -> dict[str, Any]:
|
||
"""更新远端账号(仅传入需要变更的字段)。"""
|
||
try:
|
||
resp = self._request("PUT", f"{endpoint}/{account_id}", body)
|
||
except WebsiteError as exc:
|
||
if isinstance(exc.__cause__, httpx.HTTPStatusError) and exc.__cause__.response.status_code == 404:
|
||
raise WebsiteError(f"目标账号 {account_id} 不存在或已被删除") from exc.__cause__
|
||
raise
|
||
data = _unwrap_data(resp)
|
||
return data if isinstance(data, dict) else {"value": data}
|
||
|
||
def bulk_update_accounts(self, account_ids: list[str], body: dict[str, Any], endpoint: str = "/accounts/bulk-update") -> dict[str, Any]:
|
||
"""批量更新账号。"""
|
||
int_ids = []
|
||
for aid in account_ids:
|
||
try:
|
||
int_ids.append(int(aid))
|
||
except ValueError:
|
||
pass
|
||
resp = self._request("POST", endpoint, {
|
||
"account_ids": int_ids,
|
||
**body
|
||
})
|
||
data = _unwrap_data(resp)
|
||
return data if isinstance(data, dict) else {"value": data}
|
||
|
||
def sync_account_upstream_models(self, account_id: str, endpoint: str = "/accounts") -> list[str]:
|
||
"""拉取上游真实支持模型并返回列表。"""
|
||
quoted_id = quote(account_id, safe="")
|
||
path = f"{endpoint}/{quoted_id}/models/sync-upstream"
|
||
resp = self._request("POST", path)
|
||
data = _unwrap_data(resp)
|
||
if isinstance(data, list):
|
||
return [str(m) for m in data if m]
|
||
if isinstance(data, dict):
|
||
for key in ("models", "items", "data", "list"):
|
||
val = data.get(key)
|
||
if isinstance(val, list):
|
||
return [str(m) for m in val if m]
|
||
raise WebsiteError(f"同步上游模型接口返回的数据格式不正确: {resp}")
|
||
|
||
@staticmethod
|
||
def _unwrap_list(value: dict) -> list | None:
|
||
"""递归展开嵌套的列表包装:data.items、data.data、items、accounts 等。"""
|
||
if isinstance(value, list):
|
||
return value
|
||
if not isinstance(value, dict):
|
||
return None
|
||
# 先看顶层
|
||
for key in ("items", "accounts", "records", "list", "data"):
|
||
v = value.get(key)
|
||
if isinstance(v, list):
|
||
return v
|
||
# 再看 data.items、data.records、data.list 等嵌套
|
||
data_val = value.get("data")
|
||
if isinstance(data_val, dict):
|
||
for key in ("items", "records", "list", "data", "accounts"):
|
||
v = data_val.get(key)
|
||
if isinstance(v, list):
|
||
return v
|
||
return None
|
||
|
||
def list_accounts(self, endpoint: str = "/accounts") -> list[dict[str, Any]] | None:
|
||
"""拉取远端账号列表。支持分页拉取,成功返回账号 dict 列表,失败返回 None。"""
|
||
base_path = endpoint
|
||
query_params = {}
|
||
if "?" in endpoint:
|
||
base_path, query_str = endpoint.split("?", 1)
|
||
for part in query_str.split("&"):
|
||
if "=" in part:
|
||
k, v = part.split("=", 1)
|
||
query_params[k] = v
|
||
|
||
page_size = 100
|
||
if "page_size" in query_params:
|
||
try:
|
||
page_size = int(query_params["page_size"])
|
||
except ValueError:
|
||
pass
|
||
|
||
page = 1
|
||
all_items = []
|
||
|
||
while True:
|
||
params = dict(query_params)
|
||
params["page"] = str(page)
|
||
params["page_size"] = str(page_size)
|
||
|
||
qp_str = "&".join(f"{k}={v}" for k, v in params.items())
|
||
current_path = f"{base_path}?{qp_str}"
|
||
|
||
try:
|
||
resp = self._request("GET", current_path)
|
||
except Exception:
|
||
logger.warning("account list fetch failed for %s", current_path, exc_info=True)
|
||
return None
|
||
|
||
items = self._unwrap_list(resp)
|
||
if items is None:
|
||
logger.warning("account list unexpected format for %s", current_path)
|
||
return None
|
||
|
||
all_items.extend([item for item in items if isinstance(item, dict)])
|
||
|
||
if isinstance(resp, dict):
|
||
# 尝试从顶层或嵌套的 data 字典中提取分页字段
|
||
meta = resp
|
||
data_val = resp.get("data")
|
||
if isinstance(data_val, dict) and any(k in data_val for k in ("pages", "total", "page_size")):
|
||
meta = data_val
|
||
|
||
pages = meta.get("pages")
|
||
if pages is not None:
|
||
try:
|
||
pages_val = int(pages)
|
||
if page >= pages_val:
|
||
break
|
||
page += 1
|
||
continue
|
||
except (ValueError, TypeError):
|
||
pass
|
||
|
||
total = meta.get("total")
|
||
current_page_size = meta.get("page_size")
|
||
if total is not None and current_page_size is not None:
|
||
try:
|
||
total_val = int(total)
|
||
page_size_val = int(current_page_size)
|
||
import math
|
||
calculated_pages = math.ceil(total_val / page_size_val)
|
||
if page >= calculated_pages:
|
||
break
|
||
page += 1
|
||
continue
|
||
except (ValueError, TypeError):
|
||
pass
|
||
|
||
break
|
||
|
||
return all_items
|
||
|
||
def _get_account_ids(self, endpoint: str = "/accounts") -> set[str] | None:
|
||
"""拉取远端账号列表。成功返回 ID 集合(可能为空),解析失败返回 None。"""
|
||
items = self.list_accounts(endpoint)
|
||
if items is None:
|
||
return None
|
||
ids: set[str] = set()
|
||
for item in items:
|
||
item_id = self.extract_id(item)
|
||
if item_id:
|
||
ids.add(item_id)
|
||
return ids
|
||
|
||
def account_exists(self, account_id: str, endpoint: str = "/accounts") -> bool | None:
|
||
"""检查目标账号是否存在。
|
||
|
||
优先拉取账号列表判断:
|
||
- 列表成功取到 → return account_id in ids(True=存在,False=已删除)
|
||
- 列表取不到(None)→ return None(校验失败,不清本地)
|
||
返回 True=存在,False=已删除,None=校验失败。
|
||
"""
|
||
ids = self._get_account_ids(endpoint)
|
||
if ids is None:
|
||
logger.warning("account_exists cannot verify %s: list fetch failed", account_id)
|
||
return None
|
||
return account_id in ids
|
||
|
||
def test_account(self, account_id: str, endpoint: str = "/accounts") -> dict[str, Any]:
|
||
"""测试远端账号可用性。利用 SSE 监听 test_complete 事件的 success 状态。
|
||
若 404,返回 {"status": "404", "cleanup_allowed": True, "message": "账号不存在"}
|
||
若 timeout / 无法建立连接,返回 {"status": "timeout"/"error", "cleanup_allowed": False}
|
||
"""
|
||
import json
|
||
quoted_id = quote(account_id, safe="")
|
||
path = f"{endpoint}/{quoted_id}/test"
|
||
url = self._url(path)
|
||
headers = self._headers()
|
||
# Accept text/event-stream since it's an SSE stream
|
||
headers["Accept"] = "text/event-stream"
|
||
|
||
started = time.monotonic()
|
||
status_code: int | None = None
|
||
error_type: str | None = None
|
||
error_msg: str | None = None
|
||
|
||
try:
|
||
try:
|
||
with self._client.stream("POST", url, headers=headers) as response:
|
||
status_code = response.status_code
|
||
if response.status_code == 404:
|
||
error_type = "HTTPStatus"
|
||
error_msg = "HTTP 404"
|
||
return {"status": "404", "cleanup_allowed": True, "message": "账号在远端不存在 (404)"}
|
||
if response.status_code >= 400:
|
||
response.raise_for_status()
|
||
|
||
# Read SSE stream line by line
|
||
current_event = None
|
||
success_value = None
|
||
has_error_event = False
|
||
|
||
for line in response.iter_lines():
|
||
if not line:
|
||
continue
|
||
if line.startswith("event:"):
|
||
current_event = line[6:].strip()
|
||
elif line.startswith("data:"):
|
||
data_str = line[5:].strip()
|
||
if current_event == "error":
|
||
has_error_event = True
|
||
if current_event == "test_complete" or not current_event:
|
||
try:
|
||
parsed = json.loads(data_str)
|
||
if isinstance(parsed, dict):
|
||
if "success" in parsed:
|
||
success_value = parsed["success"]
|
||
elif parsed.get("event") == "test_complete" and "success" in parsed:
|
||
success_value = parsed["success"]
|
||
except Exception:
|
||
pass
|
||
|
||
if success_value is True:
|
||
return {"status": "success", "cleanup_allowed": False, "message": "测试成功,账号有效"}
|
||
elif success_value is False:
|
||
return {"status": "failed", "cleanup_allowed": True, "message": "测试失败,账号失效"}
|
||
elif has_error_event:
|
||
error_type = "WebsiteError"
|
||
error_msg = "测试接口返回了错误事件"
|
||
return {"status": "error", "cleanup_allowed": False, "message": "测试接口返回了错误事件"}
|
||
else:
|
||
error_type = "WebsiteError"
|
||
error_msg = "未收到有效的测试完成事件"
|
||
return {"status": "invalid_sse", "cleanup_allowed": False, "message": "未收到有效的测试完成事件"}
|
||
except httpx.TimeoutException as exc:
|
||
error_type, error_msg = type(exc).__name__, str(exc)[:500]
|
||
return {"status": "timeout", "cleanup_allowed": False, "message": f"连接超时: {exc}"}
|
||
except httpx.HTTPStatusError as exc:
|
||
error_type, error_msg = type(exc).__name__, str(exc)[:500]
|
||
status_code = exc.response.status_code
|
||
return {"status": "http_error", "cleanup_allowed": False, "message": f"测试接口返回 HTTP {status_code}"}
|
||
except Exception as exc:
|
||
error_type, error_msg = type(exc).__name__, str(exc)[:500]
|
||
return {"status": "error", "cleanup_allowed": False, "message": f"请求异常: {exc}"}
|
||
finally:
|
||
elapsed = int((time.monotonic() - started) * 1000)
|
||
parsed = urlparse(url)
|
||
log_external_api_call(
|
||
direction="website",
|
||
target_type="website",
|
||
method="POST",
|
||
path=parsed.path or path,
|
||
url_host=parsed.hostname or "",
|
||
duration_ms=elapsed,
|
||
status_code=status_code,
|
||
success=error_type is None,
|
||
target_id=self.target_id,
|
||
target_name=self.target_name,
|
||
error_type=error_type,
|
||
error_message=error_msg,
|
||
)
|
||
|
||
def delete_account(self, account_id: str, endpoint: str = "/accounts") -> dict[str, Any]:
|
||
"""删除目标账号。
|
||
若 404 视为 already_deleted。
|
||
"""
|
||
quoted_id = quote(account_id, safe="")
|
||
path = f"{endpoint}/{quoted_id}"
|
||
try:
|
||
self._request("DELETE", path)
|
||
return {"status": "deleted", "message": "账号已删除"}
|
||
except WebsiteError as exc:
|
||
cause = exc.__cause__
|
||
if isinstance(cause, httpx.HTTPStatusError) and cause.response.status_code == 404:
|
||
return {"status": "already_deleted", "message": "账号不存在或已在远端被删除"}
|
||
raise
|
||
|
||
@staticmethod
|
||
def extract_id(value: Any) -> str:
|
||
return _extract_id(value)
|