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:
SmartUp Developer
2026-07-02 09:26:19 +08:00
co-authored by Claude Sonnet 4.6
parent 9c9a4b3423
commit c96d55665d
7 changed files with 447 additions and 17 deletions
+4 -4
View File
@@ -198,7 +198,7 @@ def list_generated_keys(uid: int, db: Session = Depends(get_db), _=Depends(get_c
target_id=upstream.id,
target_name=upstream.name,
) as client:
client.login()
client.ensure_authenticated()
for prefix in website_sync._fetch_remote_managed_prefixes(db, uid):
# list_api_keys 的参数在 Sub2API 中通常是 search 参数
remote_keys_list.extend(client.list_api_keys(search=prefix, status="active"))
@@ -634,7 +634,7 @@ def generate_keys_by_groups(
target_name=u.name,
) as client:
try:
client.login()
client.ensure_authenticated()
groups = client.get_available_groups(u.groups_endpoint)
# New-API/Nox 上游:批量前预热 token 列表缓存(一次 /api/token/ 拉取),
# 避免每个分组都调 /api/token/search 触发 Nox 搜索限流(10 次/60 秒)。
@@ -686,7 +686,7 @@ def _test_upstream_core(db: Session, u: Upstream) -> UpstreamBatchActionItem:
target_id=u.id,
target_name=u.name,
) as client:
client.login()
client.ensure_authenticated()
groups = client.get_available_groups(u.groups_endpoint)
# 余额(与单行 test_upstream 保持一致)
if u.balance_endpoint and u.balance_response_path:
@@ -728,7 +728,7 @@ def _check_now_core(db: Session, u: Upstream) -> tuple[str, bool]:
target_id=u.id,
target_name=u.name,
) as client:
client.login()
client.ensure_authenticated()
groups = client.get_available_groups(u.groups_endpoint)
raw_rates = client.get_group_rates(u.rate_endpoint)
snapshot = build_snapshot(uid, u.base_url, u.api_prefix, groups, raw_rates)
+2 -2
View File
@@ -82,7 +82,7 @@ def _check_upstream(upstream_id: int) -> None:
target_name=upstream.name,
) as client:
try:
client.login()
client.ensure_authenticated()
groups = client.get_available_groups(upstream.groups_endpoint)
raw_rates = client.get_group_rates(upstream.rate_endpoint)
snapshot = build_snapshot(
@@ -264,7 +264,7 @@ def _sync_upstream_keys(upstream_id: int, snapshot: dict[str, Any], captured_at:
target_id=upstream.id,
target_name=upstream.name,
) as client:
client.login()
client.ensure_authenticated()
remote_key_ids = website_sync._fetch_remote_managed_key_ids(db, client, upstream_id)
except Exception as exc:
logger.warning("sync upstream keys list failed for %s: %s", upstream_id, exc)
+110 -10
View File
@@ -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)
+1 -1
View File
@@ -930,7 +930,7 @@ def reconcile_upstream_keys_full(db: Session, upstream_id: int) -> bool:
target_id=upstream.id,
target_name=upstream.name,
) as client:
client.login()
client.ensure_authenticated()
# 获取远端 Key 列表(支持自定义 managed_prefix
remote_key_ids = _fetch_remote_managed_key_ids(db, client, upstream_id)
keys_fetched = True
+6
View File
@@ -735,6 +735,9 @@ def test_generated_keys_persists_new_api_tokens_with_plaintext(db_session, monke
def login(self):
return None
def ensure_authenticated(self):
return None
def list_api_keys(self, search="", status="active"):
assert search == "SmartUp"
return [
@@ -800,6 +803,9 @@ def test_generate_keys_allows_new_api_user_upstream(db_session, monkeypatch):
def login(self):
return None
def ensure_authenticated(self):
return None
def get_available_groups(self, endpoint):
assert endpoint == "/api/user/self/groups"
return [{"id": "vip", "name": "VIP"}]
@@ -0,0 +1,272 @@
"""测试 login_password 类型的 token 复用与 refresh 逻辑。"""
from __future__ import annotations
import time
import httpx
import pytest
from app.services import upstream_client
from app.services.upstream_client import UpstreamClient
class FakeHttpClient:
def __init__(self, handler):
self.handler = handler
self.calls = []
def request(self, method: str, url: str, **kwargs):
request = httpx.Request(method, url)
self.calls.append({
"method": method,
"path": request.url.path,
"headers": kwargs.get("headers") or {},
"json": kwargs.get("json"),
})
return self.handler(request, kwargs)
def close(self) -> None:
pass
def _install_fake_client(monkeypatch, handler) -> FakeHttpClient:
fake = FakeHttpClient(handler)
monkeypatch.setattr(upstream_client.httpx, "Client", lambda timeout=None: fake)
return fake
def _response(request: httpx.Request, status_code: int, payload=None) -> httpx.Response:
return httpx.Response(status_code, json=payload, request=request)
def test_ensure_authenticated_with_valid_token_does_not_login(monkeypatch):
"""已有未过期 token 时不调用 login。"""
def handler(request: httpx.Request, kwargs):
raise AssertionError(f"unexpected request: {request.method} {request.url}")
fake = _install_fake_client(monkeypatch, handler)
client = UpstreamClient(
base_url="http://api.local",
api_prefix="/api/v1",
auth_type="login_password",
auth_config={
"email": "user@example.com",
"password": "secret",
"token": "valid-token",
"token_expires_at": int(time.time()) + 3600,
},
)
client.ensure_authenticated()
assert len(fake.calls) == 0
assert client._token == "valid-token"
def test_ensure_authenticated_with_expired_token_refreshes(monkeypatch):
"""token 过期且 refresh 成功时只调用 refresh。"""
updates: list[dict] = []
def handler(request: httpx.Request, kwargs):
if request.url.path == "/api/v1/auth/refresh":
assert kwargs["json"] == {"refresh_token": "old-refresh"}
return _response(request, 200, {
"access_token": "new-access",
"refresh_token": "new-refresh",
"expires_in": 1800,
})
raise AssertionError(f"unexpected request: {request.method} {request.url}")
fake = _install_fake_client(monkeypatch, handler)
client = UpstreamClient(
base_url="http://api.local",
api_prefix="/api/v1",
auth_type="login_password",
auth_config={
"email": "user@example.com",
"password": "secret",
"token": "expired-token",
"refresh_token": "old-refresh",
"token_expires_at": int(time.time()) - 100,
},
on_auth_config_update=updates.append,
)
client.ensure_authenticated()
assert [(c["method"], c["path"]) for c in fake.calls] == [
("POST", "/api/v1/auth/refresh"),
]
assert client._token == "new-access"
assert updates[-1]["token"] == "new-access"
assert updates[-1]["refresh_token"] == "new-refresh"
def test_ensure_authenticated_refresh_failure_falls_back_to_login(monkeypatch):
"""refresh 失败时回退 login。"""
updates: list[dict] = []
def handler(request: httpx.Request, kwargs):
if request.url.path == "/api/v1/auth/refresh":
return _response(request, 401, {"detail": "refresh token expired"})
if request.url.path == "/api/v1/auth/login":
assert kwargs["json"] == {"email": "user@example.com", "password": "secret"}
return _response(request, 200, {
"access_token": "login-access",
"refresh_token": "login-refresh",
"expires_in": 3600,
})
raise AssertionError(f"unexpected request: {request.method} {request.url}")
fake = _install_fake_client(monkeypatch, handler)
client = UpstreamClient(
base_url="http://api.local",
api_prefix="/api/v1",
auth_type="login_password",
auth_config={
"email": "user@example.com",
"password": "secret",
"token": "expired-token",
"refresh_token": "invalid-refresh",
"token_expires_at": int(time.time()) - 100,
},
on_auth_config_update=updates.append,
)
client.ensure_authenticated()
assert [(c["method"], c["path"]) for c in fake.calls] == [
("POST", "/api/v1/auth/refresh"),
("POST", "/api/v1/auth/login"),
]
assert client._token == "login-access"
assert updates[-1]["token"] == "login-access"
assert updates[-1]["refresh_token"] == "login-refresh"
def test_request_401_triggers_refresh_then_retry(monkeypatch):
"""请求 401 时 refresh 后重试成功。"""
updates: list[dict] = []
def handler(request: httpx.Request, kwargs):
if request.url.path == "/api/v1/groups/available":
auth = kwargs["headers"].get("Authorization")
if auth == "Bearer old-token":
return _response(request, 401, {"detail": "token expired"})
assert auth == "Bearer refreshed-token"
return _response(request, 200, [{"id": "g1", "name": "G1"}])
if request.url.path == "/api/v1/auth/refresh":
assert kwargs["json"] == {"refresh_token": "old-refresh"}
return _response(request, 200, {
"access_token": "refreshed-token",
"refresh_token": "new-refresh",
"expires_in": 1800,
})
raise AssertionError(f"unexpected request: {request.method} {request.url}")
fake = _install_fake_client(monkeypatch, handler)
client = UpstreamClient(
base_url="http://api.local",
api_prefix="/api/v1",
auth_type="login_password",
auth_config={
"email": "user@example.com",
"password": "secret",
"token": "old-token",
"refresh_token": "old-refresh",
"token_expires_at": int(time.time()) + 3600,
},
on_auth_config_update=updates.append,
)
groups = client.get_available_groups("/groups/available")
assert groups == [{"id": "g1", "name": "G1"}]
assert [(c["method"], c["path"]) for c in fake.calls] == [
("GET", "/api/v1/groups/available"),
("POST", "/api/v1/auth/refresh"),
("GET", "/api/v1/groups/available"),
]
assert updates[-1]["token"] == "refreshed-token"
def test_request_401_refresh_fails_then_login_and_retry(monkeypatch):
"""请求 401 时 refresh 失败,login 并重试成功。"""
updates: list[dict] = []
def handler(request: httpx.Request, kwargs):
if request.url.path == "/api/v1/groups/available":
auth = kwargs["headers"].get("Authorization")
if auth == "Bearer old-token":
return _response(request, 401, {"detail": "token expired"})
assert auth == "Bearer login-token"
return _response(request, 200, [{"id": "g1", "name": "G1"}])
if request.url.path == "/api/v1/auth/refresh":
return _response(request, 401, {"detail": "refresh token expired"})
if request.url.path == "/api/v1/auth/login":
return _response(request, 200, {
"access_token": "login-token",
"refresh_token": "login-refresh",
"expires_in": 3600,
})
raise AssertionError(f"unexpected request: {request.method} {request.url}")
fake = _install_fake_client(monkeypatch, handler)
client = UpstreamClient(
base_url="http://api.local",
api_prefix="/api/v1",
auth_type="login_password",
auth_config={
"email": "user@example.com",
"password": "secret",
"token": "old-token",
"refresh_token": "old-refresh",
"token_expires_at": int(time.time()) + 3600,
},
on_auth_config_update=updates.append,
)
groups = client.get_available_groups("/groups/available")
assert groups == [{"id": "g1", "name": "G1"}]
assert [(c["method"], c["path"]) for c in fake.calls] == [
("GET", "/api/v1/groups/available"),
("POST", "/api/v1/auth/refresh"),
("POST", "/api/v1/auth/login"),
("GET", "/api/v1/groups/available"),
]
assert updates[-1]["token"] == "login-token"
def test_non_login_password_calls_login_directly(monkeypatch):
"""非 login_password 认证类型行为不变。"""
fake = _install_fake_client(monkeypatch, lambda r, k: _response(r, 200, []))
client = UpstreamClient(
base_url="http://api.local",
api_prefix="/api/v1",
auth_type="bearer",
auth_config={"token": "static-token"},
)
client.ensure_authenticated()
assert len(fake.calls) == 0
def test_init_loads_saved_token_from_auth_config(monkeypatch):
"""初始化时从 auth_config 加载已保存的 token。"""
_install_fake_client(monkeypatch, lambda r, k: _response(r, 200, []))
client = UpstreamClient(
base_url="http://api.local",
api_prefix="/api/v1",
auth_type="login_password",
auth_config={
"email": "user@example.com",
"password": "secret",
"token": "saved-token",
"new_api_user": "123",
},
)
assert client._token == "saved-token"
assert client._new_api_user == "123"
@@ -0,0 +1,52 @@
#!/usr/bin/env python3
"""更新 wangwang888 网站的 groups_endpoint 为 /groups/all
使用方法:
python backend/update_wangwang888_groups_endpoint.py
"""
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).parent))
from app.database import SessionLocal
from app.models.website import Website
def update_wangwang888_endpoint():
db = SessionLocal()
try:
# 查找 wangwang888 网站(通过名称或 base_url 匹配)
websites = db.query(Website).filter(
(Website.name.like('%wangwang888%')) |
(Website.name.like('%旺旺888%')) |
(Website.base_url.like('%wangwang888%'))
).all()
if not websites:
print("未找到 wangwang888 网站配置")
return
for website in websites:
old_endpoint = website.groups_endpoint
if old_endpoint == "/groups/all":
print(f"网站 {website.name} (ID={website.id}) 已配置为 /groups/all,无需更新")
continue
website.groups_endpoint = "/groups/all"
print(f"更新网站 {website.name} (ID={website.id}):")
print(f" groups_endpoint: {old_endpoint} -> /groups/all")
db.commit()
print("\n更新完成")
except Exception as e:
db.rollback()
print(f"更新失败: {e}")
raise
finally:
db.close()
if __name__ == "__main__":
update_wangwang888_endpoint()