feat: 优化上游认证复用与 wangwang888 分组接口
## 上游认证优化 ### 核心改进 - login_password 类型上游优先复用已保存 token,过期后自动 refresh,失败再重新登录 - 新增 ensure_authenticated() 方法替代直接调用 login(),减少不必要的登录请求 - 初始化时从 auth_config 自动加载已保存的 token 和 user_id - 请求遇到 401 时自动尝试 refresh 或 login 并重试一次 ### 实现细节 - UpstreamClient.__init__: 初始化时加载 auth_config 中的 token/new_api_user - _is_login_password_with_refresh(): 判断是否支持 refresh - _token_is_expired(): 检查 token 是否过期(提前 60 秒) - _refresh_login_password_token(): 刷新 login_password 类型的 token - ensure_authenticated(): 优先复用 token,过期后 refresh,失败再 login - _send_request(): 401 时针对 login_password 类型先 refresh 再 login 并重试 ### 调用点更新 - scheduler.py: _check_upstream, _sync_upstream_keys - website_sync.py: reconcile_upstream_keys_full - upstreams.py: list_generated_keys, generate_keys_by_groups, test_all, check_now ## wangwang888 分组接口优化 - 提供数据库更新脚本 update_wangwang888_groups_endpoint.py - 将 wangwang888 的 groups_endpoint 从 /groups 改为 /groups/all - 减少不必要的 /api/v1/admin/groups 统计开销 - website_client 已有 fallback 逻辑保证兼容性 ## 测试 - 新增 test_upstream_login_password_refresh.py(7 个测试用例) - 验证 token 复用、refresh、401 重试等逻辑 - 更新 test_upstream_key_sync.py 的 FakeClient 增加 ensure_authenticated - 所有现有测试保持通过(87 passed) Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Sonnet 4.6
parent
9c9a4b3423
commit
c96d55665d
@@ -451,6 +451,17 @@ class UpstreamClient:
|
||||
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()
|
||||
@@ -603,6 +614,60 @@ class UpstreamClient:
|
||||
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_sub2api_bearer_token(self) -> bool:
|
||||
if not self._is_sub2api_bearer():
|
||||
return False
|
||||
@@ -653,17 +718,30 @@ class UpstreamClient:
|
||||
getattr(resp, "status_code", None) == 401
|
||||
and auth
|
||||
and allow_refresh
|
||||
and self._is_sub2api_bearer()
|
||||
and self._refresh_sub2api_bearer_token()
|
||||
):
|
||||
resp = self._do_request(
|
||||
method,
|
||||
url,
|
||||
headers=self._headers(auth),
|
||||
cookies=self._cookies,
|
||||
**kwargs,
|
||||
)
|
||||
self._cookies.update(dict(resp.cookies))
|
||||
# sub2api bearer: 尝试 refresh
|
||||
if self._is_sub2api_bearer() and self._refresh_sub2api_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:
|
||||
@@ -969,6 +1047,28 @@ class UpstreamClient:
|
||||
return
|
||||
raise UpstreamError("login succeeded but no token or session cookie found in response")
|
||||
|
||||
def ensure_authenticated(self) -> None:
|
||||
"""确保已认证,优先复用 token,过期后 refresh,失败再重新登录。
|
||||
|
||||
对于 login_password 类型:
|
||||
- 已有未过期 token:不登录,直接使用
|
||||
- token 快过期或已过期且有 refresh_token:先 refresh
|
||||
- refresh 失败、无 token、无 refresh token:回退到现有登录流程
|
||||
其他认证类型:调用 login() 保持原有行为
|
||||
"""
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user