"""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 安全截断到剩余空间 - hash8:sha1("{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_bytes(skeleton 已包含两个 `-` 和一个 `-` 给 group_name, # 但 skeleton 里 group_name 位置是空的,所以可用字节 = max - skeleton_bytes + 1 # skeleton: prefix-uid--hash8 → prefix-uid--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 _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") 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 无法取明文 key;key 待下次回填" ) 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) 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}, }