diff --git a/backend/app/routers/upstreams.py b/backend/app/routers/upstreams.py index e42eec6..36a33a5 100644 --- a/backend/app/routers/upstreams.py +++ b/backend/app/routers/upstreams.py @@ -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) diff --git a/backend/app/services/scheduler.py b/backend/app/services/scheduler.py index f36d4ae..4395d5f 100644 --- a/backend/app/services/scheduler.py +++ b/backend/app/services/scheduler.py @@ -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) diff --git a/backend/app/services/upstream_client.py b/backend/app/services/upstream_client.py index 99e2698..92294e0 100644 --- a/backend/app/services/upstream_client.py +++ b/backend/app/services/upstream_client.py @@ -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) diff --git a/backend/app/services/website_sync.py b/backend/app/services/website_sync.py index 3ab4e10..d268acd 100644 --- a/backend/app/services/website_sync.py +++ b/backend/app/services/website_sync.py @@ -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 diff --git a/backend/test_upstream_key_sync.py b/backend/test_upstream_key_sync.py index 36812b4..fe88887 100644 --- a/backend/test_upstream_key_sync.py +++ b/backend/test_upstream_key_sync.py @@ -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"}] diff --git a/backend/test_upstream_login_password_refresh.py b/backend/test_upstream_login_password_refresh.py new file mode 100644 index 0000000..b8a21bb --- /dev/null +++ b/backend/test_upstream_login_password_refresh.py @@ -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" diff --git a/backend/update_wangwang888_groups_endpoint.py b/backend/update_wangwang888_groups_endpoint.py new file mode 100644 index 0000000..14a60f9 --- /dev/null +++ b/backend/update_wangwang888_groups_endpoint.py @@ -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()