diff --git a/README.md b/README.md index 8a0c9b1..45ee032 100644 --- a/README.md +++ b/README.md @@ -93,7 +93,7 @@ cp config.example.json config.json | `grok2api_auto_add_local` | 是否写入本地 grok2api token 池 | | `grok2api_local_token_file` | 本地 grok2api token 文件路径 | | `grok2api_auto_add_remote` | 是否写入远端 grok2api | -| `grok2api_remote_base` | 远端 grok2api 管理 API 地址 | +| `grok2api_remote_base` | 远端 grok2api 地址,可填站点根地址或 `/admin/api` 管理 API 地址 | | `grok2api_remote_app_key` | 远端 grok2api app key | ### Cloudflare 临时邮箱 admin 模式 diff --git a/grok_register_ttk.py b/grok_register_ttk.py index f7bd9c0..95c2998 100755 --- a/grok_register_ttk.py +++ b/grok_register_ttk.py @@ -316,6 +316,35 @@ def add_token_to_grok2api_local_pool(raw_token, email="", log_callback=None): return True +def get_grok2api_remote_api_bases(base): + """生成 grok2api 管理 API 候选根路径。 + + 参数: + - base str: 用户配置的 grok2api 远端地址 + + 返回: + - list[str]: 依次尝试的管理 API 根路径 + """ + normalized = str(base or "").strip().rstrip("/") + if not normalized: + return [] + lower = normalized.lower() + candidates = [normalized] + if lower.endswith("/admin/api"): + return candidates + if lower.endswith("/admin"): + candidates.append(f"{normalized}/api") + else: + candidates.append(f"{normalized}/admin/api") + seen = set() + unique = [] + for item in candidates: + if item not in seen: + unique.append(item) + seen.add(item) + return unique + + def add_token_to_grok2api_remote_pool(raw_token, email="", log_callback=None): token = _normalize_sso_token(raw_token) if not token: @@ -331,24 +360,29 @@ def add_token_to_grok2api_remote_pool(raw_token, email="", log_callback=None): query = {"app_key": app_key} pool_map = {"ssoBasic": "basic", "ssoSuper": "super"} remote_pool = pool_map.get(pool_name, "basic") + api_bases = get_grok2api_remote_api_bases(base) + add_errors = [] # 优先使用 add 接口,避免全量覆盖远端池 - try: - add_payload = {"tokens": [token], "pool": remote_pool, "tags": ["auto-register"]} - resp_add = http_post( - f"{base}/tokens/add", - headers=headers, - params=query, - json=add_payload, - timeout=30, - proxies={}, - ) - resp_add.raise_for_status() - if log_callback: - log_callback(f"[+] 已写入 grok2api 远端池: {pool_name} ({base}/tokens/add)") - return True - except Exception as add_exc: - if log_callback: - log_callback(f"[Debug] /tokens/add 写入失败,尝试 /tokens 全量模式: {add_exc}") + add_payload = {"tokens": [token], "pool": remote_pool, "tags": ["auto-register"]} + for api_base in api_bases: + endpoint = f"{api_base}/tokens/add" + try: + resp_add = http_post( + endpoint, + headers=headers, + params=query, + json=add_payload, + timeout=30, + proxies={}, + ) + resp_add.raise_for_status() + if log_callback: + log_callback(f"[+] 已写入 grok2api 远端池: {pool_name} ({endpoint})") + return True + except Exception as add_exc: + add_errors.append(f"{endpoint}: {add_exc}") + if log_callback: + log_callback(f"[Debug] /tokens/add 写入失败,尝试 /tokens 全量模式: {'; '.join(add_errors)}") # 兜底:旧版全量保存接口 current = {} diff --git a/tests/test_grok2api_remote_pool.py b/tests/test_grok2api_remote_pool.py new file mode 100644 index 0000000..783752e --- /dev/null +++ b/tests/test_grok2api_remote_pool.py @@ -0,0 +1,81 @@ +import unittest +from unittest.mock import patch + +import grok_register_ttk as app + + +class DummyResponse: + def __init__(self, payload=None, status_code=200, reason=""): + self._payload = payload or {} + self.status_code = status_code + self.reason = reason + self.text = "" + + def raise_for_status(self): + if self.status_code >= 400: + raise RuntimeError(f"HTTP Error {self.status_code}: {self.reason}") + + def json(self): + return self._payload + + +class Grok2ApiRemotePoolTests(unittest.TestCase): + def setUp(self): + self.original_config = app.config.copy() + + def tearDown(self): + app.config = self.original_config + + def test_remote_pool_falls_back_to_admin_api_prefix_when_root_tokens_add_is_404(self): + app.config.update({ + "grok2api_remote_base": "https://grok.example.com", + "grok2api_remote_app_key": "app-secret", + "grok2api_pool_name": "ssoBasic", + }) + calls = [] + + def fake_post(url, **kwargs): + calls.append((url, kwargs)) + if url == "https://grok.example.com/tokens/add": + return DummyResponse(status_code=404) + return DummyResponse({"status": "success", "count": 1}) + + with patch.object(app, "http_post", side_effect=fake_post): + ok = app.add_token_to_grok2api_remote_pool("sso=abc123", email="a@example.com") + + self.assertTrue(ok) + self.assertEqual([url for url, _ in calls], [ + "https://grok.example.com/tokens/add", + "https://grok.example.com/admin/api/tokens/add", + ]) + self.assertEqual(calls[-1][1]["params"], {"app_key": "app-secret"}) + self.assertEqual(calls[-1][1]["json"], { + "tokens": ["abc123"], + "pool": "basic", + "tags": ["auto-register"], + }) + + def test_remote_pool_does_not_duplicate_admin_api_prefix_when_base_already_points_to_admin_api(self): + app.config.update({ + "grok2api_remote_base": "https://grok.example.com/admin/api", + "grok2api_remote_app_key": "app-secret", + "grok2api_pool_name": "ssoSuper", + }) + calls = [] + + def fake_post(url, **kwargs): + calls.append((url, kwargs)) + return DummyResponse({"status": "success", "count": 1}) + + with patch.object(app, "http_post", side_effect=fake_post): + ok = app.add_token_to_grok2api_remote_pool("sso=super123", email="a@example.com") + + self.assertTrue(ok) + self.assertEqual([url for url, _ in calls], [ + "https://grok.example.com/admin/api/tokens/add", + ]) + self.assertEqual(calls[0][1]["json"]["pool"], "super") + + +if __name__ == "__main__": + unittest.main()