Files
SmartUp/backend/app/services/upstream_client.py
T

1328 lines
51 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.
"""Upstream HTTP client — ported from monitor_ai98pro_group_rates.py."""
from __future__ import annotations
import base64
import hashlib
import json
import logging
import re
import time
from typing import Any, Callable, Optional
from urllib.parse import urljoin, urlparse
import httpx
from app.utils.number import decimal_string
from app.services.external_api_logger import log_external_api_call
logger = logging.getLogger(__name__)
class UpstreamError(RuntimeError):
pass
class _PendingKeyError(UpstreamError):
"""POST /api/token/ 成功但最终无法取得明文 key 时抛出。
路由层捕获后将记录落库为 created_pending_key,而不是 failed。
"""
pass
NEW_API_DEFAULT_QUOTA_PER_UNIT = 500000
# New-API/Nox-API token name 上限(UTF-8 字节)
_NEW_API_TOKEN_NAME_MAX_BYTES = 50
def _utf8_safe_truncate(text: str, max_bytes: int) -> str:
"""按 UTF-8 字节安全截断字符串,不切坏多字节字符。"""
encoded = text.encode("utf-8")
if len(encoded) <= max_bytes:
return text
truncated = encoded[:max_bytes]
# 向后退直到合法 UTF-8 边界
return truncated.decode("utf-8", errors="ignore")
def build_new_api_token_name(
prefix: str,
upstream_id: int | str,
group_id: str,
group_name: str,
) -> str:
"""构造不超过 50 UTF-8 字节的 New-API/Nox-API token 名。
格式:{prefix}-{upstream_id}-{短分组名}-{hash8}
- 短分组名:去除非字母/数字/汉字/_/- 后,UTF-8 安全截断到剩余空间
- hash8sha1("{upstream_id}:{group_id}")[:8],保证同名分组仍能区分
- 若 {prefix}-{upstream_id}-{hash8} 本身超过 50 字节,抛出 ValueError
示例:
SmartUp-7-ClaudeMax智能分组-a1b2c3d4
"""
hash8 = hashlib.sha1(f"{upstream_id}:{group_id}".encode()).hexdigest()[:8]
# 最小骨架:{prefix}-{upstream_id}--{hash8}(含两个分隔符)
skeleton = f"{prefix}-{upstream_id}--{hash8}"
skeleton_bytes = len(skeleton.encode("utf-8"))
if skeleton_bytes > _NEW_API_TOKEN_NAME_MAX_BYTES:
raise ValueError(
f"token name prefix '{prefix}' 过长,骨架 '{skeleton}' 已达 "
f"{skeleton_bytes} 字节(上限 {_NEW_API_TOKEN_NAME_MAX_BYTES}"
)
# 可用于短分组名的字节数
safe_name = re.sub(r"[^a-zA-Z0-9\u4e00-\u9fff_-]", "", group_name or group_id)
# 骨架含一个额外分隔符,留给 group_name 的字节 = max - len(prefix-upid--hash8) - 1
# 即 max - skeleton_bytesskeleton 已包含两个 `-` 和一个 `-` 给 group_name
# 但 skeleton 里 group_name 位置是空的,所以可用字节 = max - skeleton_bytes + 1
# skeleton: prefix-uid--hash8 → prefix-uid-<group>-hash8 多出 group 部分
# 可用字节:_NEW_API_TOKEN_NAME_MAX_BYTES - len(f"{prefix}-{upstream_id}-".encode()) - len(f"-{hash8}".encode())
base_bytes = len(f"{prefix}-{upstream_id}-".encode("utf-8"))
suffix_bytes = len(f"-{hash8}".encode("utf-8"))
group_budget = _NEW_API_TOKEN_NAME_MAX_BYTES - base_bytes - suffix_bytes
short_name = _utf8_safe_truncate(safe_name, group_budget) if group_budget > 0 else ""
return f"{prefix}-{upstream_id}-{short_name}-{hash8}"
def _find_token(value: Any) -> str:
if isinstance(value, str) and value.count(".") >= 2:
return value
if isinstance(value, dict):
for key in ("token", "access_token", "accessToken", "jwt", "auth_token", "authToken"):
candidate = value.get(key)
if isinstance(candidate, str) and candidate:
return candidate
for key in ("data", "result", "user", "session"):
tok = _find_token(value.get(key))
if tok:
return tok
return ""
def _find_refresh_token(value: Any) -> str:
if isinstance(value, dict):
for key in ("refresh_token", "refreshToken"):
candidate = value.get(key)
if isinstance(candidate, str) and candidate:
return candidate
for key in ("data", "result", "user", "session"):
tok = _find_refresh_token(value.get(key))
if tok:
return tok
return ""
def _find_expires_in(value: Any) -> int | None:
if isinstance(value, dict):
for key in ("expires_in", "expiresIn"):
candidate = value.get(key)
if candidate is None:
continue
try:
expires_in = int(float(candidate))
except (TypeError, ValueError):
continue
if expires_in > 0:
return expires_in
for key in ("data", "result", "user", "session"):
expires_in = _find_expires_in(value.get(key))
if expires_in is not None:
return expires_in
return None
def _parse_jwt_exp_timestamp(token: str) -> int | None:
"""Extract the 'exp' claim (Unix seconds) from a JWT access token.
Returns None if the token isn't a JWT or has no valid exp claim.
"""
if not token or token.count(".") != 2:
return None
try:
payload_b64 = token.split(".")[1]
padding = 4 - len(payload_b64) % 4
if padding != 4:
payload_b64 += "=" * padding
claims = json.loads(base64.urlsafe_b64decode(payload_b64))
exp = claims.get("exp")
if exp is not None:
return int(exp)
except Exception:
pass
return None
def _clean_auth_header_value(value: Any, field_name: str) -> str:
text = str(value or "").strip()
if not text:
return ""
if text.startswith("Bearer "):
text = text[7:].strip()
# Try to sanitize non-latin-1 characters instead of hard-failing
try:
text.encode("latin-1")
except UnicodeEncodeError:
# Try stripping non-ASCII characters
cleaned = text.encode("ascii", errors="ignore").decode("ascii").strip()
if cleaned:
return cleaned
raise UpstreamError(
f"{field_name} 含有非 HTTP 标头字符(如中文或 emoji),"
f"请重新登录后再试"
) from None
return text
def _is_dashboard_jwt(value: Any) -> bool:
"""识别 New-API 面板登录 JWT,避免将其当作长期 API Token 使用。"""
token = str(value or "").strip()
return len(token) > 512 and token.count(".") == 2
def _find_user_id(value: Any) -> str:
if isinstance(value, dict):
for key in ("id", "user_id", "userId"):
candidate = value.get(key)
if candidate is not None:
return str(candidate)
for key in ("data", "result", "user", "session"):
user_id = _find_user_id(value.get(key))
if user_id:
return user_id
return ""
def mask_secret(value: Any) -> str:
text = str(value or "")
if not text:
return ""
if len(text) <= 8:
return text[:2] + "****" + text[-2:] if len(text) > 4 else "****"
return text[:4] + "**********" + text[-4:]
def _unwrap_data(value: Any) -> Any:
if isinstance(value, dict) and "data" in value and ("code" in value or "message" in value):
return value.get("data")
if isinstance(value, dict) and "data" in value and "success" in value:
return value.get("data")
return value
def _extract_id(value: Any) -> str:
if isinstance(value, dict):
for key in ("id", "key_id", "keyId"):
candidate = value.get(key)
if candidate is not None:
return str(candidate)
for key in ("data", "result", "key", "api_key"):
found = _extract_id(value.get(key))
if found:
return found
return ""
def _extract_key_value(value: Any) -> str:
if isinstance(value, str):
return value
if isinstance(value, dict):
for key in ("key", "api_key", "apiKey", "token", "value"):
candidate = value.get(key)
if isinstance(candidate, str) and candidate:
return candidate
for key in ("data", "result", "api_key", "key"):
found = _extract_key_value(value.get(key))
if found:
return found
return ""
def _is_success_response(value: Any) -> bool:
if not isinstance(value, dict) or "success" not in value:
return True
return value.get("success") is True
def _response_message(value: Any, fallback: str = "") -> str:
if isinstance(value, dict):
msg = value.get("message") or value.get("detail")
if msg:
return str(msg)
return fallback
def _unwrap_list(value: Any) -> Optional[list[dict[str, Any]]]:
def _normalize(lst: list) -> list[dict[str, Any]]:
out = []
for i in lst:
if isinstance(i, dict):
out.append(i)
elif isinstance(i, str):
out.append({"id": i, "name": i})
return out
if isinstance(value, list):
return _normalize(value)
if isinstance(value, dict) and not _is_success_response(value):
raise UpstreamError(_response_message(value, "upstream API returned success=false"))
if isinstance(value, dict):
for key in ("data", "items", "groups", "available_groups", "availableGroups"):
nested = value.get(key)
if isinstance(nested, list):
return _normalize(nested)
elif isinstance(nested, dict):
# Handle /api/user/self/groups where data is a dict of group_name -> { desc, ratio }
out = []
for k, item in nested.items():
row = {"id": k, "name": k}
if isinstance(item, dict):
row.update(item)
out.append(row)
return out
return None
def _group_id(group: dict[str, Any]) -> str:
for key in ("id", "group_id", "groupId"):
v = group.get(key)
if v is not None:
return str(v)
name = str(group.get("name") or group.get("group_name") or "")
platform = str(group.get("platform") or "")
return f"{platform}:{name}"
def _rate_from_group(group: dict[str, Any]) -> str:
for key in (
"user_rate_multiplier", "userRateMultiplier",
"effective_rate_multiplier", "effectiveRateMultiplier",
"rate_multiplier", "rateMultiplier",
):
r = decimal_string(group.get(key))
if r:
return r
return ""
def _extract_rates_map(raw: Any) -> dict[str, str]:
if raw is None:
return {}
# Handle one-api/new-api /api/option response where GroupRatio is in a list of options
if isinstance(raw, dict) and isinstance(raw.get("data"), list):
for item in raw["data"]:
if isinstance(item, dict) and item.get("key") == "GroupRatio":
val = item.get("value")
if isinstance(val, str):
try:
import json
parsed = json.loads(val)
if isinstance(parsed, dict):
result: dict[str, str] = {}
for k, v in parsed.items():
r = decimal_string(v)
if r:
result[str(k)] = r
return result
except Exception:
pass
elif isinstance(val, dict):
# In case it's returned as dict directly
result = {}
for k, v in val.items():
r = decimal_string(v)
if r:
result[str(k)] = r
return result
if isinstance(raw, dict):
candidates = raw
for key in ("data", "rates", "group_rates", "groupRates", "GroupRatio"):
nested = raw.get(key)
if isinstance(nested, dict):
candidates = nested
break
elif isinstance(nested, str) and key == "GroupRatio":
# Handle GroupRatio as a JSON string
try:
import json
parsed = json.loads(nested)
if isinstance(parsed, dict):
candidates = parsed
break
except Exception:
pass
result: dict[str, str] = {}
for k, v in candidates.items():
if isinstance(v, dict):
r = decimal_string(
v.get("rate_multiplier") or v.get("rateMultiplier")
or v.get("user_rate_multiplier") or v.get("userRateMultiplier")
or v.get("ratio")
)
else:
r = decimal_string(v)
if r:
result[str(k)] = r
return result
if isinstance(raw, list):
result = {}
for item in raw:
if not isinstance(item, dict):
continue
gid = _group_id(item)
rate = _rate_from_group(item)
if gid and rate:
result[gid] = rate
return result
return {}
def build_snapshot(upstream_id: int, base_url: str, api_prefix: str,
groups: list[dict[str, Any]], raw_rates: Any) -> dict[str, Any]:
from datetime import datetime, timezone
override_rates = _extract_rates_map(raw_rates)
entries: dict[str, dict[str, Any]] = {}
for g in groups:
gid = _group_id(g)
default_rate = _rate_from_group(g)
effective_rate = override_rates.get(gid, default_rate)
entries[gid] = {
"group_id": gid,
"group_name": g.get("name") or g.get("group_name") or "",
"platform": g.get("platform") or "",
"rate": effective_rate,
"default_rate": default_rate,
"override_rate": override_rates.get(gid, ""),
}
return {
"upstream_id": upstream_id,
"base_url": base_url.rstrip("/"),
"api_prefix": api_prefix,
"captured_at": datetime.now(timezone.utc).astimezone().isoformat(timespec="seconds"),
"groups": entries,
}
# 429 重试配置(秒)
_RATE_LIMIT_BACKOFFS = (2, 5, 10, 20)
_RATE_LIMIT_MAX_WAIT = 60
def _retry_on_429(fn: Callable[[], Any], backoffs: tuple[float, ...] | None = None) -> Any:
"""执行 fn(),若遇 429 则按退避序列重试,超限后抛 UpstreamError。
fn 应是 HTTP 调用;429 响应会被 httpx.HTTPStatusError 捕获(raise_for_status 之后)
或者 fn 内部已经解析为 UpstreamError 且 message 含 429。
"""
if backoffs is None:
backoffs = _RATE_LIMIT_BACKOFFS
total_waited = 0
for attempt, wait in enumerate(backoffs):
try:
return fn()
except Exception as exc:
status = None
retry_after: int | None = None
# httpx HTTPStatusError
if hasattr(exc, "response") and hasattr(exc.response, "status_code"):
status = exc.response.status_code
try:
retry_after = int(exc.response.headers.get("Retry-After", 0))
except (TypeError, ValueError):
retry_after = None
# UpstreamError wrapping a 429 text
if status != 429 and "429" in str(exc):
status = 429
if status != 429:
raise
sleep_time = retry_after if (retry_after and retry_after > 0) else wait
if total_waited + sleep_time > _RATE_LIMIT_MAX_WAIT:
raise UpstreamError(
f"rate limited (429) after {total_waited}s total wait; giving up"
) from exc
time.sleep(sleep_time)
total_waited += sleep_time
# 最后一次尝试,不捕获
return fn()
class UpstreamClient:
"""Sync HTTP client that handles all auth types."""
def __init__(
self,
base_url: str,
api_prefix: str,
auth_type: str,
auth_config: dict[str, Any],
timeout: float = 30.0,
on_auth_config_update: Callable[[dict[str, Any]], None] | None = None,
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.on_auth_config_update = on_auth_config_update
self.target_id = target_id
self.target_name = target_name
self._token: str = ""
self._cookies: dict[str, str] = {}
self._new_api_user: str = ""
self._client = httpx.Client(timeout=timeout)
# 批量生成期间的 token 列表缓存(name -> record),避免重复调用 /api/token/search
self._token_list_cache: dict[str, dict[str, Any]] | None = None
# login_password 类型初始化时从 auth_config 加载已保存的 token
if self.auth_type == "login_password":
saved_token = str(self.auth_config.get("token") or "").strip()
if saved_token:
self._token = saved_token
saved_user = (
self.auth_config.get("new_api_user", "")
or self.auth_config.get("user_id", "")
)
if saved_user:
self._new_api_user = str(saved_user)
def close(self) -> None:
self._client.close()
def __enter__(self) -> UpstreamClient:
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 _is_new_api_user_mode(self) -> bool:
login_path = str(self.auth_config.get("login_path") or "")
return (
self.api_prefix == ""
and (
bool(self.auth_config.get("new_api_user"))
or login_path == "/api/user/login"
or self.auth_type == "cookie"
or self.auth_type == "new_api_token"
or self.auth_type == "nox_token"
)
)
def _user_header_value(self) -> str:
return _clean_auth_header_value(
self.auth_config.get("user_id", "") or self.auth_config.get("new_api_user", ""),
"User header",
)
def _headers(self, auth: bool = True) -> dict[str, str]:
headers: dict[str, str] = {
"Accept": "application/json",
"User-Agent": "SmartUp/1.0",
}
if not auth:
return headers
if self.auth_type == "bearer":
token = _clean_auth_header_value(self.auth_config.get("token", ""), "Bearer token")
if token:
headers["Authorization"] = f"Bearer {token}"
elif self.auth_type == "nox_token":
token = _clean_auth_header_value(self.auth_config.get("token", ""), "Nox access token")
user_id = self._user_header_value()
if token:
headers["Authorization"] = f"Bearer {token}"
if user_id:
headers["Nox-Api-User"] = user_id
elif self.auth_type == "new_api_token":
token = _clean_auth_header_value(self.auth_config.get("token", ""), "New-API access token")
if _is_dashboard_jwt(token):
raise UpstreamError(
"New-API token 是面板登录 JWT,不是用户 API Token"
"请在上游 Token 页面重新生成 Access Token"
)
user_id = self._user_header_value()
if token:
headers["Authorization"] = f"Bearer {token}"
if user_id:
headers["New-Api-User"] = user_id
elif self.auth_type == "api_key":
key = _clean_auth_header_value(self.auth_config.get("key", ""), "API key")
header = self.auth_config.get("header", "Authorization")
if key:
headers[header] = key
elif self.auth_type == "cookie":
cookie_str = _clean_auth_header_value(self.auth_config.get("cookie_string", ""), "Cookie")
if cookie_str:
headers["Cookie"] = cookie_str
user_id = self._user_header_value()
if user_id:
headers["New-Api-User"] = user_id
headers["Nox-Api-User"] = user_id
elif self.auth_type == "login_password" and self._token:
token = _clean_auth_header_value(self._token, "Login token")
if token:
headers["Authorization"] = f"Bearer {token}"
if self.auth_type == "login_password" and self._new_api_user:
headers["New-Api-User"] = self._new_api_user
headers["Nox-Api-User"] = self._new_api_user
return headers
def _remember_auth_tokens(
self,
token: str = "",
refresh_token: str = "",
expires_in: int | None = None,
token_expires_at: int | None = None,
) -> None:
changed = False
if token and self.auth_config.get("token") != token:
self.auth_config["token"] = token
changed = True
if refresh_token and self.auth_config.get("refresh_token") != refresh_token:
self.auth_config["refresh_token"] = refresh_token
changed = True
if expires_in is not None:
if self.auth_config.get("expires_in") != expires_in:
self.auth_config["expires_in"] = expires_in
changed = True
token_expires_at = int(time.time()) + expires_in
if self.auth_config.get("token_expires_at") != token_expires_at:
self.auth_config["token_expires_at"] = token_expires_at
changed = True
elif token_expires_at is not None:
if self.auth_config.get("token_expires_at") != token_expires_at:
self.auth_config["token_expires_at"] = token_expires_at
changed = True
if changed and self.on_auth_config_update:
self.on_auth_config_update(dict(self.auth_config))
def _do_request(self, method: str, url: str, **kwargs: Any) -> httpx.Response:
"""Wrapped self._client.request() with timing and external-API logging.
Every HTTP call (initial, retry, refresh token) goes through here so
the log is accurate and complete.
A response with status_code >= 400 is logged as a failure even though
no exception was raised, so callers can check success/failure rates
accurately. Callers still use raise_for_status() for flow control.
"""
started = time.monotonic()
status_code: int | None = None
error_type: str | None = None
error_msg: str | None = None
try:
resp = self._client.request(method, url, **kwargs)
status_code = getattr(resp, "status_code", None)
if status_code is not None and status_code >= 400:
error_type = "HTTPStatus"
error_msg = f"HTTP {status_code}"
return resp
except Exception as exc:
error_type = type(exc).__name__
error_msg = str(exc)[:500]
if hasattr(exc, "response") and hasattr(exc.response, "status_code"):
status_code = exc.response.status_code
raise
finally:
elapsed = int((time.monotonic() - started) * 1000)
parsed = urlparse(url)
log_external_api_call(
direction="upstream",
target_type="upstream",
method=method,
path=parsed.path or "",
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 _is_login_password_with_refresh(self) -> bool:
"""判断是否为支持 refresh 的 login_password 上游。"""
return (
self.auth_type == "login_password"
and bool(self.auth_config.get("refresh_token"))
)
def _token_is_expired(self) -> bool:
"""判断当前 token 是否已过期或即将过期(提前 60 秒)。"""
expires_at = self.auth_config.get("token_expires_at")
if expires_at is None:
return False
try:
expires_at_int = int(expires_at)
except (TypeError, ValueError):
return False
return int(time.time()) + 60 >= expires_at_int
def _refresh_login_password_token(self) -> bool:
"""刷新 login_password 类型的 token。"""
if not self._is_login_password_with_refresh():
return False
refresh_token = str(self.auth_config.get("refresh_token") or "").strip()
if not refresh_token:
return False
default_refresh_path = "/api/user/refresh" if self.api_prefix == "" else "/auth/refresh"
refresh_path = self.auth_config.get("refresh_path") or default_refresh_path
try:
resp = self._do_request(
"POST",
self._url(refresh_path),
json={"refresh_token": refresh_token},
headers=self._headers(auth=False),
cookies=self._cookies,
)
self._cookies.update(dict(resp.cookies))
resp.raise_for_status()
payload = resp.json() if resp.content else {}
except Exception:
return False
token = _find_token(payload)
if not token:
return False
self._token = token
self._remember_auth_tokens(
token=token,
refresh_token=_find_refresh_token(payload) or refresh_token,
expires_in=_find_expires_in(payload),
)
return True
def _refresh_bearer_token(self) -> bool:
"""Refresh a Bearer access token using its refresh_token.
Default refresh path depends on api_prefix:
- api_prefix == "" → /api/user/refresh
- otherwise → /auth/refresh
Can be overridden via auth_config ``refresh_path``.
When the response has no ``expires_in``, falls back to parsing
the new access token's JWT ``exp`` claim.
"""
refresh_token = str(self.auth_config.get("refresh_token") or "").strip()
if not refresh_token:
return False
default_path = "/api/user/refresh" if self.api_prefix == "" else "/auth/refresh"
refresh_path = self.auth_config.get("refresh_path") or default_path
try:
resp = self._do_request(
"POST",
self._url(refresh_path),
json={"refresh_token": refresh_token},
headers=self._headers(auth=False),
cookies=self._cookies,
)
self._cookies.update(dict(resp.cookies))
resp.raise_for_status()
payload = resp.json() if resp.content else {}
except Exception:
return False
token = _find_token(payload)
if not token:
return False
expires_in = _find_expires_in(payload)
token_expires_at = None
if expires_in is None:
token_expires_at = _parse_jwt_exp_timestamp(token)
self._remember_auth_tokens(
token=token,
refresh_token=_find_refresh_token(payload) or refresh_token,
expires_in=expires_in,
token_expires_at=token_expires_at,
)
return True
def _ensure_bearer_token_fresh(self) -> None:
"""Proactively refresh Bearer token if it will expire within 1 hour.
Writes missing ``token_expires_at`` from JWT ``exp`` claim if absent.
Does nothing if there is no ``refresh_token`` or no expiry info.
"""
refresh_token = str(self.auth_config.get("refresh_token") or "").strip()
if not refresh_token:
return
# Try to derive token_expires_at from JWT exp if not already saved
if self.auth_config.get("token_expires_at") is None:
token = str(self.auth_config.get("token") or "").strip()
exp = _parse_jwt_exp_timestamp(token)
if exp is not None:
self._remember_auth_tokens(token_expires_at=exp)
token_expires_at = self.auth_config.get("token_expires_at")
if token_expires_at is None:
return
try:
expires_at_int = int(token_expires_at)
except (TypeError, ValueError):
return
# Refresh if within 1 hour of expiry
if int(time.time()) + 3600 >= expires_at_int:
self._refresh_bearer_token()
def _send_request(
self,
method: str,
url: str,
auth: bool = True,
allow_refresh: bool = True,
**kwargs: Any,
) -> httpx.Response:
resp = self._do_request(
method,
url,
headers=self._headers(auth),
cookies=self._cookies,
**kwargs,
)
self._cookies.update(dict(resp.cookies))
if (
getattr(resp, "status_code", None) == 401
and auth
and allow_refresh
):
# bearer: 尝试 refresh(含 Sub2API,统一走 _refresh_bearer_token
if self.auth_type == "bearer" and self._refresh_bearer_token():
resp = self._do_request(
method,
url,
headers=self._headers(auth),
cookies=self._cookies,
**kwargs,
)
self._cookies.update(dict(resp.cookies))
# login_password: 先尝试 refresh,失败再 login,然后重试一次
elif self.auth_type == "login_password":
refreshed = self._refresh_login_password_token()
if not refreshed:
self.login()
resp = self._do_request(
method,
url,
headers=self._headers(auth),
cookies=self._cookies,
**kwargs,
)
self._cookies.update(dict(resp.cookies))
return resp
def _request(self, method: str, path: str, body: Any = None, auth: bool = True) -> Any:
if auth and self.auth_type == "cookie" and "user/self" in path and not self._user_header_value():
raise UpstreamError("New-API user endpoint requires New-Api-User; re-extract the session cookie after login and save the upstream")
if auth and self.auth_type == "new_api_token" and "user/self" in path and not self._user_header_value():
raise UpstreamError("New-API endpoint requires New-Api-User; please fill user id and retry")
if auth and self.auth_type == "nox_token" and "user/self" in path and not self._user_header_value():
raise UpstreamError("Nox-API endpoint requires Nox-Api-User; please fill user id and retry")
url = self._url(path)
if body is not None:
resp = self._send_request(method, url, json=body, auth=auth)
else:
resp = self._send_request(method, url, auth=auth)
resp.raise_for_status()
ct = resp.headers.get("content-type", "")
if not resp.content:
return None
text = resp.text
if "application/json" not in ct and text.lstrip().startswith("<"):
raise UpstreamError(f"{method} {path} returned HTML, not JSON")
return resp.json()
def _ensure_api_success(self, payload: Any, action: str) -> None:
if not _is_success_response(payload):
raise UpstreamError(_response_message(payload, f"{action} failed"))
def _new_api_quota_per_unit(self) -> int:
try:
payload = self._request("GET", "/api/status", auth=False)
data = _unwrap_data(payload)
if isinstance(data, dict):
value = data.get("quota_per_unit")
quota_per_unit = int(float(value))
if quota_per_unit > 0:
return quota_per_unit
except Exception:
pass
return NEW_API_DEFAULT_QUOTA_PER_UNIT
@staticmethod
def _normalize_key_record(record: dict[str, Any]) -> dict[str, Any]:
out = dict(record)
if not out.get("name") and out.get("key_name"):
out["name"] = out.get("key_name")
if not out.get("group_id") and out.get("group"):
out["group_id"] = str(out.get("group"))
if not out.get("group_name") and out.get("group"):
out["group_name"] = str(out.get("group"))
return out
@staticmethod
def _extract_new_api_token_items(payload: Any) -> tuple[list[dict[str, Any]], dict[str, Any]]:
nested = _unwrap_data(payload)
meta: dict[str, Any] = {}
items: list[dict[str, Any]] | None = None
if isinstance(nested, list):
items = [i for i in nested if isinstance(i, dict)]
elif isinstance(nested, dict):
meta = nested
for key in ("items", "tokens", "list", "records"):
value = nested.get(key)
if isinstance(value, list):
items = [i for i in value if isinstance(i, dict)]
break
if items is None:
raise UpstreamError("unexpected New-API token list response")
return items, meta
def _request_new_api_token_list(self, path: str, params: dict[str, Any]) -> tuple[list[dict[str, Any]], dict[str, Any]]:
resp = self._send_request("GET", self._url(path), params=params)
resp.raise_for_status()
data = resp.json()
self._ensure_api_success(data, "list New-API tokens")
return self._extract_new_api_token_items(data)
def _list_all_new_api_tokens(self, page_size: int = 100, max_pages: int = 20) -> list[dict[str, Any]]:
all_items: list[dict[str, Any]] = []
for page in range(1, max_pages + 1):
items, meta = self._request_new_api_token_list(
"/api/token/",
{"p": page, "size": page_size},
)
all_items.extend(items)
total = meta.get("total") if isinstance(meta, dict) else None
if isinstance(total, int) and len(all_items) >= total:
break
if len(items) < page_size:
break
return [self._normalize_key_record(i) for i in all_items]
@staticmethod
def _matches_new_api_token_search(record: dict[str, Any], search: str) -> bool:
if not search:
return True
needle = search.strip()
if not needle:
return True
name = str(record.get("name") or record.get("key_name") or "")
key = str(record.get("key") or record.get("api_key") or record.get("apiKey") or record.get("token") or "")
return name.startswith(needle) or name == needle or key == needle
def _hydrate_new_api_token_key(self, record: dict[str, Any]) -> dict[str, Any]:
out = dict(record)
key_value = _extract_key_value(out)
if key_value and "*" not in key_value:
out["key"] = key_value
out["masked_key"] = out.get("masked_key") or mask_secret(key_value)
return out
token_id = out.get("id")
if token_id is None:
return out
try:
plaintext = self._get_new_api_token_key(token_id)
except Exception:
return out
out["key"] = plaintext
out["masked_key"] = mask_secret(plaintext)
return out
def _list_new_api_tokens(
self,
search: str = "",
group_id: str | int | None = None,
) -> list[dict[str, Any]]:
normalized: list[dict[str, Any]] = []
if search:
search_items, _ = self._request_new_api_token_list(
"/api/token/search",
{"keyword": search, "token": "", "p": 1, "size": 100},
)
normalized.extend(self._normalize_key_record(i) for i in search_items)
all_tokens = self._list_all_new_api_tokens()
seen_ids = {str(i.get("id")) for i in normalized if i.get("id") is not None}
for item in all_tokens:
item_id = item.get("id")
if item_id is not None and str(item_id) in seen_ids:
continue
if self._matches_new_api_token_search(item, search):
normalized.append(item)
if item_id is not None:
seen_ids.add(str(item_id))
if group_id is not None:
gid = str(group_id)
normalized = [i for i in normalized if str(i.get("group_id") or i.get("group") or "") == gid]
return [self._hydrate_new_api_token_key(i) for i in normalized]
def _get_new_api_token_key(self, token_id: str | int) -> str:
payload = self._request("POST", f"/api/token/{token_id}/key")
self._ensure_api_success(payload, "get New-API token key")
key_value = _extract_key_value(_unwrap_data(payload))
if not key_value:
raise UpstreamError("New-API token key response did not include key")
return key_value
def warm_token_list_cache(self) -> None:
"""预热 token 列表缓存;批量生成前调用一次,后续 find_smartup_group_key 走缓存。
缓存以 token name 为 key;同名取最新(id 最大)。
"""
all_tokens = _retry_on_429(lambda: self._list_all_new_api_tokens())
cache: dict[str, dict[str, Any]] = {}
for t in all_tokens:
name = str(t.get("name") or t.get("key_name") or "")
if not name:
continue
existing = cache.get(name)
if existing is None or (t.get("id") or 0) > (existing.get("id") or 0):
cache[name] = t
self._token_list_cache = cache
def _cache_add_token(self, record: dict[str, Any]) -> None:
"""创建成功后把新 token 追加进缓存,让同批次后续查询能感知它。"""
if self._token_list_cache is None:
return
name = str(record.get("name") or record.get("key_name") or "")
if name:
self._token_list_cache[name] = record
def _create_new_api_token(
self,
name: str,
group_id: str | int,
quota: float = 0,
expires_in_days: int | None = None,
) -> dict[str, Any]:
"""创建 Nox/New-API token。
主路径:POST 响应里直接有明文 key → 立即返回,不调用 search。
Fallback:响应无 key → 按名称查 token 列表再按 id 取明文 key;查询带 429 退避重试。
极端情况:POST 已成功但最终无法取得 key → 抛 _PendingKeyError,调用方应落库为 pending。
"""
unlimited = quota <= 0
body: dict[str, Any] = {
"name": name,
"remain_quota": 0 if unlimited else int(round(quota * self._new_api_quota_per_unit())),
"unlimited_quota": unlimited,
"expired_time": int(time.time()) + expires_in_days * 86400 if expires_in_days else -1,
"model_limits_enabled": False,
"model_limits": "",
"allow_ips": "",
"group": str(group_id),
"cross_group_retry": False,
}
payload = self._request("POST", "/api/token/", body)
self._ensure_api_success(payload, "create New-API token")
# ── 主路径:从创建响应直接取明文 key(Nox-API 当前行为:data="sk-...")──
create_data = _unwrap_data(payload)
key_from_response = _extract_key_value(create_data) if create_data else ""
token_id_from_response = _extract_id(create_data) if isinstance(create_data, dict) else ""
if key_from_response and "*" not in key_from_response:
record = {
"name": name,
"group": str(group_id),
"group_id": str(group_id),
"key": key_from_response,
"id": token_id_from_response or None,
}
if isinstance(create_data, dict):
record.update(create_data)
self._cache_add_token(record)
return {
"id": token_id_from_response or "",
"key": key_from_response,
"masked_key": mask_secret(key_from_response),
"raw": record,
}
# ── Fallback:按名称找 token,再按 id 取明文 key;带 429 退避重试 ──
try:
matches = _retry_on_429(
lambda: self._list_new_api_tokens(search=name, group_id=group_id)
)
except Exception as exc:
raise _PendingKeyError(
f"POST /api/token/ 成功,但查询 token 列表失败({exc});"
"key 待下次回填"
) from exc
token = next(
(i for i in matches if str(i.get("name") or "").strip() == name.strip()), None
)
if not token:
raise _PendingKeyError(
"POST /api/token/ 成功,但列表中找不到该 token;key 待下次回填"
)
token_id = token.get("id")
if token_id is None:
raise _PendingKeyError(
"POST /api/token/ 成功,但 token 无 id 无法取明文 keykey 待下次回填"
)
try:
key_value = _retry_on_429(lambda: self._get_new_api_token_key(token_id))
except Exception as exc:
raise _PendingKeyError(
f"POST /api/token/ 成功,但取明文 key 失败({exc});key 待下次回填"
) from exc
self._cache_add_token({**self._normalize_key_record(token), "key": key_value})
return {
"id": str(token_id),
"key": key_value,
"masked_key": mask_secret(key_value),
"raw": self._normalize_key_record(token),
}
def login(self) -> None:
if self.auth_type != "login_password":
return
email = self.auth_config.get("email", "")
password = self.auth_config.get("password", "")
default_login_path = "/api/user/login" if self.api_prefix == "" else "/auth/login"
login_path = self.auth_config.get("login_path") or default_login_path
default_username_field = "username" if login_path == "/api/user/login" else "email"
username_field = self.auth_config.get("username_field") or default_username_field
if not email or not password:
raise UpstreamError("login_password auth requires email and password in auth_config")
resp = self._request("POST", login_path, {username_field: email, "password": password}, auth=False)
data = _unwrap_data(resp)
if isinstance(data, dict) and data.get("require_2fa"):
raise UpstreamError(
"New-API 账号启用了二次验证,SmartUp 暂不支持密码登录 2FA"
"请改用 Access Token 或 Cookie"
)
if not _is_success_response(resp):
message = _response_message(resp, "New-API 登录失败")
raise UpstreamError(message)
token = _find_token(resp)
if token:
self._token = token
self._new_api_user = (
self.auth_config.get("new_api_user", "")
or self.auth_config.get("user_id", "")
or _find_user_id(resp)
)
refresh_token = _find_refresh_token(resp)
if self.api_prefix == "api/v1" or refresh_token:
self._remember_auth_tokens(
token=token,
refresh_token=refresh_token,
expires_in=_find_expires_in(resp),
)
return
if self._cookies:
self._new_api_user = self.auth_config.get("new_api_user", "") or _find_user_id(resp)
return
raise UpstreamError("login succeeded but no token or session cookie found in response")
def ensure_authenticated(self) -> None:
"""确保已认证,优先复用 token,过期后 refresh,失败再重新登录。
对于 bearer 类型:
- 有 refresh_token 时,在过期前 1 小时主动 refresh(避免上游在过期后拒绝刷新)
- 无 refresh_token 时不做任何事(静默失败,后续 API 调用会 401)
对于 login_password 类型(保持原有逻辑):
- 已有未过期 token:不登录,直接使用
- token 快过期或已过期且有 refresh_token:先 refresh
- refresh 失败、无 token、无 refresh token:回退到现有登录流程
"""
if self.auth_type == "bearer":
self._ensure_bearer_token_fresh()
return
if self.auth_type != "login_password":
self.login()
return
# 已有 token 且未过期,不需要登录
if self._token and not self._token_is_expired():
return
# token 已过期或即将过期,尝试 refresh
if self._token and self._is_login_password_with_refresh():
if self._refresh_login_password_token():
return
# refresh 失败或无 token,执行登录
self.login()
def get_available_groups(self, endpoint: str) -> list[dict[str, Any]]:
resp = self._request("GET", endpoint)
groups = _unwrap_list(resp)
if groups is None:
raise UpstreamError(f"{endpoint} did not return a list")
return groups
def get_group_rates(self, endpoint: str) -> Any:
return self._request("GET", endpoint)
def get_balance(self, endpoint: str, response_path: str) -> Optional[float]:
"""Call the balance endpoint and extract a numeric value using a dot-separated JSON path.
response_path 示例:
"balance" → resp["balance"]
"data.quota" → resp["data"]["quota"]
"data.total_balance" → resp["data"]["total_balance"]
"""
if not endpoint or not response_path:
return None
resp = self._request("GET", endpoint)
if not isinstance(resp, dict):
return None
parts = response_path.split(".")
value: Any = resp
for part in parts:
if isinstance(value, dict):
value = value.get(part)
else:
return None
if value is None:
return None
try:
return float(value)
except (ValueError, TypeError):
return None
def list_api_keys(
self,
search: str = "",
group_id: str | int | None = None,
status: str = "active",
endpoint: str = "/keys",
) -> list[dict[str, Any]]:
"""查询远端上游 Key 列表,支持按名称搜索、分组筛选、状态筛选。"""
if endpoint in {"/api/token", "/api/token/"} or (endpoint == "/keys" and self._is_new_api_user_mode()):
return self._list_new_api_tokens(search=search, group_id=group_id)
params: dict[str, Any] = {}
if search:
params["search"] = search
if group_id is not None:
params["group_id"] = int(group_id) if str(group_id).isdigit() else group_id
if status:
params["status"] = status
url = self._url(endpoint)
resp = self._send_request(
"GET",
url,
params=params if params else None,
)
resp.raise_for_status()
data = resp.json()
if isinstance(data, list):
return data
if isinstance(data, dict):
# 尝试展开常见的包装结构
for top_key in ("data", "result", "response"):
val = data.get(top_key)
if isinstance(val, list):
return val
if isinstance(val, dict):
for inner_key in ("items", "keys", "list", "records", "data"):
inner = val.get(inner_key)
if isinstance(inner, list):
return inner
# 顶层本身就是 list-like wrapper
for key in ("items", "keys", "list", "records"):
val = data.get(key)
if isinstance(val, list):
return [self._normalize_key_record(i) for i in val if isinstance(i, dict)]
raise UpstreamError(f"unexpected keys response type: {type(data).__name__}")
def delete_api_key(self, key_id: str, endpoint: str = "/keys") -> None:
"""删除远端上游上的一个 Key。"""
self._request("DELETE", f"{endpoint}/{key_id}")
def find_smartup_group_key(
self,
group_id: str | int,
expected_name: str,
prefix: str = "SmartUp",
) -> dict[str, Any] | None:
"""查找同一上游分组下是否已存在 SmartUp 前缀的 Key。
优先走本地 token 列表缓存(warm_token_list_cache 预热后有效),
避免在批量生成时高频调用 /api/token/search 触发 Nox 限流。
缓存未命中时回退到远端搜索。
"""
# ── 缓存命中路径 ──
if self._token_list_cache is not None:
record = self._token_list_cache.get(expected_name)
if record is not None:
rec_group = str(record.get("group_id") or record.get("group") or "")
if not rec_group or rec_group == str(group_id):
return self._hydrate_new_api_token_key(record)
# 缓存中无此名 → 远端也不存在(批次内强一致性假设)
return None
# ── 无缓存:走远端搜索(原有逻辑)──
gid = int(group_id) if str(group_id).isdigit() else group_id
keys = self.list_api_keys(search=prefix, group_id=gid, status="active")
for k in keys:
name = k.get("name") or k.get("key_name") or ""
if name == expected_name:
return k
# 部分后端返回的 name 可能带空格或 trimming
if name.strip() == expected_name.strip():
return k
return None
def create_api_key(
self,
name: str,
group_id: str | int,
quota: float = 0,
expires_in_days: int | None = None,
rate_limit_5h: float = 0,
rate_limit_1d: float = 0,
rate_limit_7d: float = 0,
endpoint: str = "/keys",
) -> dict[str, Any]:
if endpoint in {"/api/token", "/api/token/"} or (endpoint == "/keys" and self._is_new_api_user_mode()):
return self._create_new_api_token(
name,
group_id,
quota=quota,
expires_in_days=expires_in_days,
)
body: dict[str, Any] = {
"name": name,
"group_id": int(group_id) if str(group_id).isdigit() else group_id,
"quota": quota,
"rate_limit_5h": rate_limit_5h,
"rate_limit_1d": rate_limit_1d,
"rate_limit_7d": rate_limit_7d,
}
if expires_in_days:
body["expires_in_days"] = expires_in_days
resp = self._request("POST", endpoint, body)
data = _unwrap_data(resp)
key_value = _extract_key_value(data)
if not key_value:
raise UpstreamError("key create response did not include key")
return {
"id": _extract_id(data),
"key": key_value,
"masked_key": mask_secret(key_value),
"raw": data if isinstance(data, dict) else {"value": data},
}