Download tests/test_session_manager.py from longxing/gmn2a: direct link, hf CLI and curl.
- Browser
- Download file 10.3 kB
-
https://huggingface.co/spaces/longxing/gmn2a/resolve/main/tests/test_session_manager.py
- Command line
-
hf download hf://spaces/longxing/gmn2a/tests/test_session_manager.py
-
curl -L -o test_session_manager.py https://huggingface.co/spaces/longxing/gmn2a/resolve/main/tests/test_session_manager.py
10.3 kB
| """对 gemini_session 的行为验证:不联网,用假客户端覆盖关键自愈路径。""" | |
| import asyncio | |
| import io | |
| import json | |
| import os | |
| import sys | |
| import tempfile | |
| import time | |
| from pathlib import Path | |
| from types import SimpleNamespace | |
| sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) | |
| import gemini_session as gs # noqa: E402 | |
| PASS, FAIL = [], [] | |
| def check(name, cond, extra=""): | |
| (PASS if cond else FAIL).append(name) | |
| print(("[OK] " if cond else "[FAIL] ") + name + (f" {extra}" if extra and not cond else "")) | |
| class FakeCookies: | |
| def __init__(self, values): | |
| self.jar = [SimpleNamespace(name=k, value=v) for k, v in values.items()] | |
| def get(self, name, default=None): | |
| for c in self.jar: | |
| if c.name == name: | |
| return c.value | |
| return default | |
| class FakeClient: | |
| """行为对齐 gemini_webapi.GeminiClient 的鸭子类型替身。""" | |
| def __init__(self, psid=None, psidts=None, **kwargs): | |
| self.psid = psid | |
| self.psidts = psidts | |
| self._running = False | |
| self.client = None | |
| self.auto_refresh = False | |
| self.refresh_task = None | |
| self.activity_task = None | |
| self.account_status = SimpleNamespace(name="AVAILABLE") | |
| self.cookies = FakeCookies({"__Secure-1PSID": psid, "__Secure-1PSIDTS": psidts}) | |
| self.closed = False | |
| self.init_kwargs = {} | |
| async def init(self, **kwargs): | |
| self.init_kwargs = kwargs | |
| self._running = True | |
| self.client = object() | |
| self.auto_refresh = kwargs.get("auto_refresh", False) | |
| return None | |
| async def start_auto_refresh(self): | |
| while self._running: | |
| await asyncio.sleep(3600) | |
| async def start_activity_watchdog(self): | |
| while self._running: | |
| await asyncio.sleep(3600) | |
| async def close(self): | |
| self.closed = True | |
| self._running = False | |
| self.client = None | |
| class FakeBackup: | |
| def __init__(self): | |
| self.enabled = True | |
| self.pushed = [] | |
| async def restore(self, psid): | |
| return None | |
| async def push(self, psid, value, force=False): | |
| self.pushed.append((psid, value)) | |
| return {"dataset_uploaded": ["x"]} | |
| def make_manager(cache_dir, psid="PSID-AAA", psidts="ENV-TS"): | |
| manager = gs.GeminiSessionManager( | |
| credential_provider=lambda: (psid, psidts), | |
| cache_dir=cache_dir, | |
| validate=None, | |
| backup=None, | |
| refresh_interval=600, | |
| health_interval=30, | |
| rearm_cooldown=300, | |
| ) | |
| return manager | |
| async def main(): | |
| tmp = tempfile.mkdtemp(prefix="gs-verify-") | |
| gs.GeminiClient = FakeClient | |
| # ---------------------------------------------------------------- 1. 基于 Cookie 初始化 | |
| manager = make_manager(tmp) | |
| client = await manager.get_client() | |
| check("1a 用环境变量 Cookie 初始化成功", isinstance(client, FakeClient)) | |
| check("1b auto_refresh 已开启", client.init_kwargs.get("auto_refresh") is True) | |
| check("1c auto_close 关闭(常驻服务)", client.init_kwargs.get("auto_close") is False) | |
| check("1d refresh_interval 透传", client.init_kwargs.get("refresh_interval") == 600) | |
| check("1e 凭据来源标记为 environment", manager.status()["credential_source"] == "environment") | |
| txt = Path(tmp) / ".cached_1psidts_PSID-AAA.txt" | |
| check("1f 初始化后把 1PSIDTS 落盘为文本缓存", txt.is_file() and txt.read_text() == "ENV-TS") | |
| # ---------------------------------------------------------------- 2. 观察到轮换 -> 落盘 | |
| client.cookies = FakeCookies({"__Secure-1PSID": "PSID-AAA", "__Secure-1PSIDTS": "ROTATED-1"}) | |
| manager._observe_rotation(client) | |
| check("2a 轮换后文本缓存被更新", txt.read_text() == "ROTATED-1") | |
| check("2b rotation_count 递增", manager.status()["rotation_count"] == 1) | |
| # ---------------------------------------------------------------- 3. 重复值不重复计数 | |
| manager._observe_rotation(client) | |
| check("3a 值未变化时不重复计数", manager.status()["rotation_count"] == 1) | |
| # ---------------------------------------------------------------- 4. 后台轮换任务死后被重新拉起 | |
| client.refresh_task = None | |
| manager._rearm_until = 0 | |
| await manager._rearm_background_tasks(client) | |
| check("4a 轮换任务被重新拉起", client.refresh_task is not None and not client.refresh_task.done()) | |
| check("4b rearm_count 递增", manager.status()["rearm_count"] == 1) | |
| check("4c rotation_task_alive 为真", manager.status()["rotation_task_alive"] is True) | |
| # 冷却期内不会重复拉起(避免账号状态异常时忙循环) | |
| first_task = client.refresh_task | |
| first_task.cancel() | |
| await asyncio.sleep(0) | |
| await manager._rearm_background_tasks(client) | |
| check("4d 冷却期内不重复拉起", client.refresh_task is first_task) | |
| # ---------------------------------------------------------------- 5. 活跃度看门狗同样被拉起 | |
| manager._rearm_until = 0 | |
| client.refresh_task = asyncio.create_task(asyncio.sleep(3600)) | |
| client.activity_task = None | |
| await manager._rearm_background_tasks(client) | |
| check("5a 活跃度看门狗被拉起", client.activity_task is not None and not client.activity_task.done()) | |
| # ---------------------------------------------------------------- 6. 失效检测 | |
| check("6a 健康客户端无问题", gs.GeminiSessionManager.session_problem(client) == "") | |
| client.account_status = SimpleNamespace(name="UNAUTHENTICATED") | |
| check("6b 401 会话被判为失效", "UNAUTHENTICATED" in gs.GeminiSessionManager.session_problem(client)) | |
| client.account_status = SimpleNamespace(name="AVAILABLE") | |
| client._running = False | |
| check("6c 未运行的客户端被判为失效", "not running" in gs.GeminiSessionManager.session_problem(client)) | |
| # ---------------------------------------------------------------- 7. 失效后自动重建 | |
| old = client | |
| rebuilt = await manager.get_client() | |
| check("7a 失效后按需重建", rebuilt is not old and isinstance(rebuilt, FakeClient)) | |
| check("7b 旧客户端被关闭", old.closed is True) | |
| check("7c recycle_count 递增", manager.status()["recycle_count"] >= 1) | |
| # ---------------------------------------------------------------- 8. 监督循环:任务停止后自动恢复 | |
| manager._rearm_until = 0 | |
| rebuilt.refresh_task = None | |
| rebuilt.cookies = FakeCookies({"__Secure-1PSID": "PSID-AAA", "__Secure-1PSIDTS": "ROTATED-2"}) | |
| await manager._supervise_once() | |
| check("8a 监督循环重新拉起轮换任务", rebuilt.refresh_task is not None and not rebuilt.refresh_task.done()) | |
| check("8b 监督循环同步落盘最新值", txt.read_text() == "ROTATED-2") | |
| # ---------------------------------------------------------------- 9. 辅助用户状态误报不能关闭正在服务的客户端 | |
| rebuilt.account_status = SimpleNamespace(name="UNAUTHENTICATED") | |
| before_status_probe = manager.client | |
| await manager._supervise_once() | |
| check("9a UNAUTHENTICATED 探针误报不会触发重建", manager.client is before_status_probe) | |
| check("9b 探针误报时传输层仍健康", gs.GeminiSessionManager.session_transport_problem(manager.client) == "") | |
| # 真正的传输层关闭仍然必须触发重建。 | |
| rebuilt._running = False | |
| await manager._supervise_once() | |
| check("9c 传输层失效后自动重建", manager.client is not rebuilt and manager.client is not None) | |
| # ---------------------------------------------------------------- 10. 两套缓存取较新的一份 | |
| tmp2 = tempfile.mkdtemp(prefix="gs-verify2-") | |
| store = gs.CookieStore(tmp2) | |
| json_path = Path(tmp2) / ".cached_cookies_PSID-BBB.json" | |
| json_path.write_text(json.dumps([{"name": "__Secure-1PSIDTS", "value": "FROM-JSON"}])) | |
| text_path = Path(tmp2) / ".cached_1psidts_PSID-BBB.txt" | |
| text_path.write_text("FROM-TEXT") | |
| os.utime(json_path, (time.time() + 10, time.time() + 10)) # json 更新 | |
| value, source = store.read_1psidts("PSID-BBB") | |
| check("10a JSON 更新时以 JSON 为准", (value, source) == ("FROM-JSON", "json-cache"), f"{value}/{source}") | |
| os.utime(text_path, (time.time() + 20, time.time() + 20)) # text 更新 | |
| value, source = store.read_1psidts("PSID-BBB") | |
| check("10b 文本更新时以文本为准", (value, source) == ("FROM-TEXT", "text-cache"), f"{value}/{source}") | |
| check("10c 非法 1PSID 不生成缓存路径", store.text_path("bad/../name") is None) | |
| # ---------------------------------------------------------------- 11. 缓存优先于环境变量 | |
| tmp3 = tempfile.mkdtemp(prefix="gs-verify3-") | |
| store3 = gs.CookieStore(tmp3) | |
| store3.write_1psidts("PSID-CCC", "CACHED-TS") | |
| manager3 = make_manager(tmp3, psid="PSID-CCC", psidts="ENV-TS-STALE") | |
| client3 = await manager3.get_client() | |
| check("11a 优先使用缓存里的 1PSIDTS", client3.psidts == "CACHED-TS", client3.psidts) | |
| check("11b 来源标记为 cache", manager3.status()["credential_source"].startswith("cache"), manager3.status()["credential_source"]) | |
| # ---------------------------------------------------------------- 12. 轮换后触发备份 | |
| tmp4 = tempfile.mkdtemp(prefix="gs-verify4-") | |
| backup = FakeBackup() | |
| manager4 = gs.GeminiSessionManager( | |
| credential_provider=lambda: ("PSID-DDD", "ENV-TS"), | |
| cache_dir=tmp4, | |
| backup=backup, | |
| health_interval=30, | |
| ) | |
| client4 = await manager4.get_client() | |
| client4.cookies = FakeCookies({"__Secure-1PSID": "PSID-DDD", "__Secure-1PSIDTS": "ROTATED-9"}) | |
| manager4._observe_rotation(client4) | |
| await asyncio.sleep(0.2) | |
| check("12a 轮换后触发 HF 备份", backup.pushed == [("PSID-DDD", "ROTATED-9")], str(backup.pushed)) | |
| # ---------------------------------------------------------------- 13. 生命周期 start/stop | |
| await manager4.start() | |
| check("13a supervisor_alive 为真", manager4.status()["supervisor_alive"] is True) | |
| await manager4.stop() | |
| check("13b stop 后 supervisor 停止", manager4.status()["supervisor_alive"] is False) | |
| check("13c stop 后客户端被关闭", manager4.client is None) | |
| # ---------------------------------------------------------------- 14. 无凭据时不炸 | |
| manager5 = make_manager(tempfile.mkdtemp(prefix="gs-verify5-"), psid="", psidts="") | |
| try: | |
| await manager5.get_client() | |
| check("14a 缺凭据时抛错", False) | |
| except RuntimeError as e: | |
| check("14a 缺凭据时抛 RuntimeError", "SECURE_1PSID" in str(e), str(e)) | |
| check("14b 状态里记录 last_error", "SECURE_1PSID" in manager5.status()["last_error"]) | |
| print(f"\n通过 {len(PASS)} / {len(PASS) + len(FAIL)}") | |
| if FAIL: | |
| print("失败项:") | |
| for name in FAIL: | |
| print(" -", name) | |
| return 1 if FAIL else 0 | |
| if __name__ == "__main__": | |
| sys.exit(asyncio.run(main())) | |